diff --git a/dotnet/test/E2E/HooksE2ETests.cs b/dotnet/test/E2E/HooksE2ETests.cs index 0d9155fbc7..08366a87a4 100644 --- a/dotnet/test/E2E/HooksE2ETests.cs +++ b/dotnet/test/E2E/HooksE2ETests.cs @@ -32,13 +32,11 @@ public async Task Should_Invoke_PreToolUse_Hook_When_Model_Runs_A_Tool() // Create a file for the model to read await File.WriteAllTextAsync(Path.Join(Ctx.WorkDir, "hello.txt"), "Hello from the test!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of hello.txt and tell me what it says" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - // Should have received at least one preToolUse hook call Assert.NotEmpty(preToolUseInputs); @@ -68,13 +66,11 @@ public async Task Should_Invoke_PostToolUse_Hook_After_Model_Runs_A_Tool() // Create a file for the model to read await File.WriteAllTextAsync(Path.Join(Ctx.WorkDir, "world.txt"), "World from the test!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of world.txt and tell me what it says" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - // Should have received at least one postToolUse hook call Assert.NotEmpty(postToolUseInputs); @@ -109,13 +105,11 @@ public async Task Should_Invoke_Both_PreToolUse_And_PostToolUse_Hooks_For_Single await File.WriteAllTextAsync(Path.Join(Ctx.WorkDir, "both.txt"), "Testing both hooks!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of both.txt" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - // Both hooks should have been called Assert.NotEmpty(preToolUseInputs); Assert.NotEmpty(postToolUseInputs); @@ -149,13 +143,11 @@ public async Task Should_Deny_Tool_Execution_When_PreToolUse_Returns_Deny() var originalContent = "Original content that should not be modified"; await File.WriteAllTextAsync(Path.Join(Ctx.WorkDir, "protected.txt"), originalContent); - await session.SendAsync(new MessageOptions + var response = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Edit protected.txt and replace 'Original' with 'Modified'" }); - var response = await TestHelper.GetFinalAssistantMessageAsync(session); - // The hook should have been called Assert.NotEmpty(preToolUseInputs); diff --git a/dotnet/test/E2E/PermissionE2ETests.cs b/dotnet/test/E2E/PermissionE2ETests.cs index 2225dcba89..1de100dee1 100644 --- a/dotnet/test/E2E/PermissionE2ETests.cs +++ b/dotnet/test/E2E/PermissionE2ETests.cs @@ -111,13 +111,11 @@ public async Task Should_Deny_Permission_When_Handler_Returns_Denied() var testFilePath = Path.Combine(Ctx.WorkDir, "protected.txt"); await File.WriteAllTextAsync(testFilePath, "protected content"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Edit protected.txt and replace 'protected' with 'hacked'." }); - await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.True( userRejectedToolCall, "Expected a tool.execution_complete event whose error indicates the user rejected the call."); @@ -159,13 +157,16 @@ await session.SendAndWaitAsync(new MessageOptions public async Task Should_Work_With_Approve_All_Permission_Handler() { var session = await CreateSessionAsync(new SessionConfig()); + await AssertApproveAllPermissionHandlerAsync(session, TimeSpan.FromSeconds(120)); + } - await session.SendAsync(new MessageOptions + internal static async Task AssertApproveAllPermissionHandlerAsync(CopilotSession session, TimeSpan timeout) + { + var message = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What is 2+2?" - }); + }, timeout); - var message = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.Contains("4", message?.Data.Content ?? string.Empty); } @@ -183,13 +184,11 @@ public async Task Should_Handle_Async_Permission_Handler() } }); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Run 'echo test' and tell me what happens" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.True(permissionRequestReceived, "Permission request should have been received"); } @@ -321,13 +320,11 @@ public async Task Should_Receive_ToolCallId_In_Permission_Requests() } }); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Run 'echo test'" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.True(receivedToolCallId, "Should have received toolCallId in permission request"); } diff --git a/dotnet/test/E2E/SessionE2ETests.cs b/dotnet/test/E2E/SessionE2ETests.cs index 08781ce209..37c25095fc 100644 --- a/dotnet/test/E2E/SessionE2ETests.cs +++ b/dotnet/test/E2E/SessionE2ETests.cs @@ -54,8 +54,7 @@ public async Task Should_Create_A_Session_With_Appended_SystemMessage_Config() SystemMessage = new SystemMessageConfig { Mode = SystemMessageMode.Append, Content = systemMessageSuffix } }); - await session.SendAsync(new MessageOptions { Prompt = "What is your full name?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What is your full name?" }); Assert.NotNull(assistantMessage); var content = assistantMessage!.Data.Content ?? string.Empty; @@ -78,8 +77,7 @@ public async Task Should_Create_A_Session_With_Replaced_SystemMessage_Config() SystemMessage = new SystemMessageConfig { Mode = SystemMessageMode.Replace, Content = testSystemMessage } }); - await session.SendAsync(new MessageOptions { Prompt = "What is your full name?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What is your full name?" }); Assert.NotNull(assistantMessage); var content = assistantMessage!.Data.Content ?? string.Empty; @@ -219,8 +217,7 @@ public async Task Should_Create_Session_With_Custom_Tool() ] }); - await session.SendAsync(new MessageOptions { Prompt = "What is the secret number for key ALPHA?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What is the secret number for key ALPHA?" }); Assert.NotNull(assistantMessage); Assert.Contains("54321", assistantMessage!.Data.Content ?? string.Empty); } @@ -463,10 +460,7 @@ public async Task Should_Receive_Session_Events() // Events must be dispatched serially — never more than one handler invocation at a time. Assert.Equal(1, maxConcurrent); - // Verify the assistant response contains the expected answer. - // session.idle is ephemeral and not in getEvents(), but we already - // confirmed idle via the live event handler above. - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session, alreadyIdle: true); + var assistantMessage = observedEvents.OfType().LastOrDefault(); Assert.NotNull(assistantMessage); Assert.Contains("300", assistantMessage!.Data.Content); @@ -481,8 +475,17 @@ public async Task Send_Returns_Immediately_While_Events_Stream_In_Background() OnPermissionRequest = PermissionHandler.ApproveAll, }); var events = new ConcurrentQueue(); + AssistantMessageEvent? message = null; - session.On(evt => events.Enqueue(evt.Type)); + session.On(evt => + { + events.Enqueue(evt.Type); + if (evt is AssistantMessageEvent assistantMessage) + { + message = assistantMessage; + } + }); + var idle = TestHelper.GetNextEventOfTypeAsync(session); // Use a slow command so we can verify SendAsync() returns before completion await session.SendAsync(new MessageOptions { Prompt = "Run 'sleep 2 && echo done'" }); @@ -491,7 +494,7 @@ public async Task Send_Returns_Immediately_While_Events_Stream_In_Background() Assert.DoesNotContain("session.idle", events); // Wait for turn to complete - var message = await TestHelper.GetFinalAssistantMessageAsync(session); + await idle; Assert.Contains("done", message?.Data.Content ?? string.Empty); Assert.Contains("session.idle", events); diff --git a/dotnet/test/E2E/SystemMessageSectionsE2ETests.cs b/dotnet/test/E2E/SystemMessageSectionsE2ETests.cs index 41c46d3b9d..8e9f670c3e 100644 --- a/dotnet/test/E2E/SystemMessageSectionsE2ETests.cs +++ b/dotnet/test/E2E/SystemMessageSectionsE2ETests.cs @@ -31,8 +31,7 @@ public async Task Should_Use_Replaced_Identity_Section_In_Response() } }); - await session.SendAsync(new MessageOptions { Prompt = "Who are you?" }); - var response = await TestHelper.GetFinalAssistantMessageAsync(session); + var response = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Who are you?" }); Assert.NotNull(response); var content = response.Data.Content.ToLowerInvariant(); @@ -61,8 +60,7 @@ public async Task Should_Use_Replaced_Preamble_Section_In_Response() } }); - await session.SendAsync(new MessageOptions { Prompt = "Who are you?" }); - var response = await TestHelper.GetFinalAssistantMessageAsync(session); + var response = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Who are you?" }); Assert.NotNull(response); var content = response.Data.Content.ToLowerInvariant(); diff --git a/dotnet/test/E2E/SystemMessageTransformE2ETests.cs b/dotnet/test/E2E/SystemMessageTransformE2ETests.cs index 79210e61b3..9a8e687071 100644 --- a/dotnet/test/E2E/SystemMessageTransformE2ETests.cs +++ b/dotnet/test/E2E/SystemMessageTransformE2ETests.cs @@ -48,13 +48,11 @@ public async Task Should_Invoke_Transform_Callbacks_With_Section_Content() await File.WriteAllTextAsync(Path.Combine(Ctx.WorkDir, "test.txt"), "Hello transform!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of test.txt and tell me what it says" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.True(identityCallbackInvoked, "Expected identity transform callback to be invoked"); Assert.True(toneCallbackInvoked, "Expected tone transform callback to be invoked"); } @@ -83,13 +81,11 @@ public async Task Should_Apply_Transform_Modifications_To_Section_Content() await File.WriteAllTextAsync(Path.Combine(Ctx.WorkDir, "hello.txt"), "Hello!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of hello.txt" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - // Verify the transform result was actually applied to the system message var traffic = await Ctx.GetExchangesAsync(); Assert.NotEmpty(traffic); @@ -128,13 +124,11 @@ public async Task Should_Work_With_Static_Overrides_And_Transforms_Together() await File.WriteAllTextAsync(Path.Combine(Ctx.WorkDir, "combo.txt"), "Combo test!"); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Read the contents of combo.txt and tell me what it says" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.True(transformCallbackInvoked, "Expected identity transform callback to be invoked"); } } diff --git a/dotnet/test/E2E/TelemetryExportE2ETests.cs b/dotnet/test/E2E/TelemetryExportE2ETests.cs index e2ad447e26..f6914688fe 100644 --- a/dotnet/test/E2E/TelemetryExportE2ETests.cs +++ b/dotnet/test/E2E/TelemetryExportE2ETests.cs @@ -40,8 +40,7 @@ public async Task Should_Export_File_Telemetry_For_Sdk_Interactions() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions { Prompt = prompt }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = prompt }); Assert.NotNull(assistantMessage); Assert.Contains("TELEMETRY_E2E_DONE", assistantMessage!.Data.Content ?? string.Empty, StringComparison.Ordinal); diff --git a/dotnet/test/E2E/ToolResultsE2ETests.cs b/dotnet/test/E2E/ToolResultsE2ETests.cs index 103c7ffe2a..34aada586f 100644 --- a/dotnet/test/E2E/ToolResultsE2ETests.cs +++ b/dotnet/test/E2E/ToolResultsE2ETests.cs @@ -29,12 +29,11 @@ public async Task Should_Handle_Structured_ToolResultObject_From_Custom_Tool() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What's the weather in Paris?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Matches("(?i)sunny|72", assistantMessage!.Data.Content ?? string.Empty); @@ -56,12 +55,11 @@ public async Task Should_Handle_Tool_Result_With_Failure_ResultType() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Check the status of the service using check_status. If it fails, say 'service is down'." }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("service is down", assistantMessage!.Data.Content?.ToLowerInvariant() ?? string.Empty); @@ -84,12 +82,11 @@ public async Task Should_Preserve_ToolTelemetry_And_Not_Stringify_Structured_Res OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Analyze the file main.ts for issues." }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("no issues", assistantMessage!.Data.Content?.ToLowerInvariant() ?? string.Empty); diff --git a/dotnet/test/E2E/ToolsE2ETests.cs b/dotnet/test/E2E/ToolsE2ETests.cs index 8a786f3927..943e49f45f 100644 --- a/dotnet/test/E2E/ToolsE2ETests.cs +++ b/dotnet/test/E2E/ToolsE2ETests.cs @@ -35,12 +35,11 @@ await File.WriteAllTextAsync( OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What's the first line of README.md in this directory?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("ELIZA", assistantMessage!.Data.Content ?? string.Empty); } @@ -54,12 +53,11 @@ public async Task Invokes_Custom_Tool() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use encrypt_string to encrypt this string: Hello" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("HELLO", assistantMessage!.Data.Content ?? string.Empty); @@ -92,13 +90,11 @@ public async Task Low_Level_Tool_Definition() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "First, set the current phase to 'analyzing'. Then search for items with keyword 'copilot'. Report the phase and search results." }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); - Assert.NotNull(assistantMessage); var content = assistantMessage!.Data.Content ?? string.Empty; Assert.NotEmpty(content); @@ -133,8 +129,7 @@ public async Task Handles_Tool_Calling_Errors() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions { Prompt = "What is my location? If you can't find out, just say 'unknown'." }); - var answer = await TestHelper.GetFinalAssistantMessageAsync(session); + var answer = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "What is my location? If you can't find out, just say 'unknown'." }); // Check the underlying traffic var traffic = await Ctx.GetExchangesAsync(); @@ -175,14 +170,13 @@ public async Task Can_Receive_And_Return_Complex_Types() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Perform a DB query for the 'cities' table using IDs 12 and 19, sorting ascending. " + "Reply only with lines of the form: [cityname] [population]" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); var responseContent = assistantMessage?.Data.Content!; Assert.NotNull(assistantMessage); Assert.NotEmpty(responseContent); @@ -227,12 +221,11 @@ public async Task Overrides_Built_In_Tool_With_Custom_Tool() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use grep to search for the word 'hello'" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("CUSTOM_GREP_RESULT", assistantMessage!.Data.Content ?? string.Empty); @@ -351,12 +344,11 @@ static string SafeLookup([Description("Lookup ID")] string id) } }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use safe_lookup to look up 'test123'" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("RESULT", assistantMessage!.Data.Content ?? string.Empty); Assert.False(didRunPermissionRequest); @@ -371,12 +363,11 @@ public async Task Can_Return_Binary_Result() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use get_image. What color is the square in the image?" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("yellow", assistantMessage!.Data.Content?.ToLowerInvariant() ?? string.Empty); @@ -408,12 +399,11 @@ public async Task Invokes_Custom_Tool_With_Permission_Handler() }, }); - await session.SendAsync(new MessageOptions + var assistantMessage = await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use encrypt_string to encrypt this string: Hello" }); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); Assert.NotNull(assistantMessage); Assert.Contains("HELLO", assistantMessage!.Data.Content ?? string.Empty); @@ -438,13 +428,11 @@ public async Task Denies_Custom_Tool_When_Permission_Denied() OnPermissionRequest = async (request, invocation) => PermissionDecision.Reject(), }); - await session.SendAsync(new MessageOptions + await TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use encrypt_string to encrypt this string: Hello" }); - await TestHelper.GetFinalAssistantMessageAsync(session); - // The tool handler should NOT have been called since permission was denied Assert.False(toolHandlerCalled); @@ -472,7 +460,7 @@ public async Task Should_Execute_Multiple_Custom_Tools_In_Parallel_Single_Turn() OnPermissionRequest = PermissionHandler.ApproveAll, }); - await session.SendAsync(new MessageOptions + var assistantMessageTask = TestHelper.SendAndGetFinalAssistantMessageAsync(session, new MessageOptions { Prompt = "Use lookup_city with 'Paris' and lookup_country with 'France' at the same time, then combine both results in your reply." }); @@ -483,7 +471,7 @@ await session.SendAsync(new MessageOptions Assert.Equal("Paris", cityResult); Assert.Equal("France", countryResult); - var assistantMessage = await TestHelper.GetFinalAssistantMessageAsync(session); + var assistantMessage = await assistantMessageTask; Assert.NotNull(assistantMessage); var content = assistantMessage!.Data.Content ?? string.Empty; Assert.Contains("CITY_PARIS", content); diff --git a/dotnet/test/Harness/TestHelper.cs b/dotnet/test/Harness/TestHelper.cs index a230ddb81e..76045bfbb0 100644 --- a/dotnet/test/Harness/TestHelper.cs +++ b/dotnet/test/Harness/TestHelper.cs @@ -13,122 +13,14 @@ public static class TestHelper private static readonly TimeSpan DefaultEventTimeout = TimeSpan.FromSeconds(120); private static readonly TimeSpan DefaultPollInterval = TimeSpan.FromMilliseconds(100); - public static async Task GetFinalAssistantMessageAsync( + public static async Task SendAndGetFinalAssistantMessageAsync( CopilotSession session, - TimeSpan? timeout = null, - bool alreadyIdle = false) - { - var tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - using var cts = new CancellationTokenSource(timeout ?? DefaultEventTimeout); - - // Both `finalAssistantMessage` and `sawIdle` are set from two threads — the - // subscription callback (CLI read loop) and CheckExistingMessagesAsync (RPC reply). - // We complete only once we've observed both, regardless of which path saw which. - var stateLock = new object(); - AssistantMessageEvent? finalAssistantMessage = null; - bool sawIdle = false; - - void TryComplete() - { - AssistantMessageEvent? snapshot; - bool idle; - lock (stateLock) - { - snapshot = finalAssistantMessage; - idle = sawIdle; - } - if (snapshot != null && idle) tcs.TrySetResult(snapshot); - } - - using var subscription = session.On(evt => - { - switch (evt) - { - case AssistantMessageEvent msg: - lock (stateLock) { finalAssistantMessage = msg; } - TryComplete(); - break; - case SessionIdleEvent: - lock (stateLock) { sawIdle = true; } - TryComplete(); - break; - case SessionErrorEvent error: - tcs.TrySetException(new Exception(error.Data.Message ?? "session error")); - break; - } - }); - - // Backfill from already-delivered messages so we don't lose events that arrived - // between SendAsync returning and the subscription being installed. Run it - // concurrently with the live subscription, but keep the Task observable so any - // exception is propagated through tcs (not the unobserved-task handler) and so - // we can drain it deterministically below. Pass cts.Token so the backfill is - // bounded by the same timeout as the wait itself, and so a hung GetEventsAsync - // can't block the drain in `finally`. - var backfill = CheckExistingMessagesAsync(cts.Token); - - using var registration = cts.Token.Register( - static state => ((TaskCompletionSource)state!).TrySetException( - new TimeoutException("Timeout waiting for assistant message")), - tcs); - - try - { - return await tcs.Task; - } - finally - { - // Drain the backfill before our `using` scopes (cts, subscription) dispose. - // Any exception was already routed through tcs above, so swallow here. - try { await backfill.ConfigureAwait(false); } - catch (Exception) { /* intentionally ignored: already propagated via tcs */ } - } - - async Task CheckExistingMessagesAsync(CancellationToken cancellationToken) - { - try - { - var (existingFinal, existingIdle) = await GetExistingMessagesAsync(session, alreadyIdle, cancellationToken); - lock (stateLock) - { - // Preserve a newer message captured by the subscription in the meantime. - if (existingFinal != null && finalAssistantMessage == null) - { - finalAssistantMessage = existingFinal; - } - if (existingIdle) sawIdle = true; - } - TryComplete(); - } - catch (Exception ex) - { - tcs.TrySetException(ex); - } - } - } - - private static async Task<(AssistantMessageEvent? Final, bool SawIdle)> GetExistingMessagesAsync(CopilotSession session, bool alreadyIdle, CancellationToken cancellationToken = default) + MessageOptions options, + TimeSpan? timeout = null) { - var messages = (await session.GetEventsAsync(cancellationToken)).ToList(); - - var lastUserIdx = messages.FindLastIndex(m => m is UserMessageEvent); - var currentTurn = lastUserIdx < 0 ? messages : messages.Skip(lastUserIdx).ToList(); - - var error = currentTurn.OfType().FirstOrDefault(); - if (error != null) throw new Exception(error.Data.Message ?? "session error"); - - var idleIdx = alreadyIdle ? currentTurn.Count : currentTurn.FindIndex(m => m is SessionIdleEvent); - var sawIdle = alreadyIdle || idleIdx >= 0; - - // Find the most recent assistant message in the turn (whether idle has arrived or not). - var searchEnd = idleIdx >= 0 ? idleIdx : currentTurn.Count; - for (var i = searchEnd - 1; i >= 0; i--) - { - if (currentTurn[i] is AssistantMessageEvent msg) - return (msg, sawIdle); - } - - return (null, sawIdle); + // Subscribe before sending: session.idle is ephemeral and cannot be backfilled. + return await session.SendAndWaitAsync(options, timeout ?? DefaultEventTimeout) + ?? throw new InvalidOperationException("Session became idle without an assistant message."); } public static async Task GetNextEventOfTypeAsync( diff --git a/dotnet/test/Unit/ClientSessionLifetimeTests.cs b/dotnet/test/Unit/ClientSessionLifetimeTests.cs index dd7fdc2bbb..88de6abd0d 100644 --- a/dotnet/test/Unit/ClientSessionLifetimeTests.cs +++ b/dotnet/test/Unit/ClientSessionLifetimeTests.cs @@ -12,6 +12,7 @@ using System.Text; using System.Text.Json; using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; using Microsoft.Extensions.AI; using Xunit; @@ -1717,6 +1718,105 @@ private static void AssertMessageSource(JsonElement request, string? source) Assert.False(request.TryGetProperty("wait", out _)); } + [Theory] + [InlineData(true)] + [InlineData(false)] + public async Task Approve_All_Permission_Handler_Observes_Early_Events(bool completesBeforeReply) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll + }); + var timeout = TimeSpan.FromSeconds(5); + var sendReplied = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + server.BeforeResponseAsync = async (request, cancellationToken) => + { + if (request.Method == "session.send") + { + Assert.Equal("What is 2+2?", request.Params.GetProperty("prompt").GetString()); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["content"] = request.Params.GetProperty("prompt").GetString() + }); + await server.SendAndDrainSessionEventAsync(session, "assistant.message", new() + { + ["messageId"] = "permission-message", + ["content"] = "4" + }, timeout, cancellationToken); + if (completesBeforeReply) + { + await server.SendAndDrainSessionEventAsync(session, "session.idle", new(), timeout, cancellationToken); + } + } + }; + server.AfterResponseAsync = (request, _) => + { + if (request.Method == "session.send") + { + sendReplied.TrySetResult(); + } + return Task.CompletedTask; + }; + + // Exercise the E2E test's actual ordering and assertion, without launching a CLI. + var scenario = E2E.PermissionE2ETests.AssertApproveAllPermissionHandlerAsync(session, timeout); + await sendReplied.Task.WaitAsync(timeout); + if (!completesBeforeReply) + { + Assert.False(scenario.IsCompleted); + await server.SendAndDrainSessionEventAsync(session, "session.idle", new(), timeout); + } + await scenario; + + Assert.Single(server.Requests, request => request.Method == "session.send"); + var history = await session.GetEventsAsync(); + Assert.DoesNotContain(history, evt => evt is SessionIdleEvent); + Assert.Equal("4", Assert.Single(history.OfType()).Data.Content); + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public async Task SendAndGetFinalAssistantMessage_Requires_Current_Turn_Message(bool hasPreviousTurn) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var timeout = TimeSpan.FromSeconds(5); + server.BeforeResponseAsync = async (request, cancellationToken) => + { + if (request.Method == "session.send") + { + var prompt = request.Params.GetProperty("prompt").GetString(); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() { ["content"] = prompt }); + if (prompt == "previous turn") + { + await server.SendAndDrainSessionEventAsync(session, "assistant.message", new() + { + ["messageId"] = "previous-message", + ["content"] = "previous answer" + }, timeout, cancellationToken); + } + await server.SendAndDrainSessionEventAsync(session, "session.idle", new(), timeout, cancellationToken); + } + }; + + if (hasPreviousTurn) + { + var previous = await TestHelper.SendAndGetFinalAssistantMessageAsync( + session, new MessageOptions { Prompt = "previous turn" }, timeout); + Assert.Equal("previous answer", previous.Data.Content); + } + + var error = await Assert.ThrowsAsync(() => + TestHelper.SendAndGetFinalAssistantMessageAsync( + session, new MessageOptions { Prompt = "no assistant message" }, timeout)); + Assert.Equal("Session became idle without an assistant message.", error.Message); + Assert.DoesNotContain(await session.GetEventsAsync(), evt => evt is SessionIdleEvent); + } + [Theory] [InlineData(true)] [InlineData(false)] @@ -1738,23 +1838,23 @@ public async Task Abort_Recovery_Observes_Early_Events(bool recoveryCompletesBef }); if (sendCount == 1) { - await SendAndDrainAsync("tool.execution_start", new() + await server.SendAndDrainSessionEventAsync(session, "tool.execution_start", new() { ["toolCallId"] = "slow-tool", ["toolName"] = "shell" - }, cancellationToken); + }, timeout, cancellationToken); } else { Assert.Equal(2, sendCount); - await SendAndDrainAsync("assistant.message", new() + await server.SendAndDrainSessionEventAsync(session, "assistant.message", new() { ["messageId"] = "recovery-message", ["content"] = "4" - }, cancellationToken); + }, timeout, cancellationToken); if (recoveryCompletesBeforeReply) { - await SendAndDrainAsync("session.idle", new(), cancellationToken); + await server.SendAndDrainSessionEventAsync(session, "session.idle", new(), timeout, cancellationToken); } } } @@ -1765,14 +1865,14 @@ public async Task Abort_Recovery_Observes_Early_Events(bool recoveryCompletesBef { ["reason"] = "user" }); - await SendAndDrainAsync("session.idle", new() { ["aborted"] = true }, cancellationToken); + await server.SendAndDrainSessionEventAsync(session, "session.idle", new() { ["aborted"] = true }, timeout, cancellationToken); } }; server.AfterResponseAsync = async (request, cancellationToken) => { if (request.Method == "session.send" && sendCount == 2 && !recoveryCompletesBeforeReply) { - await SendAndDrainAsync("session.idle", new(), cancellationToken); + await server.SendAndDrainSessionEventAsync(session, "session.idle", new(), timeout, cancellationToken); } }; @@ -1786,16 +1886,6 @@ public async Task Abort_Recovery_Observes_Early_Events(bool recoveryCompletesBef var history = await session.GetEventsAsync(); Assert.DoesNotContain(history, evt => evt is SessionIdleEvent); Assert.Equal("4", Assert.Single(history.OfType()).Data.Content); - - async Task SendAndDrainAsync(string type, Dictionary data, CancellationToken cancellationToken) - { - var drained = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - using var subscription = session.On(_ => drained.TrySetResult()); - await server.SendSessionEventAsync(session.SessionId, type, data); - // A later event is a fence: every subscriber has finished handling the target event. - await server.SendSessionEventAsync(session.SessionId, "session.title_changed", new() { ["title"] = "fence" }); - await drained.Task.WaitAsync(timeout, cancellationToken); - } } [Fact] @@ -2411,6 +2501,21 @@ public Task SendSessionEventAsync(string sessionId, string type, Dictionary data, + TimeSpan timeout, + CancellationToken cancellationToken = default) + { + var drained = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var subscription = session.On(_ => drained.TrySetResult()); + await SendSessionEventAsync(session.SessionId, type, data); + // A later event is a fence: every subscriber has finished handling the target event. + await SendSessionEventAsync(session.SessionId, "session.title_changed", new() { ["title"] = "fence" }); + await drained.Task.WaitAsync(timeout, cancellationToken); + } + public async ValueTask DisposeAsync() { _allowDestroy.TrySetResult(); diff --git a/go/internal/e2e/mcp_and_agents_e2e_test.go b/go/internal/e2e/mcp_and_agents_e2e_test.go index 0cc57810ed..086f4b083f 100644 --- a/go/internal/e2e/mcp_and_agents_e2e_test.go +++ b/go/internal/e2e/mcp_and_agents_e2e_test.go @@ -34,6 +34,8 @@ func TestMCPServersE2E(t *testing.T) { waitForMCPServerStatus(t, session, "test-server", rpc.MCPServerStatusConnected) // Simple interaction to verify session works + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "What is 2+2?", }) @@ -41,7 +43,7 @@ func TestMCPServersE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - message, err := testharness.GetFinalAssistantMessage(t.Context(), session) + message, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get final message: %v", err) } @@ -205,6 +207,8 @@ func TestCustomAgentsE2E(t *testing.T) { } // Simple interaction to verify session works + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "What is 5+5?", }) @@ -212,7 +216,7 @@ func TestCustomAgentsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - message, err := testharness.GetFinalAssistantMessage(t.Context(), session) + message, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get final message: %v", err) } diff --git a/go/internal/e2e/permissions_e2e_test.go b/go/internal/e2e/permissions_e2e_test.go index 89681470e8..1b3cf7aeb1 100644 --- a/go/internal/e2e/permissions_e2e_test.go +++ b/go/internal/e2e/permissions_e2e_test.go @@ -157,6 +157,8 @@ func TestPermissionsE2E(t *testing.T) { t.Fatalf("Failed to write test file: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "Edit protected.txt and replace 'protected' with 'hacked'.", }) @@ -164,7 +166,7 @@ func TestPermissionsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - _, err = testharness.GetFinalAssistantMessage(t.Context(), session) + _, err = finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get final message: %v", err) } @@ -285,12 +287,14 @@ func TestPermissionsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 2+2?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - message, err := testharness.GetFinalAssistantMessage(t.Context(), session) + message, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get final message: %v", err) } @@ -487,6 +491,8 @@ func TestPermissionsE2E(t *testing.T) { } }) + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() go func() { _, _ = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "Run 'echo slow_handler_test'", @@ -515,9 +521,9 @@ func TestPermissionsE2E(t *testing.T) { close(releaseHandler) - message, err := testharness.GetFinalAssistantMessage(t.Context(), session) + message, err := finalMessage.Wait(t.Context()) if err != nil { - t.Fatalf("GetFinalAssistantMessage failed: %v", err) + t.Fatalf("Waiting for final assistant message failed: %v", err) } lifecycleMu.Lock() diff --git a/go/internal/e2e/resume_mcp_oauth_e2e_test.go b/go/internal/e2e/resume_mcp_oauth_e2e_test.go index db61f483a9..0c8dbd6cc1 100644 --- a/go/internal/e2e/resume_mcp_oauth_e2e_test.go +++ b/go/internal/e2e/resume_mcp_oauth_e2e_test.go @@ -29,12 +29,14 @@ func TestResumeMCPOAuthE2E(t *testing.T) { } sessionID := session1.SessionID + finalMessage := testharness.SubscribeToFinalAssistantMessage(session1) + defer finalMessage.Close() _, err = session1.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 1+1?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session1) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } diff --git a/go/internal/e2e/session_e2e_test.go b/go/internal/e2e/session_e2e_test.go index 12550e6e2d..941015ed37 100644 --- a/go/internal/e2e/session_e2e_test.go +++ b/go/internal/e2e/session_e2e_test.go @@ -1,6 +1,7 @@ package e2e import ( + "context" "encoding/base64" "os" "path/filepath" @@ -156,12 +157,14 @@ func TestSessionE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What is your full name?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - assistantMessage, err := testharness.GetFinalAssistantMessage(t.Context(), session) + assistantMessage, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -368,12 +371,14 @@ func TestSessionE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What is the secret number for key ALPHA?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - assistantMessage, err := testharness.GetFinalAssistantMessage(t.Context(), session) + assistantMessage, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -402,12 +407,14 @@ func TestSessionE2E(t *testing.T) { } sessionID := session1.SessionID + finalMessage := testharness.SubscribeToFinalAssistantMessage(session1) + defer finalMessage.Close() _, err = session1.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 1+1?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session1) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -428,7 +435,7 @@ func TestSessionE2E(t *testing.T) { t.Errorf("Expected resumed session ID to match, got %q vs %q", session2.SessionID, sessionID) } - answer2, err := testharness.GetFinalAssistantMessage(t.Context(), session2, true) + answer2, err := testharness.GetFinalAssistantMessageFromHistory(t.Context(), session2) if err != nil { t.Fatalf("Failed to get assistant message from resumed session: %v", err) } @@ -463,12 +470,14 @@ func TestSessionE2E(t *testing.T) { } sessionID := session1.SessionID + finalMessage := testharness.SubscribeToFinalAssistantMessage(session1) + defer finalMessage.Close() _, err = session1.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 1+1?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session1) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -632,28 +641,13 @@ func TestSessionE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } - // Set up event listeners BEFORE sending to avoid race conditions - toolStartCh := make(chan *copilot.SessionEvent, 1) - toolStartErrCh := make(chan error, 1) - go func() { - evt, err := testharness.GetNextEventOfType(session, copilot.SessionEventTypeToolExecutionStart, 60*time.Second) - if err != nil { - toolStartErrCh <- err - } else { - toolStartCh <- evt - } - }() - - sessionIdleCh := make(chan *copilot.SessionEvent, 1) - sessionIdleErrCh := make(chan error, 1) - go func() { - evt, err := testharness.GetNextEventOfType(session, copilot.SessionEventTypeSessionIdle, 60*time.Second) - if err != nil { - sessionIdleErrCh <- err - } else { - sessionIdleCh <- evt - } - }() + // Install subscriptions synchronously; starting a goroutine is not a fence. + abortCtx, cancelAbort := context.WithTimeout(t.Context(), 60*time.Second) + defer cancelAbort() + toolStart := testharness.SubscribeToEvent(session, copilot.SessionEventTypeToolExecutionStart) + defer toolStart.Close() + sessionIdle := testharness.SubscribeToEvent(session, copilot.SessionEventTypeSessionIdle) + defer sessionIdle.Close() // Send a message that triggers a long-running shell command _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "run the shell command 'sleep 100' (note this works on both bash and PowerShell)"}) @@ -662,10 +656,7 @@ func TestSessionE2E(t *testing.T) { } // Wait for tool.execution_start - select { - case <-toolStartCh: - // Tool execution has started - case err := <-toolStartErrCh: + if _, err := toolStart.Wait(abortCtx); err != nil { t.Fatalf("Failed waiting for tool.execution_start: %v", err) } @@ -676,10 +667,7 @@ func TestSessionE2E(t *testing.T) { } // Wait for session.idle after abort - select { - case <-sessionIdleCh: - // Session is idle - case err := <-sessionIdleErrCh: + if _, err := sessionIdle.Wait(abortCtx); err != nil { t.Fatalf("Failed waiting for session.idle after abort: %v", err) } @@ -705,26 +693,18 @@ func TestSessionE2E(t *testing.T) { } // We should be able to send another message - answerCh := make(chan *copilot.SessionEvent, 1) - answerErrCh := make(chan error, 1) - go func() { - evt, err := testharness.GetNextEventOfType(session, copilot.SessionEventTypeAssistantMessage, 60*time.Second) - if err != nil { - answerErrCh <- err - } else { - answerCh <- evt - } - }() + answerCtx, cancelAnswer := context.WithTimeout(t.Context(), 60*time.Second) + defer cancelAnswer() + answerWaiter := testharness.SubscribeToEvent(session, copilot.SessionEventTypeAssistantMessage) + defer answerWaiter.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 2+2?"}) if err != nil { t.Fatalf("Failed to send message after abort: %v", err) } - var answer *copilot.SessionEvent - select { - case answer = <-answerCh: - case err := <-answerErrCh: + answer, err := answerWaiter.Wait(answerCtx) + if err != nil { t.Fatalf("Failed waiting for assistant message after abort: %v", err) } @@ -825,7 +805,7 @@ func TestSessionE2E(t *testing.T) { // Verify the assistant response contains the expected answer. // session.idle is ephemeral and not in GetEvents(), but we already // confirmed idle via the live event handler above. - assistantMessage, err := testharness.GetFinalAssistantMessage(t.Context(), session, true) + assistantMessage, err := testharness.GetFinalAssistantMessageFromHistory(t.Context(), session) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -852,12 +832,14 @@ func TestSessionE2E(t *testing.T) { } // Session should work normally with custom config dir + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What is 1+1?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - assistantMessage, err := testharness.GetFinalAssistantMessage(t.Context(), session) + assistantMessage, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } diff --git a/go/internal/e2e/telemetry_e2e_test.go b/go/internal/e2e/telemetry_e2e_test.go index 77f8bec8ea..6c03d83819 100644 --- a/go/internal/e2e/telemetry_e2e_test.go +++ b/go/internal/e2e/telemetry_e2e_test.go @@ -52,10 +52,12 @@ func TestTelemetryE2E(t *testing.T) { } sessionID := session.SessionID + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() if _, err := session.Send(t.Context(), copilot.MessageOptions{Prompt: prompt}); err != nil { t.Fatalf("Send failed: %v", err) } - final, err := testharness.GetFinalAssistantMessage(t.Context(), session) + final, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to wait for final assistant message: %v", err) } diff --git a/go/internal/e2e/testharness/helper.go b/go/internal/e2e/testharness/helper.go index af08b2dbcc..cc508074af 100644 --- a/go/internal/e2e/testharness/helper.go +++ b/go/internal/e2e/testharness/helper.go @@ -5,7 +5,7 @@ import ( "errors" "path/filepath" "runtime" - "time" + "sync" copilot "github.com/github/copilot-sdk/go" ) @@ -29,88 +29,87 @@ func RepoPath(elem ...string) string { return filepath.Join(append([]string{repoRoot}, elem...)...) } -// GetFinalAssistantMessage waits for and returns the final assistant message from a session turn. -// If alreadyIdle is true, skip waiting for session.idle (useful for resumed sessions where the -// idle event was ephemeral and not persisted in the event history). -func GetFinalAssistantMessage(ctx context.Context, session *copilot.Session, alreadyIdle ...bool) (*copilot.SessionEvent, error) { - result := make(chan *copilot.SessionEvent, 1) - errCh := make(chan error, 1) +type eventResult struct { + event *copilot.SessionEvent + err error +} + +// EventWaiter is a synchronously installed subscription. Call Close even if the +// operation that should produce the event fails before Wait is called. +type EventWaiter struct { + result chan eventResult + once sync.Once + unsubscribe func() +} + +func (w *EventWaiter) complete(event *copilot.SessionEvent, err error) { + w.once.Do(func() { w.result <- eventResult{event: event, err: err} }) +} + +// Wait waits using only the caller's context, without imposing a default timeout. +// Events received between subscription and Wait are retained. +func (w *EventWaiter) Wait(ctx context.Context) (*copilot.SessionEvent, error) { + defer w.Close() + select { + case result := <-w.result: + return result.event, result.err + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +// Close removes the subscription and is safe to call more than once. +func (w *EventWaiter) Close() { + w.unsubscribe() +} - // Subscribe to future events +// SubscribeToFinalAssistantMessage subscribes before returning. Call it before +// Send, releasing a blocked handler, or any other operation that can finish a +// turn. Unlike durable assistant messages, session.idle is ephemeral: GetEvents +// cannot recover a missed completion. A successful wait always includes a message. +func SubscribeToFinalAssistantMessage(session *copilot.Session) *EventWaiter { + w := &EventWaiter{result: make(chan eventResult, 1)} var finalAssistantMessage *copilot.SessionEvent - unsubscribe := session.On(func(event copilot.SessionEvent) { + w.unsubscribe = session.On(func(event copilot.SessionEvent) { switch d := event.Data.(type) { case *copilot.AssistantMessageData: finalAssistantMessage = &event case *copilot.SessionIdleData: - if finalAssistantMessage != nil { - result <- finalAssistantMessage + if finalAssistantMessage == nil { + w.complete(nil, errors.New("session became idle without an assistant message")) + } else { + w.complete(finalAssistantMessage, nil) } case *copilot.SessionErrorData: - errCh <- errors.New(d.Message) + w.complete(nil, errors.New(d.Message)) } }) - defer unsubscribe() - - // Also check existing messages in case the response already arrived - isAlreadyIdle := len(alreadyIdle) > 0 && alreadyIdle[0] - go func() { - existing, err := getExistingFinalResponse(ctx, session, isAlreadyIdle) - if err != nil { - errCh <- err - return - } - if existing != nil { - result <- existing - } - }() - - select { - case msg := <-result: - return msg, nil - case err := <-errCh: - return nil, err - case <-ctx.Done(): - return nil, errors.New("timeout waiting for assistant message") - } + return w } -// GetNextEventOfType waits for and returns the next event of the specified type from a session. -func GetNextEventOfType(session *copilot.Session, eventType copilot.SessionEventType, timeout time.Duration) (*copilot.SessionEvent, error) { - result := make(chan *copilot.SessionEvent, 1) - errCh := make(chan error, 1) - - unsubscribe := session.On(func(event copilot.SessionEvent) { +// SubscribeToEvent subscribes before returning, so the triggering operation can +// run before Wait without losing events. Call Close if the operation fails. +func SubscribeToEvent(session *copilot.Session, eventType copilot.SessionEventType) *EventWaiter { + w := &EventWaiter{result: make(chan eventResult, 1)} + w.unsubscribe = session.On(func(event copilot.SessionEvent) { switch event.Type() { case eventType: - select { - case result <- &event: - default: - } + w.complete(&event, nil) case copilot.SessionEventTypeSessionError: msg := "session error" if d, ok := event.Data.(*copilot.SessionErrorData); ok { msg = d.Message } - select { - case errCh <- errors.New(msg): - default: - } + w.complete(nil, errors.New(msg)) } }) - defer unsubscribe() - - select { - case evt := <-result: - return evt, nil - case err := <-errCh: - return nil, err - case <-time.After(timeout): - return nil, errors.New("timeout waiting for event: " + string(eventType)) - } + return w } -func getExistingFinalResponse(ctx context.Context, session *copilot.Session, alreadyIdle bool) (*copilot.SessionEvent, error) { +// GetFinalAssistantMessageFromHistory reads a turn whose completion has already +// been observed independently, including after resuming an idle session. It does +// not wait for completion and must not be used as a substitute for a live waiter. +func GetFinalAssistantMessageFromHistory(ctx context.Context, session *copilot.Session) (*copilot.SessionEvent, error) { messages, err := session.GetEvents(ctx) if err != nil { return nil, err @@ -143,27 +142,11 @@ func getExistingFinalResponse(ctx context.Context, session *copilot.Session, alr } } - // Find session.idle and get last assistant message before it - sessionIdleIndex := -1 - if alreadyIdle { - sessionIdleIndex = len(currentTurnMessages) - } else { - for i, msg := range currentTurnMessages { - if msg.Type() == "session.idle" { - sessionIdleIndex = i - break - } - } - } - - if sessionIdleIndex != -1 { - // Find last assistant.message before session.idle - for i := sessionIdleIndex - 1; i >= 0; i-- { - if currentTurnMessages[i].Type() == "assistant.message" { - return ¤tTurnMessages[i], nil - } + for i := len(currentTurnMessages) - 1; i >= 0; i-- { + if currentTurnMessages[i].Type() == "assistant.message" { + return ¤tTurnMessages[i], nil } } - return nil, nil + return nil, errors.New("no assistant message in the completed turn") } diff --git a/go/internal/e2e/testharness/helper_test.go b/go/internal/e2e/testharness/helper_test.go new file mode 100644 index 0000000000..a7b378e25d --- /dev/null +++ b/go/internal/e2e/testharness/helper_test.go @@ -0,0 +1,352 @@ +package testharness + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net" + "testing" + "time" + + copilot "github.com/github/copilot-sdk/go" + "github.com/github/copilot-sdk/go/internal/jsonrpc2" + "github.com/github/copilot-sdk/go/rpc" +) + +func TestFinalAssistantMessageWaiterBeforeCompletionRPC(t *testing.T) { + for _, method := range []string{ + "session.send", + "session.tools.handlePendingToolCall", + "session.permissions.handlePendingPermissionRequest", + } { + t.Run(method, func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + assistant := copilot.SessionEvent{Data: &copilot.AssistantMessageData{MessageID: "answer", Content: "4"}} + f := newCompletionFixture(t, ctx, []copilot.SessionEvent{assistant}) + waiter := SubscribeToFinalAssistantMessage(f.session) + defer waiter.Close() + f.server.SetRequestHandler(method, f.beforeResponse(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.AssistantMessageData{MessageID: "intermediate", Content: "thinking"}}, + assistant, + {Data: &copilot.SessionIdleData{}}, + {Data: &copilot.SessionIdleData{}}, + {Data: &copilot.SessionIdleData{}}, + })) + + switch method { + case "session.send": + messageID, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "What is 2+2?"}) + if err != nil || messageID != "sent" { + t.Fatalf("Send = %q, %v; want sent, nil", messageID, err) + } + case "session.tools.handlePendingToolCall": + result, err := f.session.RPC.Tools.HandlePendingToolCall(ctx, &rpc.HandlePendingToolCallRequest{ + RequestID: "tool-request", Result: rpc.ExternalToolStringResult("4"), + }) + if err != nil || !result.Success { + t.Fatalf("HandlePendingToolCall = %+v, %v", result, err) + } + case "session.permissions.handlePendingPermissionRequest": + result, err := f.session.RPC.Permissions.HandlePendingPermissionRequest(ctx, &rpc.PermissionDecisionRequest{ + RequestID: "permission-request", Result: &rpc.PermissionDecisionApproveOnce{}, + }) + if err != nil || !result.Success { + t.Fatalf("HandlePendingPermissionRequest = %+v, %v", result, err) + } + } + + // The RPC response is withheld until all notifications have been + // processed. No goroutine has called Wait, and history contains no idle. + history, err := f.session.GetEvents(ctx) + if err != nil || len(history) != 1 || history[0].Type() != copilot.SessionEventTypeAssistantMessage { + t.Fatalf("Expected durable assistant-only history, got %+v, %v", history, err) + } + answer, err := waiter.Wait(ctx) + requireCompletionAnswer(t, answer, err, "4") + + existing, err := GetFinalAssistantMessageFromHistory(ctx, f.session) + requireCompletionAnswer(t, existing, err, "4") + + // A new waiter must not mistake a previous turn's durable answer for + // completion of a turn it never observed. + next := SubscribeToFinalAssistantMessage(f.session) + cancelled, cancelNext := context.WithCancel(ctx) + cancelNext() + if answer, err := next.Wait(cancelled); answer != nil || !errors.Is(err, context.Canceled) { + t.Fatalf("Late waiter = %+v, %v; want cancellation", answer, err) + } + }) + } +} + +func TestFinalAssistantMessageWaiterRequiresMessageAndPropagatesErrors(t *testing.T) { + for _, tc := range []struct { + name string + events []copilot.SessionEvent + wantErr string + }{ + { + name: "idle without assistant", + events: []copilot.SessionEvent{{Data: &copilot.SessionIdleData{}}}, + wantErr: "session became idle without an assistant message", + }, + { + name: "session error before idle", + events: []copilot.SessionEvent{ + {Data: &copilot.SessionErrorData{Message: "model failed"}}, + {Data: &copilot.SessionErrorData{Message: "another error"}}, + {Data: &copilot.SessionIdleData{}}, + }, + wantErr: "model failed", + }, + } { + t.Run(tc.name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + // Prior-turn messages and errors must not satisfy or fail this wait. + f := newCompletionFixture(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.UserMessageData{Content: "previous prompt"}}, + {Data: &copilot.AssistantMessageData{MessageID: "previous", Content: "previous answer"}}, + {Data: &copilot.SessionErrorData{Message: "previous error"}}, + }) + waiter := SubscribeToFinalAssistantMessage(f.session) + defer waiter.Close() + f.server.SetRequestHandler("session.send", f.beforeResponse(t, ctx, tc.events)) + if _, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "hello"}); err != nil { + t.Fatal(err) + } + answer, err := waiter.Wait(ctx) + if answer != nil || err == nil || err.Error() != tc.wantErr { + t.Fatalf("Wait = %+v, %v; want %q", answer, err, tc.wantErr) + } + }) + } +} + +func TestFinalAssistantMessageWaiterBeforeHandlerRelease(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + f := newCompletionFixture(t, ctx, nil) + waiter := SubscribeToFinalAssistantMessage(f.session) + defer waiter.Close() + complete := f.beforeResponse(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.AssistantMessageData{Content: "released"}}, + {Data: &copilot.SessionIdleData{}}, + }) + entered := make(chan struct{}) + release := make(chan struct{}) + f.server.SetRequestHandler("session.send", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + close(entered) + select { + case <-release: + return complete(params) + case <-ctx.Done(): + return nil, &jsonrpc2.Error{Code: -32000, Message: ctx.Err().Error()} + } + }) + sent := make(chan error, 1) + go func() { + _, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "use the blocked handler"}) + sent <- err + }() + select { + case <-entered: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + close(release) + select { + case err := <-sent: + if err != nil { + t.Fatal(err) + } + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + answer, err := waiter.Wait(ctx) + requireCompletionAnswer(t, answer, err, "released") +} + +func TestFinalAssistantMessageWaiterRequiresLiveIdle(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + assistant := copilot.SessionEvent{Data: &copilot.AssistantMessageData{Content: "not finished"}} + f := newCompletionFixture(t, ctx, []copilot.SessionEvent{assistant}) + waiter := SubscribeToFinalAssistantMessage(f.session) + defer waiter.Close() + f.server.SetRequestHandler("session.send", f.beforeResponse(t, ctx, []copilot.SessionEvent{assistant})) + if _, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "hello"}); err != nil { + t.Fatal(err) + } + cancelled, cancelWait := context.WithCancel(ctx) + cancelWait() + if answer, err := waiter.Wait(cancelled); answer != nil || !errors.Is(err, context.Canceled) { + t.Fatalf("Wait without idle = %+v, %v; want cancellation", answer, err) + } +} + +func TestEventWaitersBeforeAbortAndRecovery(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + f := newCompletionFixture(t, ctx, nil) + toolStart := SubscribeToEvent(f.session, copilot.SessionEventTypeToolExecutionStart) + defer toolStart.Close() + idle := SubscribeToEvent(f.session, copilot.SessionEventTypeSessionIdle) + defer idle.Close() + f.server.SetRequestHandler("session.send", f.beforeResponse(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.ToolExecutionStartData{ToolCallID: "tool", ToolName: "shell"}}, + })) + if _, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "start tool"}); err != nil { + t.Fatal(err) + } + if _, err := toolStart.Wait(ctx); err != nil { + t.Fatal(err) + } + + f.server.SetRequestHandler("session.abort", f.beforeResponse(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.SessionIdleData{}}, + })) + if err := f.session.Abort(ctx); err != nil { + t.Fatal(err) + } + if _, err := idle.Wait(ctx); err != nil { + t.Fatal(err) + } + + answerWaiter := SubscribeToEvent(f.session, copilot.SessionEventTypeAssistantMessage) + defer answerWaiter.Close() + f.server.SetRequestHandler("session.send", f.beforeResponse(t, ctx, []copilot.SessionEvent{ + {Data: &copilot.AssistantMessageData{Content: "recovered"}}, + {Data: &copilot.SessionIdleData{}}, + })) + if _, err := f.session.Send(ctx, copilot.MessageOptions{Prompt: "recover"}); err != nil { + t.Fatal(err) + } + answer, err := answerWaiter.Wait(ctx) + requireCompletionAnswer(t, answer, err, "recovered") +} + +func requireCompletionAnswer(t *testing.T, event *copilot.SessionEvent, err error, content string) { + t.Helper() + if err != nil || event == nil { + t.Fatalf("Expected assistant message, got %+v, %v", event, err) + } + if data, ok := event.Data.(*copilot.AssistantMessageData); !ok || data.Content != content { + t.Fatalf("Expected assistant content %q, got %+v", content, event.Data) + } +} + +type completionFixture struct { + session *copilot.Session + server *jsonrpc2.Client + conn net.Conn + nextFence int +} + +// Uses the same minimal JSON-RPC server approach as the client unit tests, with +// a public TCP client so these tests exercise real session event dispatch. +func newCompletionFixture(t *testing.T, ctx context.Context, history []copilot.SessionEvent) *completionFixture { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = listener.Close() }) + historyResult, err := json.Marshal(map[string]any{"events": history}) + if err != nil { + t.Fatal(err) + } + ready := make(chan *completionFixture, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + return + } + server := jsonrpc2.NewClient(conn, conn) + t.Cleanup(server.Stop) + for method, result := range map[string]string{ + "connect": `{"ok":true,"protocolVersion":3,"version":"test"}`, + "plugins.builtin.set": `{}`, + "session.create": `{"sessionId":"completion-session"}`, + "session.options.update": `{"success":true}`, + "session.detach": `{"success":true}`, + } { + server.SetRequestHandler(method, func(json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return []byte(result), nil + }) + } + server.SetRequestHandler("session.getMessages", func(json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return historyResult, nil + }) + server.Start() + ready <- &completionFixture{server: server, conn: conn} + }() + client := copilot.NewClient(&copilot.ClientOptions{ + Connection: copilot.URIConnection{URL: listener.Addr().String()}, + }) + t.Cleanup(func() { client.ForceStop() }) + session, err := client.CreateSession(ctx, &copilot.SessionConfig{ + SessionID: "completion-session", OnPermissionRequest: copilot.PermissionHandler.ApproveAll, + }) + if err != nil { + t.Fatal(err) + } + select { + case fixture := <-ready: + fixture.session = session + return fixture + case <-ctx.Done(): + t.Fatal(ctx.Err()) + return nil + } +} + +// Called after the waiters are armed. A trailing notification acknowledges that +// every preceding event has run through the session's consumer, not just reached +// its queue. All RPCs in the fixture are sequential; notifications are written +// only inside the current request handler, before its response is written. +func (f *completionFixture) beforeResponse(t *testing.T, ctx context.Context, events []copilot.SessionEvent) jsonrpc2.RequestHandler { + t.Helper() + f.nextFence++ + fenceID := fmt.Sprintf("fence-%d", f.nextFence) + delivered := make(chan struct{}, 1) + unsubscribe := f.session.On(func(event copilot.SessionEvent) { + if event.ID == fenceID { + select { + case delivered <- struct{}{}: + default: + } + } + }) + t.Cleanup(unsubscribe) + var frames bytes.Buffer + for _, event := range append(append([]copilot.SessionEvent(nil), events...), copilot.SessionEvent{ + ID: fenceID, Data: &copilot.SessionInfoData{Message: "delivery fence"}, + }) { + if event.Type() == copilot.SessionEventTypeSessionIdle { + event.Ephemeral = copilot.Bool(true) + } + data, err := json.Marshal(map[string]any{ + "jsonrpc": "2.0", "method": "session.event", + "params": map[string]any{"sessionId": f.session.SessionID, "event": event}, + }) + if err != nil { + t.Fatal(err) + } + fmt.Fprintf(&frames, "Content-Length: %d\r\n\r\n%s", len(data), data) + } + return func(json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + if _, err := f.conn.Write(frames.Bytes()); err != nil { + return nil, &jsonrpc2.Error{Code: -32000, Message: err.Error()} + } + select { + case <-delivered: + return []byte(`{"messageId":"sent","success":true}`), nil + case <-ctx.Done(): + return nil, &jsonrpc2.Error{Code: -32000, Message: ctx.Err().Error()} + } + } +} diff --git a/go/internal/e2e/tool_results_e2e_test.go b/go/internal/e2e/tool_results_e2e_test.go index 8908ffcdaf..ecd00a078f 100644 --- a/go/internal/e2e/tool_results_e2e_test.go +++ b/go/internal/e2e/tool_results_e2e_test.go @@ -37,12 +37,14 @@ func TestToolResultsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What's the weather in Paris?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -83,6 +85,8 @@ func TestToolResultsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "Check the status of the service using check_status. If it fails, say 'service is down'.", }) @@ -90,7 +94,7 @@ func TestToolResultsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -135,12 +139,14 @@ func TestToolResultsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Analyze the file main.ts for issues."}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -297,6 +303,8 @@ func TestToolResultsE2E(t *testing.T) { } }) + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "Use access_secret to get the API key. If access is denied, tell me it was 'access denied'.", }) @@ -326,7 +334,7 @@ func TestToolResultsE2E(t *testing.T) { t.Fatal("Timed out waiting for tool execution complete") } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get final assistant message: %v", err) } diff --git a/go/internal/e2e/tools_e2e_test.go b/go/internal/e2e/tools_e2e_test.go index 062d377917..75613bebdf 100644 --- a/go/internal/e2e/tools_e2e_test.go +++ b/go/internal/e2e/tools_e2e_test.go @@ -34,12 +34,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "What's the first line of README.md in this directory?"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -69,12 +71,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Use encrypt_string to encrypt this string: Hello"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -126,6 +130,8 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "First, set the current phase to 'analyzing'. Then search for items with keyword 'copilot'. Report the phase and search results.", }) @@ -133,7 +139,7 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -188,6 +194,8 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "What is my location? If you can't find out, just say 'unknown'.", }) @@ -195,7 +203,7 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -306,6 +314,8 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{ Prompt: "Perform a DB query for the 'cities' table using IDs 12 and 19, sorting ascending. " + "Reply only with lines of the form: [cityname] [population]", @@ -314,7 +324,7 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -383,12 +393,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Use safe_lookup to look up 'test123'"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -551,12 +563,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Use grep to search for the word 'hello'"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -594,12 +608,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Use encrypt_string to encrypt this string: Hello"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - answer, err := testharness.GetFinalAssistantMessage(t.Context(), session) + answer, err := finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } @@ -650,12 +666,14 @@ func TestToolsE2E(t *testing.T) { t.Fatalf("Failed to create session: %v", err) } + finalMessage := testharness.SubscribeToFinalAssistantMessage(session) + defer finalMessage.Close() _, err = session.Send(t.Context(), copilot.MessageOptions{Prompt: "Use encrypt_string to encrypt this string: Hello"}) if err != nil { t.Fatalf("Failed to send message: %v", err) } - _, err = testharness.GetFinalAssistantMessage(t.Context(), session) + _, err = finalMessage.Wait(t.Context()) if err != nil { t.Fatalf("Failed to get assistant message: %v", err) } diff --git a/nodejs/test/e2e/harness/sdkTestHelper.ts b/nodejs/test/e2e/harness/sdkTestHelper.ts index c30aa9ea6f..9cca10388b 100644 --- a/nodejs/test/e2e/harness/sdkTestHelper.ts +++ b/nodejs/test/e2e/harness/sdkTestHelper.ts @@ -53,73 +53,51 @@ function waitForChildExit(child: ChildProcess, timeoutMs: number): Promise Promise ): Promise { - // Install the live subscription (via getFutureFinalResponse) before issuing the - // existing-messages RPC so we don't miss events that arrive while that RPC is in flight. - const futurePromise = getFutureFinalResponse(session); - // We may end up returning from the existing-messages path; attach a noop handler so - // the unawaited future-response rejection doesn't surface as an unhandled rejection. - futurePromise.catch(() => {}); - - const existing = await getExistingFinalResponse(session, alreadyIdle); - if (existing) { - return existing; - } - return futurePromise; -} - -async function getExistingFinalResponse( - session: CopilotSession, - alreadyIdle: boolean = false -): Promise { - const messages = await session.getEvents(); - const finalUserMessageIndex = messages.findLastIndex((m) => m.type === "user.message"); - const currentTurnMessages = - finalUserMessageIndex < 0 ? messages : messages.slice(finalUserMessageIndex); - - const currentTurnError = currentTurnMessages.find((m) => m.type === "session.error"); - if (currentTurnError) { - const error = new Error(currentTurnError.data.message); - error.stack = currentTurnError.data.stack; - throw error; - } + type Outcome = { message: AssistantMessageEvent } | { error: Error }; + let resolveOutcome!: (outcome: Outcome) => void; + const outcomePromise = new Promise((resolve) => { + resolveOutcome = resolve; + }); + let finalAssistantMessage: AssistantMessageEvent | undefined; + + // session.idle is ephemeral: subscribe before triggering work, not after an RPC + // reply or a history lookup. Keep errors as values while the trigger is in flight. + const unsubscribe = session.on((event) => { + if (event.type === "assistant.message") { + finalAssistantMessage = event; + } else if (event.type === "session.idle" && event.data.mode !== "autopilot") { + unsubscribe(); + resolveOutcome( + finalAssistantMessage + ? { message: finalAssistantMessage } + : { + error: new Error( + "Received session.idle without a preceding assistant.message" + ), + } + ); + } else if (event.type === "session.error") { + unsubscribe(); + const error = new Error(event.data.message); + error.stack = event.data.stack; + resolveOutcome({ error }); + } + }); - const sessionIdleMessageIndex = alreadyIdle - ? currentTurnMessages.length - : currentTurnMessages.findIndex((m) => m.type === "session.idle"); - if (sessionIdleMessageIndex !== -1) { - return currentTurnMessages - .slice(0, sessionIdleMessageIndex) - .findLast((m) => m.type === "assistant.message") as AssistantMessageEvent | undefined; + try { + await trigger(); + const outcome = await outcomePromise; + if ("error" in outcome) { + throw outcome.error; + } + return outcome.message; + } finally { + unsubscribe(); } - - return undefined; -} - -function getFutureFinalResponse(session: CopilotSession): Promise { - return new Promise((resolve, reject) => { - let finalAssistantMessage: AssistantMessageEvent | undefined; - session.on((event) => { - if (event.type === "assistant.message") { - finalAssistantMessage = event; - } else if (event.type === "session.idle") { - if (!finalAssistantMessage) { - reject( - new Error("Received session.idle without a preceding assistant.message") - ); - } else { - resolve(finalAssistantMessage); - } - } else if (event.type === "session.error") { - const error = new Error(event.data.message); - error.stack = event.data.stack; - reject(error); - } - }); - }); } export async function retry( diff --git a/nodejs/test/e2e/permissions.e2e.test.ts b/nodejs/test/e2e/permissions.e2e.test.ts index d7600ed165..638ea12a33 100644 --- a/nodejs/test/e2e/permissions.e2e.test.ts +++ b/nodejs/test/e2e/permissions.e2e.test.ts @@ -15,7 +15,7 @@ import type { } from "../../src/index.js"; import { approveAll, defineTool, createAttributedPermissionResult } from "../../src/index.js"; import { createSdkTestContext, isInProcessTransport } from "./harness/sdkTestContext.js"; -import { getFinalAssistantMessage, getNextEventOfType } from "./harness/sdkTestHelper.js"; +import { withFinalAssistantMessage, getNextEventOfType } from "./harness/sdkTestHelper.js"; const isWindows = process.platform === "win32"; @@ -342,9 +342,9 @@ describe("Permission callbacks", async () => { } }); - const sessionDone = getFinalAssistantMessage(session); - - void session.send({ prompt: "Run 'echo slow_handler_test'" }); + const sessionDone = withFinalAssistantMessage(session, () => + session.send({ prompt: "Run 'echo slow_handler_test'" }) + ); // Wait for permission handler to be invoked await handlerStarted; diff --git a/nodejs/test/e2e/session.e2e.test.ts b/nodejs/test/e2e/session.e2e.test.ts index 77a1dec408..900ca4e859 100644 --- a/nodejs/test/e2e/session.e2e.test.ts +++ b/nodejs/test/e2e/session.e2e.test.ts @@ -3,7 +3,7 @@ import { describe, expect, it, onTestFinished, vi } from "vitest"; import { ParsedHttpExchange } from "../../../test/harness/replayingCapiProxy.js"; import { CopilotClient, approveAll, defineTool, RuntimeConnection } from "../../src/index.js"; import { createSdkTestContext, DEFAULT_GITHUB_TOKEN, isCI } from "./harness/sdkTestContext.js"; -import { getFinalAssistantMessage, getNextEventOfType, retry } from "./harness/sdkTestHelper.js"; +import { withFinalAssistantMessage, getNextEventOfType, retry } from "./harness/sdkTestHelper.js"; const { copilotClient: client, @@ -440,12 +440,12 @@ describe("Sessions", () => { }); expect(session2.sessionId).toBe(sessionId); - // session.idle is ephemeral and not persisted, so use alreadyIdle - // to find the assistant message from the completed session. - const answer2 = await getFinalAssistantMessage(session2, { alreadyIdle: true }); + // sendAndWait already observed idle on session1; only durable messages + // are needed to verify the completed turn survived resumption. + const messages = await session2.getEvents(); + const answer2 = messages.findLast((m) => m.type === "assistant.message"); expect(answer2?.data.content).toContain("2"); - const messages = await session2.getEvents(); expect(messages).toContainEqual(expect.objectContaining({ type: "user.message" })); expect(messages).toContainEqual(expect.objectContaining({ type: "session.resume" })); @@ -671,9 +671,9 @@ describe("Sessions", () => { expect(session.sessionId).toMatch(/^[a-f0-9-]+$/); // Session should work normally with custom config dir - await session.send({ prompt: "What is 1+1?" }); - const assistantMessage = await getFinalAssistantMessage(session); - expect(assistantMessage.data.content).toContain("2"); + const assistantMessage = await session.sendAndWait({ prompt: "What is 1+1?" }); + expect(assistantMessage).toBeDefined(); + expect(assistantMessage?.data.content).toContain("2"); }); it("should log messages at all levels and emit matching session events", async () => { @@ -969,14 +969,13 @@ describe("Send Blocking Behavior", async () => { events.push(event.type); }); - // Use a slow command so we can verify send() returns before completion - await session.send({ prompt: "Run 'sleep 2 && echo done'" }); - - // send() should return before turn completes (no session.idle yet) - expect(events).not.toContain("session.idle"); + const message = await withFinalAssistantMessage(session, async () => { + // Use a slow command so we can verify send() returns before completion. + await session.send({ prompt: "Run 'sleep 2 && echo done'" }); - // Wait for turn to complete - const message = await getFinalAssistantMessage(session); + // send() should return before turn completes (no session.idle yet). + expect(events).not.toContain("session.idle"); + }); expect(message.data.content).toContain("done"); expect(events).toContain("session.idle"); diff --git a/nodejs/test/e2e/session_lifecycle.e2e.test.ts b/nodejs/test/e2e/session_lifecycle.e2e.test.ts index fae8782736..fe7ba645e0 100644 --- a/nodejs/test/e2e/session_lifecycle.e2e.test.ts +++ b/nodejs/test/e2e/session_lifecycle.e2e.test.ts @@ -85,7 +85,7 @@ describe("Session Lifecycle", async () => { const messages = await session.getEvents(); expect(messages.length).toBeGreaterThan(0); - // Should have at least session.start, user.message, assistant.message, session.idle + // History contains durable messages, not the ephemeral session.idle event. const types = messages.map((m: SessionEvent) => m.type); expect(types).toContain("session.start"); expect(types).toContain("user.message"); diff --git a/nodejs/test/e2e/telemetry.e2e.test.ts b/nodejs/test/e2e/telemetry.e2e.test.ts index 66a0bb8cef..9fb89fc0d7 100644 --- a/nodejs/test/e2e/telemetry.e2e.test.ts +++ b/nodejs/test/e2e/telemetry.e2e.test.ts @@ -8,7 +8,6 @@ import { describe, expect, it } from "vitest"; import { z } from "zod"; import { approveAll, defineTool, RuntimeConnection } from "../../src/index.js"; import { createSdkTestContext } from "./harness/sdkTestContext.js"; -import { getFinalAssistantMessage } from "./harness/sdkTestHelper.js"; interface TelemetryEntry { type?: string; @@ -85,10 +84,9 @@ describe("Telemetry export", async () => { ], }); - await session.send({ prompt }); - const assistantMessage = await getFinalAssistantMessage(session); + const assistantMessage = await session.sendAndWait({ prompt }, 90_000); expect(assistantMessage).toBeDefined(); - expect(assistantMessage.data.content ?? "").toContain("TELEMETRY_E2E_DONE"); + expect(assistantMessage?.data.content ?? "").toContain("TELEMETRY_E2E_DONE"); await session.disconnect(); await client.stop(); diff --git a/nodejs/test/session-send-and-wait.test.ts b/nodejs/test/session-send-and-wait.test.ts index 4ee2e8eebc..d88b80c896 100644 --- a/nodejs/test/session-send-and-wait.test.ts +++ b/nodejs/test/session-send-and-wait.test.ts @@ -2,10 +2,11 @@ * Copyright (c) Microsoft Corporation. All rights reserved. *--------------------------------------------------------------------------------------------*/ -import { describe, expect, it, onTestFinished } from "vitest"; +import { describe, expect, it, onTestFinished, vi } from "vitest"; import type { MessageConnection } from "vscode-jsonrpc/node.js"; import { CopilotSession } from "../src/session.js"; -import type { SessionEvent } from "../src/generated/session-events.js"; +import type { AssistantMessageEvent, SessionEvent } from "../src/generated/session-events.js"; +import { withFinalAssistantMessage } from "./e2e/harness/sdkTestHelper.js"; function sessionEvent( type: "session.idle", @@ -32,32 +33,49 @@ function errorEvent(message: string): SessionEvent { } as SessionEvent; } -function controlledSession(): { - session: CopilotSession; - sendStarted: Promise; - resolveSend: () => void; - rejectSend: (error: Error) => void; -} { +function assistantMessage(content: string): AssistantMessageEvent { + return { + type: "assistant.message", + id: "assistant-1", + parentId: null, + timestamp: new Date().toISOString(), + data: { messageId: "message-1", content }, + }; +} + +function controlledSession( + history: SessionEvent[] = [], + onSend?: (session: CopilotSession) => void +) { let resolveSendRequest: ((value: unknown) => void) | undefined; let rejectSendRequest: ((error: Error) => void) | undefined; let markSendStarted: () => void; const sendStarted = new Promise((resolve) => { markSendStarted = resolve; }); - const connection = { - sendRequest: () => - new Promise((resolve, reject) => { - resolveSendRequest = resolve; - rejectSendRequest = reject; - markSendStarted(); - }), - } as unknown as MessageConnection; + const sendRequest = vi.fn((method: string) => { + if (method === "session.getMessages") { + return Promise.resolve({ events: history }); + } + if (method !== "session.send") { + throw new Error(`Unexpected RPC: ${method}`); + } + return new Promise((resolve, reject) => { + resolveSendRequest = resolve; + rejectSendRequest = reject; + markSendStarted(); + onSend?.(session); + }); + }); + const connection = { sendRequest } as unknown as MessageConnection; + const session = new CopilotSession("session-1", connection); return { - session: new CopilotSession("session-1", connection), + session, + sendRequest, sendStarted, resolveSend: () => resolveSendRequest?.({ messageId: "msg-1" }), - rejectSend: (error) => rejectSendRequest?.(error), + rejectSend: (error: Error) => rejectSendRequest?.(error), }; } @@ -156,3 +174,139 @@ describe("sendAndWait", () => { await expect(errorFirstPending).rejects.toThrow("first error"); }); }); + +describe("completion subscriptions", () => { + it.each(["sendAndWait", "withFinalAssistantMessage"] as const)( + "%s captures final output and ephemeral idle before the send reply", + async (waiter) => { + const finalMessage = assistantMessage("done"); + const history: SessionEvent[] = []; + const { session, sendRequest, sendStarted, resolveSend } = controlledSession( + history, + (session) => { + const events = [ + assistantMessage("working"), + finalMessage, + sessionEvent("session.idle"), + ]; + for (const event of events) { + if (!event.ephemeral) { + history.push(event); + } + session._dispatchEvent(event); + } + } + ); + + let completed = false; + const pending = + waiter === "sendAndWait" + ? session.sendAndWait({ prompt: "hi" }) + : withFinalAssistantMessage(session, () => session.send({ prompt: "hi" })); + const observed = pending.then((message) => { + completed = true; + return message; + }); + await sendStarted; + expect(completed).toBe(false); + expect(sendRequest).toHaveBeenCalledTimes(1); + + // Only durable messages can be retrieved, even though idle was delivered. + const stored = await session.getEvents(); + expect(stored).toEqual(history); + expect(stored.map((event) => event.type)).toEqual([ + "assistant.message", + "assistant.message", + ]); + + resolveSend(); + await expect(observed).resolves.toBe(finalMessage); + } + ); +}); + +describe("withFinalAssistantMessage", () => { + it("does not accept a prior turn's message when the new turn has no assistant output", async () => { + const { session, sendStarted, resolveSend } = controlledSession( + [assistantMessage("old turn")], + (session) => session._dispatchEvent(sessionEvent("session.idle")) + ); + const pending = withFinalAssistantMessage(session, () => session.send({ prompt: "hi" })); + const outcome = expect(pending).rejects.toThrow( + "Received session.idle without a preceding assistant.message" + ); + + await sendStarted; + resolveSend(); + await outcome; + }); + + it("does not complete on an assistant message or an autopilot continuation", async () => { + const { session, sendStarted, resolveSend } = controlledSession(); + const pending = withFinalAssistantMessage(session, () => session.send({ prompt: "hi" })); + await sendStarted; + resolveSend(); + + session._dispatchEvent(assistantMessage("continuing")); + session._dispatchEvent(sessionEvent("session.idle", { mode: "autopilot" })); + + const finalMessage = assistantMessage("done"); + session._dispatchEvent(finalMessage); + session._dispatchEvent(sessionEvent("session.idle", { mode: "interactive" })); + await expect(pending).resolves.toBe(finalMessage); + }); + + it.each(["idle", "session.error", "send rejection", "trigger throw"] as const)( + "removes the completion subscription after %s", + async (outcome) => { + const { session, sendStarted, resolveSend, rejectSend } = controlledSession(); + const originalOn = session.on.bind(session); + const unsubscribe = vi.fn<() => void>(); + vi.spyOn(session, "on").mockImplementation((handler) => { + unsubscribe.mockImplementation(originalOn(handler)); + return unsubscribe; + }); + + const pending = withFinalAssistantMessage(session, () => { + if (outcome === "trigger throw") { + throw new Error("trigger failed"); + } + return session.send({ prompt: "hi" }); + }); + const finalMessage = assistantMessage("done"); + const errorMessage = + outcome === "session.error" + ? "session failed" + : outcome === "send rejection" + ? "send failed" + : "trigger failed"; + const expected = + outcome === "idle" + ? expect(pending).resolves.toBe(finalMessage) + : expect(pending).rejects.toThrow(errorMessage); + + if (outcome !== "trigger throw") { + await sendStarted; + } + + if (outcome === "idle") { + session._dispatchEvent(finalMessage); + session._dispatchEvent(sessionEvent("session.idle")); + // Later events must not replace an already observed terminal outcome. + session._dispatchEvent(assistantMessage("next turn")); + session._dispatchEvent(errorEvent("later error")); + resolveSend(); + } else if (outcome === "session.error") { + session._dispatchEvent(errorEvent("session failed")); + session._dispatchEvent(sessionEvent("session.idle")); + resolveSend(); + } else if (outcome === "send rejection") { + session._dispatchEvent(errorEvent("session failed")); + rejectSend(new Error("send failed")); + } + + await expected; + expect(unsubscribe).toHaveBeenCalled(); + } + ); +}); diff --git a/python/.gitignore b/python/.gitignore index 671fe9a8bb..122f34c543 100644 --- a/python/.gitignore +++ b/python/.gitignore @@ -49,6 +49,7 @@ coverage.xml *.py,cover .hypothesis/ .pytest_cache/ +.pytest-diagnostics/ cover/ # Translations diff --git a/python/README.md b/python/README.md index 2d1d8ca153..3b7d937c44 100644 --- a/python/README.md +++ b/python/README.md @@ -1233,3 +1233,13 @@ cd python uv sync uv run pytest ``` + +Signal-based E2E failures from `pytest-timeout` include an **Async timeout diagnostics** report +section with suspended coroutine await chains, pending JSON-RPC request IDs and +methods, session/transport state, and Python thread stacks. The same report is +saved under `python/.pytest-diagnostics/`. macOS in-process timeouts also capture a +one-second native thread sample there. RPC payloads and arbitrary frame locals +are not included. After recording the timeout, the harness cancels only the +abandoned test coroutine so it does not retain locks needed by later fixture +cleanup. The original timeout failure is retained; this does not abort native +runtime work or repair a missing RPC response. diff --git a/python/_session_test_helpers.py b/python/_session_test_helpers.py new file mode 100644 index 0000000000..77eec022a7 --- /dev/null +++ b/python/_session_test_helpers.py @@ -0,0 +1,64 @@ +"""Live-event test helpers without E2E runtime or proxy initialization.""" + +import asyncio +from collections.abc import Callable + +from copilot.session import CopilotSession +from copilot.session_events import SessionErrorData, SessionEvent + + +def wait_for_event( + session: CopilotSession, + predicate: Callable[[SessionEvent], bool], + timeout: float = 30.0, + *, + fail_on_session_error: bool = False, +) -> asyncio.Task[SessionEvent]: + """Subscribe synchronously and return a task for the next matching live event. + + Call before the operation that emits the event, then await or cancel the task. + Merely scheduling an async subscriber would leave a race before it starts. + In particular, session.idle is ephemeral and cannot be recovered from history. + + Only predicate matches complete the wait by default. Set fail_on_session_error + to also fail on unmatched session errors when the caller requires that policy. + """ + loop = asyncio.get_running_loop() + result_future: asyncio.Future[SessionEvent] = loop.create_future() + + def on_event(event: SessionEvent) -> None: + if result_future.done(): + return + + if predicate(event): + result_future.set_result(event) + elif fail_on_session_error and isinstance(event.data, SessionErrorData): + result_future.set_exception(RuntimeError(event.data.message or "session error")) + + unsubscribe = session.on(on_event) + + async def wait() -> SessionEvent: + return await asyncio.wait_for(result_future, timeout=timeout) + + def cleanup(_task: asyncio.Task[SessionEvent]) -> None: + unsubscribe() + result_future.cancel() + if not result_future.cancelled(): + result_future.exception() + + task = loop.create_task(wait()) + # A task cancelled before its first step never executes a coroutine's finally. + task.add_done_callback(cleanup) + return task + + +def get_next_event_of_type( + session: CopilotSession, event_type: str, timeout: float = 30.0 +) -> asyncio.Task[SessionEvent]: + """Subscribe before an operation; fail on session errors unless waiting for that event.""" + return wait_for_event( + session, + lambda event: event.type.value == event_type, + timeout, + fail_on_session_error=True, + ) diff --git a/python/e2e/conftest.py b/python/e2e/conftest.py index d61fe0d875..309c090c3a 100644 --- a/python/e2e/conftest.py +++ b/python/e2e/conftest.py @@ -10,6 +10,8 @@ import copilot._cli_download as cli_download from .testharness import E2ETestContext, is_inprocess_transport +from .timeout_diagnostics import add_timeout_diagnostics, cancel_timed_out_test +from .timeout_diagnostics import pytest_timeout_set_timer as pytest_timeout_set_timer # Host-side auth resolution ranks HMAC above the GitHub token, so an ambient # COPILOT_HMAC_KEY (CI sets one as a job-level credential) would be picked over @@ -34,7 +36,9 @@ def pytest_runtest_makereport(item, call): """Track test failures to avoid writing corrupted snapshots.""" outcome = yield rep = outcome.get_result() + add_timeout_diagnostics(item, call, rep) if rep.when == "call" and rep.failed: + cancel_timed_out_test(call) # Store on the item's stash so the fixture can access it item.session.stash.setdefault("any_test_failed", False) item.session.stash["any_test_failed"] = True diff --git a/python/e2e/test_mode_handlers_e2e.py b/python/e2e/test_mode_handlers_e2e.py index d5182c453d..d681b403c9 100644 --- a/python/e2e/test_mode_handlers_e2e.py +++ b/python/e2e/test_mode_handlers_e2e.py @@ -20,7 +20,7 @@ SessionModelChangeData, ) -from .testharness import E2ETestContext +from .testharness import E2ETestContext, wait_for_event pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -55,22 +55,6 @@ async def mode_ctx(ctx: E2ETestContext): return ctx -async def _wait_for_event(session, predicate, timeout: float = 30.0): - """Wait for the first session event matching predicate.""" - loop = asyncio.get_event_loop() - fut: asyncio.Future = loop.create_future() - - def on_event(event): - if not fut.done() and predicate(event): - fut.set_result(event) - - unsubscribe = session.on(on_event) - try: - return await asyncio.wait_for(fut, timeout=timeout) - finally: - unsubscribe() - - class TestModeHandlers: async def test_should_invoke_exit_plan_mode_handler_when_model_uses_tool( self, mode_ctx: E2ETestContext @@ -92,27 +76,23 @@ async def on_exit_plan_mode_request(request, invocation): on_exit_plan_mode_request=on_exit_plan_mode_request, ) - try: - requested_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: ( - isinstance(event.data, ExitPlanModeRequestedData) - and event.data.summary == PLAN_SUMMARY - ), - ) - ) - completed_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: ( - isinstance(event.data, ExitPlanModeCompletedData) - and event.data.approved is True - and event.data.selected_action == ExitPlanModeAction.INTERACTIVE - ), - ) - ) + requested_event = wait_for_event( + session, + lambda event: ( + isinstance(event.data, ExitPlanModeRequestedData) + and event.data.summary == PLAN_SUMMARY + ), + ) + completed_event = wait_for_event( + session, + lambda event: ( + isinstance(event.data, ExitPlanModeCompletedData) + and event.data.approved is True + and event.data.selected_action == ExitPlanModeAction.INTERACTIVE + ), + ) + try: await session.rpc.mode.set(ModeSetRequest(mode=SessionMode.PLAN)) response = await session.send_and_wait(PLAN_PROMPT) @@ -132,6 +112,9 @@ async def on_exit_plan_mode_request(request, invocation): assert completed.data.feedback == "Approved by the Python E2E test" assert response is not None finally: + requested_event.cancel() + completed_event.cancel() + await asyncio.gather(requested_event, completed_event, return_exceptions=True) await session.disconnect() async def test_should_invoke_auto_mode_switch_handler_when_rate_limited( @@ -150,42 +133,34 @@ async def on_auto_mode_switch_request(request, invocation): on_auto_mode_switch_request=on_auto_mode_switch_request, ) - try: - requested_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: ( - isinstance(event.data, AutoModeSwitchRequestedData) - and event.data.error_code == "user_weekly_rate_limited" - and event.data.retry_after_seconds == 1 - ), - ) - ) - completed_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: ( - isinstance(event.data, AutoModeSwitchCompletedData) - and event.data.response == AutoModeSwitchResponse.YES - ), - ) - ) - model_change_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: ( - isinstance(event.data, SessionModelChangeData) - and event.data.cause == "rate_limit_auto_switch" - ), - ) - ) - idle_event = asyncio.create_task( - _wait_for_event( - session, - lambda event: isinstance(event.data, SessionIdleData), - ) - ) + requested_event = wait_for_event( + session, + lambda event: ( + isinstance(event.data, AutoModeSwitchRequestedData) + and event.data.error_code == "user_weekly_rate_limited" + and event.data.retry_after_seconds == 1 + ), + ) + completed_event = wait_for_event( + session, + lambda event: ( + isinstance(event.data, AutoModeSwitchCompletedData) + and event.data.response == AutoModeSwitchResponse.YES + ), + ) + model_change_event = wait_for_event( + session, + lambda event: ( + isinstance(event.data, SessionModelChangeData) + and event.data.cause == "rate_limit_auto_switch" + ), + ) + idle_event = wait_for_event( + session, + lambda event: isinstance(event.data, SessionIdleData), + ) + try: message_id = await session.send(AUTO_MODE_PROMPT) assert message_id @@ -206,4 +181,8 @@ async def on_auto_mode_switch_request(request, invocation): assert request["errorCode"] == "user_weekly_rate_limited" assert request["retryAfterSeconds"] == 1 finally: + waiters = [requested_event, completed_event, model_change_event, idle_event] + for waiter in waiters: + waiter.cancel() + await asyncio.gather(*waiters, return_exceptions=True) await session.disconnect() diff --git a/python/e2e/test_multi_client_e2e.py b/python/e2e/test_multi_client_e2e.py index 1938ddfe89..4840155cc6 100644 --- a/python/e2e/test_multi_client_e2e.py +++ b/python/e2e/test_multi_client_e2e.py @@ -22,7 +22,7 @@ from copilot.session import PermissionHandler, PermissionNoResult from copilot.tools import ToolInvocation -from .testharness import get_final_assistant_message +from .testharness import wait_for_event from .testharness.proxy import CapiProxy pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -187,25 +187,6 @@ async def configure_multi_test(request, mctx): yield -def wait_for_event(session, predicate, timeout: float = 30.0): - loop = asyncio.get_running_loop() - future = loop.create_future() - - def on_event(event): - if not future.done() and predicate(event): - future.set_result(event) - - unsubscribe = session.on(on_event) - - async def wait(): - try: - return await asyncio.wait_for(future, timeout=timeout) - finally: - unsubscribe() - - return loop.create_task(wait()) - - class TestMultiClientBroadcast: async def test_both_clients_see_tool_request_and_completion_events( self, mctx: MultiClientContext @@ -245,11 +226,12 @@ def magic_number(params: SeedParams, invocation: ToolInvocation) -> str: waiters = [client1_requested, client2_requested, client1_completed, client2_completed] # Send a prompt that triggers the custom tool - await session1.send( - "Use the magic_number tool with seed 'hello' and tell me the result" - ) # Use a longer timeout: first multi-client TCP test on Windows CI needs extra time - response = await get_final_assistant_message(session1, timeout=30.0) + response = await session1.send_and_wait( + "Use the magic_number tool with seed 'hello' and tell me the result", + timeout=30.0, + ) + assert response is not None assert "MAGIC_hello_42" in (response.data.content or "") # Both clients should have seen the external_tool.requested and completed events @@ -296,8 +278,10 @@ async def test_one_client_approves_permission_and_both_see_the_result( waiters = [client1_requested, client2_requested, client1_completed, client2_completed] # Send a prompt that triggers a write operation (requires permission) - await session1.send("Create a file called hello.txt containing the text 'hello world'") - response = await get_final_assistant_message(session1) + response = await session1.send_and_wait( + "Create a file called hello.txt containing the text 'hello world'", timeout=10.0 + ) + assert response is not None assert response.data.content # Client 1 should have handled permission requests @@ -415,16 +399,18 @@ def currency_lookup(params: CountryCodeParams, invocation: ToolInvocation) -> st ) # Send prompts sequentially to avoid nondeterministic tool_call ordering - await session1.send( - "Use the city_lookup tool with countryCode 'US' and tell me the result." + response1 = await session1.send_and_wait( + "Use the city_lookup tool with countryCode 'US' and tell me the result.", + timeout=10.0, ) - response1 = await get_final_assistant_message(session1) + assert response1 is not None assert "CITY_FOR_US" in (response1.data.content or "") - await session1.send( - "Now use the currency_lookup tool with countryCode 'US' and tell me the result." + response2 = await session1.send_and_wait( + "Now use the currency_lookup tool with countryCode 'US' and tell me the result.", + timeout=10.0, ) - response2 = await get_final_assistant_message(session1) + assert response2 is not None assert "CURRENCY_FOR_US" in (response2.data.content or "") await session2.disconnect() @@ -464,12 +450,16 @@ def ephemeral_tool(params: InputParams, invocation: ToolInvocation) -> str: # Verify both tools work before disconnect. # Sequential prompts avoid nondeterministic tool_call ordering. - await session1.send("Use the stable_tool with input 'test1' and tell me the result.") - stable_response = await get_final_assistant_message(session1) + stable_response = await session1.send_and_wait( + "Use the stable_tool with input 'test1' and tell me the result.", timeout=10.0 + ) + assert stable_response is not None assert "STABLE_test1" in (stable_response.data.content or "") - await session1.send("Use the ephemeral_tool with input 'test2' and tell me the result.") - ephemeral_response = await get_final_assistant_message(session1) + ephemeral_response = await session1.send_and_wait( + "Use the ephemeral_tool with input 'test2' and tell me the result.", timeout=10.0 + ) + assert ephemeral_response is not None assert "EPHEMERAL_test2" in (ephemeral_response.data.content or "") # Force disconnect client 2 without destroying the shared session @@ -487,12 +477,13 @@ def ephemeral_tool(params: InputParams, invocation: ToolInvocation) -> str: ) # Now only stable_tool should be available - await session1.send( + after_response = await session1.send_and_wait( "Use the stable_tool with input 'still_here'." " Also try using ephemeral_tool" - " if it is available." + " if it is available.", + timeout=10.0, ) - after_response = await get_final_assistant_message(session1) + assert after_response is not None assert "STABLE_still_here" in (after_response.data.content or "") # ephemeral_tool should NOT have produced a result assert "EPHEMERAL_" not in (after_response.data.content or "") diff --git a/python/e2e/test_pending_work_resume_e2e.py b/python/e2e/test_pending_work_resume_e2e.py index 5b6d978f31..0a09c37c59 100644 --- a/python/e2e/test_pending_work_resume_e2e.py +++ b/python/e2e/test_pending_work_resume_e2e.py @@ -11,7 +11,6 @@ from __future__ import annotations import asyncio -from typing import Any import pytest @@ -23,9 +22,16 @@ SessionsCheckInUseRequest, ) from copilot.session import PermissionHandler +from copilot.session_events import ExternalToolRequestedData, PermissionRequestedData from copilot.tools import Tool, ToolInvocation, ToolResult -from .testharness import DEFAULT_GITHUB_TOKEN, E2ETestContext, wait_for_condition +from .testharness import ( + DEFAULT_GITHUB_TOKEN, + E2ETestContext, + get_next_event_of_type, + wait_for_condition, + wait_for_event, +) pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -74,53 +80,6 @@ async def wrapped(invocation: ToolInvocation) -> ToolResult: ) -async def _wait_for_external_tool_requests( - session, tool_names: list[str], timeout: float = PENDING_WORK_TIMEOUT -) -> dict[str, Any]: - """Wait for ExternalToolRequested events for the named tools.""" - expected = set(tool_names) - seen: dict[str, Any] = {} - completed: asyncio.Future = asyncio.get_event_loop().create_future() - - def on_event(event): - if completed.done(): - return - if event.type.value == "external_tool.requested": - tool_name = event.data.tool_name - if tool_name in expected and tool_name not in seen: - seen[tool_name] = event - if len(seen) == len(expected): - completed.set_result(dict(seen)) - elif event.type.value == "session.error": - msg = event.data.message or "session error" - completed.set_exception(RuntimeError(msg)) - - unsubscribe = session.on(on_event) - try: - return await asyncio.wait_for(completed, timeout=timeout) - finally: - unsubscribe() - - -async def _wait_for_permission_request(session, timeout: float = PENDING_WORK_TIMEOUT) -> Any: - completed: asyncio.Future = asyncio.get_event_loop().create_future() - - def on_event(event): - if completed.done(): - return - if event.type.value == "permission.requested": - completed.set_result(event) - elif event.type.value == "session.error": - msg = event.data.message or "session error" - completed.set_exception(RuntimeError(msg)) - - unsubscribe = session.on(on_event) - try: - return await asyncio.wait_for(completed, timeout=timeout) - finally: - unsubscribe() - - async def _safe_force_stop(client: CopilotClient) -> None: try: await client.stop() @@ -160,13 +119,16 @@ def original_tool_handler(args): ) session_id = session1.session_id + permission_event_task = get_next_event_of_type( + session1, "permission.requested", timeout=PENDING_WORK_TIMEOUT + ) try: - permission_event_task = asyncio.create_task(_wait_for_permission_request(session1)) await session1.send( "Use resume_permission_tool with value 'alpha', then reply with the result." ) _ = await captured_request permission_event = await permission_event_task + assert isinstance(permission_event.data, PermissionRequestedData) # Force-stop the suspended client without releasing the in-flight # permission so the request remains pending in the runtime. @@ -204,6 +166,8 @@ def resumed_tool_handler(args): finally: await _safe_force_stop(resumed_client) finally: + permission_event_task.cancel() + await asyncio.gather(permission_event_task, return_exceptions=True) if not release_original.done(): release_original.set_result(PermissionDecisionUserNotAvailable()) finally: @@ -237,14 +201,21 @@ async def blocking_external_tool(args): ) session_id = session1.session_id + tool_request_task = wait_for_event( + session1, + lambda event: ( + isinstance(event.data, ExternalToolRequestedData) + and event.data.tool_name == "resume_external_tool" + ), + timeout=PENDING_WORK_TIMEOUT, + fail_on_session_error=True, + ) try: - tool_request_task = asyncio.create_task( - _wait_for_external_tool_requests(session1, ["resume_external_tool"]) - ) await session1.send( "Use resume_external_tool with value 'beta', then reply with the result." ) - tool_events = await tool_request_task + tool_event = await tool_request_task + assert isinstance(tool_event.data, ExternalToolRequestedData) assert (await asyncio.wait_for(tool_started, PENDING_WORK_TIMEOUT)) == "beta" await suspended_client.force_stop() @@ -263,7 +234,7 @@ async def blocking_external_tool(args): tool_result = await session2.rpc.tools.handle_pending_tool_call( HandlePendingToolCallRequest( - request_id=tool_events["resume_external_tool"].data.request_id, + request_id=tool_event.data.request_id, result="EXTERNAL_RESUMED_BETA", ) ) @@ -273,6 +244,8 @@ async def blocking_external_tool(args): finally: await _safe_force_stop(resumed_client) finally: + tool_request_task.cancel() + await asyncio.gather(tool_request_task, return_exceptions=True) if not release_original.done(): release_original.set_result("ORIGINAL_SHOULD_NOT_WIN") finally: @@ -315,17 +288,32 @@ async def tool_b(args): ) session_id = session1.session_id + tool_a_request = wait_for_event( + session1, + lambda event: ( + isinstance(event.data, ExternalToolRequestedData) + and event.data.tool_name == "pending_lookup_a" + ), + timeout=PENDING_WORK_TIMEOUT, + fail_on_session_error=True, + ) + tool_b_request = wait_for_event( + session1, + lambda event: ( + isinstance(event.data, ExternalToolRequestedData) + and event.data.tool_name == "pending_lookup_b" + ), + timeout=PENDING_WORK_TIMEOUT, + fail_on_session_error=True, + ) try: - tool_requests_task = asyncio.create_task( - _wait_for_external_tool_requests( - session1, ["pending_lookup_a", "pending_lookup_b"] - ) - ) await session1.send( "Call pending_lookup_a with value 'alpha' and " "pending_lookup_b with value 'beta', then reply with both results." ) - tool_events = await tool_requests_task + tool_a_event, tool_b_event = await asyncio.gather(tool_a_request, tool_b_request) + assert isinstance(tool_a_event.data, ExternalToolRequestedData) + assert isinstance(tool_b_event.data, ExternalToolRequestedData) await asyncio.wait_for( asyncio.gather(tool_a_started, tool_b_started), PENDING_WORK_TIMEOUT ) @@ -348,14 +336,14 @@ async def tool_b(args): result_b = await session2.rpc.tools.handle_pending_tool_call( HandlePendingToolCallRequest( - request_id=tool_events["pending_lookup_b"].data.request_id, + request_id=tool_b_event.data.request_id, result="PARALLEL_B_BETA", ) ) assert result_b.success result_a = await session2.rpc.tools.handle_pending_tool_call( HandlePendingToolCallRequest( - request_id=tool_events["pending_lookup_a"].data.request_id, + request_id=tool_a_event.data.request_id, result="PARALLEL_A_ALPHA", ) ) @@ -365,6 +353,9 @@ async def tool_b(args): finally: await _safe_force_stop(resumed_client) finally: + tool_a_request.cancel() + tool_b_request.cancel() + await asyncio.gather(tool_a_request, tool_b_request, return_exceptions=True) if not release_a.done(): release_a.set_result("ORIGINAL_A_SHOULD_NOT_WIN") if not release_b.done(): @@ -478,14 +469,21 @@ async def blocking_external_tool(args): ) session_id = session1.session_id + tool_request_task = wait_for_event( + session1, + lambda event: ( + isinstance(event.data, ExternalToolRequestedData) + and event.data.tool_name == "resume_external_tool" + ), + timeout=PENDING_WORK_TIMEOUT, + fail_on_session_error=True, + ) try: - tool_request_task = asyncio.create_task( - _wait_for_external_tool_requests(session1, ["resume_external_tool"]) - ) await session1.send( "Use resume_external_tool with value 'beta', then reply with the result." ) - tool_events = await tool_request_task + tool_event = await tool_request_task + assert isinstance(tool_event.data, ExternalToolRequestedData) assert (await asyncio.wait_for(tool_started, PENDING_WORK_TIMEOUT)) == "beta" if disconnect_original_client: @@ -564,7 +562,7 @@ async def resumed_external_tool(args): # session should still be healthy for new turns. tool_result = await session2.rpc.tools.handle_pending_tool_call( HandlePendingToolCallRequest( - request_id=tool_events["resume_external_tool"].data.request_id, + request_id=tool_event.data.request_id, result="EXTERNAL_RESUMED_BETA", ) ) @@ -582,6 +580,8 @@ async def resumed_external_tool(args): finally: await _safe_force_stop(resumed_client) finally: + tool_request_task.cancel() + await asyncio.gather(tool_request_task, return_exceptions=True) if not release_original.done(): release_original.set_result("ORIGINAL_SHOULD_NOT_WIN") await _safe_force_stop(suspended_client) diff --git a/python/e2e/test_permissions_e2e.py b/python/e2e/test_permissions_e2e.py index c6c644c934..472b958451 100644 --- a/python/e2e/test_permissions_e2e.py +++ b/python/e2e/test_permissions_e2e.py @@ -321,9 +321,10 @@ def on_event(event): add_event("tool-complete", event.data.tool_call_id) unsubscribe = session.on(on_event) + response_task = asyncio.create_task( + session.send_and_wait("Run 'echo slow_handler_test'", timeout=60.0) + ) try: - asyncio.ensure_future(session.send("Run 'echo slow_handler_test'")) - await asyncio.wait_for(handler_entered, timeout=30.0) target_id = await asyncio.wait_for(target_tool_call_id, timeout=30.0) @@ -334,9 +335,7 @@ def on_event(event): release_handler.set_result(True) - from .testharness.helper import get_final_assistant_message - - message = await get_final_assistant_message(session, timeout=60.0) + message = await response_task perm_start = next( ( @@ -386,6 +385,8 @@ def on_event(event): finally: if not release_handler.done(): release_handler.set_result(True) + response_task.cancel() + await asyncio.gather(response_task, return_exceptions=True) unsubscribe() await session.disconnect() diff --git a/python/e2e/test_rpc_event_side_effects_e2e.py b/python/e2e/test_rpc_event_side_effects_e2e.py index ce3951aacd..c759be2657 100644 --- a/python/e2e/test_rpc_event_side_effects_e2e.py +++ b/python/e2e/test_rpc_event_side_effects_e2e.py @@ -35,22 +35,6 @@ pytestmark = pytest.mark.asyncio(loop_scope="module") -async def _wait_for_event(session, predicate, timeout: float = 15.0): - """Wait for the first session event matching predicate.""" - loop = asyncio.get_event_loop() - fut: asyncio.Future = loop.create_future() - - def on_event(event): - if not fut.done() and predicate(event): - fut.set_result(event) - - unsub = session.on(on_event) - try: - return await asyncio.wait_for(fut, timeout=timeout) - finally: - unsub() - - class TestRpcEventSideEffects: async def test_should_emit_mode_changed_event_when_mode_set(self, ctx: E2ETestContext): session = await ctx.client.create_session( diff --git a/python/e2e/test_rpc_server_e2e.py b/python/e2e/test_rpc_server_e2e.py index 83dc4a01e2..044121e0cb 100644 --- a/python/e2e/test_rpc_server_e2e.py +++ b/python/e2e/test_rpc_server_e2e.py @@ -11,9 +11,10 @@ from datetime import UTC, datetime from pathlib import Path +import httpx import pytest -from copilot import CopilotClient, RuntimeConnection +from copilot import CopilotClient, CopilotRequestContext, CopilotRequestHandler, RuntimeConnection from copilot.rpc import ( AccountGetQuotaRequest, AgentsDiscoverRequest, @@ -55,7 +56,14 @@ ) from copilot.session import PermissionHandler -from .testharness import E2ETestContext, is_inprocess_transport, wait_for_condition +from ._copilot_request_helpers import ( + SYNTHETIC_TEXT, + assistant_text, + build_inference_response, + build_non_inference_response, + is_inference_url, +) +from .testharness import E2ETestContext, is_inprocess_transport pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -89,7 +97,12 @@ async def authed_ctx(ctx: E2ETestContext): return ctx -def _make_authed_client(ctx: E2ETestContext, token: str) -> CopilotClient: +def _make_authed_client( + ctx: E2ETestContext, + token: str, + *, + request_handler: CopilotRequestHandler | None = None, +) -> CopilotClient: env = ctx.get_env() env["COPILOT_DEBUG_GITHUB_API_URL"] = ctx.proxy_url return CopilotClient( @@ -97,9 +110,21 @@ def _make_authed_client(ctx: E2ETestContext, token: str) -> CopilotClient: working_directory=ctx.work_dir, env=env, github_token=token, + request_handler=request_handler, ) +class _PersistedSessionRequestHandler(CopilotRequestHandler): + """Complete the metadata fixture's real turn without a live inference request.""" + + async def send_request( + self, request: httpx.Request, ctx: CopilotRequestContext + ) -> httpx.Response: + if is_inference_url(str(request.url)): + return build_inference_response(request) + return build_non_inference_response(str(request.url), supported_endpoints=["/responses"]) + + def _make_client_with_env(ctx: E2ETestContext, env_overrides: dict[str, str]) -> CopilotClient: env = ctx.get_env() env.update(env_overrides) @@ -295,7 +320,9 @@ async def test_should_list_find_and_inspect_persisted_session_state( ): token = os.environ.get("GITHUB_TOKEN", "fakevalue") await _configure_user(authed_ctx, token) - client = _make_authed_client(authed_ctx, token) + client = _make_authed_client( + authed_ctx, token, request_handler=_PersistedSessionRequestHandler() + ) session_id = str(uuid.uuid4()) working_directory = Path(authed_ctx.work_dir) / f"server-rpc-list-{uuid.uuid4().hex}" @@ -311,33 +338,20 @@ async def test_should_list_find_and_inspect_persisted_session_state( on_permission_request=PermissionHandler.approve_all, ) - await session.send( - "Record a turn for sessions.list discriminator coverage", mode="enqueue" + # A user turn makes sessions.list nonempty. Finish a synthetic turn + # before inspecting persistence or detaching; enqueue alone leaves + # unobserved inference racing cleanup. + message = await session.send_and_wait( + "Record a turn for sessions.list discriminator coverage", timeout=60.0 ) - - listed = None - - async def session_is_listed() -> bool: - nonlocal listed - # Re-save on every attempt: on slower runners the enqueued turn is not - # necessarily recorded yet when the first save runs, so a single save - # followed by a fixed sleep races the CLI's own persistence. - save = await client.rpc.sessions.save(SessionsSaveRequest(session_id=session_id)) - assert save is not None - listed = await client.rpc.sessions.list( - SessionsListRequest( - filter=SessionListFilter(cwd=str(working_directory)), - metadata_limit=0, - ) + assert assistant_text(message) == SYNTHETIC_TEXT + save = await client.rpc.sessions.save(SessionsSaveRequest(session_id=session_id)) + assert save is not None + listed = await client.rpc.sessions.list( + SessionsListRequest( + filter=SessionListFilter(cwd=str(working_directory)), + metadata_limit=0, ) - return any(item.session_id == session_id for item in listed.sessions or []) - - await wait_for_condition( - session_is_listed, - timeout=60.0, - timeout_message=( - "Timed out waiting for the saved session to be returned by sessions.list." - ), ) assert listed is not None @@ -379,15 +393,17 @@ async def session_is_listed() -> bool: ) assert missing_session_id not in in_use.in_use finally: - if session is not None: - await session.disconnect() try: - await client.stop() - except ExceptionGroup: - # Intentional: shutting down the per-test client can race the - # CLI's own teardown and surface as an aggregated cancellation - # error from anyio. We don't want it to fail the test. - pass + if session is not None: + await session.disconnect() + finally: + try: + await client.stop() + except ExceptionGroup: + # Intentional: shutting down the per-test client can race the + # CLI's own teardown and surface as an aggregated cancellation + # error from anyio. We don't want it to fail the test. + pass async def test_should_enrich_basic_session_metadata(self, ctx: E2ETestContext): session_id = str(uuid.uuid4()) diff --git a/python/e2e/test_session_e2e.py b/python/e2e/test_session_e2e.py index f57b9f5736..160d627a0d 100644 --- a/python/e2e/test_session_e2e.py +++ b/python/e2e/test_session_e2e.py @@ -15,7 +15,6 @@ from .testharness import ( DEFAULT_GITHUB_TOKEN, E2ETestContext, - get_final_assistant_message, get_next_event_of_type, wait_for_condition, ) @@ -63,8 +62,8 @@ async def test_should_create_a_session_with_appended_systemMessage_config( system_message={"mode": "append", "content": system_message_suffix}, ) - await session.send("What is your full name?") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait("What is your full name?", timeout=10.0) + assert assistant_message is not None assert "GitHub" in assistant_message.data.content assert "Have a nice day!" in assistant_message.data.content @@ -83,8 +82,8 @@ async def test_should_create_a_session_with_replaced_systemMessage_config( system_message={"mode": "replace", "content": test_system_message}, ) - await session.send("What is your full name?") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait("What is your full name?", timeout=10.0) + assert assistant_message is not None assert "GitHub" not in assistant_message.data.content assert "Testy" in assistant_message.data.content @@ -235,8 +234,12 @@ async def test_should_resume_a_session_using_the_same_client(self, ctx: E2ETestC session_id, on_permission_request=PermissionHandler.approve_all ) assert session2.session_id == session_id - answer2 = await get_final_assistant_message(session2, already_idle=True) - assert "2" in answer2.data.content + # The completed turn's assistant message is durable; session.idle is not. + messages = await session2.get_events() + assert not any(message.type.value == "session.error" for message in messages) + answers = [message for message in messages if message.type.value == "assistant.message"] + assert answers + assert "2" in answers[-1].data.content # Can continue the conversation statefully answer3 = await session2.send_and_wait("Now if you double that, what do you get?") @@ -582,26 +585,27 @@ async def test_should_abort_a_session(self, ctx: E2ETestContext): ) # Set up event listeners BEFORE sending to avoid race conditions - wait_for_tool_start = asyncio.create_task( - get_next_event_of_type(session, "tool.execution_start", timeout=60.0) - ) - wait_for_session_idle = asyncio.create_task( - get_next_event_of_type(session, "session.idle", timeout=30.0) - ) + wait_for_tool_start = get_next_event_of_type(session, "tool.execution_start", timeout=60.0) + wait_for_session_idle = get_next_event_of_type(session, "session.idle", timeout=30.0) - # Send a message that will trigger a long-running shell command - await session.send( - "run the shell command 'sleep 100' (note this works on both bash and PowerShell)" - ) + try: + # Send a message that will trigger a long-running shell command + await session.send( + "run the shell command 'sleep 100' (note this works on both bash and PowerShell)" + ) - # Wait for the tool to start executing - _ = await wait_for_tool_start + # Wait for the tool to start executing + _ = await wait_for_tool_start - # Abort the session while the tool is running - await session.abort() + # Abort the session while the tool is running + await session.abort() - # Wait for session to become idle after abort - _ = await wait_for_session_idle + # Wait for session to become idle after abort + _ = await wait_for_session_idle + finally: + wait_for_tool_start.cancel() + wait_for_session_idle.cancel() + await asyncio.gather(wait_for_tool_start, wait_for_session_idle, return_exceptions=True) # The session should still be alive and usable after abort messages = await session.get_events() @@ -612,11 +616,8 @@ async def test_should_abort_a_session(self, ctx: E2ETestContext): assert len(abort_events) > 0, "Expected an abort event in messages" # We should be able to send another message - wait_for_answer = asyncio.create_task( - get_next_event_of_type(session, "assistant.message", timeout=60.0) - ) - await session.send("What is 2+2?") - answer = await wait_for_answer + answer = await session.send_and_wait("What is 2+2?", timeout=60.0) + assert answer is not None assert "4" in answer.data.content async def test_should_receive_session_events(self, ctx: E2ETestContext): @@ -671,11 +672,13 @@ def on_event(event): assert "assistant.message" in event_types assert "session.idle" in event_types - # Verify the assistant response contains the expected answer. - # session.idle is ephemeral and not in get_events(), but we already - # confirmed idle via the live event handler above. - assistant_message = await get_final_assistant_message(session, already_idle=True) - assert "300" in assistant_message.data.content + # Idle was observed live, so inspect the messages captured for this turn. + assert "session.error" not in event_types + assistant_messages = [ + event for event in received_events if event.type.value == "assistant.message" + ] + assert assistant_messages + assert "300" in assistant_messages[-1].data.content async def test_should_create_session_with_custom_config_dir(self, ctx: E2ETestContext): import os @@ -688,8 +691,8 @@ async def test_should_create_session_with_custom_config_dir(self, ctx: E2ETestCo assert session.session_id # Session should work normally with custom config dir - await session.send("What is 1+1?") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait("What is 1+1?", timeout=10.0) + assert assistant_message is not None assert "2" in assistant_message.data.content async def test_session_log_emits_events_at_all_levels(self, ctx: E2ETestContext): @@ -993,28 +996,35 @@ async def test_send_returns_immediately_while_events_stream_in_background( self, ctx: E2ETestContext ): """`send` returns before the session goes idle; events are streamed.""" + import asyncio + session = await ctx.client.create_session( on_permission_request=PermissionHandler.approve_all, ) - events: list[str] = [] + events = [] def on_event(event): - events.append(event.type.value) + events.append(event) - session.on(on_event) - - # Use a slow command so we can verify send() returns before completion - await session.send("Run 'sleep 2 && echo done'") - - # send() should return before turn completes (no session.idle yet) - assert "session.idle" not in events + unsubscribe = session.on(on_event) + idle_task = get_next_event_of_type(session, "session.idle", timeout=10.0) + try: + # Use a slow command so we can verify send() returns before completion + await session.send("Run 'sleep 2 && echo done'") - message = await get_final_assistant_message(session) - assert "done" in message.data.content - assert "session.idle" in events - assert "assistant.message" in events + # send() should return before turn completes (no session.idle yet) + assert not any(event.type.value == "session.idle" for event in events) - await session.disconnect() + await idle_task + messages = [event for event in events if event.type.value == "assistant.message"] + assert messages + assert "done" in messages[-1].data.content + assert any(event.type.value == "session.idle" for event in events) + finally: + idle_task.cancel() + await asyncio.gather(idle_task, return_exceptions=True) + unsubscribe() + await session.disconnect() async def test_sendandwait_blocks_until_session_idle_and_returns_final_assistant_message( self, ctx: E2ETestContext @@ -1043,24 +1053,24 @@ async def test_sendandwait_throws_on_timeout(self, ctx: E2ETestContext): on_permission_request=PermissionHandler.approve_all, ) - # Start a background wait for session.idle so we can drain after we abort. - idle_task = asyncio.create_task( - get_next_event_of_type(session, "session.idle", timeout=30.0) - ) - - with pytest.raises(TimeoutError) as exc_info: - await session.send_and_wait( - "Run 'sleep 2 && echo done'", - timeout=0.1, - ) - assert "Timeout" in str(exc_info.value) or "timed out" in str(exc_info.value).lower() - - # The timeout only cancels the client-side wait; abort the agent and wait for idle - # so leftover requests don't leak into subsequent tests. - await session.abort() - await idle_task + # Subscribe before sending so even an idle emitted before the abort reply is captured. + idle_task = get_next_event_of_type(session, "session.idle", timeout=30.0) + try: + with pytest.raises(TimeoutError) as exc_info: + await session.send_and_wait( + "Run 'sleep 2 && echo done'", + timeout=0.1, + ) + assert "Timeout" in str(exc_info.value) or "timed out" in str(exc_info.value).lower() - await session.disconnect() + # The timeout only cancels the client-side wait; abort the agent and wait for idle + # so leftover requests don't leak into subsequent tests. + await session.abort() + await idle_task + finally: + idle_task.cancel() + await asyncio.gather(idle_task, return_exceptions=True) + await session.disconnect() async def test_sendandwait_throws_operationcanceledexception_when_token_cancelled( self, ctx: E2ETestContext @@ -1072,12 +1082,8 @@ async def test_sendandwait_throws_operationcanceledexception_when_token_cancelle on_permission_request=PermissionHandler.approve_all, ) - tool_start_task = asyncio.create_task( - get_next_event_of_type(session, "tool.execution_start", timeout=60.0) - ) - idle_task = asyncio.create_task( - get_next_event_of_type(session, "session.idle", timeout=30.0) - ) + tool_start_task = get_next_event_of_type(session, "tool.execution_start", timeout=60.0) + idle_task = get_next_event_of_type(session, "session.idle", timeout=30.0) send_task = asyncio.create_task( session.send_and_wait( @@ -1086,18 +1092,23 @@ async def test_sendandwait_throws_operationcanceledexception_when_token_cancelle ) ) - # Wait for the tool to begin executing before cancelling. - await tool_start_task - - send_task.cancel() - with pytest.raises((asyncio.CancelledError, BaseException)): - await send_task + try: + # Wait for the tool to begin executing before cancelling. + await tool_start_task - # Cancelling only cancels the client-side wait; abort and wait for idle. - await session.abort() - await idle_task + send_task.cancel() + with pytest.raises((asyncio.CancelledError, BaseException)): + await send_task - await session.disconnect() + # Cancelling only cancels the client-side wait; abort and wait for idle. + await session.abort() + await idle_task + finally: + tool_start_task.cancel() + idle_task.cancel() + send_task.cancel() + await asyncio.gather(tool_start_task, idle_task, send_task, return_exceptions=True) + await session.disconnect() async def test_should_set_model_on_existing_session(self, ctx: E2ETestContext): """`set_model` emits a session.model_change event with the new model.""" diff --git a/python/e2e/test_session_todos_changed_e2e.py b/python/e2e/test_session_todos_changed_e2e.py index 8911ffb117..fd6d92f6c8 100644 --- a/python/e2e/test_session_todos_changed_e2e.py +++ b/python/e2e/test_session_todos_changed_e2e.py @@ -31,11 +31,13 @@ async def test_fires_session_todos_changed_and_exposes_rows_and_dependencies( async with await ctx.client.create_session( on_permission_request=PermissionHandler.approve_all, ) as session: - todos_changed = asyncio.create_task( - get_next_event_of_type(session, "session.todos_changed", timeout=120.0) - ) - await session.send_and_wait(PROMPT, timeout=120.0) - await todos_changed + todos_changed = get_next_event_of_type(session, "session.todos_changed", timeout=120.0) + try: + await session.send_and_wait(PROMPT, timeout=120.0) + await todos_changed + finally: + todos_changed.cancel() + await asyncio.gather(todos_changed, return_exceptions=True) result = await session.rpc.plan.read_sql_todos_with_dependencies() ids = sorted(row.id for row in result.rows if row.id) diff --git a/python/e2e/test_telemetry_e2e.py b/python/e2e/test_telemetry_e2e.py index 8b9c82abef..56031a14ec 100644 --- a/python/e2e/test_telemetry_e2e.py +++ b/python/e2e/test_telemetry_e2e.py @@ -26,7 +26,7 @@ from copilot.session import PermissionHandler from copilot.tools import Tool, ToolInvocation, ToolResult -from .testharness import E2ETestContext, get_final_assistant_message +from .testharness import E2ETestContext pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -102,8 +102,8 @@ def echo(invocation: ToolInvocation) -> ToolResult: ) session_id = session.session_id - await session.send(prompt) - answer = await get_final_assistant_message(session, timeout=60.0) + answer = await session.send_and_wait(prompt, timeout=60.0) + assert answer is not None assert "TELEMETRY_E2E_DONE" in (answer.data.content or "") await session.disconnect() diff --git a/python/e2e/test_tool_results_e2e.py b/python/e2e/test_tool_results_e2e.py index 41d1967bc9..f2f646f648 100644 --- a/python/e2e/test_tool_results_e2e.py +++ b/python/e2e/test_tool_results_e2e.py @@ -9,7 +9,7 @@ from copilot.session import PermissionHandler from copilot.tools import ToolInvocation, ToolResult -from .testharness import E2ETestContext, get_final_assistant_message +from .testharness import E2ETestContext pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -33,8 +33,10 @@ def get_weather(params: WeatherParams, invocation: ToolInvocation) -> ToolResult ) try: - await session.send("What's the weather in Paris?") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "What's the weather in Paris?", timeout=10.0 + ) + assert assistant_message is not None assert ( "sunny" in assistant_message.data.content.lower() or "72" in assistant_message.data.content @@ -87,8 +89,10 @@ def analyze_code(params: AnalyzeParams, invocation: ToolInvocation) -> ToolResul ) try: - await session.send("Analyze the file main.ts for issues.") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Analyze the file main.ts for issues.", timeout=10.0 + ) + assert assistant_message is not None assert "no issues" in assistant_message.data.content.lower() # Verify the LLM received just textResultForLlm, not stringified JSON diff --git a/python/e2e/test_tools_e2e.py b/python/e2e/test_tools_e2e.py index 1421dbaf40..303ddc3ac9 100644 --- a/python/e2e/test_tools_e2e.py +++ b/python/e2e/test_tools_e2e.py @@ -13,7 +13,7 @@ from copilot.session import PermissionHandler, PermissionNoResult from copilot.tools import Tool, ToolInvocation, ToolResult -from .testharness import E2ETestContext, get_final_assistant_message +from .testharness import E2ETestContext pytestmark = pytest.mark.asyncio(loop_scope="module") @@ -28,8 +28,10 @@ async def test_invokes_built_in_tools(self, ctx: E2ETestContext): on_permission_request=PermissionHandler.approve_all ) - await session.send("What's the first line of README.md in this directory?") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "What's the first line of README.md in this directory?", timeout=10.0 + ) + assert assistant_message is not None assert "ELIZA" in assistant_message.data.content async def test_invokes_custom_tool(self, ctx: E2ETestContext): @@ -44,8 +46,10 @@ def encrypt_string(params: EncryptParams, invocation: ToolInvocation) -> str: on_permission_request=PermissionHandler.approve_all, tools=[encrypt_string] ) - await session.send("Use encrypt_string to encrypt this string: Hello") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Use encrypt_string to encrypt this string: Hello", timeout=10.0 + ) + assert assistant_message is not None assert "HELLO" in assistant_message.data.content async def test_low_level_tool_definition(self, ctx: E2ETestContext): @@ -83,8 +87,8 @@ def search_items(params: SearchArgs, invocation: ToolInvocation) -> str: "First, set the current phase to 'analyzing'. Then search for items with " "keyword 'copilot'. Report the phase and search results." ) - await session.send(prompt) - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait(prompt, timeout=10.0) + assert assistant_message is not None content = assistant_message.data.content or "" assert content != "" assert "analyzing" in content.lower() @@ -100,8 +104,10 @@ def get_user_location() -> str: on_permission_request=PermissionHandler.approve_all, tools=[get_user_location] ) - await session.send("What is my location? If you can't find out, just say 'unknown'.") - answer = await get_final_assistant_message(session) + answer = await session.send_and_wait( + "What is my location? If you can't find out, just say 'unknown'.", timeout=10.0 + ) + assert answer is not None # Check the underlying traffic traffic = await ctx.get_exchanges() @@ -164,12 +170,13 @@ def db_query(params: DbQueryParams, invocation: ToolInvocation) -> list[City]: ) expected_session_id = session.session_id - await session.send( + assistant_message = await session.send_and_wait( "Perform a DB query for the 'cities' table using IDs 12 and 19, " - "sorting ascending. Reply only with lines of the form: [cityname] [population]" + "sorting ascending. Reply only with lines of the form: [cityname] [population]", + timeout=10.0, ) - assistant_message = await get_final_assistant_message(session) + assert assistant_message is not None response_content = assistant_message.data.content or "" assert response_content != "" @@ -201,8 +208,10 @@ def tracking_handler(request, invocation): on_permission_request=tracking_handler, tools=[safe_lookup] ) - await session.send("Use safe_lookup to look up 'test123'") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Use safe_lookup to look up 'test123'", timeout=10.0 + ) + assert assistant_message is not None assert "RESULT: test123" in assistant_message.data.content assert not did_run_permission_request @@ -222,8 +231,10 @@ def custom_grep(params: GrepParams, invocation: ToolInvocation) -> str: on_permission_request=PermissionHandler.approve_all, tools=[custom_grep] ) - await session.send("Use grep to search for the word 'hello'") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Use grep to search for the word 'hello'", timeout=10.0 + ) + assert assistant_message is not None assert "CUSTOM_GREP_RESULT" in assistant_message.data.content async def test_invokes_custom_tool_with_permission_handler(self, ctx: E2ETestContext): @@ -244,8 +255,10 @@ def on_permission_request(request, invocation): on_permission_request=on_permission_request, tools=[encrypt_string] ) - await session.send("Use encrypt_string to encrypt this string: Hello") - assistant_message = await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Use encrypt_string to encrypt this string: Hello", timeout=10.0 + ) + assert assistant_message is not None assert "HELLO" in assistant_message.data.content # Should have received a custom-tool permission request @@ -272,8 +285,10 @@ def on_permission_request(request, invocation): on_permission_request=on_permission_request, tools=[encrypt_string] ) - await session.send("Use encrypt_string to encrypt this string: Hello") - await get_final_assistant_message(session) + assistant_message = await session.send_and_wait( + "Use encrypt_string to encrypt this string: Hello", timeout=10.0 + ) + assert assistant_message is not None # The tool handler should NOT have been called since permission was denied assert not tool_handler_called @@ -328,9 +343,10 @@ def lookup_country(invocation: ToolInvocation) -> ToolResult: ) try: - await session.send( + assistant_message = await session.send_and_wait( "Use lookup_city with 'Paris' and lookup_country with 'France' at the same time," - " then combine both results in your reply." + " then combine both results in your reply.", + timeout=60.0, ) city_result = await asyncio.wait_for(city_called, timeout=60.0) @@ -338,7 +354,6 @@ def lookup_country(invocation: ToolInvocation) -> ToolResult: assert city_result == "Paris" assert country_result == "France" - assistant_message = await get_final_assistant_message(session, timeout=60.0) assert assistant_message is not None content = assistant_message.data.content or "" assert "CITY_PARIS" in content diff --git a/python/e2e/testharness/__init__.py b/python/e2e/testharness/__init__.py index 75ce76d9c5..cfc33f29f7 100644 --- a/python/e2e/testharness/__init__.py +++ b/python/e2e/testharness/__init__.py @@ -1,7 +1,7 @@ """Test harness for E2E tests.""" from .context import CLI_PATH, DEFAULT_GITHUB_TOKEN, E2ETestContext, is_inprocess_transport -from .helper import get_final_assistant_message, get_next_event_of_type, wait_for_condition +from .helper import get_next_event_of_type, wait_for_condition, wait_for_event from .proxy import CapiProxy __all__ = [ @@ -9,8 +9,8 @@ "DEFAULT_GITHUB_TOKEN", "E2ETestContext", "CapiProxy", - "get_final_assistant_message", "get_next_event_of_type", "wait_for_condition", + "wait_for_event", "is_inprocess_transport", ] diff --git a/python/e2e/testharness/helper.py b/python/e2e/testharness/helper.py index 7933dd9ec8..ed50d17efb 100644 --- a/python/e2e/testharness/helper.py +++ b/python/e2e/testharness/helper.py @@ -8,104 +8,8 @@ import time from collections.abc import Awaitable, Callable -from copilot import CopilotSession -from copilot.session_events import ( - AssistantMessageData, - SessionErrorData, - SessionIdleData, -) - - -async def get_final_assistant_message( - session: CopilotSession, timeout: float = 10.0, already_idle: bool = False -): - """ - Wait for and return the final assistant message from a session turn. - - Args: - session: The session to wait on - timeout: Maximum time to wait in seconds - - Returns: - The final assistant message event - - Raises: - TimeoutError: If no message arrives within timeout - RuntimeError: If a session error occurs - """ - result_future: asyncio.Future = asyncio.get_event_loop().create_future() - - final_assistant_message = None - - def on_event(event): - nonlocal final_assistant_message - if result_future.done(): - return - - match event.data: - case AssistantMessageData(): - final_assistant_message = event - case SessionIdleData(): - if final_assistant_message is not None: - result_future.set_result(final_assistant_message) - case SessionErrorData() as data: - msg = data.message if data.message else "session error" - result_future.set_exception(RuntimeError(msg)) - - # Subscribe to future events - unsubscribe = session.on(on_event) - - try: - # Also check existing messages in case the response already arrived - existing = await _get_existing_final_response(session, already_idle) - if existing is not None: - return existing - - return await asyncio.wait_for(result_future, timeout=timeout) - finally: - unsubscribe() - - -async def _get_existing_final_response(session: CopilotSession, already_idle: bool = False): - """Check existing messages for a final response.""" - messages = await session.get_events() - - # Find last user message - final_user_message_index = -1 - for i in range(len(messages) - 1, -1, -1): - if messages[i].type.value == "user.message": - final_user_message_index = i - break - - if final_user_message_index < 0: - current_turn_messages = messages - else: - current_turn_messages = messages[final_user_message_index:] - - # Check for errors - for msg in current_turn_messages: - match msg.data: - case SessionErrorData() as data: - err_msg = data.message if data.message else "session error" - raise RuntimeError(err_msg) - - # Find session.idle and get last assistant message before it - if already_idle: - session_idle_index = len(current_turn_messages) - else: - session_idle_index = -1 - for i, msg in enumerate(current_turn_messages): - if msg.type.value == "session.idle": - session_idle_index = i - break - - if session_idle_index != -1: - # Find last assistant.message before session.idle - for i in range(session_idle_index - 1, -1, -1): - if current_turn_messages[i].type.value == "assistant.message": - return current_turn_messages[i] - - return None +from _session_test_helpers import get_next_event_of_type as get_next_event_of_type +from _session_test_helpers import wait_for_event as wait_for_event def write_file(work_dir: str, filename: str, content: str) -> str: @@ -165,41 +69,3 @@ async def wait_for_condition( if result: return raise TimeoutError(timeout_message) - - -async def get_next_event_of_type(session: CopilotSession, event_type: str, timeout: float = 30.0): - """ - Wait for and return the next event of a specific type from a session. - - Args: - session: The session to wait on - event_type: The event type to wait for (e.g., "tool.execution_start", "session.idle") - timeout: Maximum time to wait in seconds - - Returns: - The matching event - - Raises: - TimeoutError: If no matching event arrives within timeout - RuntimeError: If a session error occurs - """ - result_future: asyncio.Future = asyncio.get_event_loop().create_future() - - def on_event(event): - if result_future.done(): - return - - if event.type.value == event_type: - result_future.set_result(event) - else: - match event.data: - case SessionErrorData() as data: - msg = data.message if data.message else "session error" - result_future.set_exception(RuntimeError(msg)) - - unsubscribe = session.on(on_event) - - try: - return await asyncio.wait_for(result_future, timeout=timeout) - finally: - unsubscribe() diff --git a/python/e2e/timeout_diagnostics.py b/python/e2e/timeout_diagnostics.py new file mode 100644 index 0000000000..c02f069246 --- /dev/null +++ b/python/e2e/timeout_diagnostics.py @@ -0,0 +1,232 @@ +"""Failure-only diagnostics for async E2E timeouts, including xdist workers.""" + +import asyncio +import gc +import inspect +import io +import os +import signal +import subprocess +import sys +import threading +import time +import traceback +import uuid +from pathlib import Path + +import pytest + +from copilot._jsonrpc import JsonRpcClient + +_TIMEOUT_DIAGNOSTICS = pytest.StashKey[tuple[str, Path | None]]() + + +@pytest.hookimpl(hookwrapper=True, optionalhook=True) +def pytest_timeout_set_timer(item, settings): + yield + if settings.method != "signal" or threading.current_thread() is not threading.main_thread(): + return + original_handler = signal.getsignal(signal.SIGALRM) + if not callable(original_handler): + return + + def capture_timeout(signum, frame): + try: + original_handler(signum, frame) + except pytest.fail.Exception as exc: + if "from pytest-timeout" in str(exc): + # Capture before fixture finalizers/Runner.close cancel the + # suspended tasks. In particular, makereport is too late for teardown. + item.stash[_TIMEOUT_DIAGNOSTICS] = _collect_timeout_diagnostics(item) + raise + + signal.signal(signal.SIGALRM, capture_timeout) + + +def _dump_awaitable(awaitable, output, seen=None): + if seen is None: + seen = set() + while awaitable is not None and id(awaitable) not in seen: + seen.add(id(awaitable)) + frame = None + next_awaitable = None + for frame_attr, await_attr in ( + ("cr_frame", "cr_await"), + ("ag_frame", "ag_await"), + ("gi_frame", "gi_yieldfrom"), + ): + frame = getattr(awaitable, frame_attr, None) + if frame is not None: + next_awaitable = getattr(awaitable, await_attr, None) + break + if frame is not None: + code = frame.f_code + print(f" {code.co_filename}:{frame.f_lineno} in {code.co_qualname}", file=output) + # Do not dump arbitrary locals, RPC payloads, prompts, tokens, or results. + if code is JsonRpcClient.request.__code__: + values = frame.f_locals + params = values.get("params") or {} + print( + f" outbound method={values.get('method')}" + f" request_id={values.get('request_id')}" + f" session_id={params.get('sessionId')}" + f" elapsed={time.perf_counter() - values['request_start']:.3f}s", + file=output, + ) + elif code is JsonRpcClient._dispatch_request.__code__: + message = frame.f_locals["message"] + print( + f" inbound method={message.get('method')} request_id={message.get('id')}", + file=output, + ) + else: + print(f" awaiting {type(awaitable).__name__}", file=output) + # Async fixture finalizers await an asend object, which hides ag_await. + if type(awaitable).__name__ in ("async_generator_asend", "async_generator_athrow"): + for referent in gc.get_referents(awaitable): + if inspect.isasyncgen(referent): + _dump_awaitable(referent, output, seen) + awaitable = next_awaitable + + +def _dump_client(client, output): + rpc = client._client + if rpc is not None: + reader = rpc._read_thread + print( + f"JSON-RPC running={rpc._running}" + f" reader_alive={reader is not None and reader.is_alive()}" + f" write_locked={rpc._write_lock.locked()}" + f" pending_locked={rpc._pending_lock.locked()}", + file=output, + ) + # A timeout may have interrupted a lock owner: snapshots must not acquire locks. + for request_id, future in rpc.pending_requests.copy().items(): + print( + f" pending request_id={request_id} done={future.done()}" + f" cancelled={future.cancelled()}", + file=output, + ) + for session_id, session in client._sessions.copy().items(): + print( + f"Session {session_id} destroyed={session._destroyed}" + f" disconnect_locked={session._disconnect_lock.locked()}", + file=output, + ) + host = client._ffi_host + if host is not None: + print( + f"FFI server_id={host._server_id} connection_id={host._connection_id}" + f" disposed={host._disposed} starting={host._starting}" + f" operation_locked={host._operation_lock.locked()}" + f" dispose_locked={host._dispose_lock.locked()}" + f" receive_closed={host._receive_buffer._closed}" + f" receive_bytes={len(host._receive_buffer._buffer)}", + file=output, + ) + + +def _sample_native_threads(path: Path, output): + try: + result = subprocess.run( + ["sample", str(os.getpid()), "1", "-file", str(path)], + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + timeout=10, + check=False, + ) + print(f"Native sample exit={result.returncode} file={path}", file=output) + if result.returncode: + print(result.stdout, file=output) + except (OSError, subprocess.TimeoutExpired) as exc: + print(f"Native sample unavailable: {type(exc).__name__}", file=output) + + +def _collect_timeout_diagnostics(item): + output = io.StringIO() + print(f"Test: {item.nodeid}\nPID: {os.getpid()}", file=output) + client = None + try: + context = item.funcargs.get("ctx") + client = getattr(context, "_client", None) + loops = set() + if client is not None: + _dump_client(client, output) + if client._client is not None and client._client._loop is not None: + loops.add(client._client._loop) + for fixture in item.funcargs.values(): + if isinstance(fixture, asyncio.Runner): + loops.add(fixture.get_loop()) + for loop in loops: + print(f"Event loop running={loop.is_running()} closed={loop.is_closed()}", file=output) + for task in sorted(asyncio.all_tasks(loop), key=lambda task: task.get_name()): + print( + f"Task {task.get_name()} done={task.done()} cancelling={task.cancelling()}", + file=output, + ) + _dump_awaitable(task.get_coro(), output) + names = {thread.ident: thread.name for thread in threading.enumerate()} + for ident, frame in sys._current_frames().items(): + print(f"Thread {ident} ({names.get(ident, 'native')})", file=output) + traceback.print_stack(frame, file=output) + except Exception as exc: + print(f"Diagnostic collection failed: {type(exc).__name__}", file=output) + + path = None + try: + directory = item.config.rootpath / ".pytest-diagnostics" + directory.mkdir(exist_ok=True) + stem = f"{os.getpid()}-{uuid.uuid4().hex}" + path = directory / f"{stem}.txt" + path.write_text(output.getvalue(), encoding="utf-8") + if sys.platform == "darwin" and getattr(client, "_ffi_host", None) is not None: + _sample_native_threads(directory / f"{stem}.sample.txt", output) + print(f"Diagnostics saved to {path}", file=output) + path.write_text(output.getvalue(), encoding="utf-8") + except Exception as exc: + print(f"Diagnostic artifact unavailable: {type(exc).__name__}", file=output) + return output.getvalue(), path + + +def add_timeout_diagnostics(item, call, report): + """Attach the timeout snapshot to a report so it survives xdist's suppressed stdout.""" + if call.excinfo is None: + return + error = str(call.excinfo.value) + if not report.failed or "Timeout (" not in error or "from pytest-timeout" not in error: + return + snapshot = item.stash.get(_TIMEOUT_DIAGNOSTICS, None) + if snapshot is None: + snapshot = _collect_timeout_diagnostics(item) + else: + del item.stash[_TIMEOUT_DIAGNOSTICS] + text, path = snapshot + text = f"Phase: {report.when}\n{text}" + if path is not None: + try: + path.write_text(text, encoding="utf-8") + except OSError as exc: + text += f"Diagnostic artifact update failed: {type(exc).__name__}\n" + report.sections.append(("Async timeout diagnostics", text)) + + +def cancel_timed_out_test(call): + """Cancel only the test task abandoned by a timeout outside its coroutine.""" + if call.when != "call" or call.excinfo is None: + return False + error = call.excinfo.value + if "Timeout (" not in str(error) or "from pytest-timeout" not in str(error): + return False + traceback_entry = error.__traceback__ + while traceback_entry is not None: + frame = traceback_entry.tb_frame + if frame.f_code is asyncio.BaseEventLoop.run_until_complete.__code__: + task = frame.f_locals.get("future") + if isinstance(task, asyncio.Task) and not task.done(): + # The signal interrupts the runner, not its task. Let cancellation + # release that task's locks when the fixture's loop next resumes. + return task.cancel() + return False + traceback_entry = traceback_entry.tb_next + return False diff --git a/python/test_rpc_server_fixture.py b/python/test_rpc_server_fixture.py new file mode 100644 index 0000000000..4467eba68c --- /dev/null +++ b/python/test_rpc_server_fixture.py @@ -0,0 +1,97 @@ +"""Regression controls for the persisted-session E2E fixture's completion fence.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from copilot.generated.session_events import AssistantMessageData +from copilot.rpc import LocalSessionMetadataValue, SessionContext +from e2e import test_rpc_server_e2e as scenario + + +@pytest.mark.asyncio +@pytest.mark.parametrize("disconnect_fails", [False, True]) +async def test_persisted_session_fixture_observes_completion_before_save_and_cleanup( + tmp_path, monkeypatch, disconnect_fails +): + state = {} + calls = [] + + async def send(*args, **kwargs): + calls.append("send-without-completion") + + async def send_and_wait(prompt, timeout): + assert prompt == "Record a turn for sessions.list discriminator coverage" + assert timeout == 60.0 + calls.append("completed-turn") + return SimpleNamespace( + data=AssistantMessageData( + content=scenario.SYNTHETIC_TEXT, message_id="fixture-assistant" + ) + ) + + async def save(request): + assert calls == ["completed-turn"], "persistence must follow observed turn completion" + assert request.session_id == state["session_id"] + calls.append("save") + return object() + + async def listed(request): + assert calls == ["completed-turn", "save"] + return SimpleNamespace( + sessions=[ + LocalSessionMetadataValue( + session_id=state["session_id"], + is_remote=False, + start_time="2026-01-01T00:00:00Z", + modified_time="2026-01-01T00:00:00Z", + context=SessionContext(cwd=state["working_directory"]), + ) + ] + ) + + async def disconnect(): + calls.append("disconnect") + if disconnect_fails: + raise RuntimeError("controlled detach failure") + + session = SimpleNamespace(send=send, send_and_wait=send_and_wait, disconnect=disconnect) + + async def create_session(**kwargs): + state.update(kwargs) + return session + + client = SimpleNamespace( + start=AsyncMock(), + stop=AsyncMock(), + create_session=create_session, + rpc=SimpleNamespace( + sessions=SimpleNamespace( + save=save, + list=listed, + find_by_prefix=AsyncMock(return_value=SimpleNamespace(session_id=None)), + find_by_task_id=AsyncMock(return_value=SimpleNamespace(session_id=None)), + get_last_for_context=AsyncMock(return_value=SimpleNamespace(session_id=None)), + get_sizes=AsyncMock(return_value=SimpleNamespace(sizes={})), + check_in_use=AsyncMock(return_value=SimpleNamespace(in_use=[])), + ) + ), + ) + + def make_client(ctx, token, **kwargs): + if "request_handler" in kwargs: + assert isinstance(kwargs["request_handler"], scenario._PersistedSessionRequestHandler) + return client + + monkeypatch.setattr(scenario, "_configure_user", AsyncMock()) + monkeypatch.setattr(scenario, "_make_authed_client", make_client) + run = scenario.TestRpcServer().test_should_list_find_and_inspect_persisted_session_state + ctx = SimpleNamespace(work_dir=str(tmp_path)) + if disconnect_fails: + with pytest.raises(RuntimeError, match="controlled detach failure"): + await run(ctx) + else: + await run(ctx) + assert calls == ["completed-turn", "save", "disconnect"] + client.stop.assert_awaited_once() diff --git a/python/test_session.py b/python/test_session.py index 8fcecba439..5750364bb5 100644 --- a/python/test_session.py +++ b/python/test_session.py @@ -8,6 +8,7 @@ import pytest +from _session_test_helpers import get_next_event_of_type, wait_for_event from copilot import AgentMessageSource, MessageSource from copilot.session import Attachment, CopilotSession from copilot.session_events import ( @@ -57,6 +58,186 @@ def _event(data, event_type: SessionEventType) -> SessionEvent: ) +@pytest.mark.parametrize("use_send_and_wait", [False, True]) +@pytest.mark.asyncio +async def test_completion_captures_live_idle_before_send_reply_without_idle_in_history( + use_send_and_wait, +): + client = Mock() + session = CopilotSession("session-1", client) + intermediate = _event( + AssistantMessageData(content="working", message_id="assistant-1"), + SessionEventType.ASSISTANT_MESSAGE, + ) + assistant = _event( + AssistantMessageData(content="done", message_id="assistant-2"), + SessionEventType.ASSISTANT_MESSAGE, + ) + idle = _event(SessionIdleData(), SessionEventType.SESSION_IDLE) + idle.ephemeral = True + history = [intermediate, assistant] + received = [] + unsubscribe = session.on(received.append) + + async def respond(method, params): + assert params["sessionId"] == session.session_id + if method == "session.getMessages": + return {"events": [event.to_dict() for event in history]} + assert method == "session.send" + # No suspension: even a create_task(async_waiter()) cannot subscribe in time. + session._dispatch_event(intermediate) + session._dispatch_event(assistant) + session._dispatch_event(idle) + return {"messageId": "message-1"} + + client.request = AsyncMock(side_effect=respond) + try: + if use_send_and_wait: + message = await session.send_and_wait("hello", timeout=1) + assert message is not None + assert message is assistant + else: + idle_task = get_next_event_of_type(session, "session.idle", timeout=1) + try: + assert await session.send("hello") == "message-1" + assert received == [intermediate, assistant, idle] + assert await idle_task is idle + messages = [ + event for event in received if event.type == SessionEventType.ASSISTANT_MESSAGE + ] + assert messages[-1] is assistant + finally: + idle_task.cancel() + await asyncio.gather(idle_task, return_exceptions=True) + + client.request.assert_awaited_once() + persisted = await session.get_events() + assert [event.type for event in persisted] == [ + SessionEventType.ASSISTANT_MESSAGE, + SessionEventType.ASSISTANT_MESSAGE, + ] + assert persisted[-1].data.content == "done" + finally: + unsubscribe() + + +@pytest.mark.asyncio +async def test_event_waiters_capture_abort_and_recovery_before_rpc_replies(): + client = Mock() + session = CopilotSession("session-1", client) + idle = _event(SessionIdleData(), SessionEventType.SESSION_IDLE) + assistant = _event( + AssistantMessageData(content="recovered", message_id="assistant-1"), + SessionEventType.ASSISTANT_MESSAGE, + ) + + async def respond(method, params): + if method == "session.abort": + session._dispatch_event(idle) + return {} + assert method == "session.send" + session._dispatch_event(assistant) + session._dispatch_event(idle) + return {"messageId": "message-1"} + + client.request = AsyncMock(side_effect=respond) + aborted = get_next_event_of_type(session, "session.idle", timeout=1) + try: + await session.abort() + assert await aborted is idle + finally: + aborted.cancel() + await asyncio.gather(aborted, return_exceptions=True) + + recovery = wait_for_event( + session, + lambda event: ( + isinstance(event.data, AssistantMessageData) and event.data.content == "recovered" + ), + timeout=1, + ) + recovered_idle = get_next_event_of_type(session, "session.idle", timeout=1) + try: + await session.send("recover") + assert await recovery is assistant + assert await recovered_idle is idle + finally: + recovery.cancel() + recovered_idle.cancel() + await asyncio.gather(recovery, recovered_idle, return_exceptions=True) + + +@pytest.mark.parametrize( + "outcome", ["success", "error", "timeout", "cancel-before-start", "cancel-after-start"] +) +@pytest.mark.asyncio +async def test_event_waiter_unsubscribes_for_every_outcome(outcome): + session = Mock(spec=CopilotSession) + unsubscribe = session.on.return_value + waiter = get_next_event_of_type( + session, "session.idle", timeout=0 if outcome == "timeout" else 1 + ) + session.on.assert_called_once() + on_event = session.on.call_args.args[0] + + if outcome == "success": + idle = _event(SessionIdleData(), SessionEventType.SESSION_IDLE) + on_event(idle) + on_event(idle) + assert await waiter is idle + elif outcome == "error": + on_event( + _event( + SessionErrorData(error_type="notification", message="turn failed"), + SessionEventType.SESSION_ERROR, + ) + ) + with pytest.raises(RuntimeError, match="turn failed"): + await waiter + elif outcome == "timeout": + with pytest.raises(TimeoutError): + await waiter + else: + if outcome == "cancel-after-start": + loop = asyncio.get_running_loop() + started = loop.create_future() + loop.call_soon(started.set_result, None) + await started + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + + unsubscribe.assert_called_once() + + +@pytest.mark.parametrize("fail_on_session_error", [None, False, True]) +@pytest.mark.asyncio +async def test_event_waiter_preserves_caller_error_policy(fail_on_session_error): + session = Mock(spec=CopilotSession) + options = ( + {} if fail_on_session_error is None else {"fail_on_session_error": fail_on_session_error} + ) + waiter = wait_for_event( + session, lambda event: isinstance(event.data, SessionIdleData), timeout=1, **options + ) + on_event = session.on.call_args.args[0] + on_event( + _event( + SessionErrorData(error_type="rate_limit", message="rate limited"), + SessionEventType.SESSION_ERROR, + ) + ) + idle = _event(SessionIdleData(), SessionEventType.SESSION_IDLE) + on_event(idle) + + if fail_on_session_error: + with pytest.raises(RuntimeError, match="rate limited"): + await waiter + else: + assert await waiter is idle + session.on.return_value.assert_called_once() + + @pytest.mark.asyncio async def test_send_omits_source_for_plain_human_prompt(monkeypatch): monkeypatch.setattr("copilot.session.get_trace_context", lambda: {}) diff --git a/python/test_timeout_cleanup.py b/python/test_timeout_cleanup.py new file mode 100644 index 0000000000..4e9e9fad46 --- /dev/null +++ b/python/test_timeout_cleanup.py @@ -0,0 +1,163 @@ +"""Regression tests for timed-out tasks retaining module-fixture session locks.""" + +import os +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from e2e.timeout_diagnostics import cancel_timed_out_test + + +@pytest.mark.parametrize( + "cancel_test,withhold_all,teardown_error", + [(False, False, True), (True, False, False), (True, True, True)], +) +def test_module_cleanup_after_timeout(tmp_path, cancel_test, withhold_all, teardown_error): + (tmp_path / "conftest.py").write_text( + """ +import asyncio +import os +import signal +from types import SimpleNamespace + +import pytest +import pytest_asyncio +import pytest_timeout + +from copilot import CopilotClient, RuntimeConnection +from copilot._jsonrpc import JsonRpcClient +from copilot.session import CopilotSession +from e2e.timeout_diagnostics import add_timeout_diagnostics, cancel_timed_out_test +from e2e.timeout_diagnostics import pytest_timeout_set_timer as pytest_timeout_set_timer + +active_item = None + +def pytest_runtest_setup(item): + global active_item + active_item = item + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + report = outcome.get_result() + add_timeout_diagnostics(item, call, report) + if os.environ["CANCEL_TEST"] == "1" and report.when == "call" and report.failed: + assert cancel_timed_out_test(call) + +@pytest_asyncio.fixture(scope="module", loop_scope="module") +async def ctx(): + loop = asyncio.get_running_loop() + rpc = JsonRpcClient(None) + rpc._loop = loop + client = CopilotClient(connection=RuntimeConnection.for_uri("localhost:1234")) + client._client = rpc + client._state = "connected" + session = CopilotSession("resumed-session", rpc) + client._sessions[session.session_id] = session + state = {"armed": False, "calls": 0} + background = asyncio.create_task(asyncio.Event().wait()) + + async def send(message): + assert message["method"] == "session.detach" + state["calls"] += 1 + if state["calls"] > 1 and os.environ["WITHHOLD_ALL"] == "0": + rpc._handle_message({"id": message["id"], "result": {"success": True}}) + else: + state["armed"] = True + + rpc._send_message = send + if not hasattr(signal, "SIGALRM"): + # Exercise the actual plugin's exception at the runner boundary on + # Windows, where the real thread timeout would terminate the process. + run_once = loop._run_once + def interrupt_blocked_loop(): + if state["armed"] and not loop._ready: + state["armed"] = False + active_item.config.hook.pytest_timeout_cancel_timer(item=active_item) + pytest_timeout.timeout_sigalrm( + active_item, pytest_timeout._get_item_settings(active_item) + ) + run_once() + loop._run_once = interrupt_blocked_loop + + yield SimpleNamespace(_client=client, session=session, state=state, background=background) + try: + state["armed"] = True + await client.stop() + assert state["calls"] == 2 + assert session._destroyed + assert not session._disconnect_lock.locked() + assert not rpc.pending_requests + finally: + background.cancel() + await asyncio.gather(background, return_exceptions=True) +""", + encoding="utf-8", + ) + (tmp_path / "test_stalled.py").write_text( + """ +import signal +import pytest + +pytestmark = [ + pytest.mark.asyncio(loop_scope="module"), + pytest.mark.timeout( + 1 if hasattr(signal, "SIGALRM") else 20, + method="signal" if hasattr(signal, "SIGALRM") else "thread", + ), +] + +async def test_first_detach_times_out(ctx): + await ctx.session.disconnect() + +async def test_later_test_passes(ctx): + assert ctx.state["calls"] == 1 + assert not ctx.background.done() +""", + encoding="utf-8", + ) + env = dict(os.environ) + env["CANCEL_TEST"] = str(int(cancel_test)) + env["WITHHOLD_ALL"] = str(int(withhold_all)) + env["PYTHONPATH"] = str(Path(__file__).parent.resolve()) + result = subprocess.run( + [ + sys.executable, + "-m", + "pytest", + "-v", + "-s", + "-n", + "0", + "--rootdir", + str(tmp_path), + "--basetemp", + str(tmp_path / "workers"), + str(tmp_path / "test_stalled.py"), + ], + cwd=tmp_path, + env=env, + capture_output=True, + text=True, + timeout=30, + check=False, + ) + output = result.stdout + result.stderr + assert result.returncode == 1, output + assert "1 failed, 1 passed" in output + assert ("1 error" in output) == teardown_error + assert "disconnect_locked=True" in output + assert "outbound method=session.detach" in output + assert "from pytest-timeout" in output + + +@pytest.mark.parametrize("when", ["setup", "call", "teardown"]) +@pytest.mark.parametrize("error", [None, AssertionError("ordinary failure")]) +def test_cleanup_ignores_other_failures(when, error): + call = SimpleNamespace( + when=when, excinfo=None if error is None else SimpleNamespace(value=error) + ) + assert not cancel_timed_out_test(call) diff --git a/python/test_timeout_diagnostics.py b/python/test_timeout_diagnostics.py new file mode 100644 index 0000000000..25a57494c7 --- /dev/null +++ b/python/test_timeout_diagnostics.py @@ -0,0 +1,317 @@ +"""Regression tests for diagnostic evidence retained after an async timeout.""" + +import asyncio +import io +import os +import signal +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from copilot._jsonrpc import JsonRpcClient +from copilot.session import CopilotSession +from e2e import timeout_diagnostics +from e2e.timeout_diagnostics import ( + _dump_awaitable, + _sample_native_threads, + add_timeout_diagnostics, +) + + +@pytest.mark.parametrize("blocked_write", [False, True]) +async def test_timeout_report_identifies_pending_rpc_without_payloads( + tmp_path, monkeypatch, blocked_write +): + rpc = JsonRpcClient(None) + rpc._loop = asyncio.get_running_loop() + sent = asyncio.Event() + release_write = asyncio.Event() + + async def send_message(message): + sent.set() + if blocked_write: + await release_write.wait() + + monkeypatch.setattr(rpc, "_send_message", send_message) + task = asyncio.create_task( + rpc.request( + "session.resume", + {"sessionId": "diagnostic-session", "githubToken": "secret-not-in-diagnostics"}, + ), + name="pending-resume", + ) + try: + await sent.wait() + client = SimpleNamespace(_client=rpc, _sessions={}, _ffi_host=None) + item = SimpleNamespace( + nodeid="test_session_config_e2e.py::test_resume", + config=SimpleNamespace(rootpath=tmp_path), + funcargs={"ctx": SimpleNamespace(_client=client)}, + stash=pytest.Stash(), + ) + call = SimpleNamespace( + excinfo=SimpleNamespace(value=Exception("Timeout (>300.0s) from pytest-timeout.")) + ) + report = SimpleNamespace(failed=True, when="call", sections=[]) + # Snapshotting must not try to acquire a potentially orphaned SDK lock. + with rpc._pending_lock: + add_timeout_diagnostics(item, call, report) + + title, text = report.sections[0] + assert title == "Async timeout diagnostics" + assert "Phase: call" in text + assert "Task pending-resume" in text + assert "JsonRpcClient.request" in text + assert "outbound method=session.resume" in text + assert "session_id=diagnostic-session" in text + assert f"pending request_id={next(iter(rpc.pending_requests))} done=False" in text + assert "pending_locked=True" in text + if blocked_write: + assert ".send_message" in text + else: + assert "awaiting FutureIter" in text + assert "secret-not-in-diagnostics" not in text + assert "githubToken" not in text + (artifact,) = (tmp_path / ".pytest-diagnostics").glob("*.txt") + assert artifact.read_text(encoding="utf-8") == text + assert report.failed + assert not task.done() + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +async def test_teardown_dump_follows_async_generator_and_disconnect_lock(monkeypatch): + rpc = JsonRpcClient(None) + rpc._loop = asyncio.get_running_loop() + sent = asyncio.Event() + + async def send_message(message): + sent.set() + + monkeypatch.setattr(rpc, "_send_message", send_message) + session = CopilotSession("diagnostic-session", rpc) + disconnect = asyncio.create_task(session.disconnect()) + teardown_started = asyncio.Event() + + async def fixture(): + yield + teardown_started.set() + await session.disconnect() + + generator = fixture() + await anext(generator) + + async def finalize(): + await anext(generator) + + finalizer = None + try: + await sent.wait() + finalizer = asyncio.create_task(finalize()) + await teardown_started.wait() + output = io.StringIO() + _dump_awaitable(finalizer.get_coro(), output) + text = output.getvalue() + assert "async_generator_asend" in text + assert ".fixture" in text + assert "CopilotSession.disconnect" in text + assert "Lock.acquire" in text + assert not finalizer.done() + assert not disconnect.done() + finally: + tasks = [disconnect] + ([finalizer] if finalizer is not None else []) + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + await generator.aclose() + + +@pytest.mark.parametrize("error", [None, AssertionError("ordinary failure")]) +def test_non_timeout_does_not_collect_diagnostics(error): + item = SimpleNamespace() + call = SimpleNamespace(excinfo=None if error is None else SimpleNamespace(value=error)) + report = SimpleNamespace(failed=error is not None, sections=[]) + add_timeout_diagnostics(item, call, report) + assert report.sections == [] + + +def test_native_sample_is_bounded_and_targets_this_worker(tmp_path, monkeypatch): + calls = [] + + def run(args, **kwargs): + calls.append((args, kwargs)) + return SimpleNamespace(returncode=0) + + monkeypatch.setattr(subprocess, "run", run) + path = tmp_path / "native.sample.txt" + output = io.StringIO() + _sample_native_threads(path, output) + args, options = calls[0] + assert args == ["sample", str(os.getpid()), "1", "-file", str(path)] + assert options["timeout"] == 10 + assert "Native sample exit=0" in output.getvalue() + + +@pytest.mark.parametrize( + "error", [FileNotFoundError(), subprocess.TimeoutExpired("sample", timeout=10)] +) +def test_native_sample_failure_preserves_diagnostics(tmp_path, monkeypatch, error): + def run(*args, **kwargs): + raise error + + monkeypatch.setattr(subprocess, "run", run) + output = io.StringIO() + _sample_native_threads(tmp_path / "native.sample.txt", output) + assert f"Native sample unavailable: {type(error).__name__}" in output.getvalue() + + +async def test_signal_snapshot_precedes_task_cleanup_and_keeps_original_failure( + tmp_path, monkeypatch +): + installed = [] + failure = pytest.fail.Exception("Timeout (>300.0s) from pytest-timeout.") + + def original_handler(signum, frame): + raise failure + + monkeypatch.setattr(signal, "SIGALRM", 12345, raising=False) + monkeypatch.setattr(signal, "getsignal", lambda signum: original_handler) + monkeypatch.setattr(signal, "signal", lambda signum, handler: installed.append(handler)) + rpc = JsonRpcClient(None) + rpc._loop = asyncio.get_running_loop() + client = SimpleNamespace(_client=rpc, _sessions={}, _ffi_host=None) + item = SimpleNamespace( + nodeid="test_teardown", + config=SimpleNamespace(rootpath=tmp_path), + funcargs={"ctx": SimpleNamespace(_client=client)}, + stash=pytest.Stash(), + ) + hook = timeout_diagnostics.pytest_timeout_set_timer(item, SimpleNamespace(method="signal")) + next(hook) + with pytest.raises(StopIteration): + next(hook) + started = asyncio.Event() + + async def teardown_waiting_for_disconnect(): + started.set() + await asyncio.Future() + + task = asyncio.create_task(teardown_waiting_for_disconnect()) + try: + await started.wait() + with pytest.raises(pytest.fail.Exception) as caught: + installed[0](signal.SIGALRM, None) + assert caught.value is failure + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + report = SimpleNamespace(failed=True, when="teardown", sections=[]) + add_timeout_diagnostics(item, SimpleNamespace(excinfo=SimpleNamespace(value=failure)), report) + text = report.sections[0][1] + assert "Phase: teardown" in text + assert ( + "in test_signal_snapshot_precedes_task_cleanup_and_keeps_original_failure.." in text + ) + assert "teardown_waiting_for_disconnect" in text + assert not item.stash + + +@pytest.mark.parametrize("phase", ["call", "teardown"] if hasattr(signal, "SIGALRM") else ["call"]) +def test_timeout_report_survives_xdist_without_stdout_capture(tmp_path, phase): + # Exercise the real signal timeout on POSIX. On Windows emulate its exception + # at the event-loop boundary, since the thread timeout terminates the worker. + (tmp_path / "conftest.py").write_text( + """ +import pytest +from e2e.timeout_diagnostics import add_timeout_diagnostics +from e2e.timeout_diagnostics import pytest_timeout_set_timer as pytest_timeout_set_timer + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + add_timeout_diagnostics(item, call, outcome.get_result()) +""", + encoding="utf-8", + ) + (tmp_path / "test_stalled.py").write_text( + """ +import asyncio +import signal +import pytest + +@pytest.fixture +def runner(): + with asyncio.Runner() as runner: + yield runner + +def stall(runner): + loop = runner.get_loop() + if not hasattr(signal, "SIGALRM"): + run_once = loop._run_once + def interrupt_loop(): + run_once() + loop._run_once = run_once + pytest.fail("Timeout (>1.0s) from pytest-timeout.") + loop._run_once = interrupt_loop + + async def stalled_rpc(): + await asyncio.Future() + + runner.run(stalled_rpc()) + +@pytest.fixture +def cleanup(runner): + yield + if PHASE == "teardown": + stall(runner) + +@pytest.mark.timeout( + 1 if hasattr(signal, "SIGALRM") else 10, + method="signal" if hasattr(signal, "SIGALRM") else "thread", +) +def test_stalled(runner, cleanup): + if PHASE == "call": + stall(runner) +""", + encoding="utf-8", + ) + with (tmp_path / "test_stalled.py").open("a", encoding="utf-8") as source: + source.write(f"\nPHASE = {phase!r}\n") + env = dict(os.environ) + env["PYTHONPATH"] = str(Path(__file__).parent.resolve()) + result = subprocess.run( + [ + sys.executable, + "-m", + "pytest", + "-v", + "-s", + "-n", + "1", + "--dist=loadfile", + "--basetemp", + str(tmp_path / "workers"), + "--rootdir", + str(tmp_path), + str(tmp_path / "test_stalled.py"), + ], + cwd=tmp_path, + env=env, + capture_output=True, + text=True, + timeout=30, + check=False, + ) + assert result.returncode == 1, result.stdout + result.stderr + assert ("1 failed" if phase == "call" else "1 error") in result.stdout + assert "Async timeout diagnostics" in result.stdout + assert f"Phase: {phase}" in result.stdout + assert "stalled_rpc" in result.stdout + assert "awaiting FutureIter" in result.stdout + (artifact,) = (tmp_path / ".pytest-diagnostics").glob("*.txt") + assert "stalled_rpc" in artifact.read_text(encoding="utf-8")