diff --git a/dotnet/src/Client.cs b/dotnet/src/Client.cs index dff4681a80..5cbca5f9d6 100644 --- a/dotnet/src/Client.cs +++ b/dotnet/src/Client.cs @@ -463,6 +463,11 @@ async Task StartCoreAsync(CancellationToken ct) "CopilotClient.StartAsync protocol verification complete. Elapsed={Elapsed}", startTimestamp); + if (_options.ExtensionLaunchProvider is not null) + { + await connection.Server.RegisterExtensionLaunchProviderAsync(ct); + } + if (_builtinPluginDirectories.Length > 0) { var request = new BuiltinPluginDirectoriesRequest(_builtinPluginDirectories); @@ -2041,8 +2046,7 @@ await Rpc.SessionFs.SetProviderAsync( /// /// Builds the client-global RPC handler bag at construction time. Registers - /// the LLM inference provider adapter and/or the GitHub telemetry adapter - /// depending on which options are configured. The GitHub token dispatcher is + /// the configured connection-level adapters. The GitHub token dispatcher is /// always registered because providers are configured per session. /// private ClientGlobalApiHandlers? BuildClientGlobalApis() @@ -2051,6 +2055,7 @@ await Rpc.SessionFs.SetProviderAsync( var onGitHubTelemetry = _options.OnGitHubTelemetry; return new ClientGlobalApiHandlers { + ExtensionLaunchProvider = _options.ExtensionLaunchProvider, LlmInference = handler is null ? null : new LlmInferenceAdapter(handler, () => _serverRpc), GitHubTelemetry = onGitHubTelemetry is null ? null : new GitHubTelemetryAdapter(onGitHubTelemetry, _logger), GitHubToken = new GitHubTokenAdapter(this), @@ -2698,6 +2703,10 @@ private async Task ConnectToServerAsync(Process? cliProcess, string? { ClientGlobalApiRegistration.RegisterClientGlobalApiHandlers(rpc, _clientGlobalApis); } + if (cliProcess is not null) + { + RegisterRpcProcessExit(cliProcess, rpc); + } rpc.StartListening(); _ = CancelExternalToolsWhenConnectionClosesAsync(rpc); LoggingHelpers.LogTiming(_logger, LogLevel.Debug, null, @@ -2729,6 +2738,35 @@ private async Task ConnectToServerAsync(Process? cliProcess, string? } } + private void RegisterRpcProcessExit(Process cliProcess, JsonRpc rpc) + { + try + { + cliProcess.EnableRaisingEvents = true; + cliProcess.Exited += (_, _) => DisposeRpcAfterProcessExit(rpc); + if (cliProcess.HasExited) + { + DisposeRpcAfterProcessExit(rpc); + } + } + catch (Exception ex) when (ex is InvalidOperationException or ObjectDisposedException) + { + _logger.LogDebug(ex, "Unable to monitor the Copilot CLI process for transport closure"); + } + } + + private void DisposeRpcAfterProcessExit(JsonRpc rpc) + { + try + { + rpc.Dispose(new ConnectionLostException()); + } + catch (Exception ex) when (IsRecoverableConnectionCleanupFailure(ex)) + { + _logger.LogDebug(ex, "Failed to dispose JSON-RPC connection after Copilot CLI process exit"); + } + } + private static bool IsRecoverableConnectionCleanupFailure(Exception exception) => exception is not OutOfMemoryException and not StackOverflowException diff --git a/dotnet/src/JsonRpc.cs b/dotnet/src/JsonRpc.cs index c5c444ea70..33a203f09f 100644 --- a/dotnet/src/JsonRpc.cs +++ b/dotnet/src/JsonRpc.cs @@ -42,7 +42,8 @@ internal sealed partial class JsonRpc : IDisposable private readonly SemaphoreSlim _writeLock = new(1, 1); private readonly CancellationTokenSource _disposeCts = new(); private long _nextId; - private bool _disposed; + private int _disposeStarted; + private Exception? _terminalError; /// /// Initializes a new . @@ -96,6 +97,11 @@ public async Task InvokeAsync(string method, object?[]? args, Cancellation CancellationTokenRegistration cancelRegistration = default; try { + if (Volatile.Read(ref _terminalError) is { } terminalError) + { + throw terminalError; + } + if (cancellationToken.CanBeCanceled) { cancelRegistration = cancellationToken.Register(static state => @@ -136,6 +142,11 @@ await SendMessageAsync(new JsonRpcRequest LogInvokeTiming(LogLevel.Debug, ex, method, id, "Canceled", timingTimestamp); throw; } + catch (ObjectDisposedException ex) when (Volatile.Read(ref _terminalError) is ConnectionLostException) + { + LogInvokeTiming(LogLevel.Warning, ex, method, id, "Failed", timingTimestamp); + throw new ConnectionLostException(); + } catch (Exception ex) { LogInvokeTiming(LogLevel.Warning, ex, method, id, "Failed", timingTimestamp); @@ -183,27 +194,25 @@ public void SetLocalRpcMethod(string methodName, Delegate handler, bool singleOb } /// - public void Dispose() + public void Dispose() => Dispose(new ObjectDisposedException(nameof(JsonRpc))); + + internal void Dispose(Exception reason) { - if (_disposed) + if (Interlocked.Exchange(ref _disposeStarted, 1) != 0) { return; } - _disposed = true; - _disposeCts.Cancel(); - - // Fail all pending requests - foreach (var kvp in _pendingRequests) + FailPendingRequests(reason); + try { - if (_pendingRequests.TryRemove(kvp.Key, out var pending)) - { - pending.TrySetException(new ObjectDisposedException(nameof(JsonRpc))); - } + _disposeCts.Cancel(); + } + finally + { + _completionSource.TrySetResult(); + _writeLock.Dispose(); } - - _completionSource.TrySetResult(); - _writeLock.Dispose(); } private async Task SendMessageAsync(T message, JsonTypeInfo typeInfo, CancellationToken cancellationToken) @@ -338,16 +347,20 @@ private async Task ReadLoopAsync(CancellationToken cancellationToken) } finally { - // Fail all pending requests - foreach (var kvp in _pendingRequests) + FailPendingRequests(new ConnectionLostException()); + _completionSource.TrySetResult(); + } + } + + private void FailPendingRequests(Exception reason) + { + var terminalError = Interlocked.CompareExchange(ref _terminalError, reason, null) ?? reason; + foreach (var kvp in _pendingRequests) + { + if (_pendingRequests.TryRemove(kvp.Key, out var pending)) { - if (_pendingRequests.TryRemove(kvp.Key, out var pending)) - { - pending.TrySetException(new ConnectionLostException()); - } + pending.TrySetException(terminalError); } - - _completionSource.TrySetResult(); } } diff --git a/dotnet/src/Types.cs b/dotnet/src/Types.cs index 4cda151f1d..5e559c4cd0 100644 --- a/dotnet/src/Types.cs +++ b/dotnet/src/Types.cs @@ -319,6 +319,7 @@ private CopilotClientOptions(CopilotClientOptions? other) OnListModels = other.OnListModels; SessionFs = other.SessionFs; RequestHandler = other.RequestHandler; + ExtensionLaunchProvider = other.ExtensionLaunchProvider; OnGitHubTelemetry = other.OnGitHubTelemetry; SessionIdleTimeoutSeconds = other.SessionIdleTimeoutSeconds; EnableRemoteSessions = other.EnableRemoteSessions; @@ -433,6 +434,14 @@ private CopilotClientOptions(CopilotClientOptions? other) [Experimental(Diagnostics.Experimental)] public CopilotRequestHandler? RequestHandler { get; set; } + /// + /// Connection-level extension launch profile provider. + /// When set, the SDK registers the provider during StartAsync() + /// before any session can be created. + /// + [Experimental(Diagnostics.Experimental)] + public IExtensionLaunchProviderHandler? ExtensionLaunchProvider { get; set; } + /// /// Experimental. Receives GitHub telemetry events the runtime forwards to this /// connection; setting a handler opts created/resumed sessions into forwarding. diff --git a/dotnet/test/E2E/ExternalToolCancellationE2ETests.cs b/dotnet/test/E2E/ExternalToolCancellationE2ETests.cs deleted file mode 100644 index b7f34fd037..0000000000 --- a/dotnet/test/E2E/ExternalToolCancellationE2ETests.cs +++ /dev/null @@ -1,62 +0,0 @@ -/*--------------------------------------------------------------------------------------------- - * Copyright (c) Microsoft Corporation. All rights reserved. - *--------------------------------------------------------------------------------------------*/ - -using Microsoft.Extensions.AI; -using System.ComponentModel; -using Xunit; -using Xunit.Abstractions; - -namespace GitHub.Copilot.Test.E2E; - -public class ExternalToolCancellationE2ETests(E2ETestFixture fixture, ITestOutputHelper output) - : E2ETestBase(fixture, "external_tool_cancellation", output) -{ - [Fact] - public async Task Should_Cancel_Tool_Handler_When_Session_Disposes() - { - var toolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - var toolCancelled = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - var releaseTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - - var session = await CreateSessionAsync(new SessionConfig - { - Tools = [AIFunctionFactory.Create(SlowTool, "slow_analysis")], - OnPermissionRequest = PermissionHandler.ApproveAll, - }); - - _ = session.SendAsync(new MessageOptions - { - Prompt = "Use slow_analysis with value 'test_abort'. Wait for the result.", - }); - - var startedValue = await toolStarted.Task.WaitAsync(TimeSpan.FromSeconds(60)); - Assert.Equal("test_abort", startedValue); - - await session.DisposeAsync(); - await toolCancelled.Task.WaitAsync(TimeSpan.FromSeconds(60)); - - releaseTool.TrySetResult("RELEASED"); - - [Description("A slow analysis tool that blocks until released")] - async Task SlowTool([Description("Value to analyze")] string value, CancellationToken cancellationToken) - { - toolStarted.TrySetResult(value); - try - { - var completed = await Task.WhenAny(releaseTool.Task, Task.Delay(Timeout.Infinite, cancellationToken)); - if (completed == releaseTool.Task) - { - return await releaseTool.Task; - } - - throw new OperationCanceledException(cancellationToken); - } - catch (OperationCanceledException) - { - toolCancelled.TrySetResult(true); - throw; - } - } - } -} diff --git a/dotnet/test/E2E/ModeHandlersE2ETests.cs b/dotnet/test/E2E/ModeHandlersE2ETests.cs deleted file mode 100644 index b9f0e69b22..0000000000 --- a/dotnet/test/E2E/ModeHandlersE2ETests.cs +++ /dev/null @@ -1,197 +0,0 @@ -/*--------------------------------------------------------------------------------------------- - * Copyright (c) Microsoft Corporation. All rights reserved. - *--------------------------------------------------------------------------------------------*/ - -using GitHub.Copilot.Test.Harness; -using Xunit; -using Xunit.Abstractions; - -namespace GitHub.Copilot.Test.E2E; - -[Trait(E2ETestTraits.Backend, E2ETestTraits.CapiOnly)] -public class ModeHandlersE2ETests(E2ETestFixture fixture, ITestOutputHelper output) - : E2ETestBase(fixture, "mode_handlers", output) -{ - private const string Token = "mode-handler-token"; - private const string AutoModePrompt = "Explain that auto mode recovered from a rate limit in one short sentence."; - - [Fact] - public async Task Should_Invoke_Exit_Plan_Mode_Handler_When_Model_Uses_Tool() - { - const string summary = "Greeting file implementation plan"; - await ConfigureAuthenticatedUserAsync(); - - var handlerTask = new TaskCompletionSource<(ExitPlanModeRequest Request, ExitPlanModeInvocation Invocation)>( - TaskCreationOptions.RunContinuationsAsynchronously); - - await using var client = CreateAuthenticatedClient(); - var session = await Ctx.CreateSessionAsync(client, new SessionConfig - { - GitHubToken = Token, - OnPermissionRequest = PermissionHandler.ApproveAll, - OnExitPlanModeRequest = (request, invocation) => - { - handlerTask.TrySetResult((request, invocation)); - return Task.FromResult(new ExitPlanModeResult - { - Approved = true, - SelectedAction = "interactive", - Feedback = "Approved by the C# E2E test", - }); - }, - }); - - var requestedEventTask = TestHelper.GetNextEventOfTypeAsync( - session, - evt => evt.Data.Summary == summary, - TimeSpan.FromSeconds(30), - timeoutDescription: "exit_plan_mode.requested event"); - var completedEventTask = TestHelper.GetNextEventOfTypeAsync( - session, - evt => evt.Data.Approved == true && evt.Data.SelectedAction.GetValueOrDefault() == ExitPlanModeAction.Interactive, - TimeSpan.FromSeconds(30), - timeoutDescription: "exit_plan_mode.completed event"); - - var response = await session.SendAndWaitAsync(new MessageOptions - { - AgentMode = AgentMode.Plan, - Prompt = "Create a brief implementation plan for adding a greeting.txt file, then request approval with exit_plan_mode.", - }, timeout: TimeSpan.FromSeconds(120)); - - var (request, invocation) = await handlerTask.Task.WaitAsync(TimeSpan.FromSeconds(30)); - Assert.Equal(session.SessionId, invocation.SessionId); - Assert.Equal(summary, request.Summary); - Assert.Equal(["autopilot", "interactive", "exit_only"], request.Actions); - Assert.Equal("interactive", request.RecommendedAction); - Assert.NotNull(request.PlanContent); - - var requestedEvent = await requestedEventTask; - Assert.Equal(request.Summary, requestedEvent.Data.Summary); - Assert.Equal(request.Actions, requestedEvent.Data.Actions.Select(action => action.Value)); - Assert.Equal(request.RecommendedAction, requestedEvent.Data.RecommendedAction.Value); - - var completedEvent = await completedEventTask; - Assert.True(completedEvent.Data.Approved); - if (completedEvent.Data.SelectedAction is not { } selectedAction) - { - Assert.Fail("Expected a selected action."); - return; - } - - Assert.Equal("interactive", selectedAction.Value); - Assert.Equal("Approved by the C# E2E test", completedEvent.Data.Feedback); - - Assert.NotNull(response); - } - - [Fact] - public async Task Should_Invoke_Auto_Mode_Switch_Handler_When_Rate_Limited() - { - await ConfigureAuthenticatedUserAsync(); - - var handlerTask = new TaskCompletionSource<(AutoModeSwitchRequest Request, AutoModeSwitchInvocation Invocation)>( - TaskCreationOptions.RunContinuationsAsynchronously); - - await using var client = CreateAuthenticatedClient(); - var session = await Ctx.CreateSessionAsync(client, new SessionConfig - { - GitHubToken = Token, - OnPermissionRequest = PermissionHandler.ApproveAll, - OnAutoModeSwitchRequest = (request, invocation) => - { - handlerTask.TrySetResult((request, invocation)); - return Task.FromResult(AutoModeSwitchResponse.Yes); - }, - }); - - const long expectedRetryAfter = 1; - var requestedEventTask = GetNextEventOfTypeAllowingRateLimitAsync( - session, - evt => evt.Data.ErrorCode == "user_weekly_rate_limited" && evt.Data.RetryAfterSeconds == expectedRetryAfter, - TimeSpan.FromSeconds(30), - timeoutDescription: "auto_mode_switch.requested event"); - var completedEventTask = GetNextEventOfTypeAllowingRateLimitAsync( - session, - evt => evt.Data.Response == AutoModeSwitchResponse.Yes, - TimeSpan.FromSeconds(30), - timeoutDescription: "auto_mode_switch.completed event"); - var modelChangeTask = GetNextEventOfTypeAllowingRateLimitAsync( - session, - evt => evt.Data.Cause == "rate_limit_auto_switch", - TimeSpan.FromSeconds(30), - timeoutDescription: "rate-limit auto-mode model change"); - var idleEventTask = GetNextEventOfTypeAllowingRateLimitAsync( - session, - static _ => true, - TimeSpan.FromSeconds(30), - timeoutDescription: "session.idle after auto-mode switch"); - - var messageId = await session.SendAsync(new MessageOptions - { - Prompt = AutoModePrompt, - }); - Assert.NotEmpty(messageId); - - var (request, invocation) = await handlerTask.Task.WaitAsync(TimeSpan.FromSeconds(30)); - Assert.Equal(session.SessionId, invocation.SessionId); - Assert.Equal("user_weekly_rate_limited", request.ErrorCode); - Assert.Equal(1, request.RetryAfterSeconds); - - var requestedEvent = await requestedEventTask; - Assert.Equal(request.ErrorCode, requestedEvent.Data.ErrorCode); - Assert.Equal(expectedRetryAfter, requestedEvent.Data.RetryAfterSeconds); - - var completedEvent = await completedEventTask; - Assert.Equal(AutoModeSwitchResponse.Yes, completedEvent.Data.Response); - - var modelChange = await modelChangeTask; - Assert.Equal("rate_limit_auto_switch", modelChange.Data.Cause); - await idleEventTask; - } - - private CopilotClient CreateAuthenticatedClient() - { - var env = new Dictionary(Ctx.GetEnvironment()) - { - ["COPILOT_DEBUG_GITHUB_API_URL"] = Ctx.ProxyUrl, - }; - - return Ctx.CreateClient(environment: env); - } - - private Task ConfigureAuthenticatedUserAsync() - { - return Ctx.SetCopilotUserByTokenAsync(Token, new CopilotUserConfig( - Login: "mode-handler-user", - CopilotPlan: "individual_pro", - Endpoints: new CopilotUserEndpoints(Api: Ctx.ProxyUrl, Telemetry: "https://localhost:1/telemetry"), - AnalyticsTrackingId: "mode-handler-tracking-id")); - } - - private static async Task GetNextEventOfTypeAllowingRateLimitAsync( - CopilotSession session, - Func predicate, - TimeSpan? timeout = null, - string? timeoutDescription = null) where T : SessionEvent - { - var tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - using var cts = new CancellationTokenSource(timeout ?? TimeSpan.FromSeconds(30)); - - using var subscription = session.On(evt => - { - if (evt is T matched && predicate(matched)) - { - tcs.TrySetResult(matched); - } - else if (evt is SessionErrorEvent { Data.ErrorType: not "rate_limit" } error) - { - tcs.TrySetException(new Exception(error.Data.Message ?? "session error")); - } - }); - - cts.Token.Register(() => tcs.TrySetException( - new TimeoutException($"Timeout waiting for {timeoutDescription ?? $"event of type '{typeof(T).Name}'"}"))); - - return await tcs.Task; - } -} diff --git a/dotnet/test/E2E/RpcSurfaceCoverageE2ETests.cs b/dotnet/test/E2E/RpcSurfaceCoverageE2ETests.cs new file mode 100644 index 0000000000..b969e970aa --- /dev/null +++ b/dotnet/test/E2E/RpcSurfaceCoverageE2ETests.cs @@ -0,0 +1,698 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +/// +/// Direct coverage for public generated RPC methods that are not exercised by another test. +/// The deterministic stdio runtime verifies request serialization and returns property-rich +/// responses so the generated result projections are validated without network or timing inputs. +/// +[Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] +public class RpcSurfaceCoverageE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : E2ETestBase(fixture, "rpc_surface_coverage", output) +{ + [Fact] + public async Task Server_Rpcs_Serialize_Requests_And_Project_Results() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + { + await client.StartAsync(); + + await client.Rpc.RegisterExtensionLaunchProviderAsync(); + + var commands = await client.Rpc.Commands.ListAsync(); + var command = Assert.Single(commands.Commands); + Assert.Equal("rpc-command", command.Name); + Assert.Equal("RPC command", command.Description); + Assert.Equal(["rpc"], command.Aliases); + Assert.True(command.AllowDuringAgentExecution); + Assert.False(command.Experimental); + Assert.Equal("", command.Input!.Hint); + Assert.False(command.Input.PreserveMultilineInput); + Assert.True(command.Input.Required); + Assert.Equal(SlashCommandKind.Builtin, command.Kind); + Assert.True(command.Schedulable); + + var hooks = await client.Rpc.Hooks.DiscoverAsync(["Q:\\rpc-project"], excludeHostHooks: true); + Assert.Empty(hooks.Hooks); + Assert.Equal(["rpc-warning"], hooks.Warnings); + Assert.Empty(hooks.Errors); + + Assert.True((await client.Rpc.LlmInference.SetProviderAsync()).Success); + + var managedSettings = await client.Rpc.ManagedSettings.ReadAsync(); + Assert.Null(managedSettings.ErrorMessage); + Assert.Equal("strict", managedSettings.SettingsJson!.Value.GetProperty("policy").GetString()); + + var install = await client.Rpc.Mcp.PlanInstallAsync( + new CatalogClientContract + { + ProtocolVersion = 3, + RequiredCapabilities = ["mcp-install-planning"], + }, + new McpPlanInstallSourceCandidate + { + CandidateHandle = "candidate-1", + SearchId = "search-1", + }, + McpPlanScope.User); + var unavailable = Assert.IsType(install); + Assert.Equal("The host does not provide installation.", unavailable.Message); + Assert.Equal("host-not-available", unavailable.Reason.Value); + + var model = Assert.Single((await client.Rpc.Models.GetBuiltInCatalogAsync()).Models); + Assert.Equal("built-in-model", model.Id); + + await client.Rpc.Plugins.Builtin.SetAsync(["Q:\\rpc-plugins"]); + + var metadataEntries = await client.Rpc.Sessions.GetClientMetadataAsync( + ["persisted-session"], + ["rpc/key"]); + var metadata = Assert.IsType(Assert.Single(metadataEntries)); + Assert.Equal("persisted-session", metadata.SessionId); + Assert.Equal("rpc-value", metadata.Metadata["rpc/key"]); + + await client.Rpc.Skills.Config.SetSkillDisabledAsync("skill-one", disabled: true); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "registerExtensionLaunchProvider", + "commands.list", + "hooks.discover", + "llmInference.setProvider", + "managedSettings.read", + "mcp.planInstall", + "models.getBuiltInCatalog", + "plugins.builtin.set", + "sessions.getClientMetadata", + "skills.config.setSkillDisabled"); + + var hooksRequest = GetParams(FindRequest(requests, "hooks.discover")); + Assert.True(hooksRequest.GetProperty("excludeHostHooks").GetBoolean()); + Assert.Equal("Q:\\rpc-project", hooksRequest.GetProperty("projectPaths")[0].GetString()); + + var installRequest = GetParams(FindRequest(requests, "mcp.planInstall")); + Assert.Equal(3, installRequest.GetProperty("contract").GetProperty("protocolVersion").GetInt64()); + Assert.Equal("candidate", installRequest.GetProperty("source").GetProperty("kind").GetString()); + Assert.Equal("user", installRequest.GetProperty("scope").GetString()); + + var plugins = GetParams(FindRequest(requests, "plugins.builtin.set")); + Assert.Equal("Q:\\rpc-plugins", plugins.GetProperty("paths")[0].GetString()); + var skill = GetParams(FindRequest(requests, "skills.config.setSkillDisabled")); + Assert.Equal("skill-one", skill.GetProperty("name").GetString()); + Assert.True(skill.GetProperty("disabled").GetBoolean()); + } + + [Fact] + public async Task Session_Control_And_State_Rpcs_Project_All_Result_Properties() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + await using (var session = await Ctx.CreateSessionAsync(client, new SessionConfig())) + { + await session.Rpc.Agent.SetPromptAsync("agent-1", "Use the RPC prompt."); + + var exclusion = await session.Rpc.ContentExclusion.CheckPathsAsync( + ["/tmp/rpc-workspace/file.txt"]); + Assert.True(exclusion.Available); + var pathCheck = Assert.Single(exclusion.Checks); + Assert.Equal("/tmp/rpc-workspace/file.txt", pathCheck.Path); + Assert.False(pathCheck.Excluded); + + var logs = await session.Rpc.Debug.CollectLogsAsync( + new DebugCollectLogsDestinationDirectory { OutputDirectory = "/tmp/rpc-debug" }, + new DebugCollectLogsInclude + { + Events = true, + ProcessLogs = false, + ShellLogs = true, + }, + [ + new DebugCollectLogsEntry + { + BundlePath = "host/diagnostic.txt", + Kind = new DebugCollectLogsEntryKind("file"), + Path = "/tmp/diagnostic.txt", + Required = true, + }, + ]); + Assert.Equal("directory", logs.Kind.Value); + Assert.Equal("/tmp/rpc-debug", logs.Path); + var includedLog = Assert.Single(logs.Entries); + Assert.Equal("host/diagnostic.txt", includedLog.BundlePath); + Assert.Equal(123, includedLog.SizeBytes); + Assert.Equal(DebugCollectLogsSource.Additional, includedLog.Source); + var skippedLog = Assert.Single(logs.SkippedEntries!); + Assert.Equal("host/missing.txt", skippedLog.BundlePath); + Assert.Equal("/tmp/missing.txt", skippedLog.Path); + Assert.Equal("not found", skippedLog.Reason); + + Assert.Equal(4, (await session.Rpc.History.ClearContextAsync("Reset context.")).MessagesCleared); + + var prediction = await session.Rpc.LimitPrediction.PredictAsync( + new SessionLimitPredictionPredictRequest + { + ClientType = new SessionLimitPredictionClientType("sdk"), + ModelId = "model-a", + }); + var unavailable = Assert.IsType(prediction); + Assert.Equal("insufficient-data", unavailable.Reason.Value); + + var clientMetadata = await session.Rpc.Metadata.GetClientMetadataAsync(); + Assert.Equal("rpc-value", clientMetadata["rpc/key"]); + Assert.Equal("other-value", clientMetadata["rpc/other"]); + + var allowed = await session.Rpc.Model.SetAllowedModelsAsync(["model-a", "model-b"]); + Assert.Equal(["model-a", "model-b"], allowed.AllowedModels); + Assert.Equal(["model-a"], allowed.EffectiveAllowedModels); + Assert.Equal("model-a", allowed.FallbackModel); + Assert.Equal("model-a", allowed.ModelId); + + var tier = await session.Rpc.Model.SwitchAutoTierAsync(AutoTier.Intelligence); + Assert.Equal("applied", tier.Status.Value); + Assert.Equal(AutoTier.Intelligence, tier.ActivatingAutoTier); + Assert.Equal(AutoTier.Intelligence, tier.EffectiveAutoTier); + Assert.Null(tier.PendingAutoTier); + Assert.Equal(AutoTier.Balance, tier.SupersededAutoTier); + + var enforcement = await session.Rpc.Sandbox.GetEnforcementStatusAsync(); + Assert.True(enforcement.Required); + Assert.False(enforcement.Blocked); + Assert.Equal("managed-policy", enforcement.Reason); + + var disabled = await session.Rpc.Sandbox.DisableForSessionAsync("sandbox-request-1"); + Assert.True(disabled.Success); + Assert.False(disabled.Enabled); + + var abort = await session.Rpc.AbortAsync(new AbortReason("user")); + Assert.True(abort.Success); + Assert.Null(abort.Error); + + Assert.True((await session.Rpc.InterruptMainTurnAsync(flushQueued: true)).Interrupted); + Assert.Equal(3, await session.Rpc.CancelAllBackgroundAgentsAsync()); + + var log = await session.Rpc.LogAsync( + "RPC log", + level: SessionLogLevel.Warning, + type: "rpc", + ephemeral: true, + url: "https://example.test/rpc", + tip: "Inspect the RPC."); + Assert.Equal(Guid.Parse("11111111-2222-3333-4444-555555555555"), log.EventId); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "session.agent.setPrompt", + "session.contentExclusion.checkPaths", + "session.debug.collectLogs", + "session.history.clearContext", + "session.limitPrediction.predict", + "session.metadata.getClientMetadata", + "session.model.setAllowedModels", + "session.model.switchAutoTier", + "session.sandbox.getEnforcementStatus", + "session.sandbox.disableForSession", + "session.abort", + "session.interruptMainTurn", + "session.cancelAllBackgroundAgents", + "session.log"); + + var logRequest = GetParams(FindRequest(requests, "session.log")); + Assert.Equal("warning", logRequest.GetProperty("level").GetString()); + Assert.True(logRequest.GetProperty("ephemeral").GetBoolean()); + Assert.Equal("https://example.test/rpc", logRequest.GetProperty("url").GetString()); + + var prompt = GetParams(FindRequest(requests, "session.agent.setPrompt")); + Assert.Equal("agent-1", prompt.GetProperty("id").GetString()); + Assert.Equal("Use the RPC prompt.", prompt.GetProperty("prompt").GetString()); + var disable = GetParams(FindRequest(requests, "session.sandbox.disableForSession")); + Assert.Equal("sandbox-request-1", disable.GetProperty("requestId").GetString()); + var interrupt = GetParams(FindRequest(requests, "session.interruptMainTurn")); + Assert.True(interrupt.GetProperty("flushQueued").GetBoolean()); + } + + [Fact] + public async Task Factory_Rpcs_Project_Run_Journal_And_Agent_State() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + await using (var session = await Ctx.CreateSessionAsync(client, new SessionConfig())) + { + var run = await session.Rpc.Factory.RunAsync( + "rpc-factory", + ParseJson("""{ "input": 42 }"""), + new RunOptions + { + Limits = new FactoryRunLimits + { + MaxAiCredits = 2.5, + MaxConcurrentSubagents = 2, + MaxTotalSubagents = 4, + TimeoutSeconds = 30, + }, + LogPhaseNames = true, + NotifyOnComplete = false, + }); + Assert.Equal("factory-run-1", run.RunId); + Assert.Equal(FactoryRunStatus.Running, run.Status); + Assert.Equal(1, run.Attempt); + Assert.Equal("running", run.Result!.Value.GetProperty("value").GetString()); + Assert.Equal(1, run.Snapshot!.Value.GetProperty("step").GetInt32()); + + var resumed = await session.Rpc.Factory.ResumeAsync( + "factory-run-1", + new FactoryRunLimits { MaxTotalSubagents = 8 }, + notifyOnComplete: true, + logPhaseNames: false); + Assert.Equal("rpc-factory", resumed.FactoryName); + Assert.Equal(FactoryRunStatus.Running, resumed.Run.Status); + Assert.Equal(2, resumed.Run.Attempt); + + var current = await session.Rpc.Factory.GetRunAsync("factory-run-1"); + Assert.Equal("factory-run-1", current.RunId); + Assert.Equal(FactoryRunStatus.Running, current.Status); + + var paused = await session.Rpc.Factory.PauseAsync("factory-run-1"); + Assert.Equal(FactoryRunStatus.Paused, paused.Status); + Assert.Equal("caller requested pause", paused.Reason); + Assert.Equal(2, paused.Snapshot!.Value.GetProperty("step").GetInt32()); + + await session.Rpc.Factory.LogAsync( + "factory-run-1", + "execution-token-1", + [ + new FactoryLogLine + { + Kind = FactoryLogLineKind.Log, + Seq = 7, + Text = "Factory progress", + }, + ]); + + var agent = await session.Rpc.Factory.AgentAsync( + "factory-run-1", + "execution-token-1", + "Complete the RPC task.", + new FactoryAgentOptions + { + Agent = "explore", + Label = "rpc-agent", + Model = "model-a", + ReasoningEffort = "high", + }); + Assert.Equal("agent-result", agent.Result!.Value.GetProperty("answer").GetString()); + + var journal = await session.Rpc.Factory.Journal.GetAsync( + "factory-run-1", + "execution-token-1", + "checkpoint"); + Assert.True(journal.Hit); + Assert.Equal(7, journal.ResultJson!.Value.GetProperty("checkpoint").GetInt32()); + + await session.Rpc.Factory.Journal.PutAsync( + "factory-run-1", + "execution-token-1", + "checkpoint", + ParseJson("""{ "checkpoint": 8 }""")); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "session.factory.run", + "session.factory.resume", + "session.factory.getRun", + "session.factory.pause", + "session.factory.log", + "session.factory.agent", + "session.factory.journal.get", + "session.factory.journal.put"); + + var runRequest = GetParams(FindRequest(requests, "session.factory.run")); + Assert.Equal(42, runRequest.GetProperty("args").GetProperty("input").GetInt32()); + Assert.Equal(2.5, runRequest.GetProperty("options").GetProperty("limits").GetProperty("maxAiCredits").GetDouble()); + Assert.True(runRequest.GetProperty("options").GetProperty("logPhaseNames").GetBoolean()); + + var factoryLog = GetParams(FindRequest(requests, "session.factory.log")); + Assert.Equal("execution-token-1", factoryLog.GetProperty("executionToken").GetString()); + var line = Assert.Single(factoryLog.GetProperty("lines").EnumerateArray()); + Assert.Equal("log", line.GetProperty("kind").GetString()); + Assert.Equal(7, line.GetProperty("seq").GetInt64()); + Assert.Equal("Factory progress", line.GetProperty("text").GetString()); + + var journalPut = GetParams(FindRequest(requests, "session.factory.journal.put")); + Assert.Equal("checkpoint", journalPut.GetProperty("key").GetString()); + Assert.Equal(8, journalPut.GetProperty("resultJson").GetProperty("checkpoint").GetInt32()); + } + + [Fact] + public async Task Mcp_Rpcs_Project_Resource_And_Oauth_State() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + await using (var session = await Ctx.CreateSessionAsync(client, new SessionConfig())) + { + Assert.True((await session.Rpc.Mcp.MoveLoadingToBackgroundAsync()).MovedToBackground); + await session.Rpc.Mcp.StartServerAsync( + "rpc-server", + ParseJson("""{ "command": "node", "args": ["server.js"] }""")); + await session.Rpc.Mcp.Oauth.AuthenticationStateChangedAsync( + "rpc-server", + refreshSessionToken: true); + Assert.True((await session.Rpc.Mcp.Oauth.RespondAsync("oauth-request-1")).Success); + + var resources = await session.Rpc.Mcp.Resources.ListAsync("rpc-server", "resource-cursor"); + Assert.Equal("resource-next", resources.NextCursor); + var resource = Assert.Single(resources.Resources); + Assert.Equal("file://rpc/resource.txt", resource.Uri); + Assert.Equal("RPC resource", resource.Name); + Assert.Equal("Resource description", resource.Description); + Assert.Equal("text/plain", resource.MimeType); + Assert.Equal(16, resource.Size); + Assert.Equal("RPC Resource", resource.Title); + + var templates = await session.Rpc.Mcp.Resources.ListTemplatesAsync("rpc-server", "template-cursor"); + Assert.Equal("template-next", templates.NextCursor); + var template = Assert.Single(templates.ResourceTemplates); + Assert.Equal("file://rpc/{name}", template.UriTemplate); + Assert.Equal("RPC template", template.Name); + Assert.Equal("Template description", template.Description); + Assert.Equal("text/plain", template.MimeType); + Assert.Equal("RPC Template", template.Title); + + var read = await session.Rpc.Mcp.Resources.ReadAsync("rpc-server", "file://rpc/resource.txt"); + var content = Assert.Single(read.Contents); + Assert.Equal("file://rpc/resource.txt", content.Uri); + Assert.Equal("text/plain", content.MimeType); + Assert.Equal("resource-content", content.Text); + Assert.Null(content.Blob); + Assert.Equal("assistant", content.Meta!["audience"].GetString()); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "session.mcp.moveLoadingToBackground", + "session.mcp.startServer", + "session.mcp.oauth.authenticationStateChanged", + "session.mcp.oauth.respond", + "session.mcp.resources.list", + "session.mcp.resources.listTemplates", + "session.mcp.resources.read"); + + var start = GetParams(FindRequest(requests, "session.mcp.startServer")); + Assert.Equal("rpc-server", start.GetProperty("serverName").GetString()); + Assert.Equal("node", start.GetProperty("config").GetProperty("command").GetString()); + var auth = GetParams(FindRequest(requests, "session.mcp.oauth.authenticationStateChanged")); + Assert.True(auth.GetProperty("refreshSessionToken").GetBoolean()); + var list = GetParams(FindRequest(requests, "session.mcp.resources.list")); + Assert.Equal("resource-cursor", list.GetProperty("cursor").GetString()); + var listTemplates = GetParams(FindRequest(requests, "session.mcp.resources.listTemplates")); + Assert.Equal("template-cursor", listTemplates.GetProperty("cursor").GetString()); + var readRequest = GetParams(FindRequest(requests, "session.mcp.resources.read")); + Assert.Equal("file://rpc/resource.txt", readRequest.GetProperty("uri").GetString()); + } + + [Fact] + public async Task Tasks_And_Tools_Rpcs_Project_And_Serialize_Complete_State() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + await using (var session = await Ctx.CreateSessionAsync(client, new SessionConfig())) + { + var registered = await session.Rpc.Tasks.RegisterAsync( + TaskClientType.Client, + "client-task-1", + "RPC task", + cancellable: true, + displayName: "RPC Task"); + Assert.True(registered.Created); + Assert.False(registered.Reclaimed); + Assert.Equal("task-1", registered.Task.Id); + Assert.Equal(TaskClientType.Client, registered.Task.Type); + Assert.Equal("client-task-1", registered.Task.ClientTaskId); + Assert.Equal("RPC task", registered.Task.Description); + Assert.Equal("RPC Task", registered.Task.DisplayName); + Assert.Equal(0, registered.Task.Sequence); + Assert.Equal("running", registered.Task.Status.Value); + Assert.Equal(500, registered.Task.ActiveTimeMs); + Assert.True(registered.Task.CanCancel); + Assert.Equal(TaskClientExecutionMode.Background, registered.Task.ExecutionMode); + Assert.Equal("RPC owner", registered.Task.Owner.DisplayName); + Assert.Equal("join-1", registered.Task.Owner.JoinId); + Assert.Equal(TaskClientOwnerKind.Sdk, registered.Task.Owner.Kind); + Assert.Equal("participant-1", registered.Task.Owner.ParticipantId); + Assert.Equal(TaskClientOwnerPresence.Connected, registered.Task.Owner.Presence); + Assert.Equal("rpc-test", registered.Task.Owner.Source); + + var updated = await session.Rpc.Tasks.UpdateAsync( + "task-1", + sequence: 1, + new TaskClientUpdateProgress + { + Message = "Halfway", + Percentage = 50, + Phase = "work", + Status = new TaskClientActiveStatus("running"), + }); + Assert.True(updated.Applied); + Assert.False(updated.Duplicate); + Assert.Equal(1, updated.Task.Sequence); + + var executed = await session.Rpc.Tools.ExecuteAsync( + "rpc_tool", + ParseJson("""{ "value": "input" }"""), + toolCallId: "tool-call-1"); + Assert.Equal("success", executed.GetProperty("resultType").GetString()); + Assert.Equal("executed", executed.GetProperty("textResultForLlm").GetString()); + + var descriptors = await session.Rpc.Tools.GetBuiltinDescriptorsAsync( + reduceUserIntervention: true, + includeAuthor: true, + skillEmbeddingEnabled: false, + shellConfig: new ToolsShellDescriptorConfig + { + DisplayName = "PowerShell", + ShellType = "powershell", + ShellToolName = "shell", + ListShellsToolName = "list_shells", + ReadShellToolName = "read_shell", + StopShellToolName = "stop_shell", + DescriptionLines = ["Runs shell commands."], + }, + shellSupportsPowerShell7Syntax: true, + shellTimeoutMs: 1234, + backgroundTaskNotificationsEnabled: true); + var descriptor = Assert.Single(descriptors.Tools); + Assert.Equal("rpc_builtin", descriptor.Name); + Assert.Equal("RPC built-in tool", descriptor.Description); + Assert.True(descriptor.HasSummariseIntention); + Assert.Equal(BuiltinToolInputSchemaType.Object, descriptor.InputSchema!.Type); + Assert.Equal("Use the RPC built-in.", descriptor.Instructions); + Assert.False(descriptor.IsTerminal); + Assert.True(descriptor.SafeForTelemetry.GetBoolean()); + Assert.Equal("RPC Built-in", descriptor.Title); + Assert.Equal("test", descriptor.Type); + + using var parameterType = JsonDocument.Parse("\"object\""); + await session.Rpc.Tools.SetAsync( + [ + new ProtocolExternalToolDefinition + { + Name = "rpc_external", + Title = "RPC External", + Description = "External RPC tool", + Parameters = new Dictionary + { + ["type"] = parameterType.RootElement.Clone(), + }, + IsTerminal = false, + OverridesBuiltInTool = false, + SkipPermission = true, + }, + ]); + + var completion = await session.Rpc.Tools.TaskCompleteEventDataAsync( + ParseJson("""{ "objectiveId": 17 }"""), + new ToolResultExpanded + { + ResultType = ToolResultType.Success, + TextResultForLlm = "RPC task complete", + SessionLog = "Completion logged.", + }); + Assert.Equal(17, completion.ObjectiveId); + Assert.Equal(TaskCompletionOutcome.Completed, completion.Outcome); + Assert.Equal("completed", completion.Reason); + Assert.True(completion.Success); + Assert.Equal("RPC task complete", completion.Summary); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "session.tasks.register", + "session.tasks.update", + "session.tools.execute", + "session.tools.getBuiltinDescriptors", + "session.tools.set", + "session.tools.taskCompleteEventData"); + + var execute = GetParams(FindRequest(requests, "session.tools.execute")); + Assert.Equal("rpc_tool", execute.GetProperty("name").GetString()); + Assert.Equal("input", execute.GetProperty("arguments").GetProperty("value").GetString()); + Assert.Equal("tool-call-1", execute.GetProperty("toolCallId").GetString()); + + var set = GetParams(FindRequest(requests, "session.tools.set")); + var tool = Assert.Single(set.GetProperty("tools").EnumerateArray()); + Assert.Equal("rpc_external", tool.GetProperty("name").GetString()); + Assert.True(tool.GetProperty("skipPermission").GetBoolean()); + + var register = GetParams(FindRequest(requests, "session.tasks.register")); + Assert.Equal("client", register.GetProperty("type").GetString()); + Assert.Equal("client-task-1", register.GetProperty("clientTaskId").GetString()); + Assert.True(register.GetProperty("cancellable").GetBoolean()); + var update = GetParams(FindRequest(requests, "session.tasks.update")); + Assert.Equal(1, update.GetProperty("sequence").GetInt64()); + Assert.Equal("progress", update.GetProperty("update").GetProperty("kind").GetString()); + + var descriptorsRequest = GetParams(FindRequest(requests, "session.tools.getBuiltinDescriptors")); + Assert.True(descriptorsRequest.GetProperty("reduceUserIntervention").GetBoolean()); + Assert.Equal(1234, descriptorsRequest.GetProperty("shellTimeoutMs").GetInt64()); + Assert.Equal("powershell", descriptorsRequest.GetProperty("shellConfig").GetProperty("shellType").GetString()); + } + + [Fact] + public async Task Workspace_Rpcs_Serialize_Mutations_And_Project_Metadata() + { + var (client, capturePath) = await CreateClientAsync(); + await using (client) + await using (var session = await Ctx.CreateSessionAsync(client, new SessionConfig())) + { + var updated = await session.Rpc.Workspaces.UpdateMetadataAsync( + ParseJson("""{ "owner": "rpc-test" }"""), + name: "Updated RPC workspace"); + Assert.Equal("/tmp/rpc-workspace", updated.Path); + Assert.Equal("workspace-1", updated.Workspace!.Id); + Assert.Equal("/tmp/rpc-workspace", updated.Workspace.Cwd); + Assert.Equal("Updated RPC workspace", updated.Workspace.Name); + Assert.Equal("rpc-branch", updated.Workspace.Branch); + Assert.Equal("rpc-client", updated.Workspace.ClientName); + Assert.Equal(DateTimeOffset.Parse("2026-09-18T11:00:00.000Z"), updated.Workspace.CreatedAt); + Assert.Equal("/tmp/rpc-workspace", updated.Workspace.GitRoot); + Assert.True(updated.Workspace.RemoteSteerable); + + var ensured = await session.Rpc.Workspaces.EnsureAsync(ParseJson("""{ "owner": "rpc-test" }""")); + Assert.Equal("/tmp/rpc-workspace", ensured.Path); + Assert.Equal("RPC workspace", ensured.Workspace!.Name); + + var stat = await session.Rpc.Workspaces.StatFileAsync("folder/file.txt"); + Assert.True(stat.IsFile); + Assert.False(stat.IsDirectory); + Assert.Equal(42, stat.Size); + Assert.Equal(1000, stat.BirthtimeMs); + Assert.Equal(2000, stat.MtimeMs); + + await session.Rpc.Workspaces.CreateDirectoryAsync("folder/nested", recursive: true); + await session.Rpc.Workspaces.RenamePathAsync("folder/file.txt", "folder/renamed.txt"); + await session.Rpc.Workspaces.RemovePathAsync("folder", recursive: true, force: true); + + var summary = await session.Rpc.Workspaces.AddSummaryAsync("RPC summary", "Summary content"); + Assert.NotNull(summary.Summary); + Assert.NotNull(summary.Workspace); + + var truncated = await session.Rpc.Workspaces.TruncateSummariesAsync(keepCount: 2); + Assert.Equal("/tmp/rpc-workspace", truncated.Path); + Assert.Equal("Truncated RPC workspace", truncated.Workspace!.Name); + } + + var requests = await RpcSurfaceTestCli.ReadRequestsAsync(capturePath); + AssertCalledExactlyOnce( + requests, + "session.workspaces.updateMetadata", + "session.workspaces.ensure", + "session.workspaces.statFile", + "session.workspaces.createDirectory", + "session.workspaces.renamePath", + "session.workspaces.removePath", + "session.workspaces.addSummary", + "session.workspaces.truncateSummaries"); + + var createDirectory = GetParams(FindRequest(requests, "session.workspaces.createDirectory")); + Assert.Equal("folder/nested", createDirectory.GetProperty("path").GetString()); + Assert.True(createDirectory.GetProperty("recursive").GetBoolean()); + + var rename = GetParams(FindRequest(requests, "session.workspaces.renamePath")); + Assert.Equal("folder/file.txt", rename.GetProperty("source").GetString()); + Assert.Equal("folder/renamed.txt", rename.GetProperty("destination").GetString()); + + var remove = GetParams(FindRequest(requests, "session.workspaces.removePath")); + Assert.True(remove.GetProperty("recursive").GetBoolean()); + Assert.True(remove.GetProperty("force").GetBoolean()); + + var updateMetadata = GetParams(FindRequest(requests, "session.workspaces.updateMetadata")); + Assert.Equal("rpc-test", updateMetadata.GetProperty("context").GetProperty("owner").GetString()); + Assert.Equal("Updated RPC workspace", updateMetadata.GetProperty("name").GetString()); + var addSummary = GetParams(FindRequest(requests, "session.workspaces.addSummary")); + Assert.Equal("RPC summary", addSummary.GetProperty("title").GetString()); + Assert.Equal("Summary content", addSummary.GetProperty("content").GetString()); + Assert.Equal(2, GetParams(FindRequest(requests, "session.workspaces.truncateSummaries")) + .GetProperty("keepCount").GetInt64()); + } + + private async Task<(CopilotClient Client, string CapturePath)> CreateClientAsync() + { + var (cliPath, capturePath) = await RpcSurfaceTestCli.CreateAsync(Ctx); + var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath]), + UseLoggedInUser = false, + }); + return (client, capturePath); + } + + private static JsonElement FindRequest(JsonElement[] requests, string method) => + Assert.Single(requests, request => request.GetProperty("method").GetString() == method); + + private static JsonElement GetParams(JsonElement request) + { + var parameters = request.GetProperty("params"); + return parameters.ValueKind == JsonValueKind.Array ? parameters[0] : parameters; + } + + private static void AssertCalledExactlyOnce(JsonElement[] requests, params string[] methods) + { + Assert.All( + methods, + method => + { + var request = Assert.Single( + requests, + request => request.GetProperty("method").GetString() == method); + if (method.StartsWith("session.", StringComparison.Ordinal)) + { + Assert.False(string.IsNullOrWhiteSpace(GetParams(request).GetProperty("sessionId").GetString())); + } + }); + } + + private static JsonElement ParseJson(string json) + { + using var document = JsonDocument.Parse(json); + return document.RootElement.Clone(); + } +} diff --git a/dotnet/test/E2E/RpcSurfaceTestCli.cs b/dotnet/test/E2E/RpcSurfaceTestCli.cs new file mode 100644 index 0000000000..9fba398804 --- /dev/null +++ b/dotnet/test/E2E/RpcSurfaceTestCli.cs @@ -0,0 +1,402 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using System.Text.Json; + +namespace GitHub.Copilot.Test.E2E; + +internal static class RpcSurfaceTestCli +{ + public static async Task<(string CliPath, string CapturePath)> CreateAsync(E2ETestContext context) + { + var cliPath = Path.Join(context.WorkDir, $"rpc-surface-cli-{Guid.NewGuid():N}.js"); + var capturePath = Path.Join(context.WorkDir, $"rpc-surface-cli-{Guid.NewGuid():N}.json"); + await File.WriteAllTextAsync(cliPath, Script); + return (cliPath, capturePath); + } + + public static async Task ReadRequestsAsync(string capturePath) + { + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(File.Exists(capturePath)), + timeout: TimeSpan.FromSeconds(10), + timeoutMessage: "Timed out waiting for the RPC surface request capture."); + + using var capture = JsonDocument.Parse(await File.ReadAllTextAsync(capturePath)); + return capture.RootElement.GetProperty("requests").EnumerateArray().Select(request => request.Clone()).ToArray(); + } + + private const string Script = """ + const fs = require("fs"); + + const captureIndex = process.argv.indexOf("--capture-file"); + const captureFile = process.argv[captureIndex + 1]; + const requests = []; + let buffer = Buffer.alloc(0); + + function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ requests })); + } + + function writeResponse(id, result) { + const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function params(message) { + return Array.isArray(message.params) ? (message.params[0] ?? {}) : (message.params ?? {}); + } + + function workspace(name = "RPC workspace") { + return { + path: "/tmp/rpc-workspace", + workspace: { + id: "workspace-1", + cwd: "/tmp/rpc-workspace", + name, + branch: "rpc-branch", + client_name: "rpc-client", + created_at: "2026-09-18T11:00:00.000Z", + git_root: "/tmp/rpc-workspace", + remote_steerable: true + } + }; + } + + function task(sequence, status) { + return { + id: "task-1", + type: "client", + clientTaskId: "client-task-1", + description: "RPC task", + displayName: "RPC Task", + activeStartedAt: "2026-09-18T12:00:00.500Z", + activeTimeMs: 500, + canCancel: true, + executionMode: "background", + owner: { + displayName: "RPC owner", + joinId: "join-1", + kind: "sdk", + participantId: "participant-1", + presence: "connected", + source: "rpc-test" + }, + sequence, + status, + startedAt: "2026-09-18T12:00:00.000Z", + updatedAt: "2026-09-18T12:00:01.000Z" + }; + } + + function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) { + return; + } + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + switch (message.method) { + case "connect": + writeResponse(message.id, { ok: true, protocolVersion: 3, version: "rpc-surface-test" }); + return; + case "session.create": + writeResponse(message.id, { + sessionId: params(message).sessionId ?? "rpc-surface-session", + workspacePath: "/tmp/rpc-workspace", + capabilities: null + }); + return; + case "commands.list": + writeResponse(message.id, { + commands: [{ + name: "rpc-command", + description: "RPC command", + aliases: ["rpc"], + allowDuringAgentExecution: true, + experimental: false, + input: { + hint: "", + preserveMultilineInput: false, + required: true + }, + kind: "builtin", + schedulable: true + }] + }); + return; + case "hooks.discover": + writeResponse(message.id, { hooks: [], warnings: ["rpc-warning"], errors: [] }); + return; + case "llmInference.setProvider": + writeResponse(message.id, { success: true }); + return; + case "managedSettings.read": + writeResponse(message.id, { settingsJson: { policy: "strict" }, errorMessage: null }); + return; + case "mcp.planInstall": + writeResponse(message.id, { + kind: "unavailable", + message: "The host does not provide installation.", + reason: "host-not-available" + }); + return; + case "models.getBuiltInCatalog": + writeResponse(message.id, { + models: [{ + id: "built-in-model", + name: "Built-in Model", + family: "test-family" + }] + }); + return; + case "sessions.getClientMetadata": + writeResponse(message.id, [{ + status: "ok", + sessionId: "persisted-session", + metadata: { "rpc/key": "rpc-value" } + }]); + return; + case "session.contentExclusion.checkPaths": + writeResponse(message.id, { + available: true, + checks: [{ path: "/tmp/rpc-workspace/file.txt", excluded: false }] + }); + return; + case "session.debug.collectLogs": + writeResponse(message.id, { + kind: "directory", + path: "/tmp/rpc-debug", + entries: [{ + bundlePath: "host/diagnostic.txt", + sizeBytes: 123, + source: "additional" + }], + skippedEntries: [{ + bundlePath: "host/missing.txt", + path: "/tmp/missing.txt", + reason: "not found" + }] + }); + return; + case "session.factory.run": + case "session.factory.getRun": + writeResponse(message.id, { + runId: "factory-run-1", + status: "running", + attempt: 1, + result: { value: "running" }, + snapshot: { step: 1 } + }); + return; + case "session.factory.pause": + writeResponse(message.id, { + runId: "factory-run-1", + status: "paused", + attempt: 1, + reason: "caller requested pause", + snapshot: { step: 2 } + }); + return; + case "session.factory.resume": + writeResponse(message.id, { + factoryName: "rpc-factory", + run: { + runId: "factory-run-1", + status: "running", + attempt: 2, + snapshot: { step: 3 } + } + }); + return; + case "session.factory.log": + case "session.factory.journal.put": + case "session.tools.set": + writeResponse(message.id, {}); + return; + case "session.factory.agent": + writeResponse(message.id, { result: { answer: "agent-result" } }); + return; + case "session.factory.journal.get": + writeResponse(message.id, { hit: true, resultJson: { checkpoint: 7 } }); + return; + case "session.history.clearContext": + writeResponse(message.id, { messagesCleared: 4 }); + return; + case "session.limitPrediction.predict": + writeResponse(message.id, { kind: "unavailable", reason: "insufficient-data" }); + return; + case "session.mcp.moveLoadingToBackground": + writeResponse(message.id, { movedToBackground: true }); + return; + case "session.mcp.oauth.respond": + writeResponse(message.id, { success: true }); + return; + case "session.mcp.resources.list": + writeResponse(message.id, { + nextCursor: "resource-next", + resources: [{ + uri: "file://rpc/resource.txt", + name: "RPC resource", + description: "Resource description", + mimeType: "text/plain", + size: 16, + title: "RPC Resource" + }] + }); + return; + case "session.mcp.resources.listTemplates": + writeResponse(message.id, { + nextCursor: "template-next", + resourceTemplates: [{ + uriTemplate: "file://rpc/{name}", + name: "RPC template", + description: "Template description", + mimeType: "text/plain", + title: "RPC Template" + }] + }); + return; + case "session.mcp.resources.read": + writeResponse(message.id, { + contents: [{ + uri: "file://rpc/resource.txt", + mimeType: "text/plain", + text: "resource-content", + _meta: { audience: "assistant" } + }] + }); + return; + case "session.metadata.getClientMetadata": + writeResponse(message.id, { "rpc/key": "rpc-value", "rpc/other": "other-value" }); + return; + case "session.model.setAllowedModels": + writeResponse(message.id, { + allowedModels: ["model-a", "model-b"], + effectiveAllowedModels: ["model-a"], + fallbackModel: "model-a", + modelId: "model-a" + }); + return; + case "session.model.switchAutoTier": + writeResponse(message.id, { + status: "applied", + activatingAutoTier: "intelligence", + effectiveAutoTier: "intelligence", + pendingAutoTier: null, + supersededAutoTier: "balance" + }); + return; + case "session.sandbox.getEnforcementStatus": + writeResponse(message.id, { required: true, blocked: false, reason: "managed-policy" }); + return; + case "session.sandbox.disableForSession": + writeResponse(message.id, { success: true, enabled: false }); + return; + case "session.abort": + writeResponse(message.id, { success: true, error: null }); + return; + case "session.interruptMainTurn": + writeResponse(message.id, { interrupted: true }); + return; + case "session.cancelAllBackgroundAgents": + writeResponse(message.id, 3); + return; + case "session.log": + writeResponse(message.id, { eventId: "11111111-2222-3333-4444-555555555555" }); + return; + case "session.tasks.register": + writeResponse(message.id, { created: true, reclaimed: false, task: task(0, "running") }); + return; + case "session.tasks.update": + writeResponse(message.id, { applied: true, duplicate: false, task: task(1, "running") }); + return; + case "session.tools.execute": + writeResponse(message.id, { resultType: "success", textResultForLlm: "executed" }); + return; + case "session.tools.getBuiltinDescriptors": + writeResponse(message.id, { + tools: [{ + name: "rpc_builtin", + description: "RPC built-in tool", + hasSummariseIntention: true, + inputSchema: { type: "object" }, + instructions: "Use the RPC built-in.", + isTerminal: false, + safeForTelemetry: true, + title: "RPC Built-in", + type: "test" + }] + }); + return; + case "session.tools.taskCompleteEventData": + writeResponse(message.id, { + objectiveId: 17, + outcome: "completed", + reason: "completed", + success: true, + summary: "RPC task complete" + }); + return; + case "session.workspaces.updateMetadata": + writeResponse(message.id, workspace("Updated RPC workspace")); + return; + case "session.workspaces.ensure": + writeResponse(message.id, workspace()); + return; + case "session.workspaces.statFile": + writeResponse(message.id, { + birthtimeMs: 1000, + isDirectory: false, + isFile: true, + mtimeMs: 2000, + size: 42 + }); + return; + case "session.workspaces.addSummary": + writeResponse(message.id, { + summary: { number: 3, title: "RPC summary", content: "Summary content" }, + workspace: { id: "workspace-1", cwd: "/tmp/rpc-workspace", name: "RPC workspace" } + }); + return; + case "session.workspaces.truncateSummaries": + writeResponse(message.id, workspace("Truncated RPC workspace")); + return; + default: + writeResponse(message.id, {}); + } + } + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) { + return; + } + + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) { + throw new Error("Missing Content-Length header"); + } + + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + Number(match[1]); + if (buffer.length < bodyEnd) { + return; + } + + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } + }); + + process.stdin.resume(); + saveCapture(); + """; +} diff --git a/dotnet/test/E2E/ScenarioTestingCallbacksE2ETests.cs b/dotnet/test/E2E/ScenarioTestingCallbacksE2ETests.cs new file mode 100644 index 0000000000..d71c7d6aae --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingCallbacksE2ETests.cs @@ -0,0 +1,456 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using System.ComponentModel; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +[Trait(E2ETestTraits.Backend, E2ETestTraits.CapiOnly)] +public class ScenarioTestingCallbacksE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_callbacks", output) +{ + private const string ModeHandlerToken = "scenario-client-mode-handler-token"; + private const string AutoModePrompt = "Explain that the scenario client recovered from a rate limit in one short sentence."; + + [Fact] + public async Task Should_Run_Scenario_Prompt_And_Tool_Hooks_With_Full_Context_And_Suppression() + { + UserPromptSubmittedHookInput? submitted = null; + UserPromptTransformedHookInput? transformed = null; + PreToolUseHookInput? preTool = null; + PostToolUseHookInput? postTool = null; + CopilotSession? session = null; + + session = await CreateSessionAsync(new SessionConfig + { + Tools = [AIFunctionFactory.Create(ScenarioHookTool, "scenario_hook_tool")], + Hooks = new SessionHooks + { + OnUserPromptSubmitted = (input, invocation) => + { + Assert.Equal(session!.SessionId, invocation.SessionId); + submitted = input; + return Task.FromResult(new UserPromptSubmittedHookOutput + { + SuppressOutput = true, + }); + }, + OnUserPromptTransformed = (input, invocation) => + { + Assert.Equal(session!.SessionId, invocation.SessionId); + transformed = input; + return Task.FromResult(new UserPromptTransformedHookOutput + { + ModifiedTransformedPrompt = + "Call scenario_hook_tool with value 'original', then reply with exactly SCENARIO_POST_RESULT.", + }); + }, + OnPreToolUse = (input, invocation) => + { + if (input.ToolName != "scenario_hook_tool") + { + return Task.FromResult(new PreToolUseHookOutput + { + PermissionDecision = "allow", + }); + } + + Assert.Equal(session!.SessionId, invocation.SessionId); + preTool = input; + return Task.FromResult(new PreToolUseHookOutput + { + PermissionDecision = "allow", + ModifiedArgs = new Dictionary { ["value"] = "pre-hook" }, + SuppressOutput = false, + }); + }, + OnPostToolUse = (input, invocation) => + { + if (input.ToolName != "scenario_hook_tool") + { + return Task.FromResult(null); + } + + Assert.Equal(session!.SessionId, invocation.SessionId); + postTool = input; + return Task.FromResult(new PostToolUseHookOutput + { + ModifiedResult = new ToolResultObject + { + TextResultForLlm = "SCENARIO_POST_RESULT", + ResultType = "success", + ToolTelemetry = new Dictionary(), + }, + SuppressOutput = false, + }); + }, + }, + }); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Original hidden scenario hook prompt.", + DisplayPrompt = "Run scenario hook pipeline", + Source = MessageSource.Agent("scenario-client"), + }); + + AssertHookContext(submitted, session.SessionId); + AssertHookContext(transformed, session.SessionId); + AssertHookContext(preTool, session.SessionId); + AssertHookContext(postTool, session.SessionId); + Assert.Equal("original", preTool!.ToolArgs!.Value.GetProperty("value").GetString()); + Assert.Equal("pre-hook", postTool!.ToolArgs!.Value.GetProperty("value").GetString()); + Assert.Contains("SCENARIO_TOOL_PRE-HOOK", postTool.ToolResult!.Value.ToString(), StringComparison.OrdinalIgnoreCase); + Assert.Contains("SCENARIO_POST_RESULT", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + [Description("Returns a scenario-owned hook value")] + static string ScenarioHookTool([Description("Value to transform")] string value) => + $"SCENARIO_TOOL_{value.ToUpperInvariant()}"; + } + + [Fact] + public async Task Should_Handle_Scenario_User_Input_And_Form_Url_Elicitation_Outcomes() + { + var events = new List(); + var allEventsReceived = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "emit-ui-events"]), + UseLoggedInUser = false, + }); + + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + OnEvent = evt => + { + lock (events) + { + events.Add(evt); + if (events.Count(evt => evt is UserInputRequestedEvent or ElicitationRequestedEvent) == 4) + { + allEventsReceived.TrySetResult(); + } + } + }, + }); + + await allEventsReceived.Task.WaitAsync(TimeSpan.FromSeconds(30)); + + UserInputRequestedEvent userInput; + List elicitations; + lock (events) + { + userInput = Assert.Single(events.OfType()); + elicitations = events.OfType().ToList(); + } + + Assert.Equal("Choose a scenario action", userInput.Data.Question); + Assert.NotNull(userInput.Data.Choices); + Assert.Equal(["Approve", "Decline"], userInput.Data.Choices); + Assert.True(userInput.Data.AllowFreeform); + Assert.True((await session.Rpc.Ui.HandlePendingUserInputAsync( + userInput.Data.RequestId, + new UIUserInputResponse { Answer = "Approve", WasFreeform = false })).Success); + + var form = Assert.Single(elicitations, evt => evt.Data.RequestId == "scenario-form-accept"); + Assert.Equal(ElicitationRequestedMode.Form, form.Data.Mode); + Assert.Equal("name", Assert.Single(form.Data.RequestedSchema!.Properties).Key); + Assert.True((await session.Rpc.Ui.HandlePendingElicitationAsync( + form.Data.RequestId, + new UIElicitationResponse + { + Action = UIElicitationResponseAction.Accept, + Content = new Dictionary + { + ["name"] = JsonDocument.Parse("\"Mona\"").RootElement.Clone(), + }, + })).Success); + + var url = Assert.Single(elicitations, evt => evt.Data.RequestId == "scenario-url-decline"); + Assert.Equal(ElicitationRequestedMode.Url, url.Data.Mode); + Assert.Equal("https://example.test/authorize", url.Data.Url); + Assert.True((await session.Rpc.Ui.HandlePendingElicitationAsync( + url.Data.RequestId, + new UIElicitationResponse { Action = UIElicitationResponseAction.Decline })).Success); + + var cancelled = Assert.Single(elicitations, evt => evt.Data.RequestId == "scenario-form-cancel"); + Assert.True((await session.Rpc.Ui.HandlePendingElicitationAsync( + cancelled.Data.RequestId, + new UIElicitationResponse { Action = UIElicitationResponseAction.Cancel })).Success); + + var stale = await session.Rpc.Ui.HandlePendingElicitationAsync( + "stale-scenario-request", + new UIElicitationResponse { Action = UIElicitationResponseAction.Cancel }); + Assert.False(stale.Success); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + var userInputResponse = RequestParameters(Assert.Single( + requests, + request => request.GetProperty("method").GetString() == "session.ui.handlePendingUserInput")); + Assert.Equal("scenario-user-input", userInputResponse.GetProperty("requestId").GetString()); + Assert.Equal("Approve", userInputResponse.GetProperty("response").GetProperty("answer").GetString()); + Assert.False(userInputResponse.GetProperty("response").GetProperty("wasFreeform").GetBoolean()); + + var elicitationResponses = requests + .Where(request => request.GetProperty("method").GetString() == "session.ui.handlePendingElicitation") + .Select(RequestParameters) + .ToDictionary(request => request.GetProperty("requestId").GetString()!); + Assert.Equal("accept", elicitationResponses["scenario-form-accept"].GetProperty("result").GetProperty("action").GetString()); + Assert.Equal( + "Mona", + elicitationResponses["scenario-form-accept"].GetProperty("result").GetProperty("content").GetProperty("name").GetString()); + Assert.Equal("decline", elicitationResponses["scenario-url-decline"].GetProperty("result").GetProperty("action").GetString()); + Assert.Equal("cancel", elicitationResponses["scenario-form-cancel"].GetProperty("result").GetProperty("action").GetString()); + Assert.Equal("cancel", elicitationResponses["stale-scenario-request"].GetProperty("result").GetProperty("action").GetString()); + } + + [Fact] + public async Task Should_Cancel_Scenario_Host_Callback_When_Channel_Disconnects() + { + var callbackStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var callbackCancelled = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var client = Ctx.CreateClient(); + var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + Tools = [AIFunctionFactory.Create(BlockingHostCallback, "scenario_host_callback")], + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + _ = session.SendAsync(new MessageOptions + { + Prompt = "Call scenario_host_callback with value 'disconnect' and wait for it.", + DisplayPrompt = "Run disconnectable scenario callback", + Source = MessageSource.Agent("scenario-client"), + }); + + await callbackStarted.Task.WaitAsync(TimeSpan.FromSeconds(60)); + await client.ForceStopAsync(); + await callbackCancelled.Task.WaitAsync(TimeSpan.FromSeconds(60)); + + [Description("Waits for scenario host channel cancellation")] + async Task BlockingHostCallback( + [Description("Callback value")] string value, + CancellationToken cancellationToken) + { + Assert.Equal("disconnect", value); + callbackStarted.TrySetResult(); + try + { + await Task.Delay(Timeout.Infinite, cancellationToken); + return "UNREACHABLE"; + } + catch (OperationCanceledException) + { + callbackCancelled.TrySetResult(); + throw; + } + } + } + + [Fact] + public async Task Should_Approve_Scenario_Exit_Plan_With_Full_Callback_And_Event_State() + { + const string summary = "scenario client implementation plan"; + await ConfigureAuthenticatedUserAsync(); + + var callback = new TaskCompletionSource<(ExitPlanModeRequest Request, ExitPlanModeInvocation Invocation)>( + TaskCreationOptions.RunContinuationsAsynchronously); + await using var client = CreateAuthenticatedClient(); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + GitHubToken = ModeHandlerToken, + OnPermissionRequest = PermissionHandler.ApproveAll, + OnExitPlanModeRequest = (request, invocation) => + { + callback.TrySetResult((request, invocation)); + return Task.FromResult(new ExitPlanModeResult + { + Approved = true, + SelectedAction = "interactive", + Feedback = "Approved by the scenario client", + }); + }, + }); + + var userMessageTask = TestHelper.GetNextEventOfTypeAsync( + session, + evt => evt.Data.Source == "agent-scenario-client", + TimeSpan.FromSeconds(30), + timeoutDescription: "scenario client exit-plan user message"); + var requestedTask = TestHelper.GetNextEventOfTypeAsync( + session, + evt => evt.Data.Summary == summary, + TimeSpan.FromSeconds(30), + timeoutDescription: "scenario client exit-plan request"); + var completedTask = TestHelper.GetNextEventOfTypeAsync( + session, + evt => evt.Data.Approved == true && + evt.Data.SelectedAction.GetValueOrDefault() == ExitPlanModeAction.Interactive, + TimeSpan.FromSeconds(30), + timeoutDescription: "scenario client exit-plan completion"); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + AgentMode = AgentMode.Plan, + Prompt = "Create a scenario client plan, then request approval with exit_plan_mode.", + DisplayPrompt = "Review proposed scenario client plan", + Source = MessageSource.Agent("scenario-client"), + }, timeout: TimeSpan.FromSeconds(120)); + + var userMessage = await userMessageTask; + Assert.Equal("Review proposed scenario client plan", userMessage.Data.Content); + Assert.Equal(UserMessageAgentMode.Plan, userMessage.Data.AgentMode); + + var (request, invocation) = await callback.Task.WaitAsync(TimeSpan.FromSeconds(30)); + Assert.Equal(session.SessionId, invocation.SessionId); + Assert.Equal(summary, request.Summary); + Assert.Equal(["autopilot", "interactive", "exit_only"], request.Actions); + Assert.Equal("interactive", request.RecommendedAction); + Assert.NotNull(request.PlanContent); + + var requested = await requestedTask; + Assert.Equal(request.Summary, requested.Data.Summary); + Assert.Equal(request.Actions, requested.Data.Actions.Select(action => action.Value)); + Assert.Equal(request.RecommendedAction, requested.Data.RecommendedAction.Value); + + var completed = await completedTask; + Assert.True(completed.Data.Approved); + Assert.Equal(ExitPlanModeAction.Interactive, completed.Data.SelectedAction); + Assert.Equal("Approved by the scenario client", completed.Data.Feedback); + Assert.NotNull(response); + } + + [Fact] + public async Task Should_Auto_Switch_Scenario_Mode_After_Rate_Limit() + { + await ConfigureAuthenticatedUserAsync(); + + var callback = new TaskCompletionSource<(AutoModeSwitchRequest Request, AutoModeSwitchInvocation Invocation)>( + TaskCreationOptions.RunContinuationsAsynchronously); + await using var client = CreateAuthenticatedClient(); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + GitHubToken = ModeHandlerToken, + OnPermissionRequest = PermissionHandler.ApproveAll, + OnAutoModeSwitchRequest = (request, invocation) => + { + callback.TrySetResult((request, invocation)); + return Task.FromResult(AutoModeSwitchResponse.Yes); + }, + }); + + const long expectedRetryAfter = 1; + var userMessageTask = GetNextEventAllowingRateLimitAsync( + session, + evt => evt.Data.Source == "agent-scenario-client", + "scenario client auto-switch user message"); + var requestedTask = GetNextEventAllowingRateLimitAsync( + session, + evt => evt.Data.ErrorCode == "user_weekly_rate_limited" && + evt.Data.RetryAfterSeconds == expectedRetryAfter, + "scenario client auto-switch request"); + var completedTask = GetNextEventAllowingRateLimitAsync( + session, + evt => evt.Data.Response == AutoModeSwitchResponse.Yes, + "scenario client auto-switch completion"); + var modelChangeTask = GetNextEventAllowingRateLimitAsync( + session, + evt => evt.Data.Cause == "rate_limit_auto_switch", + "scenario client rate-limit model change"); + var idleTask = GetNextEventAllowingRateLimitAsync( + session, + static _ => true, + "scenario client auto-switch idle"); + + var messageId = await session.SendAsync(new MessageOptions + { + Prompt = AutoModePrompt, + DisplayPrompt = "Continue scenario client request automatically", + Source = MessageSource.Agent("scenario-client"), + }); + Assert.NotEmpty(messageId); + + var userMessage = await userMessageTask; + Assert.Equal("Continue scenario client request automatically", userMessage.Data.Content); + + var (request, invocation) = await callback.Task.WaitAsync(TimeSpan.FromSeconds(30)); + Assert.Equal(session.SessionId, invocation.SessionId); + Assert.Equal("user_weekly_rate_limited", request.ErrorCode); + Assert.Equal(expectedRetryAfter, request.RetryAfterSeconds); + + var requested = await requestedTask; + Assert.Equal(request.ErrorCode, requested.Data.ErrorCode); + Assert.Equal(request.RetryAfterSeconds, requested.Data.RetryAfterSeconds); + Assert.Equal(AutoModeSwitchResponse.Yes, (await completedTask).Data.Response); + Assert.Equal("rate_limit_auto_switch", (await modelChangeTask).Data.Cause); + await idleTask; + } + + private CopilotClient CreateAuthenticatedClient() + { + var environment = new Dictionary(Ctx.GetEnvironment()) + { + ["COPILOT_DEBUG_GITHUB_API_URL"] = Ctx.ProxyUrl, + }; + + return Ctx.CreateClient(environment: environment); + } + + private Task ConfigureAuthenticatedUserAsync() => + Ctx.SetCopilotUserByTokenAsync(ModeHandlerToken, new CopilotUserConfig( + Login: "scenario-client-mode-handler-user", + CopilotPlan: "individual_pro", + Endpoints: new CopilotUserEndpoints(Api: Ctx.ProxyUrl, Telemetry: "https://localhost:1/telemetry"), + AnalyticsTrackingId: "scenario-client-mode-handler-tracking-id")); + + private static async Task GetNextEventAllowingRateLimitAsync( + CopilotSession session, + Func predicate, + string description) where T : SessionEvent + { + var result = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(30)); + using var subscription = session.On(evt => + { + if (evt is T typed && predicate(typed)) + { + result.TrySetResult(typed); + } + else if (evt is SessionErrorEvent { Data.ErrorType: not "rate_limit" } error) + { + result.TrySetException(new Exception(error.Data.Message ?? "session error")); + } + }); + + using var registration = timeout.Token.Register( + () => result.TrySetException(new TimeoutException($"Timed out waiting for {description}."))); + return await result.Task; + } + + private static void AssertHookContext(object? input, string sessionId) + { + Assert.NotNull(input); + var type = input.GetType(); + Assert.Equal(sessionId, type.GetProperty("SessionId")!.GetValue(input)); + Assert.True((DateTimeOffset)type.GetProperty("Timestamp")!.GetValue(input)! > DateTimeOffset.UnixEpoch); + Assert.False(string.IsNullOrWhiteSpace((string)type.GetProperty("WorkingDirectory")!.GetValue(input)!)); + } + + private static JsonElement RequestParameters(JsonElement request) + { + var parameters = request.GetProperty("params"); + return parameters.ValueKind == JsonValueKind.Array ? parameters[0] : parameters; + } +} diff --git a/dotnet/test/E2E/ScenarioTestingCanvasE2ETests.cs b/dotnet/test/E2E/ScenarioTestingCanvasE2ETests.cs new file mode 100644 index 0000000000..32c0c92bb2 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingCanvasE2ETests.cs @@ -0,0 +1,393 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingCanvasE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_canvas", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Run_Ordered_Scenario_Canvas_Lifecycle_With_Exact_Context_And_Snapshot() + { + var handler = new RecordingCanvasHandler(); + await using var session = await CreateSessionAsync(CreateSessionConfig(handler)); + + CanvasList list = await WaitForCanvasRegistryAsync(session); + var canvas = Assert.Single(list.Canvases); + Assert.Equal("scenario:builtin:e2e-window", canvas.ExtensionId); + Assert.Equal("scenario-inspector", canvas.CanvasId); + Assert.Equal("Scenario Inspector", canvas.DisplayName); + Assert.Equal("Displays scenario-owned state.", canvas.Description); + Assert.Equal("object", canvas.InputSchema!.Value.GetProperty("type").GetString()); + var action = Assert.Single(canvas.Actions!); + Assert.Equal("replace", action.Name); + Assert.Equal("Replaces the displayed value.", action.Description); + Assert.Equal("object", action.InputSchema!.Value.GetProperty("type").GetString()); + + var opened = await session.Rpc.Canvas.OpenAsync( + canvasId: "scenario-inspector", + instanceId: "scenario-inspector-1", + extensionId: canvas.ExtensionId, + input: new Dictionary { ["value"] = "before" }); + + Assert.Equal("ready", opened.Status); + Assert.Equal("Scenario Inspector: before", opened.Title); + Assert.Equal("https://example.test/scenario-inspector/scenario-inspector-1", opened.Url); + AssertRequest(handler.OpenRequests.Single(), session.SessionId, "scenario-inspector-1"); + Assert.Equal("before", handler.OpenRequests[0].Input!.Value.GetProperty("value").GetString()); + + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(session.OpenCanvases.Count == 1), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the scenario canvas snapshot."); + AssertOpenCanvas(Assert.Single(session.OpenCanvases), "scenario-inspector-1", "before"); + + var actionResult = await session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "scenario-inspector-1", + actionName: "replace", + input: new Dictionary { ["value"] = "after" }); + + Assert.Equal("after", actionResult.Result!.Value.GetProperty("value").GetString()); + AssertRequest(handler.ActionRequests.Single(), session.SessionId, "scenario-inspector-1"); + Assert.Equal("replace", handler.ActionRequests[0].ActionName); + Assert.Equal("after", handler.ActionRequests[0].Input!.Value.GetProperty("value").GetString()); + + var liveSnapshot = Assert.Single((await session.Rpc.Canvas.ListOpenAsync()).OpenCanvases); + AssertOpenCanvas(liveSnapshot, "scenario-inspector-1", "before"); + + await session.Rpc.Canvas.CloseAsync("scenario-inspector-1"); + + AssertRequest(handler.CloseRequests.Single(), session.SessionId, "scenario-inspector-1"); + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(session.OpenCanvases.Count == 0), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the scenario canvas to close."); + Assert.Empty((await session.Rpc.Canvas.ListOpenAsync()).OpenCanvases); + Assert.Equal( + ["open:scenario-inspector-1", "action:scenario-inspector-1:replace", "close:scenario-inspector-1"], + handler.Callbacks); + } + + [Theory] + [InlineData("open", true)] + [InlineData("action", true)] + [InlineData("close", false)] + public async Task Should_Handle_Structured_Scenario_Canvas_Error(string operation, bool surfacesToCaller) + { + var handler = new RecordingCanvasHandler { StructuredErrorOperation = operation }; + await using var session = await CreateSessionAsync(CreateSessionConfig(handler)); + var canvas = Assert.Single((await WaitForCanvasRegistryAsync(session)).Canvases); + const string instanceId = "scenario-inspector-error"; + var input = new Dictionary { ["value"] = "before" }; + + if (operation != "open") + { + await session.Rpc.Canvas.OpenAsync( + canvasId: "scenario-inspector", + instanceId, + extensionId: canvas.ExtensionId, + input); + } + + Task InvokeAsync() => operation switch + { + "open" => session.Rpc.Canvas.OpenAsync( + canvasId: "scenario-inspector", + instanceId, + extensionId: canvas.ExtensionId, + input), + "action" => session.Rpc.Canvas.Action.InvokeAsync( + instanceId, + actionName: "replace", + input: new Dictionary { ["value"] = "after" }), + "close" => session.Rpc.Canvas.CloseAsync(instanceId), + _ => throw new ArgumentOutOfRangeException(nameof(operation)), + }; + + IOException? exception = null; + if (surfacesToCaller) + { + exception = await Assert.ThrowsAsync(InvokeAsync); + } + else + { + await InvokeAsync(); + } + + var expectedCode = $"scenario_canvas_{operation}_failed"; + var expectedMessage = $"The scenario canvas {operation} operation failed."; + Assert.Equal(expectedCode, handler.ThrownError?.Code); + Assert.Equal(expectedMessage, handler.ThrownError?.Message); + if (surfacesToCaller) + { + Assert.Contains(expectedMessage, exception!.Message, StringComparison.Ordinal); + } + Assert.Equal( + operation switch + { + "open" => [$"open:{instanceId}"], + "action" => [$"open:{instanceId}", $"action:{instanceId}:replace"], + "close" => [$"open:{instanceId}", $"close:{instanceId}"], + _ => throw new ArgumentOutOfRangeException(nameof(operation)), + }, + handler.Callbacks); + } + + [Fact] + public async Task Should_Reattach_Scenario_Canvas_And_Route_All_Callbacks_After_Resume() + { + var originalHandler = new RecordingCanvasHandler(); + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync(client1, CreateSessionConfig(originalHandler)); + var sessionId = session1.SessionId; + var response = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_CANVAS_READY.", + }); + Assert.Equal("SCENARIO_CANVAS_READY", response?.Data.Content); + var canvas = Assert.Single((await WaitForCanvasRegistryAsync(session1)).Canvases); + await session1.Rpc.Canvas.OpenAsync( + canvasId: "scenario-inspector", + instanceId: "scenario-inspector-resume", + extensionId: canvas.ExtensionId, + input: new Dictionary { ["value"] = "persisted" }); + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(session1.OpenCanvases.Count == 1), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the pre-resume canvas snapshot."); + var snapshot = session1.OpenCanvases.ToList(); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.ForceStopAsync(); + + var resumedHandler = new RecordingCanvasHandler(); + var client2 = Ctx.CreateClient(); + await using var session2 = await Ctx.ResumeSessionAsync( + client2, + sessionId, + CreateResumeConfig(resumedHandler, snapshot)); + + await resumedHandler.Opened.Task.WaitAsync(EventTimeout); + AssertRequest(resumedHandler.OpenRequests.Single(), sessionId, "scenario-inspector-resume"); + Assert.Equal("persisted", resumedHandler.OpenRequests[0].Input!.Value.GetProperty("value").GetString()); + AssertOpenCanvas( + await WaitForOpenCanvasAsync(session2, "scenario-inspector-resume"), + "scenario-inspector-resume", + "persisted"); + + var result = await session2.Rpc.Canvas.Action.InvokeAsync( + instanceId: "scenario-inspector-resume", + actionName: "replace", + input: new Dictionary { ["value"] = "resumed" }); + Assert.Equal("resumed", result.Result!.Value.GetProperty("value").GetString()); + + await session2.Rpc.Canvas.CloseAsync("scenario-inspector-resume"); + Assert.Equal( + ["open:scenario-inspector-resume", "action:scenario-inspector-resume:replace", "close:scenario-inspector-resume"], + resumedHandler.Callbacks); + Assert.Empty((await session2.Rpc.Canvas.ListOpenAsync()).OpenCanvases); + } + + private static SessionConfig CreateSessionConfig(RecordingCanvasHandler handler) => new() + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + RequestCanvasRenderer = true, + CanvasProvider = CreateProvider(), + Canvases = CreateCanvases(), + CanvasHandler = handler, + }; + + private static ResumeSessionConfig CreateResumeConfig( + RecordingCanvasHandler handler, + IList openCanvases) => new() + { + Streaming = true, + ContinuePendingWork = false, + OnPermissionRequest = PermissionHandler.ApproveAll, + RequestCanvasRenderer = true, + CanvasProvider = CreateProvider(), + Canvases = CreateCanvases(), + CanvasHandler = handler, + OpenCanvases = openCanvases, + }; + + private static CanvasProviderIdentity CreateProvider() => new() + { + Id = "scenario:builtin:e2e-window", + Name = "scenario client E2E", + }; + + private static IList CreateCanvases() + { + using var inputSchema = JsonDocument.Parse( + """{"type":"object","properties":{"value":{"type":"string"}},"required":["value"]}"""); + return + [ + new CanvasDeclaration + { + Id = "scenario-inspector", + DisplayName = "Scenario Inspector", + Description = "Displays scenario-owned state.", + InputSchema = inputSchema.RootElement.Clone(), + Actions = + [ + new CanvasAction + { + Name = "replace", + Description = "Replaces the displayed value.", + InputSchema = inputSchema.RootElement.Clone(), + }, + ], + }, + ]; + } + + private static async Task WaitForCanvasRegistryAsync(CopilotSession session) + { + CanvasList? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + result = await session.Rpc.Canvas.ListAsync(); + return result.Canvases.Count == 1; + }, + timeout: EventTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: "Timed out waiting for the scenario canvas registry."); + return result!; + } + + private static async Task WaitForOpenCanvasAsync( + CopilotSession session, + string instanceId) + { + OpenCanvasInstance? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + result = (await session.Rpc.Canvas.ListOpenAsync()).OpenCanvases + .SingleOrDefault(canvas => canvas.InstanceId == instanceId); + return result is not null; + }, + timeout: EventTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: $"Timed out waiting for open scenario canvas '{instanceId}'."); + return result!; + } + + private static void AssertRequest(CanvasProviderOpenRequest request, string sessionId, string instanceId) + { + Assert.Equal(sessionId, request.SessionId); + Assert.Equal("scenario:builtin:e2e-window", request.ExtensionId); + Assert.Equal("scenario-inspector", request.CanvasId); + Assert.Equal(instanceId, request.InstanceId); + Assert.Null(request.Host); + } + + private static void AssertRequest(CanvasProviderInvokeActionRequest request, string sessionId, string instanceId) + { + Assert.Equal(sessionId, request.SessionId); + Assert.Equal("scenario:builtin:e2e-window", request.ExtensionId); + Assert.Equal("scenario-inspector", request.CanvasId); + Assert.Equal(instanceId, request.InstanceId); + Assert.Null(request.Host); + } + + private static void AssertRequest(CanvasProviderCloseRequest request, string sessionId, string instanceId) + { + Assert.Equal(sessionId, request.SessionId); + Assert.Equal("scenario:builtin:e2e-window", request.ExtensionId); + Assert.Equal("scenario-inspector", request.CanvasId); + Assert.Equal(instanceId, request.InstanceId); + Assert.Null(request.Host); + } + + private static void AssertOpenCanvas( + OpenCanvasInstance canvas, + string expectedInstanceId, + string expectedInput) + { + Assert.Equal("scenario-inspector", canvas.CanvasId); + Assert.Equal("scenario:builtin:e2e-window", canvas.ExtensionId); + Assert.Equal("scenario client E2E", canvas.ExtensionName); + Assert.Equal(expectedInstanceId, canvas.InstanceId); + Assert.Equal(expectedInput, canvas.Input!.Value.GetProperty("value").GetString()); + Assert.Equal("ready", canvas.Status); + Assert.Equal($"Scenario Inspector: {expectedInput}", canvas.Title); + Assert.StartsWith("https://example.test/scenario-inspector/", canvas.Url, StringComparison.Ordinal); + } + + private sealed class RecordingCanvasHandler : CanvasHandlerBase + { + public string? StructuredErrorOperation { get; init; } + public CanvasException? ThrownError { get; private set; } + public List Callbacks { get; } = []; + public List OpenRequests { get; } = []; + public List ActionRequests { get; } = []; + public List CloseRequests { get; } = []; + public TaskCompletionSource Opened { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public override Task OnOpenAsync( + CanvasProviderOpenRequest request, + CancellationToken cancellationToken) + { + OpenRequests.Add(request); + Callbacks.Add($"open:{request.InstanceId}"); + Opened.TrySetResult(); + ThrowStructuredError("open"); + var value = request.Input!.Value.GetProperty("value").GetString(); + return Task.FromResult(new CanvasProviderOpenResult + { + Status = "ready", + Title = $"Scenario Inspector: {value}", + Url = $"https://example.test/scenario-inspector/{request.InstanceId}", + }); + } + + public override Task OnActionAsync( + CanvasProviderInvokeActionRequest request, + CancellationToken cancellationToken) + { + ActionRequests.Add(request); + Callbacks.Add($"action:{request.InstanceId}:{request.ActionName}"); + ThrowStructuredError("action"); + + return Task.FromResult(request.Input!.Value.Clone()); + } + + public override Task OnCloseAsync( + CanvasProviderCloseRequest request, + CancellationToken cancellationToken) + { + CloseRequests.Add(request); + Callbacks.Add($"close:{request.InstanceId}"); + ThrowStructuredError("close"); + return Task.CompletedTask; + } + + private void ThrowStructuredError(string operation) + { + if (StructuredErrorOperation != operation) + { + return; + } + + ThrownError = new CanvasException( + $"scenario_canvas_{operation}_failed", + $"The scenario canvas {operation} operation failed."); + throw ThrownError; + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingCloudE2ETests.cs b/dotnet/test/E2E/ScenarioTestingCloudE2ETests.cs new file mode 100644 index 0000000000..65d3a51cb2 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingCloudE2ETests.cs @@ -0,0 +1,310 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +#pragma warning disable GHCP001 + +public class ScenarioTestingCloudE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_cloud", output) +{ + private static readonly TimeSpan TestTimeout = TimeSpan.FromSeconds(30); + + [Fact] + public async Task Should_Notify_Steerability_Then_Send_First_Message_Without_Remote_Enable() + { + await using var session = await CreateSessionAsync(); + + await session.Rpc.Remote.NotifySteerableChangedAsync(true); + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_STEERABLE_FIRST_SEND.", + }); + + Assert.Contains("SCENARIO_STEERABLE_FIRST_SEND", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var events = await session.GetEventsAsync(); + var remoteIndex = -1; + var messageIndex = -1; + for (var i = 0; i < events.Count; i++) + { + if (remoteIndex < 0 && events[i] is SessionRemoteSteerableChangedEvent { Data.RemoteSteerable: true }) + { + remoteIndex = i; + } + + if (messageIndex < 0 + && events[i] is UserMessageEvent user + && user.Data.TransformedContent?.Contains("SCENARIO_STEERABLE_FIRST_SEND", StringComparison.Ordinal) == true) + { + messageIndex = i; + } + } + + Assert.True(remoteIndex >= 0, "Expected the persisted steerability notification."); + Assert.True(messageIndex > remoteIndex, "Expected steerability to be persisted before the first send."); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Route_First_Cloud_Event_For_Server_Assigned_Session_Id() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + var firstEvent = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "cloud-assigned-event"]), + UseLoggedInUser = false, + }); + + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + Cloud = new CloudSessionOptions + { + Repository = new CloudSessionRepository + { + Owner = "github", + Name = "copilot-sdk", + Branch = "main", + }, + }, + OnEvent = evt => firstEvent.TrySetResult(evt), + }); + + Assert.Equal("server-assigned-cloud-session", session.SessionId); + var started = Assert.IsType( + await firstEvent.Task.WaitAsync(TestTimeout)); + Assert.Equal(session.SessionId, started.Data.SessionId); + + var create = Assert.Single( + await ScenarioTestingTestCli.ReadRequestsAsync(capturePath), + request => request.GetProperty("method").GetString() == "session.create") + .GetProperty("params"); + Assert.False(create.TryGetProperty("sessionId", out _)); + Assert.Equal("github", create.GetProperty("cloud").GetProperty("repository").GetProperty("owner").GetString()); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Resume_Using_Runtime_Id_Returned_By_Cloud_Connect() + { + var (cliPath, capturePath) = await CreateFakeCloudRuntimeAsync(); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--resource-id", "github/copilot-sdk#123"]), + UseLoggedInUser = false, + }); + await client.StartAsync(); + + var connection = await client.Rpc.Sessions.ConnectAsync("cloud-control-session"); + Assert.Equal("runtime-session-id", connection.SessionId); + Assert.Equal("runtime-session-id", connection.Metadata.SessionId); + Assert.Equal("github/copilot-sdk#123", connection.Metadata.ResourceId); + + await using var resumed = await Ctx.ResumeSessionAsync(client, connection.SessionId); + Assert.Equal(connection.SessionId, resumed.SessionId); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => GetRequests(root, "session.resume").Count == 1); + var connectRequest = Assert.Single(GetRequests(capture.RootElement, "sessions.connect")) + .GetProperty("params"); + var resumeRequest = Assert.Single(GetRequests(capture.RootElement, "session.resume")) + .GetProperty("params"); + Assert.Equal("cloud-control-session", connectRequest.GetProperty("sessionId").GetString()); + Assert.Equal(connection.SessionId, resumeRequest.GetProperty("sessionId").GetString()); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Expose_Cloud_Resource_Mismatch_Before_Resume() + { + const string ExpectedResourceId = "github/copilot-sdk#123"; + var (cliPath, capturePath) = await CreateFakeCloudRuntimeAsync(); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--resource-id", "github/other-repository#456"]), + UseLoggedInUser = false, + }); + await client.StartAsync(); + + var connection = await client.Rpc.Sessions.ConnectAsync("cloud-control-session"); + Assert.NotEqual(ExpectedResourceId, connection.Metadata.ResourceId); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => GetRequests(root, "sessions.connect").Count == 1); + Assert.Empty(GetRequests(capture.RootElement, "session.resume")); + } + + private async Task<(string CliPath, string CapturePath)> CreateFakeCloudRuntimeAsync() + { + var cliPath = Path.Join(Ctx.WorkDir, $"scenario-client-cloud-{Guid.NewGuid():N}.js"); + var capturePath = Path.Join(Ctx.WorkDir, $"scenario-client-cloud-{Guid.NewGuid():N}.json"); + await File.WriteAllTextAsync(cliPath, FakeCloudRuntimeScript); + return (cliPath, capturePath); + } + + private static List GetRequests(JsonElement root, string method) => + root.GetProperty("requests") + .EnumerateArray() + .Where(item => item.GetProperty("method").GetString() == method) + .ToList(); + + private static async Task WaitForCaptureAsync( + string path, + Func predicate) + { + JsonDocument? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + try + { + using var stream = new FileStream( + path, + FileMode.Open, + FileAccess.Read, + FileShare.ReadWrite | FileShare.Delete); + using var reader = new StreamReader(stream); + var document = JsonDocument.Parse(await reader.ReadToEndAsync()); + if (!predicate(document.RootElement)) + { + document.Dispose(); + return false; + } + + result = document; + return true; + } + catch (Exception ex) when (ex is IOException or JsonException) + { + return false; + } + }, + timeout: TestTimeout, + pollInterval: TimeSpan.FromMilliseconds(50), + timeoutMessage: $"Timed out waiting for fake cloud runtime capture at {path}."); + return result!; + } + + private const string FakeCloudRuntimeScript = """ + const fs = require("fs"); + + function argument(name) { + const index = process.argv.indexOf(name); + return index >= 0 ? process.argv[index + 1] : undefined; + } + + const captureFile = argument("--capture-file"); + const resourceId = argument("--resource-id"); + const requests = []; + let buffer = Buffer.alloc(0); + + function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ requests })); + } + + function write(message) { + const body = JSON.stringify(message); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function respond(id, result) { + write({ jsonrpc: "2.0", id, result }); + } + + function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "method")) return; + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + respond(message.id, { ok: true, protocolVersion: 3, version: "fake" }); + return; + } + + if (message.method === "sessions.connect") { + respond(message.id, { + sessionId: "runtime-session-id", + metadata: { + kind: "coding-agent", + modifiedTime: "2026-09-17T20:00:00Z", + name: "Cloud task", + repository: { + branch: "main", + name: "copilot-sdk", + owner: "github" + }, + resourceId, + sessionId: "runtime-session-id", + startTime: "2026-09-17T19:00:00Z", + state: "active" + } + }); + return; + } + + if (message.method === "session.resume") { + respond(message.id, { + sessionId: message.params.sessionId, + workspacePath: null, + capabilities: null + }); + return; + } + + if (message.method === "session.detach") { + respond(message.id, { success: true }); + return; + } + + if (message.method === "runtime.shutdown") { + respond(message.id, {}); + setTimeout(() => process.exit(0), 10); + return; + } + + respond(message.id, { success: true }); + } + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } + }); + + process.stdin.resume(); + saveCapture(); + setInterval(() => {}, 1000); + """; +} + +#pragma warning restore GHCP001 diff --git a/dotnet/test/E2E/ScenarioTestingCompositionE2ETests.cs b/dotnet/test/E2E/ScenarioTestingCompositionE2ETests.cs new file mode 100644 index 0000000000..b88f6225d0 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingCompositionE2ETests.cs @@ -0,0 +1,648 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using System.Collections.Concurrent; +using System.ComponentModel; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +/// +/// End-to-end coverage for representative representative SDK workflows. +/// These tests intentionally compose APIs that are otherwise covered individually. +/// +public class ScenarioTestingCompositionE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_composition", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Send_Scenario_Message_With_Metadata_And_Extension_Context() + { + using var payload = JsonDocument.Parse("""{"selection":"TRACE_SENTINEL","line":42}"""); + await using var session = await CreateSessionAsync(new SessionConfig { Streaming = true }); + var idle = TestHelper.GetNextEventOfTypeAsync(session, EventTimeout); + + var messageId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly TRACE_SENTINEL from the attached extension context.", + DisplayPrompt = "Analyze the selected trace entry", + Mode = "enqueue", + AgentMode = AgentMode.Interactive, + Source = MessageSource.Agent("trace-viewer"), + Attachments = + [ + new AttachmentExtensionContext + { + CapturedAt = DateTimeOffset.Parse("2026-09-17T20:00:00Z"), + ExtensionId = "scenario-client:trace-viewer", + CanvasId = "trace", + InstanceId = "trace-1", + Title = "Selected trace entry", + Payload = payload.RootElement.Clone(), + }, + ], + }); + + await idle; + + var events = await session.GetEventsAsync(); + var userMessage = Assert.Single( + events.OfType(), + evt => string.Equals(evt.Data.MessageId, messageId, StringComparison.Ordinal)); + Assert.Equal("Analyze the selected trace entry", userMessage.Data.Content); + Assert.Equal(UserMessageDelivery.Idle, userMessage.Data.Delivery); + Assert.Equal(UserMessageAgentMode.Interactive, userMessage.Data.AgentMode); + Assert.Equal("agent-trace-viewer", userMessage.Data.Source); + Assert.Contains("TRACE_SENTINEL", userMessage.Data.TransformedContent ?? string.Empty, StringComparison.Ordinal); + + var attachment = Assert.IsType(Assert.Single(userMessage.Data.Attachments!)); + Assert.Equal("scenario-client:trace-viewer", attachment.ExtensionId); + Assert.Equal("trace", attachment.CanvasId); + Assert.Equal("trace-1", attachment.InstanceId); + Assert.Equal("Selected trace entry", attachment.Title); + Assert.Equal("TRACE_SENTINEL", attachment.Payload!.Value.GetProperty("selection").GetString()); + + var assistantMessage = events.OfType().Last(); + Assert.Contains("TRACE_SENTINEL", assistantMessage.Data.Content ?? string.Empty, StringComparison.Ordinal); + } + + [Fact] + public async Task Should_Classify_Queued_And_Immediate_Scenario_Messages_While_Busy() + { + var toolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + Tools = [AIFunctionFactory.Create(WaitForReleaseAsync, "wait_for_scenario_release")], + }); + var userMessages = new List(); + var userMessagesLock = new object(); + using var subscription = session.On(message => + { + lock (userMessagesLock) + { + userMessages.Add(message); + } + }); + + try + { + await session.SendAsync(new MessageOptions + { + Prompt = "Call wait_for_scenario_release, then reply with its result.", + }); + await toolStarted.Task.WaitAsync(EventTimeout); + + var queuedMessageId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with QUEUED_SCENARIO_MESSAGE after the active turn.", + DisplayPrompt = "Queued follow-up", + Mode = "enqueue", + Source = MessageSource.System, + }); + var steeringMessageId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with STEERING_SCENARIO_MESSAGE instead.", + DisplayPrompt = "Immediate steering update", + Mode = "immediate", + Source = MessageSource.Agent("session-coordinator"), + }); + + var finalQueuedResponse = TestHelper.GetNextEventOfTypeAsync( + session, + message => message.Data.Content?.Contains("QUEUED_SCENARIO_MESSAGE", StringComparison.Ordinal) == true, + EventTimeout, + "the queued scenario response"); + releaseTool.TrySetResult("ACTIVE_TURN_RELEASED"); + + await TestHelper.WaitForConditionAsync( + () => + { + lock (userMessagesLock) + { + return Task.FromResult( + userMessages.Any(evt => evt.Data.MessageId == queuedMessageId) && + userMessages.Any(evt => evt.Data.MessageId == steeringMessageId)); + } + }, + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for queued and steering messages to be consumed."); + await finalQueuedResponse; + + List observedMessages; + lock (userMessagesLock) + { + observedMessages = [.. userMessages]; + } + + var queued = Assert.Single(observedMessages, evt => evt.Data.MessageId == queuedMessageId); + Assert.Equal("Queued follow-up", queued.Data.Content); + Assert.Equal(UserMessageDelivery.Queued, queued.Data.Delivery); + Assert.Equal("system", queued.Data.Source); + + var steering = Assert.Single(observedMessages, evt => evt.Data.MessageId == steeringMessageId); + Assert.Equal("Immediate steering update", steering.Data.Content); + Assert.Equal(UserMessageDelivery.Steering, steering.Data.Delivery); + Assert.Equal("agent-session-coordinator", steering.Data.Source); + } + finally + { + releaseTool.TrySetResult("RELEASED_AFTER_TEST"); + } + + [Description("Waits until the scenario releases the active turn")] + async Task WaitForReleaseAsync(CancellationToken cancellationToken) + { + toolStarted.TrySetResult(); + return await releaseTool.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + } + + [Fact] + public async Task Should_Resume_With_Reattached_Scenario_Host_State() + { + var originalCanvasHandler = new ScenarioCanvasHandler(); + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync( + client1, + CreateScenarioSessionConfig(originalCanvasHandler, includeTool: false, includeMcp: true)); + var sessionId = session1.SessionId; + await WaitForMcpServerStatusAsync(session1, "scenario-resume-mcp", McpServerStatus.Connected); + var initialResponse = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Remember SCENARIO_RESUME_MARKER and reply with exactly INITIALIZED.", + }); + Assert.Contains("INITIALIZED", initialResponse?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var canvas = Assert.Single((await session1.Rpc.Canvas.ListAsync()).Canvases); + await session1.Rpc.Canvas.OpenAsync( + canvasId: "scenario-counter", + instanceId: "scenario-counter-1", + extensionId: canvas.ExtensionId, + input: new Dictionary { ["start"] = 40 }); + await session1.LogAsync("SCENARIO_HOST_STATE_MARKER"); + + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(session1.OpenCanvases.Count == 1), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the open canvas snapshot."); + var openCanvases = session1.OpenCanvases.ToList(); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.StopAsync(); + + var resumedCanvasHandler = new ScenarioCanvasHandler(); + var client2 = Ctx.CreateClient(); + await using var session2 = await Ctx.ResumeSessionAsync( + client2, + sessionId, + CreateScenarioResumeConfig(resumedCanvasHandler, openCanvases, includeMcp: true)); + + var restoredOpenRequest = await resumedCanvasHandler.Opened.Task.WaitAsync(EventTimeout); + Assert.Equal("scenario-counter-1", restoredOpenRequest.InstanceId); + Assert.Equal(40, restoredOpenRequest.Input!.Value.GetProperty("start").GetInt32()); + + var restoredCanvas = await WaitForOpenCanvasAsync(session2, "scenario-counter-1"); + Assert.Equal("scenario-counter-1", restoredCanvas.InstanceId); + Assert.Equal("scenario-counter", restoredCanvas.CanvasId); + Assert.Equal(40, restoredCanvas.Input!.Value.GetProperty("start").GetInt32()); + + var action = await session2.Rpc.Canvas.Action.InvokeAsync( + instanceId: "scenario-counter-1", + actionName: "increment", + input: new Dictionary { ["delta"] = 2 }); + Assert.Equal(42, action.Result!.Value.GetProperty("count").GetInt32()); + Assert.Single(resumedCanvasHandler.ActionRequests); + + await WaitForMcpServerStatusAsync(session2, "scenario-resume-mcp", McpServerStatus.Connected); + var mcpTools = await session2.Rpc.Mcp.ListToolsAsync("scenario-resume-mcp"); + Assert.NotEmpty(mcpTools.Tools); + + var response = await session2.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_host_lookup with key ALPHA, then reply with exactly its result.", + }); + Assert.Contains("SCENARIO_HOST_VALUE_ALPHA", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var events = await session2.GetEventsAsync(); + Assert.Contains(events.OfType(), evt => evt.Data.Message == "SCENARIO_HOST_STATE_MARKER"); + Assert.Single(events.OfType()); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Resume_With_Reattached_Scenario_Provider() + { + var initialProviderTokenRequest = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var initialRequestHandler = new RecordingRequestHandler(); + var client1 = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio(), + RequestHandler = initialRequestHandler, + }); + var createConfig = new SessionConfig + { + Model = "scenario-resume-provider/scenario-model", + OnPermissionRequest = PermissionHandler.ApproveAll, + }; + ConfigureScenarioProvider( + createConfig, + args => + { + initialProviderTokenRequest.TrySetResult(args); + return Task.FromResult("initial-scenario-provider-token"); + }); + var session1 = await Ctx.CreateSessionAsync(client1, createConfig); + var sessionId = session1.SessionId; + var initialResponse = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Create persisted history before the provider resume.", + }); + Assert.Contains( + RecordingRequestHandler.SyntheticText, + initialResponse?.Data.Content ?? string.Empty, + StringComparison.Ordinal); + + var initialProviderRequest = await initialProviderTokenRequest.Task.WaitAsync(EventTimeout); + Assert.Equal(sessionId, initialProviderRequest.SessionId); + Assert.Equal("scenario-resume-provider", initialProviderRequest.ProviderName); + Assert.Contains( + initialRequestHandler.InferenceRequests, + request => request.Url.StartsWith("https://scenario-resume.invalid/", StringComparison.Ordinal) + && request.SessionId == sessionId); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.StopAsync(); + + var providerTokenRequest = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var resumedRequestHandler = new RecordingRequestHandler(); + var client2 = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio(), + RequestHandler = resumedRequestHandler, + }); + var resumeConfig = new ResumeSessionConfig + { + Model = "scenario-resume-provider/scenario-model", + OnPermissionRequest = PermissionHandler.ApproveAll, + }; + ConfigureScenarioProvider( + resumeConfig, + args => + { + providerTokenRequest.TrySetResult(args); + return Task.FromResult("resumed-scenario-provider-token"); + }); + await using var session2 = await Ctx.ResumeSessionAsync(client2, sessionId, resumeConfig); + + var response = await session2.SendAndWaitAsync(new MessageOptions + { + Prompt = "Use the reattached scenario provider.", + }); + Assert.Contains( + RecordingRequestHandler.SyntheticText, + response?.Data.Content ?? string.Empty, + StringComparison.Ordinal); + + var providerRequest = await providerTokenRequest.Task.WaitAsync(EventTimeout); + Assert.Equal(sessionId, providerRequest.SessionId); + Assert.Equal("scenario-resume-provider", providerRequest.ProviderName); + Assert.Contains( + resumedRequestHandler.InferenceRequests, + request => request.Url.StartsWith("https://scenario-resume.invalid/", StringComparison.Ordinal) + && request.SessionId == sessionId); + } + + [Fact] + public async Task Should_Retry_Resume_On_Replacement_Client_After_Recoverable_Setup_Failure() + { + var originalHandler = new ScenarioCanvasHandler(); + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync(client1, CreateScenarioSessionConfig(originalHandler)); + var sessionId = session1.SessionId; + var initialResponse = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_RETRY_RESUME_READY.", + }); + Assert.Contains("SCENARIO_RETRY_RESUME_READY", initialResponse?.Data.Content ?? string.Empty, StringComparison.Ordinal); + var canvas = Assert.Single((await session1.Rpc.Canvas.ListAsync()).Canvases); + await session1.Rpc.Canvas.OpenAsync( + canvasId: "scenario-counter", + instanceId: "scenario-retry-canvas", + extensionId: canvas.ExtensionId, + input: new Dictionary { ["start"] = 40 }); + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(session1.OpenCanvases.Count == 1), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the retry canvas snapshot."); + var openCanvases = session1.OpenCanvases.ToList(); + await session1.LogAsync("SCENARIO_RETRY_RESUME_HISTORY"); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.StopAsync(); + + var failingClient = Ctx.CreateClient(); + var failingConfig = CreateScenarioResumeConfig(new ScenarioCanvasHandler(), openCanvases); + failingConfig.Tools = + [ + AIFunctionFactory.Create(() => "first", "duplicate_scenario_tool"), + AIFunctionFactory.Create(() => "second", "duplicate_scenario_tool"), + ]; + await Assert.ThrowsAnyAsync(() => + Ctx.ResumeSessionAsync(failingClient, sessionId, failingConfig)); + await failingClient.ForceStopAsync(); + + var replacementHandler = new ScenarioCanvasHandler(); + var replacementClient = Ctx.CreateClient(); + await using var resumed = await Ctx.ResumeSessionAsync( + replacementClient, + sessionId, + CreateScenarioResumeConfig(replacementHandler, openCanvases)); + + var reopened = await replacementHandler.Opened.Task.WaitAsync(EventTimeout); + Assert.Equal("scenario-retry-canvas", reopened.InstanceId); + await TestHelper.WaitForConditionAsync( + async () => (await resumed.Rpc.Canvas.ListOpenAsync()).OpenCanvases.Count == 1, + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the replacement client to restore the open canvas."); + Assert.Single((await resumed.Rpc.Canvas.ListOpenAsync()).OpenCanvases); + } + + [Fact] + public async Task Should_Not_Emit_Redundant_Model_Change_When_Resuming_Same_Model() + { + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync(client1, new SessionConfig + { + Model = "claude-sonnet-5", + }); + var sessionId = session1.SessionId; + Assert.Equal("claude-sonnet-5", (await session1.Rpc.Model.GetCurrentAsync()).ModelId); + var initialResponse = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_SAME_MODEL_HISTORY_READY.", + }); + Assert.Contains("SCENARIO_SAME_MODEL_HISTORY_READY", initialResponse?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.StopAsync(); + + var earlyEvents = new ConcurrentQueue(); + var client2 = Ctx.CreateClient(); + await using var resumed = await Ctx.ResumeSessionAsync(client2, sessionId, new ResumeSessionConfig + { + Model = "claude-sonnet-5", + OnEvent = earlyEvents.Enqueue, + }); + + Assert.Equal("claude-sonnet-5", (await resumed.Rpc.Model.GetCurrentAsync()).ModelId); + var persistedEvents = await resumed.GetEventsAsync(); + Assert.DoesNotContain(earlyEvents, evt => evt is SessionModelChangeEvent); + Assert.DoesNotContain(persistedEvents, evt => evt is SessionModelChangeEvent); + } + + [Fact] + public async Task Should_Read_Persisted_Scenario_Events_Without_Resuming() + { + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync(client1); + var sessionId = session1.SessionId; + var response = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_PERSISTED_HISTORY.", + }); + Assert.Contains("SCENARIO_PERSISTED_HISTORY", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.StopAsync(); + + await using var client2 = Ctx.CreateClient(); + await client2.StartAsync(); + var persisted = await client2.Rpc.Sessions.ReadPersistedEventsAsync(sessionId); + + Assert.Equal(EventsCursorStatus.Ok, persisted.CursorStatus); + Assert.False(persisted.HasMore); + Assert.Contains( + persisted.Events.OfType(), + evt => evt.Data.TransformedContent?.Contains("SCENARIO_PERSISTED_HISTORY", StringComparison.Ordinal) == true); + Assert.Contains( + persisted.Events.OfType(), + evt => evt.Data.Content?.Contains("SCENARIO_PERSISTED_HISTORY", StringComparison.Ordinal) == true); + } + + [Fact] + public async Task Should_Propagate_Canvas_Handler_Error() + { + var handler = new ScenarioCanvasHandler { ThrowOnAction = true }; + await using var session = await CreateSessionAsync(CreateScenarioSessionConfig(handler)); + var canvas = Assert.Single((await session.Rpc.Canvas.ListAsync()).Canvases); + await session.Rpc.Canvas.OpenAsync( + canvasId: "scenario-counter", + instanceId: "scenario-counter-error", + extensionId: canvas.ExtensionId, + input: new Dictionary { ["start"] = 0 }); + + var exception = await Assert.ThrowsAnyAsync(() => + session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "scenario-counter-error", + actionName: "increment", + input: new Dictionary { ["delta"] = 1 })); + + Assert.Contains("The scenario canvas could not increment.", exception.ToString(), StringComparison.Ordinal); + } + + private static SessionConfig CreateScenarioSessionConfig( + ScenarioCanvasHandler canvasHandler, + bool includeTool = true, + bool includeMcp = false) + { + var config = new SessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + RequestCanvasRenderer = true, + CanvasProvider = new CanvasProviderIdentity + { + Id = "scenario:builtin:test-window", + Name = "scenario client", + }, + Canvases = + [ + new CanvasDeclaration + { + Id = "scenario-counter", + DisplayName = "Scenario Counter", + Description = "Represents a scenario-hosted canvas.", + Actions = + [ + new CanvasAction + { + Name = "increment", + Description = "Increments the counter.", + }, + ], + }, + ], + CanvasHandler = canvasHandler, + }; + if (includeTool) + { + config.Tools = [AIFunctionFactory.Create(ScenarioHostLookup, "scenario_host_lookup")]; + } + if (includeMcp) + { + config.McpServers = CreateTestMcpServers("scenario-resume-mcp"); + } + + return config; + } + + private static ResumeSessionConfig CreateScenarioResumeConfig( + ScenarioCanvasHandler canvasHandler, + IList openCanvases, + bool includeMcp = false) + { + var config = new ResumeSessionConfig + { + Streaming = true, + ContinuePendingWork = false, + Tools = [AIFunctionFactory.Create(ScenarioHostLookup, "scenario_host_lookup")], + OnPermissionRequest = PermissionHandler.ApproveAll, + RequestCanvasRenderer = true, + CanvasProvider = new CanvasProviderIdentity + { + Id = "scenario:builtin:test-window", + Name = "scenario client", + }, + Canvases = + [ + new CanvasDeclaration + { + Id = "scenario-counter", + DisplayName = "Scenario Counter", + Description = "Represents a scenario-hosted canvas.", + Actions = + [ + new CanvasAction + { + Name = "increment", + Description = "Increments the counter.", + }, + ], + }, + ], + CanvasHandler = canvasHandler, + OpenCanvases = openCanvases, + }; + if (includeMcp) + { + config.McpServers = CreateTestMcpServers("scenario-resume-mcp"); + } + return config; + } + + private static void ConfigureScenarioProvider( + SessionConfigBase config, + Func>? providerTokenProvider) + { + if (providerTokenProvider is null) + { + return; + } + + config.Providers = + [ + new NamedProviderConfig + { + Name = "scenario-resume-provider", + Type = "openai", + WireApi = "responses", + BaseUrl = "https://scenario-resume.invalid/v1", + BearerTokenProvider = providerTokenProvider, + }, + ]; + config.Models = + [ + new ProviderModelConfig + { + Provider = "scenario-resume-provider", + Id = "scenario-model", + WireModel = "scenario-wire-model", + }, + ]; + } + + private static async Task WaitForOpenCanvasAsync( + CopilotSession session, + string instanceId) + { + OpenCanvasInstance? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + result = (await session.Rpc.Canvas.ListOpenAsync()).OpenCanvases + .SingleOrDefault(canvas => canvas.InstanceId == instanceId); + return result is not null; + }, + timeout: EventTimeout, + timeoutMessage: $"Timed out waiting for open scenario canvas '{instanceId}'."); + return result!; + } + + [Description("Looks up scenario-owned host state")] + private static string ScenarioHostLookup([Description("Lookup key")] string key) => + $"SCENARIO_HOST_VALUE_{key.ToUpperInvariant()}"; + + private sealed class ScenarioCanvasHandler : CanvasHandlerBase + { + public bool ThrowOnAction { get; init; } + public TaskCompletionSource Opened { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + public List ActionRequests { get; } = []; + + public override Task OnOpenAsync( + CanvasProviderOpenRequest request, + CancellationToken cancellationToken) + { + Opened.TrySetResult(request); + return Task.FromResult(new CanvasProviderOpenResult + { + Status = "ready", + Title = "Scenario Counter", + Url = $"https://example.test/canvas/{request.InstanceId}", + }); + } + + public override Task OnActionAsync( + CanvasProviderInvokeActionRequest request, + CancellationToken cancellationToken) + { + if (ThrowOnAction) + { + throw new CanvasException( + "scenario_canvas_action_failed", + "The scenario canvas could not increment."); + } + + ActionRequests.Add(request); + var delta = request.Input!.Value.GetProperty("delta").GetInt32(); + using var result = JsonDocument.Parse($$"""{"count":{{40 + delta}}}"""); + return Task.FromResult(result.RootElement.Clone()); + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingControlStateE2ETests.cs b/dotnet/test/E2E/ScenarioTestingControlStateE2ETests.cs new file mode 100644 index 0000000000..051ea26bbe --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingControlStateE2ETests.cs @@ -0,0 +1,104 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using System.ComponentModel; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingControlStateE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_control_state", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Compose_Mode_Name_Plan_Client_Metadata_And_Objective_State() + { + await using var session = await CreateSessionAsync(); + const string sessionName = "Scenario control state"; + const string plan = "# Scenario plan\n- Verify control state"; + const string objective = """{"objective":"VERIFY_SCENARIO_CONTROL","status":"active"}"""; + + await session.Rpc.Mode.SetAsync(SessionMode.Plan); + await session.Rpc.Name.SetAsync(sessionName); + await session.Rpc.Plan.UpdateAsync(plan); + var metadata = await session.Rpc.Metadata.UpdateClientMetadataAsync( + set: new Dictionary + { + ["scenario-client/control-mode"] = "plan", + ["scenario-client/objective"] = "VERIFY_SCENARIO_CONTROL", + }); + var objectiveWrite = await session.Rpc.Workspaces.WriteAutopilotObjectiveAsync(objective); + + Assert.Equal("create", objectiveWrite.Operation); + Assert.True((await session.Rpc.Workspaces.AutopilotObjectiveExistsAsync()).Exists); + Assert.Equal(objective, (await session.Rpc.Workspaces.ReadAutopilotObjectiveAsync()).Content); + Assert.Equal(plan, (await session.Rpc.Plan.ReadAsync()).Content); + Assert.Equal(sessionName, (await session.Rpc.Name.GetAsync()).Name); + Assert.Equal("VERIFY_SCENARIO_CONTROL", metadata["scenario-client/objective"]); + + var snapshot = await session.Rpc.Metadata.SnapshotAsync(); + Assert.Equal(session.SessionId, snapshot.SessionId); + Assert.Equal(MetadataSnapshotCurrentMode.Plan, snapshot.CurrentMode); + Assert.Null(snapshot.InitialName); + + var deleted = await session.Rpc.Workspaces.DeleteAutopilotObjectiveAsync(); + Assert.True(deleted.Deleted); + Assert.False((await session.Rpc.Workspaces.AutopilotObjectiveExistsAsync()).Exists); + } + + [Fact] + public async Task Should_Report_Processing_While_Scenario_Tool_Is_Running() + { + var toolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + Tools = [AIFunctionFactory.Create(WaitForScenarioAsync, "wait_for_scenario_control")], + }); + + Assert.False((await session.Rpc.Metadata.IsProcessingAsync()).Processing); + + try + { + var idle = TestHelper.GetNextEventOfTypeAsync(session, EventTimeout); + await session.SendAsync(new MessageOptions + { + Prompt = "Call wait_for_scenario_control, then reply with exactly SCENARIO_CONTROL_DONE.", + }); + await toolStarted.Task.WaitAsync(EventTimeout); + + Assert.True((await session.Rpc.Metadata.IsProcessingAsync()).Processing); + var activity = await session.Rpc.Metadata.ActivityAsync(); + Assert.True(activity.HasActiveWork); + Assert.True(activity.Abortable); + + releaseTool.TrySetResult("SCENARIO_CONTROL_DONE"); + await idle; + + await TestHelper.WaitForConditionAsync( + async () => !(await session.Rpc.Metadata.IsProcessingAsync()).Processing, + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for processing metadata to return to idle."); + + Assert.False((await session.Rpc.Metadata.ActivityAsync()).HasActiveWork); + } + finally + { + releaseTool.TrySetResult("SCENARIO_CONTROL_DONE"); + } + + [Description("Waits for the scenario controller to release the active turn")] + async Task WaitForScenarioAsync(CancellationToken cancellationToken) + { + toolStarted.TrySetResult(); + return await releaseTool.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingE2ETestBase.cs b/dotnet/test/E2E/ScenarioTestingE2ETestBase.cs new file mode 100644 index 0000000000..c7175e4e2d --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingE2ETestBase.cs @@ -0,0 +1,14 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public abstract class ScenarioTestingE2ETestBase( + E2ETestFixture fixture, + string snapshotCategory, + ITestOutputHelper output) + : E2ETestBase(fixture, snapshotCategory, output, replayOnly: true); diff --git a/dotnet/test/E2E/ScenarioTestingEmptyRuntimeE2ETests.cs b/dotnet/test/E2E/ScenarioTestingEmptyRuntimeE2ETests.cs new file mode 100644 index 0000000000..81867618e0 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingEmptyRuntimeE2ETests.cs @@ -0,0 +1,40 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingEmptyRuntimeE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_empty_runtime", output) +{ + [Fact] + public async Task Empty_Mode_Minimal_Toolless_Session_Has_No_Tools() + { + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Mode = CopilotClientMode.Empty, + BaseDirectory = Ctx.HomeDir, + }); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + AvailableTools = new ToolSet(), + OnPermissionRequest = PermissionHandler.ApproveAll, + SystemMessage = new SystemMessageConfig + { + Mode = SystemMessageMode.Replace, + Content = "Reply to every request with exactly EMPTY_SCENARIO_READY.", + }, + }); + + var response = await session.SendAndWaitAsync(new MessageOptions { Prompt = "Start." }); + Assert.Contains("EMPTY_SCENARIO_READY", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var exchanges = await Ctx.GetExchangesAsync(); + Assert.Empty(GetToolNames(exchanges[^1])); + Assert.DoesNotContain("Current working directory:", GetSystemMessage(exchanges[^1]), StringComparison.OrdinalIgnoreCase); + } +} diff --git a/dotnet/test/E2E/ScenarioTestingEventSubscriptionsE2ETests.cs b/dotnet/test/E2E/ScenarioTestingEventSubscriptionsE2ETests.cs new file mode 100644 index 0000000000..c6e4649ac3 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingEventSubscriptionsE2ETests.cs @@ -0,0 +1,170 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using Microsoft.Extensions.AI; +using GitHub.Copilot.Test.Harness; +using System.ComponentModel; +using System.Reflection; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingEventSubscriptionsE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_event_subscriptions", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Deliver_Mixed_Scenario_Event_Stream_In_Order_After_Handler_Lag() + { + var handlerEntered = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var toolInvoked = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var releaseHandler = new ManualResetEventSlim(); + var events = new List(); + + await using var session = await CreateSessionAsync(new SessionConfig + { + Streaming = true, + Tools = [AIFunctionFactory.Create(ScenarioLookup, "scenario_event_lookup")], + }); + + using var subscription = session.On(evt => + { + if (evt is UserMessageEvent) + { + handlerEntered.TrySetResult(); + releaseHandler.Wait(EventTimeout); + } + + lock (events) + { + events.Add(evt); + } + }); + + var send = session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_event_lookup with key 'ordered', then reply with exactly its result.", + DisplayPrompt = "Run ordered scenario lookup", + Source = MessageSource.Agent("scenario-client"), + }, timeout: TimeSpan.FromSeconds(120)); + + await handlerEntered.Task.WaitAsync(EventTimeout); + try + { + await toolInvoked.Task.WaitAsync(EventTimeout); + } + finally + { + releaseHandler.Set(); + } + + var response = await send; + Assert.Contains("SCENARIO_EVENT_ORDERED", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + List types; + lock (events) + { + types = events.Select(evt => evt.Type).ToList(); + } + + var user = types.IndexOf("user.message"); + var toolStart = types.IndexOf("tool.execution_start"); + var toolComplete = types.IndexOf("tool.execution_complete"); + var assistant = types.LastIndexOf("assistant.message"); + var idle = types.LastIndexOf("session.idle"); + Assert.True(user < toolStart, string.Join(", ", types)); + Assert.True(toolStart < toolComplete, string.Join(", ", types)); + Assert.True(toolComplete < assistant, string.Join(", ", types)); + Assert.True(assistant < idle, string.Join(", ", types)); + + [Description("Looks up scenario-owned event data")] + string ScenarioLookup([Description("Lookup key")] string key) + { + toolInvoked.TrySetResult(); + return $"SCENARIO_EVENT_{key.ToUpperInvariant()}"; + } + } + + [Fact] + public async Task Should_Stop_Closed_And_Replaced_Scenario_Event_Sources() + { + const string connectionToken = "scenario-client-events-token"; + await using var server = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForTcp(connectionToken: connectionToken), + }); + await server.StartAsync(); + var cliUrl = $"localhost:{server.RuntimePort}"; + + var oldEventCount = 0; + string sessionId; + await using (var firstClient = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(cliUrl, connectionToken: connectionToken), + })) + { + var firstSession = await Ctx.CreateSessionAsync(firstClient, new SessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + sessionId = firstSession.SessionId; + using var oldSubscription = firstSession.On(_ => Interlocked.Increment(ref oldEventCount)); + + await firstSession.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_EVENT_SOURCE_ONE.", + Source = MessageSource.Agent("scenario-client"), + }); + await firstSession.Rpc.SuspendAsync(); + await firstSession.DisposeAsync(); + await firstClient.ForceStopAsync(); + + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(IsEventChannelClosed(firstSession)), + timeout: EventTimeout, + timeoutMessage: "Timed out waiting for the old scenario event source to close."); + } + + var countAfterClose = Volatile.Read(ref oldEventCount); + await using var secondClient = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(cliUrl, connectionToken: connectionToken), + }); + await using var secondSession = await Ctx.ResumeSessionAsync(secondClient, sessionId, new ResumeSessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + var newEvents = new List(); + using var newSubscription = secondSession.On(newEvents.Add); + var newInfo = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var infoSubscription = secondSession.On(evt => + { + if (evt.Data.Message == "SCENARIO_EVENT_SOURCE_TWO") + { + newInfo.TrySetResult(evt); + } + }); + await secondSession.LogAsync("SCENARIO_EVENT_SOURCE_TWO"); + await newInfo.Task.WaitAsync(EventTimeout); + + Assert.Equal(countAfterClose, Volatile.Read(ref oldEventCount)); + Assert.Contains(newEvents, evt => evt is SessionInfoEvent info && info.Data.Message == "SCENARIO_EVENT_SOURCE_TWO"); + } + + private static bool IsEventChannelClosed(CopilotSession session) + { + var eventChannelField = typeof(CopilotSession).GetField( + "_eventChannel", + BindingFlags.Instance | BindingFlags.NonPublic) + ?? throw new InvalidOperationException("CopilotSession._eventChannel was not found."); + var channel = eventChannelField.GetValue(session)!; + var reader = channel.GetType().GetProperty("Reader")!.GetValue(channel)!; + return ((Task)reader.GetType().GetProperty("Completion")!.GetValue(reader)!).IsCompleted; + } +} diff --git a/dotnet/test/E2E/ScenarioTestingJsExtensionBridgeE2ETests.cs b/dotnet/test/E2E/ScenarioTestingJsExtensionBridgeE2ETests.cs new file mode 100644 index 0000000000..b714f32444 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingJsExtensionBridgeE2ETests.cs @@ -0,0 +1,575 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Diagnostics; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; +using RpcExtension = GitHub.Copilot.Rpc.Extension; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingJsExtensionBridgeE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_js_extension_bridge", output) +{ + private static readonly TimeSpan ExtensionTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Run_Standalone_Empty_Mode_Extension_Canvas_Without_Conversation() + { + var fixture = await CreateExtensionFixtureAsync(); + await using var client = CreateExtensionClient(fixture, CopilotClientMode.Empty); + var config = CreateSessionConfig(fixture.ProjectDirectory); + config.AvailableTools = new ToolSet(); + config.EnableSessionStore = false; + config.SkipCustomInstructions = true; + + await using var session = await Ctx.CreateSessionAsync(client, config); + + var extension = await WaitForExtensionAsync(session, fixture.ExtensionId); + var canvas = await WaitForCanvasAsync(session, fixture.ExtensionId); + Assert.Equal(ExtensionSource.Project, extension.Source); + Assert.Equal(ExtensionStatus.Running, extension.Status); + Assert.Equal("js-scenario-canvas", canvas.CanvasId); + Assert.NotNull(extension.Pid); + + await session.Rpc.Plugins.ReloadAsync(new SessionPluginsReloadRequest + { + ReloadExtensions = false, + }); + var extensionAfterPluginReload = await WaitForExtensionAsync(session, fixture.ExtensionId); + Assert.Equal(extension.Pid, extensionAfterPluginReload.Pid); + Assert.Equal( + "js-scenario-canvas", + (await WaitForCanvasAsync(session, fixture.ExtensionId)).CanvasId); + + await WaitForTraceAsync(fixture.TraceFile, "joined"); + var opened = await OpenCanvasWhenRegisteredAsync( + session, + canvasId: "js-scenario-canvas", + instanceId: "standalone-canvas-1", + extensionId: fixture.ExtensionId, + input: new Dictionary { ["value"] = "standalone" }); + Assert.Equal("ready", opened.Status); + + var action = await session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "standalone-canvas-1", + actionName: "set-value", + input: new Dictionary { ["value"] = "updated" }); + Assert.Equal("updated", action.Result!.Value.GetProperty("value").GetString()); + + await session.Rpc.Canvas.CloseAsync("standalone-canvas-1"); + await WaitForTraceAsync(fixture.TraceFile, "close"); + + var trace = ReadTrace(fixture.TraceFile); + var joined = trace.Where(entry => GetKind(entry) == "joined").ToList(); + Assert.NotEmpty(joined); + Assert.All(joined, entry => + { + Assert.Equal(session.SessionId, entry.GetProperty("sessionId").GetString()); + Assert.Equal( + Path.GetFullPath(fixture.ProjectDirectory), + Path.GetFullPath(entry.GetProperty("workingDirectory").GetString()!)); + }); + Assert.Equal( + ["open", "action", "close"], + trace.Where(entry => GetKind(entry) != "joined").Select(GetKind)); + Assert.DoesNotContain(trace, entry => GetKind(entry) == "sent"); + Assert.Empty(await Ctx.GetExchangesAsync()); + } + + [Fact] + public async Task Should_Persist_Server_Extension_Enablement_For_Future_Sessions() + { + var fixture = await CreateExtensionFixtureAsync(ExtensionSource.User); + await using var client = CreateExtensionClient(fixture); + await using var activeSession = await Ctx.CreateSessionAsync( + client, + CreateSessionConfig(fixture.ProjectDirectory)); + + var active = await WaitForExtensionAsync(activeSession, fixture.ExtensionId); + Assert.Equal(ExtensionStatus.Running, active.Status); + + await client.Rpc.User.Settings.ReloadAsync(); + var discovered = await client.Rpc.Extensions.DiscoverAsync(); + var discoveredExtension = Assert.Single( + discovered.Extensions, + extension => extension.Id == fixture.ExtensionId); + Assert.True(discoveredExtension.Enabled); + Assert.Equal(DiscoveredExtensionSource.User, discoveredExtension.Source); + Assert.Empty((await client.Rpc.Plugins.ListAsync()).Plugins); + + await client.Rpc.Extensions.DisableAsync([fixture.ExtensionId]); + Assert.Equal( + ExtensionStatus.Running, + (await WaitForExtensionAsync(activeSession, fixture.ExtensionId)).Status); + + await using var disabledSession = await Ctx.CreateSessionAsync( + client, + CreateSessionConfig(fixture.ProjectDirectory)); + var disabled = await WaitForExtensionAsync( + disabledSession, + fixture.ExtensionId, + ExtensionStatus.Disabled); + Assert.Null(disabled.Pid); + + await client.Rpc.Extensions.EnableAsync([fixture.ExtensionId]); + Assert.Equal( + ExtensionStatus.Disabled, + (await WaitForExtensionAsync( + disabledSession, + fixture.ExtensionId, + ExtensionStatus.Disabled)).Status); + + await using var enabledSession = await Ctx.CreateSessionAsync( + client, + CreateSessionConfig(fixture.ProjectDirectory)); + var enabled = await WaitForExtensionAsync(enabledSession, fixture.ExtensionId); + Assert.Equal(ExtensionStatus.Running, enabled.Status); + Assert.NotNull(enabled.Pid); + } + + [Fact] + public async Task Should_Bridge_Js_Extension_Canvas_Context_Log_And_Session_Continuation() + { + var fixture = await CreateExtensionFixtureAsync(); + await using var client = CreateExtensionClient(fixture); + await using var session = await Ctx.CreateSessionAsync(client, CreateSessionConfig(fixture.ProjectDirectory)); + + var extension = await WaitForExtensionAsync(session, fixture.ExtensionId); + var canvas = await WaitForCanvasAsync(session, fixture.ExtensionId); + Assert.Equal(ExtensionStatus.Running, extension.Status); + Assert.Equal("js-scenario-canvas", canvas.CanvasId); + Assert.Equal("JavaScript Scenario Canvas", canvas.DisplayName); + Assert.Equal("object", canvas.InputSchema!.Value.GetProperty("type").GetString()); + Assert.Equal(["set-value", "continue", "fail"], canvas.Actions!.Select(action => action.Name)); + + await WaitForTraceAsync(fixture.TraceFile, "joined"); + var joined = ReadTrace(fixture.TraceFile) + .Where(entry => GetKind(entry) == "joined") + .ToList(); + Assert.NotEmpty(joined); + Assert.All(joined, entry => + { + Assert.Equal(session.SessionId, entry.GetProperty("sessionId").GetString()); + Assert.Equal( + Path.GetFullPath(fixture.ProjectDirectory), + Path.GetFullPath(entry.GetProperty("workingDirectory").GetString()!)); + }); + var workspacePath = joined[^1].GetProperty("workspacePath").GetString(); + Assert.False(string.IsNullOrWhiteSpace(workspacePath)); + Assert.False(string.IsNullOrEmpty(Path.GetPathRoot(workspacePath))); + + var metadata = await session.Rpc.Metadata.SnapshotAsync(); + Assert.True( + PathsEqual(fixture.ProjectDirectory, metadata.WorkingDirectory), + $"Expected working directory '{fixture.ProjectDirectory}', actual '{metadata.WorkingDirectory}'."); + + var opened = await OpenCanvasWhenRegisteredAsync( + session, + canvasId: "js-scenario-canvas", + instanceId: "js-scenario-canvas-1", + extensionId: fixture.ExtensionId, + input: new Dictionary { ["value"] = "before" }); + Assert.Equal("ready", opened.Status); + Assert.Equal("JavaScript Scenario Canvas: before", opened.Title); + + var continuation = await session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "js-scenario-canvas-1", + actionName: "continue", + input: new Dictionary()); + Assert.False(string.IsNullOrWhiteSpace( + continuation.Result!.Value.GetProperty("messageId").GetString())); + await WaitForTraceAsync(fixture.TraceFile, "sent"); + await TestHelper.WaitForConditionAsync( + async () => + { + var events = await session.GetEventsAsync(); + return events.OfType().Any(evt => evt.Data.Message == "JS_EXTENSION_LOG") + && events.OfType().Any( + evt => (evt.Data.Content ?? string.Empty).Contains( + "JS_EXTENSION_CONTINUATION", + StringComparison.Ordinal)); + }, + timeout: ExtensionTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: "Timed out waiting for the extension log and continuation."); + + var action = await session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "js-scenario-canvas-1", + actionName: "set-value", + input: new Dictionary { ["value"] = "after" }); + Assert.Equal("after", action.Result!.Value.GetProperty("value").GetString()); + + await session.Rpc.Canvas.CloseAsync("js-scenario-canvas-1"); + await WaitForTraceAsync(fixture.TraceFile, "close"); + + var trace = ReadTrace(fixture.TraceFile); + var open = Assert.Single(trace, entry => GetKind(entry) == "open"); + AssertBridgeContext(open, session.SessionId, fixture.ExtensionId, "js-scenario-canvas-1"); + Assert.Equal("before", open.GetProperty("input").GetProperty("value").GetString()); + Assert.False(open.TryGetProperty("host", out _)); + + var actions = trace.Where(entry => GetKind(entry) == "action").ToList(); + Assert.Equal(["continue", "set-value"], actions.Select(entry => entry.GetProperty("actionName").GetString())); + Assert.All(actions, entry => + AssertBridgeContext(entry, session.SessionId, fixture.ExtensionId, "js-scenario-canvas-1")); + Assert.Equal("after", actions[1].GetProperty("input").GetProperty("value").GetString()); + + var close = Assert.Single(trace, entry => GetKind(entry) == "close"); + AssertBridgeContext(close, session.SessionId, fixture.ExtensionId, "js-scenario-canvas-1"); + Assert.Equal( + ["open", "action", "sent", "action", "close"], + trace.Where(entry => GetKind(entry) != "joined").Select(GetKind)); + } + + [Fact] + public async Task Should_Surface_Structured_CanvasError_From_Js_Extension() + { + var fixture = await CreateExtensionFixtureAsync(); + await using var client = CreateExtensionClient(fixture); + await using var session = await Ctx.CreateSessionAsync(client, CreateSessionConfig(fixture.ProjectDirectory)); + + await WaitForExtensionAsync(session, fixture.ExtensionId); + await WaitForCanvasAsync(session, fixture.ExtensionId); + await OpenCanvasWhenRegisteredAsync( + session, + canvasId: "js-scenario-canvas", + instanceId: "js-scenario-canvas-error", + extensionId: fixture.ExtensionId, + input: new Dictionary { ["value"] = "before" }); + + var exception = await Assert.ThrowsAsync(() => + session.Rpc.Canvas.Action.InvokeAsync( + instanceId: "js-scenario-canvas-error", + actionName: "fail", + input: new Dictionary())); + + Assert.Contains("The JavaScript canvas action failed.", exception.Message, StringComparison.Ordinal); + await WaitForTraceAsync(fixture.TraceFile, "error"); + var error = Assert.Single(ReadTrace(fixture.TraceFile), entry => GetKind(entry) == "error"); + Assert.Equal("js_canvas_failed", error.GetProperty("code").GetString()); + Assert.Equal("The JavaScript canvas action failed.", error.GetProperty("message").GetString()); + } + + private CopilotClient CreateExtensionClient( + ExtensionFixture fixture, + CopilotClientMode mode = CopilotClientMode.CopilotCli) + { + var environment = Ctx.GetEnvironment(); + environment["COPILOT_CLI_ENABLED_FEATURE_FLAGS"] = "EXTENSIONS"; + environment["SCENARIO_EXTENSION_TRACE_FILE"] = fixture.TraceFile; + environment["SCENARIO_EXTENSION_WORKING_DIRECTORY"] = fixture.ProjectDirectory; + + return Ctx.CreateClient( + options: new CopilotClientOptions + { + Mode = mode, + BaseDirectory = mode == CopilotClientMode.Empty ? Ctx.HomeDir : null, + Connection = RuntimeConnection.ForStdio( + path: Ctx.GetLegacyCliPath(), + args: ["--yolo"]), + }, + environment: environment); + } + + private static SessionConfig CreateSessionConfig(string workingDirectory) => new() + { + EnableConfigDiscovery = true, + RequestExtensions = true, + WorkingDirectory = workingDirectory, + OnPermissionRequest = PermissionHandler.ApproveAll, + }; + + private async Task CreateExtensionFixtureAsync( + ExtensionSource source = default) + { + var extensionName = $"js-scenario-bridge-{Guid.NewGuid():N}"; + var projectDirectory = Path.Join(Ctx.WorkDir, $"js-extension-project-{Guid.NewGuid():N}"); + source = source == default ? ExtensionSource.Project : source; + var extensionDirectory = source == ExtensionSource.User + ? Path.Join(Ctx.HomeDir, "extensions", extensionName) + : Path.Join(projectDirectory, ".github", "extensions", extensionName); + var traceFile = Path.Join(Ctx.WorkDir, $"{extensionName}.jsonl"); + Directory.CreateDirectory(projectDirectory); + Directory.CreateDirectory(extensionDirectory); + await InitializeGitRepositoryAsync(projectDirectory); + File.WriteAllText(Path.Join(extensionDirectory, "extension.mjs"), ExtensionScript); + return new ExtensionFixture( + projectDirectory, + traceFile, + $"{source.Value}:{extensionName}"); + } + + private static async Task WaitForExtensionAsync( + CopilotSession session, + string extensionId, + ExtensionStatus expectedStatus = default) + { + expectedStatus = expectedStatus == default ? ExtensionStatus.Running : expectedStatus; + RpcExtension? extension = null; + await TestHelper.WaitForConditionAsync( + async () => + { + var list = await session.Rpc.Extensions.ListAsync(); + extension = list.Extensions.FirstOrDefault( + item => string.Equals(item.Id, extensionId, StringComparison.Ordinal)); + return extension?.Status == expectedStatus; + }, + timeout: ExtensionTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: $"Timed out waiting for extension '{extensionId}'.", + transientExceptionFilter: ex => + ex.ToString().Contains("Extensions not available", StringComparison.OrdinalIgnoreCase)); + return extension!; + } + + private static async Task WaitForCanvasAsync(CopilotSession session, string extensionId) + { + DiscoveredCanvas? canvas = null; + await TestHelper.WaitForConditionAsync( + async () => + { + var list = await session.Rpc.Canvas.ListAsync(); + canvas = list.Canvases.FirstOrDefault( + item => string.Equals(item.ExtensionId, extensionId, StringComparison.Ordinal) + && string.Equals(item.CanvasId, "js-scenario-canvas", StringComparison.Ordinal)); + return canvas is not null; + }, + timeout: ExtensionTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: $"Timed out waiting for canvas from extension '{extensionId}'."); + return canvas!; + } + + private static async Task OpenCanvasWhenRegisteredAsync( + CopilotSession session, + string canvasId, + string instanceId, + string extensionId, + object input) + { + OpenCanvasInstance? opened = null; + await TestHelper.WaitForConditionAsync( + async () => + { + try + { + opened = await session.Rpc.Canvas.OpenAsync( + canvasId, + instanceId, + extensionId, + input); + return true; + } + catch (IOException exception) + when (exception.Message.Contains("is not registered", StringComparison.OrdinalIgnoreCase)) + { + return false; + } + }, + timeout: ExtensionTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: $"Timed out waiting for canvas '{extensionId}/{canvasId}' to become invokable."); + return opened!; + } + + private static async Task WaitForTraceAsync(string traceFile, string kind) + { + await TestHelper.WaitForConditionAsync( + () => Task.FromResult( + File.Exists(traceFile) + && ReadTrace(traceFile).Any(entry => GetKind(entry) == kind)), + timeout: ExtensionTimeout, + pollInterval: TimeSpan.FromMilliseconds(100), + timeoutMessage: $"Timed out waiting for extension trace entry '{kind}'."); + } + + private static List ReadTrace(string traceFile) + { + for (var attempt = 0; ; attempt++) + { + if (!File.Exists(traceFile)) + { + return []; + } + + try + { + return File.ReadAllLines(traceFile) + .Where(line => !string.IsNullOrWhiteSpace(line)) + .Select(line => + { + using var document = JsonDocument.Parse(line); + return document.RootElement.Clone(); + }) + .ToList(); + } + catch (Exception exception) + when (attempt < 9 && exception is IOException or JsonException) + { + Thread.Sleep(20); + } + } + } + + private static string GetKind(JsonElement entry) => entry.GetProperty("kind").GetString()!; + + private static void AssertBridgeContext( + JsonElement entry, + string sessionId, + string extensionId, + string instanceId) + { + Assert.Equal(sessionId, entry.GetProperty("sessionId").GetString()); + Assert.Equal(extensionId, entry.GetProperty("extensionId").GetString()); + Assert.Equal("js-scenario-canvas", entry.GetProperty("canvasId").GetString()); + Assert.Equal(instanceId, entry.GetProperty("instanceId").GetString()); + } + + private static bool PathsEqual(string expected, string? actual) => + actual is not null + && string.Equals( + Path.GetFullPath(expected).TrimEnd(Path.DirectorySeparatorChar), + Path.GetFullPath(actual).TrimEnd(Path.DirectorySeparatorChar), + OperatingSystem.IsWindows() ? StringComparison.OrdinalIgnoreCase : StringComparison.Ordinal); + + private static async Task InitializeGitRepositoryAsync(string projectDirectory) + { + var startInfo = new ProcessStartInfo("git") + { + WorkingDirectory = projectDirectory, + Arguments = "init --quiet", + RedirectStandardOutput = true, + RedirectStandardError = true, + UseShellExecute = false, + }; + + // .NET Framework drops inherited environment variables with empty values. + // Remove the indexed Git config group so GIT_CONFIG_COUNT cannot reference + // a value that disappeared while ProcessStartInfo copied the environment. + foreach (var name in startInfo.Environment.Keys + .Where(name => + name.Equals("GIT_CONFIG_COUNT", StringComparison.OrdinalIgnoreCase) + || name.StartsWith("GIT_CONFIG_KEY_", StringComparison.OrdinalIgnoreCase) + || name.StartsWith("GIT_CONFIG_VALUE_", StringComparison.OrdinalIgnoreCase)) + .ToArray()) + { + startInfo.Environment.Remove(name); + } + + using var process = new Process + { + StartInfo = startInfo, + }; + + if (!process.Start()) + { + throw new InvalidOperationException("Failed to start git init."); + } + + await process.WaitForExitAsync(); + if (process.ExitCode != 0) + { + throw new InvalidOperationException( + $"git init failed with exit code {process.ExitCode}: {await process.StandardError.ReadToEndAsync()}"); + } + } + + private sealed record ExtensionFixture( + string ProjectDirectory, + string TraceFile, + string ExtensionId); + + private const string ExtensionScript = """ + import { appendFileSync } from "node:fs"; + import { CanvasError, createCanvas, joinSession } from "@github/copilot-sdk/extension"; + + const traceFile = process.env.SCENARIO_EXTENSION_TRACE_FILE; + const workingDirectory = process.env.SCENARIO_EXTENSION_WORKING_DIRECTORY; + + function record(kind, data = {}) { + appendFileSync(traceFile, `${JSON.stringify({ kind, ...data })}\n`); + } + + let session; + const canvas = createCanvas({ + id: "js-scenario-canvas", + displayName: "JavaScript Scenario Canvas", + description: "Exercises the JavaScript extension bridge.", + inputSchema: { + type: "object", + properties: { value: { type: "string" } }, + required: ["value"] + }, + actions: [ + { + name: "set-value", + description: "Sets the displayed value.", + inputSchema: { + type: "object", + properties: { value: { type: "string" } }, + required: ["value"] + }, + handler: context => { + record("action", context); + return { value: context.input.value }; + } + }, + { + name: "continue", + description: "Continues the host session.", + handler: async context => { + record("action", context); + const messageId = await session.send( + "Reply with exactly JS_EXTENSION_CONTINUATION." + ); + record("sent", { messageId }); + return { messageId }; + } + }, + { + name: "fail", + description: "Throws a structured CanvasError.", + handler: context => { + record("action", context); + const error = new CanvasError( + "js_canvas_failed", + "The JavaScript canvas action failed." + ); + record("error", { code: error.code, message: error.message }); + throw error; + } + } + ], + open: context => { + record("open", context); + return { + status: "ready", + title: `JavaScript Scenario Canvas: ${context.input.value}`, + url: `https://example.test/js-scenario-canvas/${context.instanceId}` + }; + }, + onClose: context => record("close", context) + }); + + session = await joinSession({ + workingDirectory, + tools: [], + canvases: [canvas] + }); + + record("joined", { + sessionId: session.sessionId, + workspacePath: session.workspacePath ?? null, + workingDirectory, + cwd: process.cwd() + }); + await session.log("JS_EXTENSION_LOG"); + + setInterval(() => {}, 60_000).unref?.(); + """; +} diff --git a/dotnet/test/E2E/ScenarioTestingLifecycleRecoveryE2ETests.cs b/dotnet/test/E2E/ScenarioTestingLifecycleRecoveryE2ETests.cs new file mode 100644 index 0000000000..1feb4ab4e8 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingLifecycleRecoveryE2ETests.cs @@ -0,0 +1,178 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using Microsoft.Extensions.AI; +using System.ComponentModel; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingLifecycleRecoveryE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_lifecycle_recovery", output) +{ + private static readonly TimeSpan LifecycleTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Abort_Active_Scenario_Turn_And_Remain_Usable() + { + var toolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + Streaming = true, + Tools = [AIFunctionFactory.Create(BlockingLookup, "scenario_blocking_lookup")], + }); + + _ = session.SendAsync(new MessageOptions + { + Prompt = "Call scenario_blocking_lookup with key 'abort', then reply with the result.", + DisplayPrompt = "Run cancellable scenario lookup", + Source = MessageSource.Agent("scenario-client"), + }); + + Assert.Equal("abort", await toolStarted.Task.WaitAsync(LifecycleTimeout)); + await session.AbortAsync(); + releaseTool.TrySetResult("SCENARIO_ABORTED_TOOL_RESULT"); + + var recovery = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var recoverySubscription = session.On(message => + { + if (message.Data.Content?.Contains("SCENARIO_ABORT_RECOVERY_OK", StringComparison.Ordinal) == true) + { + recovery.TrySetResult(message); + } + }); + await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_ABORT_RECOVERY_OK.", + DisplayPrompt = "Verify scenario session recovery", + Source = MessageSource.Agent("scenario-client"), + }); + Assert.Contains( + "SCENARIO_ABORT_RECOVERY_OK", + (await recovery.Task.WaitAsync(LifecycleTimeout)).Data.Content ?? string.Empty, + StringComparison.Ordinal); + + [Description("Blocks a scenario-owned lookup until released")] + async Task BlockingLookup( + [Description("Lookup key")] string key, + CancellationToken cancellationToken) + { + toolStarted.TrySetResult(key); + return await releaseTool.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + } + + [Fact] + public async Task Should_Suspend_Disconnect_And_Resume_Scenario_State_Without_Delete() + { + const string connectionToken = "scenario-client-lifecycle-token"; + await using var server = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForTcp(connectionToken: connectionToken), + }); + await server.StartAsync(); + var cliUrl = $"localhost:{server.RuntimePort}"; + + string sessionId; + await using (var firstClient = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(cliUrl, connectionToken: connectionToken), + })) + { + var firstSession = await Ctx.CreateSessionAsync(firstClient, new SessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + sessionId = firstSession.SessionId; + var initialized = await firstSession.SendAndWaitAsync(new MessageOptions + { + Prompt = "Remember SCENARIO_LIFECYCLE_MEMORY and reply with exactly SCENARIO_LIFECYCLE_INITIALIZED.", + Source = MessageSource.Agent("scenario-client"), + }); + Assert.Contains("SCENARIO_LIFECYCLE_INITIALIZED", initialized?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + await firstSession.Rpc.SuspendAsync(); + await firstSession.DisposeAsync(); + await firstClient.ForceStopAsync(); + } + + await using var secondClient = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(cliUrl, connectionToken: connectionToken), + }); + await using var resumed = await Ctx.ResumeSessionAsync(secondClient, sessionId, new ResumeSessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + var response = await resumed.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly the scenario lifecycle memory value from the earlier turn.", + Source = MessageSource.Agent("scenario-client"), + }); + Assert.Contains("SCENARIO_LIFECYCLE_MEMORY", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + Assert.Contains((await resumed.GetEventsAsync()).OfType(), _ => true); + } + + [Fact] + public async Task Should_Classify_Delete_Not_Found_For_Scenario_Cleanup() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "delete-not-found"]), + UseLoggedInUser = false, + }); + + const string missingId = "missing-scenario-client-session"; + var exception = await Assert.ThrowsAsync(() => client.DeleteSessionAsync(missingId)); + Assert.Equal( + $"Failed to delete session {missingId}: Session file not found", + exception.Message); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + var delete = Assert.Single(requests, request => request.GetProperty("method").GetString() == "session.delete"); + var parameters = delete.GetProperty("params"); + var request = parameters.ValueKind == JsonValueKind.Array ? parameters[0] : parameters; + Assert.Equal(missingId, request.GetProperty("sessionId").GetString()); + } + + [Fact] + public async Task Should_Allow_Caller_Retry_After_Preacceptance_Session_Not_Found() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "resume-not-found-once"]), + UseLoggedInUser = false, + }); + + const string sessionId = "scenario-client-retry-session"; + var first = await Assert.ThrowsAnyAsync(() => + Ctx.ResumeSessionAsync(client, sessionId, new ResumeSessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + })); + Assert.Contains("Session not found", first.ToString(), StringComparison.OrdinalIgnoreCase); + + await using var resumed = await Ctx.ResumeSessionAsync(client, sessionId, new ResumeSessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + Assert.Equal(sessionId, resumed.SessionId); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + Assert.Equal(2, requests.Count(request => request.GetProperty("method").GetString() == "session.resume")); + } +} diff --git a/dotnet/test/E2E/ScenarioTestingMcpE2ETests.cs b/dotnet/test/E2E/ScenarioTestingMcpE2ETests.cs new file mode 100644 index 0000000000..5b4d1a7543 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingMcpE2ETests.cs @@ -0,0 +1,489 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Diagnostics; +using System.Net.Http; +using System.Text.Json; +using System.Threading.Channels; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +/// +/// Representative scenario coverage for MCP lifecycle, OAuth, configuration, and MCP Apps. +/// +public class ScenarioTestingMcpE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_mcp", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + private const string ExpectedToken = "scenario-client-mcp-token"; + + [Fact] + public async Task Should_List_Reload_Restart_And_Report_Scenario_Mcp_State() + { + const string serverName = "scenario-client-lifecycle"; + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + McpServers = CreateTestMcpServers(serverName), + }); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + + var initial = await session.Rpc.Mcp.ListAsync(); + Assert.NotNull(initial.Host); + Assert.Empty(initial.Host!.FailedServers); + Assert.Empty(initial.Host.NeedsAuthServers); + Assert.Empty(initial.Host.PendingConnections); + Assert.Equal(McpServerStatus.Connected, Assert.Single(initial.Servers).Status); + + var tools = await session.Rpc.Mcp.ListToolsAsync(serverName); + Assert.Contains(tools.Tools, tool => tool.Name == "get_env"); + + var statusEvents = Channel.CreateUnbounded(); + using var subscription = session.On( + evt => statusEvents.Writer.TryWrite(evt)); + + await session.Rpc.Mcp.RestartServerAsync(serverName); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + await session.Rpc.Mcp.ReloadAsync(); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + + var connectedEvent = await ReadMatchingAsync( + statusEvents.Reader, + evt => evt.Data.ServerName == serverName && evt.Data.Status == McpServerStatus.Connected); + Assert.Equal(serverName, connectedEvent.Data.ServerName); + Assert.True((await session.Rpc.Mcp.IsServerRunningAsync(serverName)).Running); + } + + [Fact] + public async Task Should_Provide_First_Party_Scenario_Token_And_Cancel_Third_Party_Oauth() + { + await using var firstParty = await ScenarioOAuthMcpServer.StartAsync(ExpectedToken); + await using var thirdParty = await ScenarioOAuthMcpServer.StartAsync(ExpectedToken); + const string firstPartyName = "scenario-client-first-party"; + const string thirdPartyName = "scenario-client-third-party"; + var requests = Channel.CreateUnbounded(); + + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + OnMcpAuthRequest = request => + { + requests.Writer.TryWrite(request); + return Task.FromResult( + request.ServerName == firstPartyName + ? McpAuthResult.FromToken(new McpAuthToken + { + AccessToken = ExpectedToken, + TokenType = "Bearer", + ExpiresIn = 3600, + }) + : McpAuthResult.Cancel()); + }, + McpServers = new Dictionary + { + [firstPartyName] = new McpHttpServerConfig + { + Url = $"{firstParty.Url}/mcp", + Tools = ["*"], + }, + [thirdPartyName] = new McpHttpServerConfig + { + Url = $"{thirdParty.Url}/mcp", + Tools = ["*"], + }, + }, + }); + + await session.Rpc.Mcp.ReloadAsync(); + await WaitForMcpServerStatusAsync(session, firstPartyName, McpServerStatus.Connected); + await WaitForMcpServerStatusAsync(session, thirdPartyName, McpServerStatus.NeedsAuth); + + var observed = new List(); + while (observed.Select(request => request.ServerName).Distinct(StringComparer.Ordinal).Count() < 2) + { + observed.Add(await requests.Reader.ReadAsync().AsTask().WaitAsync(EventTimeout)); + } + + Assert.Contains(observed, request => + request.ServerName == firstPartyName && request.Reason == McpOauthRequestReason.Initial); + Assert.Contains(observed, request => + request.ServerName == thirdPartyName && request.Reason == McpOauthRequestReason.Initial); + + var firstPartyRequests = await firstParty.GetRequestsAsync(); + Assert.Contains(firstPartyRequests, request => request.Authorization == $"Bearer {ExpectedToken}"); + var state = await session.Rpc.Mcp.ListAsync(); + Assert.Contains(thirdPartyName, state.Host!.NeedsAuthServers.Keys); + } + + [Fact] + public async Task Should_Reconnect_With_Cached_Scenario_Token_Then_Return_Interactive_Oauth_Url() + { + await using var oauthServer = await ScenarioOAuthMcpServer.StartAsync(ExpectedToken); + const string serverName = "scenario-client-oauth-reconnect"; + var tokenRequests = 0; + + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + OnMcpAuthRequest = request => + { + Interlocked.Increment(ref tokenRequests); + return Task.FromResult(McpAuthResult.FromToken(new McpAuthToken + { + AccessToken = ExpectedToken, + TokenType = "Bearer", + ExpiresIn = 3600, + })); + }, + McpServers = new Dictionary + { + [serverName] = new McpHttpServerConfig + { + Url = $"{oauthServer.Url}/mcp", + Tools = ["*"], + }, + }, + }); + + await session.Rpc.Mcp.ReloadAsync(); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + var tokenRequestsAfterInitialConnect = Volatile.Read(ref tokenRequests); + Assert.True(tokenRequestsAfterInitialConnect >= 1); + var serverRequestsAfterInitialConnect = (await oauthServer.GetRequestsAsync()).Count; + + await session.Rpc.Mcp.RestartServerAsync(serverName); + await TestHelper.WaitForConditionAsync( + async () => (await oauthServer.GetRequestsAsync()).Count > serverRequestsAfterInitialConnect, + timeout: EventTimeout, + pollInterval: TimeSpan.FromMilliseconds(50), + timeoutMessage: "Timed out waiting for the MCP server to reconnect after restart."); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + Assert.Equal(tokenRequestsAfterInitialConnect, Volatile.Read(ref tokenRequests)); + Assert.Contains( + await oauthServer.GetRequestsAsync(), + request => request.Authorization == $"Bearer {ExpectedToken}"); + + var cached = await session.Rpc.Mcp.Oauth.ProbeAsync(serverName); + Assert.IsType(cached); + + var interactive = await session.Rpc.Mcp.Oauth.LoginAsync( + serverName, + forceReauth: true, + clientName: "scenario client", + callbackSuccessMessage: "Return to your application.", + clientId: "scenario-client-client", + publicClient: true); + Assert.NotNull(interactive.AuthorizationUrl); + Assert.StartsWith($"{oauthServer.Url}/authorize", interactive.AuthorizationUrl, StringComparison.Ordinal); + } + + [Fact] + public async Task Should_Manage_And_Discover_Scenario_Mcp_Config_Lifecycle() + { + var serverName = $"scenario-client-config-{Guid.NewGuid():N}"; + var testServer = Path.Join(FindTestHarnessDir(), "test-mcp-server.mjs"); + await Client.StartAsync(); + + try + { + await Client.Rpc.Mcp.Config.AddAsync(serverName, new McpStdioServerConfig + { + Command = "node", + Args = [testServer], + Tools = ["get_env"], + }); + + var afterAdd = await Client.Rpc.Mcp.Config.ListAsync(); + Assert.Contains(serverName, afterAdd.Servers.Keys); + var discovered = await Client.Rpc.Mcp.DiscoverAsync( + workingDirectory: Ctx.WorkDir, + includeEffectiveSource: true); + var enabled = Assert.Single(discovered.Servers, server => server.Name == serverName); + Assert.True(enabled.Enabled); + Assert.NotNull(enabled.EffectiveSource); + + await Client.Rpc.Mcp.Config.UpdateAsync(serverName, new McpStdioServerConfig + { + Command = "node", + Args = [testServer], + Env = new Dictionary { ["SCENARIO_CONFIG_VERSION"] = "2" }, + Tools = ["*"], + }); + var updated = GetServerConfig(await Client.Rpc.Mcp.Config.ListAsync(), serverName); + Assert.Equal("2", updated.GetProperty("env").GetProperty("SCENARIO_CONFIG_VERSION").GetString()); + + await Client.Rpc.Mcp.Config.DisableAsync([serverName]); + var disabled = await Client.Rpc.Mcp.DiscoverAsync(Ctx.WorkDir); + Assert.False(Assert.Single(disabled.Servers, server => server.Name == serverName).Enabled); + + await Client.Rpc.Mcp.Config.EnableAsync([serverName]); + var reenabled = await Client.Rpc.Mcp.DiscoverAsync(Ctx.WorkDir); + Assert.True(Assert.Single(reenabled.Servers, server => server.Name == serverName).Enabled); + } + finally + { + await Client.Rpc.Mcp.Config.RemoveAsync(serverName); + } + + Assert.DoesNotContain(serverName, (await Client.Rpc.Mcp.Config.ListAsync()).Servers.Keys); + } + + [Fact] + public async Task Should_List_Mcp_App_Visible_Tools_And_Read_Resource() + { + const string serverName = "scenario-mcp-app"; + const string resourceUri = "ui://scenario/app"; + var environment = Ctx.GetEnvironment(); + environment["COPILOT_MCP_APPS"] = "true"; + environment["MCP_APPS"] = "true"; + await using var client = Ctx.CreateClient(environment: environment); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + EnableMcpApps = true, + McpServers = new Dictionary + { + [serverName] = new McpStdioServerConfig + { + Command = "node", + Args = [Path.Join(FindTestHarnessDir(), "test-mcp-app-server.mjs")], + Tools = ["*"], + }, + }, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + + var tools = await session.Rpc.Mcp.Apps.ListToolsAsync( + serverName, + originServerName: serverName); + var tool = Assert.Single(tools.Tools); + Assert.Equal("app_visible", tool["name"].GetString()); + Assert.Equal( + ["model", "app"], + tool["_meta"].GetProperty("ui.visibility") + .EnumerateArray().Select(value => value.GetString())); + + using var value = JsonDocument.Parse("\"scenario-value\""); + var call = await session.Rpc.Mcp.Apps.CallToolAsync( + serverName, + "app_visible", + originServerName: serverName, + arguments: new Dictionary + { + ["value"] = value.RootElement.Clone(), + }); + Assert.Equal( + "APP_VISIBLE:scenario-value", + call["content"][0].GetProperty("text").GetString()); + + var resource = Assert.Single( + (await session.Rpc.Mcp.Apps.ReadResourceAsync(serverName, resourceUri)).Contents); + Assert.Equal(resourceUri, resource.Uri); + Assert.Equal("text/html", resource.MimeType); + Assert.Equal("SCENARIO_MCP_APP", resource.Text); + Assert.Equal( + "https://api.example.test", + Assert.Single( + resource.Meta!["ui.csp"].GetProperty("connectDomains").EnumerateArray()) + .GetString()); + } + + [Fact] + public async Task Should_Enforce_Mcp_App_Origin_Server() + { + const string serverName = "scenario-client-origin"; + const string otherServerName = "scenario-client-other-origin"; + var servers = CreateTestMcpServers(serverName, otherServerName); + ((McpStdioServerConfig)servers[serverName]).Env = + new Dictionary { ["SCENARIO_ORIGIN_VALUE"] = "origin-ok" }; + + var environment = Ctx.GetEnvironment(); + environment["COPILOT_MCP_APPS"] = "true"; + environment["MCP_APPS"] = "true"; + await using var client = Ctx.CreateClient(environment: environment); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + ClientName = "scenario-client", + EnableMcpApps = true, + McpServers = servers, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Connected); + await WaitForMcpServerStatusAsync(session, otherServerName, McpServerStatus.Connected); + + using var argument = JsonDocument.Parse("""{"name":"SCENARIO_ORIGIN_VALUE"}"""); + var sameOrigin = await session.Rpc.Mcp.Apps.CallToolAsync( + serverName, + "get_env", + originServerName: serverName, + arguments: new Dictionary + { + ["name"] = argument.RootElement.GetProperty("name").Clone(), + }); + Assert.Contains("origin-ok", sameOrigin["content"].GetRawText(), StringComparison.Ordinal); + + var crossOrigin = await Assert.ThrowsAnyAsync(() => + session.Rpc.Mcp.Apps.CallToolAsync( + serverName, + "get_env", + originServerName: otherServerName, + arguments: new Dictionary + { + ["name"] = argument.RootElement.GetProperty("name").Clone(), + })); + Assert.Contains("origin", crossOrigin.ToString(), StringComparison.OrdinalIgnoreCase); + } + + [Fact] + public async Task Should_Preserve_Disabled_Scenario_Mcp_Servers_Across_Reload_And_Resume() + { + const string enabledName = "scenario-client-enabled-mcp"; + const string disabledName = "scenario-client-disabled-mcp"; + var client1 = Ctx.CreateClient(); + var session1 = await Ctx.CreateSessionAsync(client1, new SessionConfig + { + ClientName = "scenario-client", + EnableSessionStore = true, + McpServers = CreateTestMcpServers(enabledName, disabledName), + DisabledMcpServers = [disabledName], + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + await WaitForMcpServerStatusAsync(session1, enabledName, McpServerStatus.Connected); + await WaitForMcpServerStatusAsync(session1, disabledName, McpServerStatus.Disabled); + + await session1.Rpc.Mcp.ReloadAsync(); + await WaitForMcpServerStatusAsync(session1, enabledName, McpServerStatus.Connected); + await WaitForMcpServerStatusAsync(session1, disabledName, McpServerStatus.Disabled); + Assert.Contains(disabledName, (await session1.Rpc.Mcp.ListAsync()).Host!.DisabledServers); + + var sessionId = session1.SessionId; + var response = await session1.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_MCP_DISABLED_STATE.", + }); + Assert.Contains("SCENARIO_MCP_DISABLED_STATE", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + await session1.Rpc.SuspendAsync(); + await session1.DisposeAsync(); + await client1.ForceStopAsync(); + + await using var client2 = Ctx.CreateClient(); + await using var session2 = await Ctx.ResumeSessionAsync(client2, sessionId, new ResumeSessionConfig + { + ClientName = "scenario-client", + EnableSessionStore = true, + McpServers = CreateTestMcpServers(enabledName, disabledName), + DisabledMcpServers = [disabledName], + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + await WaitForMcpServerStatusAsync(session2, enabledName, McpServerStatus.Connected); + await WaitForMcpServerStatusAsync(session2, disabledName, McpServerStatus.Disabled); + var resumed = await session2.Rpc.Mcp.ListAsync(); + Assert.Contains(disabledName, resumed.Host!.DisabledServers); + Assert.DoesNotContain(resumed.Host.PendingConnections, name => name == disabledName); + } + + private static JsonElement GetServerConfig(McpConfigList list, string serverName) + { + Assert.True(list.Servers.TryGetValue(serverName, out var config)); + return Assert.IsType(config); + } + + private static async Task ReadMatchingAsync( + ChannelReader reader, + Func predicate) + { + using var timeout = new CancellationTokenSource(EventTimeout); + while (await reader.WaitToReadAsync(timeout.Token)) + { + while (reader.TryRead(out var item)) + { + if (predicate(item)) + { + return item; + } + } + } + + throw new TimeoutException("Timed out waiting for matching MCP event."); + } + + private sealed class ScenarioOAuthMcpServer : IAsyncDisposable + { + private readonly Process _process; + private readonly HttpClient _http = new(); + + private ScenarioOAuthMcpServer(Process process, string url) + { + _process = process; + Url = url; + } + + public string Url { get; } + + public static async Task StartAsync(string expectedToken) + { + var script = Path.Join(FindTestHarnessDir(), "test-mcp-oauth-server.mjs"); + var startInfo = new ProcessStartInfo + { + FileName = "node", + Arguments = $"\"{script.Replace("\"", "\\\"")}\"", + RedirectStandardOutput = true, + RedirectStandardError = true, + UseShellExecute = false, + }; + startInfo.Environment["EXPECTED_TOKEN"] = expectedToken; + + var process = Process.Start(startInfo) + ?? throw new InvalidOperationException("Failed to start OAuth MCP server."); + var stderr = process.StandardError.ReadToEndAsync(); + using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10)); + while (!timeout.IsCancellationRequested) + { + var line = await process.StandardOutput.ReadLineAsync(timeout.Token); + if (line is null) + { + throw new InvalidOperationException($"OAuth MCP server exited before listening: {await stderr}"); + } + + if (line.StartsWith("Listening: ", StringComparison.Ordinal)) + { + return new ScenarioOAuthMcpServer(process, line["Listening: ".Length..]); + } + } + + throw new TimeoutException($"Timed out waiting for OAuth MCP server: {await stderr}"); + } + + public async Task> GetRequestsAsync() + { + var json = await _http.GetStringAsync($"{Url}/__requests"); + using var document = JsonDocument.Parse(json); + return document.RootElement.EnumerateArray() + .Select(element => new ScenarioOAuthRequest( + element.TryGetProperty("authorization", out var authorization) + && authorization.ValueKind == JsonValueKind.String + ? authorization.GetString() + : null, + element.GetProperty("path").GetString()!)) + .ToList(); + } + + public async ValueTask DisposeAsync() + { + _http.Dispose(); + if (!_process.HasExited) + { + _process.Kill(entireProcessTree: true); + await _process.WaitForExitAsync(); + } + _process.Dispose(); + } + } + + private sealed record ScenarioOAuthRequest(string? Authorization, string Path); +} diff --git a/dotnet/test/E2E/ScenarioTestingPermissionsE2ETests.cs b/dotnet/test/E2E/ScenarioTestingPermissionsE2ETests.cs new file mode 100644 index 0000000000..b9156b29bb --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingPermissionsE2ETests.cs @@ -0,0 +1,171 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using Microsoft.Extensions.AI; +using System.ComponentModel; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingPermissionsE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_permissions", output) +{ + [Theory] + [InlineData("assisted", "gpt-5.5")] + [InlineData("allow-all", null)] + public async Task Should_Set_Reset_And_Read_Authoritative_Scenario_Permission_Mode( + string modeValue, + string? assistedApprovalModel) + { + await using var session = await CreateSessionAsync(new SessionConfig + { + FeatureFlags = new Dictionary { ["AUTO_APPROVAL"] = true }, + }); + var mode = new PermissionMode(modeValue); + + Assert.Equal(PermissionMode.Manual, (await session.Rpc.Permissions.GetModeAsync()).Mode); + + var set = await session.Rpc.Permissions.SetModeAsync( + mode, + assistedApprovalModel: assistedApprovalModel, + source: PermissionModeSource.Rpc); + Assert.True(set.Success); + Assert.Equal(mode, set.Mode); + Assert.Equal(mode, (await session.Rpc.Permissions.GetModeAsync()).Mode); + + var reset = await session.Rpc.Permissions.SetModeAsync( + PermissionMode.Manual, + source: PermissionModeSource.Rpc); + Assert.True(reset.Success); + Assert.Equal(PermissionMode.Manual, reset.Mode); + Assert.Equal(PermissionMode.Manual, (await session.Rpc.Permissions.GetModeAsync()).Mode); + } + + [Fact] + public async Task Should_Report_Managed_Effective_Mode_When_Scenario_Escalation_Fails() + { + var enforced = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + EnableManagedSettings = true, + ManagedSettings = new ManagedSettings + { + Permissions = new ManagedSettingsPermissions + { + DisableBypassPermissionsMode = DisableBypassPermissionsModes.Disable, + }, + }, + OnEvent = evt => + { + if (evt is SessionManagedSettingsEnforcedEvent enforcedEvent) + { + enforced.TrySetResult(enforcedEvent); + } + }, + }); + + var resolved = await session.Rpc.ManagedSettings.GetAsync(); + Assert.True(resolved.ClientManaged); + Assert.True(resolved.BypassPermissionsDisabled); + Assert.Contains("permissions", resolved.ManagedKeys); + + var set = await session.Rpc.Permissions.SetModeAsync( + PermissionMode.AllowAll, + source: PermissionModeSource.Rpc); + Assert.False(set.Success); + Assert.NotEqual(PermissionMode.AllowAll, set.Mode); + + var authoritative = await session.Rpc.Permissions.GetModeAsync(); + Assert.Equal(set.Mode, authoritative.Mode); + + var enforcedEvent = await enforced.Task.WaitAsync(TimeSpan.FromSeconds(30)); + Assert.Equal(ManagedSettingsEnforcedAction.BypassPermissionsBlocked, enforcedEvent.Data.Action); + Assert.Equal(ManagedSettingsEnforcedEscalation.AllowAll, enforcedEvent.Data.Escalation); + Assert.Equal("permissions.disableBypassPermissionsMode", enforcedEvent.Data.Setting); + } + + [Fact] + public async Task Should_Forward_Exact_Scenario_Permission_Callback_Payload() + { + var callback = new TaskCompletionSource<(PermissionRequestCustomTool Request, PermissionInvocation Invocation)>( + TaskCreationOptions.RunContinuationsAsynchronously); + CopilotSession? session = null; + session = await CreateSessionAsync(new SessionConfig + { + Tools = + [ + AIFunctionFactory.Create( + ScenarioPermissionTool, + "scenario_permission_tool", + "Reads a scenario-owned value after user approval") + ], + OnPermissionRequest = (request, invocation) => + { + callback.TrySetResult((Assert.IsType(request), invocation)); + return Task.FromResult(PermissionDecision.ApproveOnce()); + }, + }); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_permission_tool with key 'payload', then reply with exactly its result.", + DisplayPrompt = "Run permission-gated scenario action", + Source = MessageSource.Agent("scenario-client"), + }); + + var (request, invocation) = await callback.Task.WaitAsync(TimeSpan.FromSeconds(30)); + Assert.Equal(session.SessionId, invocation.SessionId); + Assert.False(invocation.ManagedSettingsEnabled); + Assert.Equal("scenario_permission_tool", request.ToolName); + Assert.Equal("Reads a scenario-owned value after user approval", request.ToolDescription); + Assert.Equal("payload", request.Args!.Value.GetProperty("key").GetString()); + Assert.False(string.IsNullOrWhiteSpace(request.ToolCallId)); + Assert.Contains("SCENARIO_PERMISSION_PAYLOAD", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + [Description("Reads a scenario-owned value after user approval")] + static string ScenarioPermissionTool([Description("Scenario lookup key")] string key) => + $"SCENARIO_PERMISSION_{key.ToUpperInvariant()}"; + } + + [Fact] + public async Task Should_Use_Scenario_Location_And_Folder_Trust_Rpcs() + { + await using var session = await CreateSessionAsync(); + var location = Path.Join(Ctx.WorkDir, $"scenario-location-{Guid.NewGuid():N}"); + var trusted = Path.Join(Ctx.WorkDir, $"scenario-trusted-{Guid.NewGuid():N}"); + Directory.CreateDirectory(location); + Directory.CreateDirectory(trusted); + + var resolved = await session.Rpc.Permissions.Locations.ResolveAsync(location); + Assert.Equal(PermissionLocationType.Dir, resolved.LocationType); + Assert.True(PathsEqual(location, resolved.LocationKey)); + + var identifier = $"scenario-client-command-{Guid.NewGuid():N}"; + var add = await session.Rpc.Permissions.Locations.AddToolApprovalAsync( + resolved.LocationKey, + new PermissionsLocationsAddToolApprovalDetailsCommands + { + CommandIdentifiers = [identifier], + }); + Assert.True(add.Success); + + var applied = await session.Rpc.Permissions.Locations.ApplyAsync(location); + Assert.True(applied.AppliedRuleCount >= 1); + Assert.Contains(applied.AppliedRules, rule => rule.Kind == "shell" && rule.Argument == identifier); + + Assert.False((await session.Rpc.Permissions.FolderTrust.IsTrustedAsync(trusted)).Trusted); + Assert.True((await session.Rpc.Permissions.FolderTrust.AddTrustedAsync(trusted)).Success); + Assert.True((await session.Rpc.Permissions.FolderTrust.IsTrustedAsync(trusted)).Trusted); + } + + private static bool PathsEqual(string left, string right) => + string.Equals( + Path.GetFullPath(left).TrimEnd(Path.DirectorySeparatorChar, Path.AltDirectorySeparatorChar), + Path.GetFullPath(right).TrimEnd(Path.DirectorySeparatorChar, Path.AltDirectorySeparatorChar), + OperatingSystem.IsWindows() ? StringComparison.OrdinalIgnoreCase : StringComparison.Ordinal); +} diff --git a/dotnet/test/E2E/ScenarioTestingPersistenceE2ETests.cs b/dotnet/test/E2E/ScenarioTestingPersistenceE2ETests.cs new file mode 100644 index 0000000000..1872dca5fb --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingPersistenceE2ETests.cs @@ -0,0 +1,150 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingPersistenceE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_persistence", output) +{ + [Fact] + public async Task Should_Retry_From_Existing_History_With_Empty_SendMessages() + { + await using var session = await CreateSessionAsync(); + var initial = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly EMPTY_BATCH_CONTEXT_READY.", + }); + Assert.Contains("EMPTY_BATCH_CONTEXT_READY", initial?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var retry = await session.Rpc.SendMessagesAsync([], wait: true); + + Assert.Empty(retry.MessageIds); + var events = await session.GetEventsAsync(); + Assert.Single( + events.OfType(), + evt => evt.Data.Content == "Reply with exactly EMPTY_BATCH_CONTEXT_READY."); + Assert.Contains( + events.OfType(), + evt => (evt.Data.Content ?? string.Empty).Contains("EMPTY_BATCH_RETRY_DONE", StringComparison.Ordinal)); + } + + [Theory] + [InlineData(1)] + [InlineData(3)] + public async Task Should_Page_Persisted_Events_Backward_Without_Resuming(int pageSize) + { + const string firstPrompt = "Reply with exactly PERSISTED_SCENARIO_FIRST."; + const string secondPrompt = "Reply with exactly PERSISTED_SCENARIO_SECOND."; + var session = await CreateSessionAsync(); + var sessionId = session.SessionId; + + await session.SendAndWaitAsync(new MessageOptions { Prompt = firstPrompt }); + await session.SendAndWaitAsync(new MessageOptions { Prompt = secondPrompt }); + await Client.Rpc.Sessions.SaveAsync(sessionId); + await session.DisposeAsync(); + + var pages = new List(); + EventsReadResult page = await Client.Rpc.Sessions.ReadPersistedEventsAsync( + sessionId, + max: pageSize, + direction: EventsReadDirection.Backward); + pages.Add(page); + + while (page.HasMore) + { + Assert.False(string.IsNullOrWhiteSpace(page.Cursor)); + page = await Client.Rpc.Sessions.ReadPersistedEventsAsync( + sessionId, + cursor: page.Cursor, + max: pageSize); + pages.Add(page); + } + + Assert.All(pages, current => Assert.Equal(EventsCursorStatus.Ok, current.CursorStatus)); + var events = pages.SelectMany(current => current.Events).ToList(); + Assert.Equal(events.Count, events.Select(evt => evt.Id).Distinct().Count()); + + var userMessages = events + .OfType() + .Select(evt => evt.Data.Content) + .ToList(); + Assert.Contains(firstPrompt, userMessages); + Assert.Contains(secondPrompt, userMessages); + Assert.True( + userMessages.IndexOf(secondPrompt) < userMessages.IndexOf(firstPrompt), + "Backward pages should expose the newer user turn before the older turn."); + } + + [Fact] + public async Task Should_Truncate_History_And_Resend_From_Boundary() + { + const string firstPrompt = "Reply with exactly HISTORY_SCENARIO_FIRST."; + const string discardedPrompt = "Reply with exactly HISTORY_SCENARIO_DISCARDED."; + const string replacementPrompt = "Reply with exactly HISTORY_SCENARIO_REPLACEMENT."; + + await using var session = await CreateSessionAsync(); + await session.SendAndWaitAsync(new MessageOptions { Prompt = firstPrompt }); + await session.SendAndWaitAsync(new MessageOptions { Prompt = discardedPrompt }); + + var discardedEvent = (await session.GetEventsAsync()) + .OfType() + .Single(evt => evt.Data.Content == discardedPrompt); + var truncate = await session.Rpc.History.TruncateAsync(discardedEvent.Id.ToString()); + + Assert.True(truncate.EventsRemoved > 0); + Assert.NotEqual(true, truncate.CheckpointCleanupFailed); + + var replacement = await session.SendAndWaitAsync(new MessageOptions { Prompt = replacementPrompt }); + Assert.Contains("HISTORY_SCENARIO_REPLACEMENT", replacement?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + var events = await session.GetEventsAsync(); + Assert.DoesNotContain(events.OfType(), evt => evt.Data.Content == discardedPrompt); + Assert.Contains(events.OfType(), evt => evt.Data.Content == firstPrompt); + Assert.Contains(events.OfType(), evt => evt.Data.Content == replacementPrompt); + } + + [Theory] + [InlineData("session")] + [InlineData("unstaged")] + [InlineData("branch")] + public async Task Should_List_Read_And_Diff_Scenario_Workspace_State(string mode) + { + await using var session = await CreateSessionAsync(); + var workspaceFile = $"scenario-state-{Guid.NewGuid():N}.txt"; + const string workspaceContent = "SCENARIO_WORKSPACE_STATE"; + var requestedMode = new WorkspaceDiffMode(mode); + + await session.Rpc.Workspaces.CreateFileAsync(workspaceFile, workspaceContent); + + var listed = await session.Rpc.Workspaces.ListFilesAsync(); + var read = await session.Rpc.Workspaces.ReadFileAsync(workspaceFile); + var diff = await session.Rpc.Workspaces.DiffAsync(requestedMode); + + Assert.Contains(workspaceFile, listed.Files); + Assert.Equal(workspaceContent, read.Content); + Assert.Equal(requestedMode, diff.RequestedMode); + + if (requestedMode == WorkspaceDiffMode.Unstaged) + { + Assert.Equal(WorkspaceDiffMode.Unstaged, diff.Mode); + Assert.False(diff.IsFallback); + Assert.Null(diff.UnavailableReason); + } + else + { + Assert.True( + diff.Mode == requestedMode || diff.Mode == WorkspaceDiffMode.Unstaged, + $"Unexpected effective workspace diff mode: {diff.Mode}"); + Assert.Equal(diff.Mode == WorkspaceDiffMode.Unstaged, diff.IsFallback); + Assert.Equal( + requestedMode == WorkspaceDiffMode.Session && diff.IsFallback, + diff.UnavailableReason is not null); + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingProvidersE2ETests.cs b/dotnet/test/E2E/ScenarioTestingProvidersE2ETests.cs new file mode 100644 index 0000000000..13270d3e43 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingProvidersE2ETests.cs @@ -0,0 +1,441 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Collections.Concurrent; +using System.Net; +using System.Net.Http; +using System.Text; +using System.Text.Json; +using System.Text.RegularExpressions; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +/// +/// Representative scenario coverage for provider and model selection. +/// +[Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] +public class ScenarioTestingProvidersE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_providers", output) +{ + [Fact] + public async Task Should_Route_Scenario_Models_With_Provider_Auth_Headers_Wire_Ids_And_Capabilities() + { + var handler = new ScenarioProviderRequestHandler(); + await using var client = CreateProviderClient(handler); + await using var session = await Ctx.CreateSessionAsync( + client, + CreateScenarioProviderConfig("alpha/large")); + + var alphaResponse = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with the configured provider response.", + }); + Assert.Contains(ScenarioProviderRequestHandler.SyntheticText, alphaResponse?.Data.Content ?? string.Empty); + + await session.SetModelAsync("beta/fast"); + var betaResponse = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with the configured provider response again.", + }); + Assert.Contains(ScenarioProviderRequestHandler.SyntheticText, betaResponse?.Data.Content ?? string.Empty); + + var alpha = Assert.Single(handler.InferenceRequests, request => request.Host == "alpha.scenario.invalid"); + Assert.Contains("\"model\":\"alpha-wire-large\"", alpha.Body, StringComparison.Ordinal); + Assert.Equal("alpha-scenario", alpha.Headers["X-Scenario-Provider"]); + Assert.Contains("alpha-static-key", alpha.Headers["Authorization"], StringComparison.Ordinal); + + var beta = Assert.Single(handler.InferenceRequests, request => request.Host == "beta.scenario.invalid"); + Assert.Contains("\"model\":\"beta-wire-fast\"", beta.Body, StringComparison.Ordinal); + Assert.Equal("beta-scenario", beta.Headers["X-Scenario-Provider"]); + Assert.Equal("Bearer beta-static-token", beta.Headers["Authorization"]); + + var listed = await session.Rpc.Model.ListAsync(); + var alphaModel = Assert.Single( + listed.List, + model => model.GetRawText().Contains("\"id\":\"alpha/large\"", StringComparison.Ordinal)); + var alphaJson = alphaModel.GetRawText(); + Assert.Contains("\"max_context_window_tokens\":120000", alphaJson, StringComparison.Ordinal); + Assert.Contains("\"max_prompt_tokens\":100000", alphaJson, StringComparison.Ordinal); + Assert.Contains("\"reasoningEffort\":true", alphaJson, StringComparison.Ordinal); + Assert.Contains("\"vision\":false", alphaJson, StringComparison.Ordinal); + Assert.Contains( + listed.List, + model => model.GetRawText().Contains("\"id\":\"alpha/small\"", StringComparison.Ordinal)); + Assert.Contains( + listed.List, + model => model.GetRawText().Contains("\"id\":\"beta/fast\"", StringComparison.Ordinal)); + } + + [Fact] + public async Task Should_Use_Dynamic_Scenario_Bearer_Callback_For_Selected_Provider() + { + const string token = "scenario-client-dynamic-token"; + ProviderTokenArgs? observedArgs = null; + var handler = new ScenarioProviderRequestHandler(); + await using var client = CreateProviderClient(handler); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + ClientName = "scenario-client", + Model = "managed/default", + Providers = + [ + new NamedProviderConfig + { + Name = "managed", + Type = "openai", + WireApi = "completions", + BaseUrl = "https://managed.scenario.invalid/v1", + ApiKey = "must-not-win", + BearerToken = "must-not-win-either", + BearerTokenProvider = args => + { + observedArgs = args; + return Task.FromResult(token); + }, + }, + ], + Models = + [ + new ProviderModelConfig + { + Id = "default", + Provider = "managed", + ModelId = "claude-sonnet-5", + WireModel = "managed-wire-model", + }, + ], + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with the configured provider response.", + }); + + Assert.NotNull(observedArgs); + Assert.Equal("managed", observedArgs!.ProviderName); + Assert.Equal(session.SessionId, observedArgs.SessionId); + var request = Assert.Single(handler.InferenceRequests); + Assert.Equal("Bearer " + token, request.Headers["Authorization"]); + Assert.DoesNotContain("must-not-win", request.Headers["Authorization"], StringComparison.Ordinal); + } + + [Fact] + public async Task Should_Apply_Reasoning_Context_And_Auto_Atomically_Without_Implicit_Reset() + { + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Model = "claude-sonnet-5", + }); + + await session.SetModelAsync("claude-sonnet-5", new SetModelOptions + { + ReasoningEffort = "high", + ContextTier = ContextTier.LongContext, + }); + + var atomic = await session.Rpc.Model.GetCurrentAsync(); + Assert.Equal("claude-sonnet-5", atomic.ModelId); + Assert.Equal("high", atomic.ReasoningEffort); + Assert.Equal(ContextTier.LongContext, atomic.ContextTier); + + await session.SetModelAsync("auto", new SetModelOptions + { + AutoTier = AutoTier.Intelligence, + }); + var auto = await session.Rpc.Model.GetCurrentAsync(); + Assert.Equal("auto", auto.ModelId); + Assert.Equal(AutoTier.Intelligence, auto.PendingAutoTier); + + var invalid = await Assert.ThrowsAnyAsync(() => + session.SetModelAsync("claude-sonnet-5", new SetModelOptions + { + ReasoningEffort = "low", + ContextTier = ContextTier.Default, + AutoTier = AutoTier.Fast, + })); + Assert.Contains("auto", invalid.ToString(), StringComparison.OrdinalIgnoreCase); + + var afterRejected = await session.Rpc.Model.GetCurrentAsync(); + Assert.Equal("auto", afterRejected.ModelId); + Assert.Equal(AutoTier.Intelligence, afterRejected.PendingAutoTier); + + await session.SetModelAsync("auto", new SetModelOptions + { + }); + + var omittedTier = await session.Rpc.Model.GetCurrentAsync(); + Assert.Equal("auto", omittedTier.ModelId); + Assert.Equal(AutoTier.Intelligence, omittedTier.PendingAutoTier); + } + + [Fact] + public async Task Should_Resolve_Legacy_Bare_Model_Id_When_Scenario_Resumes_With_Named_Provider() + { + var initialHandler = new ScenarioProviderRequestHandler(); + var initialClient = CreateProviderClient(initialHandler); + var initialSession = await Ctx.CreateSessionAsync(initialClient, new SessionConfig + { + ClientName = "scenario-client", + Model = "legacy-scenario-model", + Provider = new ProviderConfig + { + Type = "openai", + WireApi = "completions", + BaseUrl = "https://legacy.scenario.invalid/v1", + ApiKey = "legacy-key", + ModelId = "legacy-scenario-model", + WireModel = "legacy-wire-model", + }, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + var sessionId = initialSession.SessionId; + await initialSession.SendAndWaitAsync(new MessageOptions + { + Prompt = "Persist this scenario session.", + }); + await initialSession.Rpc.SuspendAsync(); + await initialSession.DisposeAsync(); + await initialClient.ForceStopAsync(); + + var resumedHandler = new ScenarioProviderRequestHandler(); + await using var resumedClient = CreateProviderClient(resumedHandler); + await using var resumed = await Ctx.ResumeSessionAsync(resumedClient, sessionId, new ResumeSessionConfig + { + ClientName = "scenario-client", + Providers = + [ + new NamedProviderConfig + { + Name = "scenario-provider", + Type = "openai", + WireApi = "completions", + BaseUrl = "https://legacy.scenario.invalid/v1", + ApiKey = "resumed-key", + }, + ], + Models = + [ + new ProviderModelConfig + { + Id = "legacy-scenario-model", + Provider = "scenario-provider", + ModelId = "legacy-scenario-model", + WireModel = "legacy-wire-model", + }, + ], + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + Assert.Equal("legacy-scenario-model", (await resumed.Rpc.Model.GetCurrentAsync()).ModelId); + var response = await resumed.SendAndWaitAsync(new MessageOptions + { + Prompt = "Continue the legacy scenario session.", + }); + Assert.Contains(ScenarioProviderRequestHandler.SyntheticText, response?.Data.Content ?? string.Empty); + var routed = Assert.Single(resumedHandler.InferenceRequests); + Assert.NotEqual("legacy.scenario.invalid", routed.Host); + Assert.Contains("\"model\":\"claude-sonnet-5\"", routed.Body, StringComparison.Ordinal); + } + + [Fact] + public async Task Should_Ignore_Failing_Unselected_Provider_But_Surface_Selected_Provider_Failure() + { + var handler = new ScenarioProviderRequestHandler(failingHost: "offline.scenario.invalid"); + await using var client = CreateProviderClient(handler); + var config = CreateScenarioProviderConfig("alpha/large"); + config.Providers!.Add(new NamedProviderConfig + { + Name = "offline", + Type = "openai", + WireApi = "completions", + BaseUrl = "https://offline.scenario.invalid/v1", + ApiKey = "offline-key", + }); + config.Models!.Add(new ProviderModelConfig + { + Id = "broken", + Provider = "offline", + ModelId = "claude-sonnet-5", + WireModel = "offline-wire-model", + }); + + await using var session = await Ctx.CreateSessionAsync(client, config); + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with the configured provider response.", + }); + Assert.Contains(ScenarioProviderRequestHandler.SyntheticText, response?.Data.Content ?? string.Empty); + Assert.DoesNotContain(handler.InferenceRequests, request => request.Host == "offline.scenario.invalid"); + + await session.SetModelAsync("offline/broken"); + var failure = await Assert.ThrowsAnyAsync(() => + session.SendAndWaitAsync(new MessageOptions + { + Prompt = "This selected provider should fail.", + })); + Assert.Contains("offline", failure.ToString(), StringComparison.OrdinalIgnoreCase); + Assert.Contains(handler.InferenceRequests, request => request.Host == "offline.scenario.invalid"); + } + + private CopilotClient CreateProviderClient(ScenarioProviderRequestHandler handler) => + Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio(), + RequestHandler = handler, + }); + + private static SessionConfig CreateScenarioProviderConfig(string model) => new() + { + ClientName = "scenario-client", + Model = model, + Providers = + [ + new NamedProviderConfig + { + Name = "alpha", + Type = "openai", + WireApi = "completions", + BaseUrl = "https://alpha.scenario.invalid/v1", + ApiKey = "alpha-static-key", + Headers = new Dictionary { ["X-Scenario-Provider"] = "alpha-scenario" }, + }, + new NamedProviderConfig + { + Name = "beta", + Type = "openai", + WireApi = "responses", + BaseUrl = "https://beta.scenario.invalid/v1", + BearerToken = "beta-static-token", + Headers = new Dictionary { ["X-Scenario-Provider"] = "beta-scenario" }, + }, + ], + Models = + [ + new ProviderModelConfig + { + Id = "large", + Provider = "alpha", + Name = "Scenario Large", + ModelId = "claude-sonnet-5", + WireModel = "alpha-wire-large", + MaxContextWindowTokens = 120_000, + MaxPromptTokens = 100_000, + MaxOutputTokens = 8_000, + Capabilities = new ModelCapabilitiesOverride + { + Supports = new ModelCapabilitiesOverrideSupports + { + ReasoningEffort = true, + Vision = false, + }, + }, + }, + new ProviderModelConfig + { + Id = "small", + Provider = "alpha", + ModelId = "claude-sonnet-5", + WireModel = "alpha-wire-small", + }, + new ProviderModelConfig + { + Id = "fast", + Provider = "beta", + ModelId = "claude-sonnet-5", + WireModel = "beta-wire-fast", + }, + ], + OnPermissionRequest = PermissionHandler.ApproveAll, + }; +} + +internal sealed class ScenarioProviderRequestHandler(string? failingHost = null) : CopilotRequestHandler +{ + internal const string SyntheticText = "SCENARIO_PROVIDER_RESPONSE"; + private static readonly Regex WantsStreamRegex = new("\"stream\"\\s*:\\s*true", RegexOptions.Compiled); + private readonly ConcurrentQueue _requests = new(); + + internal IReadOnlyList InferenceRequests => + [.. _requests.Where(request => RecordingRequestHandler.IsInferenceUrl(request.Url))]; + + protected override async Task SendRequestAsync( + HttpRequestMessage request, + CopilotRequestContext ctx) + { + var body = request.Content is null + ? string.Empty +#if NET8_0_OR_GREATER + : await request.Content.ReadAsStringAsync(ctx.CancellationToken).ConfigureAwait(false); +#else + : await request.Content.ReadAsStringAsync().ConfigureAwait(false); +#endif + var headers = request.Headers.ToDictionary( + pair => pair.Key, + pair => string.Join(", ", pair.Value), + StringComparer.OrdinalIgnoreCase); + var uri = request.RequestUri!; + _requests.Enqueue(new ScenarioProviderRequest(uri.ToString(), uri.Host, body, headers)); + + if (string.Equals(uri.Host, failingHost, StringComparison.Ordinal)) + { + return new HttpResponseMessage(HttpStatusCode.BadGateway) + { + Content = new StringContent( + "{\"error\":{\"message\":\"offline scenario provider\"}}", + Encoding.UTF8, + "application/json"), + }; + } + + if (!RecordingRequestHandler.IsInferenceUrl(uri.ToString())) + { + return RecordingRequestHandler.BuildNonInferenceResponse(uri.ToString()); + } + + var wantsStream = WantsStreamRegex.IsMatch(body); + if (uri.AbsolutePath.EndsWith("/responses", StringComparison.OrdinalIgnoreCase)) + { + return wantsStream + ? Sse( + "event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"scenario-response\",\"object\":\"response\",\"status\":\"in_progress\",\"output\":[]}}\n\n" + + "event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"scenario-message\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n" + + "event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\"}}\n\n" + + $"event: response.output_text.delta\ndata: {{\"type\":\"response.output_text.delta\",\"output_index\":0,\"content_index\":0,\"delta\":\"{SyntheticText}\"}}\n\n" + + $"event: response.output_text.done\ndata: {{\"type\":\"response.output_text.done\",\"output_index\":0,\"content_index\":0,\"text\":\"{SyntheticText}\"}}\n\n" + + $"event: response.completed\ndata: {{\"type\":\"response.completed\",\"response\":{{\"id\":\"scenario-response\",\"object\":\"response\",\"status\":\"completed\",\"output\":[{{\"id\":\"scenario-message\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{{\"type\":\"output_text\",\"text\":\"{SyntheticText}\"}}]}}],\"usage\":{{\"input_tokens\":5,\"output_tokens\":3,\"total_tokens\":8}}}}}}\n\n") + : Json( + $"{{\"id\":\"scenario-response\",\"object\":\"response\",\"status\":\"completed\",\"output\":[{{\"id\":\"scenario-message\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{{\"type\":\"output_text\",\"text\":\"{SyntheticText}\"}}]}}],\"usage\":{{\"input_tokens\":5,\"output_tokens\":3,\"total_tokens\":8}}}}"); + } + + return wantsStream + ? Sse( + $"data: {{\"id\":\"scenario-chat\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"scenario\",\"choices\":[{{\"index\":0,\"delta\":{{\"role\":\"assistant\",\"content\":\"{SyntheticText}\"}},\"finish_reason\":null}}]}}\n\n" + + "data: {\"id\":\"scenario-chat\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"scenario\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":5,\"completion_tokens\":3,\"total_tokens\":8}}\n\n" + + "data: [DONE]\n\n") + : Json( + $"{{\"id\":\"scenario-chat\",\"object\":\"chat.completion\",\"created\":1,\"model\":\"scenario\",\"choices\":[{{\"index\":0,\"message\":{{\"role\":\"assistant\",\"content\":\"{SyntheticText}\"}},\"finish_reason\":\"stop\"}}],\"usage\":{{\"prompt_tokens\":5,\"completion_tokens\":3,\"total_tokens\":8}}}}"); + } + + private static HttpResponseMessage Json(string body) => new(HttpStatusCode.OK) + { + Content = new StringContent(body, Encoding.UTF8, "application/json"), + }; + + private static HttpResponseMessage Sse(string body) => new(HttpStatusCode.OK) + { + Content = new StringContent(body, Encoding.UTF8, "text/event-stream"), + }; +} + +internal sealed record ScenarioProviderRequest( + string Url, + string Host, + string Body, + IReadOnlyDictionary Headers); diff --git a/dotnet/test/E2E/ScenarioTestingRuntimeE2ETests.cs b/dotnet/test/E2E/ScenarioTestingRuntimeE2ETests.cs new file mode 100644 index 0000000000..0d3e3d81cb --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingRuntimeE2ETests.cs @@ -0,0 +1,468 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using System.Diagnostics; +using System.Globalization; +using System.Text.Json; +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +#pragma warning disable GHCP001 + +public class ScenarioTestingRuntimeE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_runtime", output) +{ + private static readonly TimeSpan TestTimeout = TimeSpan.FromSeconds(30); + + [Fact] + public async Task Should_Start_With_Complete_Scenario_Options_And_Extension_Launch_Provider() + { + var (cliPath, capturePath, pidPath) = await CreateFakeRuntimeAsync("normal"); + var scenarioHome = Path.Join(Ctx.WorkDir, "scenario-client-home"); + var pluginOne = Path.GetFullPath(Path.Join(Ctx.WorkDir, "plugins", "builtin-one")); + var pluginTwo = Path.GetFullPath(Path.Join(Ctx.WorkDir, "plugins", "builtin-two")); + Directory.CreateDirectory(scenarioHome); + Directory.CreateDirectory(pluginOne); + Directory.CreateDirectory(pluginTwo); + var launchProvider = new RecordingExtensionLaunchProvider(); + + await using var client = Ctx.CreateClient( + options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--pid-file", pidPath, "--behavior", "normal"]), + Mode = CopilotClientMode.Empty, + BaseDirectory = scenarioHome, + BuiltinPluginDirectories = [pluginOne, pluginTwo], + GitHubToken = "scenario-client-runtime-token", + UseLoggedInUser = false, + LogLevel = CopilotLogLevel.Debug, + SessionIdleTimeoutSeconds = 23, + EnableRemoteSessions = true, + Telemetry = new TelemetryConfig + { + OtlpEndpoint = "http://127.0.0.1:4318", + OtlpProtocol = "http/protobuf", + FilePath = Path.Join(Ctx.WorkDir, "scenario-client-telemetry.jsonl"), + ExporterType = "file", + SourceName = "scenario-client", + CaptureContent = true, + }, + ClientInfo = new CopilotClientInfo + { + ApplicationName = "scenario-client", + ApplicationVersion = "1.2.3", + IntegrationName = "copilot-sdk", + IntegrationVersion = "4.5.6", + }, + ExtensionLaunchProvider = launchProvider, + }); + + await client.StartAsync(); + + var launchRequest = await launchProvider.Request.Task.WaitAsync(TestTimeout); + Assert.Equal("project:runtime-e2e", launchRequest.Id); + Assert.Equal("runtime-e2e", launchRequest.Name); + Assert.Equal(ExtensionSource.Project, launchRequest.Source); + Assert.Equal(Path.GetFullPath(Path.Join(Ctx.WorkDir, "extension.mjs")), launchRequest.ModulePath); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => root.GetProperty("clientResponses").GetArrayLength() == 1); + var root = capture.RootElement; + var args = root.GetProperty("args").EnumerateArray().Select(item => item.GetString()).ToArray(); + var environment = root.GetProperty("env"); + var requests = root.GetProperty("requests").EnumerateArray().ToList(); + + Assert.Contains("--stdio", args); + Assert.Contains("--remote", args); + AssertArgumentValue(args, "--log-level", "debug"); + AssertArgumentValue(args, "--auth-token-env", "COPILOT_SDK_AUTH_TOKEN"); + AssertArgumentValue(args, "--session-idle-timeout", "23"); + Assert.Contains("--no-auto-login", args); + Assert.Equal(scenarioHome, environment.GetProperty("COPILOT_HOME").GetString()); + Assert.Equal("scenario-client-runtime-token", environment.GetProperty("COPILOT_SDK_AUTH_TOKEN").GetString()); + Assert.Equal("true", environment.GetProperty("COPILOT_OTEL_ENABLED").GetString()); + Assert.Equal("scenario-client", environment.GetProperty("COPILOT_OTEL_SOURCE_NAME").GetString()); + + Assert.Equal( + ["connect", "registerExtensionLaunchProvider", "plugins.builtin.set"], + requests.Select(request => request.GetProperty("method").GetString()!).ToArray()); + + var connect = requests[0].GetProperty("params"); + var clientInfo = connect.GetProperty("clientInfo"); + Assert.Equal("scenario-client", clientInfo.GetProperty("editorName").GetString()); + Assert.Equal("1.2.3", clientInfo.GetProperty("editorVersion").GetString()); + Assert.Equal("copilot-sdk", clientInfo.GetProperty("extensionName").GetString()); + Assert.Equal("4.5.6", clientInfo.GetProperty("extensionVersion").GetString()); + + var pluginPaths = requests[2] + .GetProperty("params") + .GetProperty("paths") + .EnumerateArray() + .Select(item => item.GetString()!) + .ToArray(); + Assert.Equal([pluginOne, pluginTwo], pluginPaths); + + var launchResponse = root.GetProperty("clientResponses")[0].GetProperty("result").GetProperty("launch"); + Assert.Equal("node", launchResponse.GetProperty("executable").GetString()); + Assert.Equal("extension-host", launchResponse.GetProperty("args")[0].GetString()); + Assert.Equal("scenario-client", launchResponse.GetProperty("env").GetProperty("HOST_KIND").GetString()); + } + + [Fact] + public async Task Should_Cancel_Externally_When_Startup_Handshake_Hangs() + { + var (cliPath, capturePath, pidPath) = await CreateFakeRuntimeAsync("hang-connect"); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--pid-file", pidPath, "--behavior", "hang-connect"]), + UseLoggedInUser = false, + }); + using var cancellation = new CancellationTokenSource(TimeSpan.FromMilliseconds(500)); + + await Assert.ThrowsAnyAsync(() => client.StartAsync(cancellation.Token)); + + var pid = int.Parse(await File.ReadAllTextAsync(pidPath), CultureInfo.InvariantCulture); + await AssertProcessExitedAsync(pid); + await client.ForceStopAsync(); + } + + [Fact] + public async Task Should_Ping_Then_Reuse_Client_Across_Two_Sessions() + { + await using var client = Ctx.CreateClient(); + await client.StartAsync(); + + var ping = await client.PingAsync("scenario-client-reuse"); + Assert.Equal("pong: scenario-client-reuse", ping.Message); + + string firstSessionId; + await using (var first = await Ctx.CreateSessionAsync(client)) + { + firstSessionId = first.SessionId; + var response = await first.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly FIRST_SCENARIO_SESSION.", + }); + Assert.Contains("FIRST_SCENARIO_SESSION", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + } + + await using (var second = await Ctx.CreateSessionAsync(client)) + { + Assert.NotEqual(firstSessionId, second.SessionId); + var response = await second.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SECOND_SCENARIO_SESSION.", + }); + Assert.Contains("SECOND_SCENARIO_SESSION", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + } + } + + [Fact] + public async Task Should_Bound_Graceful_Stop_Then_Force_Stop() + { + var (cliPath, capturePath, pidPath) = await CreateFakeRuntimeAsync("hang-detach"); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--pid-file", pidPath, "--behavior", "hang-detach"]), + UseLoggedInUser = false, + }); + await using var session = await Ctx.CreateSessionAsync(client); + + var stopTask = client.StopAsync(); + var completed = await Task.WhenAny(stopTask, Task.Delay(TimeSpan.FromMilliseconds(500))); + Assert.NotSame(stopTask, completed); + + await client.ForceStopAsync(); + await stopTask.WaitAsync(TestTimeout); + + var pid = int.Parse(await File.ReadAllTextAsync(pidPath), CultureInfo.InvariantCulture); + await AssertProcessExitedAsync(pid); + } + + [Fact] + public async Task Should_Fail_Fast_After_Transport_Failure() + { + var (cliPath, capturePath, pidPath) = await CreateFakeRuntimeAsync("exit-after-create"); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--pid-file", pidPath, "--behavior", "exit-after-create"]), + UseLoggedInUser = false, + }); + await using var session = await Ctx.CreateSessionAsync(client); + + var pid = int.Parse(await File.ReadAllTextAsync(pidPath), CultureInfo.InvariantCulture); + await AssertProcessExitedAsync(pid); + + Exception sendException; + Exception pingException; + try + { + sendException = await Assert.ThrowsAnyAsync( + () => session.SendAsync(new MessageOptions { Prompt = "This transport is already gone." }) + .WaitAsync(TimeSpan.FromSeconds(3))); + pingException = await Assert.ThrowsAnyAsync( + () => client.PingAsync("after-failure").WaitAsync(TimeSpan.FromSeconds(3))); + } + finally + { + await client.ForceStopAsync(); + } + + Assert.IsNotType(sendException); + Assert.IsNotType(pingException); + } + + private async Task<(string CliPath, string CapturePath, string PidPath)> CreateFakeRuntimeAsync(string behavior) + { + var cliPath = Path.Join(Ctx.WorkDir, $"scenario-client-runtime-{behavior}-{Guid.NewGuid():N}.js"); + var capturePath = Path.Join(Ctx.WorkDir, $"scenario-client-runtime-{behavior}-{Guid.NewGuid():N}.json"); + var pidPath = Path.Join(Ctx.WorkDir, $"scenario-client-runtime-{behavior}-{Guid.NewGuid():N}.pid"); + await File.WriteAllTextAsync(cliPath, FakeRuntimeScript); + return (cliPath, capturePath, pidPath); + } + + private static async Task WaitForCaptureAsync( + string path, + Func predicate) + { + JsonDocument? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + try + { + using var stream = new FileStream( + path, + FileMode.Open, + FileAccess.Read, + FileShare.ReadWrite | FileShare.Delete); + using var reader = new StreamReader(stream); + var json = await reader.ReadToEndAsync(); + var document = JsonDocument.Parse(json); + if (!predicate(document.RootElement)) + { + document.Dispose(); + return false; + } + + result = document; + return true; + } + catch (Exception ex) when (ex is IOException or JsonException) + { + return false; + } + }, + timeout: TestTimeout, + pollInterval: TimeSpan.FromMilliseconds(50), + timeoutMessage: $"Timed out waiting for fake runtime capture at {path}."); + return result!; + } + + private static void AssertArgumentValue(string?[] args, string name, string expectedValue) + { + var index = Array.IndexOf(args, name); + Assert.True(index >= 0, $"Expected argument '{name}' was not present."); + Assert.True(index + 1 < args.Length, $"Expected argument '{name}' to have a value."); + Assert.Equal(expectedValue, args[index + 1]); + } + + private static async Task AssertProcessExitedAsync(int pid) + { + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(!IsProcessRunning(pid)), + timeout: TestTimeout, + pollInterval: TimeSpan.FromMilliseconds(50), + timeoutMessage: $"Expected process {pid} to exit."); + } + + private static bool IsProcessRunning(int pid) + { + try + { + using var process = Process.GetProcessById(pid); + return !process.HasExited; + } + catch (Exception ex) when (ex is ArgumentException or InvalidOperationException) + { + return false; + } + } + + private sealed class RecordingExtensionLaunchProvider : IExtensionLaunchProviderHandler + { + public TaskCompletionSource Request { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task ResolveAsync( + ExtensionLaunchProviderResolveRequest request, + CancellationToken cancellationToken = default) + { + Request.TrySetResult(request); + return Task.FromResult(new ExtensionLaunchProviderResolveResult + { + Launch = new ExtensionLaunchProfile + { + Executable = "node", + Args = ["extension-host", request.ModulePath], + Env = new Dictionary { ["HOST_KIND"] = "scenario-client" }, + }, + }); + } + } + + private const string FakeRuntimeScript = """ + const fs = require("fs"); + + function argument(name) { + const index = process.argv.indexOf(name); + return index >= 0 ? process.argv[index + 1] : undefined; + } + + const captureFile = argument("--capture-file"); + const pidFile = argument("--pid-file"); + const behavior = argument("--behavior") || "normal"; + const requests = []; + const clientResponses = []; + let nextRequestId = 1000; + let buffer = Buffer.alloc(0); + + fs.writeFileSync(pidFile, String(process.pid)); + + function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ + args: process.argv.slice(2), + requests, + clientResponses, + env: { + COPILOT_HOME: process.env.COPILOT_HOME, + COPILOT_SDK_AUTH_TOKEN: process.env.COPILOT_SDK_AUTH_TOKEN, + COPILOT_OTEL_ENABLED: process.env.COPILOT_OTEL_ENABLED, + COPILOT_OTEL_SOURCE_NAME: process.env.COPILOT_OTEL_SOURCE_NAME + } + })); + } + + function write(message) { + const body = JSON.stringify(message); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function respond(id, result) { + write({ jsonrpc: "2.0", id, result }); + } + + function request(method, params) { + const id = nextRequestId++; + write({ jsonrpc: "2.0", id, method, params }); + return id; + } + + function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "method")) { + clientResponses.push(message); + saveCapture(); + return; + } + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + if (behavior !== "hang-connect") { + respond(message.id, { ok: true, protocolVersion: 3, version: "fake" }); + } + return; + } + + if (message.method === "registerExtensionLaunchProvider") { + respond(message.id, {}); + setTimeout(() => request("extensionLaunchProvider.resolve", { + id: "project:runtime-e2e", + modulePath: require("path").resolve(process.cwd(), "extension.mjs"), + name: "runtime-e2e", + source: "project" + }), 10); + return; + } + + if (message.method === "plugins.builtin.set") { + respond(message.id, {}); + return; + } + + if (message.method === "ping") { + respond(message.id, { + message: `pong: ${message.params?.message ?? ""}`, + timestamp: new Date().toISOString(), + protocolVersion: 3 + }); + return; + } + + if (message.method === "session.create") { + const sessionId = message.params?.sessionId ?? "fake-session"; + respond(message.id, { sessionId, workspacePath: null, capabilities: null }); + if (behavior === "exit-after-create") { + setTimeout(() => process.exit(17), 25); + } + return; + } + + if (message.method === "session.detach" && behavior === "hang-detach") { + return; + } + + if (message.method === "session.detach") { + respond(message.id, { success: true }); + return; + } + + if (message.method === "runtime.shutdown") { + respond(message.id, {}); + setTimeout(() => process.exit(0), 10); + return; + } + + respond(message.id, {}); + } + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } + }); + + process.stdin.resume(); + saveCapture(); + setInterval(() => {}, 1000); + """; +} + +#pragma warning restore GHCP001 diff --git a/dotnet/test/E2E/ScenarioTestingSendsE2ETests.cs b/dotnet/test/E2E/ScenarioTestingSendsE2ETests.cs new file mode 100644 index 0000000000..25dcae6c32 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingSendsE2ETests.cs @@ -0,0 +1,340 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using System.ComponentModel; +using System.Diagnostics; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingSendsE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_sends", output) +{ + private static readonly TimeSpan SendTimeout = TimeSpan.FromSeconds(60); + + [Fact] + public async Task Should_Send_Complete_Scenario_Message_Wire_Shape() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "normal"]), + UseLoggedInUser = false, + }); + + using var activity = new Activity("scenario-client-send"); + activity.SetIdFormat(ActivityIdFormat.W3C); + activity.TraceStateString = "scenario-client=send"; + activity.Start(); + + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + Streaming = true, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + var filePath = Path.Join(Ctx.WorkDir, "scenario-wire-file.txt"); + var directoryPath = Path.Join(Ctx.WorkDir, "scenario-wire-directory"); + var selectionPath = Path.Join(Ctx.WorkDir, "Program.cs"); + using var payload = JsonDocument.Parse("""{"selection":"SCENARIO_SELECTION","line":17}"""); + var messageId = await session.SendAsync(new MessageOptions + { + Prompt = "Use the hidden scenario context.", + DisplayPrompt = "Review selected scenario context", + Mode = "enqueue", + AgentMode = AgentMode.Interactive, + Source = MessageSource.Agent("scenario-client"), + Attachments = + [ + new AttachmentFile + { + DisplayName = "scenario-wire-file.txt", + Path = filePath, + LineRange = new AttachmentFileLineRange { Start = 3, End = 9 }, + }, + new AttachmentDirectory + { + DisplayName = "scenario-wire-directory", + Path = directoryPath, + }, + new AttachmentSelection + { + DisplayName = "Program.cs", + FilePath = selectionPath, + Text = "SCENARIO_SELECTION", + Selection = new AttachmentSelectionDetails + { + Start = new AttachmentSelectionDetailsStart { Line = 16, Character = 0 }, + End = new AttachmentSelectionDetailsEnd { Line = 16, Character = 13 }, + }, + }, + new AttachmentGitHubReference + { + Number = 610, + ReferenceType = AttachmentGitHubReferenceType.Pr, + State = "open", + Title = "Scenario-shaped E2E coverage", + Url = "https://github.com/github/copilot-sdk/pull/610", + }, + new AttachmentBlob + { + Data = "QVBQX0JMT0I=", + MimeType = "text/plain", + DisplayName = "scenario-wire-blob.txt", + }, + new AttachmentExtensionContext + { + CapturedAt = DateTimeOffset.Parse("2026-09-17T20:00:00Z"), + ExtensionId = "scenario-client:code-review", + CanvasId = "diff", + InstanceId = "diff-17", + Title = "Selected change", + Payload = payload.RootElement.Clone(), + }, + ], + }); + + Assert.Equal("scenario-client-message", messageId); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + var send = Assert.Single(requests, request => request.GetProperty("method").GetString() == "session.send"); + var parameters = send.GetProperty("params"); + + Assert.Equal("Use the hidden scenario context.", parameters.GetProperty("prompt").GetString()); + Assert.Equal("Review selected scenario context", parameters.GetProperty("displayPrompt").GetString()); + Assert.Equal("enqueue", parameters.GetProperty("mode").GetString()); + Assert.Equal("interactive", parameters.GetProperty("agentMode").GetString()); + Assert.Equal("agent-scenario-client", parameters.GetProperty("source").GetString()); + Assert.Equal(activity.Id, parameters.GetProperty("traceparent").GetString()); + Assert.Equal("scenario-client=send", parameters.GetProperty("tracestate").GetString()); + + var attachments = parameters.GetProperty("attachments").EnumerateArray().ToArray(); + Assert.Equal( + ["file", "directory", "selection", "github_reference", "blob", "extension_context"], + attachments.Select(item => item.GetProperty("type").GetString())); + Assert.Equal(filePath, attachments[0].GetProperty("path").GetString()); + Assert.Equal(3, attachments[0].GetProperty("lineRange").GetProperty("start").GetInt32()); + Assert.Equal(directoryPath, attachments[1].GetProperty("path").GetString()); + Assert.Equal("SCENARIO_SELECTION", attachments[2].GetProperty("text").GetString()); + Assert.Equal(selectionPath, attachments[2].GetProperty("filePath").GetString()); + Assert.Equal(610, attachments[3].GetProperty("number").GetInt32()); + Assert.Equal("pr", attachments[3].GetProperty("referenceType").GetString()); + Assert.Equal("QVBQX0JMT0I=", attachments[4].GetProperty("data").GetString()); + Assert.Equal("text/plain", attachments[4].GetProperty("mimeType").GetString()); + Assert.Equal("scenario-client:code-review", attachments[5].GetProperty("extensionId").GetString()); + Assert.Equal("SCENARIO_SELECTION", attachments[5].GetProperty("payload").GetProperty("selection").GetString()); + } + + [Theory] + [InlineData(null)] + [InlineData("enqueue")] + [InlineData("immediate")] + public async Task Should_Not_Invoke_Send_When_Scenario_Cancels_Before_Dispatch(string? mode) + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "normal"]), + UseLoggedInUser = false, + }); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + using var cancellation = new CancellationTokenSource(); + cancellation.Cancel(); + + await Assert.ThrowsAnyAsync(() => + session.SendAsync( + new MessageOptions + { + Prompt = "This message must never be invoked.", + DisplayPrompt = "Cancelled scenario message", + Mode = mode, + Source = MessageSource.Agent("scenario-client"), + }, + cancellation.Token)); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + Assert.DoesNotContain(requests, request => request.GetProperty("method").GetString() == "session.send"); + } + + [Theory] + [InlineData(null)] + [InlineData("enqueue")] + [InlineData("immediate")] + public async Task Should_Not_Replay_Scenario_Send_After_Ambiguous_Transport_Loss(string? mode) + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "drop-after-send"]), + UseLoggedInUser = false, + }); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + using var cancellation = new CancellationTokenSource(TimeSpan.FromSeconds(10)); + await Assert.ThrowsAnyAsync(() => + session.SendAsync(new MessageOptions + { + Prompt = "AMBIGUOUS_SCENARIO_SEND", + DisplayPrompt = "Ambiguous scenario send", + Mode = mode, + Source = MessageSource.Agent("scenario-client"), + }, cancellation.Token)); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + var send = Assert.Single(requests, request => request.GetProperty("method").GetString() == "session.send"); + var parameters = send.GetProperty("params"); + if (mode is null) + { + Assert.False(parameters.TryGetProperty("mode", out _)); + } + else + { + Assert.Equal(mode, parameters.GetProperty("mode").GetString()); + } + } + + [Fact] + public async Task Should_Order_Idle_Queued_And_Immediate_Scenario_Delivery() + { + var firstToolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var secondToolStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseFirstTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseSecondTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var toolInvocationCount = 0; + var messages = new List(); + + await using var session = await CreateSessionAsync(new SessionConfig + { + Tools = [AIFunctionFactory.Create(BlockingTurn, "scenario_send_blocker")], + }); + using var subscription = session.On(message => + { + lock (messages) + { + messages.Add(message); + } + }); + + var idleEnqueue = TestHelper.GetNextEventOfTypeAsync(session, SendTimeout); + var idleEnqueueId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly IDLE_ENQUEUE.", + Mode = "enqueue", + Source = MessageSource.Agent("scenario-client"), + }); + await idleEnqueue; + + var idleImmediate = TestHelper.GetNextEventOfTypeAsync(session, SendTimeout); + var idleImmediateId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly IDLE_IMMEDIATE.", + Mode = "immediate", + Source = MessageSource.Agent("scenario-client"), + }); + await idleImmediate; + + await session.SendAsync(new MessageOptions + { + Prompt = "Call scenario_send_blocker, then reply with its result.", + Source = MessageSource.Agent("scenario-client"), + }); + await firstToolStarted.Task.WaitAsync(SendTimeout); + + var steeringId = await session.SendAsync(new MessageOptions + { + Prompt = "Call scenario_send_blocker again, then reply with exactly FIRST_STEERING.", + Mode = "immediate", + Source = MessageSource.Agent("scenario-client"), + }); + releaseFirstTool.TrySetResult("SCENARIO_SEND_BLOCKER_RELEASED"); + await secondToolStarted.Task.WaitAsync(SendTimeout); + + var immediateBehindSteeringId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly SECOND_IMMEDIATE.", + Mode = "immediate", + Source = MessageSource.Agent("scenario-client"), + }); + var queuedId = await session.SendAsync(new MessageOptions + { + Prompt = "Reply with exactly FINAL_QUEUED.", + Mode = "enqueue", + Source = MessageSource.Agent("scenario-client"), + }); + + var finalQueuedResponse = TestHelper.GetNextEventOfTypeAsync( + session, + message => message.Data.Content?.Contains("FINAL_QUEUED", StringComparison.Ordinal) == true, + SendTimeout, + "the final queued scenario response"); + releaseSecondTool.TrySetResult("SCENARIO_SEND_BLOCKER_RELEASED_AGAIN"); + + await TestHelper.WaitForConditionAsync( + () => + { + lock (messages) + { + return Task.FromResult( + messages.Any(message => message.Data.MessageId == steeringId) && + messages.Any(message => message.Data.MessageId == immediateBehindSteeringId) && + messages.Any(message => message.Data.MessageId == queuedId)); + } + }, + timeout: SendTimeout, + timeoutMessage: "Timed out waiting for all scenario delivery classifications."); + await finalQueuedResponse; + + List observed; + lock (messages) + { + observed = [.. messages]; + } + + Assert.Equal(UserMessageDelivery.Idle, Find(idleEnqueueId).Data.Delivery); + Assert.Equal(UserMessageDelivery.Idle, Find(idleImmediateId).Data.Delivery); + Assert.Equal(UserMessageDelivery.Steering, Find(steeringId).Data.Delivery); + Assert.Equal(UserMessageDelivery.Steering, Find(immediateBehindSteeringId).Data.Delivery); + Assert.Equal(UserMessageDelivery.Queued, Find(queuedId).Data.Delivery); + + var steeringIndex = observed.FindIndex(message => message.Data.MessageId == steeringId); + var behindIndex = observed.FindIndex(message => message.Data.MessageId == immediateBehindSteeringId); + var queuedIndex = observed.FindIndex(message => message.Data.MessageId == queuedId); + Assert.True(steeringIndex < behindIndex, "The second immediate update must remain ordered behind the first steering update."); + Assert.True(behindIndex < queuedIndex, "The second immediate message must retain its position ahead of the later enqueue."); + + UserMessageEvent Find(string id) => + Assert.Single(observed, message => string.Equals(message.Data.MessageId, id, StringComparison.Ordinal)); + + [Description("Blocks an active scenario turn until delivery ordering is staged")] + async Task BlockingTurn(CancellationToken cancellationToken) + { + if (Interlocked.Increment(ref toolInvocationCount) == 1) + { + firstToolStarted.TrySetResult(); + return await releaseFirstTool.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + + secondToolStarted.TrySetResult(); + return await releaseSecondTool.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingServerControlE2ETests.cs b/dotnet/test/E2E/ScenarioTestingServerControlE2ETests.cs new file mode 100644 index 0000000000..bf3e2cb7e8 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingServerControlE2ETests.cs @@ -0,0 +1,232 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingServerControlE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_server_control", output) +{ + [Theory] + [InlineData("all")] + [InlineData("mcp")] + [InlineData("skills")] + public async Task Should_Search_Server_Catalog_With_Category_Contract(string category) + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: CreateFakeCliOptions(cliPath, capturePath)); + await client.StartAsync(); + + var kinds = category switch + { + "all" => new[] { CatalogCandidateKind.McpServer, CatalogCandidateKind.AiSkill }, + "mcp" => [CatalogCandidateKind.McpServer], + "skills" => [CatalogCandidateKind.AiSkill], + _ => throw new ArgumentOutOfRangeException(nameof(category)), + }; + var capabilities = category switch + { + "all" => new[] { "mcp-server-card", "ai-skill-discovery" }, + "mcp" => ["mcp-server-card"], + "skills" => ["ai-skill-discovery"], + _ => throw new ArgumentOutOfRangeException(nameof(category)), + }; + + var result = await client.Rpc.Catalog.SearchAsync( + new CatalogClientContract + { + ProtocolVersion = 3, + RequiredCapabilities = capabilities, + }, + query: "scenario search", + limit: 50, + kinds: kinds); + + var succeeded = Assert.IsType(result); + Assert.Empty(succeeded.Candidates); + Assert.Equal("scenario-search", succeeded.SearchId); + Assert.False(succeeded.Truncated); + Assert.Equal(3, succeeded.Negotiated.RuntimeProtocolVersion); + Assert.Equal(capabilities, succeeded.Negotiated.GrantedCapabilities.Select(capability => capability.Value)); + + var request = Assert.Single( + await ReadRequestsAsync(capturePath, "catalog.search")).GetProperty("params"); + Assert.Equal("scenario search", request.GetProperty("query").GetString()); + Assert.Equal(50, request.GetProperty("limit").GetInt32()); + Assert.Equal(3, request.GetProperty("contract").GetProperty("protocolVersion").GetInt32()); + Assert.Equal( + capabilities, + request.GetProperty("contract").GetProperty("requiredCapabilities") + .EnumerateArray().Select(item => item.GetString())); + Assert.Equal( + kinds.Select(kind => kind.Value), + request.GetProperty("kinds").EnumerateArray().Select(item => item.GetString())); + } + + [Fact] + public async Task Should_Observe_Page_And_Cancel_Factory_Run() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: CreateFakeCliOptions(cliPath, capturePath)); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig()); + + var runs = await session.Rpc.Factory.ListRunsAsync(afterSeq: 3, beforeSeq: 20, limit: 10); + var summary = Assert.Single(runs.Runs); + Assert.Equal("factory-run-1", summary.RunId); + Assert.Equal("scenario-factory", summary.FactoryName); + Assert.Equal(FactoryRunStatus.Running, summary.Status); + Assert.Equal(7, runs.OldestSeq); + Assert.Equal(7, runs.NewestSeq); + Assert.False(runs.HasMoreNewer); + + var detail = await session.Rpc.Factory.GetRunDetailAsync(summary.RunId); + Assert.Equal(summary.RunId, detail.RunId); + Assert.Equal(summary.FactoryName, detail.FactoryName); + Assert.Equal(FactoryRunStatus.Running, detail.Status); + Assert.Equal(4, detail.Revision); + + var progress = await session.Rpc.Factory.GetRunProgressAsync( + summary.RunId, + phaseId: "verify", + afterSeq: 5, + beforeSeq: 20, + limit: 25); + var line = Assert.Single(progress.Records); + Assert.Equal(12, line.Seq); + Assert.Equal("verify", line.PhaseId); + Assert.Equal(FactoryLogLineKind.Log, line.Kind); + Assert.Equal("Validation complete", line.Text); + + var cancelled = await session.Rpc.Factory.CancelAsync(summary.RunId); + Assert.Equal(summary.RunId, cancelled.RunId); + Assert.Equal(FactoryRunStatus.Cancelled, cancelled.Status); + Assert.Equal("cancelled by user", cancelled.Reason); + + var requests = await ScenarioTestingTestCli.ReadRequestsAsync(capturePath); + var list = Assert.Single(requests, request => GetMethod(request) == "session.factory.listRuns") + .GetProperty("params"); + Assert.Equal(3, list.GetProperty("afterSeq").GetInt64()); + Assert.Equal(20, list.GetProperty("beforeSeq").GetInt64()); + Assert.Equal(10, list.GetProperty("limit").GetInt32()); + + var progressRequest = Assert.Single( + requests, + request => GetMethod(request) == "session.factory.getRunProgress").GetProperty("params"); + Assert.Equal("factory-run-1", progressRequest.GetProperty("runId").GetString()); + Assert.Equal("verify", progressRequest.GetProperty("phaseId").GetString()); + Assert.Equal(5, progressRequest.GetProperty("afterSeq").GetInt64()); + Assert.Equal(20, progressRequest.GetProperty("beforeSeq").GetInt64()); + Assert.Equal(25, progressRequest.GetProperty("limit").GetInt32()); + } + + [Theory] + [InlineData("on", true)] + [InlineData("export", false)] + public async Task Should_Read_Autopilot_State_And_Enable_Remote_Mode( + string remoteMode, + bool expectedSteerable) + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: CreateFakeCliOptions(cliPath, capturePath)); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig()); + + var objective = (await session.Rpc.AutopilotObjective.GetStateAsync()).State; + Assert.NotNull(objective); + Assert.Equal(17, objective.Id); + Assert.Equal("Ship the scenario.", objective.Objective); + Assert.Equal(AutopilotObjectiveStatus.Active, objective.Status); + Assert.Equal(3, objective.TurnCount); + Assert.Equal("1250000000", objective.CreditCountNanoAiu); + Assert.Equal(5, objective.CreditLimit!.Credits); + Assert.Equal(1.25, objective.CreditLimit.CreditsUsed); + Assert.Equal("1250000000", objective.CreditLimit.CreditsUsedNanoAiu); + + var enabled = await session.Rpc.Remote.EnableAsync(new RemoteSessionMode(remoteMode)); + Assert.Equal(expectedSteerable, enabled.RemoteSteerable); + Assert.Equal($"https://example.test/sessions/{session.SessionId}", enabled.Url); + + var remoteRequest = Assert.Single( + await ReadRequestsAsync(capturePath, "session.remote.enable")).GetProperty("params"); + Assert.Equal(session.SessionId, remoteRequest.GetProperty("sessionId").GetString()); + Assert.Equal(remoteMode, remoteRequest.GetProperty("mode").GetString()); + } + + [Fact] + public async Task Should_Edit_Reorder_Duplicate_Remove_And_Send_Queued_Items() + { + var (cliPath, capturePath) = await ScenarioTestingTestCli.CreateAsync(Ctx); + await using var client = Ctx.CreateClient(options: CreateFakeCliOptions(cliPath, capturePath)); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig()); + + await session.Rpc.Queue.SetDrainPausedAsync(true); + var first = await session.Rpc.Queue.InsertAtAsync( + 0, + new QueueInsertMessage + { + Prompt = "First hidden prompt", + DisplayPrompt = "First visible prompt", + AgentMode = SendAgentMode.Interactive, + }); + var second = await session.Rpc.Queue.InsertAtAsync( + 1, + new QueueInsertMessage + { + Prompt = "Second hidden prompt", + DisplayPrompt = "Second visible prompt", + AgentMode = SendAgentMode.Plan, + }); + + Assert.True(await UpdateTextAsync()); + var duplicate = await session.Rpc.Queue.DuplicateAtAsync(first.Id); + Assert.NotEqual(first.Id, duplicate.Id); + Assert.True((await session.Rpc.Queue.MoveItemAsync(second.Id, 0)).Changed); + + var reordered = await session.Rpc.Queue.PendingItemsAsync(); + Assert.Equal([second.Id, first.Id, duplicate.Id], reordered.Items.Select(item => item.Id)); + Assert.Equal("Updated visible prompt", reordered.Items[1].DisplayText); + Assert.Equal(SendAgentMode.Interactive, reordered.Items[1].AgentMode); + + Assert.True((await session.Rpc.Queue.SendNowAsync(second.Id)).Steered); + Assert.True((await session.Rpc.Queue.RemoveAtAsync(duplicate.Id)).Removed); + + var remaining = Assert.Single((await session.Rpc.Queue.PendingItemsAsync()).Items); + Assert.Equal(first.Id, remaining.Id); + Assert.Equal("Updated visible prompt", remaining.DisplayText); + await session.Rpc.Queue.SetDrainPausedAsync(false); + + var pauseRequests = await ReadRequestsAsync(capturePath, "session.queue.setDrainPaused"); + Assert.Equal([true, false], pauseRequests.Select( + request => request.GetProperty("params").GetProperty("paused").GetBoolean())); + + async Task UpdateTextAsync() + { + var result = await session.Rpc.Queue.UpdateTextAsync( + first.Id, + "Updated hidden prompt", + "Updated visible prompt"); + return result.Updated; + } + } + + private static CopilotClientOptions CreateFakeCliOptions(string cliPath, string capturePath) => new() + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "control-rpcs"]), + UseLoggedInUser = false, + }; + + private static async Task ReadRequestsAsync(string capturePath, string method) => + (await ScenarioTestingTestCli.ReadRequestsAsync(capturePath)) + .Where(request => GetMethod(request) == method) + .ToArray(); + + private static string? GetMethod(JsonElement request) => + request.GetProperty("method").GetString(); +} diff --git a/dotnet/test/E2E/ScenarioTestingSessionSetupE2ETests.cs b/dotnet/test/E2E/ScenarioTestingSessionSetupE2ETests.cs new file mode 100644 index 0000000000..df0431055e --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingSessionSetupE2ETests.cs @@ -0,0 +1,810 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using System.Collections.Concurrent; +using System.ComponentModel; +using System.Text.Json; +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +#pragma warning disable GHCP001 + +public class ScenarioTestingSessionSetupE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_session_setup", output) +{ + private static readonly TimeSpan TestTimeout = TimeSpan.FromSeconds(60); + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Round_Trip_Full_Composed_Scenario_Session_Config() + { + var (cliPath, capturePath) = await CreateFakeRuntimeAsync("capture"); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "capture"]), + UseLoggedInUser = false, + }); + + var sessionId = $"scenario-client-composed-{Guid.NewGuid():N}"; + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + SessionId = sessionId, + ClientName = "scenario-client", + Model = "claude-sonnet-5", + ReasoningEffort = "high", + ReasoningSummary = ReasoningSummary.Detailed, + ContextTier = ContextTier.LongContext, + Streaming = true, + IncludeSubAgentStreamingEvents = false, + SystemMessage = new SystemMessageConfig + { + Mode = SystemMessageMode.Append, + Content = "SCENARIO_COMPOSED_SYSTEM_MESSAGE", + }, + EnableConfigDiscovery = true, + EnableSessionTelemetry = false, + EnableExperimentalMode = true, + SkipCustomInstructions = false, + CustomAgentsLocalOnly = true, + CoauthorEnabled = false, + ManageScheduleEnabled = false, + SkipEmbeddingRetrieval = true, + EmbeddingCacheStorage = EmbeddingCacheStorageMode.InMemory, + OrganizationCustomInstructions = "SCENARIO_ORG_INSTRUCTIONS", + EnableOnDemandInstructionDiscovery = false, + EnableFileHooks = false, + EnableHostGitOperations = false, + EnableSessionStore = false, + EnableSkills = false, + AvailableTools = ["scenario_tool"], + ExcludedTools = ["shell"], + Tools = [AIFunctionFactory.Create(() => "unused", "scenario_tool")], + Commands = + [ + new CommandDefinition + { + Name = "scenario-command", + Description = "Scenario command", + Handler = _ => Task.CompletedTask, + }, + ], + McpServers = new Dictionary + { + ["scenario-mcp"] = new McpStdioServerConfig + { + Command = "node", + Args = ["scenario-mcp.mjs"], + Tools = ["*"], + }, + }, + CustomAgents = + [ + new CustomAgentConfig + { + Name = "scenario-agent", + DisplayName = "Scenario Agent", + Description = "scenario client agent", + Prompt = "Act as the scenario agent.", + Tools = ["scenario_tool"], + }, + ], + DefaultAgent = new DefaultAgentConfig { ExcludedTools = ["edit"] }, + Agent = "scenario-agent", + Providers = + [ + new NamedProviderConfig + { + Name = "scenario-provider", + Type = "openai", + WireApi = "responses", + BaseUrl = "https://provider.example.test/v1", + BearerTokenProvider = _ => Task.FromResult("scenario-provider-token"), + }, + ], + Models = + [ + new ProviderModelConfig + { + Provider = "scenario-provider", + Id = "scenario-model", + ModelId = "claude-sonnet-5", + WireModel = "scenario-wire-model", + }, + ], + RemoteSession = RemoteSessionMode.Export, + EnableMcpApps = true, + GitHubMcpToolConfig = new GitHubMcpToolConfig + { + EnableAllTools = false, + AdditionalTools = ["issues.get"], + DisableFormDeferral = true, + }, + RequestCanvasRenderer = true, + RequestExtensions = true, + ExtensionSdkPath = "scenario-extension-sdk", + ExtensionInfo = new ExtensionInfo { Source = "scenario-client", Name = "desktop" }, + CanvasProvider = new CanvasProviderIdentity { Id = "scenario:builtin:desktop", Name = "scenario client" }, + Canvases = + [ + new CanvasDeclaration + { + Id = "scenario-canvas", + DisplayName = "Scenario Canvas", + Description = "Scenario-hosted canvas", + }, + ], + CanvasHandler = new NoOpCanvasHandler(), + OnPermissionRequest = PermissionHandler.ApproveAll, + OnUserInputRequest = (_, _) => Task.FromResult(new UserInputResponse { Answer = "yes" }), + OnElicitationRequest = _ => Task.FromResult(new ElicitationResult { Action = UIElicitationResponseAction.Accept }), + OnExitPlanModeRequest = (_, _) => Task.FromResult(new ExitPlanModeResult { Approved = true }), + OnAutoModeSwitchRequest = (_, _) => Task.FromResult(AutoModeSwitchResponse.No), + OnMcpAuthRequest = _ => Task.FromResult(McpAuthResult.Cancel()), + OnEvent = _ => { }, + }); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => GetRequests(root, "session.create").Count == 1 + && GetRequests(root, "session.options.update").Count == 1); + var request = Assert.Single(GetRequests(capture.RootElement, "session.create")).GetProperty("params"); + var optionsUpdate = Assert.Single(GetRequests(capture.RootElement, "session.options.update")).GetProperty("params"); + + Assert.Equal(sessionId, request.GetProperty("sessionId").GetString()); + Assert.Equal("scenario-client", request.GetProperty("clientName").GetString()); + Assert.Equal("claude-sonnet-5", request.GetProperty("model").GetString()); + Assert.Equal("high", request.GetProperty("reasoningEffort").GetString()); + Assert.Equal("detailed", request.GetProperty("reasoningSummary").GetString()); + Assert.Equal("long_context", request.GetProperty("contextTier").GetString()); + Assert.True(request.GetProperty("streaming").GetBoolean()); + Assert.False(request.GetProperty("includeSubAgentStreamingEvents").GetBoolean()); + Assert.Equal("SCENARIO_COMPOSED_SYSTEM_MESSAGE", request.GetProperty("systemMessage").GetProperty("content").GetString()); + Assert.True(request.GetProperty("enableConfigDiscovery").GetBoolean()); + Assert.False(request.GetProperty("enableSessionTelemetry").GetBoolean()); + Assert.True(request.GetProperty("isExperimentalMode").GetBoolean()); + Assert.True(request.GetProperty("customAgentsLocalOnly").GetBoolean()); + Assert.True(request.GetProperty("skipEmbeddingRetrieval").GetBoolean()); + Assert.Equal("in-memory", request.GetProperty("embeddingCacheStorage").GetString()); + Assert.Equal("SCENARIO_ORG_INSTRUCTIONS", request.GetProperty("organizationCustomInstructions").GetString()); + Assert.False(request.GetProperty("enableOnDemandInstructionDiscovery").GetBoolean()); + Assert.False(request.GetProperty("enableFileHooks").GetBoolean()); + Assert.False(request.GetProperty("enableHostGitOperations").GetBoolean()); + Assert.False(request.GetProperty("enableSessionStore").GetBoolean()); + Assert.False(request.GetProperty("enableSkills").GetBoolean()); + Assert.Equal("scenario_tool", request.GetProperty("availableTools")[0].GetString()); + Assert.Equal("shell", request.GetProperty("excludedTools")[0].GetString()); + Assert.Equal("scenario_tool", request.GetProperty("tools")[0].GetProperty("name").GetString()); + Assert.Equal("scenario-command", request.GetProperty("commands")[0].GetProperty("name").GetString()); + Assert.Equal("node", request.GetProperty("mcpServers").GetProperty("scenario-mcp").GetProperty("command").GetString()); + Assert.Equal("scenario-agent", request.GetProperty("customAgents")[0].GetProperty("name").GetString()); + Assert.Equal("scenario-agent", request.GetProperty("agent").GetString()); + Assert.Equal("edit", request.GetProperty("defaultAgent").GetProperty("excludedTools")[0].GetString()); + Assert.Equal("scenario-provider", request.GetProperty("providers")[0].GetProperty("name").GetString()); + Assert.True(request.GetProperty("providers")[0].GetProperty("hasBearerTokenProvider").GetBoolean()); + Assert.Equal("scenario-model", request.GetProperty("models")[0].GetProperty("id").GetString()); + Assert.Equal("export", request.GetProperty("remoteSession").GetString()); + Assert.True(request.GetProperty("requestMcpApps").GetBoolean()); + Assert.False(request.GetProperty("githubMcpToolConfig").GetProperty("enableAllTools").GetBoolean()); + Assert.True(request.GetProperty("githubMcpToolConfig").GetProperty("disableFormDeferral").GetBoolean()); + Assert.True(request.GetProperty("requestCanvasRenderer").GetBoolean()); + Assert.True(request.GetProperty("requestExtensions").GetBoolean()); + Assert.Equal("scenario-extension-sdk", request.GetProperty("extensionSdkPath").GetString()); + Assert.Equal("desktop", request.GetProperty("extensionInfo").GetProperty("name").GetString()); + Assert.Equal("scenario:builtin:desktop", request.GetProperty("canvasProvider").GetProperty("id").GetString()); + Assert.Equal("scenario-canvas", request.GetProperty("canvases")[0].GetProperty("id").GetString()); + Assert.True(request.GetProperty("requestPermission").GetBoolean()); + Assert.True(request.GetProperty("requestUserInput").GetBoolean()); + Assert.True(request.GetProperty("requestElicitation").GetBoolean()); + Assert.True(request.GetProperty("requestExitPlanMode").GetBoolean()); + Assert.True(request.GetProperty("requestAutoModeSwitch").GetBoolean()); + Assert.Equal(sessionId, optionsUpdate.GetProperty("sessionId").GetString()); + Assert.False(optionsUpdate.GetProperty("skipCustomInstructions").GetBoolean()); + Assert.True(optionsUpdate.GetProperty("customAgentsLocalOnly").GetBoolean()); + Assert.False(optionsUpdate.GetProperty("coauthorEnabled").GetBoolean()); + Assert.False(optionsUpdate.GetProperty("manageScheduleEnabled").GetBoolean()); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Preserve_Omitted_Versus_Disabled_Scenario_Semantics() + { + var (cliPath, capturePath) = await CreateFakeRuntimeAsync("capture"); + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "capture"]), + UseLoggedInUser = false, + }); + + await using var sparse = await Ctx.CreateSessionAsync(client, new SessionConfig + { + SessionId = "scenario-sparse", + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + await using var disabled = await Ctx.CreateSessionAsync(client, new SessionConfig + { + SessionId = "scenario-disabled", + EnableSessionTelemetry = false, + EnableExperimentalMode = false, + SkipCustomInstructions = false, + CustomAgentsLocalOnly = false, + CoauthorEnabled = false, + ManageScheduleEnabled = false, + EnableConfigDiscovery = false, + SkipEmbeddingRetrieval = false, + EnableOnDemandInstructionDiscovery = false, + EnableFileHooks = false, + EnableHostGitOperations = false, + EnableSessionStore = false, + EnableSkills = false, + RequestCanvasRenderer = false, + RequestExtensions = false, + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => GetRequests(root, "session.create").Count == 2 + && GetRequests(root, "session.options.update").Count == 1); + var requests = GetRequests(capture.RootElement, "session.create") + .Select(item => item.GetProperty("params")) + .ToDictionary(item => item.GetProperty("sessionId").GetString()!, StringComparer.Ordinal); + var sparseRequest = requests["scenario-sparse"]; + var disabledRequest = requests["scenario-disabled"]; + + string[] fields = + [ + "enableSessionTelemetry", + "isExperimentalMode", + "customAgentsLocalOnly", + "enableConfigDiscovery", + "skipEmbeddingRetrieval", + "enableOnDemandInstructionDiscovery", + "enableFileHooks", + "enableHostGitOperations", + "enableSessionStore", + "enableSkills", + "requestCanvasRenderer", + "requestExtensions", + ]; + + Assert.All(fields, field => Assert.False(sparseRequest.TryGetProperty(field, out _))); + Assert.All(fields, field => Assert.False(disabledRequest.GetProperty(field).GetBoolean())); + + var optionsUpdate = Assert.Single(GetRequests(capture.RootElement, "session.options.update")) + .GetProperty("params"); + Assert.Equal("scenario-disabled", optionsUpdate.GetProperty("sessionId").GetString()); + Assert.False(optionsUpdate.GetProperty("skipCustomInstructions").GetBoolean()); + Assert.False(optionsUpdate.GetProperty("customAgentsLocalOnly").GetBoolean()); + Assert.False(optionsUpdate.GetProperty("coauthorEnabled").GetBoolean()); + Assert.False(optionsUpdate.GetProperty("manageScheduleEnabled").GetBoolean()); + } + + [Fact] + public async Task Should_Use_Preallocated_Id_For_First_Subscribed_Event() + { + var requestedSessionId = Guid.NewGuid().ToString(); + var firstEvent = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + SessionId = requestedSessionId, + OnEvent = evt => firstEvent.TrySetResult(evt), + }); + + var observed = await firstEvent.Task.WaitAsync(TestTimeout); + var start = Assert.IsType(observed); + Assert.Equal(requestedSessionId, session.SessionId); + Assert.Equal(requestedSessionId, start.Data.SessionId); + } + + [Fact] + [Trait(E2ETestTraits.Backend, E2ETestTraits.SelfConfiguredBackend)] + public async Task Should_Invoke_All_Scenario_Handler_Kinds() + { + var (cliPath, capturePath) = await CreateFakeRuntimeAsync("callbacks"); + var observed = new ConcurrentDictionary(StringComparer.Ordinal); + var allObserved = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + string[] expected = + [ + "event", + "permission", + "user-input", + "elicitation", + "exit-plan", + "auto-mode", + "mcp-auth", + "tool", + "command", + "canvas", + "provider-token", + ]; + + void Mark(string name) + { + observed.TryAdd(name, 0); + if (expected.All(observed.ContainsKey)) + { + allObserved.TrySetResult(); + } + } + + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio( + path: cliPath, + args: ["--capture-file", capturePath, "--behavior", "callbacks"]), + UseLoggedInUser = false, + }); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + SessionId = "scenario-handler-session", + Tools = [AIFunctionFactory.Create(() => { Mark("tool"); return "tool-result"; }, "scenario_tool")], + Commands = + [ + new CommandDefinition + { + Name = "scenario-command", + Handler = _ => + { + Mark("command"); + return Task.CompletedTask; + }, + }, + ], + Providers = + [ + new NamedProviderConfig + { + Name = "scenario-provider", + Type = "openai", + BaseUrl = "https://provider.example.test/v1", + BearerTokenProvider = _ => + { + Mark("provider-token"); + return Task.FromResult("provider-token"); + }, + }, + ], + Models = + [ + new ProviderModelConfig + { + Provider = "scenario-provider", + Id = "scenario-model", + ModelId = "claude-sonnet-5", + }, + ], + Canvases = [new CanvasDeclaration { Id = "scenario-canvas", DisplayName = "Scenario Canvas" }], + CanvasProvider = new CanvasProviderIdentity { Id = "scenario:builtin:desktop", Name = "scenario client" }, + CanvasHandler = new CallbackCanvasHandler(() => Mark("canvas")), + OnPermissionRequest = (_, _) => + { + Mark("permission"); + return Task.FromResult(PermissionDecision.ApproveOnce()); + }, + OnUserInputRequest = (_, _) => + { + Mark("user-input"); + return Task.FromResult(new UserInputResponse { Answer = "approved" }); + }, + OnElicitationRequest = _ => + { + Mark("elicitation"); + return Task.FromResult(new ElicitationResult + { + Action = UIElicitationResponseAction.Accept, + Content = new Dictionary { ["value"] = "accepted" }, + }); + }, + OnExitPlanModeRequest = (_, _) => + { + Mark("exit-plan"); + return Task.FromResult(new ExitPlanModeResult + { + Approved = true, + SelectedAction = "interactive", + }); + }, + OnAutoModeSwitchRequest = (_, _) => + { + Mark("auto-mode"); + return Task.FromResult(AutoModeSwitchResponse.No); + }, + OnMcpAuthRequest = _ => + { + Mark("mcp-auth"); + return Task.FromResult(McpAuthResult.Cancel()); + }, + OnEvent = evt => + { + if (evt is SessionInfoEvent { Data.Message: "SCENARIO_HANDLER_EVENT" }) + { + Mark("event"); + } + }, + }); + + await allObserved.Task.WaitAsync(TestTimeout); + Assert.Equal( + expected.OrderBy(value => value, StringComparer.Ordinal), + observed.Keys.OrderBy(value => value, StringComparer.Ordinal)); + + using var capture = await WaitForCaptureAsync( + capturePath, + root => root.GetProperty("clientResponses").GetArrayLength() == 5 + && GetRequests(root, "session.permissions.handlePendingPermissionRequest").Count == 1 + && GetRequests(root, "session.ui.handlePendingElicitation").Count == 1 + && GetRequests(root, "session.mcp.oauth.handlePendingRequest").Count == 1 + && GetRequests(root, "session.tools.handlePendingToolCall").Count == 1 + && GetRequests(root, "session.commands.handlePendingCommand").Count == 1); + var responses = capture.RootElement.GetProperty("clientResponses") + .EnumerateArray() + .ToDictionary(item => item.GetProperty("id").GetInt32()); + + Assert.Equal("approved", responses[1000].GetProperty("result").GetProperty("answer").GetString()); + Assert.False(responses[1000].GetProperty("result").GetProperty("wasFreeform").GetBoolean()); + Assert.True(responses[1001].GetProperty("result").GetProperty("approved").GetBoolean()); + Assert.Equal("interactive", responses[1001].GetProperty("result").GetProperty("selectedAction").GetString()); + Assert.Equal("no", responses[1002].GetProperty("result").GetProperty("response").GetString()); + Assert.Equal("ready", responses[1003].GetProperty("result").GetProperty("status").GetString()); + Assert.Equal("Scenario Canvas", responses[1003].GetProperty("result").GetProperty("title").GetString()); + Assert.Equal("provider-token", responses[1004].GetProperty("result").GetProperty("token").GetString()); + + Assert.Equal( + ["permission-1", "elicitation-1", "mcp-auth-1", "tool-1", "command-1"], + new[] + { + GetRequests(capture.RootElement, "session.permissions.handlePendingPermissionRequest").Single(), + GetRequests(capture.RootElement, "session.ui.handlePendingElicitation").Single(), + GetRequests(capture.RootElement, "session.mcp.oauth.handlePendingRequest").Single(), + GetRequests(capture.RootElement, "session.tools.handlePendingToolCall").Single(), + GetRequests(capture.RootElement, "session.commands.handlePendingCommand").Single(), + }.Select(item => item.GetProperty("params").GetProperty("requestId").GetString())); + } + + [Fact] + public async Task Should_Create_Then_Reload_Mcp_In_Order() + { + const string ServerName = "scenario-client-reload"; + var milestones = new List(); + var milestonesLock = new object(); + var startObserved = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + McpServers = CreateTestMcpServers(ServerName), + OnEvent = evt => + { + if (evt is SessionStartEvent) + { + lock (milestonesLock) + { + milestones.Add("session-start"); + } + startObserved.TrySetResult(); + } + }, + }); + + await startObserved.Task.WaitAsync(TestTimeout); + lock (milestonesLock) + { + milestones.Add("create-returned"); + milestones.Add("reload-requested"); + } + + await session.Rpc.Mcp.ReloadAsync(); + await WaitForMcpServerStatusAsync(session, ServerName, McpServerStatus.Connected); + + lock (milestonesLock) + { + milestones.Add("reload-completed"); + Assert.Equal( + ["session-start", "create-returned", "reload-requested", "reload-completed"], + milestones); + } + } + + private async Task<(string CliPath, string CapturePath)> CreateFakeRuntimeAsync(string behavior) + { + var cliPath = Path.Join(Ctx.WorkDir, $"scenario-client-session-{behavior}-{Guid.NewGuid():N}.js"); + var capturePath = Path.Join(Ctx.WorkDir, $"scenario-client-session-{behavior}-{Guid.NewGuid():N}.json"); + await File.WriteAllTextAsync(cliPath, FakeRuntimeScript); + return (cliPath, capturePath); + } + + private static List GetRequests(JsonElement root, string method) => + root.GetProperty("requests") + .EnumerateArray() + .Where(item => item.GetProperty("method").GetString() == method) + .ToList(); + + private static async Task WaitForCaptureAsync( + string path, + Func predicate) + { + JsonDocument? result = null; + await TestHelper.WaitForConditionAsync( + async () => + { + try + { + using var stream = new FileStream( + path, + FileMode.Open, + FileAccess.Read, + FileShare.ReadWrite | FileShare.Delete); + using var reader = new StreamReader(stream); + var document = JsonDocument.Parse(await reader.ReadToEndAsync()); + if (!predicate(document.RootElement)) + { + document.Dispose(); + return false; + } + + result = document; + return true; + } + catch (Exception ex) when (ex is IOException or JsonException) + { + return false; + } + }, + timeout: TestTimeout, + pollInterval: TimeSpan.FromMilliseconds(50), + timeoutMessage: $"Timed out waiting for fake runtime capture at {path}."); + return result!; + } + + private sealed class NoOpCanvasHandler : CanvasHandlerBase + { + public override Task OnOpenAsync( + CanvasProviderOpenRequest request, + CancellationToken cancellationToken) => + Task.FromResult(new CanvasProviderOpenResult { Status = "ready" }); + } + + private sealed class CallbackCanvasHandler(Action callback) : CanvasHandlerBase + { + public override Task OnOpenAsync( + CanvasProviderOpenRequest request, + CancellationToken cancellationToken) + { + callback(); + return Task.FromResult(new CanvasProviderOpenResult { Status = "ready", Title = "Scenario Canvas" }); + } + } + + private const string FakeRuntimeScript = """ + const fs = require("fs"); + + function argument(name) { + const index = process.argv.indexOf(name); + return index >= 0 ? process.argv[index + 1] : undefined; + } + + const captureFile = argument("--capture-file"); + const behavior = argument("--behavior") || "capture"; + const requests = []; + const clientResponses = []; + let nextRequestId = 1000; + let callbackSessionId; + let callbacksStarted = false; + let buffer = Buffer.alloc(0); + + function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ requests, clientResponses })); + } + + function write(message) { + const body = JSON.stringify(message); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function respond(id, result) { + write({ jsonrpc: "2.0", id, result }); + } + + function request(method, params) { + write({ jsonrpc: "2.0", id: nextRequestId++, method, params }); + } + + function notify(method, params) { + write({ jsonrpc: "2.0", method, params }); + } + + function event(type, data, ordinal) { + return { + id: `00000000-0000-0000-0000-${String(ordinal).padStart(12, "0")}`, + timestamp: "2026-09-17T20:00:00Z", + parentId: null, + type, + data + }; + } + + function fireCallbacks() { + if (callbacksStarted || behavior !== "callbacks") return; + callbacksStarted = true; + const sessionId = callbackSessionId; + + request("userInput.request", { + sessionId, + question: "Continue?", + choices: ["approved", "declined"], + allowFreeform: false + }); + request("exitPlanMode.request", { + sessionId, + summary: "Scenario plan", + planContent: "# Scenario plan", + actions: ["interactive", "exit_only"], + recommendedAction: "interactive" + }); + request("autoModeSwitch.request", { + sessionId, + errorCode: "scenario-rate-limit", + retryAfterSeconds: 1 + }); + request("canvas.open", { + sessionId, + canvasId: "scenario-canvas", + extensionId: "scenario:builtin:desktop", + instanceId: "scenario-canvas-1", + input: { start: 1 } + }); + request("providerToken.getToken", { + sessionId, + providerName: "scenario-provider" + }); + + notify("session.event", { + sessionId, + event: event("session.info", { + infoType: "notification", + message: "SCENARIO_HANDLER_EVENT" + }, 1) + }); + notify("session.event", { + sessionId, + event: event("permission.requested", { + requestId: "permission-1", + permissionRequest: { + kind: "read", + intention: "Read the scenario README", + path: "README.md" + } + }, 2) + }); + notify("session.event", { + sessionId, + event: event("elicitation.requested", { + requestId: "elicitation-1", + message: "Provide a value", + mode: "form", + requestedSchema: { + type: "object", + properties: { value: { type: "string" } }, + required: ["value"] + } + }, 3) + }); + notify("session.event", { + sessionId, + event: event("mcp.oauth_required", { + requestId: "mcp-auth-1", + reason: "initial", + serverName: "scenario-mcp", + serverUrl: "https://example.test/mcp" + }, 4) + }); + notify("session.event", { + sessionId, + event: event("external_tool.requested", { + requestId: "tool-1", + sessionId, + toolCallId: "tool-call-1", + toolName: "scenario_tool", + arguments: {} + }, 5) + }); + notify("session.event", { + sessionId, + event: event("command.execute", { + requestId: "command-1", + commandName: "scenario-command", + command: "/scenario-command value", + args: "value" + }, 6) + }); + } + + function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "method")) { + clientResponses.push(message); + saveCapture(); + return; + } + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + respond(message.id, { ok: true, protocolVersion: 3, version: "fake" }); + return; + } + + if (message.method === "session.create") { + callbackSessionId = message.params?.sessionId ?? "fake-session"; + respond(message.id, { + sessionId: callbackSessionId, + workspacePath: null, + capabilities: { ui: { elicitation: true } } + }); + return; + } + + if (message.method === "session.eventLog.registerInterest") { + respond(message.id, { handle: "scenario-handler-interest" }); + setTimeout(fireCallbacks, 10); + return; + } + + if (message.method === "session.detach") { + respond(message.id, { success: true }); + return; + } + + if (message.method === "runtime.shutdown") { + respond(message.id, {}); + setTimeout(() => process.exit(0), 10); + return; + } + + if (message.method === "ping") { + respond(message.id, { + message: "pong", + timestamp: new Date().toISOString(), + protocolVersion: 3 + }); + return; + } + + respond(message.id, { success: true }); + } + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } + }); + + process.stdin.resume(); + saveCapture(); + setInterval(() => {}, 1000); + """; +} + +#pragma warning restore GHCP001 diff --git a/dotnet/test/E2E/ScenarioTestingSkillsAndAgentsE2ETests.cs b/dotnet/test/E2E/ScenarioTestingSkillsAndAgentsE2ETests.cs new file mode 100644 index 0000000000..d9090221ce --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingSkillsAndAgentsE2ETests.cs @@ -0,0 +1,208 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using System.Text.Json; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingSkillsAndAgentsE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_skills_and_agents", output) +{ + [Fact] + public async Task Should_Reload_Atomically_Replaced_Skill_And_Replay_It_On_Resume() + { + const string skillName = "scenario-reloadable-skill"; + var skillsDirectory = Path.Join(Ctx.WorkDir, "scenario-skills", Guid.NewGuid().ToString("N")); + var skillFile = WriteSkill( + skillsDirectory, + skillName, + "Scenario skill version one.", + "Use SCENARIO_SKILL_VERSION_ONE."); + + await using var session1 = await CreateSessionAsync(new SessionConfig + { + SkillDirectories = [skillsDirectory], + }); + + AssertSkill( + await session1.Rpc.Skills.ListAsync(), + skillName, + "Scenario skill version one.", + skillFile); + + var replacement = Path.Join(Path.GetDirectoryName(skillFile)!, "SKILL.replacement.md"); + File.WriteAllText( + replacement, + CreateSkillContent( + skillName, + "Scenario skill version two.", + "Use SCENARIO_SKILL_VERSION_TWO.")); + File.Replace(replacement, skillFile, destinationBackupFileName: null); + + await session1.Rpc.Skills.ReloadAsync(); + AssertSkill( + await session1.Rpc.Skills.ListAsync(), + skillName, + "Scenario skill version two.", + skillFile); + + var sessionId = session1.SessionId; + await SuspendAndUntrackSessionForResumeAsync(session1); + + await using var session2 = await ResumeSessionAsync(sessionId, new ResumeSessionConfig + { + ContinuePendingWork = false, + SkillDirectories = [skillsDirectory], + }); + + AssertSkill( + await session2.Rpc.Skills.ListAsync(), + skillName, + "Scenario skill version two.", + skillFile); + } + + [Fact] + public async Task Should_Classify_Agent_Method_Not_Found_As_Remote_Protocol_Error() + { + var scriptPath = Path.Join( + Path.GetTempPath(), + $"copilot-agent-method-not-found-{Guid.NewGuid():N}.cjs"); + File.WriteAllText(scriptPath, FakeAgentMethodNotFoundCliScript); + + try + { + await using var client = Ctx.CreateClient(options: new CopilotClientOptions + { + Connection = RuntimeConnection.ForStdio(path: "node", args: [scriptPath]), + }); + await using var session = await Ctx.CreateSessionAsync(client, new SessionConfig + { + OnPermissionRequest = PermissionHandler.ApproveAll, + }); + + var exception = await Assert.ThrowsAsync(() => session.Rpc.Agent.ReloadAsync()); + + Assert.Contains("Method not found: session.agent.reload", exception.Message, StringComparison.Ordinal); + Assert.NotNull(exception.InnerException); + Assert.Equal("RemoteRpcException", exception.InnerException!.GetType().Name); + var errorCode = exception.InnerException.GetType().GetProperty("ErrorCode")!.GetValue(exception.InnerException); + Assert.Equal(-32601, Assert.IsType(errorCode)); + Assert.DoesNotContain("Unhandled method", exception.ToString(), StringComparison.OrdinalIgnoreCase); + } + finally + { + File.Delete(scriptPath); + } + } + + private static string WriteSkill( + string skillsDirectory, + string skillName, + string description, + string body) + { + var skillDirectory = Path.Join(skillsDirectory, skillName); + Directory.CreateDirectory(skillDirectory); + var skillFile = Path.Join(skillDirectory, "SKILL.md"); + File.WriteAllText(skillFile, CreateSkillContent(skillName, description, body)); + return skillFile; + } + + private static string CreateSkillContent(string skillName, string description, string body) => + $""" + --- + name: {skillName} + description: {description} + --- + + # Scenario Reloadable Skill + + {body} + """.ReplaceLineEndings("\n"); + + private static void AssertSkill( + SkillList list, + string skillName, + string description, + string expectedPath) + { + var skill = Assert.Single( + list.Skills, + skill => string.Equals(skill.Name, skillName, StringComparison.Ordinal)); + Assert.True(skill.Enabled); + Assert.Equal(description, skill.Description); + Assert.Equal(expectedPath, skill.Path); + } + + private const string FakeAgentMethodNotFoundCliScript = """ + let buffer = Buffer.alloc(0); + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); + }); + process.stdin.resume(); + + function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length header"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handleMessage(JSON.parse(body)); + } + } + + function handleMessage(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) return; + if (message.method === "connect") { + writeResult(message.id, { ok: true, protocolVersion: 3, version: "fake" }); + return; + } + if (message.method === "ping") { + writeResult(message.id, { message: "pong", protocolVersion: 3 }); + return; + } + if (message.method === "session.create") { + const params = Array.isArray(message.params) ? message.params[0] : message.params; + writeResult(message.id, { + sessionId: params?.sessionId ?? "fake-agent-session", + workspacePath: null, + capabilities: null, + openCanvases: [] + }); + return; + } + if (message.method === "session.agent.reload") { + writeError(message.id, -32601, "Method not found: session.agent.reload"); + return; + } + writeResult(message.id, {}); + } + + function writeResult(id, result) { + writeMessage({ jsonrpc: "2.0", id, result }); + } + + function writeError(id, code, message) { + writeMessage({ jsonrpc: "2.0", id, error: { code, message } }); + } + + function writeMessage(message) { + const body = JSON.stringify(message); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + """; +} diff --git a/dotnet/test/E2E/ScenarioTestingTestCli.cs b/dotnet/test/E2E/ScenarioTestingTestCli.cs new file mode 100644 index 0000000000..12c6bd8e7a --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingTestCli.cs @@ -0,0 +1,424 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using System.Text.Json; + +namespace GitHub.Copilot.Test.E2E; + +internal static class ScenarioTestingTestCli +{ + public static async Task<(string CliPath, string CapturePath)> CreateAsync(E2ETestContext context) + { + var cliPath = Path.Join(context.WorkDir, $"scenario-client-test-cli-{Guid.NewGuid():N}.js"); + var capturePath = Path.Join(context.WorkDir, $"scenario-client-test-cli-{Guid.NewGuid():N}.json"); + await File.WriteAllTextAsync(cliPath, Script); + return (cliPath, capturePath); + } + + public static async Task ReadRequestsAsync(string capturePath) + { + await TestHelper.WaitForConditionAsync( + () => Task.FromResult(File.Exists(capturePath)), + timeout: TimeSpan.FromSeconds(10), + timeoutMessage: "Timed out waiting for the fake CLI request capture."); + + using var capture = JsonDocument.Parse(await File.ReadAllTextAsync(capturePath)); + return capture.RootElement.GetProperty("requests").EnumerateArray().Select(request => request.Clone()).ToArray(); + } + + private const string Script = """ + const fs = require("fs"); + + const captureIndex = process.argv.indexOf("--capture-file"); + const behaviorIndex = process.argv.indexOf("--behavior"); + const captureFile = process.argv[captureIndex + 1]; + const behavior = process.argv[behaviorIndex + 1]; + const requests = []; + let resumeAttempts = 0; + let nextQueueId = 1; + let queueItems = []; + let buffer = Buffer.alloc(0); + + function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ requests })); + } + + function writeResponse(id, result) { + const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function writeError(id, code, message) { + const body = JSON.stringify({ jsonrpc: "2.0", id, error: { code, message } }); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function writeSessionEvent(sessionId, type, data) { + const body = JSON.stringify({ + jsonrpc: "2.0", + method: "session.event", + params: { + sessionId, + event: { + id: "00000000-0000-0000-0000-" + String(requests.length).padStart(12, "0"), + timestamp: "2026-09-17T20:00:00.000Z", + parentId: null, + type, + data + } + } + }); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); + } + + function getParams(message) { + return Array.isArray(message.params) ? (message.params[0] ?? {}) : (message.params ?? {}); + } + + function handleMessage(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) { + return; + } + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + writeResponse(message.id, { ok: true, protocolVersion: 3, version: "scenario-client-test" }); + return; + } + + if (message.method === "session.create") { + const sessionId = behavior === "cloud-assigned-event" + ? "server-assigned-cloud-session" + : getParams(message).sessionId ?? "scenario-client-session"; + writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); + if (behavior === "cloud-assigned-event") { + writeSessionEvent(sessionId, "session.start", { + sessionId, + copilotVersion: "scenario-runtime", + producer: "scenario-test-cli", + startTime: "2026-09-17T20:00:00.000Z", + version: 1 + }); + } + if (behavior === "emit-ui-events") { + setTimeout(() => { + writeSessionEvent(sessionId, "user_input.requested", { + requestId: "scenario-user-input", + question: "Choose a scenario action", + choices: ["Approve", "Decline"], + allowFreeform: true, + toolCallId: "tool-user-input" + }); + writeSessionEvent(sessionId, "elicitation.requested", { + requestId: "scenario-form-accept", + message: "Provide scenario settings", + mode: "form", + requestedSchema: { + type: "object", + properties: { name: { type: "string" } }, + required: ["name"] + }, + toolCallId: "tool-form" + }); + writeSessionEvent(sessionId, "elicitation.requested", { + requestId: "scenario-url-decline", + message: "Authorize the scenario", + mode: "url", + url: "https://example.test/authorize", + toolCallId: "tool-url" + }); + writeSessionEvent(sessionId, "elicitation.requested", { + requestId: "scenario-form-cancel", + message: "Optional scenario settings", + mode: "form", + requestedSchema: { + type: "object", + properties: {}, + required: [] + }, + toolCallId: "tool-cancel" + }); + }, 10); + } + return; + } + + if (message.method === "session.resume") { + resumeAttempts++; + if (behavior === "resume-not-found-once" && resumeAttempts === 1) { + writeError(message.id, -32001, "Session not found"); + return; + } + + const sessionId = getParams(message).sessionId ?? "scenario-client-session"; + writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); + return; + } + + if (message.method === "catalog.search") { + const params = getParams(message); + writeResponse(message.id, { + kind: "succeeded", + candidates: [], + negotiated: { + runtimeProtocolVersion: params.contract.protocolVersion, + grantedCapabilities: params.contract.requiredCapabilities + }, + searchId: "scenario-search", + truncated: false + }); + return; + } + + if (message.method === "session.autopilotObjective.getState") { + writeResponse(message.id, { + state: { + id: 17, + objective: "Ship the scenario.", + status: "active", + turnCount: 3, + creditCountNanoAiu: "1250000000", + creditLimit: { + credits: 5, + creditsUsed: 1.25, + creditsUsedNanoAiu: "1250000000" + } + } + }); + return; + } + + if (message.method === "session.remote.enable") { + const params = getParams(message); + writeResponse(message.id, { + url: `https://example.test/sessions/${params.sessionId}`, + remoteSteerable: params.mode === "on" + }); + return; + } + + if (message.method === "session.factory.listRuns") { + writeResponse(message.id, { + runs: [{ + runId: "factory-run-1", + factoryName: "scenario-factory", + description: "Scenario factory", + status: "running", + revision: 4, + createdAt: 1000, + updatedAt: 2000, + observedAt: 2100, + canResume: false, + declaredLimits: {}, + consumed: {} + }], + oldestSeq: 7, + newestSeq: 7, + hasMoreNewer: false, + omittedOlder: 0 + }); + return; + } + + if (message.method === "session.factory.getRunDetail") { + writeResponse(message.id, { + runId: "factory-run-1", + factoryName: "scenario-factory", + description: "Scenario factory", + status: "running", + revision: 4, + createdAt: 1000, + updatedAt: 2000, + observedAt: 2100, + canResume: false, + declaredLimits: {}, + consumed: {}, + progress: { + records: [], + revision: 4, + hasMoreOlder: false, + hasMoreNewer: false + } + }); + return; + } + + if (message.method === "session.factory.getRunProgress") { + writeResponse(message.id, { + records: [{ + seq: 12, + attempt: 1, + phaseId: "verify", + kind: "log", + text: "Validation complete", + recordedAt: 2000 + }], + oldestSeq: 12, + newestSeq: 12, + revision: 4, + hasMoreOlder: false, + hasMoreNewer: false + }); + return; + } + + if (message.method === "session.factory.cancel") { + writeResponse(message.id, { + runId: "factory-run-1", + status: "cancelled", + reason: "cancelled by user", + attempt: 1 + }); + return; + } + + if (message.method === "session.queue.setDrainPaused") { + writeResponse(message.id, {}); + return; + } + + if (message.method === "session.queue.insertAt") { + const params = getParams(message); + const id = `queue-${nextQueueId++}`; + const item = { + id, + messageId: `message-${id}`, + kind: "message", + displayText: params.message.displayPrompt ?? params.message.prompt, + prompt: params.message.prompt, + agentMode: params.message.agentMode ?? "interactive" + }; + const position = Math.max(0, Math.min(Number(params.position), queueItems.length)); + queueItems.splice(position, 0, item); + writeResponse(message.id, { id }); + return; + } + + if (message.method === "session.queue.pendingItems") { + writeResponse(message.id, { + items: queueItems.map(({ prompt, ...item }) => item), + steeringMessages: [], + inFlightSteeringCount: 0 + }); + return; + } + + if (message.method === "session.queue.updateText") { + const params = getParams(message); + const item = queueItems.find(candidate => candidate.id === params.id); + if (item) { + item.prompt = params.prompt; + item.displayText = params.displayPrompt ?? params.prompt; + } + writeResponse(message.id, { updated: Boolean(item) }); + return; + } + + if (message.method === "session.queue.duplicateAt") { + const params = getParams(message); + const index = queueItems.findIndex(candidate => candidate.id === params.id); + const id = `queue-${nextQueueId++}`; + if (index >= 0) { + queueItems.splice(index + 1, 0, { + ...queueItems[index], + id, + messageId: `message-${id}` + }); + } + writeResponse(message.id, { id }); + return; + } + + if (message.method === "session.queue.moveItem") { + const params = getParams(message); + const index = queueItems.findIndex(candidate => candidate.id === params.id); + if (index < 0) { + writeResponse(message.id, { changed: false }); + return; + } + const [item] = queueItems.splice(index, 1); + const target = Math.max(0, Math.min(Number(params.toPosition), queueItems.length)); + queueItems.splice(target, 0, item); + writeResponse(message.id, { changed: index !== target }); + return; + } + + if (message.method === "session.queue.removeAt") { + const params = getParams(message); + const index = queueItems.findIndex(candidate => candidate.id === params.id); + if (index >= 0) { + queueItems.splice(index, 1); + } + writeResponse(message.id, { removed: index >= 0 }); + return; + } + + if (message.method === "session.queue.sendNow") { + const params = getParams(message); + const index = queueItems.findIndex(candidate => candidate.id === params.id); + if (index >= 0) { + queueItems.splice(index, 1); + } + writeResponse(message.id, { steered: index >= 0 }); + return; + } + + if (message.method === "session.send" && behavior === "drop-after-send") { + process.stdout.end(); + return; + } + + if (message.method === "session.send") { + writeResponse(message.id, { messageId: "scenario-client-message" }); + return; + } + + if (message.method === "session.delete" && behavior === "delete-not-found") { + writeResponse(message.id, { success: false, error: "Session file not found" }); + return; + } + + if (message.method === "session.ui.handlePendingElicitation" || + message.method === "session.ui.handlePendingUserInput") { + const requestId = message.params?.requestId ?? message.params?.[0]?.requestId; + writeResponse(message.id, { success: requestId !== "stale-scenario-request" }); + return; + } + + writeResponse(message.id, { success: true }); + } + + process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) { + return; + } + + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) { + throw new Error("Missing Content-Length header"); + } + + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + Number(match[1]); + if (buffer.length < bodyEnd) { + return; + } + + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handleMessage(JSON.parse(body)); + } + }); + + process.stdin.resume(); + saveCapture(); + """; +} diff --git a/dotnet/test/E2E/ScenarioTestingToolsE2ETests.cs b/dotnet/test/E2E/ScenarioTestingToolsE2ETests.cs new file mode 100644 index 0000000000..ae817fac0c --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingToolsE2ETests.cs @@ -0,0 +1,256 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using Microsoft.Extensions.AI; +using System.ComponentModel; +using System.Text.Json; +using System.Text.Json.Nodes; +using System.Text.Json.Serialization; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +/// +/// Representative scenario coverage for host-owned tools. +/// +public partial class ScenarioTestingToolsE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_tools", output) +{ + private static readonly TimeSpan EventTimeout = TimeSpan.FromSeconds(60); + + [JsonSourceGenerationOptions(JsonSerializerDefaults.Web)] + [JsonSerializable(typeof(ToolResultAIContent))] + [JsonSerializable(typeof(ToolResultObject))] + [JsonSerializable(typeof(JsonElement))] + private partial class ScenarioToolsJsonContext : JsonSerializerContext; + + [Fact] + public async Task Should_Advertise_Scenario_Tool_Schema_Override_And_Availability() + { + var hiddenToolCalled = false; + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Tools = + [ + CopilotTool.DefineTool( + (Func)LookupIssue, + factoryOptions: new AIFunctionFactoryOptions + { + Name = "scenario_lookup_issue", + Description = "Looks up an issue in the scenario client installation.", + }), + CopilotTool.DefineTool( + (Func)ScenarioGrep, + new CopilotToolOptions { OverridesBuiltInTool = true }, + new AIFunctionFactoryOptions + { + Name = "grep", + Description = "Searches the scenario-owned index.", + }), + CopilotTool.DefineTool( + (Func)HiddenAdminTool, + factoryOptions: new AIFunctionFactoryOptions { Name = "scenario_hidden_admin" }), + ], + AvailableTools = new ToolSet() + .AddCustom("scenario_lookup_issue") + .AddCustom("grep"), + ExcludedTools = new ToolSet().AddCustom("scenario_hidden_admin"), + }); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_lookup_issue for owner octo and issue number 42. Reply with its result.", + }); + + Assert.Contains("SCENARIO_ISSUE_octo_42", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + Assert.False(hiddenToolCalled); + + var exchange = (await Ctx.GetExchangesAsync()).Last(); + var names = GetToolNames(exchange); + Assert.Contains("scenario_lookup_issue", names); + Assert.Contains("grep", names); + Assert.DoesNotContain("scenario_hidden_admin", names); + Assert.Equal(1, names.Count(name => name == "grep")); + + var lookup = Assert.Single(exchange.Request.Tools!, tool => tool.Function.Name == "scenario_lookup_issue"); + Assert.Equal("Looks up an issue in the scenario client installation.", lookup.Function.Description); + var parameters = lookup.Function.Parameters!.Value; + Assert.Equal("object", parameters.GetProperty("type").GetString()); + Assert.Equal("string", parameters.GetProperty("properties").GetProperty("owner").GetProperty("type").GetString()); + Assert.Equal("integer", parameters.GetProperty("properties").GetProperty("number").GetProperty("type").GetString()); + + static string LookupIssue( + [Description("Repository owner")] string owner, + [Description("Issue number")] int number) => + $"SCENARIO_ISSUE_{owner}_{number}"; + + static string ScenarioGrep([Description("Search query")] string query) => $"SCENARIO_GREP_{query}"; + + string HiddenAdminTool() + { + hiddenToolCalled = true; + return "SHOULD_NOT_RUN"; + } + } + + [Fact] + public async Task Should_Preserve_Scenario_Tool_Invocation_Identity_Arguments_And_Text() + { + ToolInvocation? observedInvocation = null; + var toolCompleted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Tools = + [ + CopilotTool.DefineTool( + (Func)SearchPullRequests, + factoryOptions: new AIFunctionFactoryOptions + { + Name = "scenario_search_pull_requests", + Description = "Searches pull requests visible to the scenario client.", + }), + ], + }); + using var subscription = session.On(evt => + toolCompleted.TrySetResult(evt)); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_search_pull_requests with query is:open label:bug. Reply with its result.", + }); + var completed = await toolCompleted.Task.WaitAsync(EventTimeout); + + Assert.NotNull(observedInvocation); + Assert.Equal(session.SessionId, observedInvocation!.SessionId); + Assert.Equal("scenario_search_pull_requests", observedInvocation.ToolName); + Assert.False(string.IsNullOrWhiteSpace(observedInvocation.ToolCallId)); + Assert.Equal("is:open label:bug", observedInvocation.Arguments!.Value.GetProperty("query").GetString()); + Assert.Equal(observedInvocation.ToolCallId, completed.Data.ToolCallId); + Assert.True(completed.Data.Success); + Assert.Contains("SCENARIO_SEARCH_TEXT", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + TextContent SearchPullRequests( + [Description("GitHub search query")] string query, + ToolInvocation invocation) + { + observedInvocation = invocation; + return new TextContent($"SCENARIO_SEARCH_TEXT:{query}"); + } + } + + [Fact] + public async Task Should_Deliver_Expanded_Scenario_Tool_Result_To_The_Model() + { + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Tools = + [ + AIFunctionFactory.Create( + GetDeployment, + "scenario_get_deployment", + serializerOptions: ScenarioToolsJsonContext.Default.Options), + ], + }); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_get_deployment for environment production. Reply with its result.", + }); + + Assert.Contains("SCENARIO_DEPLOYMENT_READY", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + var exchange = (await Ctx.GetExchangesAsync()).Last(); + var toolResult = Assert.Single(exchange.Request.Messages, message => message.Role == "tool"); + Assert.Equal("SCENARIO_DEPLOYMENT_READY:production", toolResult.StringContent); + Assert.DoesNotContain("toolTelemetry", toolResult.StringContent, StringComparison.Ordinal); + Assert.DoesNotContain("resultType", toolResult.StringContent, StringComparison.Ordinal); + + [Description("Gets deployment state from the scenario client")] + static ToolResultAIContent GetDeployment([Description("Deployment environment")] string environment) => + new(new ToolResultObject + { + TextResultForLlm = $"SCENARIO_DEPLOYMENT_READY:{environment}", + ResultType = "success", + SessionLog = "scenario client deployment lookup completed.", + ToolTelemetry = new Dictionary + { + ["source"] = JsonValue.Create("scenario-client")!, + }, + }); + } + + [Fact] + public async Task Should_Isolate_Scenario_Tool_Handler_Error() + { + var toolCompleted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + await using var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Tools = [AIFunctionFactory.Create(FailingLookup, "scenario_failing_lookup")], + }); + using var subscription = session.On(evt => + toolCompleted.TrySetResult(evt)); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Call scenario_failing_lookup. If it fails, reply with exactly SCENARIO_LOOKUP_UNAVAILABLE.", + }); + var completed = await toolCompleted.Task.WaitAsync(EventTimeout); + + Assert.False(completed.Data.Success); + Assert.DoesNotContain("SCENARIO_PRIVATE_HANDLER_DETAIL", completed.Data.Error?.Message ?? string.Empty, StringComparison.Ordinal); + Assert.Contains("SCENARIO_LOOKUP_UNAVAILABLE", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + Assert.DoesNotContain("SCENARIO_PRIVATE_HANDLER_DETAIL", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + + static string FailingLookup() => throw new InvalidOperationException("SCENARIO_PRIVATE_HANDLER_DETAIL"); + } + + [Fact] + public async Task Should_Cancel_Scenario_Tool_Handler_When_Session_Disposes() + { + var started = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var cancelled = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + var session = await CreateSessionAsync(new SessionConfig + { + ClientName = "scenario-client", + Tools = [AIFunctionFactory.Create(WaitForScenarioAsync, "scenario_wait_for_operation")], + }); + + _ = session.SendAsync(new MessageOptions + { + Prompt = "Call scenario_wait_for_operation with operation sync-installation.", + }); + + Assert.Equal("sync-installation", await started.Task.WaitAsync(EventTimeout)); + await session.DisposeAsync(); + await cancelled.Task.WaitAsync(EventTimeout); + release.TrySetResult("RELEASED_AFTER_DISPOSE"); + + [Description("Waits for a scenario-owned operation")] + async Task WaitForScenarioAsync( + [Description("Operation name")] string operation, + CancellationToken cancellationToken) + { + started.TrySetResult(operation); + try + { + return await release.Task.WaitAsync(Timeout.InfiniteTimeSpan, cancellationToken); + } + catch (OperationCanceledException) + { + cancelled.TrySetResult(); + throw; + } + } + } +} diff --git a/dotnet/test/E2E/ScenarioTestingUtilityE2ETests.cs b/dotnet/test/E2E/ScenarioTestingUtilityE2ETests.cs new file mode 100644 index 0000000000..0a9bbcdcb4 --- /dev/null +++ b/dotnet/test/E2E/ScenarioTestingUtilityE2ETests.cs @@ -0,0 +1,47 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Test.Harness; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public class ScenarioTestingUtilityE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : ScenarioTestingE2ETestBase(fixture, "scenario_testing_utility", output) +{ + [Fact] + public async Task Should_Send_Wait_Observe_Idle_Events_And_Delete_Suggestion_Session() + { + var sessionId = Guid.NewGuid().ToString(); + var session = await CreateSessionAsync(new SessionConfig { SessionId = sessionId }); + var idle = TestHelper.GetNextEventOfTypeAsync( + session, + TimeSpan.FromSeconds(60)); + + var response = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Reply with exactly SCENARIO_SUGGESTION_ACCEPTED.", + DisplayPrompt = "Apply suggested response", + Mode = "enqueue", + Source = MessageSource.Agent("suggestions"), + }); + await idle; + + Assert.Contains("SCENARIO_SUGGESTION_ACCEPTED", response?.Data.Content ?? string.Empty, StringComparison.Ordinal); + var events = await session.GetEventsAsync(); + var userMessage = Assert.Single( + events.OfType(), + evt => evt.Data.Content == "Apply suggested response"); + Assert.Equal("agent-suggestions", userMessage.Data.Source); + Assert.Equal(UserMessageDelivery.Idle, userMessage.Data.Delivery); + Assert.Contains( + events.OfType(), + evt => (evt.Data.Content ?? string.Empty).Contains("SCENARIO_SUGGESTION_ACCEPTED", StringComparison.Ordinal)); + + await session.DisposeAsync(); + await Client.DeleteSessionAsync(sessionId); + Assert.Null(await Client.GetSessionMetadataAsync(sessionId)); + } +} diff --git a/dotnet/test/Harness/E2ETestBase.cs b/dotnet/test/Harness/E2ETestBase.cs index f664812d58..852e6a694a 100644 --- a/dotnet/test/Harness/E2ETestBase.cs +++ b/dotnet/test/Harness/E2ETestBase.cs @@ -17,15 +17,21 @@ public abstract class E2ETestBase : IClassFixture, IAsyncLifetim private readonly E2ETestFixture _fixture; private readonly string _snapshotCategory; private readonly string _testName; + private readonly bool _replayOnly; protected E2ETestContext Ctx => _fixture.Ctx; protected CopilotClient Client => _fixture.Client; - protected E2ETestBase(E2ETestFixture fixture, string snapshotCategory, ITestOutputHelper output) + protected E2ETestBase( + E2ETestFixture fixture, + string snapshotCategory, + ITestOutputHelper output, + bool replayOnly = false) { _fixture = fixture; _snapshotCategory = snapshotCategory; _testName = GetTestName(output); + _replayOnly = replayOnly; Logger = new XunitLogger(output); // Wire logger into the shared context so all clients created via Ctx.CreateClient get it. @@ -61,7 +67,7 @@ public async Task InitializeAsync() { Ctx.PrepareForTest(); await Ctx.CleanupAfterTestAsync(); - await Ctx.ConfigureForTestAsync(_snapshotCategory, _testName); + await Ctx.ConfigureForTestAsync(_snapshotCategory, _testName, _replayOnly); } public Task DisposeAsync() diff --git a/dotnet/test/Harness/E2ETestContext.cs b/dotnet/test/Harness/E2ETestContext.cs index 6f03cbaf37..7b529cc591 100644 --- a/dotnet/test/Harness/E2ETestContext.cs +++ b/dotnet/test/Harness/E2ETestContext.cs @@ -207,7 +207,10 @@ private static string PrepareCliPath(string repoRoot, string option) return cliPath; } - public async Task ConfigureForTestAsync(string testFile, [CallerMemberName] string? testName = null) + public async Task ConfigureForTestAsync( + string testFile, + [CallerMemberName] string? testName = null, + bool replayOnly = false) { // Convert test method names to lowercase snake_case for snapshot filenames // to avoid case collisions on case-insensitive filesystems (macOS/Windows) @@ -216,7 +219,8 @@ public async Task ConfigureForTestAsync(string testFile, [CallerMemberName] stri await _proxy.ConfigureAsync( snapshotPath, WorkDir, - E2ETestBackendConfiguration.Current.ToWireName()); + E2ETestBackendConfiguration.Current.ToWireName(), + replayOnly); } public Task> GetExchangesAsync() diff --git a/dotnet/test/Harness/ReplayProxy.cs b/dotnet/test/Harness/ReplayProxy.cs index 1f906d72a5..8c981988ce 100644 --- a/dotnet/test/Harness/ReplayProxy.cs +++ b/dotnet/test/Harness/ReplayProxy.cs @@ -150,19 +150,19 @@ public async Task StopAsync(bool skipWritingCache = false) _startupTask = null; } - public async Task ConfigureAsync(string filePath, string workDir, string backend) + public async Task ConfigureAsync(string filePath, string workDir, string backend, bool replayOnly = false) { var url = await (_startupTask ?? throw new InvalidOperationException("Proxy not started")); using var client = new HttpClient(); var response = await client.PostAsJsonAsync( $"{url}/config", - new ConfigureRequest(filePath, workDir, backend), + new ConfigureRequest(filePath, workDir, backend, replayOnly), ReplayProxyJsonContext.Default.ConfigureRequest); response.EnsureSuccessStatusCode(); } - private record ConfigureRequest(string FilePath, string WorkDir, string Backend); + private record ConfigureRequest(string FilePath, string WorkDir, string Backend, bool ReplayOnly); private record ProxyStartupMetadata(string? ConnectProxyUrl, string? CaFilePath); diff --git a/dotnet/test/Unit/CloneTests.cs b/dotnet/test/Unit/CloneTests.cs index 20d5a1c296..0234431892 100644 --- a/dotnet/test/Unit/CloneTests.cs +++ b/dotnet/test/Unit/CloneTests.cs @@ -8,9 +8,11 @@ namespace GitHub.Copilot.Test.Unit; public class CloneTests { +#pragma warning disable GHCP001 [Fact] public void CopilotClientOptions_Clone_CopiesAllProperties() { + var extensionLaunchProvider = new TestExtensionLaunchProvider(); var original = new CopilotClientOptions { Connection = RuntimeConnection.ForTcp(port: 8080, connectionToken: "tok", path: "/usr/bin/copilot", args: ["--verbose", "--debug"]), @@ -23,6 +25,7 @@ public void CopilotClientOptions_Clone_CopiesAllProperties() BuiltinPluginDirectories = ["/plugins/core", "/plugins/github"], EnableRemoteSessions = true, SessionIdleTimeoutSeconds = 600, + ExtensionLaunchProvider = extensionLaunchProvider, ClientInfo = new CopilotClientInfo { ApplicationName = "example-app", @@ -45,8 +48,10 @@ public void CopilotClientOptions_Clone_CopiesAllProperties() Assert.NotSame(original.BuiltinPluginDirectories, clone.BuiltinPluginDirectories); Assert.Equal(original.EnableRemoteSessions, clone.EnableRemoteSessions); Assert.Equal(original.SessionIdleTimeoutSeconds, clone.SessionIdleTimeoutSeconds); + Assert.Same(extensionLaunchProvider, clone.ExtensionLaunchProvider); Assert.Same(original.ClientInfo, clone.ClientInfo); } +#pragma warning restore GHCP001 [Fact] public void CopilotClientOptions_Clone_ConnectionIsShared() @@ -70,6 +75,14 @@ public void CopilotClientOptions_Clone_EnvironmentIsShared() Assert.Same(original.Environment, clone.Environment); } + private sealed class TestExtensionLaunchProvider : GitHub.Copilot.Rpc.IExtensionLaunchProviderHandler + { + public Task ResolveAsync( + GitHub.Copilot.Rpc.ExtensionLaunchProviderResolveRequest request, + CancellationToken cancellationToken = default) => + Task.FromResult(new GitHub.Copilot.Rpc.ExtensionLaunchProviderResolveResult()); + } + [Fact] public void SessionConfig_Clone_CopiesAllProperties() { diff --git a/dotnet/test/Unit/JsonRpcTests.cs b/dotnet/test/Unit/JsonRpcTests.cs index f4acfd3555..8e9a6aaacc 100644 --- a/dotnet/test/Unit/JsonRpcTests.cs +++ b/dotnet/test/Unit/JsonRpcTests.cs @@ -94,6 +94,46 @@ public async Task JsonRpc_Cancels_And_Disposes_Pending_Requests() await Assert.ThrowsAnyAsync(() => pending); } + [Fact] + public async Task JsonRpc_Dispose_Completes_Cleanup_When_Cancellation_Callback_Throws() + { + using var pair = JsonRpcReflectionPair.Create(startServer: false); + using var registration = pair.Client.RegisterDisposeCallback( + () => throw new InvalidOperationException("callback failed")); + var pending = pair.Client.InvokeAsync("stillPending", args: null); + + var exception = Assert.Throws(() => pair.Client.Dispose()); + + Assert.Contains( + exception.InnerExceptions, + inner => inner is InvalidOperationException { Message: "callback failed" }); + await Assert.ThrowsAnyAsync(() => pending); + Assert.True(pair.Client.Completion.IsCompleted); + Assert.False(pair.Client.Completion.IsFaulted); + Assert.False(pair.Client.Completion.IsCanceled); + pair.Client.Dispose(); + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public async Task JsonRpc_Process_Exit_Reports_Connection_Lost(bool requestBeforeExit) + { + using var pair = JsonRpcReflectionPair.Create(startServer: false); + await using var client = new CopilotClient(); + var pending = requestBeforeExit + ? pair.Client.InvokeAsync("pendingAtExit", args: null) + : null; + + pair.Client.NotifyProcessExit(client); + pending ??= pair.Client.InvokeAsync("afterExit", args: null); + + var exception = await Assert.ThrowsAnyAsync(() => pending); + Assert.Equal("ConnectionLostException", exception.GetType().Name); + Assert.Equal("The JSON-RPC connection was lost.", exception.Message); + Assert.True(pair.Client.Completion.IsCompleted); + } + [Fact] public async Task JsonRpc_Does_Not_Retain_Oversized_Receive_Buffer() { @@ -236,11 +276,21 @@ public JsonRpcReflection(Stream sendStream, Stream receiveStream) culture: null)!; } + public Task Completion => (Task)JsonRpcType.GetProperty(nameof(Completion))!.GetValue(_instance)!; + public void StartListening() => JsonRpcType.GetMethod(nameof(StartListening))!.Invoke(_instance, null); public void SetLocalRpcMethod(string methodName, Delegate handler, bool singleObjectParam = false) => JsonRpcType.GetMethod("SetLocalRpcMethod")!.Invoke(_instance, [methodName, handler, singleObjectParam]); + public CancellationTokenRegistration RegisterDisposeCallback(Action callback) + { + var disposeCts = (CancellationTokenSource)JsonRpcType + .GetField("_disposeCts", BindingFlags.Instance | BindingFlags.NonPublic)! + .GetValue(_instance)!; + return disposeCts.Token.Register(callback); + } + public async Task InvokeAsync(string methodName, object?[]? args, CancellationToken cancellationToken = default) { var method = JsonRpcType @@ -253,6 +303,11 @@ public async Task InvokeAsync(string methodName, object?[]? args, Cancella } public void Dispose() => ((IDisposable)_instance).Dispose(); + + public void NotifyProcessExit(CopilotClient client) => + typeof(CopilotClient) + .GetMethod("DisposeRpcAfterProcessExit", BindingFlags.Instance | BindingFlags.NonPublic)! + .Invoke(client, [_instance]); } private sealed class CoalescedFramesThenWaitStream : Stream diff --git a/go/README.md b/go/README.md index fd8bdd245a..b6d80cf6d0 100644 --- a/go/README.md +++ b/go/README.md @@ -209,6 +209,7 @@ Event types: `SessionLifecycleCreated`, `SessionLifecycleDeleted`, `SessionLifec `StdioConnection` and `TCPConnection` accept an optional connection-level `Env`. Set environment variables via **either** the client-level `Env` option or the connection's `Env`, not both (setting both panics); prefer the connection-level `Env`. - `WorkingDirectory` (string): Working directory for the runtime process (default: current process working directory) - `BaseDirectory` (string): Base directory for Copilot data (session state, config, etc.). Sets `COPILOT_HOME` on the spawned runtime. When empty, the runtime defaults to `~/.copilot`. Ignored with `URIConnection`. This does **not** affect where the Go SDK extracts the embedded CLI binary; use `embeddedcli.Config.Dir` for the extraction/cache location. +- `ExtensionLaunchProvider` (ExtensionLaunchProvider): Experimental connection-level resolver for extension launch profiles. `Start` installs the reverse-RPC handler and registers the provider before sessions can be created. - `LogLevel` (string): Log level. When empty (default), the runtime uses its own default level (the SDK does not pass `--log-level`). - `Env` ([]string): Environment variables for the runtime process (default: inherits from current process) - `GitHubToken` (string): GitHub token for authentication. When provided, takes priority over other auth methods. diff --git a/go/client.go b/go/client.go index e450c595f2..4623872a53 100644 --- a/go/client.go +++ b/go/client.go @@ -499,6 +499,19 @@ func (c *Client) Start(ctx context.Context) error { return errors.Join(err, killErr) } + if c.options.ExtensionLaunchProvider != nil { + if _, err := c.RPC.RegisterExtensionLaunchProvider(ctx); err != nil { + c.client.Stop() + c.client = nil + c.conn = nil + c.RPC = nil + c.internalRPC = nil + killErr := c.killProcess() + c.state = stateError + return errors.Join(err, killErr) + } + } + if len(c.options.BuiltinPluginDirectories) > 0 { if _, err := c.client.Request(ctx, "plugins.builtin.set", map[string]any{ "paths": c.options.BuiltinPluginDirectories, @@ -2486,8 +2499,9 @@ func (c *Client) setupNotificationHandler() { // payload's sessionId. Always register the global handlers so the generated // hooks.invoke handler is wired to our dispatcher. handlers := &rpc.ClientGlobalAPIHandlers{ - Hooks: &hooksAdapter{client: c}, - GitHubToken: &gitHubTokenAdapter{client: c}, + ExtensionLaunchProvider: c.options.ExtensionLaunchProvider, + Hooks: &hooksAdapter{client: c}, + GitHubToken: &gitHubTokenAdapter{client: c}, } if c.options.RequestHandler != nil { diff --git a/go/internal/e2e/client_options_e2e_test.go b/go/internal/e2e/client_options_e2e_test.go index 0b8423649c..739f1a58da 100644 --- a/go/internal/e2e/client_options_e2e_test.go +++ b/go/internal/e2e/client_options_e2e_test.go @@ -128,6 +128,7 @@ func TestClientOptionsE2E(t *testing.T) { } opts.UseLoggedInUser = copilot.Bool(false) }) + t.Cleanup(func() { client.ForceStop() }) if err := client.Start(t.Context()); err != nil { @@ -241,6 +242,74 @@ func TestClientOptionsE2E(t *testing.T) { t.Fatalf("session.resume request was not captured. Captured requests: %+v", resumedCapture.Requests) }) + t.Run("should register and invoke extension launch provider during startup", func(t *testing.T) { + if testharness.RunInIsolatedProcess(t) { + return + } + ctx := testharness.NewTestContext(t) + cliPath := filepath.Join(ctx.WorkDir, "fake-cli-extension-provider-"+randomHex(t)+".js") + capturePath := filepath.Join(ctx.WorkDir, "fake-cli-extension-provider-"+randomHex(t)+".json") + if err := os.WriteFile(cliPath, []byte(fakeStdioCliScript), 0644); err != nil { + t.Fatalf("Failed to write fake CLI script: %v", err) + } + + requests := make(chan *copilot.ExtensionLaunchProviderResolveRequest, 1) + client := ctx.NewClient(func(opts *copilot.ClientOptions) { + opts.Connection = copilot.StdioConnection{ + Path: cliPath, + Args: []string{"--capture-file", capturePath}, + } + opts.GitHubToken = "" + opts.UseLoggedInUser = copilot.Bool(false) + opts.ExtensionLaunchProvider = extensionLaunchProviderFunc(func( + request *copilot.ExtensionLaunchProviderResolveRequest, + ) (*copilot.ExtensionLaunchProviderResolveResult, error) { + requests <- request + return &copilot.ExtensionLaunchProviderResolveResult{ + Launch: &copilot.ExtensionLaunchProfile{ + Executable: "go", + Args: []string{"extension-host"}, + Env: map[string]string{"EXTENSION_SOURCE": "go"}, + }, + }, nil + }) + }) + t.Cleanup(func() { client.ForceStop() }) + + if err := client.Start(t.Context()); err != nil { + t.Fatalf("Start failed: %v", err) + } + + request := <-requests + if request.ID != "project:go-e2e" || request.Name != "go-e2e" || + request.ModulePath != "/extensions/go-e2e.go" || request.Source != rpc.ExtensionSourceProject { + t.Fatalf("Unexpected extension launch request: %+v", request) + } + + capture := readCapture(t, capturePath) + foundRegistration := false + for _, captured := range capture.Requests { + if captured.Method == "registerExtensionLaunchProvider" { + foundRegistration = true + break + } + } + if !foundRegistration { + t.Fatalf("registerExtensionLaunchProvider request was not captured: %+v", capture.Requests) + } + if len(capture.ClientResponses) != 1 { + t.Fatalf("Expected one extension launch response, got %+v", capture.ClientResponses) + } + response := capture.ClientResponses[0] + if response.ID != 9001 { + t.Fatalf("Expected response id 9001, got %d", response.ID) + } + launch := response.Result["launch"].(map[string]any) + if launch["executable"] != "go" { + t.Fatalf("Expected Go launch profile, got %+v", launch) + } + }) + t.Run("should send empty-mode custom agent locality defaults in initial requests", func(t *testing.T) { if testharness.RunInIsolatedProcess(t) { return @@ -761,10 +830,11 @@ func assertArgValue(t *testing.T, args []string, name, expected string) { // capturedCli mirrors the JSON file written by the fake stdio CLI script. type capturedCli struct { - Args []string `json:"args"` - WorkingDirectory string `json:"cwd"` - Requests []capturedRequest `json:"requests"` - Env map[string]string `json:"env"` + Args []string `json:"args"` + WorkingDirectory string `json:"cwd"` + Requests []capturedRequest `json:"requests"` + ClientResponses []capturedResponse `json:"clientResponses"` + Env map[string]string `json:"env"` } type capturedRequest struct { @@ -772,6 +842,21 @@ type capturedRequest struct { Params any `json:"params"` } +type capturedResponse struct { + ID int `json:"id"` + Result map[string]any `json:"result"` +} + +type extensionLaunchProviderFunc func( + request *copilot.ExtensionLaunchProviderResolveRequest, +) (*copilot.ExtensionLaunchProviderResolveResult, error) + +func (f extensionLaunchProviderFunc) Resolve( + request *copilot.ExtensionLaunchProviderResolveRequest, +) (*copilot.ExtensionLaunchProviderResolveResult, error) { + return f(request) +} + func readCapture(t *testing.T, path string) capturedCli { t.Helper() data, err := os.ReadFile(path) @@ -821,6 +906,8 @@ const fs = require("fs"); const captureIndex = process.argv.indexOf("--capture-file"); const captureFile = captureIndex >= 0 ? process.argv[captureIndex + 1] : undefined; const requests = []; +const clientResponses = []; +let extensionRegistrationId; function saveCapture() { if (!captureFile) { @@ -830,6 +917,7 @@ function saveCapture() { args: process.argv.slice(2), cwd: process.cwd(), requests, + clientResponses, env: { COPILOT_HOME: process.env.COPILOT_HOME, COPILOT_SDK_AUTH_TOKEN: process.env.COPILOT_SDK_AUTH_TOKEN, @@ -875,6 +963,15 @@ function handleMessage(message) { if (!Object.prototype.hasOwnProperty.call(message, "id")) { return; } + if (!message.method) { + clientResponses.push(message); + saveCapture(); + if (message.id === 9001 && extensionRegistrationId !== undefined) { + writeResponse(extensionRegistrationId, {}); + extensionRegistrationId = undefined; + } + return; + } requests.push({ method: message.method, params: message.params }); saveCapture(); if (message.method === "connect") { @@ -885,6 +982,16 @@ function handleMessage(message) { writeResponse(message.id, { message: "pong", protocolVersion: 3, timestamp: Date.now() }); return; } + if (message.method === "registerExtensionLaunchProvider") { + extensionRegistrationId = message.id; + writeRequest(9001, "extensionLaunchProvider.resolve", { + id: "project:go-e2e", + name: "go-e2e", + modulePath: "/extensions/go-e2e.go", + source: "project", + }); + return; + } if (message.method === "session.create" || message.method === "session.resume") { const sessionId = (message.params && message.params.sessionId) || "fake-session"; writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); @@ -906,4 +1013,9 @@ function writeResponse(id, result) { const body = JSON.stringify({ jsonrpc: "2.0", id, result }); process.stdout.write("Content-Length: " + Buffer.byteLength(body, "utf8") + "\r\n\r\n" + body); } + +function writeRequest(id, method, params) { + const body = JSON.stringify({ jsonrpc: "2.0", id, method, params }); + process.stdout.write("Content-Length: " + Buffer.byteLength(body, "utf8") + "\r\n\r\n" + body); +} ` diff --git a/go/internal/e2e/rpc_generated_surface_gaps_e2e_test.go b/go/internal/e2e/rpc_generated_surface_gaps_e2e_test.go new file mode 100644 index 0000000000..99d9c0a4b1 --- /dev/null +++ b/go/internal/e2e/rpc_generated_surface_gaps_e2e_test.go @@ -0,0 +1,517 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package e2e + +import ( + "context" + "encoding/json" + "net" + "reflect" + "runtime" + "strings" + "testing" + "time" + + copilot "github.com/github/copilot-sdk/go" + "github.com/github/copilot-sdk/go/internal/jsonrpc2" + "github.com/github/copilot-sdk/go/rpc" +) + +type generatedRPCGapCase struct { + name string + wire string + signature any + receiver func(*generatedRPCFixture) any + rpcError bool +} + +func TestGeneratedRPCSurfaceGapsE2E(t *testing.T) { + ctx, cancel := testContextWithTimeout(t, 20*time.Second) + defer cancel() + f := newGeneratedRPCFixture(t, ctx) + + server := func(get func(*rpc.ServerRPC) any) func(*generatedRPCFixture) any { + return func(f *generatedRPCFixture) any { return get(f.client.RPC) } + } + session := func(get func(*rpc.SessionRPC) any) func(*generatedRPCFixture) any { + return func(f *generatedRPCFixture) any { return get(f.session.RPC) } + } + + cases := []generatedRPCGapCase{ + {"server catalog search", "catalog.search", (*rpc.ServerCatalogAPI).Search, server(func(r *rpc.ServerRPC) any { return r.Catalog }), true}, + {"server extensions discover", "extensions.discover", (*rpc.ServerExtensionsAPI).Discover, server(func(r *rpc.ServerRPC) any { return r.Extensions }), false}, + {"server hooks discover", "hooks.discover", (*rpc.ServerHooksAPI).Discover, server(func(r *rpc.ServerRPC) any { return r.Hooks }), false}, + {"server llm inference set provider", "llmInference.setProvider", (*rpc.ServerLlmInferenceAPI).SetProvider, server(func(r *rpc.ServerRPC) any { return r.LlmInference }), false}, + {"server managed settings read", "managedSettings.read", (*rpc.ServerManagedSettingsAPI).Read, server(func(r *rpc.ServerRPC) any { return r.ManagedSettings }), false}, + {"server mcp plan install", "mcp.planInstall", (*rpc.ServerMCPAPI).PlanInstall, server(func(r *rpc.ServerRPC) any { return r.MCP }), false}, + {"server models built in catalog", "models.getBuiltInCatalog", (*rpc.ServerModelsAPI).GetBuiltInCatalog, server(func(r *rpc.ServerRPC) any { return r.Models }), false}, + {"server plugins builtin set", "plugins.builtin.set", (*rpc.ServerPluginsBuiltinAPI).Set, server(func(r *rpc.ServerRPC) any { return r.Plugins.Builtin() }), false}, + {"server sessions client metadata", "sessions.getClientMetadata", (*rpc.ServerSessionsAPI).GetClientMetadata, server(func(r *rpc.ServerRPC) any { return r.Sessions }), false}, + {"server sessions persisted events", "sessions.readPersistedEvents", (*rpc.ServerSessionsAPI).ReadPersistedEvents, server(func(r *rpc.ServerRPC) any { return r.Sessions }), false}, + {"server skills disabled", "skills.config.setSkillDisabled", (*rpc.ServerSkillsConfigAPI).SetSkillDisabled, server(func(r *rpc.ServerRPC) any { return r.Skills.Config() }), false}, + {"server extension launch provider", "registerExtensionLaunchProvider", (*rpc.ServerRPC).RegisterExtensionLaunchProvider, server(func(r *rpc.ServerRPC) any { return r }), false}, + {"session agent prompt", "session.agent.setPrompt", (*rpc.AgentAPI).SetPrompt, session(func(r *rpc.SessionRPC) any { return r.Agent }), false}, + {"session autopilot objective state", "session.autopilotObjective.getState", (*rpc.AutopilotObjectiveAPI).GetState, session(func(r *rpc.SessionRPC) any { return r.AutopilotObjective }), false}, + {"session canvas list open", "session.canvas.listOpen", (*rpc.CanvasAPI).ListOpen, session(func(r *rpc.SessionRPC) any { return r.Canvas }), false}, + {"session completion triggers", "session.completions.getTriggerCharacters", (*rpc.CompletionsAPI).GetTriggerCharacters, session(func(r *rpc.SessionRPC) any { return r.Completions }), false}, + {"session content exclusion paths", "session.contentExclusion.checkPaths", (*rpc.ContentExclusionAPI).CheckPaths, session(func(r *rpc.SessionRPC) any { return r.ContentExclusion }), false}, + {"session debug logs", "session.debug.collectLogs", (*rpc.DebugAPI).CollectLogs, session(func(r *rpc.SessionRPC) any { return r.Debug }), false}, + {"session factory agent", "session.factory.agent", (*rpc.FactoryAPI).Agent, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory cancel", "session.factory.cancel", (*rpc.FactoryAPI).Cancel, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory get run", "session.factory.getRun", (*rpc.FactoryAPI).GetRun, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory detail", "session.factory.getRunDetail", (*rpc.FactoryAPI).GetRunDetail, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory progress", "session.factory.getRunProgress", (*rpc.FactoryAPI).GetRunProgress, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory list runs", "session.factory.listRuns", (*rpc.FactoryAPI).ListRuns, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory log", "session.factory.log", (*rpc.FactoryAPI).Log, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory pause", "session.factory.pause", (*rpc.FactoryAPI).Pause, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory resume", "session.factory.resume", (*rpc.FactoryAPI).Resume, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory run", "session.factory.run", (*rpc.FactoryAPI).Run, session(func(r *rpc.SessionRPC) any { return r.Factory }), false}, + {"session factory journal get", "session.factory.journal.get", (*rpc.FactoryJournalAPI).Get, session(func(r *rpc.SessionRPC) any { return r.Factory.Journal() }), false}, + {"session factory journal put", "session.factory.journal.put", (*rpc.FactoryJournalAPI).Put, session(func(r *rpc.SessionRPC) any { return r.Factory.Journal() }), false}, + {"session history clear context", "session.history.clearContext", (*rpc.HistoryAPI).ClearContext, session(func(r *rpc.SessionRPC) any { return r.History }), false}, + {"session limit prediction", "session.limitPrediction.predict", (*rpc.LimitPredictionAPI).Predict, session(func(r *rpc.SessionRPC) any { return r.LimitPrediction }), false}, + {"session mcp loading background", "session.mcp.moveLoadingToBackground", (*rpc.MCPAPI).MoveLoadingToBackground, session(func(r *rpc.SessionRPC) any { return r.MCP }), false}, + {"session mcp restart", "session.mcp.restartServer", (*rpc.MCPAPI).RestartServer, session(func(r *rpc.SessionRPC) any { return r.MCP }), false}, + {"session mcp start", "session.mcp.startServer", (*rpc.MCPAPI).StartServer, session(func(r *rpc.SessionRPC) any { return r.MCP }), false}, + {"session mcp oauth state changed", "session.mcp.oauth.authenticationStateChanged", (*rpc.MCPOauthAPI).AuthenticationStateChanged, session(func(r *rpc.SessionRPC) any { return r.MCP.Oauth() }), false}, + {"session mcp oauth probe", "session.mcp.oauth.probe", (*rpc.MCPOauthAPI).Probe, session(func(r *rpc.SessionRPC) any { return r.MCP.Oauth() }), true}, + {"session mcp oauth respond", "session.mcp.oauth.respond", (*rpc.MCPOauthAPI).Respond, session(func(r *rpc.SessionRPC) any { return r.MCP.Oauth() }), false}, + {"session mcp resources list", "session.mcp.resources.list", (*rpc.MCPResourcesAPI).List, session(func(r *rpc.SessionRPC) any { return r.MCP.Resources() }), false}, + {"session mcp resource templates", "session.mcp.resources.listTemplates", (*rpc.MCPResourcesAPI).ListTemplates, session(func(r *rpc.SessionRPC) any { return r.MCP.Resources() }), false}, + {"session mcp resource read", "session.mcp.resources.read", (*rpc.MCPResourcesAPI).Read, session(func(r *rpc.SessionRPC) any { return r.MCP.Resources() }), false}, + {"session metadata get", "session.metadata.getClientMetadata", (*rpc.MetadataAPI).GetClientMetadata, session(func(r *rpc.SessionRPC) any { return r.Metadata }), false}, + {"session metadata update", "session.metadata.updateClientMetadata", (*rpc.MetadataAPI).UpdateClientMetadata, session(func(r *rpc.SessionRPC) any { return r.Metadata }), false}, + {"session allowed models", "session.model.setAllowedModels", (*rpc.ModelAPI).SetAllowedModels, session(func(r *rpc.SessionRPC) any { return r.Model }), false}, + {"session auto tier", "session.model.switchAutoTier", (*rpc.ModelAPI).SwitchAutoTier, session(func(r *rpc.SessionRPC) any { return r.Model }), false}, + {"session queue duplicate", "session.queue.duplicateAt", (*rpc.QueueAPI).DuplicateAt, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue insert", "session.queue.insertAt", (*rpc.QueueAPI).InsertAt, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue move", "session.queue.moveItem", (*rpc.QueueAPI).MoveItem, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue remove", "session.queue.removeAt", (*rpc.QueueAPI).RemoveAt, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue send now", "session.queue.sendNow", (*rpc.QueueAPI).SendNow, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue drain pause", "session.queue.setDrainPaused", (*rpc.QueueAPI).SetDrainPaused, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session queue update", "session.queue.updateText", (*rpc.QueueAPI).UpdateText, session(func(r *rpc.SessionRPC) any { return r.Queue }), false}, + {"session sandbox disable", "session.sandbox.disableForSession", (*rpc.SandboxAPI).DisableForSession, session(func(r *rpc.SessionRPC) any { return r.Sandbox }), false}, + {"session sandbox status", "session.sandbox.getEnforcementStatus", (*rpc.SandboxAPI).GetEnforcementStatus, session(func(r *rpc.SessionRPC) any { return r.Sandbox }), false}, + {"session tasks register", "session.tasks.register", (*rpc.TasksAPI).Register, session(func(r *rpc.SessionRPC) any { return r.Tasks }), false}, + {"session tasks update", "session.tasks.update", (*rpc.TasksAPI).Update, session(func(r *rpc.SessionRPC) any { return r.Tasks }), false}, + {"session tools execute", "session.tools.execute", (*rpc.ToolsAPI).Execute, session(func(r *rpc.SessionRPC) any { return r.Tools }), false}, + {"session builtin tool descriptors", "session.tools.getBuiltinDescriptors", (*rpc.ToolsAPI).GetBuiltinDescriptors, session(func(r *rpc.SessionRPC) any { return r.Tools }), false}, + {"session tools set", "session.tools.set", (*rpc.ToolsAPI).Set, session(func(r *rpc.SessionRPC) any { return r.Tools }), false}, + {"session task complete event data", "session.tools.taskCompleteEventData", (*rpc.ToolsAPI).TaskCompleteEventData, session(func(r *rpc.SessionRPC) any { return r.Tools }), false}, + {"session workspace summary", "session.workspaces.addSummary", (*rpc.WorkspacesAPI).AddSummary, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace objective exists", "session.workspaces.autopilotObjectiveExists", (*rpc.WorkspacesAPI).AutopilotObjectiveExists, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace directory", "session.workspaces.createDirectory", (*rpc.WorkspacesAPI).CreateDirectory, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace objective delete", "session.workspaces.deleteAutopilotObjective", (*rpc.WorkspacesAPI).DeleteAutopilotObjective, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace ensure", "session.workspaces.ensure", (*rpc.WorkspacesAPI).Ensure, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace objective read", "session.workspaces.readAutopilotObjective", (*rpc.WorkspacesAPI).ReadAutopilotObjective, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace remove", "session.workspaces.removePath", (*rpc.WorkspacesAPI).RemovePath, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace rename", "session.workspaces.renamePath", (*rpc.WorkspacesAPI).RenamePath, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace stat", "session.workspaces.statFile", (*rpc.WorkspacesAPI).StatFile, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace truncate summaries", "session.workspaces.truncateSummaries", (*rpc.WorkspacesAPI).TruncateSummaries, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace metadata", "session.workspaces.updateMetadata", (*rpc.WorkspacesAPI).UpdateMetadata, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session workspace objective write", "session.workspaces.writeAutopilotObjective", (*rpc.WorkspacesAPI).WriteAutopilotObjective, session(func(r *rpc.SessionRPC) any { return r.Workspaces }), false}, + {"session abort", "session.abort", (*rpc.SessionRPC).Abort, session(func(r *rpc.SessionRPC) any { return r }), false}, + {"session cancel background agents", "session.cancelAllBackgroundAgents", (*rpc.SessionRPC).CancelAllBackgroundAgents, session(func(r *rpc.SessionRPC) any { return r }), false}, + {"session interrupt main turn", "session.interruptMainTurn", (*rpc.SessionRPC).InterruptMainTurn, session(func(r *rpc.SessionRPC) any { return r }), false}, + {"session log", "session.log", (*rpc.SessionRPC).Log, session(func(r *rpc.SessionRPC) any { return r }), false}, + {"session send", "session.send", (*rpc.SessionRPC).Send, session(func(r *rpc.SessionRPC) any { return r }), false}, + {"session send messages", "session.sendMessages", (*rpc.SessionRPC).SendMessages, session(func(r *rpc.SessionRPC) any { return r }), false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var captured json.RawMessage + f.server.SetRequestHandler(tc.wire, func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + captured = append(captured[:0], params...) + if tc.rpcError { + return nil, &jsonrpc2.Error{Code: -32042, Message: "synthetic " + tc.wire} + } + return json.RawMessage(generatedRPCResponse(tc.wire)), nil + }) + + method := reflect.ValueOf(tc.receiver(f)).MethodByName(methodName(tc.signature)) + if !method.IsValid() { + t.Fatalf("Method expression %T did not resolve on receiver %T", tc.signature, tc.receiver(f)) + } + args := []reflect.Value{reflect.ValueOf(ctx)} + requestJSON := generatedRPCRequest(tc.wire) + if method.Type().NumIn() == 2 { + if requestJSON == "" { + if !method.Type().IsVariadic() { + args = append(args, reflect.Zero(method.Type().In(1))) + } + } else { + requestType := method.Type().In(1) + if method.Type().IsVariadic() { + requestType = requestType.Elem() + } + request := reflect.New(requestType.Elem()) + if err := json.Unmarshal([]byte(requestJSON), request.Interface()); err != nil { + t.Fatalf("Decode %s typed request: %v", tc.wire, err) + } + args = append(args, request) + } + } + results := method.Call(args) + if len(results) != 2 { + t.Fatalf("Expected two return values, got %d", len(results)) + } + err, _ := results[1].Interface().(error) + if tc.rpcError { + if err == nil || !strings.Contains(err.Error(), "synthetic "+tc.wire) { + t.Fatalf("Expected synthetic RPC error, got %v", err) + } + } else { + if err != nil { + t.Fatalf("%s returned error: %v", tc.wire, err) + } + if results[0].Kind() == reflect.Pointer && results[0].IsNil() { + t.Fatalf("%s returned a nil typed result", tc.wire) + } + assertGeneratedRPCResult(t, tc.wire, results[0].Interface()) + } + + var request map[string]any + if len(captured) != 0 && string(captured) != "null" { + if err := json.Unmarshal(captured, &request); err != nil { + t.Fatalf("Decode %s request %s: %v", tc.wire, captured, err) + } + } + if strings.HasPrefix(tc.wire, "session.") { + if request["sessionId"] != f.session.SessionID { + t.Fatalf("%s sessionId = %#v, want %q; request=%s", tc.wire, request["sessionId"], f.session.SessionID, captured) + } + } else if _, exists := request["sessionId"]; exists { + t.Fatalf("%s unexpectedly serialized sessionId; request=%s", tc.wire, captured) + } + if requestJSON != "" { + var expected map[string]any + if err := json.Unmarshal([]byte(requestJSON), &expected); err != nil { + t.Fatal(err) + } + assertJSONSubset(t, tc.wire+" request", expected, request) + } + }) + } +} + +func generatedRPCRequest(wire string) string { + switch wire { + case "hooks.discover": + return `{"projectPaths":["Q:\\rpc-project"],"excludeHostHooks":true}` + case "mcp.planInstall": + return `{"contract":{"protocolVersion":3,"requiredCapabilities":["mcp-install-planning"]},"source":{"kind":"candidate","candidateHandle":"candidate-1","searchId":"search-1"},"scope":"user"}` + case "plugins.builtin.set": + return `{"paths":["Q:\\rpc-plugins"]}` + case "sessions.getClientMetadata": + return `{"sessionIds":["persisted-session"],"keys":["rpc/key"]}` + case "skills.config.setSkillDisabled": + return `{"name":"skill-one","disabled":true}` + case "session.agent.setPrompt": + return `{"id":"agent-1","prompt":"Use the RPC prompt."}` + case "session.contentExclusion.checkPaths": + return `{"paths":["/tmp/rpc-workspace/file.txt"]}` + case "session.debug.collectLogs": + return `{"destination":{"kind":"directory","outputDirectory":"/tmp/rpc-debug"},"include":{"events":true,"processLogs":false,"shellLogs":true},"additionalEntries":[{"bundlePath":"host/diagnostic.txt","kind":"file","path":"/tmp/diagnostic.txt","required":true}]}` + case "session.factory.run": + return `{"name":"rpc-factory","args":{"input":42},"options":{"limits":{"maxAiCredits":2.5,"maxConcurrentSubagents":2,"maxTotalSubagents":4,"timeoutSeconds":30},"logPhaseNames":true,"notifyOnComplete":false}}` + case "session.factory.resume": + return `{"runId":"factory-run-1","limits":{"maxTotalSubagents":8},"notifyOnComplete":true,"logPhaseNames":false}` + case "session.factory.getRun", "session.factory.pause": + return `{"runId":"factory-run-1"}` + case "session.factory.log": + return `{"runId":"factory-run-1","executionToken":"execution-token-1","lines":[{"kind":"log","seq":7,"text":"Factory progress"}]}` + case "session.factory.agent": + return `{"factoryRunId":"factory-run-1","executionToken":"execution-token-1","prompt":"Complete the RPC task.","opts":{"agent":"explore","label":"rpc-agent","model":"model-a","reasoningEffort":"high"}}` + case "session.factory.journal.get": + return `{"runId":"factory-run-1","executionToken":"execution-token-1","key":"checkpoint"}` + case "session.factory.journal.put": + return `{"runId":"factory-run-1","executionToken":"execution-token-1","key":"checkpoint","resultJson":{"checkpoint":8}}` + case "session.history.clearContext": + return `{"prompt":"Reset context."}` + case "session.limitPrediction.predict": + return `{"clientType":"sdk","modelId":"model-a"}` + case "session.mcp.startServer": + return `{"serverName":"rpc-server","config":{"type":"stdio","command":"node","args":["server.js"]}}` + case "session.mcp.oauth.authenticationStateChanged": + return `{"serverName":"rpc-server","refreshSessionToken":true}` + case "session.mcp.oauth.respond": + return `{"requestId":"oauth-request-1"}` + case "session.mcp.resources.list": + return `{"serverName":"rpc-server","cursor":"resource-cursor"}` + case "session.mcp.resources.listTemplates": + return `{"serverName":"rpc-server","cursor":"template-cursor"}` + case "session.mcp.resources.read": + return `{"serverName":"rpc-server","uri":"file://rpc/resource.txt"}` + case "session.model.setAllowedModels": + return `{"allowedModels":["model-a","model-b"]}` + case "session.model.switchAutoTier": + return `{"autoTier":"intelligence"}` + case "session.sandbox.disableForSession": + return `{"requestId":"sandbox-request-1"}` + case "session.abort": + return `{"reason":"user"}` + case "session.interruptMainTurn": + return `{"flushQueued":true}` + case "session.log": + return `{"message":"RPC log","level":"warning","type":"rpc","ephemeral":true,"url":"https://example.test/rpc","tip":"Inspect the RPC."}` + case "session.tasks.register": + return `{"type":"client","clientTaskId":"client-task-1","description":"RPC task","cancellable":true,"displayName":"RPC Task"}` + case "session.tasks.update": + return `{"id":"task-1","sequence":1,"update":{"kind":"progress","message":"Halfway","percentage":50,"phase":"work","status":"running"}}` + case "session.tools.execute": + return `{"name":"rpc_tool","arguments":{"value":"input"},"toolCallId":"tool-call-1"}` + case "session.tools.getBuiltinDescriptors": + return `{"reduceUserIntervention":true,"includeAuthor":true,"skillEmbeddingEnabled":false,"shellConfig":{"displayName":"PowerShell","shellType":"powershell","shellToolName":"shell","listShellsToolName":"list_shells","readShellToolName":"read_shell","stopShellToolName":"stop_shell","descriptionLines":["Runs shell commands."]},"shellSupportsPowerShell7Syntax":true,"shellTimeoutMs":1234,"backgroundTaskNotificationsEnabled":true}` + case "session.tools.set": + return `{"tools":[{"name":"rpc_external","title":"RPC External","description":"External RPC tool","parameters":{"type":"object"},"isTerminal":false,"overridesBuiltInTool":false,"skipPermission":true}]}` + case "session.tools.taskCompleteEventData": + return `{"toolArgs":{"objectiveId":17},"finalResult":{"resultType":"success","textResultForLlm":"RPC task complete","sessionLog":"Completion logged."}}` + case "session.workspaces.updateMetadata": + return `{"context":{"owner":"rpc-test"},"name":"Updated RPC workspace"}` + case "session.workspaces.ensure": + return `{"context":{"owner":"rpc-test"}}` + case "session.workspaces.statFile": + return `{"path":"folder/file.txt"}` + case "session.workspaces.createDirectory": + return `{"path":"folder/nested","recursive":true}` + case "session.workspaces.renamePath": + return `{"source":"folder/file.txt","destination":"folder/renamed.txt"}` + case "session.workspaces.removePath": + return `{"path":"folder","recursive":true,"force":true}` + case "session.workspaces.addSummary": + return `{"title":"RPC summary","content":"Summary content"}` + case "session.workspaces.truncateSummaries": + return `{"keepCount":2}` + default: + return "" + } +} + +func generatedRPCResponse(wire string) string { + switch wire { + case "hooks.discover": + return `{"hooks":[],"warnings":["rpc-warning"],"errors":[]}` + case "llmInference.setProvider": + return `{"success":true}` + case "managedSettings.read": + return `{"settingsJson":{"policy":"strict"}}` + case "mcp.planInstall": + return `{"kind":"unavailable","message":"The host does not provide installation.","reason":"host-not-available"}` + case "models.getBuiltInCatalog": + return `{"models":[{"id":"built-in-model"}]}` + case "sessions.getClientMetadata": + return `[{"status":"ok","sessionId":"persisted-session","metadata":{"rpc/key":"rpc-value"}}]` + case "session.contentExclusion.checkPaths": + return `{"available":true,"checks":[{"path":"/tmp/rpc-workspace/file.txt","excluded":false}]}` + case "session.debug.collectLogs": + return `{"kind":"directory","path":"/tmp/rpc-debug","entries":[{"bundlePath":"host/diagnostic.txt","sizeBytes":123,"source":"additional"}],"skippedEntries":[{"bundlePath":"host/missing.txt","path":"/tmp/missing.txt","reason":"not found"}]}` + case "session.factory.run", "session.factory.getRun": + return `{"runId":"factory-run-1","status":"running","attempt":1,"result":{"value":"running"},"snapshot":{"step":1}}` + case "session.factory.pause": + return `{"runId":"factory-run-1","status":"paused","attempt":1,"reason":"caller requested pause","snapshot":{"step":2}}` + case "session.factory.resume": + return `{"factoryName":"rpc-factory","run":{"runId":"factory-run-1","status":"running","attempt":2,"snapshot":{"step":3}}}` + case "session.factory.agent": + return `{"result":{"answer":"agent-result"}}` + case "session.factory.journal.get": + return `{"hit":true,"resultJson":{"checkpoint":7}}` + case "session.history.clearContext": + return `{"messagesCleared":4}` + case "session.limitPrediction.predict": + return `{"kind":"unavailable","reason":"insufficient-data"}` + case "session.mcp.moveLoadingToBackground": + return `{"movedToBackground":true}` + case "session.mcp.oauth.respond": + return `{"success":true}` + case "session.mcp.resources.list": + return `{"nextCursor":"resource-next","resources":[{"uri":"file://rpc/resource.txt","name":"RPC resource","description":"Resource description","mimeType":"text/plain","size":16,"title":"RPC Resource"}]}` + case "session.mcp.resources.listTemplates": + return `{"nextCursor":"template-next","resourceTemplates":[{"uriTemplate":"file://rpc/{name}","name":"RPC template","description":"Template description","mimeType":"text/plain","title":"RPC Template"}]}` + case "session.mcp.resources.read": + return `{"contents":[{"uri":"file://rpc/resource.txt","mimeType":"text/plain","text":"resource-content","_meta":{"audience":"assistant"}}]}` + case "session.metadata.getClientMetadata": + return `{"rpc/key":"rpc-value","rpc/other":"other-value"}` + case "session.model.setAllowedModels": + return `{"allowedModels":["model-a","model-b"],"effectiveAllowedModels":["model-a"],"fallbackModel":"model-a","modelId":"model-a"}` + case "session.model.switchAutoTier": + return `{"status":"applied","activatingAutoTier":"intelligence","effectiveAutoTier":"intelligence","supersededAutoTier":"balance"}` + case "session.sandbox.getEnforcementStatus": + return `{"required":true,"blocked":false,"reason":"managed-policy"}` + case "session.sandbox.disableForSession": + return `{"success":true,"enabled":false}` + case "session.abort": + return `{"success":true}` + case "session.interruptMainTurn": + return `{"interrupted":true}` + case "session.cancelAllBackgroundAgents": + return `3` + case "session.log": + return `{"eventId":"11111111-2222-3333-4444-555555555555"}` + case "session.tasks.register": + return `{"created":true,"reclaimed":false,"task":{"id":"task-1","type":"client","clientTaskId":"client-task-1","description":"RPC task","displayName":"RPC Task","activeTimeMs":500,"canCancel":true,"executionMode":"background","owner":{"displayName":"RPC owner","joinId":"join-1","kind":"sdk","participantId":"participant-1","presence":"connected","source":"rpc-test"},"sequence":0,"status":"running"}}` + case "session.tasks.update": + return `{"applied":true,"duplicate":false,"task":{"id":"task-1","type":"client","clientTaskId":"client-task-1","description":"RPC task","displayName":"RPC Task","activeTimeMs":500,"canCancel":true,"executionMode":"background","owner":{"displayName":"RPC owner","joinId":"join-1","kind":"sdk","participantId":"participant-1","presence":"connected","source":"rpc-test"},"sequence":1,"status":"running"}}` + case "session.tools.execute": + return `{"resultType":"success","textResultForLlm":"executed"}` + case "session.tools.getBuiltinDescriptors": + return `{"tools":[{"name":"rpc_builtin","description":"RPC built-in tool","hasSummariseIntention":true,"inputSchema":{"type":"object"},"instructions":"Use the RPC built-in.","isTerminal":false,"safeForTelemetry":true,"title":"RPC Built-in","type":"test"}]}` + case "session.tools.taskCompleteEventData": + return `{"objectiveId":17,"outcome":"completed","reason":"completed","success":true,"summary":"RPC task complete"}` + case "session.workspaces.updateMetadata": + return `{"path":"/tmp/rpc-workspace","workspace":{"id":"workspace-1","cwd":"/tmp/rpc-workspace","name":"Updated RPC workspace","branch":"rpc-branch","client_name":"rpc-client","created_at":"2026-09-18T11:00:00Z","git_root":"/tmp/rpc-workspace","remote_steerable":true}}` + case "session.workspaces.ensure": + return `{"path":"/tmp/rpc-workspace","workspace":{"id":"workspace-1","cwd":"/tmp/rpc-workspace","name":"RPC workspace"}}` + case "session.workspaces.statFile": + return `{"birthtimeMs":1000,"isDirectory":false,"isFile":true,"mtimeMs":2000,"size":42}` + case "session.workspaces.addSummary": + return `{"summary":{"number":3,"title":"RPC summary","content":"Summary content"},"workspace":{"id":"workspace-1","cwd":"/tmp/rpc-workspace","name":"RPC workspace"}}` + case "session.workspaces.truncateSummaries": + return `{"path":"/tmp/rpc-workspace","workspace":{"id":"workspace-1","cwd":"/tmp/rpc-workspace","name":"Truncated RPC workspace"}}` + default: + return `{}` + } +} + +func assertGeneratedRPCResult(t *testing.T, wire string, result any) { + t.Helper() + if wire == "sessions.getClientMetadata" { + entries, ok := result.(*rpc.SessionsGetClientMetadataResult) + if !ok || len(*entries) != 1 { + t.Fatalf("%s result = %#v, want one metadata entry", wire, result) + } + entry, ok := (*entries)[0].(*rpc.SessionsClientMetadataEntryOk) + if !ok || entry.SessionID != "persisted-session" || entry.Metadata["rpc/key"] != "rpc-value" { + t.Fatalf("%s result entry = %#v", wire, (*entries)[0]) + } + return + } + expectedJSON := generatedRPCResponse(wire) + if expectedJSON == "{}" { + return + } + var expected any + if err := json.Unmarshal([]byte(expectedJSON), &expected); err != nil { + t.Fatal(err) + } + actualJSON, err := json.Marshal(result) + if err != nil { + t.Fatalf("Marshal %s typed result: %v", wire, err) + } + var actual any + if err := json.Unmarshal(actualJSON, &actual); err != nil { + t.Fatalf("Decode %s typed result %s: %v", wire, actualJSON, err) + } + assertJSONSubset(t, wire+" result", expected, actual) +} + +func assertJSONSubset(t *testing.T, label string, expected, actual any) { + t.Helper() + switch expected := expected.(type) { + case map[string]any: + actual, ok := actual.(map[string]any) + if !ok { + t.Fatalf("%s = %#v, want object", label, actual) + } + for key, value := range expected { + actualValue, exists := actual[key] + if !exists { + t.Fatalf("%s missing %q in %#v", label, key, actual) + } + assertJSONSubset(t, label+"."+key, value, actualValue) + } + case []any: + actual, ok := actual.([]any) + if !ok || len(actual) != len(expected) { + t.Fatalf("%s = %#v, want %d items", label, actual, len(expected)) + } + for i := range expected { + assertJSONSubset(t, label, expected[i], actual[i]) + } + default: + if !reflect.DeepEqual(expected, actual) { + t.Fatalf("%s = %#v, want %#v", label, actual, expected) + } + } +} + +func methodName(signature any) string { + value := reflect.ValueOf(signature) + function := runtime.FuncForPC(value.Pointer()) + if function == nil { + panic("method expression has no runtime function") + } + name := function.Name() + if dot := strings.LastIndexByte(name, '.'); dot >= 0 { + return strings.TrimSuffix(name[dot+1:], "-fm") + } + panic("unexpected method expression name " + name) +} + +func testContextWithTimeout(t *testing.T, timeout time.Duration) (context.Context, context.CancelFunc) { + t.Helper() + return context.WithTimeout(t.Context(), timeout) +} + +type generatedRPCFixture struct { + client *copilot.Client + session *copilot.Session + server *jsonrpc2.Client + conn net.Conn +} + +func newGeneratedRPCFixture(t *testing.T, ctx context.Context) *generatedRPCFixture { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = listener.Close() }) + + type serverConnection struct { + server *jsonrpc2.Client + conn net.Conn + } + ready := make(chan serverConnection, 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":"generated-rpc-surface"}`, + "session.options.update": `{"success":true}`, + "session.detach": `{"success":true}`, + } { + server.SetRequestHandler(method, func(json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return json.RawMessage(result), nil + }) + } + server.Start() + ready <- serverConnection{server: server, conn: conn} + }() + + client := copilot.NewClient(&copilot.ClientOptions{ + Connection: copilot.URIConnection{URL: listener.Addr().String()}, + }) + t.Cleanup(client.ForceStop) + session, err := client.CreateSession(ctx, &copilot.SessionConfig{ + SessionID: "generated-rpc-surface", + OnPermissionRequest: copilot.PermissionHandler.ApproveAll, + }) + if err != nil { + t.Fatal(err) + } + + select { + case connection := <-ready: + return &generatedRPCFixture{client: client, session: session, server: connection.server, conn: connection.conn} + case <-ctx.Done(): + t.Fatal(ctx.Err()) + return nil + } +} diff --git a/go/internal/e2e/scenario_testing_cloud_e2e_test.go b/go/internal/e2e/scenario_testing_cloud_e2e_test.go new file mode 100644 index 0000000000..dcb42e87f8 --- /dev/null +++ b/go/internal/e2e/scenario_testing_cloud_e2e_test.go @@ -0,0 +1,480 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package e2e + +import ( + "bufio" + "encoding/json" + "fmt" + "io" + "net" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/github/copilot-sdk/go" + "github.com/github/copilot-sdk/go/internal/jsonrpc2" + "github.com/github/copilot-sdk/go/rpc" +) + +func TestScenarioTestingCloudE2E(t *testing.T) { + t.Run("notifies steerability before first send", func(t *testing.T) { + var mu sync.Mutex + var events []copilot.SessionEvent + messageID := "message-1" + fixture := newGeneratedRPCFixture(t, t.Context()) + fixture.server.SetRequestHandler("session.remote.notifySteerableChanged", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + events = append(events, scenarioEvent("remote-1", &copilot.SessionRemoteSteerableChangedData{ + RemoteSteerable: true, + })) + mu.Unlock() + return json.RawMessage(`{}`), nil + }) + fixture.server.SetRequestHandler("session.send", func(req json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + Prompt string `json:"prompt"` + } + if err := json.Unmarshal(req, ¶ms); err != nil { + t.Errorf("unmarshal session.send: %v", err) + } + mu.Lock() + events = append(events, scenarioEvent("message-1", &copilot.UserMessageData{ + Content: params.Prompt, + MessageID: &messageID, + TransformedContent: ¶ms.Prompt, + })) + mu.Unlock() + return mustJSON(t, map[string]any{"messageId": messageID}), nil + }) + fixture.server.SetRequestHandler("session.getMessages", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + defer mu.Unlock() + return mustJSON(t, map[string]any{"events": append([]copilot.SessionEvent(nil), events...)}), nil + }) + + session := fixture.session + + if _, err := session.RPC.Remote.NotifySteerableChanged(t.Context(), &rpc.RemoteNotifySteerableChangedRequest{ + RemoteSteerable: true, + }); err != nil { + t.Fatalf("NotifySteerableChanged failed: %v", err) + } + const prompt = "SCENARIO_STEERABLE_FIRST_SEND" + if _, err := session.Send(t.Context(), copilot.MessageOptions{Prompt: prompt}); err != nil { + t.Fatalf("Send failed: %v", err) + } + + got, err := session.GetEvents(t.Context()) + if err != nil { + t.Fatalf("GetEvents failed: %v", err) + } + if len(got) != 2 { + t.Fatalf("Expected two persisted events, got %d: %#v", len(got), got) + } + remote, ok := got[0].Data.(*copilot.SessionRemoteSteerableChangedData) + if !ok || !remote.RemoteSteerable { + t.Fatalf("Expected first event to persist remote steerability, got %#v", got[0].Data) + } + message, ok := got[1].Data.(*copilot.UserMessageData) + if !ok || message.TransformedContent == nil || !strings.Contains(*message.TransformedContent, prompt) { + t.Fatalf("Expected second event to contain first message, got %#v", got[1].Data) + } + }) + + t.Run("routes first event for server assigned session id", func(t *testing.T) { + fixture := newAssignedCloudSessionFixture(t) + defer fixture.Close() + + client := copilot.NewClient(&copilot.ClientOptions{ + Connection: copilot.URIConnection{URL: fixture.URL()}, + }) + if err := client.Start(t.Context()); err != nil { + t.Fatalf("Start failed: %v", err) + } + defer client.Stop() + + firstEvent := make(chan copilot.SessionEvent, 1) + session, err := client.CreateSession(t.Context(), &copilot.SessionConfig{ + Cloud: &copilot.CloudSessionOptions{ + Repository: &copilot.CloudSessionRepository{ + Owner: "github", + Name: "copilot-sdk", + Branch: "main", + }, + }, + OnEvent: func(event copilot.SessionEvent) { + select { + case firstEvent <- event: + default: + } + }, + }) + if err != nil { + t.Fatalf("CreateSession failed: %v", err) + } + defer session.Disconnect() + + if session.SessionID != "server-assigned-cloud-session" { + t.Fatalf("Expected server-assigned id, got %q", session.SessionID) + } + select { + case event := <-firstEvent: + start, ok := event.Data.(*copilot.SessionStartData) + if !ok { + t.Fatalf("Expected session.start event, got %T", event.Data) + } + if start.SessionID != session.SessionID { + t.Fatalf("Expected event session id %q, got %q", session.SessionID, start.SessionID) + } + case <-t.Context().Done(): + t.Fatal("Test context ended before first cloud event was routed") + } + + create := fixture.CreateRequest() + if _, ok := create["sessionId"]; ok { + t.Fatalf("Cloud session.create unexpectedly sent sessionId: %#v", create) + } + cloud, ok := create["cloud"].(map[string]any) + if !ok { + t.Fatalf("Expected cloud request object, got %#v", create["cloud"]) + } + repository, ok := cloud["repository"].(map[string]any) + if !ok || repository["owner"] != "github" || repository["name"] != "copilot-sdk" || repository["branch"] != "main" { + t.Fatalf("Unexpected cloud repository: %#v", cloud["repository"]) + } + }) + + t.Run("resumes using runtime id returned by cloud connect", func(t *testing.T) { + var mu sync.Mutex + var calls []capturedRPCRequest + fixture := newGeneratedRPCFixture(t, t.Context()) + setHandler := func(method string, result any) { + fixture.server.SetRequestHandler(method, func(req json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + calls = append(calls, capturedRPCRequest{Method: method, Request: append(json.RawMessage(nil), req...)}) + mu.Unlock() + return mustJSON(t, result), nil + }) + } + setHandler("sessions.connect", remoteSessionConnection("github/copilot-sdk#123")) + setHandler("session.resume", map[string]any{"sessionId": "runtime-session-id", "workspacePath": "C:\\workspace"}) + client := fixture.client + + connection, err := client.RPC.Sessions.Connect(t.Context(), &rpc.ConnectRemoteSessionParams{ + SessionID: "cloud-control-session", + }) + if err != nil { + t.Fatalf("Sessions.Connect failed: %v", err) + } + if connection.SessionID != "runtime-session-id" || connection.Metadata.SessionID != "runtime-session-id" { + t.Fatalf("Unexpected runtime ids: %#v", connection) + } + if connection.Metadata.ResourceID == nil || *connection.Metadata.ResourceID != "github/copilot-sdk#123" { + t.Fatalf("Unexpected resource id: %#v", connection.Metadata.ResourceID) + } + + resumed, err := client.ResumeSession(t.Context(), connection.SessionID, nil) + if err != nil { + t.Fatalf("ResumeSession failed: %v", err) + } + defer resumed.Disconnect() + if resumed.SessionID != connection.SessionID { + t.Fatalf("Expected resumed id %q, got %q", connection.SessionID, resumed.SessionID) + } + + mu.Lock() + defer mu.Unlock() + assertCapturedSessionID(t, calls, "sessions.connect", "cloud-control-session") + assertCapturedSessionID(t, calls, "session.resume", "runtime-session-id") + }) + + t.Run("exposes cloud resource mismatch before resume", func(t *testing.T) { + var mu sync.Mutex + var methods []string + fixture := newGeneratedRPCFixture(t, t.Context()) + fixture.server.SetRequestHandler("sessions.connect", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + methods = append(methods, "sessions.connect") + mu.Unlock() + return mustJSON(t, remoteSessionConnection("github/other-repository#456")), nil + }) + fixture.server.SetRequestHandler("session.resume", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + methods = append(methods, "session.resume") + mu.Unlock() + return json.RawMessage(`{"sessionId":"runtime-session-id"}`), nil + }) + client := fixture.client + + connection, err := client.RPC.Sessions.Connect(t.Context(), &rpc.ConnectRemoteSessionParams{ + SessionID: "cloud-control-session", + }) + if err != nil { + t.Fatalf("Sessions.Connect failed: %v", err) + } + if connection.Metadata.ResourceID == nil || *connection.Metadata.ResourceID == "github/copilot-sdk#123" { + t.Fatalf("Expected mismatched resource id, got %#v", connection.Metadata.ResourceID) + } + + mu.Lock() + defer mu.Unlock() + for _, method := range methods { + if method == "session.resume" { + t.Fatal("Resource mismatch must be exposed before session.resume") + } + } + }) +} + +func mustJSON(t *testing.T, value any) json.RawMessage { + t.Helper() + result, err := json.Marshal(value) + if err != nil { + t.Fatalf("Marshal fixture result failed: %v", err) + } + return result +} + +func scenarioEvent(id string, data copilot.SessionEventData) copilot.SessionEvent { + return copilot.SessionEvent{ + Data: data, + ID: id, + Timestamp: time.Date(2026, 9, 17, 19, 0, 0, 0, time.UTC), + } +} + +func remoteSessionConnection(resourceID string) map[string]any { + return map[string]any{ + "sessionId": "runtime-session-id", + "metadata": map[string]any{ + "kind": "coding-agent", + "modifiedTime": "2026-09-17T20:00:00Z", + "name": "Cloud task", + "repository": map[string]any{ + "branch": "main", + "name": "copilot-sdk", + "owner": "github", + }, + "resourceId": resourceID, + "sessionId": "runtime-session-id", + "startTime": "2026-09-17T19:00:00Z", + "state": "active", + }, + } +} + +func assertCapturedSessionID(t *testing.T, calls []capturedRPCRequest, method, want string) { + t.Helper() + for _, call := range calls { + if call.Method != method { + continue + } + var params map[string]any + if err := json.Unmarshal(call.Request, ¶ms); err != nil { + t.Fatalf("Unmarshal %s request failed: %v", method, err) + } + if got := params["sessionId"]; got != want { + t.Fatalf("%s sessionId: got %#v, want %q", method, got, want) + } + return + } + t.Fatalf("Did not capture %s", method) +} + +type capturedRPCRequest struct { + Method string + Request json.RawMessage +} + +type assignedCloudSessionFixture struct { + t *testing.T + listener net.Listener + done chan struct{} + mu sync.Mutex + createRequest map[string]any +} + +func newAssignedCloudSessionFixture(t *testing.T) *assignedCloudSessionFixture { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("Listen failed: %v", err) + } + f := &assignedCloudSessionFixture{ + t: t, + listener: listener, + done: make(chan struct{}), + } + go f.serve() + return f +} + +func (f *assignedCloudSessionFixture) URL() string { + return "http://" + f.listener.Addr().String() +} + +func (f *assignedCloudSessionFixture) CreateRequest() map[string]any { + f.mu.Lock() + defer f.mu.Unlock() + result := make(map[string]any, len(f.createRequest)) + for key, value := range f.createRequest { + result[key] = value + } + return result +} + +func (f *assignedCloudSessionFixture) Close() { + _ = f.listener.Close() + <-f.done +} + +func (f *assignedCloudSessionFixture) serve() { + defer close(f.done) + conn, err := f.listener.Accept() + if err != nil { + return + } + defer conn.Close() + + reader := bufio.NewReader(conn) + for { + body, err := readRPCFrame(reader) + if err != nil { + if err != io.EOF { + f.t.Errorf("Read fake cloud RPC frame: %v", err) + } + return + } + var request struct { + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params"` + } + if err := json.Unmarshal(body, &request); err != nil { + f.t.Errorf("Unmarshal fake cloud request: %v", err) + return + } + switch request.Method { + case "connect": + if err := writeRPCFrame(conn, map[string]any{ + "jsonrpc": "2.0", + "id": request.ID, + "result": map[string]any{ + "ok": true, + "protocolVersion": 3, + "version": "fake", + }, + }); err != nil { + f.t.Errorf("Write connect response: %v", err) + return + } + case "session.create": + var params map[string]any + if err := json.Unmarshal(request.Params, ¶ms); err != nil { + f.t.Errorf("Unmarshal session.create params: %v", err) + return + } + f.mu.Lock() + f.createRequest = params + f.mu.Unlock() + if err := writeRPCFrame(conn, map[string]any{ + "jsonrpc": "2.0", + "id": request.ID, + "result": map[string]any{ + "sessionId": "server-assigned-cloud-session", + "workspacePath": "C:\\cloud-workspace", + }, + }); err != nil { + f.t.Errorf("Write session.create response: %v", err) + return + } + event := scenarioEvent("start-1", &copilot.SessionStartData{ + CopilotVersion: "fake", + Producer: "scenario-test", + SessionID: "server-assigned-cloud-session", + StartTime: time.Date(2026, 9, 17, 19, 0, 0, 0, time.UTC), + Version: 1, + }) + if err := writeRPCFrame(conn, map[string]any{ + "jsonrpc": "2.0", + "method": "session.event", + "params": map[string]any{ + "sessionId": "server-assigned-cloud-session", + "event": event, + }, + }); err != nil { + f.t.Errorf("Write first session event: %v", err) + return + } + case "session.options.update", "session.detach": + if err := writeRPCFrame(conn, map[string]any{ + "jsonrpc": "2.0", + "id": request.ID, + "result": map[string]any{}, + }); err != nil { + f.t.Errorf("Write %s response: %v", request.Method, err) + return + } + default: + if len(request.ID) == 0 { + continue + } + if err := writeRPCFrame(conn, map[string]any{ + "jsonrpc": "2.0", + "id": request.ID, + "error": map[string]any{ + "code": -32601, + "message": "unexpected method " + request.Method, + }, + }); err != nil { + f.t.Errorf("Write error response: %v", err) + return + } + } + } +} + +func readRPCFrame(reader *bufio.Reader) ([]byte, error) { + contentLength := -1 + for { + line, err := reader.ReadString('\n') + if err != nil { + return nil, err + } + line = strings.TrimSpace(line) + if line == "" { + break + } + name, value, ok := strings.Cut(line, ":") + if !ok { + return nil, fmt.Errorf("invalid RPC header %q", line) + } + if name == "Content-Length" { + contentLength, err = strconv.Atoi(strings.TrimSpace(value)) + if err != nil { + return nil, fmt.Errorf("parse content length: %w", err) + } + } + } + if contentLength < 0 { + return nil, fmt.Errorf("missing Content-Length") + } + body := make([]byte, contentLength) + _, err := io.ReadFull(reader, body) + return body, err +} + +func writeRPCFrame(writer io.Writer, message any) error { + body, err := json.Marshal(message) + if err != nil { + return err + } + if _, err := fmt.Fprintf(writer, "Content-Length: %d\r\n\r\n", len(body)); err != nil { + return err + } + _, err = writer.Write(body) + return err +} diff --git a/go/internal/e2e/scenario_testing_control_state_e2e_test.go b/go/internal/e2e/scenario_testing_control_state_e2e_test.go new file mode 100644 index 0000000000..5047038948 --- /dev/null +++ b/go/internal/e2e/scenario_testing_control_state_e2e_test.go @@ -0,0 +1,133 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package e2e + +import ( + "encoding/json" + "fmt" + "net" + "sync/atomic" + "testing" + + "github.com/github/copilot-sdk/go" + "github.com/github/copilot-sdk/go/internal/jsonrpc2" +) + +func TestScenarioTestingControlStateE2E(t *testing.T) { + t.Run("reports processing while scenario tool is running", func(t *testing.T) { + fixture := newGeneratedRPCFixture(t, t.Context()) + fixture.server.SetRequestHandler("session.create", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return json.RawMessage(`{"sessionId":"processing-session"}`), nil + }) + + var processing atomic.Bool + fixture.server.SetRequestHandler("session.metadata.isProcessing", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return mustJSON(t, map[string]any{"processing": processing.Load()}), nil + }) + fixture.server.SetRequestHandler("session.metadata.activity", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + active := processing.Load() + return mustJSON(t, map[string]any{ + "hasActiveWork": active, + "abortable": active, + }), nil + }) + + toolStarted := make(chan struct{}) + releaseTool := make(chan struct{}) + toolCompleted := make(chan struct{}) + fixture.server.SetRequestHandler("session.tools.handlePendingToolCall", func(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + RequestID string `json:"requestId"` + Result any `json:"result"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + t.Errorf("Unmarshal tool result failed: %v", err) + } + if params.RequestID != "processing-request" || params.Result == nil { + t.Errorf("Unexpected tool completion: %#v", params) + } + processing.Store(false) + close(toolCompleted) + return json.RawMessage(`{"success":true}`), nil + }) + + session, err := fixture.client.CreateSession(t.Context(), &copilot.SessionConfig{ + SessionID: "processing-session", + Tools: []copilot.Tool{ + copilot.DefineTool("wait_for_scenario_control", "Waits for the scenario controller", + func(_ struct{}, _ copilot.ToolInvocation) (string, error) { + close(toolStarted) + <-releaseTool + return "SCENARIO_CONTROL_DONE", nil + }), + }, + }) + if err != nil { + t.Fatalf("CreateSession failed: %v", err) + } + defer session.Disconnect() + + assertProcessingState(t, session, false) + + processing.Store(true) + writeScenarioNotification(t, fixture.conn, "session.event", map[string]any{ + "sessionId": session.SessionID, + "event": scenarioEvent("tool-request", &copilot.ExternalToolRequestedData{ + Arguments: map[string]any{}, + RequestID: "processing-request", + SessionID: session.SessionID, + ToolCallID: "processing-tool-call", + ToolName: "wait_for_scenario_control", + }), + }) + + select { + case <-toolStarted: + case <-t.Context().Done(): + t.Fatal("Test context ended before tool handler started") + } + assertProcessingState(t, session, true) + + close(releaseTool) + select { + case <-toolCompleted: + case <-t.Context().Done(): + t.Fatal("Test context ended before tool completion was handled") + } + assertProcessingState(t, session, false) + }) +} + +func writeScenarioNotification(t *testing.T, conn net.Conn, method string, params any) { + t.Helper() + message, err := json.Marshal(map[string]any{ + "jsonrpc": "2.0", + "method": method, + "params": params, + }) + if err != nil { + t.Fatal(err) + } + frame := append([]byte(fmt.Sprintf("Content-Length: %d\r\n\r\n", len(message))), message...) + if _, err := conn.Write(frame); err != nil { + t.Fatal(err) + } +} + +func assertProcessingState(t *testing.T, session *copilot.Session, want bool) { + t.Helper() + state, err := session.RPC.Metadata.IsProcessing(t.Context()) + if err != nil { + t.Fatalf("Metadata.IsProcessing failed: %v", err) + } + if state.Processing != want { + t.Fatalf("Processing = %t, want %t", state.Processing, want) + } + activity, err := session.RPC.Metadata.Activity(t.Context()) + if err != nil { + t.Fatalf("Metadata.Activity failed: %v", err) + } + if activity.HasActiveWork != want || activity.Abortable != want { + t.Fatalf("Activity = %#v, want active/abortable %t", activity, want) + } +} diff --git a/go/internal/e2e/scenario_testing_sends_e2e_test.go b/go/internal/e2e/scenario_testing_sends_e2e_test.go new file mode 100644 index 0000000000..7c46c33662 --- /dev/null +++ b/go/internal/e2e/scenario_testing_sends_e2e_test.go @@ -0,0 +1,308 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package e2e + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sync" + "sync/atomic" + "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 TestScenarioTestingSendsE2E(t *testing.T) { + t.Run("should send complete scenario message wire shape", func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + f := newGeneratedRPCFixture(t, ctx) + + var captured map[string]any + f.server.SetRequestHandler("session.send", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + if err := json.Unmarshal(params, &captured); err != nil { + return nil, &jsonrpc2.Error{Code: -32000, Message: err.Error()} + } + return json.RawMessage(`{"messageId":"scenario-client-message"}`), nil + }) + + canvasID := "diff" + instanceID := "diff-17" + blobData := "QVBQX0JMT0I=" + blobName := "scenario-wire-blob.txt" + messageID, err := f.session.Send(ctx, copilot.MessageOptions{ + Prompt: "Use the hidden scenario context.", + DisplayPrompt: "Review selected scenario context", + Mode: "enqueue", + AgentMode: copilot.AgentModeInteractive, + Source: copilot.MessageSourceAgent("scenario-client"), + RequestHeaders: map[string]string{ + "x-scenario-request": "wire-shape", + }, + Attachments: []copilot.Attachment{ + rpc.AttachmentFile{ + DisplayName: "scenario-wire-file.txt", + Path: `Q:\scenario-wire-file.txt`, + LineRange: &rpc.AttachmentFileLineRange{Start: 3, End: 9}, + }, + rpc.AttachmentDirectory{ + DisplayName: "scenario-wire-directory", + Path: `Q:\scenario-wire-directory`, + }, + rpc.AttachmentSelection{ + DisplayName: "Program.cs", + FilePath: `Q:\Program.cs`, + Text: "SCENARIO_SELECTION", + Selection: rpc.AttachmentSelectionDetails{ + Start: rpc.AttachmentSelectionDetailsStart{Line: 16, Character: 0}, + End: rpc.AttachmentSelectionDetailsEnd{Line: 16, Character: 13}, + }, + }, + rpc.AttachmentGitHubReference{ + Number: 610, + ReferenceType: rpc.AttachmentGitHubReferenceTypePr, + State: "open", + Title: "Scenario-shaped E2E coverage", + URL: "https://github.com/github/copilot-sdk/pull/610", + }, + rpc.AttachmentBlob{ + Data: &blobData, + MIMEType: "text/plain", + DisplayName: &blobName, + }, + rpc.AttachmentExtensionContext{ + CapturedAt: time.Date(2026, 9, 17, 20, 0, 0, 0, time.UTC), + ExtensionID: "scenario-client:code-review", + CanvasID: &canvasID, + InstanceID: &instanceID, + Title: "Selected change", + Payload: map[string]any{"selection": "SCENARIO_SELECTION", "line": float64(17)}, + }, + }, + }) + if err != nil { + t.Fatal(err) + } + if messageID != "scenario-client-message" { + t.Fatalf("Message ID = %q", messageID) + } + + assertJSONSubset(t, "session.send request", map[string]any{ + "sessionId": f.session.SessionID, + "prompt": "Use the hidden scenario context.", + "displayPrompt": "Review selected scenario context", + "mode": "enqueue", + "agentMode": "interactive", + "source": "agent-scenario-client", + "requestHeaders": map[string]any{ + "x-scenario-request": "wire-shape", + }, + }, captured) + + attachments, ok := captured["attachments"].([]any) + if !ok || len(attachments) != 6 { + t.Fatalf("attachments = %#v", captured["attachments"]) + } + wantTypes := []string{"file", "directory", "selection", "github_reference", "blob", "extension_context"} + for i, want := range wantTypes { + attachment := attachments[i].(map[string]any) + if attachment["type"] != want { + t.Fatalf("attachment %d type = %#v, want %q", i, attachment["type"], want) + } + } + file := attachments[0].(map[string]any) + assertJSONSubset(t, "file attachment", map[string]any{ + "path": `Q:\scenario-wire-file.txt`, + "lineRange": map[string]any{ + "start": float64(3), + "end": float64(9), + }, + }, file) + extension := attachments[5].(map[string]any) + assertJSONSubset(t, "extension attachment", map[string]any{ + "extensionId": "scenario-client:code-review", + "canvasId": "diff", + "instanceId": "diff-17", + "payload": map[string]any{ + "selection": "SCENARIO_SELECTION", + "line": float64(17), + }, + }, extension) + }) + + for _, mode := range []string{"", "enqueue", "immediate"} { + name := mode + if name == "" { + name = "default" + } + t.Run("should not invoke send when scenario cancels before dispatch "+name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + f := newGeneratedRPCFixture(t, ctx) + var calls atomic.Int64 + f.server.SetRequestHandler("session.send", func(json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + calls.Add(1) + return json.RawMessage(`{"messageId":"unexpected"}`), nil + }) + + cancelled, cancelSend := context.WithCancel(ctx) + cancelSend() + _, err := f.session.Send(cancelled, copilot.MessageOptions{ + Prompt: "This message must never be invoked.", + DisplayPrompt: "Cancelled scenario message", + Mode: mode, + Source: copilot.MessageSourceAgent("scenario-client"), + }) + if !errors.Is(err, context.Canceled) { + t.Fatalf("Send error = %v, want context cancellation", err) + } + if calls.Load() != 0 { + t.Fatalf("session.send calls = %d, want 0", calls.Load()) + } + }) + + t.Run("should not replay scenario send after ambiguous transport loss "+name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + f := newGeneratedRPCFixture(t, ctx) + var calls atomic.Int64 + capturedRequests := make(chan map[string]any, 1) + f.server.SetRequestHandler("session.send", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + calls.Add(1) + var captured map[string]any + if err := json.Unmarshal(params, &captured); err != nil { + return nil, &jsonrpc2.Error{Code: -32000, Message: err.Error()} + } + capturedRequests <- captured + _ = f.conn.Close() + return nil, nil + }) + + _, err := f.session.Send(ctx, copilot.MessageOptions{ + Prompt: "AMBIGUOUS_SCENARIO_SEND", + DisplayPrompt: "Ambiguous scenario send", + Mode: mode, + Source: copilot.MessageSourceAgent("scenario-client"), + }) + if err == nil { + t.Fatal("Expected transport loss") + } + if errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("Send waited for its deadline instead of reporting transport loss: %v", err) + } + if calls.Load() != 1 { + t.Fatalf("session.send calls = %d, want 1", calls.Load()) + } + var captured map[string]any + select { + case captured = <-capturedRequests: + case <-ctx.Done(): + t.Fatalf("Timed out waiting for captured session.send request: %v", ctx.Err()) + } + if mode == "" { + if _, exists := captured["mode"]; exists { + t.Fatalf("Default mode should be omitted: %#v", captured) + } + } else if captured["mode"] != mode { + t.Fatalf("mode = %#v, want %q", captured["mode"], mode) + } + }) + } + + t.Run("should order idle queued and immediate scenario delivery", func(t *testing.T) { + f := newGeneratedRPCFixture(t, t.Context()) + var mu sync.Mutex + var events []copilot.SessionEvent + var sends int + f.server.SetRequestHandler("session.send", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var request struct { + Prompt string `json:"prompt"` + Mode string `json:"mode"` + } + if err := json.Unmarshal(params, &request); err != nil { + t.Errorf("Unmarshal session.send failed: %v", err) + } + mu.Lock() + defer mu.Unlock() + sends++ + messageID := fmt.Sprintf("scenario-message-%d", sends) + delivery := copilot.UserMessageDeliveryIdle + switch sends { + case 4, 5: + delivery = copilot.UserMessageDeliverySteering + case 6: + delivery = copilot.UserMessageDeliveryQueued + } + events = append(events, scenarioEvent(messageID, &copilot.UserMessageData{ + Content: request.Prompt, + Delivery: &delivery, + MessageID: &messageID, + TransformedContent: &request.Prompt, + })) + return mustJSON(t, map[string]any{"messageId": messageID}), nil + }) + f.server.SetRequestHandler("session.getMessages", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + mu.Lock() + defer mu.Unlock() + return mustJSON(t, map[string]any{"events": append([]copilot.SessionEvent(nil), events...)}), nil + }) + + send := func(prompt, mode string) string { + t.Helper() + result, err := f.session.Send(t.Context(), copilot.MessageOptions{ + Prompt: prompt, + Mode: mode, + Source: copilot.MessageSourceAgent("scenario-client"), + }) + if err != nil { + t.Fatalf("Send(%q, %q) failed: %v", prompt, mode, err) + } + return result + } + + idleEnqueueID := send("IDLE_ENQUEUE", "enqueue") + idleImmediateID := send("IDLE_IMMEDIATE", "immediate") + _ = send("START_BLOCKING_TURN", "") + steeringID := send("FIRST_STEERING", "immediate") + immediateBehindID := send("SECOND_IMMEDIATE", "immediate") + queuedID := send("FINAL_QUEUED", "enqueue") + + observed, err := f.session.GetEvents(t.Context()) + if err != nil { + t.Fatalf("GetEvents failed: %v", err) + } + find := func(id string) (int, *copilot.UserMessageData) { + t.Helper() + for i, event := range observed { + data, ok := event.Data.(*copilot.UserMessageData) + if ok && data.MessageID != nil && *data.MessageID == id { + return i, data + } + } + t.Fatalf("Did not find user.message for %q", id) + return -1, nil + } + assertDelivery := func(id string, want copilot.UserMessageDelivery) int { + t.Helper() + index, data := find(id) + if data.Delivery == nil || *data.Delivery != want { + t.Fatalf("Delivery for %q = %#v, want %q", id, data.Delivery, want) + } + return index + } + + assertDelivery(idleEnqueueID, copilot.UserMessageDeliveryIdle) + assertDelivery(idleImmediateID, copilot.UserMessageDeliveryIdle) + steeringIndex := assertDelivery(steeringID, copilot.UserMessageDeliverySteering) + behindIndex := assertDelivery(immediateBehindID, copilot.UserMessageDeliverySteering) + queuedIndex := assertDelivery(queuedID, copilot.UserMessageDeliveryQueued) + if steeringIndex >= behindIndex || behindIndex >= queuedIndex { + t.Fatalf("Unexpected delivery order: steering=%d behind=%d queued=%d", steeringIndex, behindIndex, queuedIndex) + } + }) +} diff --git a/go/internal/e2e/scenario_testing_server_control_e2e_test.go b/go/internal/e2e/scenario_testing_server_control_e2e_test.go new file mode 100644 index 0000000000..656bcf82f2 --- /dev/null +++ b/go/internal/e2e/scenario_testing_server_control_e2e_test.go @@ -0,0 +1,612 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package e2e + +import ( + "encoding/json" + "reflect" + "sync" + "testing" + + "github.com/github/copilot-sdk/go/internal/jsonrpc2" + "github.com/github/copilot-sdk/go/rpc" +) + +func TestScenarioTestingServerControlE2E(t *testing.T) { + t.Run("searches server catalog with category contract", func(t *testing.T) { + tests := []struct { + name string + kinds []rpc.CatalogCandidateKind + capabilities []string + }{ + { + name: "all", + kinds: []rpc.CatalogCandidateKind{rpc.CatalogCandidateKindMCPServer, rpc.CatalogCandidateKindAiSkill}, + capabilities: []string{"mcp-server-card", "ai-skill-discovery"}, + }, + { + name: "mcp", + kinds: []rpc.CatalogCandidateKind{rpc.CatalogCandidateKindMCPServer}, + capabilities: []string{"mcp-server-card"}, + }, + { + name: "skills", + kinds: []rpc.CatalogCandidateKind{rpc.CatalogCandidateKindAiSkill}, + capabilities: []string{"ai-skill-discovery"}, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + fixture := newGeneratedRPCFixture(t, t.Context()) + var captured map[string]any + fixture.server.SetRequestHandler("catalog.search", func(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + if err := json.Unmarshal(request, &captured); err != nil { + t.Errorf("Unmarshal catalog request failed: %v", err) + } + return mustJSON(t, map[string]any{ + "kind": "succeeded", + "candidates": []any{}, + "negotiated": map[string]any{ + "grantedCapabilities": test.capabilities, + "runtimeProtocolVersion": 3, + }, + "searchId": "scenario-search", + "truncated": false, + }), nil + }) + + limit := int32(50) + result, err := fixture.client.RPC.Catalog.Search(t.Context(), &rpc.CatalogSearchRequest{ + Contract: rpc.CatalogClientContract{ + ProtocolVersion: 3, + RequiredCapabilities: test.capabilities, + }, + Kinds: test.kinds, + Limit: &limit, + Query: "scenario search", + }) + if err != nil { + t.Fatalf("Catalog.Search failed: %v", err) + } + succeeded, ok := result.(*rpc.CatalogSearchSucceeded) + if !ok { + t.Fatalf("Expected CatalogSearchSucceeded, got %T", result) + } + if len(succeeded.Candidates) != 0 || succeeded.SearchID != "scenario-search" || succeeded.Truncated { + t.Fatalf("Unexpected catalog result: %#v", succeeded) + } + if succeeded.Negotiated.RuntimeProtocolVersion != 3 || + !reflect.DeepEqual(succeeded.Negotiated.GrantedCapabilities, catalogCapabilities(test.capabilities)) { + t.Fatalf("Unexpected negotiated contract: %#v", succeeded.Negotiated) + } + + assertJSONSubset(t, "catalog.search", map[string]any{ + "query": "scenario search", + "limit": float64(50), + "contract": map[string]any{ + "protocolVersion": float64(3), + "requiredCapabilities": stringsToAny(test.capabilities), + }, + "kinds": catalogKindsToAny(test.kinds), + }, captured) + }) + } + }) + + t.Run("observes pages and cancels factory run", func(t *testing.T) { + fixture := newGeneratedRPCFixture(t, t.Context()) + var mu sync.Mutex + captured := map[string]map[string]any{} + setFactoryHandler := func(method string, result any) { + fixture.server.SetRequestHandler(method, func(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params map[string]any + if err := json.Unmarshal(request, ¶ms); err != nil { + t.Errorf("Unmarshal %s failed: %v", method, err) + } + mu.Lock() + captured[method] = params + mu.Unlock() + return mustJSON(t, result), nil + }) + } + setFactoryHandler("session.factory.listRuns", map[string]any{ + "runs": []any{map[string]any{ + "runId": "factory-run-1", + "factoryName": "scenario-factory", + "status": "running", + }}, + "oldestSeq": 7, + "newestSeq": 7, + "hasMoreNewer": false, + }) + setFactoryHandler("session.factory.getRunDetail", map[string]any{ + "runId": "factory-run-1", + "factoryName": "scenario-factory", + "status": "running", + "revision": 4, + }) + setFactoryHandler("session.factory.getRunProgress", map[string]any{ + "records": []any{map[string]any{ + "attempt": 1, + "kind": "log", + "phaseId": "verify", + "recordedAt": 1234, + "seq": 12, + "text": "Validation complete", + }}, + "revision": 4, + }) + setFactoryHandler("session.factory.cancel", map[string]any{ + "runId": "factory-run-1", + "status": "cancelled", + "reason": "cancelled by user", + }) + + after, before, limit := int64(3), int64(20), int32(10) + runs, err := fixture.session.RPC.Factory.ListRuns(t.Context(), &rpc.FactoryListRunsRequest{ + AfterSeq: &after, + BeforeSeq: &before, + Limit: &limit, + }) + if err != nil { + t.Fatalf("Factory.ListRuns failed: %v", err) + } + if len(runs.Runs) != 1 || runs.Runs[0].RunID != "factory-run-1" || + runs.Runs[0].FactoryName != "scenario-factory" || runs.Runs[0].Status != rpc.FactoryRunStatusRunning || + runs.OldestSeq == nil || *runs.OldestSeq != 7 || runs.NewestSeq == nil || *runs.NewestSeq != 7 || + runs.HasMoreNewer == nil || *runs.HasMoreNewer { + t.Fatalf("Unexpected factory run page: %#v", runs) + } + + detail, err := fixture.session.RPC.Factory.GetRunDetail(t.Context(), &rpc.FactoryGetRunRequest{RunID: "factory-run-1"}) + if err != nil { + t.Fatalf("Factory.GetRunDetail failed: %v", err) + } + if detail.RunID != "factory-run-1" || detail.FactoryName != "scenario-factory" || + detail.Status != rpc.FactoryRunStatusRunning || detail.Revision != 4 { + t.Fatalf("Unexpected factory detail: %#v", detail) + } + + progressAfter, progressBefore, progressLimit := int64(5), int64(20), int32(25) + phaseID := "verify" + progress, err := fixture.session.RPC.Factory.GetRunProgress(t.Context(), &rpc.FactoryGetRunProgressRequest{ + RunID: "factory-run-1", + PhaseID: &phaseID, + AfterSeq: &progressAfter, + BeforeSeq: &progressBefore, + Limit: &progressLimit, + }) + if err != nil { + t.Fatalf("Factory.GetRunProgress failed: %v", err) + } + if len(progress.Records) != 1 || progress.Records[0].Seq != 12 || + progress.Records[0].PhaseID == nil || *progress.Records[0].PhaseID != "verify" || + progress.Records[0].Kind != rpc.FactoryLogLineKindLog || + progress.Records[0].Text != "Validation complete" { + t.Fatalf("Unexpected factory progress: %#v", progress) + } + + cancelled, err := fixture.session.RPC.Factory.Cancel(t.Context(), &rpc.FactoryCancelRequest{RunID: "factory-run-1"}) + if err != nil { + t.Fatalf("Factory.Cancel failed: %v", err) + } + if cancelled.RunID != "factory-run-1" || cancelled.Status != rpc.FactoryRunStatusCancelled || + cancelled.Reason == nil || *cancelled.Reason != "cancelled by user" { + t.Fatalf("Unexpected cancelled factory run: %#v", cancelled) + } + + mu.Lock() + defer mu.Unlock() + assertJSONSubset(t, "session.factory.listRuns", map[string]any{ + "afterSeq": float64(3), + "beforeSeq": float64(20), + "limit": float64(10), + }, captured["session.factory.listRuns"]) + assertJSONSubset(t, "session.factory.getRunProgress", map[string]any{ + "runId": "factory-run-1", + "phaseId": "verify", + "afterSeq": float64(5), + "beforeSeq": float64(20), + "limit": float64(25), + }, captured["session.factory.getRunProgress"]) + }) + + t.Run("reads autopilot state and enables remote mode", func(t *testing.T) { + tests := []struct { + mode rpc.RemoteSessionMode + steerable bool + }{ + {mode: rpc.RemoteSessionModeOn, steerable: true}, + {mode: rpc.RemoteSessionModeExport, steerable: false}, + } + for _, test := range tests { + t.Run(string(test.mode), func(t *testing.T) { + fixture := newGeneratedRPCFixture(t, t.Context()) + fixture.server.SetRequestHandler("session.autopilotObjective.getState", func(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + return json.RawMessage(`{"state":{"id":17,"objective":"Ship the scenario.","status":"active","turnCount":3,"creditCountNanoAiu":"1250000000","creditLimit":{"credits":5,"creditsUsed":1.25,"creditsUsedNanoAiu":"1250000000"}}}`), nil + }) + var remoteRequest map[string]any + fixture.server.SetRequestHandler("session.remote.enable", func(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + if err := json.Unmarshal(request, &remoteRequest); err != nil { + t.Errorf("Unmarshal remote enable failed: %v", err) + } + return mustJSON(t, map[string]any{ + "remoteSteerable": test.steerable, + "url": "https://example.test/sessions/" + fixture.session.SessionID, + }), nil + }) + + state, err := fixture.session.RPC.AutopilotObjective.GetState(t.Context()) + if err != nil { + t.Fatalf("AutopilotObjective.GetState failed: %v", err) + } + objective := state.State + if objective == nil || objective.ID != 17 || objective.Objective != "Ship the scenario." || + objective.Status != rpc.AutopilotObjectiveStatusActive || objective.TurnCount != 3 || + objective.CreditCountNanoAiu != "1250000000" || objective.CreditLimit == nil || + objective.CreditLimit.Credits == nil || *objective.CreditLimit.Credits != 5 || + objective.CreditLimit.CreditsUsed != 1.25 || + objective.CreditLimit.CreditsUsedNanoAiu != "1250000000" { + t.Fatalf("Unexpected autopilot objective: %#v", objective) + } + + enabled, err := fixture.session.RPC.Remote.Enable(t.Context(), &rpc.RemoteEnableRequest{Mode: &test.mode}) + if err != nil { + t.Fatalf("Remote.Enable failed: %v", err) + } + expectedURL := "https://example.test/sessions/" + fixture.session.SessionID + if enabled.RemoteSteerable != test.steerable || enabled.URL == nil || *enabled.URL != expectedURL { + t.Fatalf("Unexpected remote result: %#v", enabled) + } + assertJSONSubset(t, "session.remote.enable", map[string]any{ + "sessionId": fixture.session.SessionID, + "mode": string(test.mode), + }, remoteRequest) + }) + } + }) + + t.Run("edits reorders duplicates removes and sends queued items", func(t *testing.T) { + fixture := newGeneratedRPCFixture(t, t.Context()) + queue := newScenarioQueue(t, fixture.server) + + if _, err := fixture.session.RPC.Queue.SetDrainPaused(t.Context(), &rpc.QueueSetDrainPausedRequest{Paused: true}); err != nil { + t.Fatalf("Queue.SetDrainPaused(true) failed: %v", err) + } + firstDisplay := "First visible prompt" + first, err := fixture.session.RPC.Queue.InsertAt(t.Context(), &rpc.QueueInsertAtRequest{ + Position: 0, + Message: rpc.QueueInsertMessage{ + Prompt: "First hidden prompt", + DisplayPrompt: &firstDisplay, + AgentMode: ptr(rpc.SendAgentModeInteractive), + }, + }) + if err != nil { + t.Fatalf("Queue.InsertAt(first) failed: %v", err) + } + secondDisplay := "Second visible prompt" + second, err := fixture.session.RPC.Queue.InsertAt(t.Context(), &rpc.QueueInsertAtRequest{ + Position: 1, + Message: rpc.QueueInsertMessage{ + Prompt: "Second hidden prompt", + DisplayPrompt: &secondDisplay, + AgentMode: ptr(rpc.SendAgentModePlan), + }, + }) + if err != nil { + t.Fatalf("Queue.InsertAt(second) failed: %v", err) + } + + updatedDisplay := "Updated visible prompt" + updated, err := fixture.session.RPC.Queue.UpdateText(t.Context(), &rpc.QueueUpdateTextRequest{ + ID: first.ID, + Prompt: "Updated hidden prompt", + DisplayPrompt: &updatedDisplay, + }) + if err != nil || !updated.Updated { + t.Fatalf("Queue.UpdateText: result=%#v err=%v", updated, err) + } + duplicate, err := fixture.session.RPC.Queue.DuplicateAt(t.Context(), &rpc.QueueDuplicateAtRequest{ID: first.ID}) + if err != nil || duplicate.ID == first.ID { + t.Fatalf("Queue.DuplicateAt: result=%#v err=%v", duplicate, err) + } + moved, err := fixture.session.RPC.Queue.MoveItem(t.Context(), &rpc.QueueMoveItemRequest{ + ID: second.ID, + ToPosition: 0, + }) + if err != nil || !moved.Changed { + t.Fatalf("Queue.MoveItem: result=%#v err=%v", moved, err) + } + + reordered, err := fixture.session.RPC.Queue.PendingItems(t.Context()) + if err != nil { + t.Fatalf("Queue.PendingItems failed: %v", err) + } + gotIDs := []string{reordered.Items[0].ID, reordered.Items[1].ID, reordered.Items[2].ID} + wantIDs := []string{second.ID, first.ID, duplicate.ID} + if !reflect.DeepEqual(gotIDs, wantIDs) { + t.Fatalf("Queue order: got %v, want %v", gotIDs, wantIDs) + } + if reordered.Items[1].DisplayText != updatedDisplay || + reordered.Items[1].AgentMode != rpc.SendAgentModeInteractive { + t.Fatalf("Unexpected edited queue item: %#v", reordered.Items[1]) + } + + sent, err := fixture.session.RPC.Queue.SendNow(t.Context(), &rpc.QueueSendNowRequest{ID: second.ID}) + if err != nil || !sent.Steered { + t.Fatalf("Queue.SendNow: result=%#v err=%v", sent, err) + } + removed, err := fixture.session.RPC.Queue.RemoveAt(t.Context(), &rpc.QueueRemoveAtRequest{ID: duplicate.ID}) + if err != nil || !removed.Removed { + t.Fatalf("Queue.RemoveAt: result=%#v err=%v", removed, err) + } + remaining, err := fixture.session.RPC.Queue.PendingItems(t.Context()) + if err != nil { + t.Fatalf("Queue.PendingItems after edits failed: %v", err) + } + if len(remaining.Items) != 1 || remaining.Items[0].ID != first.ID || + remaining.Items[0].DisplayText != updatedDisplay || + len(remaining.SteeringMessages) != 1 || remaining.SteeringMessages[0] != secondDisplay { + t.Fatalf("Unexpected remaining queue state: %#v", remaining) + } + if _, err := fixture.session.RPC.Queue.SetDrainPaused(t.Context(), &rpc.QueueSetDrainPausedRequest{Paused: false}); err != nil { + t.Fatalf("Queue.SetDrainPaused(false) failed: %v", err) + } + if !reflect.DeepEqual(queue.pauseRequests, []bool{true, false}) { + t.Fatalf("Pause requests: got %v, want [true false]", queue.pauseRequests) + } + }) +} + +func catalogCapabilities(values []string) []rpc.CatalogCapability { + result := make([]rpc.CatalogCapability, len(values)) + for i, value := range values { + result[i] = rpc.CatalogCapability(value) + } + return result +} + +func stringsToAny(values []string) []any { + result := make([]any, len(values)) + for i, value := range values { + result[i] = value + } + return result +} + +func catalogKindsToAny(values []rpc.CatalogCandidateKind) []any { + result := make([]any, len(values)) + for i, value := range values { + result[i] = string(value) + } + return result +} + +func ptr[T any](value T) *T { + return &value +} + +type scenarioQueueItem struct { + id string + prompt string + displayText string + agentMode rpc.SendAgentMode +} + +type scenarioQueue struct { + t *testing.T + mu sync.Mutex + items []scenarioQueueItem + steeringMessages []string + pauseRequests []bool + nextID int +} + +func newScenarioQueue(t *testing.T, server *jsonrpc2.Client) *scenarioQueue { + q := &scenarioQueue{t: t} + server.SetRequestHandler("session.queue.setDrainPaused", q.setDrainPaused) + server.SetRequestHandler("session.queue.insertAt", q.insertAt) + server.SetRequestHandler("session.queue.updateText", q.updateText) + server.SetRequestHandler("session.queue.duplicateAt", q.duplicateAt) + server.SetRequestHandler("session.queue.moveItem", q.moveItem) + server.SetRequestHandler("session.queue.pendingItems", q.pendingItems) + server.SetRequestHandler("session.queue.sendNow", q.sendNow) + server.SetRequestHandler("session.queue.removeAt", q.removeAt) + return q +} + +func (q *scenarioQueue) setDrainPaused(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + Paused bool `json:"paused"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal setDrainPaused: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + q.pauseRequests = append(q.pauseRequests, params.Paused) + return json.RawMessage(`{}`), nil +} + +func (q *scenarioQueue) insertAt(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + Position int `json:"position"` + Message struct { + Prompt string `json:"prompt"` + DisplayPrompt *string `json:"displayPrompt"` + AgentMode *rpc.SendAgentMode `json:"agentMode"` + } `json:"message"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal insertAt: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + q.nextID++ + item := scenarioQueueItem{ + id: "queue-" + string(rune('0'+q.nextID)), + prompt: params.Message.Prompt, + displayText: params.Message.Prompt, + agentMode: rpc.SendAgentModeInteractive, + } + if params.Message.DisplayPrompt != nil { + item.displayText = *params.Message.DisplayPrompt + } + if params.Message.AgentMode != nil { + item.agentMode = *params.Message.AgentMode + } + position := params.Position + if position < 0 { + position = 0 + } + if position > len(q.items) { + position = len(q.items) + } + q.items = append(q.items, scenarioQueueItem{}) + copy(q.items[position+1:], q.items[position:]) + q.items[position] = item + return mustJSON(q.t, map[string]any{"id": item.id}), nil +} + +func (q *scenarioQueue) updateText(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + ID string `json:"id"` + Prompt string `json:"prompt"` + DisplayPrompt *string `json:"displayPrompt"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal updateText: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + for i := range q.items { + if q.items[i].id == params.ID { + q.items[i].prompt = params.Prompt + q.items[i].displayText = params.Prompt + if params.DisplayPrompt != nil { + q.items[i].displayText = *params.DisplayPrompt + } + return json.RawMessage(`{"updated":true}`), nil + } + } + return json.RawMessage(`{"updated":false}`), nil +} + +func (q *scenarioQueue) duplicateAt(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + ID string `json:"id"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal duplicateAt: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + for i, item := range q.items { + if item.id == params.ID { + q.nextID++ + duplicate := item + duplicate.id = "queue-" + string(rune('0'+q.nextID)) + q.items = append(q.items, scenarioQueueItem{}) + copy(q.items[i+2:], q.items[i+1:]) + q.items[i+1] = duplicate + return mustJSON(q.t, map[string]any{"id": duplicate.id}), nil + } + } + return nil, &jsonrpc2.Error{Code: -32602, Message: "queue item not found"} +} + +func (q *scenarioQueue) moveItem(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + ID string `json:"id"` + ToPosition int `json:"toPosition"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal moveItem: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + from := -1 + for i := range q.items { + if q.items[i].id == params.ID { + from = i + break + } + } + if from < 0 { + return nil, &jsonrpc2.Error{Code: -32602, Message: "queue item not found"} + } + to := params.ToPosition + if to < 0 { + to = 0 + } + if to >= len(q.items) { + to = len(q.items) - 1 + } + if from == to { + return json.RawMessage(`{"changed":false}`), nil + } + item := q.items[from] + q.items = append(q.items[:from], q.items[from+1:]...) + q.items = append(q.items, scenarioQueueItem{}) + copy(q.items[to+1:], q.items[to:]) + q.items[to] = item + return json.RawMessage(`{"changed":true}`), nil +} + +func (q *scenarioQueue) pendingItems(_ json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + q.mu.Lock() + defer q.mu.Unlock() + items := make([]map[string]any, len(q.items)) + for i, item := range q.items { + messageID := "message-" + item.id + items[i] = map[string]any{ + "agentMode": item.agentMode, + "displayText": item.displayText, + "id": item.id, + "kind": "message", + "messageId": messageID, + } + } + return mustJSON(q.t, map[string]any{ + "items": items, + "steeringMessages": append([]string(nil), q.steeringMessages...), + }), nil +} + +func (q *scenarioQueue) sendNow(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + ID string `json:"id"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal sendNow: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + for i, item := range q.items { + if item.id == params.ID { + q.items = append(q.items[:i], q.items[i+1:]...) + q.steeringMessages = append(q.steeringMessages, item.displayText) + return json.RawMessage(`{"steered":true}`), nil + } + } + return json.RawMessage(`{"steered":false}`), nil +} + +func (q *scenarioQueue) removeAt(request json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + var params struct { + ID string `json:"id"` + } + if err := json.Unmarshal(request, ¶ms); err != nil { + q.t.Errorf("Unmarshal removeAt: %v", err) + } + q.mu.Lock() + defer q.mu.Unlock() + for i, item := range q.items { + if item.id == params.ID { + q.items = append(q.items[:i], q.items[i+1:]...) + return json.RawMessage(`{"removed":true}`), nil + } + } + return json.RawMessage(`{"removed":false}`), nil +} diff --git a/go/internal/e2e/testharness/inprocess_cleanup_disabled.go b/go/internal/e2e/testharness/inprocess_cleanup_disabled.go index d5fb4db466..f423efdac4 100644 --- a/go/internal/e2e/testharness/inprocess_cleanup_disabled.go +++ b/go/internal/e2e/testharness/inprocess_cleanup_disabled.go @@ -8,3 +8,8 @@ func waitForInProcessCleanup() error { // PrepareForProcessWait is a no-op when the in-process runtime is unavailable. func PrepareForProcessWait() {} + +// ProtectProcessWait is a no-op when the in-process runtime is unavailable. +func ProtectProcessWait() func() { + return func() {} +} diff --git a/go/internal/e2e/testharness/inprocess_cleanup_enabled.go b/go/internal/e2e/testharness/inprocess_cleanup_enabled.go index 83ca9f584a..4491f67a59 100644 --- a/go/internal/e2e/testharness/inprocess_cleanup_enabled.go +++ b/go/internal/e2e/testharness/inprocess_cleanup_enabled.go @@ -22,3 +22,8 @@ func waitForInProcessCleanup() error { func PrepareForProcessWait() { ffihost.PrepareForChildProcessWait() } + +// ProtectProcessWait keeps signal handlers compatible while a child exits. +func ProtectProcessWait() func() { + return ffihost.ProtectChildProcessWait() +} diff --git a/go/internal/e2e/testharness/inprocess_isolation.go b/go/internal/e2e/testharness/inprocess_isolation.go index 945077a6d3..4803526308 100644 --- a/go/internal/e2e/testharness/inprocess_isolation.go +++ b/go/internal/e2e/testharness/inprocess_isolation.go @@ -123,6 +123,7 @@ func runIsolatedProcess(ctx context.Context, name, selector string, timeout time command := exec.CommandContext(ctx, executable, isolatedTestArgs(os.Args[1:], selector, false, timeout)...) command.Env = setEnvironmentValue(os.Environ(), isolatedInProcessTestEnv, name) command.WaitDelay = 5 * time.Second + PrepareForProcessWait() output, err := command.CombinedOutput() fmt.Print(string(output)) if err != nil { diff --git a/go/internal/e2e/testharness/proxy.go b/go/internal/e2e/testharness/proxy.go index f3dc1c4c09..e6a595ffe2 100644 --- a/go/internal/e2e/testharness/proxy.go +++ b/go/internal/e2e/testharness/proxy.go @@ -130,7 +130,8 @@ func (p *CapiProxy) StopWithOptions(skipWritingCache bool) error { p.proxyURL = "" }() - PrepareForProcessWait() + releaseSignalGuard := ProtectProcessWait() + defer releaseSignalGuard() // Send stop request to the server if p.proxyURL != "" { @@ -148,6 +149,7 @@ func (p *CapiProxy) StopWithOptions(skipWritingCache bool) error { exited := make(chan struct{}, 1) go func() { + PrepareForProcessWait() _ = cmd.Wait() exited <- struct{}{} }() diff --git a/go/internal/ffihost/ffihost.go b/go/internal/ffihost/ffihost.go index 5d5161e833..2bb5dd678d 100644 --- a/go/internal/ffihost/ffihost.go +++ b/go/internal/ffihost/ffihost.go @@ -217,6 +217,12 @@ func PrepareForChildProcessWait() { rearmForeignSignalHandlers(0) } +// ProtectChildProcessWait keeps SIGCHLD compatible with the Go runtime while a +// child process is being stopped and reaped. +func ProtectChildProcessWait() func() { + return protectChildProcessSignalHandler() +} + // Create resolves the native library and prepares the host. environment and // args contain SDK-managed runtime options. func Create(runtimeEntrypoint, cliEntrypoint string, environment map[string]string, args []string) (*Host, error) { diff --git a/go/internal/ffihost/sigonstack_darwin.go b/go/internal/ffihost/sigonstack_darwin.go index 0663f17063..50da1d0489 100644 --- a/go/internal/ffihost/sigonstack_darwin.go +++ b/go/internal/ffihost/sigonstack_darwin.go @@ -66,6 +66,13 @@ func rearmForeignSignalHandlers(_ uintptr) { } } +func protectChildProcessSignalHandler() func() { + rearmForeignSignalHandlers(0) + return func() { + rearmForeignSignalHandlers(0) + } +} + // bindSigaction resolves libc's sigaction into fn, converting the panic // RegisterLibFunc raises on a missing symbol into a false return. func bindSigaction(handle uintptr, fn *func(sig int32, act, oact unsafe.Pointer) int32) (ok bool) { diff --git a/go/internal/ffihost/sigonstack_linux.go b/go/internal/ffihost/sigonstack_linux.go index 6c422fd96a..8a02e36af4 100644 --- a/go/internal/ffihost/sigonstack_linux.go +++ b/go/internal/ffihost/sigonstack_linux.go @@ -5,6 +5,7 @@ package ffihost import ( + "runtime" "syscall" "unsafe" ) @@ -13,6 +14,7 @@ const ( linuxSaOnStack = 0x08000000 linuxSigDfl = 0 linuxSigIgn = 1 + linuxSigChild = 17 linuxMaxSignal = 31 ) @@ -39,18 +41,52 @@ type linuxSigaction struct { // is the pre-existing crash. func rearmForeignSignalHandlers(_ uintptr) { for sig := 1; sig <= linuxMaxSignal; sig++ { - var action linuxSigaction - if !linuxGetSigaction(sig, &action) { - continue - } - if action.handler == linuxSigDfl || action.handler == linuxSigIgn { - continue - } - if action.flags&linuxSaOnStack != 0 { - continue + rearmLinuxSignalHandler(sig) + } +} + +func rearmLinuxSignalHandler(sig int) { + var action linuxSigaction + if !linuxGetSigaction(sig, &action) { + return + } + if action.handler == linuxSigDfl || action.handler == linuxSigIgn { + return + } + if action.flags&linuxSaOnStack != 0 { + return + } + action.flags |= linuxSaOnStack + linuxSetSigaction(sig, &action) +} + +func protectChildProcessSignalHandler() func() { + return protectLinuxSignalHandler(linuxSigChild) +} + +func protectLinuxSignalHandler(sig int) func() { + stop := make(chan struct{}) + ready := make(chan struct{}) + stopped := make(chan struct{}) + go func() { + defer close(stopped) + rearmLinuxSignalHandler(sig) + close(ready) + for { + select { + case <-stop: + return + default: + rearmLinuxSignalHandler(sig) + runtime.Gosched() + } } - action.flags |= linuxSaOnStack - linuxSetSigaction(sig, &action) + }() + <-ready + return func() { + close(stop) + <-stopped + rearmLinuxSignalHandler(sig) } } diff --git a/go/internal/ffihost/sigonstack_linux_test.go b/go/internal/ffihost/sigonstack_linux_test.go index 393d590d0c..919b23b381 100644 --- a/go/internal/ffihost/sigonstack_linux_test.go +++ b/go/internal/ffihost/sigonstack_linux_test.go @@ -5,8 +5,10 @@ package ffihost import ( "os" "os/signal" + "runtime" "syscall" "testing" + "time" "unsafe" ) @@ -38,6 +40,39 @@ func TestRearmForeignSignalHandlersAddsOnStack(t *testing.T) { } } +func TestProtectLinuxSignalHandlerRearmsConcurrentReplacement(t *testing.T) { + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGUSR1) + defer signal.Stop(signals) + + var original linuxSigaction + if !linuxGetSigaction(int(syscall.SIGUSR1), &original) { + t.Fatal("failed to read SIGUSR1 action") + } + defer linuxSetSigaction(int(syscall.SIGUSR1), &original) + + release := protectLinuxSignalHandler(int(syscall.SIGUSR1)) + defer release() + + withoutOnStack := original + withoutOnStack.flags &^= linuxSaOnStack + if !linuxSetSigaction(int(syscall.SIGUSR1), &withoutOnStack) { + t.Fatal("failed to clear SA_ONSTACK") + } + + deadline := time.Now().Add(time.Second) + for { + var action linuxSigaction + if linuxGetSigaction(int(syscall.SIGUSR1), &action) && action.flags&linuxSaOnStack != 0 { + return + } + if time.Now().After(deadline) { + t.Fatal("signal guard did not restore SA_ONSTACK") + } + runtime.Gosched() + } +} + func TestHostRearmsSignalHandlersAroundNativeOperations(t *testing.T) { for _, entrypoint := range []string{"", "copilot"} { t.Run("entrypoint="+entrypoint, func(t *testing.T) { diff --git a/go/internal/ffihost/sigonstack_other.go b/go/internal/ffihost/sigonstack_other.go index 6f40992883..3d988e9c7f 100644 --- a/go/internal/ffihost/sigonstack_other.go +++ b/go/internal/ffihost/sigonstack_other.go @@ -8,3 +8,7 @@ package ffihost // linux. Only those Unix platforms deliver the SA_ONSTACK-less SIGCHLD handler // installed by Tokio that the Go runtime rejects; Windows is unaffected. func rearmForeignSignalHandlers(_ uintptr) {} + +func protectChildProcessSignalHandler() func() { + return func() {} +} diff --git a/go/internal/jsonrpc2/jsonrpc2.go b/go/internal/jsonrpc2/jsonrpc2.go index 09364057c3..fa0b45d43c 100644 --- a/go/internal/jsonrpc2/jsonrpc2.go +++ b/go/internal/jsonrpc2/jsonrpc2.go @@ -72,6 +72,8 @@ type Client struct { requestHandlers map[string]RequestHandler running atomic.Bool stopChan chan struct{} + connectionClosed chan struct{} + connectionClosedOnce sync.Once wg sync.WaitGroup processDone chan struct{} // closed when the underlying process exits processErrorPtr *error // points to the process error @@ -89,6 +91,7 @@ func NewClient(stdin io.WriteCloser, stdout io.ReadCloser) *Client { pendingInlineCallbacks: make(map[string]func(json.RawMessage) error), requestHandlers: make(map[string]RequestHandler), stopChan: make(chan struct{}), + connectionClosed: make(chan struct{}), } c.writer <- newHeaderWriter(stdin) return c @@ -254,6 +257,12 @@ func (c *Client) RequestWithInlineResponse(ctx context.Context, method string, p default: // Process still running, continue } + } else { + select { + case <-c.connectionClosed: + return nil, fmt.Errorf("connection closed") + default: + } } var paramsData json.RawMessage @@ -309,6 +318,8 @@ func (c *Client) RequestWithInlineResponse(ctx context.Context, method string, p return nil, response.Error } return response.Result, nil + case <-c.connectionClosed: + return nil, fmt.Errorf("connection closed") case <-c.stopChan: return nil, fmt.Errorf("client stopped") } @@ -356,6 +367,7 @@ func (c *Client) readLoop() { if c.onClose != nil && c.running.Load() { c.onClose() } + c.connectionClosedOnce.Do(func() { close(c.connectionClosed) }) }() for c.running.Load() { diff --git a/go/internal/jsonrpc2/jsonrpc2_test.go b/go/internal/jsonrpc2/jsonrpc2_test.go index 2c7bb3f566..258992e101 100644 --- a/go/internal/jsonrpc2/jsonrpc2_test.go +++ b/go/internal/jsonrpc2/jsonrpc2_test.go @@ -77,6 +77,39 @@ func TestOnCloseNotCalledOnIntentionalStop(t *testing.T) { } } +func TestRequestReturnsWhenConnectionCloses(t *testing.T) { + stdinR, stdinW := io.Pipe() + stdoutR, stdoutW := io.Pipe() + defer stdinR.Close() + defer stdoutR.Close() + + client := NewClient(stdinW, stdoutR) + client.Start() + defer client.Stop() + + result := make(chan error, 1) + go func() { + _, err := client.Request(context.Background(), "test.method", nil) + result <- err + }() + + if _, err := newHeaderReader(stdinR).Read(); err != nil { + t.Fatalf("Read request failed: %v", err) + } + if err := stdoutW.Close(); err != nil { + t.Fatalf("Close response stream failed: %v", err) + } + + select { + case err := <-result: + if err == nil || err.Error() != "connection closed" { + t.Fatalf("Request error = %v, want connection closed", err) + } + case <-time.After(time.Second): + t.Fatal("Request did not return when the connection closed") + } +} + // TestSetProcessDone_ErrorAvailableImmediately validates that getProcessError() // returns the correct error immediately after processDone is closed. // The current implementation stores a pointer to the process error diff --git a/go/rpc/sessions_client_metadata_json.go b/go/rpc/sessions_client_metadata_json.go new file mode 100644 index 0000000000..b9bac9aa16 --- /dev/null +++ b/go/rpc/sessions_client_metadata_json.go @@ -0,0 +1,32 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package rpc + +import ( + "encoding/json" + "fmt" +) + +// UnmarshalJSON decodes the discriminated metadata result entries. +func (r *SessionsGetClientMetadataResult) UnmarshalJSON(data []byte) error { + var entries []json.RawMessage + if err := json.Unmarshal(data, &entries); err != nil { + return err + } + if entries == nil { + *r = nil + return nil + } + + result := make(SessionsGetClientMetadataResult, 0, len(entries)) + for index, entry := range entries { + value, err := unmarshalSessionsClientMetadataEntry(entry) + if err != nil { + return fmt.Errorf("decode sessions client metadata entry %d: %w", index, err) + } + result = append(result, value) + } + + *r = result + return nil +} diff --git a/go/rpc/sessions_client_metadata_json_test.go b/go/rpc/sessions_client_metadata_json_test.go new file mode 100644 index 0000000000..8aece6a8ae --- /dev/null +++ b/go/rpc/sessions_client_metadata_json_test.go @@ -0,0 +1,98 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. + +package rpc + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestSessionsGetClientMetadataResultUnmarshalJSON(t *testing.T) { + const input = `[ + {"status":"corrupt","sessionId":"corrupt-session"}, + {"status":"notFound","sessionId":"missing-session"}, + {"status":"ok","sessionId":"ok-session","metadata":{"branch":"main","owner":"github"}}, + {"status":"unavailable","sessionId":"locked-session","code":"EBUSY","message":"metadata is locked"}, + {"status":"unsupportedVersion","sessionId":"future-session"} + ]` + + var result SessionsGetClientMetadataResult + if err := json.Unmarshal([]byte(input), &result); err != nil { + t.Fatalf("Unmarshal failed: %v", err) + } + if len(result) != 5 { + t.Fatalf("Expected five entries, got %d", len(result)) + } + + corrupt, ok := result[0].(*SessionsClientMetadataEntryCorrupt) + if !ok || corrupt.SessionID != "corrupt-session" { + t.Fatalf("Unexpected corrupt entry: %#v", result[0]) + } + notFound, ok := result[1].(*SessionsClientMetadataEntryNotFound) + if !ok || notFound.SessionID != "missing-session" { + t.Fatalf("Unexpected not-found entry: %#v", result[1]) + } + metadata, ok := result[2].(*SessionsClientMetadataEntryOk) + if !ok || metadata.SessionID != "ok-session" || metadata.Metadata["branch"] != "main" || metadata.Metadata["owner"] != "github" { + t.Fatalf("Unexpected ok entry: %#v", result[2]) + } + unavailable, ok := result[3].(*SessionsClientMetadataEntryUnavailable) + if !ok || unavailable.SessionID != "locked-session" || unavailable.Code != "EBUSY" || unavailable.Message != "metadata is locked" { + t.Fatalf("Unexpected unavailable entry: %#v", result[3]) + } + unsupported, ok := result[4].(*SessionsClientMetadataEntryUnsupportedVersion) + if !ok || unsupported.SessionID != "future-session" { + t.Fatalf("Unexpected unsupported-version entry: %#v", result[4]) + } +} + +func TestSessionsGetClientMetadataResultUnmarshalJSONRejectsInvalidEntries(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + { + name: "malformed entry", + input: `[{"status":42,"sessionId":"session"}]`, + want: "cannot unmarshal number", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var result SessionsGetClientMetadataResult + err := json.Unmarshal([]byte(test.input), &result) + if err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("Expected error containing %q, got %v", test.want, err) + } + }) + } +} + +func TestSessionsGetClientMetadataResultUnmarshalJSONPreservesUnknownEntries(t *testing.T) { + const input = `[{"status":"future-status","sessionId":"future-session","detail":{"attempts":2}}]` + var result SessionsGetClientMetadataResult + if err := json.Unmarshal([]byte(input), &result); err != nil { + t.Fatalf("Unmarshal failed: %v", err) + } + if len(result) != 1 { + t.Fatalf("Expected one entry, got %d", len(result)) + } + raw, ok := result[0].(*RawSessionsClientMetadataEntryData) + if !ok || raw.Discriminator != "future-status" || string(raw.Raw) != input[1:len(input)-1] { + t.Fatalf("Unexpected raw entry: %#v", result[0]) + } +} + +func TestSessionsGetClientMetadataResultUnmarshalJSONPreservesNull(t *testing.T) { + result := SessionsGetClientMetadataResult{ + &SessionsClientMetadataEntryNotFound{SessionID: "existing"}, + } + if err := json.Unmarshal([]byte("null"), &result); err != nil { + t.Fatalf("Unmarshal failed: %v", err) + } + if result != nil { + t.Fatalf("Expected nil result, got %#v", result) + } +} diff --git a/go/types.go b/go/types.go index cfbf5b5e8f..e9d9c9a684 100644 --- a/go/types.go +++ b/go/types.go @@ -134,6 +134,10 @@ type ClientOptions struct { // runtime's complete trusted built-in plugin directory set before sessions // can be created. BuiltinPluginDirectories []string + // ExtensionLaunchProvider resolves launch profiles for extension entrypoints + // discovered by the runtime. When non-nil, Start registers the provider + // before any sessions can be created. + ExtensionLaunchProvider ExtensionLaunchProvider // LogLevel for the runtime. When empty (the default), the runtime // uses its own default level; the SDK does not pass --log-level. // Recognized values: "none", "error", "warning", "info", "debug", "all". @@ -211,6 +215,23 @@ type ClientOptions struct { Mode ClientMode } +// ExtensionLaunchProvider resolves launch profiles for extension entrypoints +// discovered by the runtime. +// +// Experimental: this API may change or be removed. +type ExtensionLaunchProvider = rpc.ExtensionLaunchProviderHandler + +// ExtensionLaunchProviderResolveRequest describes a discovered extension +// entrypoint that may need a host-provided launch profile. +type ExtensionLaunchProviderResolveRequest = rpc.ExtensionLaunchProviderResolveRequest + +// ExtensionLaunchProviderResolveResult contains the optional host-provided +// launch profile for an extension entrypoint. +type ExtensionLaunchProviderResolveResult = rpc.ExtensionLaunchProviderResolveResult + +// ExtensionLaunchProfile describes how the runtime should launch an extension. +type ExtensionLaunchProfile = rpc.ExtensionLaunchProfile + // ClientInfo identifies the integrating application on the `server.connect` handshake. // // Declaring it lets the telemetry the runtime emits on the connection be diff --git a/java/README.md b/java/README.md index f3381e3255..d9b783afa1 100644 --- a/java/README.md +++ b/java/README.md @@ -178,6 +178,11 @@ directly. `CopilotClientOptions.setCwd(...)` sets the runtime process working directory, which otherwise inherits the current process working directory. `SessionConfig.setWorkingDirectory(...)` sets the session working directory, which otherwise defaults to the runtime process working directory. +`CopilotClientOptions.setExtensionLaunchProvider(...)` configures an experimental +connection-level resolver for extension launch profiles. The client installs the +reverse-RPC handler and registers the provider during startup before sessions can +be created. + `SessionConfig.setAskUserVariant(AskUserVariant.ELICITATION)` selects the structured form-based `ask_user` tool when an elicitation handler is also set. The default is `AskUserVariant.LEGACY`. Re-supply the option and handler through diff --git a/java/sdk/src/main/java/com/github/copilot/CopilotClient.java b/java/sdk/src/main/java/com/github/copilot/CopilotClient.java index 9f7d8ebcf8..ae41477c15 100644 --- a/java/sdk/src/main/java/com/github/copilot/CopilotClient.java +++ b/java/sdk/src/main/java/com/github/copilot/CopilotClient.java @@ -29,6 +29,7 @@ import com.github.copilot.ffi.NativeRuntimeLoader; import com.github.copilot.rpc.CopilotClientMode; import com.github.copilot.rpc.CopilotClientOptions; +import com.github.copilot.rpc.ExtensionLaunchProvider; import com.github.copilot.rpc.InProcessRuntimeConnection; import com.github.copilot.rpc.RuntimeConnection; import com.github.copilot.rpc.StdioRuntimeConnection; @@ -577,11 +578,20 @@ private Connection startCoreBody() { telemetryAdapter.registerHandlers(connectedRpc); } + ExtensionLaunchProvider extensionLaunchProvider = this.options.getExtensionLaunchProvider(); + if (extensionLaunchProvider != null) { + new ExtensionLaunchProviderAdapter(extensionLaunchProvider).registerHandlers(connectedRpc); + } + // Verify protocol version verifyProtocolVersion(connection); LoggingHelpers.logTiming(LOG, Level.FINE, "CopilotClient.start protocol verification complete. Elapsed={Elapsed}", startNanos); + if (extensionLaunchProvider != null) { + connection.serverRpc().registerExtensionLaunchProvider().join(); + } + var builtinPluginDirectories = options.getBuiltinPluginDirectories(); if (builtinPluginDirectories != null && !builtinPluginDirectories.isEmpty()) { var paths = new ArrayList(builtinPluginDirectories.size()); diff --git a/java/sdk/src/main/java/com/github/copilot/ExtensionLaunchProviderAdapter.java b/java/sdk/src/main/java/com/github/copilot/ExtensionLaunchProviderAdapter.java new file mode 100644 index 0000000000..7db53aecf1 --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/ExtensionLaunchProviderAdapter.java @@ -0,0 +1,86 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import java.io.IOException; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import java.util.logging.Level; +import java.util.logging.Logger; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveRequest; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveResult; +import com.github.copilot.rpc.ExtensionLaunchProvider; + +/** + * Bridges {@code extensionLaunchProvider.resolve} reverse RPC calls to the + * configured extension launch provider. + */ +final class ExtensionLaunchProviderAdapter { + + private static final Logger LOG = Logger.getLogger(ExtensionLaunchProviderAdapter.class.getName()); + private static final ObjectMapper MAPPER = JsonRpcClient.getObjectMapper(); + + private final ExtensionLaunchProvider provider; + + ExtensionLaunchProviderAdapter(ExtensionLaunchProvider provider) { + this.provider = provider; + } + + void registerHandlers(JsonRpcClient rpc) { + rpc.registerMethodHandler("extensionLaunchProvider.resolve", + (rpcId, params) -> handleResolve(rpc, rpcId, params)); + } + + private void handleResolve(JsonRpcClient rpc, String rpcId, JsonNode params) { + if (rpcId == null) { + return; + } + + try { + ExtensionLaunchProviderResolveRequest request = MAPPER.treeToValue(params, + ExtensionLaunchProviderResolveRequest.class); + CompletableFuture resolution = provider.resolve(request); + if (resolution == null) { + sendError(rpc, rpcId, "Extension launch provider returned a null future"); + return; + } + resolution.whenComplete((result, error) -> { + if (error != null) { + Throwable cause = error instanceof CompletionException && error.getCause() != null + ? error.getCause() + : error; + sendError(rpc, rpcId, cause.getMessage() != null ? cause.getMessage() : cause.toString()); + return; + } + try { + rpc.sendResponse(parseRpcId(rpcId), result); + } catch (IOException e) { + LOG.log(Level.FINE, "Failed to send extension launch provider response", e); + } + }); + } catch (Exception e) { + sendError(rpc, rpcId, e.getMessage() != null ? e.getMessage() : e.toString()); + } + } + + private static void sendError(JsonRpcClient rpc, String rpcId, String message) { + try { + rpc.sendErrorResponse(parseRpcId(rpcId), -32603, message); + } catch (IOException e) { + LOG.log(Level.FINE, "Failed to send extension launch provider error", e); + } + } + + private static Object parseRpcId(String rpcId) { + try { + return Long.valueOf(rpcId); + } catch (NumberFormatException ignored) { + return rpcId; + } + } +} diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/CopilotClientOptions.java b/java/sdk/src/main/java/com/github/copilot/rpc/CopilotClientOptions.java index c67947bf48..fa00ea1820 100644 --- a/java/sdk/src/main/java/com/github/copilot/rpc/CopilotClientOptions.java +++ b/java/sdk/src/main/java/com/github/copilot/rpc/CopilotClientOptions.java @@ -63,6 +63,7 @@ public class CopilotClientOptions { private String gitHubToken; private String logLevel = "info"; private CopilotClientMode mode = CopilotClientMode.COPILOT_CLI; + private ExtensionLaunchProvider extensionLaunchProvider; private Supplier>> onListModels; private CopilotRequestHandler requestHandler; private Function> onGitHubTelemetry; @@ -158,6 +159,34 @@ public CopilotClientOptions setBuiltinPluginDirectories(List paths) { return this; } + /** + * Gets the connection-level extension launch profile provider. + * + * @return the provider, or {@code null} if not set + */ + @JsonIgnore + @CopilotExperimental + public ExtensionLaunchProvider getExtensionLaunchProvider() { + return extensionLaunchProvider; + } + + /** + * Sets the connection-level extension launch profile provider. + *

+ * When provided, the client registers the provider during startup before any + * session can be created. + * + * @param extensionLaunchProvider + * the provider (must not be {@code null}) + * @return this options instance for method chaining + */ + @CopilotExperimental + public CopilotClientOptions setExtensionLaunchProvider(ExtensionLaunchProvider extensionLaunchProvider) { + this.extensionLaunchProvider = Objects.requireNonNull(extensionLaunchProvider, + "extensionLaunchProvider must not be null"); + return this; + } + /** * Gets the extra CLI arguments. *

@@ -866,6 +895,7 @@ public CopilotClientOptions clone() { copy.cwd = this.cwd; copy.environment = this.environment != null ? new java.util.HashMap<>(this.environment) : null; copy.executor = this.executor; + copy.extensionLaunchProvider = this.extensionLaunchProvider; copy.gitHubToken = this.gitHubToken; copy.logLevel = this.logLevel; copy.onListModels = this.onListModels; diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionContextAttachment.java b/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionContextAttachment.java new file mode 100644 index 0000000000..27a2d41eeb --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionContextAttachment.java @@ -0,0 +1,175 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot.rpc; + +import com.fasterxml.jackson.annotation.JsonInclude; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * Structured context contributed by an extension. + * + * @see MessageOptions#setAttachments(java.util.List) + * @since 1.0.0 + */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public final class ExtensionContextAttachment implements MessageAttachment { + + @JsonProperty("type") + private final String type = "extension_context"; + + @JsonProperty("capturedAt") + private String capturedAt; + + @JsonProperty("extensionId") + private String extensionId; + + @JsonProperty("title") + private String title; + + @JsonProperty("canvasId") + private String canvasId; + + @JsonProperty("instanceId") + private String instanceId; + + @JsonProperty("payload") + private Object payload; + + /** + * Returns the attachment type, always {@code "extension_context"}. + * + * @return {@code "extension_context"} + */ + @Override + public String getType() { + return type; + } + + /** + * Gets the ISO 8601 capture timestamp. + * + * @return the capture timestamp + */ + public String getCapturedAt() { + return capturedAt; + } + + /** + * Sets the ISO 8601 capture timestamp. + * + * @param capturedAt + * the capture timestamp + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setCapturedAt(String capturedAt) { + this.capturedAt = capturedAt; + return this; + } + + /** + * Gets the owning extension identifier. + * + * @return the extension identifier + */ + public String getExtensionId() { + return extensionId; + } + + /** + * Sets the owning extension identifier. + * + * @param extensionId + * the extension identifier + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setExtensionId(String extensionId) { + this.extensionId = extensionId; + return this; + } + + /** + * Gets the human-readable context title. + * + * @return the context title + */ + public String getTitle() { + return title; + } + + /** + * Sets the human-readable context title. + * + * @param title + * the context title + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setTitle(String title) { + this.title = title; + return this; + } + + /** + * Gets the provider-local canvas identifier. + * + * @return the canvas identifier, or {@code null} + */ + public String getCanvasId() { + return canvasId; + } + + /** + * Sets the provider-local canvas identifier. + * + * @param canvasId + * the canvas identifier + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setCanvasId(String canvasId) { + this.canvasId = canvasId; + return this; + } + + /** + * Gets the open canvas instance identifier. + * + * @return the instance identifier, or {@code null} + */ + public String getInstanceId() { + return instanceId; + } + + /** + * Sets the open canvas instance identifier. + * + * @param instanceId + * the instance identifier + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setInstanceId(String instanceId) { + this.instanceId = instanceId; + return this; + } + + /** + * Gets the extension-defined structured payload. + * + * @return the structured payload, or {@code null} + */ + public Object getPayload() { + return payload; + } + + /** + * Sets the extension-defined structured payload. + * + * @param payload + * a JSON-serializable payload + * @return this attachment for method chaining + */ + public ExtensionContextAttachment setPayload(Object payload) { + this.payload = payload; + return this; + } +} diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionLaunchProvider.java b/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionLaunchProvider.java new file mode 100644 index 0000000000..1428563526 --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/rpc/ExtensionLaunchProvider.java @@ -0,0 +1,30 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot.rpc; + +import java.util.concurrent.CompletableFuture; + +import com.github.copilot.CopilotExperimental; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveRequest; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveResult; + +/** + * Resolves launch profiles for extension entrypoints discovered by the runtime. + * + * @since 1.0.0 + */ +@FunctionalInterface +@CopilotExperimental +public interface ExtensionLaunchProvider { + + /** + * Resolves an optional launch profile for a discovered extension entrypoint. + * + * @param request + * the discovered extension entrypoint + * @return a future containing the launch resolution + */ + CompletableFuture resolve(ExtensionLaunchProviderResolveRequest request); +} diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/MessageAttachment.java b/java/sdk/src/main/java/com/github/copilot/rpc/MessageAttachment.java index 9b2af3ee08..46101bf15e 100644 --- a/java/sdk/src/main/java/com/github/copilot/rpc/MessageAttachment.java +++ b/java/sdk/src/main/java/com/github/copilot/rpc/MessageAttachment.java @@ -15,13 +15,15 @@ * * @see Attachment * @see BlobAttachment + * @see ExtensionContextAttachment * @see MessageOptions#setAttachments(java.util.List) * @since 1.0.0 */ @JsonTypeInfo(use = JsonTypeInfo.Id.NAME, property = "type") @JsonSubTypes({@JsonSubTypes.Type(value = Attachment.class, name = "file"), - @JsonSubTypes.Type(value = BlobAttachment.class, name = "blob")}) -public sealed interface MessageAttachment permits Attachment, BlobAttachment { + @JsonSubTypes.Type(value = BlobAttachment.class, name = "blob"), + @JsonSubTypes.Type(value = ExtensionContextAttachment.class, name = "extension_context")}) +public sealed interface MessageAttachment permits Attachment, BlobAttachment, ExtensionContextAttachment { /** * Returns the attachment type discriminator (e.g., "file", "blob"). diff --git a/java/sdk/src/test/java/com/github/copilot/ExtensionLaunchProviderTest.java b/java/sdk/src/test/java/com/github/copilot/ExtensionLaunchProviderTest.java new file mode 100644 index 0000000000..5785b0cf60 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/ExtensionLaunchProviderTest.java @@ -0,0 +1,166 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertSame; + +import java.io.IOException; +import java.net.ServerSocket; +import java.net.Socket; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; + +import org.junit.jupiter.api.Test; + +import com.github.copilot.generated.rpc.ExtensionLaunchProfile; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveRequest; +import com.github.copilot.generated.rpc.ExtensionLaunchProviderResolveResult; +import com.github.copilot.rpc.CopilotClientOptions; +import com.github.copilot.rpc.ExtensionLaunchProvider; + +@AllowCopilotExperimental +class ExtensionLaunchProviderTest { + + @Test + void configuredProviderRegistersAndHandlesResolveDuringStartup() throws Exception { + var observed = new CompletableFuture(); + ExtensionLaunchProvider provider = request -> { + observed.complete(request); + return CompletableFuture.completedFuture(new ExtensionLaunchProviderResolveResult( + new ExtensionLaunchProfile("java", List.of("extension-host"), Map.of("EXTENSION_SOURCE", "java")))); + }; + + try (var server = new FakeRuntimeServer(); + var client = new CopilotClient( + new CopilotClientOptions().setCliUrl(server.url()).setExtensionLaunchProvider(provider))) { + client.start().get(15, TimeUnit.SECONDS); + + var request = observed.get(15, TimeUnit.SECONDS); + assertEquals("project:java-e2e", request.id()); + assertEquals("java-e2e", request.name()); + assertEquals("/extensions/java-e2e.jar", request.modulePath()); + assertEquals("project", request.source().getValue()); + + var result = server.resolveResult().get(15, TimeUnit.SECONDS); + assertEquals("java", result.launch().executable()); + assertEquals(List.of("extension-host"), result.launch().args()); + assertEquals(Map.of("EXTENSION_SOURCE", "java"), result.launch().env()); + assertEquals(1, server.registrationCount()); + } + } + + @Test + void clonePreservesProvider() { + ExtensionLaunchProvider provider = request -> CompletableFuture + .completedFuture(new ExtensionLaunchProviderResolveResult(null)); + var clone = new CopilotClientOptions().setExtensionLaunchProvider(provider).clone(); + + assertSame(provider, clone.getExtensionLaunchProvider()); + } + + private static final class FakeRuntimeServer implements AutoCloseable { + + private final ServerSocket serverSocket; + private final Thread acceptThread; + private final CompletableFuture ready = new CompletableFuture<>(); + private final CompletableFuture resolveResult = new CompletableFuture<>(); + private final AtomicInteger registrationCount = new AtomicInteger(); + + FakeRuntimeServer() throws IOException { + serverSocket = new ServerSocket(0); + acceptThread = new Thread(this::acceptLoop, "extension-launch-provider-runtime"); + acceptThread.setDaemon(true); + acceptThread.start(); + } + + String url() { + return "127.0.0.1:" + serverSocket.getLocalPort(); + } + + int registrationCount() { + return registrationCount.get(); + } + + CompletableFuture resolveResult() { + return resolveResult; + } + + private void acceptLoop() { + try { + Socket socket = serverSocket.accept(); + JsonRpcClient server = JsonRpcClient.fromSocket(socket); + server.registerMethodHandler("connect", (id, params) -> respond(server, id, + Map.of("ok", true, "protocolVersion", 3, "version", "test"))); + server.registerMethodHandler("registerExtensionLaunchProvider", (id, params) -> { + registrationCount.incrementAndGet(); + server.invoke("extensionLaunchProvider.resolve", + Map.of("id", "project:java-e2e", "name", "java-e2e", "modulePath", + "/extensions/java-e2e.jar", "source", "project"), + ExtensionLaunchProviderResolveResult.class).whenComplete((result, error) -> { + if (error != null) { + resolveResult.completeExceptionally(error); + sendError(server, id, error); + return; + } + resolveResult.complete(result); + respond(server, id, Map.of()); + }); + }); + ready.complete(server); + } catch (IOException e) { + ready.completeExceptionally(e); + resolveResult.completeExceptionally(e); + } + } + + private static void respond(JsonRpcClient server, String id, Object result) { + if (id == null) { + return; + } + try { + server.sendResponse(parseRpcId(id), result); + } catch (IOException e) { + throw new IllegalStateException("Failed to send fake runtime response", e); + } + } + + private static void sendError(JsonRpcClient server, String id, Throwable error) { + if (id == null) { + return; + } + try { + server.sendErrorResponse(parseRpcId(id), -32603, error.getMessage()); + } catch (IOException e) { + resolveResultFailure(error, e); + } + } + + private static Object parseRpcId(String id) { + try { + return Long.valueOf(id); + } catch (NumberFormatException ignored) { + return id; + } + } + + private static void resolveResultFailure(Throwable original, IOException responseFailure) { + original.addSuppressed(responseFailure); + } + + @Override + public void close() throws Exception { + JsonRpcClient server = ready.getNow(null); + if (server != null) { + server.close(); + } + serverSocket.close(); + acceptThread.join(TimeUnit.SECONDS.toMillis(5)); + } + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/MessageAttachmentTest.java b/java/sdk/src/test/java/com/github/copilot/MessageAttachmentTest.java index 27e9f56cc3..35ed4e64e4 100644 --- a/java/sdk/src/test/java/com/github/copilot/MessageAttachmentTest.java +++ b/java/sdk/src/test/java/com/github/copilot/MessageAttachmentTest.java @@ -7,6 +7,7 @@ import static org.junit.jupiter.api.Assertions.*; import java.util.List; +import java.util.Map; import org.junit.jupiter.api.Test; @@ -14,6 +15,7 @@ import com.github.copilot.rpc.Attachment; import com.github.copilot.rpc.BlobAttachment; +import com.github.copilot.rpc.ExtensionContextAttachment; import com.github.copilot.rpc.MessageAttachment; import com.github.copilot.rpc.MessageOptions; import com.github.copilot.rpc.SendMessageRequest; @@ -45,6 +47,13 @@ void blobAttachmentImplementsMessageAttachment() { assertEquals("blob", blob.getType()); } + @Test + void extensionContextAttachmentImplementsMessageAttachment() { + ExtensionContextAttachment context = extensionContextAttachment(); + assertInstanceOf(MessageAttachment.class, context); + assertEquals("extension_context", context.getType()); + } + // ========================================================================= // MessageOptions type safety // ========================================================================= @@ -73,12 +82,13 @@ void setAttachmentsAcceptsListOfBlobAttachment() { void setAttachmentsAcceptsMixedList() { MessageOptions options = new MessageOptions(); List mixed = List.of(new Attachment("file", "/a.java", "A"), - new BlobAttachment().setData("ZGF0YQ==").setMimeType("image/png")); + new BlobAttachment().setData("ZGF0YQ==").setMimeType("image/png"), extensionContextAttachment()); options.setAttachments(mixed); - assertEquals(2, options.getAttachments().size()); + assertEquals(3, options.getAttachments().size()); assertInstanceOf(Attachment.class, options.getAttachments().get(0)); assertInstanceOf(BlobAttachment.class, options.getAttachments().get(1)); + assertInstanceOf(ExtensionContextAttachment.class, options.getAttachments().get(2)); } @Test @@ -132,15 +142,29 @@ void serializeBlobAttachmentIncludesType() throws Exception { assertTrue(json.contains("\"mimeType\":\"image/png\"")); } + @Test + void serializeExtensionContextAttachmentIncludesContext() throws Exception { + String json = MAPPER.writeValueAsString(extensionContextAttachment()); + assertTrue(json.contains("\"type\":\"extension_context\"")); + assertTrue(json.contains("\"capturedAt\":\"2026-09-18T20:00:00Z\"")); + assertTrue(json.contains("\"extensionId\":\"scenario-extension\"")); + assertTrue(json.contains("\"canvasId\":\"diff\"")); + assertTrue(json.contains("\"instanceId\":\"diff-17\"")); + assertTrue(json.contains("\"selection\":\"active\"")); + } + @Test void serializeMessageOptionsWithMixedAttachments() throws Exception { MessageOptions options = new MessageOptions().setPrompt("Describe") .setAttachments(List.of(new Attachment("file", "/a.java", "A"), - new BlobAttachment().setData("ZGF0YQ==").setMimeType("image/png").setDisplayName("img.png"))); + new BlobAttachment().setData("ZGF0YQ==").setMimeType("image/png").setDisplayName("img.png"), + extensionContextAttachment())); String json = MAPPER.writeValueAsString(options); assertTrue(json.contains("\"type\":\"file\"")); assertTrue(json.contains("\"type\":\"blob\"")); + assertTrue(json.contains("\"type\":\"extension_context\"")); + assertTrue(json.contains("\"selection\":\"active\"")); } @Test @@ -155,4 +179,10 @@ void cloneMessageOptionsPreservesAttachments() { // Verify clone is independent assertNotSame(original.getAttachments(), cloned.getAttachments()); } + + private static ExtensionContextAttachment extensionContextAttachment() { + return new ExtensionContextAttachment().setCapturedAt("2026-09-18T20:00:00Z") + .setExtensionId("scenario-extension").setTitle("Selected change").setCanvasId("diff") + .setInstanceId("diff-17").setPayload(Map.of("selection", "active")); + } } diff --git a/java/sdk/src/test/java/com/github/copilot/RpcSurfaceParityE2ETest.java b/java/sdk/src/test/java/com/github/copilot/RpcSurfaceParityE2ETest.java new file mode 100644 index 0000000000..aee166f536 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/RpcSurfaceParityE2ETest.java @@ -0,0 +1,1176 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.*; + +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; +import java.lang.reflect.Modifier; +import java.lang.reflect.RecordComponent; +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.util.Arrays; +import java.util.Comparator; +import java.util.IdentityHashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.TreeMap; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.TimeUnit; +import java.util.stream.Collectors; + +import org.junit.jupiter.api.Test; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.github.copilot.generated.rpc.*; +import com.github.copilot.rpc.CopilotClientOptions; +import com.github.copilot.rpc.PermissionHandler; +import com.github.copilot.rpc.RuntimeConnection; +import com.github.copilot.rpc.SessionConfig; + +@AllowCopilotExperimental +class RpcSurfaceParityE2ETest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + private static final long TIMEOUT_SECONDS = 30; + private static final int EXPECTED_RPC_METHOD_COUNT = 403; + private static final String EXPECTED_RPC_SIGNATURE_SHA256 = "24f93e3b9142caf0641a506217ef53727ab87f518ddbacb41852d913d7ad5988"; + private static final Map EXPECTED_METHODS_BY_DECLARING_TYPE = Map.ofEntries( + Map.entry("RpcCaller", 2), Map.entry("ServerAccountApi", 6), Map.entry("ServerAgentRegistryApi", 1), + Map.entry("ServerAgentsApi", 2), Map.entry("ServerCatalogApi", 2), Map.entry("ServerCommandsApi", 1), + Map.entry("ServerExtensionsApi", 3), Map.entry("ServerHooksApi", 1), Map.entry("ServerInstructionsApi", 2), + Map.entry("ServerLlmInferenceApi", 3), Map.entry("ServerManagedSettingsApi", 2), + Map.entry("ServerMcpApi", 2), Map.entry("ServerMcpConfigApi", 7), Map.entry("ServerModelsApi", 3), + Map.entry("ServerPluginsApi", 7), Map.entry("ServerPluginsBuiltinApi", 1), + Map.entry("ServerPluginsMarketplacesApi", 6), Map.entry("ServerRpc", 3), Map.entry("ServerRuntimeApi", 1), + Map.entry("ServerSecretsApi", 1), Map.entry("ServerSessionFsApi", 1), Map.entry("ServerSessionsApi", 34), + Map.entry("ServerSkillsApi", 2), Map.entry("ServerSkillsConfigApi", 2), Map.entry("ServerToolsApi", 1), + Map.entry("ServerUserSettingsApi", 3), Map.entry("SessionAgentApi", 7), + Map.entry("SessionAutopilotObjectiveApi", 1), Map.entry("SessionCanvasActionApi", 1), + Map.entry("SessionCanvasApi", 4), Map.entry("SessionCanvasProviderApi", 2), + Map.entry("SessionCommandsApi", 8), Map.entry("SessionCompletionsApi", 2), + Map.entry("SessionContentExclusionApi", 1), Map.entry("SessionDebugApi", 1), + Map.entry("SessionEventLogApi", 4), Map.entry("SessionExtensionsApi", 5), + Map.entry("SessionFactoryApi", 13), Map.entry("SessionFactoryJournalApi", 2), + Map.entry("SessionFleetApi", 1), Map.entry("SessionGitHubAuthApi", 10), Map.entry("SessionHistoryApi", 10), + Map.entry("SessionInstructionsApi", 1), Map.entry("SessionLimitPredictionApi", 2), + Map.entry("SessionManagedSettingsApi", 1), Map.entry("SessionLspApi", 1), Map.entry("SessionMcpApi", 18), + Map.entry("SessionMcpAppsApi", 6), Map.entry("SessionMcpHeadersApi", 1), Map.entry("SessionMcpOauthApi", 5), + Map.entry("SessionMcpResourcesApi", 3), Map.entry("SessionMetadataApi", 11), Map.entry("SessionModeApi", 2), + Map.entry("SessionModelApi", 8), Map.entry("SessionNameApi", 3), Map.entry("SessionOptionsApi", 1), + Map.entry("SessionPermissionsApi", 10), Map.entry("SessionPermissionsFolderTrustApi", 2), + Map.entry("SessionPermissionsLocationsApi", 3), Map.entry("SessionPermissionsPathsApi", 5), + Map.entry("SessionPermissionsUrlsApi", 1), Map.entry("SessionPlanApi", 5), + Map.entry("SessionPluginsApi", 8), Map.entry("SessionPluginsMarketplacesApi", 6), + Map.entry("SessionProviderApi", 3), Map.entry("SessionQueueApi", 20), Map.entry("SessionRemoteApi", 3), + Map.entry("SessionRpc", 9), Map.entry("SessionSandboxApi", 2), Map.entry("SessionScheduleApi", 9), + Map.entry("SessionSettingsApi", 2), Map.entry("SessionShellApi", 4), Map.entry("SessionSkillsApi", 6), + Map.entry("SessionTasksApi", 13), Map.entry("SessionTelemetryApi", 2), Map.entry("SessionToolsApi", 8), + Map.entry("SessionUiApi", 10), Map.entry("SessionUsageApi", 1), Map.entry("SessionVisibilityApi", 2), + Map.entry("SessionWorkflowApi", 13), Map.entry("SessionWorkflowJournalApi", 2), + Map.entry("SessionWorkspacesApi", 20)); + + @Test + void everyGeneratedRpcMethodHasRequestCaptureCoverageAndStableStructuralInventory() throws Exception { + var caller = new RpcSurfaceTestCli.RecordingCaller(); + var targets = new LinkedHashMap, RpcTarget>(); + collectTargets(new ServerRpc(caller), "", targets, new IdentityHashMap<>()); + collectTargets(new SessionRpc(caller, "surface-session"), "session", targets, new IdentityHashMap<>()); + + var methods = targets.values().stream() + .flatMap(target -> rpcMethods(target.instance().getClass()).stream() + .map(method -> new TargetMethod(target, method))) + .sorted(Comparator.comparing(TargetMethod::signature)).toList(); + var callerMethods = rpcMethods(RpcCaller.class); + Map counts = methods.stream() + .collect(Collectors.groupingBy(method -> method.method().getDeclaringClass().getSimpleName(), + TreeMap::new, Collectors.summingInt(ignored -> 1))); + counts.put(RpcCaller.class.getSimpleName(), callerMethods.size()); + assertEquals(EXPECTED_METHODS_BY_DECLARING_TYPE, counts, + "Generated public RPC methods changed; map each new signature to a capture test or documented exclusion"); + assertEquals(EXPECTED_RPC_METHOD_COUNT, methods.size() + callerMethods.size()); + assertEquals(EXPECTED_RPC_SIGNATURE_SHA256, + sha256(java.util.stream.Stream + .concat(methods.stream().map(TargetMethod::signature), + callerMethods.stream().map(RpcSurfaceParityE2ETest::signature)) + .sorted().collect(Collectors.joining("\n")))); + + for (TargetMethod targetMethod : methods) { + caller.clear(); + var method = targetMethod.method(); + var arguments = Arrays.stream(method.getParameterTypes()).map(RpcSurfaceParityE2ETest::fixture).toArray(); + var future = assertInstanceOf(CompletableFuture.class, + invoke(method, targetMethod.target().instance(), arguments), targetMethod.signature()); + assertNotNull(future); + + var call = assertSingleCall(caller, targetMethod.signature()); + var expectedMethod = targetMethod.target().prefix().isEmpty() + ? method.getName() + : targetMethod.target().prefix() + "." + method.getName(); + assertEquals(expectedMethod, call.method(), targetMethod.signature()); + assertNotNull(call.resultType(), targetMethod.signature()); + if (expectedMethod.startsWith("session.")) { + assertEquals("surface-session", MAPPER.valueToTree(call.params()).path("sessionId").asText(), + targetMethod.signature()); + } + } + } + + @Test + void rpcCallerOverloadsHaveDirectContractCoverage() throws Exception { + var calls = new java.util.concurrent.CopyOnWriteArrayList(); + RpcCaller caller = new RpcCaller() { + @Override + public CompletableFuture invoke(String method, Object params, Class resultType) { + calls.add(new RpcSurfaceTestCli.RecordingCaller.Call(method, params, resultType)); + if (resultType == JsonNode.class) { + return CompletableFuture.completedFuture(resultType.cast(json(""" + {"value":"deserialized"} + """))); + } + return CompletableFuture.completedFuture(null); + } + }; + + caller.invoke("contract.class", Map.of("kind", "class"), Void.class).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var result = caller + .invoke("contract.javaType", Map.of("kind", "javaType"), + MAPPER.getTypeFactory().constructMapType(Map.class, String.class, String.class)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + caller.invoke("contract.javaTypeVoid", Map.of("kind", "javaTypeVoid"), + MAPPER.getTypeFactory().constructType(Void.class)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + assertEquals(Map.of("value", "deserialized"), result); + assertEquals(List.of("contract.class", "contract.javaType", "contract.javaTypeVoid"), + calls.stream().map(RpcSurfaceTestCli.RecordingCaller.Call::method).toList()); + assertEquals(List.of(Void.class, JsonNode.class, Void.class), + calls.stream().map(RpcSurfaceTestCli.RecordingCaller.Call::resultType).toList()); + } + + @Test + void omittedNamespaceMethodsUseExactGeneratedEntryPoints() { + var caller = new RpcSurfaceTestCli.RecordingCaller(); + var rpc = new SessionRpc(caller, "direct-session"); + + rpc.permissions.paths.list(); + rpc.permissions.paths.add(params(SessionPermissionsPathsAddParams.class, "{}")); + rpc.permissions.paths.updatePrimary(params(SessionPermissionsPathsUpdatePrimaryParams.class, "{}")); + rpc.permissions.paths.isPathWithinAllowedDirectories( + params(SessionPermissionsPathsIsPathWithinAllowedDirectoriesParams.class, "{}")); + rpc.permissions.paths + .isPathWithinWorkspace(params(SessionPermissionsPathsIsPathWithinWorkspaceParams.class, "{}")); + rpc.permissions.urls.setUnrestrictedMode(params(SessionPermissionsUrlsSetUnrestrictedModeParams.class, "{}")); + + rpc.plan.read(); + rpc.plan.update(params(SessionPlanUpdateParams.class, "{}")); + rpc.plan.delete(); + rpc.plan.readSqlTodos(); + rpc.plan.readSqlTodosWithDependencies(); + + rpc.provider.getEndpoint(); + rpc.provider.getEndpoint(params(SessionProviderGetEndpointParams.class, "{}")); + rpc.provider.add(params(SessionProviderAddParams.class, "{}")); + + rpc.queue.pendingItems(); + rpc.queue.snapshot(); + rpc.queue.moveItem(params(SessionQueueMoveItemParams.class, "{}")); + rpc.queue.insertAt(params(SessionQueueInsertAtParams.class, "{}")); + rpc.queue.removeAt(params(SessionQueueRemoveAtParams.class, "{}")); + rpc.queue.updateText(params(SessionQueueUpdateTextParams.class, "{}")); + rpc.queue.duplicateAt(params(SessionQueueDuplicateAtParams.class, "{}")); + rpc.queue.setDrainPaused(params(SessionQueueSetDrainPausedParams.class, "{}")); + rpc.queue.sendNow(params(SessionQueueSendNowParams.class, "{}")); + rpc.queue.hasPending(); + rpc.queue.beginDeferredIdleDrain(params(SessionQueueBeginDeferredIdleDrainParams.class, "{}")); + rpc.queue.finishDeferredIdleDrain(params(SessionQueueFinishDeferredIdleDrainParams.class, "{}")); + rpc.queue.deferSessionIdle(params(SessionQueueDeferSessionIdleParams.class, "{}")); + rpc.queue.removeMostRecent(); + rpc.queue.clear(); + rpc.queue.consumeSystemNotifications(params(SessionQueueConsumeSystemNotificationsParams.class, "{}")); + rpc.queue.enqueueResumePending(); + rpc.queue.process(); + + rpc.remote.enable(params(SessionRemoteEnableParams.class, "{}")); + rpc.remote.disable(); + rpc.remote.notifySteerableChanged(params(SessionRemoteNotifySteerableChangedParams.class, "{}")); + rpc.sandbox.getEnforcementStatus(); + rpc.sandbox.disableForSession(params(SessionSandboxDisableForSessionParams.class, "{}")); + rpc.settings.snapshot(); + rpc.settings.evaluatePredicate(params(SessionSettingsEvaluatePredicateParams.class, "{}")); + rpc.shell.exec(params(SessionShellExecParams.class, "{}")); + rpc.shell.kill(params(SessionShellKillParams.class, "{}")); + rpc.shell.executeUserRequested(params(SessionShellExecuteUserRequestedParams.class, "{}")); + rpc.shell.cancelUserRequested(params(SessionShellCancelUserRequestedParams.class, "{}")); + + rpc.skills.list(); + rpc.skills.getInvoked(); + rpc.skills.enable(params(SessionSkillsEnableParams.class, "{}")); + rpc.skills.disable(params(SessionSkillsDisableParams.class, "{}")); + rpc.skills.reload(); + rpc.skills.ensureLoaded(); + + rpc.tasks.startAgent(params(SessionTasksStartAgentParams.class, "{}")); + rpc.tasks.list(); + rpc.tasks.register(params(SessionTasksRegisterParams.class, "{}")); + rpc.tasks.update(params(SessionTasksUpdateParams.class, "{}")); + rpc.tasks.refresh(); + rpc.tasks.waitForPending(); + rpc.tasks.getProgress(params(SessionTasksGetProgressParams.class, "{}")); + rpc.tasks.getCurrentPromotable(); + rpc.tasks.promoteToBackground(params(SessionTasksPromoteToBackgroundParams.class, "{}")); + rpc.tasks.promoteCurrentToBackground(); + rpc.tasks.cancel(params(SessionTasksCancelParams.class, "{}")); + rpc.tasks.remove(params(SessionTasksRemoveParams.class, "{}")); + rpc.tasks.sendMessage(params(SessionTasksSendMessageParams.class, "{}")); + + rpc.telemetry.getEngagementId(); + rpc.telemetry.setFeatureOverrides(params(SessionTelemetrySetFeatureOverridesParams.class, "{}")); + rpc.tools.execute(params(SessionToolsExecuteParams.class, "{}")); + rpc.tools.getBuiltinDescriptors(params(SessionToolsGetBuiltinDescriptorsParams.class, "{}")); + rpc.tools.taskCompleteEventData(params(SessionToolsTaskCompleteEventDataParams.class, "{}")); + rpc.tools.handlePendingToolCall(params(SessionToolsHandlePendingToolCallParams.class, "{}")); + rpc.tools.initializeAndValidate(); + rpc.tools.getCurrentMetadata(); + rpc.tools.set(params(SessionToolsSetParams.class, "{}")); + rpc.tools.updateSubagentSettings(params(SessionToolsUpdateSubagentSettingsParams.class, "{}")); + + rpc.ui.ephemeralQuery(params(SessionUiEphemeralQueryParams.class, "{}")); + rpc.ui.elicitation(params(SessionUiElicitationParams.class, "{}")); + rpc.ui.handlePendingElicitation(params(SessionUiHandlePendingElicitationParams.class, "{}")); + rpc.ui.handlePendingUserInput(params(SessionUiHandlePendingUserInputParams.class, "{}")); + rpc.ui.handlePendingSampling(params(SessionUiHandlePendingSamplingParams.class, "{}")); + rpc.ui.handlePendingAutoModeSwitch(params(SessionUiHandlePendingAutoModeSwitchParams.class, "{}")); + rpc.ui.handlePendingSessionLimitsExhausted( + params(SessionUiHandlePendingSessionLimitsExhaustedParams.class, "{}")); + rpc.ui.handlePendingExitPlanMode(params(SessionUiHandlePendingExitPlanModeParams.class, "{}")); + rpc.ui.registerDirectAutoModeSwitchHandler(); + rpc.ui.unregisterDirectAutoModeSwitchHandler( + params(SessionUiUnregisterDirectAutoModeSwitchHandlerParams.class, "{}")); + + rpc.visibility.get(); + rpc.visibility.set(params(SessionVisibilitySetParams.class, "{}")); + rpc.workspaces.getWorkspace(); + rpc.workspaces.updateMetadata(params(SessionWorkspacesUpdateMetadataParams.class, "{}")); + rpc.workspaces.ensure(params(SessionWorkspacesEnsureParams.class, "{}")); + rpc.workspaces.listFiles(); + rpc.workspaces.readFile(params(SessionWorkspacesReadFileParams.class, "{}")); + rpc.workspaces.createFile(params(SessionWorkspacesCreateFileParams.class, "{}")); + rpc.workspaces.statFile(params(SessionWorkspacesStatFileParams.class, "{}")); + rpc.workspaces.createDirectory(params(SessionWorkspacesCreateDirectoryParams.class, "{}")); + rpc.workspaces.removePath(params(SessionWorkspacesRemovePathParams.class, "{}")); + rpc.workspaces.renamePath(params(SessionWorkspacesRenamePathParams.class, "{}")); + rpc.workspaces.listCheckpoints(); + rpc.workspaces.readCheckpoint(params(SessionWorkspacesReadCheckpointParams.class, "{}")); + rpc.workspaces.addSummary(params(SessionWorkspacesAddSummaryParams.class, "{}")); + rpc.workspaces.truncateSummaries(params(SessionWorkspacesTruncateSummariesParams.class, "{}")); + rpc.workspaces.readAutopilotObjective(); + rpc.workspaces.writeAutopilotObjective(params(SessionWorkspacesWriteAutopilotObjectiveParams.class, "{}")); + rpc.workspaces.deleteAutopilotObjective(); + rpc.workspaces.autopilotObjectiveExists(); + rpc.workspaces.saveLargePaste(params(SessionWorkspacesSaveLargePasteParams.class, "{}")); + rpc.workspaces.diff(params(SessionWorkspacesDiffParams.class, "{}")); + + assertEquals(104, caller.calls().size()); + assertTrue(caller.calls().stream().allMatch(call -> call.method().startsWith("session."))); + assertTrue(caller.calls().stream().allMatch( + call -> "direct-session".equals(MAPPER.valueToTree(call.params()).path("sessionId").asText()))); + var methods = caller.calls().stream().map(RpcSurfaceTestCli.RecordingCaller.Call::method) + .collect(Collectors.toSet()); + assertTrue(methods.containsAll(Set.of("session.permissions.paths.list", + "session.permissions.urls.setUnrestrictedMode", "session.plan.readSqlTodosWithDependencies", + "session.provider.add", "session.queue.process", "session.remote.notifySteerableChanged", + "session.sandbox.disableForSession", "session.settings.evaluatePredicate", + "session.shell.cancelUserRequested", "session.skills.ensureLoaded", "session.tasks.sendMessage", + "session.telemetry.setFeatureOverrides", "session.tools.updateSubagentSettings", + "session.ui.unregisterDirectAutoModeSwitchHandler", "session.visibility.set", + "session.workspaces.diff"))); + } + + @Test + void remainingGeneratedMethodsUseExactEntryPoints() { + var caller = new RpcSurfaceTestCli.RecordingCaller(); + var server = new ServerRpc(caller); + var rpc = new SessionRpc(caller, "remaining-session"); + + rpc.agent.list(params(SessionAgentListParams.class, "{}")); + rpc.autopilotObjective.getState(); + rpc.canvas.action.invoke(params(SessionCanvasActionInvokeParams.class, "{}")); + rpc.canvas.close(params(SessionCanvasCloseParams.class, "{}")); + rpc.canvas.list(); + rpc.canvas.listOpen(); + rpc.canvas.open(params(SessionCanvasOpenParams.class, "{}")); + rpc.canvas.provider.register(params(SessionCanvasProviderRegisterParams.class, "{}")); + rpc.canvas.provider.unregister(params(SessionCanvasProviderUnregisterParams.class, "{}")); + rpc.commands.enqueue(params(SessionCommandsEnqueueParams.class, "{}")); + rpc.commands.execute(params(SessionCommandsExecuteParams.class, "{}")); + rpc.commands.finalizeInvocationEffect(params(SessionCommandsFinalizeInvocationEffectParams.class, "{}")); + rpc.commands.list(params(SessionCommandsListParams.class, "{}")); + rpc.commands.respondToQueuedCommand(params(SessionCommandsRespondToQueuedCommandParams.class, "{}")); + rpc.completions.getTriggerCharacters(); + rpc.eventLog.registerInterest(params(SessionEventLogRegisterInterestParams.class, "{}")); + rpc.eventLog.releaseInterest(params(SessionEventLogReleaseInterestParams.class, "{}")); + rpc.eventLog.tail(); + rpc.extensions.sendAttachmentsToMessage(params(SessionExtensionsSendAttachmentsToMessageParams.class, "{}")); + rpc.factory.cancel(params(SessionFactoryCancelParams.class, "{}")); + rpc.factory.getRunDetail(params(SessionFactoryGetRunDetailParams.class, "{}")); + rpc.factory.getRunProgress(params(SessionFactoryGetRunProgressParams.class, "{}")); + rpc.factory.listRuns(params(SessionFactoryListRunsParams.class, "{}")); + rpc.factory.pauseAtCheckpoint(params(SessionFactoryPauseAtCheckpointParams.class, "{}")); + rpc.factory.resumeFromTool(params(SessionFactoryResumeFromToolParams.class, "{}")); + rpc.factory.runFromTool(params(SessionFactoryRunFromToolParams.class, "{}")); + rpc.gitHubAuth.getAllAuthAvailable(); + rpc.gitHubAuth.getCurrentAuthInfo(); + rpc.gitHubAuth.lastAuthErrors(); + rpc.gitHubAuth.login(params(SessionGitHubAuthLoginParams.class, "{}")); + rpc.gitHubAuth.logout(); + rpc.gitHubAuth.logoutUser(params(SessionGitHubAuthLogoutUserParams.class, "{}")); + rpc.gitHubAuth.refreshCopilotUser(); + rpc.gitHubAuth.setCredentials(params(SessionGitHubAuthSetCredentialsParams.class, "{}")); + rpc.gitHubAuth.switchToAuth(params(SessionGitHubAuthSwitchToAuthParams.class, "{}")); + rpc.history.abortManualCompaction(); + rpc.history.cancelBackgroundCompaction(); + rpc.history.compact(params(SessionHistoryCompactParams.class, "{}")); + rpc.history.summarizeForHandoff(); + rpc.instructions.getSources(); + rpc.limitPrediction.predict(); + rpc.lsp.initialize(params(SessionLspInitializeParams.class, "{}")); + rpc.mcp.apps.diagnose(params(SessionMcpAppsDiagnoseParams.class, "{}")); + rpc.mcp.apps.getHostContext(); + rpc.mcp.apps.listTools(params(SessionMcpAppsListToolsParams.class, "{}")); + rpc.mcp.apps.readResource(params(SessionMcpAppsReadResourceParams.class, "{}")); + rpc.mcp.apps.setHostContext(params(SessionMcpAppsSetHostContextParams.class, "{}")); + rpc.mcp.cancelSamplingExecution(params(SessionMcpCancelSamplingExecutionParams.class, "{}")); + rpc.mcp.configureGitHub(params(SessionMcpConfigureGitHubParams.class, "{}")); + rpc.mcp.executeSampling(params(SessionMcpExecuteSamplingParams.class, "{}")); + rpc.mcp.isServerRunning(params(SessionMcpIsServerRunningParams.class, "{}")); + rpc.mcp.oauth.probe(params(SessionMcpOauthProbeParams.class, "{}")); + rpc.mcp.registerExternalClient(params(SessionMcpRegisterExternalClientParams.class, "{}")); + rpc.mcp.reloadWithConfig(params(SessionMcpReloadWithConfigParams.class, "{}")); + rpc.mcp.removeGitHub(); + rpc.mcp.restartServer(params(SessionMcpRestartServerParams.class, "{}")); + rpc.mcp.setEnvValueMode(params(SessionMcpSetEnvValueModeParams.class, "{}")); + rpc.mcp.stopServer(params(SessionMcpStopServerParams.class, "{}")); + rpc.mcp.unregisterExternalClient(params(SessionMcpUnregisterExternalClientParams.class, "{}")); + rpc.metadata.activity(); + rpc.metadata.contextInfo(params(SessionMetadataContextInfoParams.class, "{}")); + rpc.metadata.isProcessing(); + rpc.metadata.recomputeContextTokens(params(SessionMetadataRecomputeContextTokensParams.class, "{}")); + rpc.metadata.recordContextChange(params(SessionMetadataRecordContextChangeParams.class, "{}")); + rpc.metadata.setWorkingDirectory(params(SessionMetadataSetWorkingDirectoryParams.class, "{}")); + rpc.metadata.snapshot(); + rpc.metadata.updateClientMetadata(params(SessionMetadataUpdateClientMetadataParams.class, "{}")); + rpc.model.applyStartupOverlay(params(SessionModelApplyStartupOverlayParams.class, "{}")); + rpc.model.list(); + rpc.model.list(params(SessionModelListParams.class, "{}")); + rpc.model.setReasoningEffort(params(SessionModelSetReasoningEffortParams.class, "{}")); + rpc.name.get(); + rpc.name.set(params(SessionNameSetParams.class, "{}")); + rpc.name.setAuto(params(SessionNameSetAutoParams.class, "{}")); + rpc.options.update(params(SessionOptionsUpdateParams.class, "{}")); + rpc.permissions.configure(params(SessionPermissionsConfigureParams.class, "{}")); + rpc.permissions.folderTrust.addTrusted(params(SessionPermissionsFolderTrustAddTrustedParams.class, "{}")); + rpc.permissions.folderTrust.isTrusted(params(SessionPermissionsFolderTrustIsTrustedParams.class, "{}")); + rpc.permissions.getMode(); + rpc.permissions.locations.addToolApproval(params(SessionPermissionsLocationsAddToolApprovalParams.class, "{}")); + rpc.permissions.locations.apply(params(SessionPermissionsLocationsApplyParams.class, "{}")); + rpc.permissions.locations.resolve(params(SessionPermissionsLocationsResolveParams.class, "{}")); + rpc.permissions.modifyRules(params(SessionPermissionsModifyRulesParams.class, "{}")); + rpc.permissions.notifyPromptShown(params(SessionPermissionsNotifyPromptShownParams.class, "{}")); + rpc.permissions.pendingRequests(); + rpc.permissions.resetSessionApprovals(params(SessionPermissionsResetSessionApprovalsParams.class, "{}")); + rpc.permissions.setMode(params(SessionPermissionsSetModeParams.class, "{}")); + rpc.permissions.setRequired(params(SessionPermissionsSetRequiredParams.class, "{}")); + rpc.plugins.reload(); + rpc.plugins.reload(params(SessionPluginsReloadParams.class, "{}")); + rpc.schedule.add(params(SessionScheduleAddParams.class, "{}")); + rpc.schedule.addAt(params(SessionScheduleAddAtParams.class, "{}")); + rpc.schedule.addCron(params(SessionScheduleAddCronParams.class, "{}")); + rpc.schedule.addSelfPaced(params(SessionScheduleAddSelfPacedParams.class, "{}")); + rpc.schedule.hasSelfPaced(); + rpc.schedule.hydrate(); + rpc.schedule.list(); + rpc.schedule.rearmSelfPaced(params(SessionScheduleRearmSelfPacedParams.class, "{}")); + rpc.schedule.stop(params(SessionScheduleStopParams.class, "{}")); + rpc.sendMessages(params(SessionSendMessagesParams.class, "{}")); + rpc.sendSystemNotification(params(SessionSendSystemNotificationParams.class, "{}")); + rpc.shutdown(params(SessionShutdownParams.class, "{}")); + rpc.suspend(); + + server.account.getQuota(params(AccountGetQuotaParams.class, "{}")); + server.agentRegistry.spawn(params(AgentRegistrySpawnParams.class, "{}")); + server.connect(params(ConnectParams.class, "{}")); + server.extensions.disable(params(ExtensionsDisableParams.class, "{}")); + server.extensions.discover(); + server.extensions.enable(params(ExtensionsEnableParams.class, "{}")); + server.mcp.config.disable(params(McpConfigDisableParams.class, "{}")); + server.mcp.config.enable(params(McpConfigEnableParams.class, "{}")); + server.mcp.config.reload(); + server.models.list(params(ModelsListParams.class, "{}")); + server.plugins.disable(params(PluginsDisableParams.class, "{}")); + server.plugins.enable(params(PluginsEnableParams.class, "{}")); + server.plugins.install(params(PluginsInstallParams.class, "{}")); + server.plugins.list(); + server.plugins.marketplaces.add(params(PluginsMarketplacesAddParams.class, "{}")); + server.plugins.marketplaces.browse(params(PluginsMarketplacesBrowseParams.class, "{}")); + server.plugins.marketplaces.list(); + server.plugins.marketplaces.refresh(); + server.plugins.marketplaces.refresh(params(PluginsMarketplacesRefreshParams.class, "{}")); + server.plugins.marketplaces.remove(params(PluginsMarketplacesRemoveParams.class, "{}")); + server.plugins.uninstall(params(PluginsUninstallParams.class, "{}")); + server.plugins.update(params(PluginsUpdateParams.class, "{}")); + server.plugins.updateAll(); + server.runtime.shutdown(); + server.sessions.configureSessionExtensions(params(SessionsConfigureSessionExtensionsParams.class, "{}")); + server.sessions.delete(params(SessionsDeleteParams.class, "{}")); + server.sessions.getBoardEntryCount(params(SessionsGetBoardEntryCountParams.class, "{}")); + server.sessions.getMetadata(params(SessionsGetMetadataParams.class, "{}")); + server.sessions.getRemoteControlStatus(); + server.sessions.list(params(SessionsListParams.class, "{}")); + server.sessions.listNonEmptySessionIds(params(SessionsListNonEmptySessionIdsParams.class, "{}")); + server.sessions.open((SessionsOpenParams) fixture(SessionsOpenParams.class)); + server.sessions.readPersistedEvents(params(SessionsReadPersistedEventsParams.class, "{}")); + server.sessions.setRemoteControlSteering(params(SessionsSetRemoteControlSteeringParams.class, "{}")); + server.sessions.startRemoteControl(params(SessionsStartRemoteControlParams.class, "{}")); + server.sessions.stopRemoteControl(); + server.sessions.stopRemoteControl(params(SessionsStopRemoteControlParams.class, "{}")); + server.sessions.transferRemoteControl(params(SessionsTransferRemoteControlParams.class, "{}")); + + assertEquals(141, caller.calls().size()); + assertTrue(caller.calls().stream().filter(call -> call.method().startsWith("session.")).allMatch( + call -> "remaining-session".equals(MAPPER.valueToTree(call.params()).path("sessionId").asText()))); + } + + @Test + void protocolErrorsPreserveCodeAndMessage() throws Exception { + try (var runtime = new RpcSurfaceTestCli(request -> { + if ("connect".equals(request.path("method").asText())) { + return json(""" + {"ok":true,"protocolVersion":3,"version":"rpc-surface-test"} + """); + } + if ("runtime.shutdown".equals(request.path("method").asText())) { + return MAPPER.createObjectNode(); + } + throw RpcSurfaceTestCli.error(-32042, "rpc surface rejected", json(""" + {"reason":"policy","retryable":false} + """)); + }); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + var failure = assertThrows(ExecutionException.class, () -> client.getRpc().catalog + .search(params(CatalogSearchParams.class, "{}")).get(TIMEOUT_SECONDS, TimeUnit.SECONDS)); + var rpcFailure = assertInstanceOf(JsonRpcException.class, failure.getCause()); + assertEquals(-32042, rpcFailure.getCode()); + assertEquals("rpc surface rejected", rpcFailure.getMessage()); + assertEquals(1, runtime.requestCount("catalog.search")); + assertTrue(parameters(runtime, "catalog.search").isObject()); + } + } + + @Test + void serverRpcsSerializeRequestsAndProjectNestedResults() throws Exception { + try (var runtime = new RpcSurfaceTestCli(RpcSurfaceParityE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var rpc = client.getRpc(); + + rpc.registerExtensionLaunchProvider().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + var command = rpc.commands.list().get(TIMEOUT_SECONDS, TimeUnit.SECONDS).commands().get(0); + assertEquals("rpc-command", command.name()); + assertEquals(List.of("rpc"), command.aliases()); + assertTrue(command.allowDuringAgentExecution()); + assertTrue(command.schedulable()); + + var hooks = rpc.hooks.discover(params(HooksDiscoverParams.class, """ + {"projectPaths":["Q:\\\\rpc-project"],"excludeHostHooks":true} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(List.of("rpc-warning"), hooks.warnings()); + assertTrue(hooks.hooks().isEmpty()); + assertTrue(hooks.errors().isEmpty()); + + assertTrue(rpc.llmInference.setProvider().get(TIMEOUT_SECONDS, TimeUnit.SECONDS).success()); + assertEquals("strict", + ((Map) rpc.managedSettings.read().get(TIMEOUT_SECONDS, TimeUnit.SECONDS).settingsJson()) + .get("policy")); + + var plan = assertInstanceOf(McpPlanInstallPlanned.class, + rpc.mcp.planInstall(params(McpPlanInstallParams.class, """ + { + "contract":{"protocolVersion":3,"requiredCapabilities":["mcp-install-planning"]}, + "source":{"kind":"candidate","candidateHandle":"candidate-1","searchId":"search-1"}, + "scope":"user" + } + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS)); + assertEquals("plan-1", plan.getPlan().planHandle()); + assertTrue(plan.getPlan().reloadRequired()); + assertEquals(3L, plan.getNegotiated().runtimeProtocolVersion()); + + assertEquals("built-in-model", + rpc.models.getBuiltInCatalog().get(TIMEOUT_SECONDS, TimeUnit.SECONDS).models().get(0).id()); + rpc.plugins.builtin.set(new PluginsBuiltinSetParams(List.of("Q:\\rpc-plugins"))).get(TIMEOUT_SECONDS, + TimeUnit.SECONDS); + + var metadata = rpc.sessions + .getClientMetadata( + new SessionsGetClientMetadataParams(List.of("persisted-session"), List.of("rpc/key"))) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("ok", ((Map) metadata.get(0)).get("status")); + assertEquals("rpc-value", ((Map) ((Map) metadata.get(0)).get("metadata")).get("rpc/key")); + + rpc.skills.config.setSkillDisabled(new SkillsConfigSetSkillDisabledParams("skill-one", true)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + assertCalledOnce(runtime, "registerExtensionLaunchProvider", "commands.list", "hooks.discover", + "llmInference.setProvider", "managedSettings.read", "mcp.planInstall", "models.getBuiltInCatalog", + "plugins.builtin.set", "sessions.getClientMetadata", "skills.config.setSkillDisabled"); + assertTrue(parameters(runtime, "hooks.discover").path("excludeHostHooks").asBoolean()); + assertEquals("candidate", parameters(runtime, "mcp.planInstall").path("source").path("kind").asText()); + assertEquals("skill-one", parameters(runtime, "skills.config.setSkillDisabled").path("name").asText()); + } + } + + @Test + void sessionControlRpcsSerializeRequestsAndProjectUnionsAndState() throws Exception { + try (var runtime = new RpcSurfaceTestCli(RpcSurfaceParityE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + try (var session = client + .createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + var rpc = session.getRpc(); + + rpc.agent.setPrompt(new SessionAgentSetPromptParams(null, "agent-1", "Use the RPC prompt.")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + var exclusion = rpc.contentExclusion + .checkPaths( + new SessionContentExclusionCheckPathsParams(null, List.of("/rpc-workspace/file.txt"))) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(exclusion.available()); + assertFalse(exclusion.checks().get(0).excluded()); + + var logs = rpc.debug.collectLogs(params(SessionDebugCollectLogsParams.class, """ + { + "destination":{"kind":"directory","outputDirectory":"/rpc-debug"}, + "include":{"events":true,"processLogs":false,"shellLogs":true}, + "additionalEntries":[{"bundlePath":"host/diagnostic.txt","kind":"file", + "path":"/diagnostic.txt","required":true}] + } + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(DebugCollectLogsResultKind.DIRECTORY, logs.kind()); + assertEquals(123L, logs.entries().get(0).sizeBytes()); + assertEquals("not found", logs.skippedEntries().get(0).reason()); + + assertEquals(4L, rpc.history.clearContext(new SessionHistoryClearContextParams(null, "Reset context.")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).messagesCleared()); + + var prediction = assertInstanceOf(SessionLimitPredictionResultUnavailable.class, + rpc.limitPrediction.predict(params(SessionLimitPredictionPredictParams.class, """ + {"request":{"clientType":"sdk","modelId":"model-a"}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS)); + assertEquals(SessionLimitPredictionUnavailableReason.AUTO_UNRESOLVED, prediction.getReason()); + + var metadata = rpc.metadata.getClientMetadata().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertNotNull(metadata); + + var allowed = rpc.model + .setAllowedModels(new SessionModelSetAllowedModelsParams(null, List.of("model-a", "model-b"))) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(List.of("model-a", "model-b"), allowed.allowedModels()); + assertEquals("model-a", allowed.fallbackModel()); + + var tier = rpc.model + .switchAutoTier(new SessionModelSwitchAutoTierParams(null, AutoTier.INTELLIGENCE, null)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(ModelSwitchAutoTierStatus.PENDING, tier.status()); + assertEquals(AutoTier.INTELLIGENCE, tier.effectiveAutoTier()); + assertEquals(AutoTier.BALANCE, tier.supersededAutoTier()); + + var enforcement = rpc.sandbox.getEnforcementStatus().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(enforcement.required()); + assertFalse(enforcement.blocked()); + assertEquals("managed-policy", enforcement.reason()); + + var disabled = rpc.sandbox + .disableForSession(new SessionSandboxDisableForSessionParams(null, "sandbox-request-1", null)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(disabled.success()); + assertFalse(disabled.enabled()); + + assertTrue(rpc.abort(new SessionAbortParams(null, AbortReason.USER_INITIATED)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).success()); + assertTrue(rpc.interruptMainTurn(new SessionInterruptMainTurnParams(null, true)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).interrupted()); + rpc.cancelAllBackgroundAgents().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + var log = rpc.log(params(SessionLogParams.class, """ + {"message":"RPC log","level":"warning","type":"rpc","ephemeral":true, + "url":"https://example.test/rpc","tip":"Inspect the RPC."} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("11111111-2222-3333-4444-555555555555", log.eventId().toString()); + + assertCalledOnce(runtime, "session.agent.setPrompt", "session.contentExclusion.checkPaths", + "session.debug.collectLogs", "session.history.clearContext", "session.limitPrediction.predict", + "session.metadata.getClientMetadata", "session.model.setAllowedModels", + "session.model.switchAutoTier", "session.sandbox.getEnforcementStatus", + "session.sandbox.disableForSession", "session.abort", "session.interruptMainTurn", + "session.cancelAllBackgroundAgents", "session.log"); + assertEquals(session.getSessionId(), parameters(runtime, "session.log").path("sessionId").asText()); + assertTrue(parameters(runtime, "session.interruptMainTurn").path("flushQueued").asBoolean()); + } + } + } + + @Test + void factoryAndMcpRpcsSerializeRequestsAndProjectStateTransitions() throws Exception { + try (var runtime = new RpcSurfaceTestCli(RpcSurfaceParityE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + try (var session = client + .createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + var rpc = session.getRpc(); + + var run = rpc.factory.run(params(SessionFactoryRunParams.class, """ + {"name":"rpc-factory","args":{"input":42}, + "options":{"limits":{"maxAiCredits":2.5,"maxConcurrentSubagents":2, + "maxTotalSubagents":4,"timeoutSeconds":30}, + "logPhaseNames":true,"notifyOnComplete":false}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("factory-run-1", run.runId()); + assertEquals(FactoryRunStatus.RUNNING, run.status()); + assertEquals(1L, run.attempt()); + + var resumed = rpc.factory.resume(params(SessionFactoryResumeParams.class, """ + {"runId":"factory-run-1","limits":{"maxTotalSubagents":8}, + "notifyOnComplete":true,"logPhaseNames":false} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("rpc-factory", resumed.factoryName()); + assertEquals(2L, resumed.run().attempt()); + + assertEquals(FactoryRunStatus.RUNNING, + rpc.factory.getRun(new SessionFactoryGetRunParams(null, "factory-run-1")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).status()); + assertEquals(FactoryRunStatus.PAUSED, + rpc.factory.pause(new SessionFactoryPauseParams(null, "factory-run-1")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).status()); + + rpc.factory.log(params(SessionFactoryLogParams.class, """ + {"runId":"factory-run-1","executionToken":"execution-token-1", + "lines":[{"kind":"log","seq":7,"text":"Factory progress"}]} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var agent = rpc.factory.agent(params(SessionFactoryAgentParams.class, """ + {"runId":"factory-run-1","executionToken":"execution-token-1", + "prompt":"Complete the RPC task.", + "options":{"agent":"explore","label":"rpc-agent","model":"model-a", + "reasoningEffort":"high"}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("agent-result", ((Map) agent.result()).get("answer")); + + var journal = rpc.factory.journal.get( + new SessionFactoryJournalGetParams(null, "factory-run-1", "execution-token-1", "checkpoint")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(journal.hit()); + assertEquals(7, ((Map) journal.resultJson()).get("checkpoint")); + rpc.factory.journal.put(params(SessionFactoryJournalPutParams.class, """ + {"runId":"factory-run-1","executionToken":"execution-token-1", + "key":"checkpoint","resultJson":{"checkpoint":8}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + assertTrue( + rpc.mcp.moveLoadingToBackground().get(TIMEOUT_SECONDS, TimeUnit.SECONDS).movedToBackground()); + rpc.mcp.startServer(params(SessionMcpStartServerParams.class, """ + {"serverName":"rpc-server","config":{"command":"node","args":["server.js"]}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + rpc.mcp.oauth + .authenticationStateChanged( + new SessionMcpOauthAuthenticationStateChangedParams(null, "rpc-server", true)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(rpc.mcp.oauth.respond(new SessionMcpOauthRespondParams(null, "oauth-request-1")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).success()); + + var resources = rpc.mcp.resources + .list(new SessionMcpResourcesListParams(null, "rpc-server", "resource-cursor")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("resource-next", resources.nextCursor()); + assertEquals("RPC resource", resources.resources().get(0).name()); + var templates = rpc.mcp.resources + .listTemplates( + new SessionMcpResourcesListTemplatesParams(null, "rpc-server", "template-cursor")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("file://rpc/{name}", templates.resourceTemplates().get(0).uriTemplate()); + var content = rpc.mcp.resources + .read(new SessionMcpResourcesReadParams(null, "rpc-server", "file://rpc/resource.txt")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).contents().get(0); + assertEquals("resource-content", content.text()); + assertEquals("assistant", content.meta().get("audience")); + + assertCalledOnce(runtime, "session.factory.run", "session.factory.resume", "session.factory.getRun", + "session.factory.pause", "session.factory.log", "session.factory.agent", + "session.factory.journal.get", "session.factory.journal.put", + "session.mcp.moveLoadingToBackground", "session.mcp.startServer", + "session.mcp.oauth.authenticationStateChanged", "session.mcp.oauth.respond", + "session.mcp.resources.list", "session.mcp.resources.listTemplates", + "session.mcp.resources.read"); + assertEquals(42, parameters(runtime, "session.factory.run").path("args").path("input").asInt()); + assertEquals("resource-cursor", + parameters(runtime, "session.mcp.resources.list").path("cursor").asText()); + } + } + } + + @Test + void taskToolAndWorkspaceRpcsSerializeMutationsAndProjectResults() throws Exception { + try (var runtime = new RpcSurfaceTestCli(RpcSurfaceParityE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + try (var session = client + .createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + var rpc = session.getRpc(); + + var registered = rpc.tasks.register(params(SessionTasksRegisterParams.class, """ + {"type":"client","clientTaskId":"client-task-1","description":"RPC task", + "cancellable":true,"displayName":"RPC Task"} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(registered.created()); + assertEquals("task-1", registered.task().id()); + assertEquals("RPC owner", registered.task().owner().displayName()); + + var updated = rpc.tasks.update(params(SessionTasksUpdateParams.class, """ + {"id":"task-1","sequence":1, + "update":{"kind":"progress","message":"Halfway","percentage":50, + "phase":"work","status":"running"}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(updated.applied()); + assertEquals(1L, updated.task().sequence()); + + rpc.tools.execute(params(SessionToolsExecuteParams.class, """ + {"name":"rpc_tool","arguments":{"value":"input"},"toolCallId":"tool-call-1"} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var descriptors = rpc.tools + .getBuiltinDescriptors(params(SessionToolsGetBuiltinDescriptorsParams.class, """ + {"reduceUserIntervention":true,"includeAuthor":true,"skillEmbeddingEnabled":false, + "shellConfig":{"displayName":"PowerShell","shellType":"powershell", + "shellToolName":"shell","listShellsToolName":"list_shells", + "readShellToolName":"read_shell","stopShellToolName":"stop_shell", + "descriptionLines":["Runs shell commands."]}, + "shellSupportsPowerShell7Syntax":true,"shellTimeoutMs":1234, + "backgroundTaskNotificationsEnabled":true} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("rpc_builtin", descriptors.tools().get(0).name()); + assertEquals(BuiltinToolInputSchemaType.OBJECT, descriptors.tools().get(0).inputSchema().type()); + + rpc.tools.set(params(SessionToolsSetParams.class, """ + {"tools":[{"name":"rpc_external","title":"RPC External", + "description":"External RPC tool","parameters":{"type":"object"}, + "isTerminal":false,"overridesBuiltInTool":false,"skipPermission":true}]} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var completion = rpc.tools + .taskCompleteEventData(params(SessionToolsTaskCompleteEventDataParams.class, """ + {"arguments":{"objectiveId":17}, + "result":{"resultType":"success","textResultForLlm":"RPC task complete", + "sessionLog":"Completion logged."}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(17L, completion.objectiveId()); + assertEquals(TaskCompletionOutcome.COMPLETED, completion.outcome()); + assertTrue(completion.success()); + + var workspace = rpc.workspaces.updateMetadata(params(SessionWorkspacesUpdateMetadataParams.class, """ + {"context":{"owner":"rpc-test"},"name":"Updated RPC workspace"} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("/rpc-workspace", workspace.path()); + assertEquals("Updated RPC workspace", workspace.workspace().name()); + assertEquals("RPC workspace", rpc.workspaces.ensure(params(SessionWorkspacesEnsureParams.class, """ + {"context":{"owner":"rpc-test"}} + """)).get(TIMEOUT_SECONDS, TimeUnit.SECONDS).workspace().name()); + var stat = rpc.workspaces.statFile(new SessionWorkspacesStatFileParams(null, "folder/file.txt")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertTrue(stat.isFile()); + assertEquals(42L, stat.size()); + + rpc.workspaces.createDirectory(new SessionWorkspacesCreateDirectoryParams(null, "folder/nested", true)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + rpc.workspaces + .renamePath( + new SessionWorkspacesRenamePathParams(null, "folder/file.txt", "folder/renamed.txt")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + rpc.workspaces.removePath(new SessionWorkspacesRemovePathParams(null, "folder", true, true)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("RPC summary", rpc.workspaces + .addSummary(new SessionWorkspacesAddSummaryParams(null, "RPC summary", "Summary content")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).summary().get("title")); + assertEquals("Truncated RPC workspace", + rpc.workspaces.truncateSummaries(new SessionWorkspacesTruncateSummariesParams(null, 2L)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS).workspace().name()); + + assertCalledOnce(runtime, "session.tasks.register", "session.tasks.update", "session.tools.execute", + "session.tools.getBuiltinDescriptors", "session.tools.set", + "session.tools.taskCompleteEventData", "session.workspaces.updateMetadata", + "session.workspaces.ensure", "session.workspaces.statFile", + "session.workspaces.createDirectory", "session.workspaces.renamePath", + "session.workspaces.removePath", "session.workspaces.addSummary", + "session.workspaces.truncateSummaries"); + assertTrue(parameters(runtime, "session.tools.set").path("tools").get(0).path("skipPermission") + .asBoolean()); + assertTrue(parameters(runtime, "session.workspaces.removePath").path("force").asBoolean()); + } + } + } + + private static CopilotClient createClient(RpcSurfaceTestCli runtime) { + var client = new CopilotClient( + new CopilotClientOptions().setConnection(RuntimeConnection.forInProcess()).setUseLoggedInUser(false)); + client.setInProcessTransportFactory(options -> runtime.open()); + return client; + } + + private static JsonNode handle(JsonNode request) { + String method = request.path("method").asText(); + return switch (method) { + case "connect" -> json(""" + {"ok":true,"protocolVersion":3,"version":"rpc-surface-test"} + """); + case "session.create" -> sessionCreate(request); + case "session.detach" -> json(""" + {"success":true} + """); + case "commands.list" -> json(""" + {"commands":[{"name":"rpc-command","description":"RPC command","aliases":["rpc"], + "allowDuringAgentExecution":true,"experimental":false, + "input":{"hint":"","preserveMultilineInput":false,"required":true}, + "kind":"builtin","schedulable":true}]} + """); + case "hooks.discover" -> json(""" + {"hooks":[],"warnings":["rpc-warning"],"errors":[]} + """); + case "llmInference.setProvider" -> json(""" + {"success":true} + """); + case "managedSettings.read" -> json(""" + {"settingsJson":{"policy":"strict"},"errorMessage":null} + """); + case "mcp.planInstall" -> json(""" + {"kind":"planned","plan":{"planHandle":"plan-1","planHandleExpiresAt":"2026-09-18T12:00:00Z", + "transportChoices":[],"configurationChanges":[],"reloadRequired":true, + "requiresInteractiveConfiguration":false}, + "negotiated":{"runtimeProtocolVersion":3,"grantedCapabilities":[]}} + """); + case "models.getBuiltInCatalog" -> json(""" + {"models":[{"id":"built-in-model","name":"Built-in Model","family":"test-family"}]} + """); + case "sessions.getClientMetadata" -> json(""" + [{"status":"ok","sessionId":"persisted-session","metadata":{"rpc/key":"rpc-value"}}] + """); + case "session.contentExclusion.checkPaths" -> json(""" + {"available":true,"checks":[{"path":"/rpc-workspace/file.txt","excluded":false}]} + """); + case "session.debug.collectLogs" -> json(""" + {"kind":"directory","path":"/rpc-debug", + "entries":[{"bundlePath":"host/diagnostic.txt","sizeBytes":123,"source":"additional"}], + "skippedEntries":[{"bundlePath":"host/missing.txt","path":"/missing.txt","reason":"not found"}]} + """); + case "session.history.clearContext" -> json(""" + {"messagesCleared":4} + """); + case "session.limitPrediction.predict" -> json(""" + {"kind":"unavailable","reason":"auto_unresolved"} + """); + case "session.metadata.getClientMetadata" -> json(""" + {"rpc/key":"rpc-value","rpc/other":"other-value"} + """); + case "session.model.setAllowedModels" -> json(""" + {"allowedModels":["model-a","model-b"],"effectiveAllowedModels":["model-a"], + "fallbackModel":"model-a","modelId":"model-a"} + """); + case "session.model.switchAutoTier" -> json(""" + {"status":"pending","activatingAutoTier":"intelligence","effectiveAutoTier":"intelligence", + "pendingAutoTier":null,"supersededAutoTier":"balance"} + """); + case "session.sandbox.getEnforcementStatus" -> json(""" + {"required":true,"blocked":false,"reason":"managed-policy"} + """); + case "session.sandbox.disableForSession" -> json(""" + {"success":true,"enabled":false} + """); + case "session.abort" -> json(""" + {"success":true,"error":null} + """); + case "session.interruptMainTurn" -> json(""" + {"interrupted":true} + """); + case "session.log" -> json(""" + {"eventId":"11111111-2222-3333-4444-555555555555"} + """); + case "session.factory.run", "session.factory.getRun" -> json(""" + {"runId":"factory-run-1","status":"running","attempt":1, + "result":{"value":"running"},"snapshot":{"step":1}} + """); + case "session.factory.resume" -> json(""" + {"factoryName":"rpc-factory","run":{"runId":"factory-run-1","status":"running", + "attempt":2,"snapshot":{"step":3}}} + """); + case "session.factory.pause" -> json(""" + {"runId":"factory-run-1","status":"paused","attempt":1, + "reason":"caller requested pause","snapshot":{"step":2}} + """); + case "session.factory.agent" -> json(""" + {"result":{"answer":"agent-result"}} + """); + case "session.factory.journal.get" -> json(""" + {"hit":true,"resultJson":{"checkpoint":7}} + """); + case "session.mcp.moveLoadingToBackground" -> json(""" + {"movedToBackground":true} + """); + case "session.mcp.oauth.respond" -> json(""" + {"success":true} + """); + case "session.mcp.resources.list" -> json(""" + {"nextCursor":"resource-next","resources":[{"uri":"file://rpc/resource.txt", + "name":"RPC resource","description":"Resource description","mimeType":"text/plain", + "size":16,"title":"RPC Resource"}]} + """); + case "session.mcp.resources.listTemplates" -> json(""" + {"nextCursor":"template-next","resourceTemplates":[{"uriTemplate":"file://rpc/{name}", + "name":"RPC template","description":"Template description","mimeType":"text/plain", + "title":"RPC Template"}]} + """); + case "session.mcp.resources.read" -> json(""" + {"contents":[{"uri":"file://rpc/resource.txt","mimeType":"text/plain", + "text":"resource-content","blob":null,"_meta":{"audience":"assistant"}}]} + """); + case "session.tasks.register" -> json(""" + {"created":true,"reclaimed":false,"task":{"id":"task-1","type":"client", + "clientTaskId":"client-task-1","description":"RPC task","displayName":"RPC Task", + "activeStartedAt":"2026-09-18T12:00:00.500Z","activeTimeMs":500,"canCancel":true, + "executionMode":"background","owner":{"displayName":"RPC owner","joinId":"join-1", + "kind":"sdk","participantId":"participant-1","presence":"connected","source":"rpc-test"}, + "sequence":0,"status":"running","startedAt":"2026-09-18T12:00:00Z", + "updatedAt":"2026-09-18T12:00:01Z"}} + """); + case "session.tasks.update" -> json(""" + {"applied":true,"duplicate":false,"task":{"id":"task-1","type":"client", + "clientTaskId":"client-task-1","description":"RPC task","displayName":"RPC Task", + "activeTimeMs":500,"canCancel":true,"executionMode":"background", + "owner":{"displayName":"RPC owner","joinId":"join-1","kind":"sdk", + "participantId":"participant-1","presence":"connected","source":"rpc-test"}, + "sequence":1,"status":"running","startedAt":"2026-09-18T12:00:00Z", + "updatedAt":"2026-09-18T12:00:01Z"}} + """); + case "session.tools.getBuiltinDescriptors" -> json(""" + {"tools":[{"name":"rpc_builtin","description":"RPC built-in tool", + "hasSummariseIntention":true,"inputSchema":{"type":"object"}, + "instructions":"Use the RPC built-in.","isTerminal":false,"safeForTelemetry":true, + "title":"RPC Built-in","type":"test"}]} + """); + case "session.tools.taskCompleteEventData" -> json(""" + {"objectiveId":17,"outcome":"completed","reason":"completed","success":true, + "summary":"RPC task complete"} + """); + case "session.workspaces.updateMetadata" -> workspace("Updated RPC workspace"); + case "session.workspaces.ensure" -> workspace("RPC workspace"); + case "session.workspaces.statFile" -> json(""" + {"birthtimeMs":1000,"isDirectory":false,"isFile":true,"mtimeMs":2000,"size":42} + """); + case "session.workspaces.addSummary" -> json(""" + {"summary":{"number":3,"title":"RPC summary","content":"Summary content"}, + "workspace":{"id":"workspace-1","cwd":"/rpc-workspace","name":"RPC workspace"}} + """); + case "session.workspaces.truncateSummaries" -> workspace("Truncated RPC workspace"); + default -> MAPPER.createObjectNode(); + }; + } + + private static JsonNode workspace(String name) { + return json(""" + {"path":"/rpc-workspace","workspace":{"id":"workspace-1","cwd":"/rpc-workspace", + "git_root":"/rpc-workspace","branch":"rpc-branch","name":"%s","client_name":"rpc-client", + "created_at":"2026-09-18T11:00:00Z","remote_steerable":true}} + """.formatted(name)); + } + + private static JsonNode sessionCreate(JsonNode request) { + var result = MAPPER.createObjectNode(); + result.put("sessionId", request.path("params").path("sessionId").asText()); + result.put("workspacePath", "/rpc-workspace"); + result.putNull("capabilities"); + return result; + } + + private static void collectTargets(Object instance, String prefix, Map, RpcTarget> targets, + IdentityHashMap visited) throws IllegalAccessException { + if (visited.put(instance, Boolean.TRUE) != null) { + return; + } + targets.put(instance.getClass(), new RpcTarget(instance, prefix)); + for (var field : instance.getClass().getFields()) { + if (field.getType().getPackageName().equals(ServerRpc.class.getPackageName()) + && field.getType().getSimpleName().endsWith("Api")) { + var child = field.get(instance); + var childPrefix = prefix.isEmpty() ? field.getName() : prefix + "." + field.getName(); + collectTargets(child, childPrefix, targets, visited); + } + } + } + + private static List rpcMethods(Class type) { + return Arrays.stream(type.getDeclaredMethods()).filter(method -> Modifier.isPublic(method.getModifiers())) + .filter(method -> method.getReturnType() == CompletableFuture.class) + .sorted(Comparator.comparing(RpcSurfaceParityE2ETest::signature)).toList(); + } + + private static Object invoke(Method method, Object instance, Object[] arguments) { + try { + return method.invoke(instance, arguments); + } catch (IllegalAccessException e) { + throw new AssertionError("Could not invoke " + signature(method), e); + } catch (InvocationTargetException e) { + throw new AssertionError("Generated wrapper failed before dispatch for " + signature(method), e.getCause()); + } + } + + private static Object fixture(Class type) { + if (type == String.class) { + return "caller-supplied-session"; + } + if (type == boolean.class || type == Boolean.class) { + return true; + } + if (type == byte.class || type == Byte.class) { + return (byte) 1; + } + if (type == short.class || type == Short.class) { + return (short) 1; + } + if (type == int.class || type == Integer.class) { + return 1; + } + if (type == long.class || type == Long.class) { + return 1L; + } + if (type == float.class || type == Float.class) { + return 1F; + } + if (type == double.class || type == Double.class) { + return 1D; + } + if (type.isEnum()) { + return type.getEnumConstants()[0]; + } + if (JsonNode.class.isAssignableFrom(type)) { + return MAPPER.createObjectNode().put("fixture", true); + } + if (List.class.isAssignableFrom(type)) { + return List.of(); + } + if (Map.class.isAssignableFrom(type)) { + return Map.of(); + } + if (type == Object.class) { + return Map.of("fixture", "value"); + } + if (type.isRecord()) { + try { + RecordComponent[] components = type.getRecordComponents(); + var constructor = type.getDeclaredConstructor( + Arrays.stream(components).map(RecordComponent::getType).toArray(Class[]::new)); + var arguments = Arrays.stream(components).map(component -> fixtureComponent(component.getType())) + .toArray(); + return constructor.newInstance(arguments); + } catch (ReflectiveOperationException e) { + throw new AssertionError("Could not construct RPC params " + type.getName(), e); + } + } + var subTypes = type.getAnnotation(com.fasterxml.jackson.annotation.JsonSubTypes.class); + if (subTypes != null && subTypes.value().length > 0) { + return fixture(subTypes.value()[0].value()); + } + if (!Modifier.isAbstract(type.getModifiers()) && !type.isInterface()) { + try { + return type.getDeclaredConstructor().newInstance(); + } catch (ReflectiveOperationException e) { + throw new AssertionError("Could not construct RPC params " + type.getName(), e); + } + } + throw new AssertionError("Unmapped generated RPC parameter type " + type.getName()); + } + + private static Object fixtureComponent(Class type) { + if (Map.class.isAssignableFrom(type)) { + return null; + } + if (type.isPrimitive() || type == String.class || Number.class.isAssignableFrom(type) || type == Boolean.class + || type.isEnum() || List.class.isAssignableFrom(type) || type == Object.class + || JsonNode.class.isAssignableFrom(type)) { + return fixture(type); + } + return null; + } + + private static RpcSurfaceTestCli.RecordingCaller.Call assertSingleCall(RpcSurfaceTestCli.RecordingCaller caller, + String signature) { + assertEquals(1, caller.calls().size(), signature); + return caller.calls().get(0); + } + + private static String signature(Method method) { + return method.getDeclaringClass().getSimpleName() + "#" + method.getName() + "(" + + Arrays.stream(method.getParameterTypes()).map(Class::getSimpleName).collect(Collectors.joining(",")) + + ")"; + } + + private static String sha256(String value) { + try { + return java.util.HexFormat.of() + .formatHex(MessageDigest.getInstance("SHA-256").digest(value.getBytes(StandardCharsets.UTF_8))); + } catch (java.security.NoSuchAlgorithmException e) { + throw new AssertionError(e); + } + } + + private record RpcTarget(Object instance, String prefix) { + } + + private record TargetMethod(RpcTarget target, Method method) { + + String signature() { + return RpcSurfaceParityE2ETest.signature(method); + } + } + + private static T params(Class type, String json) { + try { + return MAPPER.readValue(json, type); + } catch (JsonProcessingException e) { + throw new AssertionError("Invalid test parameters for " + type.getSimpleName(), e); + } + } + + private static JsonNode json(String json) { + try { + return MAPPER.readTree(json); + } catch (JsonProcessingException e) { + throw new AssertionError("Invalid fake runtime JSON", e); + } + } + + private static JsonNode parameters(RpcSurfaceTestCli runtime, String method) { + return runtime.request(method).path("params"); + } + + private static void assertCalledOnce(RpcSurfaceTestCli runtime, String... methods) { + for (String method : methods) { + assertEquals(1, runtime.requestCount(method), "Unexpected call count for " + method); + if (method.startsWith("session.")) { + assertFalse(parameters(runtime, method).path("sessionId").asText().isBlank(), + "Expected sessionId for " + method); + } + } + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/RpcSurfaceTestCli.java b/java/sdk/src/test/java/com/github/copilot/RpcSurfaceTestCli.java new file mode 100644 index 0000000000..d51abec631 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/RpcSurfaceTestCli.java @@ -0,0 +1,247 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.fail; + +import java.io.ByteArrayOutputStream; +import java.io.Closeable; +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.nio.channels.Channels; +import java.nio.channels.Pipe; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Locale; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Function; + +import com.fasterxml.jackson.databind.JavaType; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.NullNode; +import com.fasterxml.jackson.databind.node.ObjectNode; +import com.github.copilot.generated.rpc.RpcCaller; + +final class RpcSurfaceTestCli implements AutoCloseable { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + record RpcError(int code, String message, JsonNode data) { + } + + static final class RecordingCaller implements RpcCaller { + + record Call(String method, Object params, Object resultType) { + } + + private final List calls = new CopyOnWriteArrayList<>(); + + @Override + public CompletableFuture invoke(String method, Object params, Class resultType) { + calls.add(new Call(method, params, resultType)); + return CompletableFuture.completedFuture(null); + } + + @Override + public CompletableFuture invoke(String method, Object params, JavaType resultType) { + calls.add(new Call(method, params, resultType)); + return CompletableFuture.completedFuture(null); + } + + List calls() { + return List.copyOf(calls); + } + + void clear() { + calls.clear(); + } + } + + private final Function handler; + private final List requests = new CopyOnWriteArrayList<>(); + private final AtomicBoolean closed = new AtomicBoolean(); + private final AtomicReference responderFailure = new AtomicReference<>(); + private final BytePipe toClient; + private final BytePipe toRuntime; + private final InputStream runtimeInput; + private final OutputStream runtimeOutput; + private final Thread responder; + + RpcSurfaceTestCli(Function handler) throws IOException { + this.handler = handler; + this.toClient = new BytePipe(); + this.toRuntime = new BytePipe(); + this.runtimeInput = toRuntime.inputStream(); + this.runtimeOutput = toClient.outputStream(); + this.responder = new Thread(this::respondToRequests, "fake-rpc-runtime"); + this.responder.setDaemon(true); + this.responder.start(); + } + + CopilotClient.InProcessTransport open() { + return new CopilotClient.InProcessTransport(toClient.inputStream(), toRuntime.outputStream(), this::close); + } + + JsonNode request(String method) { + return requests.stream().filter(request -> method.equals(request.path("method").asText())).findFirst() + .orElseGet(() -> { + fail("Expected request for " + method + "; captured methods were " + + requests.stream().map(request -> request.path("method").asText()).toList()); + return NullNode.getInstance(); + }); + } + + long requestCount(String method) { + return requests.stream().filter(request -> method.equals(request.path("method").asText())).count(); + } + + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + assertResponderHealthy(); + return; + } + toRuntime.close(); + toClient.close(); + responder.interrupt(); + assertResponderHealthy(); + } + + private void respondToRequests() { + try { + while (!closed.get()) { + JsonNode request = readMessage(runtimeInput); + if (request == null) { + return; + } + requests.add(request.deepCopy()); + if (!request.hasNonNull("id")) { + continue; + } + + ObjectNode response = MAPPER.createObjectNode(); + response.put("jsonrpc", "2.0"); + response.set("id", request.get("id")); + try { + JsonNode result = handler.apply(request); + response.set("result", result == null ? NullNode.getInstance() : result); + } catch (RpcErrorException error) { + ObjectNode errorNode = response.putObject("error"); + errorNode.put("code", error.error.code()); + errorNode.put("message", error.error.message()); + if (error.error.data() != null) { + errorNode.set("data", error.error.data()); + } + } catch (Throwable failure) { + responderFailure.compareAndSet(null, failure); + ObjectNode errorNode = response.putObject("error"); + errorNode.put("code", -32603); + errorNode.put("message", "Fake runtime handler failed: " + failure.getMessage()); + } + writeMessage(runtimeOutput, response); + } + } catch (IOException e) { + if (!closed.get()) { + responderFailure.compareAndSet(null, e); + } + } + } + + private void assertResponderHealthy() { + Throwable failure = responderFailure.get(); + if (failure != null) { + throw new AssertionError("Fake runtime failed", failure); + } + } + + static RuntimeException error(int code, String message, JsonNode data) { + return new RpcErrorException(new RpcError(code, message, data)); + } + + private static JsonNode readMessage(InputStream in) throws IOException { + int contentLength = -1; + var line = new ByteArrayOutputStream(); + while (true) { + int b = in.read(); + if (b == -1) { + return null; + } + if (b == '\n') { + String header = line.toString(StandardCharsets.UTF_8).trim(); + line.reset(); + if (header.isEmpty()) { + break; + } + if (header.toLowerCase(Locale.ROOT).startsWith("content-length:")) { + try { + contentLength = Integer.parseInt(header.substring(header.indexOf(':') + 1).trim()); + } catch (NumberFormatException e) { + throw new IOException("Invalid Content-Length header: " + header, e); + } + } + } else if (b != '\r') { + line.write(b); + } + } + if (contentLength < 0) { + throw new IOException("Missing Content-Length header"); + } + byte[] body = in.readNBytes(contentLength); + return body.length == contentLength ? MAPPER.readTree(body) : null; + } + + private static void writeMessage(OutputStream out, JsonNode message) throws IOException { + byte[] body = MAPPER.writeValueAsBytes(message); + out.write(("Content-Length: " + body.length + "\r\n\r\n").getBytes(StandardCharsets.UTF_8)); + out.write(body); + out.flush(); + } + + private static final class RpcErrorException extends RuntimeException { + + private static final long serialVersionUID = 1L; + private final RpcError error; + + RpcErrorException(RpcError error) { + super(error.message()); + this.error = error; + } + } + + private static final class BytePipe { + + private final Pipe pipe; + + BytePipe() throws IOException { + this.pipe = Pipe.open(); + } + + InputStream inputStream() { + return Channels.newInputStream(pipe.source()); + } + + OutputStream outputStream() { + return Channels.newOutputStream(pipe.sink()); + } + + void close() { + closeQuietly(pipe.sink()); + closeQuietly(pipe.source()); + } + + private static void closeQuietly(Closeable closeable) { + try { + closeable.close(); + } catch (IOException e) { + // Nothing useful to do while tearing down a test pipe. + } + } + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/ScenarioCoverageE2ETest.java b/java/sdk/src/test/java/com/github/copilot/ScenarioCoverageE2ETest.java new file mode 100644 index 0000000000..658d0d5ac1 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/ScenarioCoverageE2ETest.java @@ -0,0 +1,230 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.*; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.TimeUnit; + +import org.junit.jupiter.api.Test; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.github.copilot.rpc.AgentMode; +import com.github.copilot.rpc.CopilotClientOptions; +import com.github.copilot.rpc.MessageOptions; +import com.github.copilot.rpc.MessageSource; +import com.github.copilot.rpc.PermissionHandler; +import com.github.copilot.rpc.ResumeSessionConfig; +import com.github.copilot.rpc.RuntimeConnection; +import com.github.copilot.rpc.SessionConfig; + +class ScenarioCoverageE2ETest { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + private static final long TIMEOUT_SECONDS = 30; + + @Test + void publicSessionScenarioComposesMessagesAndRemainsUsableAfterAbort() throws Exception { + try (var runtime = new ScenarioTestCli(ScenarioCoverageE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + try (var session = client + .createSession(new SessionConfig().setClientName("scenario-java").setModel("model-a") + .setReasoningEffort("high").setStreaming(true).setWorkingDirectory("Q:\\scenario-work") + .setAvailableTools(List.of("view", "grep")).setExcludedTools(List.of("shell")) + .setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + String queuedId = session + .send(new MessageOptions().setPrompt("Queued scenario message") + .setDisplayPrompt("Queued display").setMode("enqueue").setSource(MessageSource.USER) + .setAgentMode(AgentMode.PLAN).setRequestHeaders(Map.of("x-scenario", "queued"))) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("message-queued", queuedId); + + session.abort().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + session.setModel("model-b", "xhigh").get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + session.log("Scenario recovered after abort", "warning", true, "https://example.test/scenario") + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + String immediateId = session + .send(new MessageOptions().setPrompt("Immediate scenario message").setMode("immediate") + .setSource(MessageSource.USER).setAgentMode(AgentMode.INTERACTIVE) + .setRequestHeaders(Map.of("x-scenario", "immediate"))) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("message-immediate", immediateId); + + JsonNode create = parameters(runtime, "session.create"); + assertEquals("scenario-java", create.path("clientName").asText()); + assertEquals("model-a", create.path("model").asText()); + assertEquals("high", create.path("reasoningEffort").asText()); + assertTrue(create.path("streaming").asBoolean()); + assertEquals("Q:\\scenario-work", create.path("workingDirectory").asText()); + assertEquals(List.of("view", "grep"), MAPPER.convertValue(create.path("availableTools"), List.class)); + assertEquals("shell", create.path("excludedTools").get(0).asText()); + + JsonNode queued = requestParameters(runtime, "session.send", "Queued scenario message"); + assertEquals("enqueue", queued.path("mode").asText()); + assertEquals("plan", queued.path("agentMode").asText()); + assertEquals("queued", queued.path("requestHeaders").path("x-scenario").asText()); + assertEquals("Queued display", queued.path("displayPrompt").asText()); + + JsonNode immediate = requestParameters(runtime, "session.send", "Immediate scenario message"); + assertEquals("immediate", immediate.path("mode").asText()); + assertEquals("interactive", immediate.path("agentMode").asText()); + assertEquals("immediate", immediate.path("requestHeaders").path("x-scenario").asText()); + + assertEquals(1, runtime.requestCount("session.abort")); + assertEquals("model-b", parameters(runtime, "session.model.switchTo").path("modelId").asText()); + assertEquals("xhigh", parameters(runtime, "session.model.switchTo").path("reasoningEffort").asText()); + JsonNode log = parameters(runtime, "session.log"); + assertEquals("Scenario recovered after abort", log.path("message").asText()); + assertEquals("warning", log.path("level").asText()); + assertTrue(log.path("ephemeral").asBoolean()); + } + } + } + + @Test + void publicClientScenarioListsResumesAndDeletesPersistedSession() throws Exception { + try (var runtime = new ScenarioTestCli(ScenarioCoverageE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + + var sessions = client.listSessions().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals(1, sessions.size()); + assertEquals("persisted-scenario", sessions.get(0).getSessionId()); + assertEquals("Persisted Java scenario", sessions.get(0).getSummary()); + assertTrue(sessions.get(0).isRemote()); + + try (var resumed = client + .resumeSession("persisted-scenario", + new ResumeSessionConfig().setWorkingDirectory("Q:\\resumed-scenario") + .setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + assertEquals("persisted-scenario", resumed.getSessionId()); + assertEquals("message-resumed", + resumed.send(new MessageOptions().setPrompt("Continue persisted scenario")).get(TIMEOUT_SECONDS, + TimeUnit.SECONDS)); + } + + client.deleteSession("persisted-scenario").get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + JsonNode resume = parameters(runtime, "session.resume"); + assertEquals("persisted-scenario", resume.path("sessionId").asText()); + assertEquals("Q:\\resumed-scenario", resume.path("workingDirectory").asText()); + assertEquals("persisted-scenario", parameters(runtime, "session.delete").path("sessionId").asText()); + } + } + + @Test + void publicClientScenarioPingsAndReusesOneRuntimeAcrossSessions() throws Exception { + try (var runtime = new ScenarioTestCli(ScenarioCoverageE2ETest::handle); var client = createClient(runtime)) { + client.start().get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var ping = client.ping("scenario-ping").get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + assertEquals("scenario-ping", ping.message()); + assertEquals(3, ping.protocolVersion()); + + try (var first = client + .createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS); + var second = client + .createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + assertNotEquals(first.getSessionId(), second.getSessionId()); + assertEquals("message-session-one", first.send(new MessageOptions().setPrompt("Session one scenario")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)); + assertEquals("message-session-two", second.send(new MessageOptions().setPrompt("Session two scenario")) + .get(TIMEOUT_SECONDS, TimeUnit.SECONDS)); + } + + assertEquals(1, runtime.requestCount("connect")); + assertEquals(1, runtime.requestCount("ping")); + assertEquals(2, runtime.requestCount("session.create")); + assertEquals(2, runtime.requestCount("session.send")); + assertEquals("scenario-ping", parameters(runtime, "ping").path("message").asText()); + } + } + + private static CopilotClient createClient(ScenarioTestCli runtime) { + var client = new CopilotClient( + new CopilotClientOptions().setConnection(RuntimeConnection.forInProcess()).setUseLoggedInUser(false)); + client.setInProcessTransportFactory(options -> runtime.open()); + return client; + } + + private static JsonNode handle(JsonNode request) { + String method = request.path("method").asText(); + JsonNode params = request.path("params"); + return switch (method) { + case "connect" -> json(""" + {"ok":true,"protocolVersion":3,"version":"scenario-parity-test"} + """); + case "ping" -> json(""" + {"message":"scenario-ping","timestamp":"2026-09-18T12:00:00Z","protocolVersion":3} + """); + case "session.create" -> sessionResult(params.path("sessionId").asText()); + case "session.resume" -> sessionResult(params.path("sessionId").asText()); + case "session.send" -> { + String prompt = params.path("prompt").asText(); + String id = switch (prompt) { + case "Queued scenario message" -> "message-queued"; + case "Immediate scenario message" -> "message-immediate"; + case "Continue persisted scenario" -> "message-resumed"; + case "Session one scenario" -> "message-session-one"; + case "Session two scenario" -> "message-session-two"; + default -> throw new AssertionError("Unexpected scenario prompt: " + prompt); + }; + yield json(""" + {"messageId":"%s"} + """.formatted(id)); + } + case "session.abort" -> MAPPER.createObjectNode(); + case "session.detach" -> json(""" + {"success":true} + """); + case "session.model.switchTo" -> json(""" + {"modelId":"model-b","deferred":false,"status":"applied","deprecationWarnings":[]} + """); + case "session.log" -> json(""" + {"eventId":"11111111-2222-3333-4444-555555555555"} + """); + case "session.list" -> json(""" + {"sessions":[{"sessionId":"persisted-scenario","startTime":"2026-09-18T10:00:00Z", + "modifiedTime":"2026-09-18T11:00:00Z","summary":"Persisted Java scenario", + "isRemote":true,"context":{"cwd":"Q:\\\\scenario-work","repository":"github/copilot-sdk", + "branch":"scenario-parity"}}]} + """); + case "session.delete" -> json(""" + {"success":true} + """); + case "runtime.shutdown" -> MAPPER.createObjectNode(); + default -> throw new AssertionError("Unexpected scenario RPC method: " + method); + }; + } + + private static JsonNode sessionResult(String sessionId) { + return json(""" + {"sessionId":"%s","workspacePath":"Q:\\\\scenario-work","capabilities":null,"openCanvases":[]} + """.formatted(sessionId)); + } + + private static JsonNode parameters(ScenarioTestCli runtime, String method) { + return runtime.request(method).path("params"); + } + + private static JsonNode requestParameters(ScenarioTestCli runtime, String method, String prompt) { + return runtime.requests(method).stream().map(request -> request.path("params")) + .filter(params -> prompt.equals(params.path("prompt").asText())).findFirst() + .orElseThrow(() -> new AssertionError("Missing " + method + " request for prompt " + prompt)); + } + + private static JsonNode json(String value) { + try { + return MAPPER.readTree(value); + } catch (Exception e) { + throw new AssertionError("Invalid scenario fake-runtime JSON", e); + } + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/ScenarioTestCli.java b/java/sdk/src/test/java/com/github/copilot/ScenarioTestCli.java new file mode 100644 index 0000000000..58d1677582 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/ScenarioTestCli.java @@ -0,0 +1,191 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.fail; + +import java.io.ByteArrayOutputStream; +import java.io.Closeable; +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.nio.channels.Channels; +import java.nio.channels.Pipe; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Locale; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Function; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.NullNode; +import com.fasterxml.jackson.databind.node.ObjectNode; + +final class ScenarioTestCli implements AutoCloseable { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private final Function handler; + private final List requests = new CopyOnWriteArrayList<>(); + private final AtomicBoolean closed = new AtomicBoolean(); + private final AtomicReference responderFailure = new AtomicReference<>(); + private final BytePipe toClient = new BytePipe(); + private final BytePipe toRuntime = new BytePipe(); + private final InputStream runtimeInput = toRuntime.inputStream(); + private final OutputStream runtimeOutput = toClient.outputStream(); + private final Thread responder; + + ScenarioTestCli(Function handler) throws IOException { + this.handler = handler; + this.responder = new Thread(this::respondToRequests, "scenario-fake-runtime"); + this.responder.setDaemon(true); + this.responder.start(); + } + + CopilotClient.InProcessTransport open() { + return new CopilotClient.InProcessTransport(toClient.inputStream(), toRuntime.outputStream(), this::close); + } + + JsonNode request(String method) { + return requests(method).stream().findFirst().orElseGet(() -> { + fail("Expected request for " + method + "; captured methods were " + + requests.stream().map(request -> request.path("method").asText()).toList()); + return NullNode.getInstance(); + }); + } + + List requests(String method) { + return requests.stream().filter(request -> method.equals(request.path("method").asText())) + .map(request -> (JsonNode) request.deepCopy()).toList(); + } + + long requestCount(String method) { + return requests.stream().filter(request -> method.equals(request.path("method").asText())).count(); + } + + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + assertResponderHealthy(); + return; + } + toRuntime.close(); + toClient.close(); + responder.interrupt(); + assertResponderHealthy(); + } + + private void respondToRequests() { + try { + while (!closed.get()) { + JsonNode request = readMessage(runtimeInput); + if (request == null) { + return; + } + requests.add(request.deepCopy()); + if (!request.hasNonNull("id")) { + continue; + } + + ObjectNode response = MAPPER.createObjectNode(); + response.put("jsonrpc", "2.0"); + response.set("id", request.get("id")); + try { + JsonNode result = handler.apply(request); + response.set("result", result == null ? NullNode.getInstance() : result); + } catch (Throwable failure) { + responderFailure.compareAndSet(null, failure); + ObjectNode errorNode = response.putObject("error"); + errorNode.put("code", -32603); + errorNode.put("message", "Scenario fake runtime handler failed: " + failure.getMessage()); + } + writeMessage(runtimeOutput, response); + } + } catch (IOException e) { + if (!closed.get()) { + responderFailure.compareAndSet(null, e); + } + } + } + + private void assertResponderHealthy() { + Throwable failure = responderFailure.get(); + if (failure != null) { + throw new AssertionError("Scenario fake runtime failed", failure); + } + } + + private static JsonNode readMessage(InputStream in) throws IOException { + int contentLength = -1; + var line = new ByteArrayOutputStream(); + while (true) { + int value = in.read(); + if (value == -1) { + return null; + } + if (value == '\n') { + String header = line.toString(StandardCharsets.UTF_8).trim(); + line.reset(); + if (header.isEmpty()) { + break; + } + if (header.toLowerCase(Locale.ROOT).startsWith("content-length:")) { + try { + contentLength = Integer.parseInt(header.substring(header.indexOf(':') + 1).trim()); + } catch (NumberFormatException e) { + throw new IOException("Invalid Content-Length header: " + header, e); + } + } + } else if (value != '\r') { + line.write(value); + } + } + if (contentLength < 0) { + throw new IOException("Missing Content-Length header"); + } + byte[] body = in.readNBytes(contentLength); + return body.length == contentLength ? MAPPER.readTree(body) : null; + } + + private static void writeMessage(OutputStream out, JsonNode message) throws IOException { + byte[] body = MAPPER.writeValueAsBytes(message); + out.write(("Content-Length: " + body.length + "\r\n\r\n").getBytes(StandardCharsets.UTF_8)); + out.write(body); + out.flush(); + } + + private static final class BytePipe { + + private final Pipe pipe; + + BytePipe() throws IOException { + this.pipe = Pipe.open(); + } + + InputStream inputStream() { + return Channels.newInputStream(pipe.source()); + } + + OutputStream outputStream() { + return Channels.newOutputStream(pipe.sink()); + } + + void close() { + closeQuietly(pipe.sink()); + closeQuietly(pipe.source()); + } + + private static void closeQuietly(Closeable closeable) { + try { + closeable.close(); + } catch (IOException e) { + // Nothing useful to do while tearing down a test pipe. + } + } + } +} diff --git a/nodejs/README.md b/nodejs/README.md index 93b6160808..d69151738c 100644 --- a/nodejs/README.md +++ b/nodejs/README.md @@ -118,6 +118,7 @@ new CopilotClient(options?: CopilotClientOptions) - `mode?: "empty" | "copilot-cli"` - Defaulting strategy. Use `"empty"` for multi-user server mode; defaults to `"copilot-cli"`. - `workingDirectory?: string` - Working directory for the runtime process (default: current process cwd). - `baseDirectory?: string` - Base directory for Copilot data (session state, config, etc.). Sets `COPILOT_HOME` on the spawned runtime. When not set, the runtime defaults to `~/.copilot`. Ignored when connecting via `RuntimeConnection.forUri`. +- `extensionLaunchProvider?: ExtensionLaunchProvider` - Experimental connection-level resolver for extension launch profiles. The client installs the reverse-RPC handler and registers the provider during startup before sessions can be created. - `logLevel?: "none" | "error" | "warning" | "info" | "debug" | "all"` - Log level. When omitted, the runtime uses its own default (currently `"info"`). - `env?: Record` - Environment variables for the runtime process. When omitted, inherits `process.env`. - `gitHubToken?: string` - GitHub token for authentication. When provided, takes priority over other auth methods. diff --git a/nodejs/src/client.ts b/nodejs/src/client.ts index 8438033e18..4b610888b6 100644 --- a/nodejs/src/client.ts +++ b/nodejs/src/client.ts @@ -60,6 +60,7 @@ import type { CustomAgentConfig, ExitPlanModeRequest, ExitPlanModeResult, + ExtensionLaunchProvider, ExtensionJoinOptions, ForegroundSessionInfo, GetAuthStatusResponse, @@ -468,6 +469,7 @@ export class CopilotClient { /** Connection-level session filesystem config, set via constructor option. */ private sessionFsConfig: SessionFsConfig | null = null; private requestHandler: CopilotRequestHandler | null = null; + private extensionLaunchProvider?: ExtensionLaunchProvider; private builtinPluginDirectories: string[] = []; private onGitHubTelemetry?: (notification: GitHubTelemetryNotification) => void | Promise; private clientGlobalHandlers: import("./generated/rpc.js").ClientGlobalApiHandlers = {}; @@ -666,6 +668,7 @@ export class CopilotClient { this.onGetTraceContext = options.onGetTraceContext; this.sessionFsConfig = options.sessionFs ?? null; this.requestHandler = options.requestHandler ?? null; + this.extensionLaunchProvider = options.extensionLaunchProvider; this.onGitHubTelemetry = options.onGitHubTelemetry; this.setupClientGlobalHandlers(); @@ -811,6 +814,7 @@ export class CopilotClient { private setupClientGlobalHandlers(): void { const handlers: import("./generated/rpc.js").ClientGlobalApiHandlers = {}; + handlers.extensionLaunchProvider = this.extensionLaunchProvider; if (this.requestHandler) { handlers.llmInference = createCopilotRequestAdapter(this.requestHandler, () => { if (!this.connection) { @@ -947,6 +951,10 @@ export class CopilotClient { // Verify protocol version compatibility await this.verifyProtocolVersion(); + if (this.extensionLaunchProvider) { + await this.rpc.registerExtensionLaunchProvider(); + } + if (this.builtinPluginDirectories.length > 0) { try { await this.connection!.sendRequest("plugins.builtin.set", { diff --git a/nodejs/src/index.ts b/nodejs/src/index.ts index a9244e9e48..f1879e8f1b 100644 --- a/nodejs/src/index.ts +++ b/nodejs/src/index.ts @@ -91,6 +91,10 @@ export type { ExitPlanModeHandler, ExitPlanModeRequest, ExitPlanModeResult, + ExtensionLaunchProfile, + ExtensionLaunchProvider, + ExtensionLaunchProviderResolveRequest, + ExtensionLaunchProviderResolveResult, ExtensionInfo, ForegroundSessionInfo, GetAuthStatusResponse, @@ -115,6 +119,7 @@ export type { MCPServerConfig, DefaultAgentConfig, BearerTokenProvider, + ExtensionContextAttachment, MessageOptions, ResponseSchema, MessageSource, diff --git a/nodejs/src/types.ts b/nodejs/src/types.ts index de332d4a3a..6a19116a37 100644 --- a/nodejs/src/types.ts +++ b/nodejs/src/types.ts @@ -11,6 +11,7 @@ import type { Canvas } from "./canvas.js"; import type { SessionFsProvider } from "./sessionFsProvider.js"; import type { CopilotRequestHandler } from "./copilotRequestHandler.js"; import type { + AttachmentExtensionContext as GeneratedExtensionContextAttachment, AutoTier, PermissionRequest as GeneratedPermissionRequest, PermissionRequestedData as GeneratedPermissionRequestedData, @@ -22,6 +23,7 @@ import type { import type { CopilotSession } from "./session.js"; import type { FactoryJsonSchema, JsonValue } from "./factory.js"; import type { + ExtensionLaunchProviderHandler as GeneratedExtensionLaunchProvider, GitHubTokenAcquireRequest, GitHubTokenAcquireResult, GitHubTelemetryNotification, @@ -34,6 +36,9 @@ import type { ToolSet } from "./toolSet.js"; export type { RemoteSessionMode } from "./generated/rpc.js"; export type { CurrentToolMetadata } from "./generated/rpc.js"; export type { + ExtensionLaunchProfile, + ExtensionLaunchProviderResolveRequest, + ExtensionLaunchProviderResolveResult, GitHubTokenAcquireReason, GitHubTokenAcquireResult, GitHubTelemetryNotification, @@ -385,6 +390,15 @@ export interface CopilotClientOptions { */ builtinPluginDirectories?: readonly string[]; + /** + * Connection-level extension launch profile provider. + * When set, the client registers the provider during startup before any + * session can be created. + * + * @experimental + */ + extensionLaunchProvider?: ExtensionLaunchProvider; + /** * Log level for the Copilot runtime. When omitted, the runtime uses its * own default (currently `"info"`). @@ -531,6 +545,9 @@ export interface CopilotClientOptions { _internalConnection?: InternalRuntimeConnection; } +/** Resolves launch profiles for extension entrypoints discovered by the runtime. */ +export type ExtensionLaunchProvider = GeneratedExtensionLaunchProvider; + /** * Configuration for creating a session */ @@ -3336,6 +3353,9 @@ export interface ProviderModelConfig { */ export type MessageSource = "user" | "system" | `agent-${string}`; +/** Structured context contributed by an extension. */ +export type ExtensionContextAttachment = GeneratedExtensionContextAttachment; + export interface MessageOptions { /** * The prompt/message to send @@ -3350,7 +3370,7 @@ export interface MessageOptions { source?: MessageSource; /** - * File, directory, selection, or blob attachments + * File, directory, selection, blob, or extension context attachments */ attachments?: Array< | { @@ -3379,6 +3399,7 @@ export interface MessageOptions { mimeType: string; displayName?: string; } + | ExtensionContextAttachment >; /** diff --git a/nodejs/test/e2e/abort.e2e.test.ts b/nodejs/test/e2e/abort.e2e.test.ts index 89877387c3..c5b67cb1a9 100644 --- a/nodejs/test/e2e/abort.e2e.test.ts +++ b/nodejs/test/e2e/abort.e2e.test.ts @@ -123,6 +123,11 @@ describe("Abort", async () => { // Wait for the tool to start executing const toolValue = await withTimeout(toolStarted, 60_000, "slow_analysis start"); expect(toolValue).toBe("test_abort"); + expect((await session.rpc.metadata.isProcessing()).processing).toBe(true); + expect(await session.rpc.metadata.activity()).toMatchObject({ + hasActiveWork: true, + abortable: true, + }); // Abort while the tool is running await session.abort(); @@ -135,13 +140,21 @@ describe("Abort", async () => { const recoveryReceived = new Promise((resolve) => { recoveryResolve = resolve; }); + let recoveryIdleResolve!: (value: void) => void; + const recoveryIdle = new Promise((resolve) => { + recoveryIdleResolve = resolve; + }); + let recoveryMessageSeen = false; session.on((event) => { if ( event.type === "assistant.message" && event.data.content?.includes("tool_abort_recovery_ok") ) { + recoveryMessageSeen = true; recoveryResolve(); + } else if (event.type === "session.idle" && recoveryMessageSeen) { + recoveryIdleResolve(); } }); @@ -150,6 +163,9 @@ describe("Abort", async () => { }); await withTimeout(recoveryReceived, 60_000, "tool abort recovery message"); + await withTimeout(recoveryIdle, 60_000, "tool abort recovery idle"); + expect((await session.rpc.metadata.isProcessing()).processing).toBe(false); + expect((await session.rpc.metadata.activity()).hasActiveWork).toBe(false); await session.disconnect(); }); diff --git a/nodejs/test/e2e/canvas.e2e.test.ts b/nodejs/test/e2e/canvas.e2e.test.ts index 23d75b6581..cd157a07bd 100644 --- a/nodejs/test/e2e/canvas.e2e.test.ts +++ b/nodejs/test/e2e/canvas.e2e.test.ts @@ -3,7 +3,7 @@ *--------------------------------------------------------------------------------------------*/ import { describe, expect, it } from "vitest"; -import { approveAll, createCanvas } from "../../src/index.js"; +import { approveAll, CanvasError, createCanvas } from "../../src/index.js"; import { createSdkTestContext } from "./harness/sdkTestContext.js"; describe("Canvas RPC", async () => { @@ -178,4 +178,117 @@ describe("Canvas RPC", async () => { await session.disconnect(); }); + + it.each(["open", "action", "close"] as const)( + "preserves structured canvas errors from %s handlers", + async (operation) => { + const errorCanvas = createCanvas({ + id: `error-${operation}`, + displayName: `Error ${operation}`, + description: `Throws a structured error from ${operation}.`, + actions: [ + { + name: "fail", + handler: () => { + if (operation === "action") { + throw new CanvasError( + `scenario_canvas_${operation}`, + `scenario ${operation} failed` + ); + } + return null; + }, + }, + ], + open: () => { + if (operation === "open") { + throw new CanvasError( + `scenario_canvas_${operation}`, + `scenario ${operation} failed` + ); + } + return { url: "https://example.test/error-canvas" }; + }, + onClose: () => { + if (operation === "close") { + throw new CanvasError( + `scenario_canvas_${operation}`, + `scenario ${operation} failed` + ); + } + }, + }); + const session = await client.createSession({ + onPermissionRequest: approveAll, + canvases: [errorCanvas], + }); + + try { + const open = () => + session.rpc.canvas.open({ + canvasId: `error-${operation}`, + instanceId: `error-${operation}-1`, + }); + if (operation === "open") { + await expect(open()).rejects.toSatisfy((error: unknown) => { + expect(error).toMatchObject({ code: -32603 }); + expect(String(error)).toContain("scenario open failed"); + return true; + }); + } else { + await open(); + } + + const action = () => + session.rpc.canvas.action.invoke({ + instanceId: `error-${operation}-1`, + actionName: "fail", + }); + const close = () => + session.rpc.canvas.close({ instanceId: `error-${operation}-1` }); + if (operation === "action") { + await expect(action()).rejects.toSatisfy((error: unknown) => { + expect(error).toMatchObject({ code: -32603 }); + expect(String(error)).toContain("scenario action failed"); + return true; + }); + } else if (operation === "close") { + await expect(close()).resolves.toBeNull(); + } + + const callback = + operation === "open" + ? session.clientSessionApis.canvas!.open({ + sessionId: session.sessionId, + extensionId: "typescript-sdk-tests", + canvasId: `error-${operation}`, + instanceId: `error-${operation}-callback`, + }) + : operation === "action" + ? session.clientSessionApis.canvas!.invoke({ + sessionId: session.sessionId, + extensionId: "typescript-sdk-tests", + canvasId: `error-${operation}`, + instanceId: `error-${operation}-1`, + actionName: "fail", + }) + : session.clientSessionApis.canvas!.close({ + sessionId: session.sessionId, + extensionId: "typescript-sdk-tests", + canvasId: `error-${operation}`, + instanceId: `error-${operation}-1`, + }); + await expect(callback).rejects.toMatchObject({ + code: -32603, + message: `scenario ${operation} failed`, + data: { + code: `scenario_canvas_${operation}`, + message: `scenario ${operation} failed`, + }, + }); + } finally { + await session.disconnect(); + } + } + ); }); diff --git a/nodejs/test/e2e/client_options.e2e.test.ts b/nodejs/test/e2e/client_options.e2e.test.ts index 4d261bea52..bbcd35e174 100644 --- a/nodejs/test/e2e/client_options.e2e.test.ts +++ b/nodejs/test/e2e/client_options.e2e.test.ts @@ -14,6 +14,8 @@ const FAKE_STDIO_CLI_SCRIPT = `const fs = require("fs"); const captureIndex = process.argv.indexOf("--capture-file"); const captureFile = captureIndex >= 0 ? process.argv[captureIndex + 1] : undefined; const requests = []; +const clientResponses = []; +let extensionRegistrationId; function saveCapture() { if (!captureFile) { @@ -24,6 +26,7 @@ function saveCapture() { args: process.argv.slice(2), cwd: process.cwd(), requests, + clientResponses, env: { COPILOT_HOME: process.env.COPILOT_HOME, COPILOT_SDK_AUTH_TOKEN: process.env.COPILOT_SDK_AUTH_TOKEN, @@ -80,6 +83,16 @@ function handleMessage(message) { return; } + if (!message.method) { + clientResponses.push(message); + saveCapture(); + if (message.id === 9001 && extensionRegistrationId !== undefined) { + writeResponse(extensionRegistrationId, {}); + extensionRegistrationId = undefined; + } + return; + } + requests.push({ method: message.method, params: message.params }); saveCapture(); @@ -93,6 +106,17 @@ function handleMessage(message) { return; } + if (message.method === "registerExtensionLaunchProvider") { + extensionRegistrationId = message.id; + writeRequest(9001, "extensionLaunchProvider.resolve", { + id: "project:node-e2e", + name: "node-e2e", + modulePath: "/extensions/node-e2e.mjs", + source: "project" + }); + return; + } + if (message.method === "session.create" || message.method === "session.resume") { const sessionId = message.params?.sessionId ?? message.params?.[0]?.sessionId ?? "fake-session"; writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); @@ -122,6 +146,11 @@ function writeResponse(id, result) { const body = JSON.stringify({ jsonrpc: "2.0", id, result }); process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); } + +function writeRequest(id, method, params) { + const body = JSON.stringify({ jsonrpc: "2.0", id, method, params }); + process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); +} `; async function getAvailableTcpPort(): Promise { @@ -375,6 +404,76 @@ describe("Client options", async () => { await resumed.disconnect(); }); + it("should register and invoke an extension launch provider during startup", async () => { + const cliPath = path.join(workDir, `fake-cli-extension-provider-${Date.now()}.js`); + const capturePath = path.join(workDir, `fake-cli-extension-provider-${Date.now()}.json`); + fs.writeFileSync(cliPath, FAKE_STDIO_CLI_SCRIPT); + + let observedRequest: + | { + id: string; + name: string; + modulePath: string; + source: string; + } + | undefined; + const client = new CopilotClient({ + workingDirectory: workDir, + connection: RuntimeConnection.forStdio({ + path: cliPath, + args: ["--capture-file", capturePath], + }), + useLoggedInUser: false, + extensionLaunchProvider: { + resolve: async (request) => { + observedRequest = request; + return { + launch: { + executable: "node", + args: ["extension-host"], + env: { EXTENSION_SOURCE: "node" }, + }, + }; + }, + }, + }); + onTestFinished(async () => { + await client.forceStop(); + }); + + await client.start(); + + expect(observedRequest).toEqual({ + id: "project:node-e2e", + name: "node-e2e", + modulePath: "/extensions/node-e2e.mjs", + source: "project", + }); + const capture = JSON.parse(fs.readFileSync(capturePath, "utf8")) as { + requests: { method: string }[]; + clientResponses: { + id: number; + result: { + launch: { executable: string; args: string[]; env: Record }; + }; + }[]; + }; + expect(capture.requests.map((request) => request.method)).toContain( + "registerExtensionLaunchProvider" + ); + expect(capture.clientResponses).toContainEqual({ + jsonrpc: "2.0", + id: 9001, + result: { + launch: { + executable: "node", + args: ["extension-host"], + env: { EXTENSION_SOURCE: "node" }, + }, + }, + }); + }); + it("should send empty-mode custom agent locality defaults in initial requests", async () => { const cliPath = path.join( workDir, @@ -682,7 +781,7 @@ describe("Client options", async () => { const session = await client.resumeSession("advanced-resume-session", { clientName: "advanced-resume-client", - model: "claude-haiku-4.5", + model: "legacy-resume-model", reasoningEffort: "low", reasoningSummary: "none", contextTier: "default", @@ -697,6 +796,22 @@ describe("Client options", async () => { }, memory: { enabled: false }, remoteSession: "on", + providers: [ + { + name: "resume-provider", + type: "openai", + wireApi: "responses", + baseUrl: "https://resume-provider.example.test/v1", + bearerTokenProvider: async () => "resume-provider-token", + }, + ], + models: [ + { + provider: "resume-provider", + id: "legacy-resume-model", + wireModel: "resume-wire-model", + }, + ], openCanvases: [ { canvasId: "resume-canvas", @@ -715,7 +830,7 @@ describe("Client options", async () => { const resumeRequest = getCapturedRequest(capturePath, "session.resume"); expect(resumeRequest.sessionId).toBe("advanced-resume-session"); expect(resumeRequest.clientName).toBe("advanced-resume-client"); - expect(resumeRequest.model).toBe("claude-haiku-4.5"); + expect(resumeRequest.model).toBe("legacy-resume-model"); expect(resumeRequest.reasoningEffort).toBe("low"); expect(resumeRequest.reasoningSummary).toBe("none"); expect(resumeRequest.contextTier).toBe("default"); @@ -728,6 +843,20 @@ describe("Client options", async () => { expect(getObject(resumeRequest.largeOutput).outputDir).toBe(outputDirectory); expect(getObject(resumeRequest.memory).enabled).toBe(false); expect(resumeRequest.remoteSession).toBe("on"); + const provider = getObject(getArray(resumeRequest.providers)[0]); + expect(provider).toMatchObject({ + name: "resume-provider", + type: "openai", + wireApi: "responses", + baseUrl: "https://resume-provider.example.test/v1", + hasBearerTokenProvider: true, + }); + expect(provider).not.toHaveProperty("bearerTokenProvider"); + expect(getObject(getArray(resumeRequest.models)[0])).toMatchObject({ + provider: "resume-provider", + id: "legacy-resume-model", + wireModel: "resume-wire-model", + }); const openCanvas = getObject(getArray(resumeRequest.openCanvases)[0]); expect(openCanvas.canvasId).toBe("resume-canvas"); diff --git a/nodejs/test/e2e/permissions.e2e.test.ts b/nodejs/test/e2e/permissions.e2e.test.ts index 638ea12a33..6de6c4a516 100644 --- a/nodejs/test/e2e/permissions.e2e.test.ts +++ b/nodejs/test/e2e/permissions.e2e.test.ts @@ -278,15 +278,13 @@ describe("Permission callbacks", async () => { }); it("should receive toolCallId in permission requests", async () => { - let receivedToolCallId = false; + let receivedRequest: PermissionRequest | undefined; + let receivedInvocation: { sessionId: string; managedSettingsEnabled?: boolean } | undefined; const session = await client.createSession({ - onPermissionRequest: (request) => { - if (request.toolCallId) { - receivedToolCallId = true; - expect(typeof request.toolCallId).toBe("string"); - expect(request.toolCallId.length).toBeGreaterThan(0); - } + onPermissionRequest: (request, invocation) => { + receivedRequest = request; + receivedInvocation = invocation; return { kind: "approve-once" }; }, }); @@ -295,7 +293,33 @@ describe("Permission callbacks", async () => { prompt: "Run 'echo test'", }); - expect(receivedToolCallId).toBe(true); + expect(receivedInvocation).toEqual({ + sessionId: session.sessionId, + managedSettingsEnabled: false, + }); + expect(receivedRequest).toMatchObject({ + kind: "shell", + canOfferSessionApproval: expect.any(Boolean), + commands: expect.arrayContaining([ + { + identifier: expect.any(String), + readOnly: expect.any(Boolean), + }, + ]), + fullCommandText: expect.stringContaining("echo test"), + hasWriteFileRedirection: false, + intention: expect.any(String), + possiblePaths: expect.any(Array), + possibleUrls: expect.any(Array), + toolCallId: expect.any(String), + }); + if (receivedRequest?.kind === "shell") { + expect(receivedRequest.toolCallId?.length).toBeGreaterThan(0); + expect(receivedRequest.intention.trim().length).toBeGreaterThan(0); + expect(receivedRequest.commands.every((command) => command.identifier.trim())).toBe( + true + ); + } await session.disconnect(); }); diff --git a/nodejs/test/e2e/rpc_surface_coverage.e2e.test.ts b/nodejs/test/e2e/rpc_surface_coverage.e2e.test.ts new file mode 100644 index 0000000000..68382dad73 --- /dev/null +++ b/nodejs/test/e2e/rpc_surface_coverage.e2e.test.ts @@ -0,0 +1,412 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { mkdtempSync, readFileSync, writeFileSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { fileURLToPath } from "node:url"; +import { rimraf } from "rimraf"; +import ts from "typescript"; +import { describe, expect, it, onTestFinished } from "vitest"; +import { approveAll, CopilotClient, RuntimeConnection } from "../../src/index.js"; + +const FAKE_RPC_CLI = `const fs = require("fs"); + +const captureIndex = process.argv.indexOf("--capture-file"); +const captureFile = process.argv[captureIndex + 1]; +const requests = []; +let mode = "interactive"; +let buffer = Buffer.alloc(0); + +function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify(requests)); +} + +function writeResponse(id, result) { + const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); +} + +function writeError(id, code, message, data) { + const body = JSON.stringify({ jsonrpc: "2.0", id, error: { code, message, data } }); + process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); +} + +function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) { + return; + } + + requests.push({ method: message.method, params: message.params ?? null }); + saveCapture(); + + if (message.method === "connect") { + writeResponse(message.id, { ok: true, protocolVersion: 3, version: "fake-rpc-surface" }); + return; + } + + if (message.method === "session.create" || message.method === "session.resume") { + writeResponse(message.id, { + sessionId: message.params?.sessionId ?? "rpc-surface-session", + workspacePath: null, + capabilities: { supportsStreaming: true } + }); + return; + } + + if (message.params?.__forceError === true) { + writeError(message.id, -32042, "deterministic fake failure", { + method: message.method, + nested: { retryable: false } + }); + return; + } + + if (message.method === "ping") { + writeResponse(message.id, { + message: \`pong: \${message.params?.message ?? ""}\`, + timestamp: "2026-09-18T15:00:00.000Z", + protocolVersion: 3 + }); + return; + } + + if (message.method === "models.list") { + writeResponse(message.id, { + models: [{ + id: "fake/model", + name: "Fake Model", + capabilities: { + supports: { vision: true, reasoningEffort: true }, + limits: { + maxContextWindowTokens: 128000, + maxPromptTokens: 120000, + maxOutputTokens: 8000 + } + }, + billing: { multiplier: 1 } + }] + }); + return; + } + + if (message.method === "session.mode.set") { + mode = message.params.mode; + writeResponse(message.id, null); + return; + } + + if (message.method === "session.mode.get") { + writeResponse(message.id, mode); + return; + } + + writeResponse(message.id, { + method: message.method, + params: message.params ?? null, + state: { + phase: "covered", + nested: { + items: [ + { kind: "success", value: 42 }, + { kind: "empty", value: null } + ] + } + } + }); +} + +function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\\r\\n\\r\\n"); + if (headerEnd < 0) { + return; + } + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\\s*(\\d+)/i.exec(header); + if (!match) { + throw new Error("Missing Content-Length header"); + } + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) { + return; + } + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } +} + +process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); +}); +process.stdin.resume(); +`; + +type RpcMethod = { + scope: "server" | "session"; + path: string; + wireMethod: string; + parameterCount: number; + returnType: string; +}; + +type RpcFunction = ((params?: Record) => Promise) & { + length: number; +}; + +function getPropertyName(node: ts.PropertyName): string { + if (ts.isIdentifier(node) || ts.isStringLiteral(node) || ts.isNumericLiteral(node)) { + return node.text; + } + throw new Error(`Unsupported generated RPC property name: ${node.getText()}`); +} + +function findSendRequestMethod(initializer: ts.ArrowFunction): string { + let wireMethod: string | undefined; + const visit = (node: ts.Node): void => { + if ( + ts.isCallExpression(node) && + ts.isPropertyAccessExpression(node.expression) && + node.expression.name.text === "sendRequest" && + node.arguments.length > 0 && + ts.isStringLiteral(node.arguments[0]) + ) { + wireMethod = node.arguments[0].text; + } + ts.forEachChild(node, visit); + }; + visit(initializer.body); + if (!wireMethod) { + throw new Error( + `Generated RPC method does not call connection.sendRequest: ${initializer.getText()}` + ); + } + return wireMethod; +} + +function collectRpcMethods( + object: ts.ObjectLiteralExpression, + scope: RpcMethod["scope"], + prefix: string[] = [] +): RpcMethod[] { + const methods: RpcMethod[] = []; + for (const property of object.properties) { + if (!ts.isPropertyAssignment(property)) { + continue; + } + const path = [...prefix, getPropertyName(property.name)]; + if (ts.isObjectLiteralExpression(property.initializer)) { + methods.push(...collectRpcMethods(property.initializer, scope, path)); + } else if (ts.isArrowFunction(property.initializer)) { + methods.push({ + scope, + path: path.join("."), + wireMethod: findSendRequestMethod(property.initializer), + parameterCount: property.initializer.parameters.length, + returnType: property.initializer.type?.getText() ?? "unknown", + }); + } + } + return methods; +} + +function getGeneratedRpcInventory(): RpcMethod[] { + const path = fileURLToPath(new URL("../../src/generated/rpc.ts", import.meta.url)); + const source = ts.createSourceFile( + path, + readFileSync(path, "utf8"), + ts.ScriptTarget.Latest, + true, + ts.ScriptKind.TS + ); + const methods: RpcMethod[] = []; + + for (const statement of source.statements) { + if ( + !ts.isFunctionDeclaration(statement) || + !statement.name || + (statement.name.text !== "createServerRpc" && + statement.name.text !== "createSessionRpc") || + !statement.body + ) { + continue; + } + const returnStatement = statement.body.statements.find(ts.isReturnStatement); + if ( + !returnStatement?.expression || + !ts.isObjectLiteralExpression(returnStatement.expression) + ) { + throw new Error(`${statement.name.text} does not return an object literal`); + } + methods.push( + ...collectRpcMethods( + returnStatement.expression, + statement.name.text === "createServerRpc" ? "server" : "session" + ) + ); + } + + return methods.sort((left, right) => + `${left.scope}.${left.path}`.localeCompare(`${right.scope}.${right.path}`) + ); +} + +function collectRuntimeFunctions( + value: object, + scope: RpcMethod["scope"], + prefix: string[] = [] +): Map { + const functions = new Map(); + for (const [name, member] of Object.entries(value)) { + const path = [...prefix, name]; + if (typeof member === "function") { + functions.set(`${scope}.${path.join(".")}`, member as RpcFunction); + } else if (member && typeof member === "object") { + for (const [nestedPath, nestedFunction] of collectRuntimeFunctions( + member, + scope, + path + )) { + functions.set(nestedPath, nestedFunction); + } + } + } + return functions; +} + +describe("Generated RPC surface coverage", () => { + it("serializes and projects every public generated RPC method", async () => { + const directory = mkdtempSync(join(tmpdir(), "copilot-node-rpc-surface-")); + const cliPath = join(directory, "fake-rpc-cli.js"); + const capturePath = join(directory, "capture.json"); + writeFileSync(cliPath, FAKE_RPC_CLI); + + const client = new CopilotClient({ + workingDirectory: directory, + baseDirectory: directory, + env: { ...process.env, COPILOT_HOME: directory }, + connection: RuntimeConnection.forStdio({ + path: cliPath, + args: ["--capture-file", capturePath], + }), + useLoggedInUser: false, + }); + onTestFinished(async () => { + await client.forceStop(); + await rimraf(directory, { maxRetries: 10, retryDelay: 100 }); + }); + + await client.start(); + const session = await client.createSession({ + sessionId: "rpc-surface-session", + onPermissionRequest: approveAll, + }); + + const inventory = getGeneratedRpcInventory(); + const runtimeFunctions = new Map([ + ...collectRuntimeFunctions(client.rpc, "server"), + ...collectRuntimeFunctions(session.rpc, "session"), + ]); + + expect(inventory).toHaveLength(340); + expect([...runtimeFunctions.keys()].sort()).toEqual( + inventory.map((method) => `${method.scope}.${method.path}`) + ); + + const invokedWireMethods = new Set(); + for (const method of inventory) { + const fullPath = `${method.scope}.${method.path}`; + const fn = runtimeFunctions.get(fullPath); + expect(fn, `Missing runtime RPC function ${fullPath}`).toBeDefined(); + expect(fn!.length, `Signature arity changed for ${fullPath}`).toBe( + method.parameterCount + ); + expect(method.returnType, `Missing generated return type for ${fullPath}`).not.toBe( + "unknown" + ); + + const marker = { __coveragePath: fullPath }; + const result = method.parameterCount === 0 ? await fn!() : await fn!(marker); + invokedWireMethods.add(method.wireMethod); + + if ( + method.wireMethod === "ping" || + method.wireMethod === "models.list" || + method.wireMethod === "session.mode.get" || + method.wireMethod === "session.mode.set" + ) { + continue; + } + + expect(result).toMatchObject({ + method: method.wireMethod, + state: { + phase: "covered", + nested: { + items: [ + { kind: "success", value: 42 }, + { kind: "empty", value: null }, + ], + }, + }, + }); + const params = (result as { params: Record | null }).params; + if (method.scope === "session") { + expect(params).toMatchObject({ sessionId: "rpc-surface-session" }); + } + if (method.parameterCount === 1) { + expect(params).toMatchObject(marker); + } + } + + const ping = await client.rpc.ping({ message: "typed projection" }); + expect(ping).toEqual({ + message: "pong: typed projection", + timestamp: "2026-09-18T15:00:00.000Z", + protocolVersion: 3, + }); + + const models = await client.rpc.models.list({}); + expect(models.models[0]).toMatchObject({ + id: "fake/model", + capabilities: { + supports: { vision: true, reasoningEffort: true }, + limits: { maxContextWindowTokens: 128000 }, + }, + billing: { multiplier: 1 }, + }); + + await session.rpc.mode.set({ mode: "plan" }); + expect(await session.rpc.mode.get()).toBe("plan"); + await session.rpc.mode.set({ mode: "interactive" }); + expect(await session.rpc.mode.get()).toBe("interactive"); + + await expect( + client.rpc.ping({ message: "error", __forceError: true } as never) + ).rejects.toMatchObject({ + code: -32042, + message: "deterministic fake failure", + data: { + method: "ping", + nested: { retryable: false }, + }, + }); + + const captured = JSON.parse(readFileSync(capturePath, "utf8")) as Array<{ + method: string; + params: Record | null; + }>; + expect(invokedWireMethods).toEqual(new Set(inventory.map((method) => method.wireMethod))); + const capturedWireMethods = new Set(captured.map((request) => request.method)); + for (const wireMethod of invokedWireMethods) { + expect(capturedWireMethods.has(wireMethod), `Missing captured ${wireMethod}`).toBe( + true + ); + } + }); +}); diff --git a/nodejs/test/e2e/rpc_workspace_checkpoints.e2e.test.ts b/nodejs/test/e2e/rpc_workspace_checkpoints.e2e.test.ts index 78a820f67a..f1d8050925 100644 --- a/nodejs/test/e2e/rpc_workspace_checkpoints.e2e.test.ts +++ b/nodejs/test/e2e/rpc_workspace_checkpoints.e2e.test.ts @@ -32,22 +32,27 @@ describe("Session workspace checkpoint RPC", async () => { } }); - it("should return typed workspace diff result", async () => { - const session = await client.createSession({ onPermissionRequest: approveAll }); - try { - const result = await session.rpc.workspaces.diff({ mode: "unstaged" }); - expect(result.requestedMode).toBe("unstaged"); - expect(["unstaged", "branch"]).toContain(result.mode); - expect(Array.isArray(result.changes)).toBe(true); - for (const change of result.changes) { - expect(change.path.trim()).toBeTruthy(); - expect(["added", "modified", "deleted", "renamed"]).toContain(change.changeType); - expect(typeof change.diff).toBe("string"); + it.each(["session", "unstaged", "branch"] as const)( + "should return typed workspace diff result for %s mode", + async (mode) => { + const session = await client.createSession({ onPermissionRequest: approveAll }); + try { + const result = await session.rpc.workspaces.diff({ mode }); + expect(result.requestedMode).toBe(mode); + expect(["session", "unstaged", "branch"]).toContain(result.mode); + expect(Array.isArray(result.changes)).toBe(true); + for (const change of result.changes) { + expect(change.path.trim()).toBeTruthy(); + expect(["added", "modified", "deleted", "renamed"]).toContain( + change.changeType + ); + expect(typeof change.diff).toBe("string"); + } + } finally { + await session.disconnect(); } - } finally { - await session.disconnect(); } - }); + ); it("should save large paste and expose readable content", async () => { const session = await client.createSession({ onPermissionRequest: approveAll }); diff --git a/nodejs/test/e2e/scenario_testing_composition.e2e.test.ts b/nodejs/test/e2e/scenario_testing_composition.e2e.test.ts new file mode 100644 index 0000000000..c8be623b7d --- /dev/null +++ b/nodejs/test/e2e/scenario_testing_composition.e2e.test.ts @@ -0,0 +1,42 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { describe, expect, it } from "vitest"; +import type { SessionEvent } from "../../src/index.js"; +import { approveAll } from "../../src/index.js"; +import { createSdkTestContext } from "./harness/sdkTestContext.js"; + +describe("Scenario testing composition", async () => { + const { copilotClient: client } = await createSdkTestContext(); + + it("should not emit redundant model change when resuming same model", async () => { + const original = await client.createSession({ + model: "claude-sonnet-5", + onPermissionRequest: approveAll, + }); + const sessionId = original.sessionId; + + try { + const response = await original.sendAndWait({ + prompt: "Reply with exactly SCENARIO_SAME_MODEL_HISTORY_READY.", + }); + expect(response?.data.content).toBe("SCENARIO_SAME_MODEL_HISTORY_READY"); + } finally { + await original.disconnect(); + } + + const resumeEvents: SessionEvent[] = []; + const resumed = await client.resumeSession(sessionId, { + model: "claude-sonnet-5", + onPermissionRequest: approveAll, + onEvent: (event) => resumeEvents.push(event), + }); + try { + expect(resumeEvents.some((event) => event.type === "session.model_change")).toBe(false); + expect((await resumed.rpc.model.getCurrent()).modelId).toBe("claude-sonnet-5"); + } finally { + await resumed.disconnect(); + } + }); +}); diff --git a/nodejs/test/e2e/scenario_testing_persistence.e2e.test.ts b/nodejs/test/e2e/scenario_testing_persistence.e2e.test.ts new file mode 100644 index 0000000000..6d776ac57c --- /dev/null +++ b/nodejs/test/e2e/scenario_testing_persistence.e2e.test.ts @@ -0,0 +1,36 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { describe, expect, it } from "vitest"; +import { approveAll } from "../../src/index.js"; +import { createSdkTestContext } from "./harness/sdkTestContext.js"; + +describe("Scenario testing persistence", async () => { + const { copilotClient: client } = await createSdkTestContext(); + + it("should retry from existing history with empty sendmessages", async () => { + const session = await client.createSession({ + model: "claude-sonnet-5", + onPermissionRequest: approveAll, + }); + + try { + const initial = await session.sendAndWait({ + prompt: "Reply with exactly EMPTY_BATCH_CONTEXT_READY.", + }); + expect(initial?.data.content).toBe("EMPTY_BATCH_CONTEXT_READY"); + + const result = await session.rpc.sendMessages({ messages: [], wait: true }); + expect(result.messageIds).toEqual([]); + + const events = await session.getEvents(); + const finalAssistantMessage = [...events] + .reverse() + .find((event) => event.type === "assistant.message"); + expect(finalAssistantMessage?.data.content).toBe("EMPTY_BATCH_RETRY_DONE"); + } finally { + await session.disconnect(); + } + }); +}); diff --git a/nodejs/test/e2e/scenario_testing_recovery.e2e.test.ts b/nodejs/test/e2e/scenario_testing_recovery.e2e.test.ts new file mode 100644 index 0000000000..99f92f0905 --- /dev/null +++ b/nodejs/test/e2e/scenario_testing_recovery.e2e.test.ts @@ -0,0 +1,262 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { mkdtempSync, readFileSync, writeFileSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { rimraf } from "rimraf"; +import { describe, expect, it } from "vitest"; +import { approveAll, CopilotClient, RuntimeConnection } from "../../src/index.js"; + +const FAKE_RECOVERY_CLI = `const fs = require("fs"); + +const captureIndex = process.argv.indexOf("--capture-file"); +const captureFile = process.argv[captureIndex + 1]; +const modeIndex = process.argv.indexOf("--scenario"); +const scenario = process.argv[modeIndex + 1]; +const requests = []; +let resumeAttempts = 0; +let buffer = Buffer.alloc(0); + +function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify(requests)); +} + +function writeResponse(id, result) { + const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); +} + +function writeError(id, code, message, data) { + const body = JSON.stringify({ jsonrpc: "2.0", id, error: { code, message, data } }); + process.stdout.write(\`Content-Length: \${Buffer.byteLength(body, "utf8")}\\r\\n\\r\\n\${body}\`); +} + +function handle(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) { + return; + } + requests.push({ method: message.method, params: message.params ?? null }); + saveCapture(); + + if (message.method === "connect") { + writeResponse(message.id, { ok: true, protocolVersion: 3, version: "fake-recovery" }); + return; + } + + if (message.method === "session.create") { + writeResponse(message.id, { + sessionId: message.params.sessionId, + workspacePath: null, + capabilities: null + }); + return; + } + + if (message.method === "session.resume") { + resumeAttempts += 1; + if (scenario === "resume-retry" && resumeAttempts === 1) { + writeError(message.id, -32001, "Session not found before acceptance", { + kind: "not_found", + phase: "preacceptance", + sessionId: message.params.sessionId + }); + return; + } + writeResponse(message.id, { + sessionId: message.params.sessionId, + workspacePath: null, + capabilities: null + }); + return; + } + + if (message.method === "session.delete" && scenario === "delete-not-found") { + writeError(message.id, -32001, "Session not found during cleanup", { + kind: "not_found", + operation: "delete", + sessionId: message.params.sessionId + }); + return; + } + + if (message.method === "session.send" && scenario === "ambiguous-send") { + writeError(message.id, -32098, "Transport lost after request acceptance", { + kind: "ambiguous_transport_loss", + accepted: true + }); + return; + } + + if (message.method === "session.detach") { + writeResponse(message.id, { success: true }); + return; + } + + writeResponse(message.id, null); +} + +function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\\r\\n\\r\\n"); + if (headerEnd < 0) { + return; + } + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\\s*(\\d+)/i.exec(header); + if (!match) { + throw new Error("Missing Content-Length header"); + } + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) { + return; + } + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handle(JSON.parse(body)); + } +} + +process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); +}); +process.stdin.resume(); +`; + +type FakeClient = { + client: CopilotClient; + capturePath: string; + directory: string; +}; + +function createFakeClient(scenario: string): FakeClient { + const directory = mkdtempSync(join(tmpdir(), "copilot-node-scenario-recovery-")); + const cliPath = join(directory, "fake-recovery-cli.js"); + const capturePath = join(directory, "capture.json"); + writeFileSync(cliPath, FAKE_RECOVERY_CLI); + return { + client: new CopilotClient({ + workingDirectory: directory, + baseDirectory: directory, + env: { ...process.env, COPILOT_HOME: directory }, + connection: RuntimeConnection.forStdio({ + path: cliPath, + args: ["--capture-file", capturePath, "--scenario", scenario], + }), + useLoggedInUser: false, + }), + capturePath, + directory, + }; +} + +async function disposeFakeClient(fake: FakeClient): Promise { + await fake.client.forceStop(); + await rimraf(fake.directory, { maxRetries: 10, retryDelay: 100 }); +} + +function capturedRequests(fake: FakeClient): Array<{ + method: string; + params: Record | null; +}> { + return JSON.parse(readFileSync(fake.capturePath, "utf8")) as Array<{ + method: string; + params: Record | null; + }>; +} + +describe("Scenario testing lifecycle recovery", () => { + it("should allow caller retry after preacceptance session not found", async () => { + const fake = createFakeClient("resume-retry"); + try { + await fake.client.start(); + await expect( + fake.client.resumeSession("retry-session", { + onPermissionRequest: approveAll, + }) + ).rejects.toMatchObject({ + code: -32001, + data: { + kind: "not_found", + phase: "preacceptance", + sessionId: "retry-session", + }, + }); + + const resumed = await fake.client.resumeSession("retry-session", { + onPermissionRequest: approveAll, + }); + expect(resumed.sessionId).toBe("retry-session"); + + const resumeRequests = capturedRequests(fake).filter( + (request) => request.method === "session.resume" + ); + expect(resumeRequests).toHaveLength(2); + await resumed.disconnect(); + } finally { + await disposeFakeClient(fake); + } + }); + + it("should classify delete not found for scenario cleanup", async () => { + const fake = createFakeClient("delete-not-found"); + try { + await fake.client.start(); + await expect( + fake.client.deleteSession("missing-cleanup-session") + ).rejects.toMatchObject({ + code: -32001, + data: { + kind: "not_found", + operation: "delete", + sessionId: "missing-cleanup-session", + }, + }); + expect( + capturedRequests(fake).filter((request) => request.method === "session.delete") + ).toHaveLength(1); + } finally { + await disposeFakeClient(fake); + } + }); + + it.each([undefined, "enqueue", "immediate"] as const)( + "should not replay scenario send after ambiguous transport loss (%s)", + async (mode) => { + const fake = createFakeClient("ambiguous-send"); + try { + await fake.client.start(); + const session = await fake.client.createSession({ + onPermissionRequest: approveAll, + }); + await expect( + session.send({ + prompt: "AMBIGUOUS_SEND_MUST_NOT_REPLAY", + ...(mode === undefined ? {} : { mode }), + }) + ).rejects.toMatchObject({ + code: -32098, + data: { + kind: "ambiguous_transport_loss", + accepted: true, + }, + }); + + const sends = capturedRequests(fake).filter( + (request) => request.method === "session.send" + ); + expect(sends).toHaveLength(1); + expect(sends[0].params).toMatchObject({ + prompt: "AMBIGUOUS_SEND_MUST_NOT_REPLAY", + ...(mode === undefined ? {} : { mode }), + }); + } finally { + await disposeFakeClient(fake); + } + } + ); +}); diff --git a/nodejs/test/message-source.test.ts b/nodejs/test/message-source.test.ts index 9c2460aa2a..394e42163c 100644 --- a/nodejs/test/message-source.test.ts +++ b/nodejs/test/message-source.test.ts @@ -10,7 +10,12 @@ import { StreamMessageReader, StreamMessageWriter, } from "vscode-jsonrpc/node.js"; -import type { MessageOptions, MessageSource, SessionEvent } from "../src/index.js"; +import type { + ExtensionContextAttachment, + MessageOptions, + MessageSource, + SessionEvent, +} from "../src/index.js"; import { CopilotSession } from "../src/session.js"; function sessionPair(traceContextProvider?: ConstructorParameters[3]) { @@ -72,12 +77,21 @@ describe.each(sources)("message source %s", (source) => { tracestate: "vendor=source", }; const { session, server } = sessionPair(() => trace); + const extensionContext = { + type: "extension_context", + capturedAt: "2026-09-18T20:00:00Z", + extensionId: "scenario-extension", + title: "Selected change", + canvasId: "diff", + instanceId: "diff-17", + payload: { selection: "active" }, + } satisfies ExtensionContextAttachment; const options: MessageOptions = { prompt: "context updated", source, mode: "immediate", agentMode: "plan", - attachments: [{ type: "blob", data: "aGk=", mimeType: "text/plain" }], + attachments: [{ type: "blob", data: "aGk=", mimeType: "text/plain" }, extensionContext], displayPrompt: "Context updated", requestHeaders: { "X-Tag": "context" }, }; diff --git a/nodejs/test/session-event-codegen.test.ts b/nodejs/test/session-event-codegen.test.ts index ec3c1d2169..60906720de 100644 --- a/nodejs/test/session-event-codegen.test.ts +++ b/nodejs/test/session-event-codegen.test.ts @@ -169,6 +169,28 @@ class ProbeResult: expect(processed).not.toContain("class ExternalRefMCPOauthHTTPResponse"); }); + it("uses external discriminated union loaders for deserialization", () => { + const code = `@dataclass +class PendingRequest: + request: PermissionPromptRequest + + @staticmethod + def from_dict(obj: Any) -> 'PendingRequest': + request = PermissionPromptRequest.from_dict(obj.get("request")) + return PendingRequest(request) +`; + + const processed = postProcessExternalRefsForPython( + code, + new Map([["__ExternalRef_PermissionPromptRequest", "PermissionPromptRequest"]]), + new Set(), + new Set(["PermissionPromptRequest"]) + ); + + expect(processed).toContain('request = _load_PermissionPromptRequest(obj.get("request"))'); + expect(processed).not.toContain("PermissionPromptRequest.from_dict"); + }); + it("maps special schema formats to the expected Python types", () => { const schema: JSONSchema7 = { definitions: { diff --git a/python/README.md b/python/README.md index 8debd30676..6802682ee3 100644 --- a/python/README.md +++ b/python/README.md @@ -238,6 +238,7 @@ All options are kw-only parameters: - `env` (dict | None): Environment variables for the CLI process. - `github_token` (str | None): GitHub token for authentication. When provided, takes priority over other auth methods. - `base_directory` (str | None): Base directory for Copilot data (session state, config, etc.). Sets `COPILOT_HOME` on the spawned CLI process. When `None`, the CLI defaults to `~/.copilot`. Useful in restricted environments where only specific directories are writable. Ignored when using a `UriRuntimeConnection`. +- `extension_launch_provider` (ExtensionLaunchProviderHandler | None): Experimental connection-level resolver for extension launch profiles. The client installs the reverse-RPC handler and registers the provider during startup before sessions can be created. - `use_logged_in_user` (bool | None): Whether to use logged-in user for authentication (default: True, but False when `github_token` is provided). - `telemetry` (dict | None): OpenTelemetry configuration for the CLI process. Providing this enables telemetry — no separate flag needed. See [Telemetry](#telemetry) below. - `session_fs` (dict | None): Connection-level session filesystem provider configuration. diff --git a/python/copilot/__init__.py b/python/copilot/__init__.py index e2326e1135..ce0baaab39 100644 --- a/python/copilot/__init__.py +++ b/python/copilot/__init__.py @@ -92,6 +92,10 @@ from .generated.rpc import ( CurrentModel, CurrentToolMetadata, + ExtensionLaunchProfile, + ExtensionLaunchProviderHandler, + ExtensionLaunchProviderResolveRequest, + ExtensionLaunchProviderResolveResult, GitHubTelemetryClientInfo, GitHubTelemetryEvent, GitHubTelemetryNotification, @@ -285,6 +289,10 @@ "ExitPlanModeHandler", "ExitPlanModeRequest", "ExitPlanModeResult", + "ExtensionLaunchProfile", + "ExtensionLaunchProviderHandler", + "ExtensionLaunchProviderResolveRequest", + "ExtensionLaunchProviderResolveResult", "ExtensionInfo", "CopilotWebSocketForwarder", "DisableBypassPermissionsModes", diff --git a/python/copilot/_jsonrpc.py b/python/copilot/_jsonrpc.py index 11baf0bd10..c992c0d91e 100644 --- a/python/copilot/_jsonrpc.py +++ b/python/copilot/_jsonrpc.py @@ -257,13 +257,18 @@ async def _send_message(self, message: dict): loop = self._loop or asyncio.get_event_loop() def write(): + if hasattr(self.process, "poll") and self.process.poll() is not None: + raise ProcessExitedError(self._get_process_exit_error()) content = json.dumps(message, separators=(",", ":")) content_bytes = content.encode("utf-8") header = f"Content-Length: {len(content_bytes)}\r\n\r\n" with self._write_lock: - self.process.stdin.write(header.encode("utf-8")) - self.process.stdin.write(content_bytes) - self.process.stdin.flush() + try: + self.process.stdin.write(header.encode("utf-8")) + self.process.stdin.write(content_bytes) + self.process.stdin.flush() + except (BrokenPipeError, OSError, ValueError) as exc: + raise ProcessExitedError(self._get_process_exit_error()) from exc # Run in thread pool to avoid blocking await loop.run_in_executor(None, write) diff --git a/python/copilot/client.py b/python/copilot/client.py index 4d9d4dfb76..e8110d656a 100644 --- a/python/copilot/client.py +++ b/python/copilot/client.py @@ -70,6 +70,7 @@ from .generated.rpc import ( ClientGlobalApiHandlers, ClientSessionApiHandlers, + ExtensionLaunchProviderHandler, GitHubTelemetryNotification, GitHubTokenAcquireReason, GitHubTokenAcquireRequest, @@ -812,6 +813,7 @@ class _CopilotClientOptions: github_token: str | None = None base_directory: str | None = None builtin_plugin_directories: tuple[str, ...] = () + extension_launch_provider: ExtensionLaunchProviderHandler | None = None use_logged_in_user: bool | None = None telemetry: TelemetryConfig | None = None session_fs: SessionFsConfig | None = None @@ -1567,6 +1569,7 @@ def __init__( github_token: str | None = None, base_directory: str | None = None, builtin_plugin_directories: Sequence[str] | None = None, + extension_launch_provider: ExtensionLaunchProviderHandler | None = None, use_logged_in_user: bool | None = None, telemetry: TelemetryConfig | None = None, session_fs: SessionFsConfig | None = None, @@ -1606,6 +1609,9 @@ def __init__( builtin_plugin_directories: Absolute paths to trusted plugin directories bundled by the host. When non-empty, the complete set is registered during startup before sessions can be created. + extension_launch_provider: Connection-level extension launch profile + provider. When set, it is registered during startup before any + session can be created. use_logged_in_user: Use the logged-in user for authentication. ``None`` (default) resolves to ``True`` unless ``github_token`` is set. @@ -1659,6 +1665,7 @@ def __init__( github_token=github_token, base_directory=base_directory, builtin_plugin_directories=tuple(builtin_plugin_directories or ()), + extension_launch_provider=extension_launch_provider, use_logged_in_user=use_logged_in_user, telemetry=telemetry, session_fs=session_fs, @@ -1975,6 +1982,9 @@ async def _start(self) -> None: start_time, ) + if self._options.extension_launch_provider is not None: + await self.rpc.register_extension_launch_provider() + if self._options.builtin_plugin_directories: assert self._client is not None try: @@ -4865,6 +4875,7 @@ def _register_client_global_handlers(self) -> None: self._client, ClientGlobalApiHandlers( hooks=_HooksAdapter(self._get_session), + extension_launch_provider=self._options.extension_launch_provider, llm_inference=llm_inference_adapter, git_hub_telemetry=github_telemetry_adapter, git_hub_token=self._github_token_provider_adapter, diff --git a/python/copilot/generated/rpc.py b/python/copilot/generated/rpc.py index b45ada8420..11999692a7 100644 --- a/python/copilot/generated/rpc.py +++ b/python/copilot/generated/rpc.py @@ -6,7 +6,7 @@ from typing import ClassVar, TYPE_CHECKING -from .session_events import AbortReason, AgentModelPolicy, Attachment, AutoTier, ContextTier, EmbeddedBlobResourceContents, EmbeddedTextResourceContents, IndexedSearchState, ManagedSettingsResolvedSource, McpOauthHttpResponse, McpOauthWWWAuthenticateParams, McpServerMetadata, McpServerSource, McpServerStatus, ModelChangeSource, PermissionDecisionSource, PermissionMode, PermissionPromptRequest, PermissionRule, ReasoningSummary, RemediationAction, SessionEvent, SessionLimitsConfig, SessionMode, ShutdownType, SkillSource, TaskBlocker, TaskCompletionOutcome, UserToolSessionApproval, Verbosity +from .session_events import AbortReason, AgentModelPolicy, Attachment, AutoTier, ContextTier, EmbeddedBlobResourceContents, EmbeddedTextResourceContents, IndexedSearchState, ManagedSettingsResolvedSource, McpOauthHttpResponse, McpOauthWWWAuthenticateParams, McpServerMetadata, McpServerSource, McpServerStatus, ModelChangeSource, PermissionDecisionSource, PermissionMode, PermissionPromptRequest, PermissionRule, ReasoningSummary, RemediationAction, SessionEvent, SessionLimitsConfig, SessionMode, ShutdownType, SkillSource, TaskBlocker, TaskCompletionOutcome, UserToolSessionApproval, Verbosity, _load_Attachment, _load_PermissionPromptRequest, _load_UserToolSessionApproval if TYPE_CHECKING: from .._jsonrpc import JsonRpcClient @@ -3853,7 +3853,7 @@ class FleetStartRequest: @staticmethod def from_dict(obj: Any) -> 'FleetStartRequest': assert isinstance(obj, dict) - attachments = from_union([lambda x: from_list(Attachment.from_dict, x), from_none], obj.get("attachments")) + attachments = from_union([lambda x: from_list(_load_Attachment, x), from_none], obj.get("attachments")) billable = from_union([from_bool, from_none], obj.get("billable")) prompt = from_union([from_str, from_none], obj.get("prompt")) wait = from_union([from_bool, from_none], obj.get("wait")) @@ -7902,7 +7902,7 @@ class PendingPermissionRequest: @staticmethod def from_dict(obj: Any) -> 'PendingPermissionRequest': assert isinstance(obj, dict) - request = PermissionPromptRequest.from_dict(obj.get("request")) + request = _load_PermissionPromptRequest(obj.get("request")) request_id = from_str(obj.get("requestId")) recovery_episode_id = from_union([from_str, from_none], obj.get("recoveryEpisodeId")) return PendingPermissionRequest(request, request_id, recovery_episode_id) @@ -10982,7 +10982,7 @@ class SendMessageItem: def from_dict(obj: Any) -> 'SendMessageItem': assert isinstance(obj, dict) prompt = from_str(obj.get("prompt")) - attachments = from_union([lambda x: from_list(Attachment.from_dict, x), from_none], obj.get("attachments")) + attachments = from_union([lambda x: from_list(_load_Attachment, x), from_none], obj.get("attachments")) billable = from_union([from_bool, from_none], obj.get("billable")) display_prompt = from_union([from_str, from_none], obj.get("displayPrompt")) required_tool = from_union([from_str, from_none], obj.get("requiredTool")) @@ -23281,7 +23281,7 @@ class PermissionDecisionApprovedForLocation: @staticmethod def from_dict(obj: Any) -> 'PermissionDecisionApprovedForLocation': assert isinstance(obj, dict) - approval = UserToolSessionApproval.from_dict(obj.get("approval")) + approval = _load_UserToolSessionApproval(obj.get("approval")) location_key = from_str(obj.get("locationKey")) return PermissionDecisionApprovedForLocation(approval, location_key) @@ -23307,7 +23307,7 @@ class PermissionDecisionApprovedForSession: @staticmethod def from_dict(obj: Any) -> 'PermissionDecisionApprovedForSession': assert isinstance(obj, dict) - approval = UserToolSessionApproval.from_dict(obj.get("approval")) + approval = _load_UserToolSessionApproval(obj.get("approval")) return PermissionDecisionApprovedForSession(approval) def to_dict(self) -> dict: @@ -25176,7 +25176,7 @@ class QueueAppendSteeringRequest: def from_dict(obj: Any) -> 'QueueAppendSteeringRequest': assert isinstance(obj, dict) agent_mode = SendAgentMode(obj.get("agentMode")) - attachments = from_list(Attachment.from_dict, obj.get("attachments")) + attachments = from_list(_load_Attachment, obj.get("attachments")) display_prompt = from_str(obj.get("displayPrompt")) expected_prompt = from_str(obj.get("expectedPrompt")) message_id = from_str(obj.get("messageId")) @@ -25248,7 +25248,7 @@ def from_dict(obj: Any) -> 'QueueInsertMessage': assert isinstance(obj, dict) prompt = from_str(obj.get("prompt")) agent_mode = from_union([SendAgentMode, from_none], obj.get("agentMode")) - attachments = from_union([lambda x: from_list(Attachment.from_dict, x), from_none], obj.get("attachments")) + attachments = from_union([lambda x: from_list(_load_Attachment, x), from_none], obj.get("attachments")) billable = from_union([from_bool, from_none], obj.get("billable")) delivery = from_union([from_str, from_none], obj.get("delivery")) display_prompt = from_union([from_str, from_none], obj.get("displayPrompt")) @@ -32243,7 +32243,7 @@ def from_dict(obj: Any) -> 'SendRequest': assert isinstance(obj, dict) prompt = from_str(obj.get("prompt")) agent_mode = from_union([SendAgentMode, from_none], obj.get("agentMode")) - attachments = from_union([lambda x: from_list(Attachment.from_dict, x), from_none], obj.get("attachments")) + attachments = from_union([lambda x: from_list(_load_Attachment, x), from_none], obj.get("attachments")) billable = from_union([from_bool, from_none], obj.get("billable")) display_prompt = from_union([from_str, from_none], obj.get("displayPrompt")) mode = from_union([SendMode, from_none], obj.get("mode")) diff --git a/python/copilot/session.py b/python/copilot/session.py index 4ce1862f48..6534e4cbae 100644 --- a/python/copilot/session.py +++ b/python/copilot/session.py @@ -278,7 +278,25 @@ class BlobAttachment(TypedDict): displayName: NotRequired[str] -Attachment = FileAttachment | DirectoryAttachment | SelectionAttachment | BlobAttachment +class ExtensionContextAttachment(TypedDict): + """Structured context contributed by an extension.""" + + type: Literal["extension_context"] + capturedAt: str + extensionId: str + title: str + canvasId: NotRequired[str] + instanceId: NotRequired[str] + payload: NotRequired[Any] + + +Attachment = ( + FileAttachment + | DirectoryAttachment + | SelectionAttachment + | BlobAttachment + | ExtensionContextAttachment +) @dataclass(frozen=True) @@ -1758,7 +1776,8 @@ async def send( Args: prompt: The message text to send. - attachments: Optional file, directory, or selection attachments. + attachments: Optional file, directory, selection, blob, or extension-context + attachments. source: Optional message provenance (``"user"``, ``"system"``, or :class:`AgentMessageSource` for an identified agent). Omitted when None, preserving the runtime's default for user messages. diff --git a/python/e2e/_scenario_fake_cli.py b/python/e2e/_scenario_fake_cli.py new file mode 100644 index 0000000000..f5ffecff60 --- /dev/null +++ b/python/e2e/_scenario_fake_cli.py @@ -0,0 +1,362 @@ +"""Deterministic bidirectional JSON-RPC CLI used by scenario-parity E2Es.""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from copilot import CopilotClient, RuntimeConnection + +from .testharness import DEFAULT_GITHUB_TOKEN, E2ETestContext + +SCENARIO_FAKE_CLI_SCRIPT = r""" +const fs = require("fs"); + +const scenarioIndex = process.argv.indexOf("--scenario"); +const captureIndex = process.argv.indexOf("--capture-file"); +const scenario = scenarioIndex >= 0 ? process.argv[scenarioIndex + 1] : ""; +const captureFile = captureIndex >= 0 ? process.argv[captureIndex + 1] : undefined; +const capture = { scenario, requests: [], callbackResponses: [] }; + +let pendingCreateId; +let pendingResumeId; +let resumeCallbackStep = 0; +let resumeAttempts = 0; +let buffer = Buffer.alloc(0); + +function saveCapture() { + if (captureFile) { + fs.writeFileSync(captureFile, JSON.stringify(capture)); + } +} + +function writeMessage(message) { + const body = JSON.stringify(message); + process.stdout.write( + `Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}` + ); +} + +function writeResponse(id, result) { + writeMessage({ jsonrpc: "2.0", id, result }); +} + +function writeError(id, code, message, data = null) { + writeMessage({ jsonrpc: "2.0", id, error: { code, message, data } }); +} + +function writeNotification(method, params) { + writeMessage({ jsonrpc: "2.0", method, params }); +} + +function writeRequest(id, method, params) { + writeMessage({ jsonrpc: "2.0", id, method, params }); +} + +function canvasParams(operation) { + const params = { + sessionId: "scenario-session", + canvasId: "counter", + extensionId: "python-scenario-tests", + instanceId: `scenario-${operation}`, + host: { capabilities: { canvases: true } }, + session: { + workingDirectory: "Q:\\scenario-workspace", + }, + }; + if (operation === "open") { + params.input = { startValue: 7 }; + } else if (operation === "action") { + params.actionName = "increment"; + params.input = { amount: 5 }; + } + return params; +} + +function sendCanvasCallback(operation, id = "canvas-callback") { + const method = + operation === "action" ? "canvas.action.invoke" : `canvas.${operation}`; + writeRequest(id, method, canvasParams(operation)); +} + +function sessionStartEvent(sessionId) { + return { + id: "11111111-1111-4111-8111-111111111111", + parentId: null, + timestamp: "2026-01-02T03:04:05Z", + type: "session.start", + data: { + copilotVersion: "fake", + producer: "scenario-fake-cli", + sessionId, + startTime: "2026-01-02T03:04:05Z", + version: 1, + remoteSteerable: false, + }, + }; +} + +function assistantEvent() { + return { + id: "22222222-2222-4222-8222-222222222222", + parentId: null, + timestamp: "2026-01-02T03:04:06Z", + type: "assistant.message", + data: { + content: "scenario response", + messageId: "assistant-message", + }, + }; +} + +function idleEvent() { + return { + id: "33333333-3333-4333-8333-333333333333", + parentId: null, + timestamp: "2026-01-02T03:04:07Z", + type: "session.idle", + data: { mode: "interactive" }, + }; +} + +function remoteSteerableEvent() { + return { + id: "44444444-4444-4444-8444-444444444444", + parentId: null, + timestamp: "2026-01-02T03:04:05Z", + type: "session.remote_steerable_changed", + data: { remoteSteerable: true }, + }; +} + +function connectedMetadata() { + return { + kind: "coding-agent", + modifiedTime: "2026-01-02T03:04:05Z", + repository: { + owner: "github", + name: "copilot-sdk", + branch: "scenario-branch", + }, + sessionId: "remote-resource-id", + startTime: "2026-01-01T00:00:00Z", + name: "Scenario cloud session", + pullRequestNumber: 42, + resourceId: "remote-resource-id", + state: "running", + summary: "Remote task summary", + }; +} + +function handleRequest(message) { + capture.requests.push({ method: message.method, params: message.params }); + saveCapture(); + + switch (message.method) { + case "connect": + writeResponse(message.id, { + ok: true, + protocolVersion: 3, + version: "scenario-fake", + }); + return; + case "ping": + writeResponse(message.id, { + message: message.params?.message ?? "pong", + protocolVersion: 3, + timestamp: "2026-01-02T03:04:05Z", + }); + return; + case "session.create": { + const sessionId = + message.params?.sessionId ?? + (scenario === "cloud" ? "cloud-runtime-session" : "scenario-session"); + if (scenario.startsWith("canvas-error-")) { + pendingCreateId = message.id; + sendCanvasCallback(scenario.slice("canvas-error-".length)); + return; + } + if (scenario === "preallocated-event") { + writeNotification("session.event", { + sessionId, + event: sessionStartEvent(sessionId), + }); + } + writeResponse(message.id, { + sessionId, + workspacePath: null, + capabilities: null, + }); + if (scenario === "cloud") { + writeNotification("session.event", { + sessionId, + event: sessionStartEvent(sessionId), + }); + } + return; + } + case "session.resume": + if (scenario === "canvas-resume") { + pendingResumeId = message.id; + sendCanvasCallback("open", "resume-open"); + return; + } + if (scenario === "resume-retry") { + resumeAttempts += 1; + if (resumeAttempts === 1) { + writeError( + message.id, + -32001, + "Session not found before acceptance", + { recoverable: true } + ); + return; + } + } + if (scenario === "resume-fail") { + writeError( + message.id, + -32001, + "Session not found before acceptance", + { recoverable: true } + ); + return; + } + writeResponse(message.id, { + sessionId: message.params.sessionId, + workspacePath: null, + capabilities: null, + openCanvases: message.params.openCanvases ?? [], + }); + return; + case "sessions.connect": + writeResponse(message.id, { + sessionId: "runtime-session-id", + metadata: connectedMetadata(), + }); + return; + case "session.remote.notifySteerableChanged": + writeResponse(message.id, {}); + writeNotification("session.event", { + sessionId: message.params.sessionId, + event: remoteSteerableEvent(), + }); + return; + case "session.send": + if (scenario === "send-fail") { + saveCapture(); + process.exit(23); + return; + } + writeResponse(message.id, { messageId: "user-message" }); + writeNotification("session.event", { + sessionId: message.params.sessionId, + event: assistantEvent(), + }); + writeNotification("session.event", { + sessionId: message.params.sessionId, + event: idleEvent(), + }); + return; + case "session.options.update": + case "session.detach": + writeResponse(message.id, { success: true }); + return; + default: + writeResponse(message.id, {}); + } +} + +function handleResponse(message) { + capture.callbackResponses.push(message); + saveCapture(); + + if (pendingCreateId) { + const createId = pendingCreateId; + pendingCreateId = undefined; + writeResponse(createId, { + sessionId: "scenario-session", + workspacePath: null, + capabilities: null, + }); + return; + } + + if (pendingResumeId) { + resumeCallbackStep += 1; + if (resumeCallbackStep === 1) { + sendCanvasCallback("action", "resume-action"); + } else if (resumeCallbackStep === 2) { + sendCanvasCallback("close", "resume-close"); + } else { + const resumeId = pendingResumeId; + pendingResumeId = undefined; + writeResponse(resumeId, { + sessionId: "scenario-session", + workspacePath: null, + capabilities: null, + openCanvases: [], + }); + } + } +} + +function handleMessage(message) { + if (Object.prototype.hasOwnProperty.call(message, "method")) { + handleRequest(message); + } else if (Object.prototype.hasOwnProperty.call(message, "id")) { + handleResponse(message); + } +} + +function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length header"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handleMessage(JSON.parse(body)); + } +} + +saveCapture(); +process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); +}); +process.stdin.resume(); +""" + + +def create_scenario_client( + ctx: E2ETestContext, + scenario: str, +) -> tuple[CopilotClient, Path]: + """Create a client backed by the deterministic scenario fake CLI.""" + cli_path = Path(ctx.work_dir, f"scenario-fake-{scenario}.js") + capture_path = Path(ctx.work_dir, f"scenario-fake-{scenario}.json") + cli_path.write_text(SCENARIO_FAKE_CLI_SCRIPT, encoding="utf-8") + client = CopilotClient( + connection=RuntimeConnection.for_stdio( + path=str(cli_path), + args=("--scenario", scenario, "--capture-file", str(capture_path)), + ), + working_directory=ctx.work_dir, + env=ctx.get_env(), + github_token=DEFAULT_GITHUB_TOKEN, + use_logged_in_user=False, + ) + return client, capture_path + + +def read_scenario_capture(path: Path) -> dict[str, Any]: + """Read the fake CLI's latest request and callback-response capture.""" + return json.loads(path.read_text(encoding="utf-8")) diff --git a/python/e2e/test_client_e2e.py b/python/e2e/test_client_e2e.py index 1e8ea82e55..a42c31885c 100644 --- a/python/e2e/test_client_e2e.py +++ b/python/e2e/test_client_e2e.py @@ -212,9 +212,16 @@ async def test_should_report_error_with_stderr_when_cli_fails_to_start(self): on_permission_request=PermissionHandler.approve_all ) await session.send("test") - # Error message varies by platform (EINVAL on Windows, EPIPE on Linux) + # A completed process preserves stderr even if its exit code was not yet + # available on the first failure; a broken transport can report EINVAL/EPIPE. error_msg = str(exc_info2.value).lower() - assert "invalid" in error_msg or "pipe" in error_msg or "closed" in error_msg + if "cli process exited with code" in error_msg: + assert ( + error_msg.partition("stderr:")[2] + == error_message.lower().partition("stderr:")[2] + ) + else: + assert "invalid" in error_msg or "pipe" in error_msg or "closed" in error_msg finally: await client.force_stop() diff --git a/python/e2e/test_client_options_e2e.py b/python/e2e/test_client_options_e2e.py index b07e9a5402..f6b2cd37cc 100644 --- a/python/e2e/test_client_options_e2e.py +++ b/python/e2e/test_client_options_e2e.py @@ -26,6 +26,9 @@ CloudSessionRepository, CopilotClient, ExtensionInfo, + ExtensionLaunchProfile, + ExtensionLaunchProviderResolveRequest, + ExtensionLaunchProviderResolveResult, OpenCanvasInstance, RemoteSessionMode, RuntimeConnection, @@ -85,6 +88,8 @@ def _get_available_port() -> int: const captureIndex = process.argv.indexOf("--capture-file"); const captureFile = captureIndex >= 0 ? process.argv[captureIndex + 1] : undefined; const requests = []; +const clientResponses = []; +let extensionRegistrationId; function saveCapture() { if (!captureFile) { @@ -94,6 +99,7 @@ def _get_available_port() -> int: args: process.argv.slice(2), cwd: process.cwd(), requests, + clientResponses, env: { COPILOT_HOME: process.env.COPILOT_HOME, COPILOT_SDK_AUTH_TOKEN: process.env.COPILOT_SDK_AUTH_TOKEN, @@ -139,6 +145,15 @@ def _get_available_port() -> int: if (!Object.prototype.hasOwnProperty.call(message, "id")) { return; } + if (!message.method) { + clientResponses.push(message); + saveCapture(); + if (message.id === 9001 && extensionRegistrationId !== undefined) { + writeResponse(extensionRegistrationId, {}); + extensionRegistrationId = undefined; + } + return; + } requests.push({ method: message.method, params: message.params }); saveCapture(); if (message.method === "connect") { @@ -149,6 +164,16 @@ def _get_available_port() -> int: writeResponse(message.id, { message: "pong", protocolVersion: 3, timestamp: Date.now() }); return; } + if (message.method === "registerExtensionLaunchProvider") { + extensionRegistrationId = message.id; + writeRequest(9001, "extensionLaunchProvider.resolve", { + id: "project:python-e2e", + name: "python-e2e", + modulePath: "/extensions/python-e2e.py", + source: "project", + }); + return; + } if (message.method === "session.create") { const sessionId = message.params?.sessionId ?? message.params?.session_id ?? "fake-session"; writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); @@ -179,6 +204,11 @@ def _get_available_port() -> int: const body = JSON.stringify({ jsonrpc: "2.0", id, result }); process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); } + +function writeRequest(id, method, params) { + const body = JSON.stringify({ jsonrpc: "2.0", id, method, params }); + process.stdout.write(`Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}`); +} """ @@ -329,6 +359,68 @@ async def test_should_propagate_process_options_to_spawned_cli(self, ctx: E2ETes except Exception: await client.force_stop() + async def test_should_register_and_invoke_extension_launch_provider(self, ctx: E2ETestContext): + cli_path = os.path.join(ctx.work_dir, "fake-cli-extension-provider.js") + capture_path = os.path.join(ctx.work_dir, "fake-cli-extension-provider-capture.json") + with open(cli_path, "w") as f: + f.write(FAKE_STDIO_CLI_SCRIPT) + + class RecordingProvider: + request: ExtensionLaunchProviderResolveRequest | None = None + + async def resolve( + self, params: ExtensionLaunchProviderResolveRequest + ) -> ExtensionLaunchProviderResolveResult: + self.request = params + return ExtensionLaunchProviderResolveResult( + launch=ExtensionLaunchProfile( + executable="python", + args=["extension-host"], + env={"EXTENSION_SOURCE": "python"}, + ) + ) + + provider = RecordingProvider() + client = CopilotClient( + **_make_options( + ctx, + cli_path=cli_path, + cli_args=["--capture-file", capture_path], + extension_launch_provider=provider, + github_token=None, + use_logged_in_user=False, + ) + ) + try: + await client.start() + + assert provider.request is not None + assert provider.request.id == "project:python-e2e" + assert provider.request.name == "python-e2e" + assert provider.request.module_path == "/extensions/python-e2e.py" + assert provider.request.source.value == "project" + + with open(capture_path) as f: + capture = json.load(f) + assert "registerExtensionLaunchProvider" in [ + request["method"] for request in capture["requests"] + ] + assert capture["clientResponses"] == [ + { + "jsonrpc": "2.0", + "id": 9001, + "result": { + "launch": { + "executable": "python", + "args": ["extension-host"], + "env": {"EXTENSION_SOURCE": "python"}, + } + }, + } + ] + finally: + await client.force_stop() + async def test_should_send_empty_mode_custom_agent_locality_defaults(self, ctx: E2ETestContext): cli_path = os.path.join(ctx.work_dir, "fake-cli-empty.js") capture_path = os.path.join(ctx.work_dir, "fake-cli-empty-capture.json") diff --git a/python/e2e/test_rpc_generated_surface_e2e.py b/python/e2e/test_rpc_generated_surface_e2e.py new file mode 100644 index 0000000000..bc5d0ae52c --- /dev/null +++ b/python/e2e/test_rpc_generated_surface_e2e.py @@ -0,0 +1,857 @@ +"""Offline E2E coverage for generated outbound RPC methods.""" + +from __future__ import annotations + +import dataclasses +import enum +import inspect +import json +import os +import types +import typing +from datetime import UTC, datetime +from pathlib import Path +from uuid import UUID + +import pytest + +from copilot import CopilotClient, RuntimeConnection, rpc +from copilot._jsonrpc import JsonRpcError +from copilot.generated import session_events as generated_session_events +from copilot.session import CopilotSession + +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + +_SESSION_ID = "generated-rpc-surface-session" +_SAMPLE_UUID = UUID("12345678-1234-5678-1234-567812345678") +_OBJECTIVE_METHODS = { + "session.workspaces.readAutopilotObjective", + "session.workspaces.writeAutopilotObjective", + "session.workspaces.deleteAutopilotObjective", + "session.workspaces.autopilotObjectiveExists", +} + +_GAP_METHODS = { + "registerExtensionLaunchProvider": rpc.ServerRpc.register_extension_launch_provider, + "hooks.discover": rpc.ServerHooksApi.discover, + "models.getBuiltInCatalog": rpc.ServerModelsApi.get_built_in_catalog, + "mcp.planInstall": rpc.ServerMcpApi.plan_install, + "extensions.discover": rpc.ServerExtensionsApi.discover, + "extensions.enable": rpc.ServerExtensionsApi.enable, + "extensions.disable": rpc.ServerExtensionsApi.disable, + "catalog.search": rpc.ServerCatalogApi.search, + "plugins.builtin.set": rpc.ServerPluginsBuiltinApi.set, + "skills.config.setSkillDisabled": rpc.ServerSkillsConfigApi.set_skill_disabled, + "commands.list": rpc.ServerCommandsApi.list, + "managedSettings.read": rpc.ServerManagedSettingsApi.read, + "llmInference.setProvider": rpc.ServerLlmInferenceApi.set_provider, + "sessions.getClientMetadata": rpc.ServerSessionsApi.get_client_metadata, + "sessions.readPersistedEvents": rpc.ServerSessionsApi.read_persisted_events, + "session.send": rpc.SessionRpc.send, + "session.sendMessages": rpc.SessionRpc.send_messages, + "session.abort": rpc.SessionRpc.abort, + "session.interruptMainTurn": rpc.SessionRpc.interrupt_main_turn, + "session.cancelAllBackgroundAgents": rpc.SessionRpc.cancel_all_background_agents, + "session.log": rpc.SessionRpc.log, + "session.sandbox.getEnforcementStatus": rpc.SandboxApi.get_enforcement_status, + "session.sandbox.disableForSession": rpc.SandboxApi.disable_for_session, + "session.debug.collectLogs": rpc.DebugApi.collect_logs, + "session.factory.run": rpc.FactoryApi.run, + "session.factory.resume": rpc.FactoryApi.resume, + "session.factory.getRun": rpc.FactoryApi.get_run, + "session.factory.listRuns": rpc.FactoryApi.list_runs, + "session.factory.getRunDetail": rpc.FactoryApi.get_run_detail, + "session.factory.getRunProgress": rpc.FactoryApi.get_run_progress, + "session.factory.cancel": rpc.FactoryApi.cancel, + "session.factory.pause": rpc.FactoryApi.pause, + "session.factory.log": rpc.FactoryApi.log, + "session.factory.agent": rpc.FactoryApi.agent, + "session.factory.journal.get": rpc.FactoryJournalApi.get, + "session.factory.journal.put": rpc.FactoryJournalApi.put, + "session.model.switchAutoTier": rpc.ModelApi.switch_auto_tier, + "session.model.setAllowedModels": rpc.ModelApi.set_allowed_models, + "session.workspaces.updateMetadata": rpc.WorkspacesApi.update_metadata, + "session.workspaces.ensure": rpc.WorkspacesApi.ensure, + "session.workspaces.statFile": rpc.WorkspacesApi.stat_file, + "session.workspaces.createDirectory": rpc.WorkspacesApi.create_directory, + "session.workspaces.removePath": rpc.WorkspacesApi.remove_path, + "session.workspaces.renamePath": rpc.WorkspacesApi.rename_path, + "session.workspaces.addSummary": rpc.WorkspacesApi.add_summary, + "session.workspaces.truncateSummaries": rpc.WorkspacesApi.truncate_summaries, + "session.workspaces.readAutopilotObjective": rpc.WorkspacesApi.read_autopilot_objective, + "session.workspaces.writeAutopilotObjective": rpc.WorkspacesApi.write_autopilot_objective, + "session.workspaces.deleteAutopilotObjective": rpc.WorkspacesApi.delete_autopilot_objective, + "session.workspaces.autopilotObjectiveExists": (rpc.WorkspacesApi.autopilot_objective_exists), + "session.autopilotObjective.getState": rpc.AutopilotObjectiveApi.get_state, + "session.agent.setPrompt": rpc.AgentApi.set_prompt, + "session.tasks.register": rpc.TasksApi.register, + "session.tasks.update": rpc.TasksApi.update, + "session.mcp.moveLoadingToBackground": rpc.McpApi.move_loading_to_background, + "session.mcp.startServer": rpc.McpApi.start_server, + "session.mcp.restartServer": rpc.McpApi.restart_server, + "session.mcp.oauth.authenticationStateChanged": (rpc.McpOauthApi.authentication_state_changed), + "session.mcp.oauth.probe": rpc.McpOauthApi.probe, + "session.mcp.oauth.respond": rpc.McpOauthApi.respond, + "session.mcp.resources.read": rpc.McpResourcesApi.read, + "session.mcp.resources.list": rpc.McpResourcesApi.list, + "session.mcp.resources.listTemplates": rpc.McpResourcesApi.list_templates, + "session.tools.execute": rpc.ToolsApi.execute, + "session.tools.getBuiltinDescriptors": rpc.ToolsApi.get_builtin_descriptors, + "session.tools.taskCompleteEventData": rpc.ToolsApi.task_complete_event_data, + "session.tools.set": rpc.ToolsApi.set, + "session.permissions.configure": rpc.PermissionsApi.configure, + "session.permissions.pendingRequests": rpc.PermissionsApi.pending_requests, + "session.permissions.modifyRules": rpc.PermissionsApi.modify_rules, + "session.permissions.setRequired": rpc.PermissionsApi.set_required, + "session.permissions.notifyPromptShown": rpc.PermissionsApi.notify_prompt_shown, + "session.permissions.paths.list": rpc.PermissionsPathsApi.list, + "session.permissions.paths.add": rpc.PermissionsPathsApi.add, + "session.permissions.paths.updatePrimary": rpc.PermissionsPathsApi.update_primary, + "session.permissions.paths.isPathWithinAllowedDirectories": ( + rpc.PermissionsPathsApi.is_path_within_allowed_directories + ), + "session.permissions.paths.isPathWithinWorkspace": ( + rpc.PermissionsPathsApi.is_path_within_workspace + ), + "session.permissions.locations.resolve": rpc.PermissionsLocationsApi.resolve, + "session.permissions.locations.apply": rpc.PermissionsLocationsApi.apply, + "session.permissions.locations.addToolApproval": ( + rpc.PermissionsLocationsApi.add_tool_approval + ), + "session.permissions.folderTrust.isTrusted": rpc.PermissionsFolderTrustApi.is_trusted, + "session.permissions.folderTrust.addTrusted": rpc.PermissionsFolderTrustApi.add_trusted, + "session.permissions.urls.setUnrestrictedMode": (rpc.PermissionsUrlsApi.set_unrestricted_mode), + "session.metadata.getClientMetadata": rpc.MetadataApi.get_client_metadata, + "session.metadata.updateClientMetadata": rpc.MetadataApi.update_client_metadata, + "session.contentExclusion.checkPaths": rpc.ContentExclusionApi.check_paths, + "session.history.clearContext": rpc.HistoryApi.clear_context, + "session.queue.moveItem": rpc.QueueApi.move_item, + "session.queue.insertAt": rpc.QueueApi.insert_at, + "session.queue.removeAt": rpc.QueueApi.remove_at, + "session.queue.updateText": rpc.QueueApi.update_text, + "session.queue.duplicateAt": rpc.QueueApi.duplicate_at, + "session.queue.setDrainPaused": rpc.QueueApi.set_drain_paused, + "session.queue.sendNow": rpc.QueueApi.send_now, + "session.limitPrediction.predict": rpc.LimitPredictionApi.predict, +} + +_FAKE_CLI = r""" +const fs = require("fs"); + +function argValue(name) { + const index = process.argv.indexOf(name); + return index >= 0 ? process.argv[index + 1] : undefined; +} + +const captureFile = argValue("--capture-file"); +const responses = JSON.parse(fs.readFileSync(argValue("--responses-file"), "utf8")); +const requests = []; +let objective = null; +let buffer = Buffer.alloc(0); + +function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify({ requests })); +} + +function writeResponse(id, result) { + const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + process.stdout.write( + `Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}` + ); +} + +function writeError(id, code, message, data) { + const body = JSON.stringify({ jsonrpc: "2.0", id, error: { code, message, data } }); + process.stdout.write( + `Content-Length: ${Buffer.byteLength(body, "utf8")}\r\n\r\n${body}` + ); +} + +function handleMessage(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) return; + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + writeResponse(message.id, { ok: true, protocolVersion: 3, version: "fake" }); + return; + } + if (message.method === "ping") { + writeResponse(message.id, { + message: "pong", + protocolVersion: 3, + timestamp: 1770000000000, + }); + return; + } + if (message.method === "catalog.search" && message.params.query === "raise-jsonrpc-error") { + writeError(message.id, -32077, "deterministic catalog failure", { + retryable: false, + source: "fake-cli", + }); + return; + } + if (message.method === "session.workspaces.writeAutopilotObjective") { + const operation = objective === null ? "created" : "updated"; + objective = message.params.content; + writeResponse(message.id, { operation }); + return; + } + if (message.method === "session.workspaces.readAutopilotObjective") { + writeResponse(message.id, { content: objective }); + return; + } + if (message.method === "session.workspaces.autopilotObjectiveExists") { + writeResponse(message.id, { exists: objective !== null }); + return; + } + if (message.method === "session.workspaces.deleteAutopilotObjective") { + const deleted = objective !== null; + objective = null; + writeResponse(message.id, { deleted }); + return; + } + + writeResponse(message.id, responses[message.method] ?? {}); +} + +function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length header"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handleMessage(JSON.parse(body)); + } +} + +saveCapture(); +process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); +}); +process.stdin.resume(); +""" + + +def _type_globals() -> dict[str, object]: + return { + **vars(generated_session_events), + **vars(rpc), + } + + +def _sample_value( + annotation: object, + *, + depth: int = 0, + stack: tuple[object, ...] = (), +) -> object: + if annotation in (inspect.Signature.empty, typing.Any, object): + return {"sample": "value"} + if annotation in (None, type(None)): + return None + + origin = typing.get_origin(annotation) + arguments = typing.get_args(annotation) + if origin in (typing.Union, types.UnionType): + choices = [item for item in arguments if item is not type(None)] + return _sample_value(choices[0] if choices else type(None), depth=depth, stack=stack) + if origin is typing.Literal: + return arguments[0] + if origin is list: + return [_sample_value(arguments[0], depth=depth + 1, stack=stack)] + if origin is dict: + return {"key": _sample_value(arguments[1], depth=depth + 1, stack=stack)} + if origin is tuple: + return tuple( + _sample_value(item, depth=depth + 1, stack=stack) + for item in arguments + if item is not Ellipsis + ) + if origin is typing.Annotated: + return _sample_value(arguments[0], depth=depth, stack=stack) + + if isinstance(annotation, type) and issubclass(annotation, enum.Enum): + return next(iter(annotation)) + if annotation is str: + return "sample-value" + if annotation is bool: + return True + if annotation is int: + return 7 + if annotation is float: + return 1.5 + if annotation is datetime: + return datetime(2026, 1, 2, 3, 4, 5, tzinfo=UTC) + if annotation is UUID: + return _SAMPLE_UUID + if annotation is list: + return [{"sample": "value"}] + if annotation is dict: + return {"key": "value"} + + if dataclasses.is_dataclass(annotation): + if annotation in stack: + return None + hints = typing.get_type_hints( + annotation, + globalns=_type_globals(), + localns=_type_globals(), + ) + values = {} + for field in dataclasses.fields(annotation): + required = ( + field.default is dataclasses.MISSING + and field.default_factory is dataclasses.MISSING + ) + if required or depth < 4: + values[field.name] = _sample_value( + hints.get(field.name, field.type), + depth=depth + 1, + stack=(*stack, annotation), + ) + constructor = typing.cast(typing.Callable[..., object], annotation) + return constructor(**values) + + return {"sample": "value"} + + +def _method_hints(method: object) -> dict[str, object]: + return typing.get_type_hints( + method, + globalns=_type_globals(), + localns=_type_globals(), + ) + + +def _request_for(method: object) -> typing.Any: + return _sample_value(_method_hints(method)["params"]) + + +def _json_value(value: object) -> object: + if hasattr(value, "to_dict"): + serializer = typing.cast(typing.Callable[[], object], getattr(value, "to_dict")) + return serializer() + if isinstance(value, enum.Enum): + return value.value + if isinstance(value, datetime): + return value.isoformat() + if isinstance(value, UUID): + return str(value) + if isinstance(value, list): + return [_json_value(item) for item in value] + if isinstance(value, dict): + return {key: _json_value(item) for key, item in value.items()} + return value + + +def _response_payloads() -> dict[str, object]: + return { + rpc_method: _json_value(_sample_value(_method_hints(method)["return"])) + for rpc_method, method in _GAP_METHODS.items() + } + + +def _assert_result_matches_payload(result: object, payload: object) -> None: + assert _json_value(result) == payload + + +def _expected_params(method: object, *, session_scoped: bool) -> dict[str, object]: + hints = _method_hints(method) + params = _request_for(method).to_dict() if "params" in hints else {} + if session_scoped: + params["sessionId"] = _SESSION_ID + return params + + +def _assert_request_serialization(capture_path: Path) -> None: + capture = json.loads(capture_path.read_text(encoding="utf-8")) + requests = capture["requests"] + + for rpc_method, method in _GAP_METHODS.items(): + matching = [request for request in requests if request["method"] == rpc_method] + assert matching, f"Missing captured request for {rpc_method}" + expected = _expected_params(method, session_scoped=rpc_method.startswith("session.")) + if rpc_method == "session.workspaces.writeAutopilotObjective": + expected["content"] = "# Deterministic objective\n\nCover generated RPC methods." + assert expected in [request["params"] for request in matching] + + catalog_error = next( + request + for request in requests + if request["method"] == "catalog.search" + and request["params"]["query"] == "raise-jsonrpc-error" + ) + assert catalog_error["params"]["query"] == "raise-jsonrpc-error" + + +async def test_generated_rpc_gap_methods_round_trip_over_fake_cli( + ctx: E2ETestContext, +) -> None: + work_dir = Path(ctx.work_dir) + suffix = str(os.getpid()) + cli_path = work_dir / f"generated-rpc-fake-cli-{suffix}.js" + capture_path = work_dir / f"generated-rpc-capture-{suffix}.json" + responses_path = work_dir / f"generated-rpc-responses-{suffix}.json" + cli_path.write_text(_FAKE_CLI, encoding="utf-8") + responses = _response_payloads() + responses_path.write_text(json.dumps(responses), encoding="utf-8") + + client = CopilotClient( + connection=RuntimeConnection.for_stdio( + path=str(cli_path), + args=[ + "--capture-file", + str(capture_path), + "--responses-file", + str(responses_path), + ], + ), + working_directory=ctx.work_dir, + env=ctx.get_env(), + use_logged_in_user=False, + ) + results: dict[str, object] = {} + + try: + await client.start() + assert client._client is not None + session = CopilotSession(_SESSION_ID, client._client) + + results[ + "registerExtensionLaunchProvider" + ] = await client.rpc.register_extension_launch_provider() + results["hooks.discover"] = await client.rpc.hooks.discover( + _request_for(rpc.ServerHooksApi.discover) + ) + results["models.getBuiltInCatalog"] = await client.rpc.models.get_built_in_catalog() + results["mcp.planInstall"] = await client.rpc.mcp.plan_install( + _request_for(rpc.ServerMcpApi.plan_install) + ) + results["extensions.discover"] = await client.rpc.extensions.discover() + results["extensions.enable"] = await client.rpc.extensions.enable( + _request_for(rpc.ServerExtensionsApi.enable) + ) + results["extensions.disable"] = await client.rpc.extensions.disable( + _request_for(rpc.ServerExtensionsApi.disable) + ) + results["catalog.search"] = await client.rpc.catalog.search( + _request_for(rpc.ServerCatalogApi.search) + ) + results["plugins.builtin.set"] = await client.rpc.plugins.builtin.set( + _request_for(rpc.ServerPluginsBuiltinApi.set) + ) + results[ + "skills.config.setSkillDisabled" + ] = await client.rpc.skills.config.set_skill_disabled( + _request_for(rpc.ServerSkillsConfigApi.set_skill_disabled) + ) + results["commands.list"] = await client.rpc.commands.list() + results["managedSettings.read"] = await client.rpc.managed_settings.read() + results["llmInference.setProvider"] = await client.rpc.llm_inference.set_provider() + results["sessions.getClientMetadata"] = await client.rpc.sessions.get_client_metadata( + _request_for(rpc.ServerSessionsApi.get_client_metadata) + ) + results["sessions.readPersistedEvents"] = await client.rpc.sessions.read_persisted_events( + _request_for(rpc.ServerSessionsApi.read_persisted_events) + ) + + results["session.send"] = await session.rpc.send(_request_for(rpc.SessionRpc.send)) + results["session.sendMessages"] = await session.rpc.send_messages( + _request_for(rpc.SessionRpc.send_messages) + ) + results["session.abort"] = await session.rpc.abort(_request_for(rpc.SessionRpc.abort)) + results["session.interruptMainTurn"] = await session.rpc.interrupt_main_turn( + _request_for(rpc.SessionRpc.interrupt_main_turn) + ) + results[ + "session.cancelAllBackgroundAgents" + ] = await session.rpc.cancel_all_background_agents() + results["session.log"] = await session.rpc.log(_request_for(rpc.SessionRpc.log)) + results[ + "session.sandbox.getEnforcementStatus" + ] = await session.rpc.sandbox.get_enforcement_status() + results[ + "session.sandbox.disableForSession" + ] = await session.rpc.sandbox.disable_for_session( + _request_for(rpc.SandboxApi.disable_for_session) + ) + results["session.debug.collectLogs"] = await session.rpc.debug.collect_logs( + _request_for(rpc.DebugApi.collect_logs) + ) + + results["session.factory.run"] = await session.rpc.factory.run( + _request_for(rpc.FactoryApi.run) + ) + results["session.factory.resume"] = await session.rpc.factory.resume( + _request_for(rpc.FactoryApi.resume) + ) + results["session.factory.getRun"] = await session.rpc.factory.get_run( + _request_for(rpc.FactoryApi.get_run) + ) + results["session.factory.listRuns"] = await session.rpc.factory.list_runs( + _request_for(rpc.FactoryApi.list_runs) + ) + results["session.factory.getRunDetail"] = await session.rpc.factory.get_run_detail( + _request_for(rpc.FactoryApi.get_run_detail) + ) + results["session.factory.getRunProgress"] = await session.rpc.factory.get_run_progress( + _request_for(rpc.FactoryApi.get_run_progress) + ) + results["session.factory.cancel"] = await session.rpc.factory.cancel( + _request_for(rpc.FactoryApi.cancel) + ) + results["session.factory.pause"] = await session.rpc.factory.pause( + _request_for(rpc.FactoryApi.pause) + ) + results["session.factory.log"] = await session.rpc.factory.log( + _request_for(rpc.FactoryApi.log) + ) + results["session.factory.agent"] = await session.rpc.factory.agent( + _request_for(rpc.FactoryApi.agent) + ) + results["session.factory.journal.get"] = await session.rpc.factory.journal.get( + _request_for(rpc.FactoryJournalApi.get) + ) + results["session.factory.journal.put"] = await session.rpc.factory.journal.put( + _request_for(rpc.FactoryJournalApi.put) + ) + + results["session.model.switchAutoTier"] = await session.rpc.model.switch_auto_tier( + _request_for(rpc.ModelApi.switch_auto_tier) + ) + results["session.model.setAllowedModels"] = await session.rpc.model.set_allowed_models( + _request_for(rpc.ModelApi.set_allowed_models) + ) + + results["session.workspaces.updateMetadata"] = await session.rpc.workspaces.update_metadata( + _request_for(rpc.WorkspacesApi.update_metadata) + ) + results["session.workspaces.ensure"] = await session.rpc.workspaces.ensure( + _request_for(rpc.WorkspacesApi.ensure) + ) + results["session.workspaces.statFile"] = await session.rpc.workspaces.stat_file( + _request_for(rpc.WorkspacesApi.stat_file) + ) + results[ + "session.workspaces.createDirectory" + ] = await session.rpc.workspaces.create_directory( + _request_for(rpc.WorkspacesApi.create_directory) + ) + results["session.workspaces.removePath"] = await session.rpc.workspaces.remove_path( + _request_for(rpc.WorkspacesApi.remove_path) + ) + results["session.workspaces.renamePath"] = await session.rpc.workspaces.rename_path( + _request_for(rpc.WorkspacesApi.rename_path) + ) + results["session.workspaces.addSummary"] = await session.rpc.workspaces.add_summary( + _request_for(rpc.WorkspacesApi.add_summary) + ) + results[ + "session.workspaces.truncateSummaries" + ] = await session.rpc.workspaces.truncate_summaries( + _request_for(rpc.WorkspacesApi.truncate_summaries) + ) + + initial_objective = await session.rpc.workspaces.read_autopilot_objective() + assert initial_objective.content is None + initial_exists = await session.rpc.workspaces.autopilot_objective_exists() + assert initial_exists.exists is False + + objective_request = _request_for(rpc.WorkspacesApi.write_autopilot_objective) + objective_request.content = "# Deterministic objective\n\nCover generated RPC methods." + results[ + "session.workspaces.writeAutopilotObjective" + ] = await session.rpc.workspaces.write_autopilot_objective(objective_request) + assert results["session.workspaces.writeAutopilotObjective"].operation == "created" + + saved_objective = await session.rpc.workspaces.read_autopilot_objective() + assert saved_objective.content == objective_request.content + saved_exists = await session.rpc.workspaces.autopilot_objective_exists() + assert saved_exists.exists is True + results["session.workspaces.readAutopilotObjective"] = saved_objective + results["session.workspaces.autopilotObjectiveExists"] = saved_exists + + results[ + "session.workspaces.deleteAutopilotObjective" + ] = await session.rpc.workspaces.delete_autopilot_objective() + assert results["session.workspaces.deleteAutopilotObjective"].deleted is True + assert (await session.rpc.workspaces.read_autopilot_objective()).content is None + + results[ + "session.autopilotObjective.getState" + ] = await session.rpc.autopilot_objective.get_state() + results["session.agent.setPrompt"] = await session.rpc.agent.set_prompt( + _request_for(rpc.AgentApi.set_prompt) + ) + results["session.tasks.register"] = await session.rpc.tasks.register( + _request_for(rpc.TasksApi.register) + ) + results["session.tasks.update"] = await session.rpc.tasks.update( + _request_for(rpc.TasksApi.update) + ) + + results[ + "session.mcp.moveLoadingToBackground" + ] = await session.rpc.mcp.move_loading_to_background() + results["session.mcp.startServer"] = await session.rpc.mcp.start_server( + _request_for(rpc.McpApi.start_server) + ) + results["session.mcp.restartServer"] = await session.rpc.mcp.restart_server( + _request_for(rpc.McpApi.restart_server) + ) + results[ + "session.mcp.oauth.authenticationStateChanged" + ] = await session.rpc.mcp.oauth.authentication_state_changed( + _request_for(rpc.McpOauthApi.authentication_state_changed) + ) + results["session.mcp.oauth.probe"] = await session.rpc.mcp.oauth.probe( + _request_for(rpc.McpOauthApi.probe) + ) + results["session.mcp.oauth.respond"] = await session.rpc.mcp.oauth.respond( + _request_for(rpc.McpOauthApi.respond) + ) + results["session.mcp.resources.read"] = await session.rpc.mcp.resources.read( + _request_for(rpc.McpResourcesApi.read) + ) + results["session.mcp.resources.list"] = await session.rpc.mcp.resources.list( + _request_for(rpc.McpResourcesApi.list) + ) + results[ + "session.mcp.resources.listTemplates" + ] = await session.rpc.mcp.resources.list_templates( + _request_for(rpc.McpResourcesApi.list_templates) + ) + + results["session.tools.execute"] = await session.rpc.tools.execute( + _request_for(rpc.ToolsApi.execute) + ) + results[ + "session.tools.getBuiltinDescriptors" + ] = await session.rpc.tools.get_builtin_descriptors( + _request_for(rpc.ToolsApi.get_builtin_descriptors) + ) + results[ + "session.tools.taskCompleteEventData" + ] = await session.rpc.tools.task_complete_event_data( + _request_for(rpc.ToolsApi.task_complete_event_data) + ) + results["session.tools.set"] = await session.rpc.tools.set(_request_for(rpc.ToolsApi.set)) + + results["session.permissions.configure"] = await session.rpc.permissions.configure( + _request_for(rpc.PermissionsApi.configure) + ) + results[ + "session.permissions.pendingRequests" + ] = await session.rpc.permissions.pending_requests() + results["session.permissions.modifyRules"] = await session.rpc.permissions.modify_rules( + _request_for(rpc.PermissionsApi.modify_rules) + ) + results["session.permissions.setRequired"] = await session.rpc.permissions.set_required( + _request_for(rpc.PermissionsApi.set_required) + ) + results[ + "session.permissions.notifyPromptShown" + ] = await session.rpc.permissions.notify_prompt_shown( + _request_for(rpc.PermissionsApi.notify_prompt_shown) + ) + results["session.permissions.paths.list"] = await session.rpc.permissions.paths.list() + results["session.permissions.paths.add"] = await session.rpc.permissions.paths.add( + _request_for(rpc.PermissionsPathsApi.add) + ) + results[ + "session.permissions.paths.updatePrimary" + ] = await session.rpc.permissions.paths.update_primary( + _request_for(rpc.PermissionsPathsApi.update_primary) + ) + results[ + "session.permissions.paths.isPathWithinAllowedDirectories" + ] = await session.rpc.permissions.paths.is_path_within_allowed_directories( + _request_for(rpc.PermissionsPathsApi.is_path_within_allowed_directories) + ) + results[ + "session.permissions.paths.isPathWithinWorkspace" + ] = await session.rpc.permissions.paths.is_path_within_workspace( + _request_for(rpc.PermissionsPathsApi.is_path_within_workspace) + ) + results[ + "session.permissions.locations.resolve" + ] = await session.rpc.permissions.locations.resolve( + _request_for(rpc.PermissionsLocationsApi.resolve) + ) + results[ + "session.permissions.locations.apply" + ] = await session.rpc.permissions.locations.apply( + _request_for(rpc.PermissionsLocationsApi.apply) + ) + results[ + "session.permissions.locations.addToolApproval" + ] = await session.rpc.permissions.locations.add_tool_approval( + _request_for(rpc.PermissionsLocationsApi.add_tool_approval) + ) + results[ + "session.permissions.folderTrust.isTrusted" + ] = await session.rpc.permissions.folder_trust.is_trusted( + _request_for(rpc.PermissionsFolderTrustApi.is_trusted) + ) + results[ + "session.permissions.folderTrust.addTrusted" + ] = await session.rpc.permissions.folder_trust.add_trusted( + _request_for(rpc.PermissionsFolderTrustApi.add_trusted) + ) + results[ + "session.permissions.urls.setUnrestrictedMode" + ] = await session.rpc.permissions.urls.set_unrestricted_mode( + _request_for(rpc.PermissionsUrlsApi.set_unrestricted_mode) + ) + + results[ + "session.metadata.getClientMetadata" + ] = await session.rpc.metadata.get_client_metadata() + results[ + "session.metadata.updateClientMetadata" + ] = await session.rpc.metadata.update_client_metadata( + _request_for(rpc.MetadataApi.update_client_metadata) + ) + results[ + "session.contentExclusion.checkPaths" + ] = await session.rpc.content_exclusion.check_paths( + _request_for(rpc.ContentExclusionApi.check_paths) + ) + results["session.history.clearContext"] = await session.rpc.history.clear_context( + _request_for(rpc.HistoryApi.clear_context) + ) + + results["session.queue.moveItem"] = await session.rpc.queue.move_item( + _request_for(rpc.QueueApi.move_item) + ) + results["session.queue.insertAt"] = await session.rpc.queue.insert_at( + _request_for(rpc.QueueApi.insert_at) + ) + results["session.queue.removeAt"] = await session.rpc.queue.remove_at( + _request_for(rpc.QueueApi.remove_at) + ) + results["session.queue.updateText"] = await session.rpc.queue.update_text( + _request_for(rpc.QueueApi.update_text) + ) + results["session.queue.duplicateAt"] = await session.rpc.queue.duplicate_at( + _request_for(rpc.QueueApi.duplicate_at) + ) + results["session.queue.setDrainPaused"] = await session.rpc.queue.set_drain_paused( + _request_for(rpc.QueueApi.set_drain_paused) + ) + results["session.queue.sendNow"] = await session.rpc.queue.send_now( + _request_for(rpc.QueueApi.send_now) + ) + results["session.limitPrediction.predict"] = await session.rpc.limit_prediction.predict( + _request_for(rpc.LimitPredictionApi.predict) + ) + + category_requests = [ + rpc.CatalogSearchRequest( + contract=rpc.CatalogClientContract( + protocol_version=3, + required_capabilities=["catalog-search"], + ), + query="all candidates", + kinds=None, + limit=5, + ), + rpc.CatalogSearchRequest( + contract=rpc.CatalogClientContract( + protocol_version=3, + required_capabilities=["catalog-search"], + ), + query="MCP candidates", + kinds=[rpc.CatalogCandidateKind.MCP_SERVER], + limit=5, + ), + rpc.CatalogSearchRequest( + contract=rpc.CatalogClientContract( + protocol_version=3, + required_capabilities=["catalog-search"], + ), + query="skill candidates", + kinds=[rpc.CatalogCandidateKind.AI_SKILL], + limit=5, + ), + ] + for category_request in category_requests: + category_result = await client.rpc.catalog.search(category_request) + assert isinstance(category_result, rpc.CatalogSearchSucceeded) + assert category_result.search_id == "sample-value" + + error_request = _request_for(rpc.ServerCatalogApi.search) + error_request.query = "raise-jsonrpc-error" + with pytest.raises(JsonRpcError) as exc_info: + await client.rpc.catalog.search(error_request) + assert exc_info.value.code == -32077 + assert exc_info.value.message == "deterministic catalog failure" + assert exc_info.value.data == {"retryable": False, "source": "fake-cli"} + + for rpc_method, payload in responses.items(): + if rpc_method not in _OBJECTIVE_METHODS: + _assert_result_matches_payload(results[rpc_method], payload) + + planned = results["mcp.planInstall"] + assert isinstance(planned, rpc.MCPPlanInstallPlanned) + assert planned.plan.transport_choices[0].transport + assert planned.plan.transport_choices[0].required_values[0].key == "sample-value" + + catalog = results["catalog.search"] + assert isinstance(catalog, rpc.CatalogSearchSucceeded) + assert catalog.candidates[0].installability.value + assert catalog.candidates[0].provenance.authority == "sample-value" + + factory_detail = results["session.factory.getRunDetail"] + assert factory_detail.consumed.active_ms == 7 + assert factory_detail.agents[0].agent_id == "sample-value" + assert factory_detail.progress.records[0].seq == 7 + + permission_requests = results["session.permissions.pendingRequests"] + assert len(permission_requests.items) == 1 + assert permission_requests.items[0].request_id == "sample-value" + assert isinstance( + permission_requests.items[0].request, + generated_session_events.PermissionPromptRequestCommands, + ) + assert permission_requests.items[0].request.kind == "commands" + assert permission_requests.items[0].request.command_identifiers == ["sample-value"] + + mcp_resources = results["session.mcp.resources.read"] + assert mcp_resources.contents[0].uri == "sample-value" + assert mcp_resources.contents[0].mime_type == "sample-value" + + tool_result = typing.cast(dict[str, typing.Any], results["session.tools.execute"]) + assert tool_result["resultType"] == "denied" + assert tool_result["binaryResultsForLlm"][0]["metadata"]["key"]["sample"] == "value" + assert tool_result["taskCompletionDecision"]["reviewerResultMeta"]["sample"] == "value" + + content_checks = results["session.contentExclusion.checkPaths"] + assert content_checks.available is True + assert content_checks.checks[0].excluded is True + + _assert_request_serialization(capture_path) + captured_catalog_params = [ + request["params"] + for request in json.loads(capture_path.read_text(encoding="utf-8"))["requests"] + if request["method"] == "catalog.search" + ] + for category_request in category_requests: + assert category_request.to_dict() in captured_catalog_params + finally: + await client.force_stop() + cli_path.unlink(missing_ok=True) + capture_path.unlink(missing_ok=True) + responses_path.unlink(missing_ok=True) diff --git a/python/e2e/test_scenario_canvas_e2e.py b/python/e2e/test_scenario_canvas_e2e.py new file mode 100644 index 0000000000..f0fee17c4a --- /dev/null +++ b/python/e2e/test_scenario_canvas_e2e.py @@ -0,0 +1,230 @@ +"""Scenario-parity E2Es for canvas provider callback routing.""" + +from __future__ import annotations + +from collections.abc import Sequence + +import pytest + +from copilot import ( + CanvasAction, + CanvasDeclaration, + CanvasError, + CanvasHandler, + OpenCanvasInstance, +) +from copilot.rpc import ( + CanvasProviderCloseRequest, + CanvasProviderInvokeActionRequest, + CanvasProviderOpenRequest, + CanvasProviderOpenResult, +) +from copilot.session import PermissionHandler + +from ._scenario_fake_cli import create_scenario_client, read_scenario_capture +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +def _counter_canvas() -> CanvasDeclaration: + return CanvasDeclaration( + id="counter", + display_name="Counter", + description="Scenario counter canvas", + input_schema={"type": "object"}, + actions=[ + CanvasAction( + name="increment", + description="Increment the counter", + input_schema={"type": "object"}, + ) + ], + ) + + +class _ScenarioCanvasHandler(CanvasHandler): + def __init__(self, failing_operation: str | None = None, *, structured: bool = True) -> None: + self.failing_operation = failing_operation + self.structured = structured + self.open_calls: list[CanvasProviderOpenRequest] = [] + self.action_calls: list[CanvasProviderInvokeActionRequest] = [] + self.close_calls: list[CanvasProviderCloseRequest] = [] + + def _fail_if_requested(self, operation: str) -> None: + if self.failing_operation != operation: + return + if self.structured: + raise CanvasError("scenario_canvas_error", f"{operation} failed") + raise RuntimeError(f"{operation} failed unexpectedly") + + async def on_open(self, ctx: CanvasProviderOpenRequest) -> CanvasProviderOpenResult: + self.open_calls.append(ctx) + self._fail_if_requested("open") + return CanvasProviderOpenResult( + status="ready", + title="Scenario Counter", + url="https://example.test/scenario-counter", + ) + + async def on_action(self, ctx: CanvasProviderInvokeActionRequest) -> dict[str, int]: + self.action_calls.append(ctx) + self._fail_if_requested("action") + return {"newValue": 42} + + async def on_close(self, ctx: CanvasProviderCloseRequest) -> None: + self.close_calls.append(ctx) + self._fail_if_requested("close") + + +def _operation_calls( + handler: _ScenarioCanvasHandler, + operation: str, +) -> Sequence[ + CanvasProviderOpenRequest | CanvasProviderInvokeActionRequest | CanvasProviderCloseRequest +]: + if operation == "open": + return handler.open_calls + if operation == "action": + return handler.action_calls + return handler.close_calls + + +class TestScenarioCanvas: + @pytest.mark.parametrize("operation", ["open", "action", "close"]) + async def test_should_preserve_structured_canvas_error_envelope( + self, + ctx: E2ETestContext, + operation: str, + ): + client, capture_path = create_scenario_client(ctx, f"canvas-error-{operation}") + handler = _ScenarioCanvasHandler(operation) + try: + session = await client.create_session( + session_id="scenario-session", + canvases=[_counter_canvas()], + canvas_handler=handler, + on_permission_request=PermissionHandler.approve_all, + ) + try: + calls = _operation_calls(handler, operation) + assert len(calls) == 1 + assert calls[0].session_id == "scenario-session" + assert calls[0].canvas_id == "counter" + assert calls[0].extension_id == "python-scenario-tests" + assert calls[0].instance_id == f"scenario-{operation}" + assert calls[0].host is not None + assert calls[0].host.capabilities is not None + assert calls[0].host.capabilities.canvases is True + + capture = read_scenario_capture(capture_path) + assert capture["callbackResponses"] == [ + { + "jsonrpc": "2.0", + "id": "canvas-callback", + "error": { + "code": -32603, + "message": f"{operation} failed", + "data": { + "code": "scenario_canvas_error", + "message": f"{operation} failed", + }, + }, + } + ] + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_wrap_unexpected_canvas_handler_error( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "canvas-error-open") + handler = _ScenarioCanvasHandler("open", structured=False) + try: + session = await client.create_session( + session_id="scenario-session", + canvases=[_counter_canvas()], + canvas_handler=handler, + on_permission_request=PermissionHandler.approve_all, + ) + try: + capture = read_scenario_capture(capture_path) + assert capture["callbackResponses"][0]["error"] == { + "code": -32603, + "message": "open failed unexpectedly", + "data": { + "code": "canvas_handler_error", + "message": "open failed unexpectedly", + }, + } + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_reattach_canvas_and_route_all_callbacks_before_resume_completes( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "canvas-resume") + first = await client.create_session( + session_id="scenario-session", + canvases=[_counter_canvas()], + canvas_handler=_ScenarioCanvasHandler(), + on_permission_request=PermissionHandler.approve_all, + ) + await first.disconnect() + + handler = _ScenarioCanvasHandler() + resumed = await client.resume_session( + "scenario-session", + canvases=[_counter_canvas()], + canvas_handler=handler, + open_canvases=[ + OpenCanvasInstance( + canvas_id="counter", + extension_id="python-scenario-tests", + instance_id="reattached-counter", + input={"startValue": 3}, + status="ready", + ) + ], + on_permission_request=PermissionHandler.approve_all, + ) + try: + assert len(handler.open_calls) == 1 + assert handler.open_calls[0].input == {"startValue": 7} + assert len(handler.action_calls) == 1 + assert handler.action_calls[0].action_name == "increment" + assert handler.action_calls[0].input == {"amount": 5} + assert len(handler.close_calls) == 1 + assert handler.close_calls[0].instance_id == "scenario-close" + + capture = read_scenario_capture(capture_path) + assert capture["callbackResponses"] == [ + { + "jsonrpc": "2.0", + "id": "resume-open", + "result": { + "status": "ready", + "title": "Scenario Counter", + "url": "https://example.test/scenario-counter", + }, + }, + { + "jsonrpc": "2.0", + "id": "resume-action", + "result": {"newValue": 42}, + }, + { + "jsonrpc": "2.0", + "id": "resume-close", + "result": None, + }, + ] + finally: + await resumed.disconnect() + await client.stop() diff --git a/python/e2e/test_scenario_cloud_e2e.py b/python/e2e/test_scenario_cloud_e2e.py new file mode 100644 index 0000000000..d02e2d8ea6 --- /dev/null +++ b/python/e2e/test_scenario_cloud_e2e.py @@ -0,0 +1,159 @@ +"""Scenario-parity E2Es for cloud connection and remote-steering workflows.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from copilot import CloudSessionOptions, CloudSessionRepository +from copilot.rpc import ConnectRemoteSessionParams, RemoteNotifySteerableChangedRequest +from copilot.session import PermissionHandler +from copilot.session_events import ( + AssistantMessageData, + SessionRemoteSteerableChangedData, + SessionStartData, +) + +from ._scenario_fake_cli import create_scenario_client, read_scenario_capture +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +class TestScenarioCloud: + async def test_should_notify_steerability_before_first_send_without_remote_enable( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "send") + events = [] + steerability_received = asyncio.Event() + + def on_event(event) -> None: + events.append(event) + if isinstance(event.data, SessionRemoteSteerableChangedData): + steerability_received.set() + + try: + session = await client.create_session( + session_id="scenario-session", + on_event=on_event, + on_permission_request=PermissionHandler.approve_all, + ) + try: + await session.rpc.remote.notify_steerable_changed( + RemoteNotifySteerableChangedRequest(remote_steerable=True) + ) + response = await session.send_and_wait("Send the first cloud message.") + await asyncio.wait_for(steerability_received.wait(), timeout=5) + + assert response is not None + assert isinstance(response.data, AssistantMessageData) + assert response.data.content == "scenario response" + remote_event = next( + event + for event in events + if isinstance(event.data, SessionRemoteSteerableChangedData) + ) + assert remote_event.data.remote_steerable is True + + methods = [ + request["method"] for request in read_scenario_capture(capture_path)["requests"] + ] + assert methods.index("session.remote.notifySteerableChanged") < methods.index( + "session.send" + ) + assert "session.remote.enable" not in methods + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_route_first_cloud_event_for_server_assigned_session_id( + self, + ctx: E2ETestContext, + ): + client, _capture_path = create_scenario_client(ctx, "cloud") + events = [] + first_event_received = asyncio.Event() + + def on_event(event) -> None: + events.append(event) + first_event_received.set() + + try: + session = await client.create_session( + cloud=CloudSessionOptions( + repository=CloudSessionRepository( + owner="github", + name="copilot-sdk", + branch="scenario-branch", + ) + ), + on_event=on_event, + on_permission_request=PermissionHandler.approve_all, + ) + try: + await asyncio.wait_for(first_event_received.wait(), timeout=5) + assert session.session_id == "cloud-runtime-session" + assert len(events) == 1 + assert isinstance(events[0].data, SessionStartData) + assert events[0].data.session_id == session.session_id + assert events[0].data.producer == "scenario-fake-cli" + assert events[0].data.remote_steerable is False + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_resume_using_runtime_id_returned_by_cloud_connect( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "cloud-connect") + try: + await client.start() + connection = await client.rpc.sessions.connect( + ConnectRemoteSessionParams(session_id="remote-resource-id") + ) + session = await client.resume_session( + connection.session_id, + on_permission_request=PermissionHandler.approve_all, + ) + try: + assert connection.session_id == "runtime-session-id" + assert session.session_id == "runtime-session-id" + requests = read_scenario_capture(capture_path)["requests"] + resume = next( + request for request in requests if request["method"] == "session.resume" + ) + assert resume["params"]["sessionId"] == "runtime-session-id" + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_expose_cloud_resource_mismatch_before_resume( + self, + ctx: E2ETestContext, + ): + client, _capture_path = create_scenario_client(ctx, "cloud-connect") + try: + await client.start() + connection = await client.rpc.sessions.connect( + ConnectRemoteSessionParams(session_id="remote-resource-id") + ) + + assert connection.session_id == "runtime-session-id" + assert connection.metadata.session_id == "remote-resource-id" + assert connection.metadata.resource_id == "remote-resource-id" + assert connection.metadata.session_id != connection.session_id + assert connection.metadata.repository.owner == "github" + assert connection.metadata.repository.name == "copilot-sdk" + assert connection.metadata.repository.branch == "scenario-branch" + assert connection.metadata.pull_request_number == 42 + assert connection.metadata.state == "running" + assert connection.metadata.summary == "Remote task summary" + finally: + await client.stop() diff --git a/python/e2e/test_scenario_lifecycle_recovery_e2e.py b/python/e2e/test_scenario_lifecycle_recovery_e2e.py new file mode 100644 index 0000000000..ed87781539 --- /dev/null +++ b/python/e2e/test_scenario_lifecycle_recovery_e2e.py @@ -0,0 +1,84 @@ +"""Scenario-parity E2Es for recoverable session setup failures.""" + +from __future__ import annotations + +import pytest + +from copilot.session import PermissionHandler + +from ._scenario_fake_cli import create_scenario_client, read_scenario_capture +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +class TestScenarioLifecycleRecovery: + async def test_should_allow_retry_after_preacceptance_session_not_found( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "resume-retry") + try: + with pytest.raises(Exception) as exc_info: + await client.resume_session( + "scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + + assert getattr(exc_info.value, "code", None) == -32001 + assert "Session not found before acceptance" in str(exc_info.value) + assert getattr(exc_info.value, "data", None) == {"recoverable": True} + + session = await client.resume_session( + "scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + assert session.session_id == "scenario-session" + assert await session.send("Retry succeeded.") == "user-message" + + requests = read_scenario_capture(capture_path)["requests"] + resume_requests = [ + request for request in requests if request["method"] == "session.resume" + ] + assert len(resume_requests) == 2 + assert all( + request["params"]["sessionId"] == "scenario-session" + for request in resume_requests + ) + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_retry_resume_on_replacement_client_after_recoverable_failure( + self, + ctx: E2ETestContext, + ): + failed_client, _failed_capture = create_scenario_client(ctx, "resume-fail") + with pytest.raises(Exception) as exc_info: + await failed_client.resume_session( + "scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + assert getattr(exc_info.value, "data", None) == {"recoverable": True} + await failed_client.stop() + + replacement_client, replacement_capture = create_scenario_client(ctx, "send") + try: + session = await replacement_client.resume_session( + "scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + assert await session.send("Replacement client recovered.") == "user-message" + requests = read_scenario_capture(replacement_capture)["requests"] + assert [request["method"] for request in requests] == [ + "connect", + "session.resume", + "session.send", + ] + finally: + await session.disconnect() + finally: + await replacement_client.stop() diff --git a/python/e2e/test_scenario_sends_e2e.py b/python/e2e/test_scenario_sends_e2e.py new file mode 100644 index 0000000000..6b9cf77e3d --- /dev/null +++ b/python/e2e/test_scenario_sends_e2e.py @@ -0,0 +1,180 @@ +"""Scenario-parity E2Es for send serialization and cancellation boundaries.""" + +from __future__ import annotations + +import asyncio +from typing import Literal + +import pytest + +from copilot.session import AgentMessageSource, PermissionHandler + +from ._scenario_fake_cli import create_scenario_client, read_scenario_capture +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +def _send_requests(capture_path) -> list[dict]: + return [ + request + for request in read_scenario_capture(capture_path)["requests"] + if request["method"] == "session.send" + ] + + +class TestScenarioSends: + async def test_should_send_complete_message_wire_shape( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "send") + try: + session = await client.create_session( + session_id="scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + message_id = await session.send( + "Run the extension workflow.", + attachments=[ + { + "type": "selection", + "filePath": "src/example.py", + "displayName": "example.py:4-6", + "selection": { + "start": {"line": 4, "character": 2}, + "end": {"line": 6, "character": 8}, + }, + "text": "selected text", + }, + { + "type": "extension_context", + "capturedAt": "2026-01-02T03:04:05.000Z", + "extensionId": "scenario-extension", + "title": "Scenario context", + "canvasId": "scenario-canvas", + "instanceId": "scenario-instance", + "payload": { + "metadata": { + "source": "scenario", + "priority": 7, + } + }, + }, + ], + source=AgentMessageSource("extension-agent"), + mode="immediate", + agent_mode="plan", + request_headers={"X-Scenario": "complete-wire-shape"}, + display_prompt="Visible extension prompt", + ) + + assert message_id == "user-message" + assert _send_requests(capture_path) == [ + { + "method": "session.send", + "params": { + "sessionId": "scenario-session", + "prompt": "Run the extension workflow.", + "attachments": [ + { + "type": "selection", + "filePath": "src/example.py", + "displayName": "example.py:4-6", + "selection": { + "start": {"line": 4, "character": 2}, + "end": {"line": 6, "character": 8}, + }, + "text": "selected text", + }, + { + "type": "extension_context", + "capturedAt": "2026-01-02T03:04:05.000Z", + "extensionId": "scenario-extension", + "title": "Scenario context", + "canvasId": "scenario-canvas", + "instanceId": "scenario-instance", + "payload": { + "metadata": { + "source": "scenario", + "priority": 7, + } + }, + }, + ], + "source": "agent-extension-agent", + "mode": "immediate", + "agentMode": "plan", + "requestHeaders": { + "X-Scenario": "complete-wire-shape", + }, + "displayPrompt": "Visible extension prompt", + }, + } + ] + finally: + await session.disconnect() + finally: + await client.stop() + + @pytest.mark.parametrize("mode", [None, "enqueue", "immediate"]) + async def test_should_not_dispatch_pre_cancelled_send( + self, + ctx: E2ETestContext, + mode: Literal["enqueue", "immediate"] | None, + ): + client, capture_path = create_scenario_client(ctx, "send") + try: + session = await client.create_session( + session_id="scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + send_task = asyncio.create_task( + session.send("This must not be dispatched.", mode=mode) + ) + assert send_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await send_task + + assert _send_requests(capture_path) == [] + finally: + await session.disconnect() + finally: + await client.stop() + + @pytest.mark.parametrize("mode", [None, "enqueue", "immediate"]) + async def test_should_not_replay_send_after_ambiguous_transport_loss( + self, + ctx: E2ETestContext, + mode: Literal["enqueue", "immediate"] | None, + ): + client, capture_path = create_scenario_client(ctx, "send-fail") + session = await client.create_session( + session_id="scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + with pytest.raises(Exception) as exc_info: + await session.send("Lose the transport after accepting this.", mode=mode) + + assert type(exc_info.value).__name__ == "ProcessExitedError" + requests = _send_requests(capture_path) + assert len(requests) == 1 + assert requests[0]["params"]["prompt"] == ("Lose the transport after accepting this.") + if mode is None: + assert "mode" not in requests[0]["params"] + else: + assert requests[0]["params"]["mode"] == mode + + with pytest.raises(Exception) as retry_exc_info: + await asyncio.wait_for( + session.send("Fail immediately after transport loss."), + timeout=1, + ) + assert type(retry_exc_info.value).__name__ == "ProcessExitedError" + assert len(_send_requests(capture_path)) == 1 + finally: + await client.force_stop() diff --git a/python/e2e/test_scenario_session_setup_e2e.py b/python/e2e/test_scenario_session_setup_e2e.py new file mode 100644 index 0000000000..4c6df76cf4 --- /dev/null +++ b/python/e2e/test_scenario_session_setup_e2e.py @@ -0,0 +1,71 @@ +"""Scenario-parity E2Es for session setup ordering.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from copilot.session import PermissionHandler +from copilot.session_events import SessionStartData + +from ._scenario_fake_cli import create_scenario_client, read_scenario_capture +from .testharness import E2ETestContext + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +class TestScenarioSessionSetup: + async def test_should_route_first_subscribed_event_for_preallocated_session_id( + self, + ctx: E2ETestContext, + ): + client, _capture_path = create_scenario_client(ctx, "preallocated-event") + events = [] + event_received = asyncio.Event() + + def on_event(event) -> None: + events.append(event) + event_received.set() + + try: + session = await client.create_session( + session_id="scenario-session", + on_event=on_event, + on_permission_request=PermissionHandler.approve_all, + ) + try: + await asyncio.wait_for(event_received.wait(), timeout=5) + assert session.session_id == "scenario-session" + assert len(events) == 1 + assert isinstance(events[0].data, SessionStartData) + assert events[0].data.session_id == "scenario-session" + assert events[0].data.producer == "scenario-fake-cli" + finally: + await session.disconnect() + finally: + await client.stop() + + async def test_should_create_then_reload_mcp_in_order( + self, + ctx: E2ETestContext, + ): + client, capture_path = create_scenario_client(ctx, "send") + try: + session = await client.create_session( + session_id="scenario-session", + on_permission_request=PermissionHandler.approve_all, + ) + try: + await session.rpc.mcp.reload() + assert [ + request["method"] for request in read_scenario_capture(capture_path)["requests"] + ] == [ + "connect", + "session.create", + "session.mcp.reload", + ] + finally: + await session.disconnect() + finally: + await client.stop() diff --git a/python/test_jsonrpc.py b/python/test_jsonrpc.py index 2f2ecafce9..5e4a4d05be 100644 --- a/python/test_jsonrpc.py +++ b/python/test_jsonrpc.py @@ -29,6 +29,24 @@ def poll(self): return self.returncode +@pytest.mark.asyncio +async def test_send_message_supports_streams_without_process_poll(): + class StreamProcess: + def __init__(self): + self.stdin = io.BytesIO() + self.stdout = io.BytesIO() + self.stderr = None + + process = StreamProcess() + client = JsonRpcClient(process) + + await client._send_message({"jsonrpc": "2.0", "method": "ping"}) + + assert process.stdin.getvalue() == ( + b'Content-Length: 33\r\n\r\n{"jsonrpc":"2.0","method":"ping"}' + ) + + class ShortReadStream: """ Mock stream that simulates short reads from a pipe. diff --git a/rust/src/types.rs b/rust/src/types.rs index 131dd24570..b3dd113948 100644 --- a/rust/src/types.rs +++ b/rust/src/types.rs @@ -5155,6 +5155,25 @@ pub enum Attachment { #[serde(skip_serializing_if = "Option::is_none")] display_name: Option, }, + /// Context captured from an extension-owned canvas. + #[serde(rename = "extension_context")] + ExtensionContext { + /// ISO 8601 timestamp when the context was captured. + captured_at: String, + /// Extension that owns the canvas. + extension_id: String, + /// Canvas declaration identifier when the context is bound to a canvas. + #[serde(skip_serializing_if = "Option::is_none")] + canvas_id: Option, + /// Open canvas instance identifier when the context is bound to a canvas. + #[serde(skip_serializing_if = "Option::is_none")] + instance_id: Option, + /// Human-readable context title. + title: String, + /// Extension-defined structured context payload. + #[serde(skip_serializing_if = "Option::is_none")] + payload: Option, + }, /// A reference to a GitHub issue, PR, or discussion. #[serde(rename = "github_reference")] GitHubReference { @@ -5299,7 +5318,8 @@ impl Attachment { | Self::GitHubTreeComparison { .. } | Self::GitHubUrl { .. } | Self::GitHubFile { .. } - | Self::GitHubSnippet { .. } => None, + | Self::GitHubSnippet { .. } + | Self::ExtensionContext { .. } => None, } } @@ -5319,6 +5339,9 @@ impl Attachment { } else { title.trim().to_string() }), + Self::ExtensionContext { title, .. } if !title.trim().is_empty() => { + Some(title.trim().to_string()) + } _ => self.derived_display_name(), } } @@ -5351,7 +5374,8 @@ impl Attachment { | Self::GitHubTreeComparison { .. } | Self::GitHubUrl { .. } | Self::GitHubFile { .. } - | Self::GitHubSnippet { .. } => {} + | Self::GitHubSnippet { .. } + | Self::ExtensionContext { .. } => {} } } @@ -5371,7 +5395,8 @@ impl Attachment { | Self::GitHubTreeComparison { .. } | Self::GitHubUrl { .. } | Self::GitHubFile { .. } - | Self::GitHubSnippet { .. } => None, + | Self::GitHubSnippet { .. } + | Self::ExtensionContext { .. } => None, } } } @@ -7896,11 +7921,17 @@ mod tests { "referenceType": "issue", "state": "open", "url": "https://github.com/example/repo/issues/42" + }, + { + "type": "extension_context", + "capturedAt": "2026-09-18T11:00:00Z", + "extensionId": "example:extension", + "title": "Unbound context" } ])) .expect("attachments should deserialize"); - assert_eq!(attachments.len(), 5); + assert_eq!(attachments.len(), 6); assert!(matches!( &attachments[0], Attachment::File { @@ -7947,6 +7978,28 @@ mod tests { && state == "open" && url == "https://github.com/example/repo/issues/42" )); + assert!(matches!( + &attachments[5], + Attachment::ExtensionContext { + captured_at, + extension_id, + canvas_id: None, + instance_id: None, + title, + payload: None, + } if captured_at == "2026-09-18T11:00:00Z" + && extension_id == "example:extension" + && title == "Unbound context" + )); + assert_eq!( + serde_json::to_value(&attachments[5]).expect("serialize extension context"), + json!({ + "type": "extension_context", + "capturedAt": "2026-09-18T11:00:00Z", + "extensionId": "example:extension", + "title": "Unbound context" + }) + ); } #[test] diff --git a/rust/tests/e2e.rs b/rust/tests/e2e.rs index cae2f92eb9..cb11464366 100644 --- a/rust/tests/e2e.rs +++ b/rust/tests/e2e.rs @@ -110,6 +110,8 @@ mod rpc_shell_and_fleet; mod rpc_shell_edge_cases; #[path = "e2e/rpc_shell_user_requested.rs"] mod rpc_shell_user_requested; +#[path = "e2e/rpc_surface_coverage.rs"] +mod rpc_surface_coverage; #[path = "e2e/rpc_tasks_and_handlers.rs"] mod rpc_tasks_and_handlers; #[path = "e2e/rpc_ui_ephemeral_query.rs"] diff --git a/rust/tests/e2e/canvas.rs b/rust/tests/e2e/canvas.rs index a2873eed23..b0fd4c5111 100644 --- a/rust/tests/e2e/canvas.rs +++ b/rust/tests/e2e/canvas.rs @@ -1,28 +1,65 @@ use std::sync::Arc; +use std::time::Duration; use async_trait::async_trait; -use github_copilot_sdk::canvas::{CanvasDeclaration, CanvasHandler, CanvasResult}; +use github_copilot_sdk::ResumeSessionConfig; +use github_copilot_sdk::canvas::{CanvasDeclaration, CanvasError, CanvasHandler, CanvasResult}; use github_copilot_sdk::rpc::{ CanvasAction, CanvasProviderCloseRequest, CanvasProviderInvokeActionRequest, CanvasProviderOpenRequest, CanvasProviderOpenResult, }; -use github_copilot_sdk::types::ExtensionInfo; +use github_copilot_sdk::types::{CanvasProviderIdentity, ExtensionInfo}; use parking_lot::Mutex; use serde_json::{Value, json}; +use tokio::sync::Notify; struct TestCanvasHandler { open_calls: Mutex>, close_calls: Mutex>, action_calls: Mutex>, + callback_order: Mutex>, + open_calls_changed: Notify, + error_operation: Option<&'static str>, } impl TestCanvasHandler { fn new() -> Self { + Self::with_error(None) + } + + fn failing(operation: &'static str) -> Self { + Self::with_error(Some(operation)) + } + + fn with_error(error_operation: Option<&'static str>) -> Self { Self { open_calls: Mutex::new(Vec::new()), close_calls: Mutex::new(Vec::new()), action_calls: Mutex::new(Vec::new()), + callback_order: Mutex::new(Vec::new()), + open_calls_changed: Notify::new(), + error_operation, + } + } + + async fn wait_for_open_calls(&self, count: usize) { + loop { + let changed = self.open_calls_changed.notified(); + if self.open_calls.lock().len() >= count { + return; + } + changed.await; + } + } + + fn fail_if_configured(&self, operation: &'static str) -> CanvasResult<()> { + if self.error_operation == Some(operation) { + return Err(CanvasError::new( + format!("scenario_canvas_{operation}_failed"), + format!("The scenario canvas {operation} operation failed."), + )); } + Ok(()) } } @@ -33,20 +70,39 @@ impl CanvasHandler for TestCanvasHandler { ctx: CanvasProviderOpenRequest, ) -> CanvasResult { self.open_calls.lock().push(ctx.clone()); + self.open_calls_changed.notify_one(); + self.callback_order + .lock() + .push(format!("open:{}", ctx.instance_id)); + self.fail_if_configured("open")?; + let value = ctx + .input + .as_ref() + .and_then(|input| input.get("value")) + .and_then(Value::as_str) + .unwrap_or_default(); Ok(CanvasProviderOpenResult { url: Some(format!("https://example.com/counter/{}", ctx.instance_id)), - title: Some(format!("Counter {}", ctx.instance_id)), + title: Some(format!("Counter: {value}")), status: Some("ready".to_string()), }) } async fn on_action(&self, ctx: CanvasProviderInvokeActionRequest) -> CanvasResult { self.action_calls.lock().push(ctx.clone()); - Ok(json!({ "newValue": 42 })) + self.callback_order + .lock() + .push(format!("action:{}:{}", ctx.instance_id, ctx.action_name)); + self.fail_if_configured("action")?; + Ok(ctx.input.unwrap_or(Value::Null)) } async fn on_close(&self, ctx: CanvasProviderCloseRequest) -> CanvasResult<()> { self.close_calls.lock().push(ctx.clone()); + self.callback_order + .lock() + .push(format!("close:{}", ctx.instance_id)); + self.fail_if_configured("close")?; Ok(()) } } @@ -65,6 +121,10 @@ fn canvas_session_config( ctx.approve_all_session_config() .with_request_canvas_renderer(true) .with_extension_info(ExtensionInfo::new("rust-sdk-tests", "canvas-provider")) + .with_canvas_provider( + CanvasProviderIdentity::new("scenario:builtin:rust-canvas") + .with_name("Rust canvas E2E"), + ) .with_canvases([decl]) .with_canvas_handler(handler) } @@ -128,7 +188,7 @@ async fn canvas_open_round_trip() { .expect("open canvas"); assert_eq!(open_result.instance_id, "counter-1"); - assert_eq!(open_result.title.as_deref(), Some("Counter counter-1")); + assert_eq!(open_result.title.as_deref(), Some("Counter: ")); assert_eq!(open_result.status.as_deref(), Some("ready")); assert_eq!( open_result.url.as_deref(), @@ -201,7 +261,7 @@ async fn canvas_invoke_action_round_trip() { .await .expect("invoke action"); - assert_eq!(result.result, Some(json!({ "newValue": 42 }))); + assert_eq!(result.result, Some(json!({ "delta": 1 }))); { let actions = handler.action_calls.lock(); @@ -279,4 +339,319 @@ async fn canvas_close_round_trip() { }) .await; } + +#[tokio::test] +async fn structured_canvas_open_error_surfaces_to_caller() { + super::support::with_dedicated_e2e_context( + "scenario_testing_canvas", + "should_handle_structured_scenario_canvas_error", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let client = ctx.start_client().await; + let handler = Arc::new(TestCanvasHandler::failing("open")); + let session = client + .create_session(canvas_session_config(ctx, handler.clone())) + .await + .expect("create session"); + let canvas = session + .rpc() + .canvas() + .list() + .await + .expect("list canvases") + .canvases + .into_iter() + .next() + .expect("declared canvas"); + + let error = session + .rpc() + .canvas() + .open(github_copilot_sdk::rpc::CanvasOpenRequest { + canvas_id: "counter".to_string(), + instance_id: "counter-error".to_string(), + extension_id: Some(canvas.extension_id), + input: Some(json!({ "value": "before" })), + }) + .await + .expect_err("open error should surface"); + + assert_eq!(error.rpc_code(), Some(-32603)); + assert!( + error + .to_string() + .contains("The scenario canvas open operation failed.") + ); + assert_eq!( + handler.callback_order.lock().as_slice(), + ["open:counter-error"] + ); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn structured_canvas_action_error_surfaces_to_caller() { + super::support::with_dedicated_e2e_context( + "scenario_testing_canvas", + "should_handle_structured_scenario_canvas_error", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let client = ctx.start_client().await; + let handler = Arc::new(TestCanvasHandler::failing("action")); + let session = client + .create_session(canvas_session_config(ctx, handler.clone())) + .await + .expect("create session"); + let canvas = session + .rpc() + .canvas() + .list() + .await + .expect("list canvases") + .canvases + .into_iter() + .next() + .expect("declared canvas"); + session + .rpc() + .canvas() + .open(github_copilot_sdk::rpc::CanvasOpenRequest { + canvas_id: "counter".to_string(), + instance_id: "counter-error".to_string(), + extension_id: Some(canvas.extension_id), + input: Some(json!({ "value": "before" })), + }) + .await + .expect("open canvas"); + + let error = session + .rpc() + .canvas() + .action() + .invoke(github_copilot_sdk::rpc::CanvasActionInvokeRequest { + instance_id: "counter-error".to_string(), + action_name: "increment".to_string(), + input: Some(json!({ "value": "after" })), + }) + .await + .expect_err("action error should surface"); + + assert_eq!(error.rpc_code(), Some(-32603)); + assert!( + error + .to_string() + .contains("The scenario canvas action operation failed.") + ); + assert_eq!( + handler.callback_order.lock().as_slice(), + ["open:counter-error", "action:counter-error:increment"] + ); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn structured_canvas_close_error_is_best_effort() { + super::support::with_dedicated_e2e_context( + "scenario_testing_canvas", + "should_handle_structured_scenario_canvas_error", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let client = ctx.start_client().await; + let handler = Arc::new(TestCanvasHandler::failing("close")); + let session = client + .create_session(canvas_session_config(ctx, handler.clone())) + .await + .expect("create session"); + let canvas = session + .rpc() + .canvas() + .list() + .await + .expect("list canvases") + .canvases + .into_iter() + .next() + .expect("declared canvas"); + session + .rpc() + .canvas() + .open(github_copilot_sdk::rpc::CanvasOpenRequest { + canvas_id: "counter".to_string(), + instance_id: "counter-error".to_string(), + extension_id: Some(canvas.extension_id), + input: Some(json!({ "value": "before" })), + }) + .await + .expect("open canvas"); + + session + .rpc() + .canvas() + .close(github_copilot_sdk::rpc::CanvasCloseRequest { + instance_id: "counter-error".to_string(), + }) + .await + .expect("close remains best effort"); + assert_eq!( + handler.callback_order.lock().as_slice(), + ["open:counter-error", "close:counter-error"] + ); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn resumed_canvas_reattaches_and_routes_all_callbacks() { + super::support::with_dedicated_e2e_context( + "scenario_testing_canvas", + "should_reattach_scenario_canvas_and_route_all_callbacks_after_resume", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let client = ctx.start_client().await; + let original_handler = Arc::new(TestCanvasHandler::new()); + let session = client + .create_session(canvas_session_config(ctx, original_handler)) + .await + .expect("create session"); + let session_id = session.id().clone(); + let canvas = session + .rpc() + .canvas() + .list() + .await + .expect("list canvases") + .canvases + .into_iter() + .next() + .expect("declared canvas"); + session + .rpc() + .canvas() + .open(github_copilot_sdk::rpc::CanvasOpenRequest { + canvas_id: "counter".to_string(), + instance_id: "counter-resume".to_string(), + extension_id: Some(canvas.extension_id), + input: Some(json!({ "value": "persisted" })), + }) + .await + .expect("open canvas"); + let snapshots = session.open_canvases(); + assert_eq!(snapshots.len(), 1); + + session.rpc().suspend().await.expect("suspend session"); + session.stop_event_loop().await; + drop(session); + + let resumed_handler = Arc::new(TestCanvasHandler::new()); + let mut declaration = + CanvasDeclaration::new("counter", "Counter", "Tracks a counter value."); + declaration.actions = Some(vec![CanvasAction { + name: "increment".to_string(), + description: Some("Increments the counter.".to_string()), + input_schema: None, + }]); + let resumed = client + .resume_session( + ResumeSessionConfig::new(session_id.clone()) + .with_github_token(super::support::DEFAULT_TEST_TOKEN) + .approve_all_permissions() + .with_request_canvas_renderer(true) + .with_canvas_provider( + CanvasProviderIdentity::new("scenario:builtin:rust-canvas") + .with_name("Rust canvas E2E"), + ) + .with_canvases([declaration]) + .with_canvas_handler(resumed_handler.clone()) + .with_open_canvases(snapshots), + ) + .await + .expect("resume session"); + + tokio::time::timeout( + Duration::from_secs(10), + resumed_handler.wait_for_open_calls(1), + ) + .await + .expect("reattached canvas open callback"); + { + let opens = resumed_handler.open_calls.lock(); + assert_eq!(opens.len(), 1); + assert_eq!(opens[0].session_id, session_id); + assert_eq!(opens[0].instance_id, "counter-resume"); + assert_eq!(opens[0].input, Some(json!({ "value": "persisted" }))); + } + // The renderer callback precedes the runtime's authoritative opened event. + let mut events = resumed.subscribe(); + tokio::time::timeout(Duration::from_secs(10), async { + while resumed.open_canvases().is_empty() { + events.recv().await.expect("resumed canvas event"); + } + }) + .await + .expect("reattached canvas snapshot"); + let resumed_snapshots = resumed.open_canvases(); + assert_eq!(resumed_snapshots.len(), 1); + assert_eq!(resumed_snapshots[0].instance_id, "counter-resume"); + + let action = resumed + .rpc() + .canvas() + .action() + .invoke(github_copilot_sdk::rpc::CanvasActionInvokeRequest { + instance_id: "counter-resume".to_string(), + action_name: "increment".to_string(), + input: Some(json!({ "value": "resumed" })), + }) + .await + .expect("invoke resumed action"); + assert_eq!(action.result, Some(json!({ "value": "resumed" }))); + resumed + .rpc() + .canvas() + .close(github_copilot_sdk::rpc::CanvasCloseRequest { + instance_id: "counter-resume".to_string(), + }) + .await + .expect("close resumed canvas"); + assert_eq!( + resumed_handler.callback_order.lock().as_slice(), + [ + "open:counter-resume", + "action:counter-resume:increment", + "close:counter-resume" + ] + ); + assert!(resumed.open_canvases().is_empty()); + + resumed + .disconnect() + .await + .expect("disconnect resumed session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} static E2E: super::support::SharedE2eGroup = super::support::SharedE2eGroup::standard("canvas", 4); diff --git a/rust/tests/e2e/client_options.rs b/rust/tests/e2e/client_options.rs index 51880803d3..b39506ec02 100644 --- a/rust/tests/e2e/client_options.rs +++ b/rust/tests/e2e/client_options.rs @@ -1,12 +1,18 @@ use std::collections::HashMap; use std::path::PathBuf; +use std::time::Duration; use github_copilot_sdk::canvas::CanvasDeclaration; -use github_copilot_sdk::rpc::{OpenCanvasInstance, RemoteSessionMode}; -use github_copilot_sdk::session_events::{ReasoningSummary, SessionLimitsConfig}; +use github_copilot_sdk::rpc::{ConnectRemoteSessionParams, OpenCanvasInstance, RemoteSessionMode}; +use github_copilot_sdk::session_events::{ + ReasoningSummary, SessionEventType, SessionLimitsConfig, SessionStartData, +}; use github_copilot_sdk::{ - CliProgram, Client, ClientOptions, CopilotExpAssignmentResponse, ExtensionInfo, ProviderConfig, - ResumeSessionConfig, SessionConfig, SessionId, Transport, + AgentMode, Attachment, AttachmentLineRange, AttachmentSelectionPosition, + AttachmentSelectionRange, CliProgram, Client, ClientOptions, CloudSessionOptions, + CloudSessionRepository, CopilotExpAssignmentResponse, DeliveryMode, ExtensionInfo, + GitHubReferenceType, MessageOptions, MessageSource, ProviderConfig, ResumeSessionConfig, + SessionConfig, SessionId, Transport, }; use serde::Deserialize; use serde_json::{Value, json}; @@ -328,6 +334,323 @@ async fn should_forward_advanced_session_resume_options_to_the_cli() { ); } +#[tokio::test] +async fn should_send_complete_message_wire_shape() { + let fake = FakeCli::new(); + let client = Client::start(fake.client_options("message-wire-client-token")) + .await + .expect("start fake CLI client"); + let session = client + .create_session(SessionConfig::default()) + .await + .expect("create session"); + let file_path = fake.path("message-file").join("scenario.txt"); + let directory_path = fake.path("message-directory"); + let selection_path = fake.path("selection").join("Program.rs"); + + let message_id = session + .send( + MessageOptions::new("Use the hidden scenario context.") + .with_display_prompt("Review selected scenario context") + .with_mode(DeliveryMode::Enqueue) + .with_agent_mode(AgentMode::Interactive) + .with_source(MessageSource::Agent("scenario-client".to_string())) + .with_traceparent("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01") + .with_tracestate("scenario-client=send") + .with_attachments(vec![ + Attachment::File { + path: file_path.clone(), + display_name: Some("scenario.txt".to_string()), + line_range: Some(AttachmentLineRange { start: 3, end: 9 }), + }, + Attachment::Directory { + path: directory_path.clone(), + display_name: Some("message-directory".to_string()), + }, + Attachment::Selection { + file_path: selection_path.clone(), + text: "SCENARIO_SELECTION".to_string(), + display_name: Some("Program.rs".to_string()), + selection: AttachmentSelectionRange { + start: AttachmentSelectionPosition { + line: 17, + character: 0, + }, + end: AttachmentSelectionPosition { + line: 17, + character: 18, + }, + }, + }, + Attachment::GitHubReference { + number: 610, + reference_type: GitHubReferenceType::Pr, + state: "open".to_string(), + title: "Scenario-shaped E2E coverage".to_string(), + url: "https://github.com/github/copilot-sdk/pull/610".to_string(), + }, + Attachment::Blob { + data: "QVBQX0JMT0I=".to_string(), + mime_type: "text/plain".to_string(), + display_name: Some("scenario-wire-blob.txt".to_string()), + }, + Attachment::ExtensionContext { + captured_at: "2026-09-17T20:00:00Z".to_string(), + extension_id: "scenario-client:code-review".to_string(), + canvas_id: Some("diff".to_string()), + instance_id: Some("diff-17".to_string()), + title: "Selected change".to_string(), + payload: Some(json!({ "selection": "SCENARIO_SELECTION", "line": 17 })), + }, + ]), + ) + .await + .expect("send complete message"); + assert_eq!(message_id, "scenario-client-message"); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + + let send = fake.captured_request("session.send"); + let params = send.params.as_object().expect("session.send params"); + assert_json_values( + params, + [ + ("prompt", json!("Use the hidden scenario context.")), + ("displayPrompt", json!("Review selected scenario context")), + ("mode", json!("enqueue")), + ("agentMode", json!("interactive")), + ("source", json!("agent-scenario-client")), + ( + "traceparent", + json!("00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"), + ), + ("tracestate", json!("scenario-client=send")), + ], + ); + let attachments = params["attachments"].as_array().expect("attachments"); + assert_eq!( + attachments + .iter() + .map(|attachment| attachment["type"].as_str().expect("attachment type")) + .collect::>(), + [ + "file", + "directory", + "selection", + "github_reference", + "blob", + "extension_context" + ] + ); + assert_eq!(attachments[0]["path"], json!(path_string(&file_path))); + assert_eq!(attachments[0]["lineRange"], json!({ "start": 3, "end": 9 })); + assert_eq!(attachments[1]["path"], json!(path_string(&directory_path))); + assert_eq!( + attachments[2]["filePath"], + json!(path_string(&selection_path)) + ); + assert_eq!(attachments[2]["text"], json!("SCENARIO_SELECTION")); + assert_eq!(attachments[3]["number"], json!(610)); + assert_eq!(attachments[3]["referenceType"], json!("pr")); + assert_eq!(attachments[4]["data"], json!("QVBQX0JMT0I=")); + assert_eq!(attachments[4]["mimeType"], json!("text/plain")); + assert_eq!( + attachments[5]["extensionId"], + json!("scenario-client:code-review") + ); + assert_eq!( + attachments[5]["payload"]["selection"], + json!("SCENARIO_SELECTION") + ); +} + +#[tokio::test] +async fn dropping_unpolled_send_never_dispatches_for_any_delivery_mode() { + let fake = FakeCli::new(); + let client = Client::start(fake.client_options("cancelled-send-client-token")) + .await + .expect("start fake CLI client"); + let session = client + .create_session(SessionConfig::default()) + .await + .expect("create session"); + + for mode in [ + None, + Some(DeliveryMode::Enqueue), + Some(DeliveryMode::Immediate), + ] { + let mut message = MessageOptions::new("This message must never be invoked.") + .with_display_prompt("Cancelled scenario message") + .with_source(MessageSource::Agent("scenario-client".to_string())); + message.mode = mode; + drop(session.send(message)); + } + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + assert!( + fake.capture() + .requests + .iter() + .all(|request| request.method != "session.send") + ); +} + +#[tokio::test] +async fn transport_loss_never_replays_send_for_any_delivery_mode() { + for (mode, expected_mode) in [ + (None, None), + (Some(DeliveryMode::Enqueue), Some("enqueue")), + (Some(DeliveryMode::Immediate), Some("immediate")), + ] { + let fake = FakeCli::new(); + let client = Client::start( + fake.client_options_with_behavior("ambiguous-send-client-token", "drop-after-send"), + ) + .await + .expect("start fake CLI client"); + let session = client + .create_session(SessionConfig::default()) + .await + .expect("create session"); + let mut message = MessageOptions::new("AMBIGUOUS_SCENARIO_SEND") + .with_display_prompt("Ambiguous scenario send") + .with_source(MessageSource::Agent("scenario-client".to_string())); + message.mode = mode; + + let error = session + .send(message) + .await + .expect_err("transport loss should fail send"); + assert!(error.is_transport_failure()); + client.force_stop(); + + let sends = fake + .capture() + .requests + .into_iter() + .filter(|request| request.method == "session.send") + .collect::>(); + assert_eq!(sends.len(), 1); + assert_eq!( + sends[0].params.get("mode").and_then(Value::as_str), + expected_mode + ); + } +} + +#[tokio::test] +async fn cloud_create_routes_first_event_for_server_assigned_session_id() { + let fake = FakeCli::new(); + let client = Client::start(fake.client_options("cloud-create-client-token")) + .await + .expect("start fake CLI client"); + let prepared = client + .prepare_session( + SessionConfig::default().with_cloud(CloudSessionOptions::with_repository( + CloudSessionRepository::new("github", "copilot-sdk").with_branch("main"), + )), + ) + .expect("prepare cloud session"); + let mut events = prepared.subscribe(); + let session = prepared.start().await.expect("start cloud session"); + + assert_eq!(session.id().as_str(), "server-assigned-cloud-session"); + let event = tokio::time::timeout(Duration::from_secs(5), events.recv()) + .await + .expect("first cloud event timed out") + .expect("cloud event stream closed"); + assert_eq!(event.parsed_type(), SessionEventType::SessionStart); + assert_eq!( + event + .typed_data::() + .expect("session.start data") + .session_id, + session.id().clone() + ); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + let create = fake.captured_request("session.create"); + assert!(create.params.get("sessionId").is_none()); + assert_eq!( + create.params["cloud"]["repository"]["owner"], + json!("github") + ); +} + +#[tokio::test] +async fn remote_connect_runtime_id_is_used_for_resume() { + let fake = FakeCli::new(); + let client = Client::start(fake.client_options("cloud-connect-client-token")) + .await + .expect("start fake CLI client"); + + let connection = client + .rpc() + .sessions() + .connect(ConnectRemoteSessionParams { + session_id: SessionId::from("cloud-control-session"), + }) + .await + .expect("connect remote session"); + assert_eq!(connection.session_id.as_str(), "runtime-session-id"); + assert_eq!(connection.metadata.session_id, connection.session_id); + assert_eq!( + connection.metadata.resource_id.as_deref(), + Some("github/copilot-sdk#123") + ); + let resumed = client + .resume_session(ResumeSessionConfig::new(connection.session_id.clone())) + .await + .expect("resume connected runtime session"); + assert_eq!(resumed.id(), &connection.session_id); + + resumed.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + assert_eq!( + fake.captured_request("sessions.connect").params["sessionId"], + json!("cloud-control-session") + ); + assert_eq!( + fake.captured_request("session.resume").params["sessionId"], + json!("runtime-session-id") + ); +} + +#[tokio::test] +async fn remote_resource_mismatch_is_observable_before_resume() { + let fake = FakeCli::new(); + let client = Client::start( + fake.client_options_with_behavior("cloud-mismatch-client-token", "resource-mismatch"), + ) + .await + .expect("start fake CLI client"); + + let connection = client + .rpc() + .sessions() + .connect(ConnectRemoteSessionParams { + session_id: SessionId::from("cloud-control-session"), + }) + .await + .expect("connect remote session"); + assert_ne!( + connection.metadata.resource_id.as_deref(), + Some("github/copilot-sdk#123") + ); + + client.stop().await.expect("stop client"); + assert!( + fake.capture() + .requests + .iter() + .all(|request| request.method != "session.resume") + ); +} + struct FakeCli { _dir: TempDir, script_path: PathBuf, @@ -352,6 +675,10 @@ impl FakeCli { } fn client_options(&self, token: &str) -> ClientOptions { + self.client_options_with_behavior(token, "normal") + } + + fn client_options_with_behavior(&self, token: &str, behavior: &str) -> ClientOptions { ClientOptions::new() .with_program(CliProgram::Path(PathBuf::from("node"))) .with_prefix_args([self.script_path.as_os_str().to_owned()]) @@ -359,6 +686,8 @@ impl FakeCli { .with_extra_args([ "--capture-file".to_string(), self.capture_path.to_string_lossy().into_owned(), + "--behavior".to_string(), + behavior.to_string(), ]) .with_github_token(token) .with_use_logged_in_user(false) @@ -421,6 +750,8 @@ const fs = require("fs"); const captureIndex = process.argv.indexOf("--capture-file"); const captureFile = captureIndex >= 0 ? process.argv[captureIndex + 1] : undefined; +const behaviorIndex = process.argv.indexOf("--behavior"); +const behavior = behaviorIndex >= 0 ? process.argv[behaviorIndex + 1] : "normal"; const requests = []; function saveCapture() { @@ -478,8 +809,57 @@ function handleMessage(message) { return; } if (message.method === "session.create") { - const sessionId = (message.params && message.params.sessionId) || "fake-session"; + const isCloud = Boolean(message.params && message.params.cloud); + const sessionId = (message.params && message.params.sessionId) + || (isCloud ? "server-assigned-cloud-session" : "fake-session"); writeResponse(message.id, { sessionId, workspacePath: null, capabilities: null }); + if (isCloud) { + writeMessage({ + jsonrpc: "2.0", + method: "session.event", + params: { + sessionId, + event: { + id: "cloud-start-event", + timestamp: "2026-09-18T00:00:00Z", + parentId: null, + type: "session.start", + data: { + sessionId, + version: 1, + producer: "fake-cli", + copilotVersion: "fake", + startTime: "2026-09-18T00:00:00Z", + } + } + } + }); + } + return; + } + if (message.method === "sessions.connect") { + const resourceId = behavior === "resource-mismatch" + ? "github/other-repository#456" + : "github/copilot-sdk#123"; + writeResponse(message.id, { + sessionId: "runtime-session-id", + metadata: { + kind: "coding_agent", + modifiedTime: "2026-09-18T00:00:00Z", + repository: { owner: "github", name: "copilot-sdk", branch: "main" }, + resourceId, + sessionId: "runtime-session-id", + startTime: "2026-09-18T00:00:00Z" + } + }); + return; + } + if (message.method === "session.send") { + if (behavior === "drop-after-send") { + process.exit(0); + return; + } + writeResponse(message.id, { messageId: "scenario-client-message" }); return; } if (message.method === "session.resume") { @@ -499,7 +879,11 @@ function handleMessage(message) { } function writeResponse(id, result) { - const body = JSON.stringify({ jsonrpc: "2.0", id, result }); + writeMessage({ jsonrpc: "2.0", id, result }); +} + +function writeMessage(message) { + const body = JSON.stringify(message); process.stdout.write("Content-Length: " + Buffer.byteLength(body, "utf8") + "\r\n\r\n" + body); } "#; diff --git a/rust/tests/e2e/event_fidelity.rs b/rust/tests/e2e/event_fidelity.rs index 7176a7e669..eab8e622e0 100644 --- a/rust/tests/e2e/event_fidelity.rs +++ b/rust/tests/e2e/event_fidelity.rs @@ -1,9 +1,18 @@ +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use async_trait::async_trait; use github_copilot_sdk::session_events::{ AssistantMessageData, AssistantUsageData, SessionEventType, SessionUsageInfoData, - ToolExecutionCompleteData, ToolExecutionStartData, UserMessageData, + ToolExecutionCompleteData, ToolExecutionStartData, UserMessageData, UserMessageDelivery, +}; +use github_copilot_sdk::tool::ToolHandler; +use github_copilot_sdk::{ + DeliveryMode, Error, MessageOptions, MessageSource, Tool, ToolInvocation, ToolResult, }; +use tokio::sync::{Mutex, mpsc}; -use super::support::{collect_until_idle, event_types}; +use super::support::{collect_until_idle, event_types, recv_with_timeout, wait_for_event}; #[tokio::test] async fn should_include_valid_fields_on_all_events() { @@ -374,5 +383,193 @@ async fn should_preserve_message_order_in_getmessages_after_tool_use() { ) .await; } + +#[tokio::test] +async fn should_order_idle_queued_and_immediate_delivery_while_busy() { + super::support::with_dedicated_e2e_context( + "scenario_testing_sends", + "should_order_idle_queued_and_immediate_scenario_delivery", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let (started_tx, mut started_rx) = mpsc::unbounded_channel(); + let (release_tx, release_rx) = mpsc::channel(2); + let client = ctx.start_client().await; + let session = client + .create_session( + ctx.approve_all_session_config().with_tools(vec![ + Tool::new("scenario_send_blocker") + .with_description("Blocks the active turn until released") + .with_handler(Arc::new(SequencedBlockingTool { + started_tx, + release_rx: Mutex::new(release_rx), + invocation_count: AtomicUsize::new(0), + })), + ]), + ) + .await + .expect("create session"); + + let idle_enqueue = tokio::spawn(wait_for_event( + session.subscribe(), + "idle enqueue completion", + |event| event.parsed_type() == SessionEventType::SessionIdle, + )); + let idle_enqueue_id = session + .send( + MessageOptions::new("Reply with exactly IDLE_ENQUEUE.") + .with_mode(DeliveryMode::Enqueue) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("send idle enqueue"); + idle_enqueue.await.expect("idle enqueue task"); + + let idle_immediate = tokio::spawn(wait_for_event( + session.subscribe(), + "idle immediate completion", + |event| event.parsed_type() == SessionEventType::SessionIdle, + )); + let idle_immediate_id = session + .send( + MessageOptions::new("Reply with exactly IDLE_IMMEDIATE.") + .with_mode(DeliveryMode::Immediate) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("send idle immediate"); + idle_immediate.await.expect("idle immediate task"); + + session + .send( + MessageOptions::new( + "Call scenario_send_blocker, then reply with its result.", + ) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("start blocking turn"); + assert_eq!( + recv_with_timeout(&mut started_rx, "first blocker invocation").await, + 1 + ); + + let steering_id = session + .send( + MessageOptions::new( + "Call scenario_send_blocker again, then reply with exactly FIRST_STEERING.", + ) + .with_mode(DeliveryMode::Immediate) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("send steering message"); + release_tx + .send("SCENARIO_SEND_BLOCKER_RELEASED".to_string()) + .await + .expect("release first blocker"); + assert_eq!( + recv_with_timeout(&mut started_rx, "second blocker invocation").await, + 2 + ); + + let second_immediate_id = session + .send( + MessageOptions::new("Reply with exactly SECOND_IMMEDIATE.") + .with_mode(DeliveryMode::Immediate) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("send second immediate"); + let queued_id = session + .send( + MessageOptions::new("Reply with exactly FINAL_QUEUED.") + .with_mode(DeliveryMode::Enqueue) + .with_source(MessageSource::Agent("scenario-client".to_string())), + ) + .await + .expect("send queued message"); + let final_queued = tokio::spawn(wait_for_event( + session.subscribe(), + "final queued response", + |event| { + event.parsed_type() == SessionEventType::AssistantMessage + && event + .typed_data::() + .is_some_and(|data| data.content.contains("FINAL_QUEUED")) + }, + )); + release_tx + .send("SCENARIO_SEND_BLOCKER_RELEASED_AGAIN".to_string()) + .await + .expect("release second blocker"); + final_queued.await.expect("final queued task"); + + let events = session.get_events().await.expect("get events"); + let messages = events + .iter() + .filter_map(|event| { + (event.parsed_type() == SessionEventType::UserMessage) + .then(|| event.typed_data::()) + .flatten() + }) + .collect::>(); + let find = |id: &str| { + messages + .iter() + .find(|message| message.message_id.as_deref() == Some(id)) + .expect("user message by id") + }; + assert_eq!( + find(&idle_enqueue_id).delivery, + Some(UserMessageDelivery::Idle) + ); + assert_eq!( + find(&idle_immediate_id).delivery, + Some(UserMessageDelivery::Idle) + ); + assert_eq!( + find(&steering_id).delivery, + Some(UserMessageDelivery::Steering) + ); + assert_eq!( + find(&second_immediate_id).delivery, + Some(UserMessageDelivery::Steering) + ); + assert_eq!(find(&queued_id).delivery, Some(UserMessageDelivery::Queued)); + + let position = |id: &str| { + messages + .iter() + .position(|message| message.message_id.as_deref() == Some(id)) + .expect("message position") + }; + assert!(position(&steering_id) < position(&second_immediate_id)); + assert!(position(&second_immediate_id) < position(&queued_id)); + + session.disconnect().await.expect("disconnect session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} + +struct SequencedBlockingTool { + started_tx: mpsc::UnboundedSender, + release_rx: Mutex>, + invocation_count: AtomicUsize, +} + +#[async_trait] +impl ToolHandler for SequencedBlockingTool { + async fn call(&self, _invocation: ToolInvocation) -> Result { + let mut release_rx = self.release_rx.lock().await; + let invocation = self.invocation_count.fetch_add(1, Ordering::SeqCst) + 1; + let _ = self.started_tx.send(invocation); + let result = release_rx.recv().await.expect("tool release value"); + Ok(ToolResult::Text(result)) + } +} static E2E: super::support::SharedE2eGroup = super::support::SharedE2eGroup::standard("event_fidelity", 8); diff --git a/rust/tests/e2e/rpc_surface_coverage.rs b/rust/tests/e2e/rpc_surface_coverage.rs new file mode 100644 index 0000000000..fda704eca6 --- /dev/null +++ b/rust/tests/e2e/rpc_surface_coverage.rs @@ -0,0 +1,1491 @@ +use std::collections::HashMap; +use std::path::PathBuf; + +use github_copilot_sdk::rpc::*; +use github_copilot_sdk::{CliProgram, Client, ClientOptions, ErrorKind, SessionConfig, Transport}; +use serde::Serialize; +use serde_json::{Value, json}; +use tempfile::TempDir; + +macro_rules! rpc_ok { + ($call:expr) => { + $call + .await + .unwrap_or_else(|error| panic!("{} failed: {error}", stringify!($call))) + }; +} + +#[tokio::test] +async fn client_rpc_surface_uses_typed_namespaces_and_round_trips_results() { + let mut results = ResponseMap::default(); + results.insert( + "account.getQuota", + AccountGetQuotaResult { + quota_snapshots: HashMap::from([( + "premium_interactions".to_string(), + AccountQuotaSnapshot { + entitlement_requests: 300, + remaining_percentage: 62.5, + reset_date: Some("2026-10-01T00:00:00Z".to_string()), + used_requests: 112, + ..Default::default() + }, + )]), + }, + ); + results.insert( + "catalog.search", + CatalogSearchResult::Succeeded(CatalogSearchSucceeded { + search_id: "search-17".to_string(), + truncated: true, + ..Default::default() + }), + ); + results.insert_default::("extensions.discover"); + results.insert_default::("hooks.discover"); + results.insert_default::("llmInference.setProvider"); + results.insert_default::("managedSettings.read"); + results.insert( + "mcp.planInstall", + McpPlanInstallResult::Unavailable(CatalogUnavailableError { + message: "catalog intentionally offline".to_string(), + ..Default::default() + }), + ); + results.insert_default::("models.list"); + results.insert_default::("models.getBuiltInCatalog"); + results.insert_default::("plugins.marketplaces.refresh"); + results.insert_value( + "sessions.getClientMetadata", + json!([{"clientName": "fixture-host", "capabilities": ["rpc"]}]), + ); + results.insert_default::("sessions.readPersistedEvents"); + + let fake = FakeCli::new(results, ErrorMap::default()); + let client = fake.start_client().await; + + rpc_ok!(client.rpc().register_extension_launch_provider()); + let quota = rpc_ok!( + client + .rpc() + .account() + .get_quota_with_params(AccountGetQuotaRequest { + git_hub_token: Some("quota-token".to_string()), + selection_id: Some("account-42".to_string()), + }) + ); + let snapshot = quota + .quota_snapshots + .get("premium_interactions") + .expect("premium quota snapshot"); + assert_eq!(snapshot.entitlement_requests, 300); + assert_eq!(snapshot.used_requests, 112); + assert_eq!(snapshot.remaining_percentage, 62.5); + assert_eq!(snapshot.reset_date.as_deref(), Some("2026-10-01T00:00:00Z")); + + let search = rpc_ok!(client.rpc().catalog().search(CatalogSearchRequest { + contract: CatalogClientContract { + protocol_version: 7, + required_capabilities: vec!["install-plans".to_string()], + }, + kinds: None, + limit: Some(4), + page: None, + query: "offline catalog".to_string(), + })); + let CatalogSearchResult::Succeeded(search) = search else { + panic!("expected successful catalog search"); + }; + assert_eq!(search.search_id, "search-17"); + assert!(search.truncated); + + rpc_ok!(client.rpc().extensions().discover()); + rpc_ok!(client.rpc().hooks().discover(HooksDiscoverRequest { + exclude_host_hooks: Some(true), + project_paths: Some(vec!["project-a".to_string(), "project-b".to_string()]), + })); + rpc_ok!(client.rpc().llm_inference().set_provider()); + rpc_ok!(client.rpc().managed_settings().read()); + + let plan = rpc_ok!(client.rpc().mcp().plan_install(McpPlanInstallRequest { + contract: CatalogClientContract { + protocol_version: 7, + required_capabilities: vec!["install-plans".to_string()], + }, + scope: None, + source: McpPlanInstallSource::Candidate(McpPlanInstallSourceCandidate { + candidate_handle: "candidate-handle".to_string(), + kind: McpPlanInstallSourceCandidateKind::Candidate, + search_id: "search-17".to_string(), + }), + })); + let McpPlanInstallResult::Unavailable(unavailable) = plan else { + panic!("expected typed unavailable install plan"); + }; + assert_eq!(unavailable.message, "catalog intentionally offline"); + + rpc_ok!( + client + .rpc() + .models() + .list_with_params(ModelsListRequest::default()) + ); + rpc_ok!(client.rpc().models().get_built_in_catalog()); + rpc_ok!( + client + .rpc() + .plugins() + .builtin() + .set(PluginsBuiltinSetRequest::default()) + ); + rpc_ok!(client.rpc().plugins().marketplaces().refresh()); + let metadata = rpc_ok!( + client + .rpc() + .sessions() + .get_client_metadata(SessionsGetClientMetadataRequest::default()) + ); + assert_eq!(metadata[0]["clientName"], "fixture-host"); + assert_eq!(metadata[0]["capabilities"][0], "rpc"); + rpc_ok!( + client + .rpc() + .sessions() + .read_persisted_events(SessionsReadPersistedEventsRequest::default()) + ); + rpc_ok!( + client + .rpc() + .skills() + .config() + .set_skill_disabled(SkillsConfigSetSkillDisabledRequest::default()) + ); + + client.stop().await.expect("stop fake CLI"); + + fake.assert_target_methods(&[ + "registerExtensionLaunchProvider", + "account.getQuota", + "catalog.search", + "extensions.discover", + "hooks.discover", + "llmInference.setProvider", + "managedSettings.read", + "mcp.planInstall", + "models.list", + "models.getBuiltInCatalog", + "plugins.builtin.set", + "plugins.marketplaces.refresh", + "sessions.getClientMetadata", + "sessions.readPersistedEvents", + "skills.config.setSkillDisabled", + ]); + fake.assert_params( + "account.getQuota", + 0, + json!({ + "gitHubToken": "quota-token", + "selectionId": "account-42" + }), + ); + fake.assert_params( + "catalog.search", + 0, + json!({ + "contract": { + "protocolVersion": 7, + "requiredCapabilities": ["install-plans"] + }, + "limit": 4, + "query": "offline catalog" + }), + ); + fake.assert_params( + "mcp.planInstall", + 0, + json!({ + "contract": { + "protocolVersion": 7, + "requiredCapabilities": ["install-plans"] + }, + "source": { + "candidateHandle": "candidate-handle", + "kind": "candidate", + "searchId": "search-17" + } + }), + ); +} + +#[tokio::test] +async fn session_lifecycle_factory_and_history_rpc_surface_is_typed() { + let mut results = ResponseMap::default(); + results.insert_default::("session.sendMessages"); + results.insert_default::("session.abort"); + results.insert_default::("session.interruptMainTurn"); + results.insert_value("session.cancelAllBackgroundAgents", json!(3)); + results.insert_default::("session.agent.list"); + results.insert( + "session.autopilotObjective.getState", + AutopilotObjectiveGetStateResult { + state: Some(AutopilotObjectiveState::default()), + }, + ); + results.insert_default::( + "session.completions.getTriggerCharacters", + ); + results + .insert_default::("session.contentExclusion.checkPaths"); + results.insert( + "session.factory.run", + FactoryRunResult { + attempt: Some(2), + result: Some(json!({"artifact": "factory-output"})), + run_id: "run-123".to_string(), + ..Default::default() + }, + ); + results.insert( + "session.factory.resume", + FactoryResumeResult { + factory_name: "coverage-factory".to_string(), + run: FactoryRunResult { + run_id: "run-123".to_string(), + ..Default::default() + }, + }, + ); + results.insert_default::("session.factory.getRun"); + results.insert_default::("session.factory.listRuns"); + results.insert_default::("session.factory.getRunDetail"); + results.insert_default::("session.factory.getRunProgress"); + results.insert_default::("session.factory.cancel"); + results.insert_default::("session.factory.pause"); + results.insert_default::("session.factory.log"); + results.insert_default::("session.factory.agent"); + results.insert_default::("session.factory.journal.get"); + results.insert_default::("session.factory.journal.put"); + results.insert_default::("session.fleet.start"); + results.insert( + "session.history.compact", + HistoryCompactResult { + messages_removed: 8, + success: true, + summary_content: Some("deterministic summary".to_string()), + tokens_removed: 144, + ..Default::default() + }, + ); + results.insert( + "session.history.clearContext", + HistoryClearContextResult { + messages_cleared: 5, + }, + ); + results.insert( + "session.limitPrediction.predict", + SessionLimitPredictionResult::Unavailable(SessionLimitPredictionResultUnavailable { + reason: SessionLimitPredictionUnavailableReason::NoModel, + ..Default::default() + }), + ); + + let fake = FakeCli::new(results, ErrorMap::default()); + let client = fake.start_client().await; + let session = fake.create_session(&client).await; + + rpc_ok!(session.rpc().suspend()); + rpc_ok!(session.rpc().send_messages(SendMessagesRequest::default())); + rpc_ok!(session.rpc().abort(AbortRequest::default())); + rpc_ok!( + session + .rpc() + .interrupt_main_turn(InterruptMainTurnRequest::default()) + ); + let cancelled = rpc_ok!(session.rpc().cancel_all_background_agents()); + assert_eq!(cancelled, 3); + rpc_ok!( + session + .rpc() + .agent() + .list_with_params(AgentListRequest::default()) + ); + rpc_ok!( + session + .rpc() + .agent() + .set_prompt(AgentSetPromptRequest::default()) + ); + let objective = rpc_ok!(session.rpc().autopilot_objective().get_state()); + assert!(objective.state.is_some()); + rpc_ok!(session.rpc().completions().get_trigger_characters()); + rpc_ok!( + session + .rpc() + .content_exclusion() + .check_paths(ContentExclusionCheckPathsRequest { + paths: vec![ + "C:\\workspace\\one.rs".to_string(), + "/workspace/two.rs".to_string() + ], + }) + ); + + let run = rpc_ok!(session.rpc().factory().run(FactoryRunRequest { + args: json!({"mode": "offline", "count": 2}), + name: "coverage-factory".to_string(), + options: None, + })); + assert_eq!(run.run_id, "run-123"); + assert_eq!(run.attempt, Some(2)); + assert_eq!(run.result, Some(json!({"artifact": "factory-output"}))); + let resumed = rpc_ok!(session.rpc().factory().resume(FactoryResumeRequest { + run_id: "run-123".to_string(), + notify_on_complete: Some(false), + ..Default::default() + })); + assert_eq!(resumed.factory_name, "coverage-factory"); + assert_eq!(resumed.run.run_id, "run-123"); + rpc_ok!( + session + .rpc() + .factory() + .get_run(FactoryGetRunRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .list_runs(FactoryListRunsRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .get_run_detail(FactoryGetRunRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .get_run_progress(FactoryGetRunProgressRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .cancel(FactoryCancelRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .pause(FactoryPauseRequest::default()) + ); + rpc_ok!(session.rpc().factory().log(FactoryLogRequest::default())); + rpc_ok!( + session + .rpc() + .factory() + .agent(FactoryAgentRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .journal() + .get(FactoryJournalGetRequest::default()) + ); + rpc_ok!( + session + .rpc() + .factory() + .journal() + .put(FactoryJournalPutRequest::default()) + ); + rpc_ok!(session.rpc().fleet().start(FleetStartRequest::default())); + + let compact = rpc_ok!(session.rpc().history().compact()); + assert!(compact.success); + assert_eq!(compact.messages_removed, 8); + assert_eq!(compact.tokens_removed, 144); + assert_eq!( + compact.summary_content.as_deref(), + Some("deterministic summary") + ); + rpc_ok!( + session + .rpc() + .history() + .compact_with_params(HistoryCompactRequest::default()) + ); + let cleared = rpc_ok!( + session + .rpc() + .history() + .clear_context(HistoryClearContextRequest::default()) + ); + assert_eq!(cleared.messages_cleared, 5); + let prediction = rpc_ok!(session.rpc().limit_prediction().predict()); + assert!(matches!( + prediction, + SessionLimitPredictionResult::Unavailable(SessionLimitPredictionResultUnavailable { + reason: SessionLimitPredictionUnavailableReason::NoModel, + .. + }) + )); + rpc_ok!( + session + .rpc() + .limit_prediction() + .predict_with_params(SessionLimitPredictionRequest { + model_id: Some("fixture-model".to_string()), + ..Default::default() + }) + ); + + session.disconnect().await.expect("disconnect fake session"); + client.stop().await.expect("stop fake CLI"); + + fake.assert_target_methods(&[ + "session.suspend", + "session.sendMessages", + "session.abort", + "session.interruptMainTurn", + "session.cancelAllBackgroundAgents", + "session.agent.list", + "session.agent.setPrompt", + "session.autopilotObjective.getState", + "session.completions.getTriggerCharacters", + "session.contentExclusion.checkPaths", + "session.factory.run", + "session.factory.resume", + "session.factory.getRun", + "session.factory.listRuns", + "session.factory.getRunDetail", + "session.factory.getRunProgress", + "session.factory.cancel", + "session.factory.pause", + "session.factory.log", + "session.factory.agent", + "session.factory.journal.get", + "session.factory.journal.put", + "session.fleet.start", + "session.history.compact", + "session.history.compact", + "session.history.clearContext", + "session.limitPrediction.predict", + "session.limitPrediction.predict", + ]); + fake.assert_params( + "session.factory.run", + 0, + json!({ + "sessionId": "rpc-surface-session", + "args": {"mode": "offline", "count": 2}, + "name": "coverage-factory" + }), + ); + fake.assert_params( + "session.limitPrediction.predict", + 1, + json!({ + "sessionId": "rpc-surface-session", + "modelId": "fixture-model" + }), + ); +} + +#[tokio::test] +async fn session_mcp_metadata_model_and_permission_rpc_surface_is_typed() { + let mut results = ResponseMap::default(); + results + .insert_default::("session.mcp.moveLoadingToBackground"); + results.insert_default::("session.mcp.apps.readResource"); + results.insert( + "session.mcp.oauth.probe", + McpOauthProbeResult::Failed(McpOauthProbeResultFailed { + error: "offline probe fixture".to_string(), + ..Default::default() + }), + ); + results.insert( + "session.mcp.oauth.respond", + McpOauthRespondResult { success: true }, + ); + results.insert_default::("session.mcp.resources.list"); + results + .insert_default::("session.mcp.resources.listTemplates"); + results.insert_value( + "session.metadata.getClientMetadata", + json!({"clientName": "rust-fixture", "version": "1"}), + ); + results.insert_value( + "session.metadata.updateClientMetadata", + json!({"clientName": "updated-fixture", "version": "2"}), + ); + results.insert_default::("session.model.switchAutoTier"); + results.insert_default::("session.model.list"); + results.insert_default::("session.permissions.configure"); + results.insert_default::("session.permissions.pendingRequests"); + results.insert_default::("session.permissions.modifyRules"); + results.insert_default::("session.permissions.setRequired"); + results.insert_default::( + "session.permissions.notifyPromptShown", + ); + results.insert_default::("session.permissions.folderTrust.isTrusted"); + results.insert_default::( + "session.permissions.folderTrust.addTrusted", + ); + results + .insert_default::("session.permissions.locations.resolve"); + results.insert( + "session.permissions.locations.apply", + PermissionLocationApplyResult { + applied_directory_count: 2, + applied_rule_count: 3, + changed: true, + location_key: "repo-key".to_string(), + location_type: PermissionLocationType::Repo, + ..Default::default() + }, + ); + results.insert_default::( + "session.permissions.locations.addToolApproval", + ); + results.insert_default::("session.permissions.paths.list"); + results.insert_default::("session.permissions.paths.add"); + results.insert_default::( + "session.permissions.paths.updatePrimary", + ); + results.insert_default::( + "session.permissions.paths.isPathWithinAllowedDirectories", + ); + results.insert_default::( + "session.permissions.paths.isPathWithinWorkspace", + ); + results.insert_default::( + "session.permissions.urls.setUnrestrictedMode", + ); + results.insert( + "session.provider.getEndpoint", + ProviderEndpoint { + api_key: Some("fixture-api-key".to_string()), + base_url: "https://offline.invalid/v1".to_string(), + headers: HashMap::from([("x-fixture".to_string(), "rust".to_string())]), + ..Default::default() + }, + ); + + let mut errors = ErrorMap::default(); + errors.insert( + "session.debug.collectLogs", + -32077, + "fixture diagnostics unavailable", + json!({"retryable": false, "source": "offline"}), + ); + let fake = FakeCli::new(results, errors); + let client = fake.start_client().await; + let session = fake.create_session(&client).await; + + let error = session + .rpc() + .debug() + .collect_logs(DebugCollectLogsRequest { + additional_entries: None, + destination: DebugCollectLogsDestination::Directory( + DebugCollectLogsDestinationDirectory { + output_directory: "fixture-debug-output".to_string(), + ..Default::default() + }, + ), + include: None, + }) + .await + .expect_err("debug collection should return fixture RPC error"); + assert_eq!(error.rpc_code(), Some(-32077)); + assert_eq!(error.kind(), &ErrorKind::Rpc { code: -32077 }); + assert!( + error + .to_string() + .contains("fixture diagnostics unavailable") + ); + + rpc_ok!(session.rpc().mcp().move_loading_to_background()); + rpc_ok!( + session + .rpc() + .mcp() + .start_server(McpStartServerRequest::default()) + ); + rpc_ok!( + session + .rpc() + .mcp() + .restart_server(McpRestartServerRequest::default()) + ); + rpc_ok!( + session + .rpc() + .mcp() + .apps() + .read_resource(McpAppsReadResourceRequest::default()) + ); + rpc_ok!( + session + .rpc() + .mcp() + .oauth() + .authentication_state_changed(McpOauthAuthenticationStateChangedRequest::default()) + ); + let probe = rpc_ok!( + session + .rpc() + .mcp() + .oauth() + .probe(McpOauthProbeRequest::default()) + ); + let McpOauthProbeResult::Failed(failed) = probe else { + panic!("expected typed failed OAuth probe"); + }; + assert_eq!(failed.error, "offline probe fixture"); + let responded = rpc_ok!( + session + .rpc() + .mcp() + .oauth() + .respond(McpOauthRespondRequest::default()) + ); + assert!(responded.success); + rpc_ok!( + session + .rpc() + .mcp() + .resources() + .list(McpResourcesListRequest::default()) + ); + rpc_ok!( + session + .rpc() + .mcp() + .resources() + .list_templates(McpResourcesListTemplatesRequest::default()) + ); + + let metadata = rpc_ok!(session.rpc().metadata().get_client_metadata()); + assert_eq!(metadata["clientName"], "rust-fixture"); + let updated = rpc_ok!( + session + .rpc() + .metadata() + .update_client_metadata(MetadataUpdateClientMetadataRequest::default()) + ); + assert_eq!(updated["clientName"], "updated-fixture"); + rpc_ok!( + session + .rpc() + .model() + .switch_auto_tier(ModelSwitchAutoTierRequest::default()) + ); + rpc_ok!( + session + .rpc() + .model() + .list_with_params(ModelListRequest::default()) + ); + + rpc_ok!( + session + .rpc() + .permissions() + .configure(PermissionsConfigureParams::default()) + ); + rpc_ok!(session.rpc().permissions().pending_requests()); + rpc_ok!( + session + .rpc() + .permissions() + .modify_rules(PermissionsModifyRulesParams::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .set_required(PermissionsSetRequiredRequest::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .notify_prompt_shown(PermissionPromptShownNotification::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .folder_trust() + .is_trusted(FolderTrustCheckParams::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .folder_trust() + .add_trusted(FolderTrustAddParams::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .locations() + .resolve(PermissionLocationResolveParams { + working_directory: "fixture-worktree".to_string(), + }) + ); + let applied = rpc_ok!(session.rpc().permissions().locations().apply( + PermissionLocationApplyParams { + working_directory: "fixture-worktree".to_string(), + } + )); + assert!(applied.changed); + assert_eq!(applied.applied_directory_count, 2); + assert_eq!(applied.applied_rule_count, 3); + assert_eq!(applied.location_key, "repo-key"); + assert_eq!(applied.location_type, PermissionLocationType::Repo); + rpc_ok!(session.rpc().permissions().locations().add_tool_approval( + PermissionLocationAddToolApprovalParams { + approval: PermissionsLocationsAddToolApprovalDetails::Read( + PermissionsLocationsAddToolApprovalDetailsRead::default(), + ), + location_key: "repo-key".to_string(), + } + )); + rpc_ok!(session.rpc().permissions().paths().list()); + rpc_ok!( + session + .rpc() + .permissions() + .paths() + .add(PermissionPathsAddParams { + path: "allowed-dir".to_string(), + }) + ); + rpc_ok!( + session + .rpc() + .permissions() + .paths() + .update_primary(PermissionPathsUpdatePrimaryParams::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .paths() + .is_path_within_allowed_directories(PermissionPathsAllowedCheckParams { + path: "allowed-dir/file.rs".to_string(), + }) + ); + rpc_ok!( + session + .rpc() + .permissions() + .paths() + .is_path_within_workspace(PermissionPathsWorkspaceCheckParams::default()) + ); + rpc_ok!( + session + .rpc() + .permissions() + .urls() + .set_unrestricted_mode(PermissionUrlsSetUnrestrictedModeParams::default()) + ); + rpc_ok!( + session + .rpc() + .plugins() + .reload_with_params(PluginsReloadRequest::default()) + ); + let endpoint = rpc_ok!(session.rpc().provider().get_endpoint_with_params( + ProviderGetEndpointRequest { + model_id: Some("fixture-model".to_string()), + } + )); + assert_eq!(endpoint.base_url, "https://offline.invalid/v1"); + assert_eq!(endpoint.api_key.as_deref(), Some("fixture-api-key")); + assert_eq!(endpoint.headers["x-fixture"], "rust"); + + session.disconnect().await.expect("disconnect fake session"); + client.stop().await.expect("stop fake CLI"); + + fake.assert_target_methods(&[ + "session.debug.collectLogs", + "session.mcp.moveLoadingToBackground", + "session.mcp.startServer", + "session.mcp.restartServer", + "session.mcp.apps.readResource", + "session.mcp.oauth.authenticationStateChanged", + "session.mcp.oauth.probe", + "session.mcp.oauth.respond", + "session.mcp.resources.list", + "session.mcp.resources.listTemplates", + "session.metadata.getClientMetadata", + "session.metadata.updateClientMetadata", + "session.model.switchAutoTier", + "session.model.list", + "session.permissions.configure", + "session.permissions.pendingRequests", + "session.permissions.modifyRules", + "session.permissions.setRequired", + "session.permissions.notifyPromptShown", + "session.permissions.folderTrust.isTrusted", + "session.permissions.folderTrust.addTrusted", + "session.permissions.locations.resolve", + "session.permissions.locations.apply", + "session.permissions.locations.addToolApproval", + "session.permissions.paths.list", + "session.permissions.paths.add", + "session.permissions.paths.updatePrimary", + "session.permissions.paths.isPathWithinAllowedDirectories", + "session.permissions.paths.isPathWithinWorkspace", + "session.permissions.urls.setUnrestrictedMode", + "session.plugins.reload", + "session.provider.getEndpoint", + ]); + fake.assert_params( + "session.permissions.locations.addToolApproval", + 0, + json!({ + "sessionId": "rpc-surface-session", + "approval": {"kind": "read"}, + "locationKey": "repo-key" + }), + ); + fake.assert_params( + "session.provider.getEndpoint", + 0, + json!({ + "sessionId": "rpc-surface-session", + "modelId": "fixture-model" + }), + ); +} + +#[tokio::test] +async fn session_queue_tasks_tools_ui_and_workspace_rpc_surface_is_typed() { + let task = || TaskClientInfo { + execution_mode: TaskClientExecutionMode::Background, + owner: TaskClientOwner { + kind: TaskClientOwnerKind::Sdk, + presence: TaskClientOwnerPresence::Connected, + ..Default::default() + }, + status: TaskClientStatus::Running, + r#type: TaskClientType::Client, + ..Default::default() + }; + let mut results = ResponseMap::default(); + results.insert_default::("session.queue.moveItem"); + results.insert( + "session.queue.insertAt", + QueueInsertAtResult { + id: "queue-item-9".to_string(), + }, + ); + results.insert_default::("session.queue.removeAt"); + results.insert_default::("session.queue.updateText"); + results.insert_default::("session.queue.duplicateAt"); + results.insert( + "session.queue.sendNow", + QueueSendNowResult { steered: true }, + ); + results.insert_default::("session.sandbox.getEnforcementStatus"); + results.insert_default::("session.sandbox.disableForSession"); + results.insert( + "session.tasks.register", + TasksRegisterResult { + created: true, + task: task(), + ..Default::default() + }, + ); + results.insert( + "session.tasks.update", + TasksUpdateResult { + applied: true, + task: task(), + ..Default::default() + }, + ); + results + .insert_default::("session.tools.getBuiltinDescriptors"); + results.insert( + "session.tools.taskCompleteEventData", + TaskCompleteData { + objective_id: Some(41), + success: Some(true), + summary: Some("coverage complete".to_string()), + ..Default::default() + }, + ); + results.insert_default::("session.tools.set"); + results.insert( + "session.ui.elicitation", + UIElicitationResponse { + action: UIElicitationResponseAction::Accept, + content: Some(HashMap::from([( + "answer".to_string(), + json!("deterministic"), + )])), + ..Default::default() + }, + ); + results.insert( + "session.workspaces.updateMetadata", + workspace_result("workspace-updated"), + ); + results.insert( + "session.workspaces.ensure", + workspace_result("workspace-ensured"), + ); + results.insert( + "session.workspaces.statFile", + WorkspacesStatFileResult { + is_file: true, + mtime_ms: 1_234.0, + size: 88.0, + ..Default::default() + }, + ); + results.insert( + "session.workspaces.addSummary", + WorkspacesAddSummaryResult { + summary: Some(json!({"title": "offline summary", "number": 4})), + workspace: Some(json!({"id": "workspace-ensured"})), + }, + ); + results.insert( + "session.workspaces.truncateSummaries", + workspace_result("workspace-truncated"), + ); + results.insert( + "session.workspaces.readAutopilotObjective", + WorkspacesReadAutopilotObjectiveResult { + content: Some("Ship deterministic coverage".to_string()), + }, + ); + results.insert( + "session.workspaces.writeAutopilotObjective", + WorkspacesWriteAutopilotObjectiveResult { + operation: "created".to_string(), + }, + ); + results.insert_default::( + "session.workspaces.deleteAutopilotObjective", + ); + results.insert( + "session.workspaces.autopilotObjectiveExists", + WorkspacesAutopilotObjectiveExistsResult { exists: true }, + ); + + let fake = FakeCli::new(results, ErrorMap::default()); + let client = fake.start_client().await; + let session = fake.create_session(&client).await; + + rpc_ok!(session.rpc().queue().move_item(QueueMoveItemRequest { + id: "queue-item-1".to_string(), + to_position: 2, + })); + let inserted = rpc_ok!(session.rpc().queue().insert_at(QueueInsertAtRequest { + message: QueueInsertMessage { + billable: Some(false), + display_prompt: Some("Fixture display".to_string()), + prompt: "Queue this deterministically".to_string(), + request_headers: Some(HashMap::from([( + "x-test".to_string(), + "rpc-surface".to_string(), + )])), + ..Default::default() + }, + position: 1, + })); + assert_eq!(inserted.id, "queue-item-9"); + rpc_ok!( + session + .rpc() + .queue() + .remove_at(QueueRemoveAtRequest::default()) + ); + rpc_ok!( + session + .rpc() + .queue() + .update_text(QueueUpdateTextRequest::default()) + ); + rpc_ok!( + session + .rpc() + .queue() + .duplicate_at(QueueDuplicateAtRequest::default()) + ); + rpc_ok!( + session + .rpc() + .queue() + .set_drain_paused(QueueSetDrainPausedRequest { paused: true }) + ); + assert!( + rpc_ok!( + session + .rpc() + .queue() + .send_now(QueueSendNowRequest::default()) + ) + .steered + ); + rpc_ok!(session.rpc().sandbox().get_enforcement_status()); + rpc_ok!( + session + .rpc() + .sandbox() + .disable_for_session(SandboxDisableForSessionRequest::default()) + ); + + let registered = rpc_ok!(session.rpc().tasks().register(TasksRegisterRequest { + cancellable: true, + client_task_id: "client-task-7".to_string(), + description: "deterministic external work".to_string(), + display_name: Some("Coverage task".to_string()), + expected_sequence: Some(0), + r#type: TaskClientType::Client, + })); + assert!(registered.created); + let updated = rpc_ok!(session.rpc().tasks().update(TasksUpdateRequest { + id: "task-7".to_string(), + sequence: 1, + update: TaskClientUpdate::Completed(TaskClientUpdateCompleted { + message: Some("done".to_string()), + result: Some(json!({"files": 2})), + ..Default::default() + }), + })); + assert!(updated.applied); + + rpc_ok!( + session + .rpc() + .tools() + .get_builtin_descriptors(ToolsGetBuiltinDescriptorsRequest::default()) + ); + let completed = rpc_ok!( + session + .rpc() + .tools() + .task_complete_event_data(ToolsTaskCompleteEventDataRequest::default()) + ); + assert_eq!(completed.objective_id, Some(41)); + assert_eq!(completed.success, Some(true)); + assert_eq!(completed.summary.as_deref(), Some("coverage complete")); + rpc_ok!(session.rpc().tools().set(ToolsSetRequest::default())); + let elicitation = rpc_ok!( + session + .rpc() + .ui() + .elicitation(UIElicitationRequest::default()) + ); + assert_eq!(elicitation.action, UIElicitationResponseAction::Accept); + assert_eq!( + elicitation + .content + .as_ref() + .and_then(|content| content.get("answer")), + Some(&json!("deterministic")) + ); + + let updated_workspace = rpc_ok!( + session + .rpc() + .workspaces() + .update_metadata(WorkspacesUpdateMetadataRequest::default()) + ); + assert_eq!( + updated_workspace + .workspace + .as_ref() + .expect("updated workspace") + .id, + "workspace-updated" + ); + rpc_ok!( + session + .rpc() + .workspaces() + .ensure(WorkspacesEnsureRequest::default()) + ); + let stat = rpc_ok!( + session + .rpc() + .workspaces() + .stat_file(WorkspacesStatFileRequest::default()) + ); + assert!(stat.is_file); + assert_eq!(stat.size, 88.0); + assert_eq!(stat.mtime_ms, 1_234.0); + rpc_ok!( + session + .rpc() + .workspaces() + .create_directory(WorkspacesCreateDirectoryRequest { + path: "nested/output".to_string(), + recursive: Some(true), + }) + ); + rpc_ok!( + session + .rpc() + .workspaces() + .remove_path(WorkspacesRemovePathRequest::default()) + ); + rpc_ok!( + session + .rpc() + .workspaces() + .rename_path(WorkspacesRenamePathRequest::default()) + ); + let summary = rpc_ok!( + session + .rpc() + .workspaces() + .add_summary(WorkspacesAddSummaryRequest::default()) + ); + assert_eq!(summary.summary.as_ref().expect("summary")["number"], 4); + rpc_ok!( + session + .rpc() + .workspaces() + .truncate_summaries(WorkspacesTruncateSummariesRequest { keep_count: 2 }) + ); + let objective = rpc_ok!(session.rpc().workspaces().read_autopilot_objective()); + assert_eq!( + objective.content.as_deref(), + Some("Ship deterministic coverage") + ); + let write = rpc_ok!(session.rpc().workspaces().write_autopilot_objective( + WorkspacesWriteAutopilotObjectiveRequest { + content: "Updated deterministic objective".to_string(), + } + )); + assert_eq!(write.operation, "created"); + rpc_ok!(session.rpc().workspaces().delete_autopilot_objective()); + assert!(rpc_ok!(session.rpc().workspaces().autopilot_objective_exists()).exists); + + session.disconnect().await.expect("disconnect fake session"); + client.stop().await.expect("stop fake CLI"); + + fake.assert_target_methods(&[ + "session.queue.moveItem", + "session.queue.insertAt", + "session.queue.removeAt", + "session.queue.updateText", + "session.queue.duplicateAt", + "session.queue.setDrainPaused", + "session.queue.sendNow", + "session.sandbox.getEnforcementStatus", + "session.sandbox.disableForSession", + "session.tasks.register", + "session.tasks.update", + "session.tools.getBuiltinDescriptors", + "session.tools.taskCompleteEventData", + "session.tools.set", + "session.ui.elicitation", + "session.workspaces.updateMetadata", + "session.workspaces.ensure", + "session.workspaces.statFile", + "session.workspaces.createDirectory", + "session.workspaces.removePath", + "session.workspaces.renamePath", + "session.workspaces.addSummary", + "session.workspaces.truncateSummaries", + "session.workspaces.readAutopilotObjective", + "session.workspaces.writeAutopilotObjective", + "session.workspaces.deleteAutopilotObjective", + "session.workspaces.autopilotObjectiveExists", + ]); + fake.assert_params( + "session.queue.insertAt", + 0, + json!({ + "sessionId": "rpc-surface-session", + "message": { + "billable": false, + "displayPrompt": "Fixture display", + "prompt": "Queue this deterministically", + "requestHeaders": {"x-test": "rpc-surface"} + }, + "position": 1 + }), + ); + fake.assert_params( + "session.tasks.update", + 0, + json!({ + "sessionId": "rpc-surface-session", + "id": "task-7", + "sequence": 1, + "update": { + "kind": "completed", + "message": "done", + "result": {"files": 2} + } + }), + ); + fake.assert_params( + "session.workspaces.writeAutopilotObjective", + 0, + json!({ + "sessionId": "rpc-surface-session", + "content": "Updated deterministic objective" + }), + ); +} + +fn workspace_result(id: &str) -> WorkspacesGetWorkspaceResult { + WorkspacesGetWorkspaceResult { + path: Some("fixture-workspace".to_string()), + workspace: Some(WorkspacesGetWorkspaceResultWorkspace { + id: id.to_string(), + ..Default::default() + }), + } +} + +#[derive(Default)] +struct ResponseMap(HashMap<&'static str, Value>); + +impl ResponseMap { + fn insert(&mut self, method: &'static str, result: T) { + self.insert_value( + method, + serde_json::to_value(result).expect("serialize fake RPC result"), + ); + } + + fn insert_default(&mut self, method: &'static str) { + self.insert(method, T::default()); + } + + fn insert_value(&mut self, method: &'static str, result: Value) { + self.0.insert(method, result); + } +} + +#[derive(Default)] +struct ErrorMap(HashMap<&'static str, Value>); + +impl ErrorMap { + fn insert(&mut self, method: &'static str, code: i32, message: &str, data: Value) { + self.0.insert( + method, + json!({ + "code": code, + "message": message, + "data": data + }), + ); + } +} + +struct FakeCli { + _dir: TempDir, + script_path: PathBuf, + capture_path: PathBuf, + config_path: PathBuf, + work_dir: PathBuf, +} + +impl FakeCli { + fn new(results: ResponseMap, errors: ErrorMap) -> Self { + let dir = tempfile::tempdir().expect("create fake CLI temp dir"); + let script_path = dir.path().join("fake-rpc-cli.js"); + let capture_path = dir.path().join("captured-requests.json"); + let config_path = dir.path().join("responses.json"); + let work_dir = dir.path().join("cwd"); + std::fs::create_dir(&work_dir).expect("create fake CLI cwd"); + std::fs::write(&script_path, FAKE_STDIO_CLI_SCRIPT).expect("write fake CLI script"); + std::fs::write( + &config_path, + serde_json::to_vec(&json!({ + "results": results.0, + "errors": errors.0, + })) + .expect("serialize fake CLI config"), + ) + .expect("write fake CLI config"); + Self { + _dir: dir, + script_path, + capture_path, + config_path, + work_dir, + } + } + + async fn start_client(&self) -> Client { + Client::start( + ClientOptions::new() + .with_program(CliProgram::Path(PathBuf::from("node"))) + .with_prefix_args([self.script_path.as_os_str().to_owned()]) + .with_cwd(&self.work_dir) + .with_extra_args([ + "--capture-file".to_string(), + self.capture_path.to_string_lossy().into_owned(), + "--response-config".to_string(), + self.config_path.to_string_lossy().into_owned(), + ]) + .with_github_token("offline-rpc-token") + .with_use_logged_in_user(false) + .with_transport(Transport::Stdio), + ) + .await + .expect("start fake CLI client") + } + + async fn create_session(&self, client: &Client) -> github_copilot_sdk::session::Session { + client + .create_session( + SessionConfig::default() + .with_session_id("rpc-surface-session") + .with_working_directory(&self.work_dir), + ) + .await + .expect("create fake session") + } + + fn assert_target_methods(&self, expected: &[&str]) { + let actual: Vec<_> = self + .capture() + .into_iter() + .filter(|request| { + !matches!( + request.method.as_str(), + "connect" | "runtime.shutdown" | "session.create" | "session.detach" + ) + }) + .map(|request| request.method) + .collect(); + assert_eq!(actual, expected); + } + + fn assert_params(&self, method: &str, occurrence: usize, expected: Value) { + let request = self + .capture() + .into_iter() + .filter(|request| request.method == method) + .nth(occurrence) + .unwrap_or_else(|| panic!("missing occurrence {occurrence} of {method}")); + assert_eq!(request.params, expected, "unexpected params for {method}"); + } + + fn capture(&self) -> Vec { + let bytes = std::fs::read(&self.capture_path).expect("read fake CLI capture"); + serde_json::from_slice(&bytes).expect("parse fake CLI capture") + } +} + +#[derive(serde::Deserialize)] +struct CapturedRequest { + method: String, + #[serde(default)] + params: Value, +} + +const FAKE_STDIO_CLI_SCRIPT: &str = r#" +const fs = require("fs"); + +function argument(name) { + const index = process.argv.indexOf(name); + return index >= 0 ? process.argv[index + 1] : undefined; +} + +const captureFile = argument("--capture-file"); +const config = JSON.parse(fs.readFileSync(argument("--response-config"), "utf8")); +const requests = []; + +function saveCapture() { + fs.writeFileSync(captureFile, JSON.stringify(requests)); +} + +saveCapture(); + +let buffer = Buffer.alloc(0); +process.stdin.on("data", chunk => { + buffer = Buffer.concat([buffer, chunk]); + processBuffer(); +}); +process.stdin.resume(); + +function processBuffer() { + while (true) { + const headerEnd = buffer.indexOf("\r\n\r\n"); + if (headerEnd < 0) return; + const header = buffer.subarray(0, headerEnd).toString("utf8"); + const match = /Content-Length:\s*(\d+)/i.exec(header); + if (!match) throw new Error("Missing Content-Length header"); + const length = Number(match[1]); + const bodyStart = headerEnd + 4; + const bodyEnd = bodyStart + length; + if (buffer.length < bodyEnd) return; + const body = buffer.subarray(bodyStart, bodyEnd).toString("utf8"); + buffer = buffer.subarray(bodyEnd); + handleMessage(JSON.parse(body)); + } +} + +function handleMessage(message) { + if (!Object.prototype.hasOwnProperty.call(message, "id")) return; + + requests.push({ method: message.method, params: message.params }); + saveCapture(); + + if (message.method === "connect") { + writeResult(message.id, { ok: true, protocolVersion: 3, version: "offline-fixture" }); + return; + } + if (message.method === "session.create") { + writeResult(message.id, { + sessionId: message.params.sessionId, + workspacePath: null, + capabilities: null, + }); + return; + } + if (message.method === "session.detach") { + writeResult(message.id, { success: true }); + return; + } + if (Object.prototype.hasOwnProperty.call(config.errors, message.method)) { + writeError(message.id, config.errors[message.method]); + return; + } + + const result = Object.prototype.hasOwnProperty.call(config.results, message.method) + ? config.results[message.method] + : {}; + writeResult(message.id, result); +} + +function writeResult(id, result) { + writeMessage({ jsonrpc: "2.0", id, result }); +} + +function writeError(id, error) { + writeMessage({ jsonrpc: "2.0", id, error }); +} + +function writeMessage(message) { + const body = JSON.stringify(message); + process.stdout.write( + "Content-Length: " + Buffer.byteLength(body, "utf8") + "\r\n\r\n" + body, + ); +} +"#; diff --git a/rust/tests/e2e/skills.rs b/rust/tests/e2e/skills.rs index 769b28b5f9..99ba6ba948 100644 --- a/rust/tests/e2e/skills.rs +++ b/rust/tests/e2e/skills.rs @@ -1,6 +1,8 @@ use std::path::{Path, PathBuf}; +use std::sync::Arc; -use github_copilot_sdk::CustomAgentConfig; +use github_copilot_sdk::handler::ApproveAllHandler; +use github_copilot_sdk::{CustomAgentConfig, ResumeSessionConfig}; use super::support::{assert_uuid_like, assistant_message_content}; @@ -164,6 +166,113 @@ async fn should_not_provide_skills_to_agent_without_skills_field() { #[tokio::test] async fn should_apply_skill_on_session_resume_with_skilldirectories() {} +#[tokio::test] +async fn should_reload_replaced_skill_and_replay_it_on_resume() { + super::support::with_dedicated_e2e_context( + "scenario_testing_skills_and_agents", + "should_reload_atomically_replaced_skill_and_replay_it_on_resume", + |ctx| { + Box::pin(async move { + ctx.set_default_copilot_user(); + let skill_name = "scenario-reloadable-skill"; + let skills_dir = ctx.work_dir().join("scenario-reloadable-skills"); + let skill_file = write_versioned_skill( + &skills_dir, + skill_name, + "Scenario skill version one.", + "SCENARIO_SKILL_VERSION_ONE", + ); + let client = ctx.start_client().await; + let session = client + .create_session( + ctx.approve_all_session_config() + .with_enable_session_store(true) + .with_skill_directories([skills_dir.clone()]), + ) + .await + .expect("create session"); + let session_id = session.id().clone(); + + assert_versioned_skill( + session.rpc().skills().list().await.expect("list v1"), + skill_name, + "Scenario skill version one.", + &skill_file, + ); + + let replacement = skill_file.with_file_name("SKILL.replacement.md"); + std::fs::write( + &replacement, + skill_contents( + skill_name, + "Scenario skill version two.", + "SCENARIO_SKILL_VERSION_TWO", + ), + ) + .expect("write replacement skill"); + std::fs::rename(&replacement, &skill_file).expect("replace skill"); + session + .rpc() + .skills() + .reload() + .await + .expect("reload replaced skill"); + assert_versioned_skill( + session.rpc().skills().list().await.expect("list v2"), + skill_name, + "Scenario skill version two.", + &skill_file, + ); + + session + .log("SCENARIO_SKILL_RELOAD_READY", None) + .await + .expect("persist skill session"); + client + .rpc() + .sessions() + .save(github_copilot_sdk::rpc::SessionsSaveRequest { + session_id: session_id.clone(), + }) + .await + .expect("save session"); + session.rpc().suspend().await.expect("suspend session"); + session.stop_event_loop().await; + drop(session); + + let resumed = client + .resume_session( + ResumeSessionConfig::new(session_id) + .with_github_token(super::support::DEFAULT_TEST_TOKEN) + .with_permission_handler(Arc::new(ApproveAllHandler)) + .with_continue_pending_work(false) + .with_skill_directories([skills_dir]), + ) + .await + .expect("resume session"); + assert_versioned_skill( + resumed + .rpc() + .skills() + .list() + .await + .expect("list resumed skill"), + skill_name, + "Scenario skill version two.", + &skill_file, + ); + + resumed + .disconnect() + .await + .expect("disconnect resumed session"); + client.stop().await.expect("stop client"); + }) + }, + ) + .await; +} + fn create_skill_dir(work_dir: &Path) -> PathBuf { let skills_dir = work_dir.join(".test_skills"); let skill_subdir = skills_dir.join("test-skill"); @@ -180,4 +289,45 @@ fn create_skill_dir(work_dir: &Path) -> PathBuf { .expect("write skill file"); skills_dir } + +fn write_versioned_skill( + skills_dir: &Path, + name: &str, + description: &str, + marker: &str, +) -> PathBuf { + let skill_dir = skills_dir.join(name); + std::fs::create_dir_all(&skill_dir).expect("create versioned skill dir"); + let skill_file = skill_dir.join("SKILL.md"); + std::fs::write(&skill_file, skill_contents(name, description, marker)) + .expect("write versioned skill"); + skill_file +} + +fn skill_contents(name: &str, description: &str, marker: &str) -> String { + format!( + "---\nname: {name}\ndescription: {description}\n---\n\n\ + # Scenario Reloadable Skill\n\nUse {marker}.\n" + ) +} + +fn assert_versioned_skill( + list: github_copilot_sdk::rpc::SkillList, + name: &str, + description: &str, + path: &Path, +) { + let skill = list + .skills + .iter() + .find(|skill| skill.name == name) + .expect("versioned skill"); + assert!(skill.enabled); + assert_eq!(skill.description, description); + assert_eq!( + skill.path.as_deref().map(Path::new), + Some(path), + "unexpected skill path" + ); +} static E2E: super::support::SharedE2eGroup = super::support::SharedE2eGroup::standard("skills", 4); diff --git a/scripts/codegen/python.ts b/scripts/codegen/python.ts index ca312f3859..f10110c490 100644 --- a/scripts/codegen/python.ts +++ b/scripts/codegen/python.ts @@ -146,7 +146,8 @@ function placeholderToQuicktypeIdentifiers(placeholder: string): string[] { export function postProcessExternalRefsForPython( code: string, placeholderToReal: Map, - externalEnumNames: Set = new Set() + externalEnumNames: Set = new Set(), + externalDiscriminatedUnionNames: Set = new Set() ): string { for (const [placeholder, realName] of placeholderToReal) { const markerProperty = `__externalRefMarker_${placeholder}`; @@ -182,12 +183,49 @@ export function postProcessExternalRefsForPython( new RegExp(`to_class\\(${realName},\\s*([^)]+)\\)`, "g"), `to_enum(${realName}, $1)` ); + } else if (externalDiscriminatedUnionNames.has(realName)) { + code = code.replace(new RegExp(`\\b${realName}\\.from_dict\\b`, "g"), `_load_${realName}`); } } return code.replace(/\n{3,}/g, "\n\n"); } +function collectPythonExternalDiscriminatedUnionNames( + schema: JSONSchema7 | undefined, + placeholderToReal: Map +): Set { + const unionNames = new Set(); + if (!schema) return unionNames; + + const definitions = collectDefinitionCollections(schema as Record); + for (const realName of placeholderToReal.values()) { + // SessionEvent is emitted as a wrapper class with its own from_dict dispatcher, + // not as a union alias with a _load_* helper. + if (realName === "SessionEvent") continue; + const definition = definitions.definitions[realName] ?? definitions.$defs[realName]; + if (!definition) continue; + + const variants = definition.anyOf ?? definition.oneOf; + if (!Array.isArray(variants) || variants.length < 2) continue; + const resolvedVariants = variants.map((variant) => + typeof variant === "object" && variant !== null + ? resolveObjectSchema(variant, definitions) ?? + resolveSchema(variant, definitions) ?? + variant + : undefined + ); + if ( + resolvedVariants.every((variant) => variant?.properties !== undefined) && + findPyDiscriminator(resolvedVariants as JSONSchema7[]) + ) { + unionNames.add(realName); + } + } + + return unionNames; +} + function collectPythonExternalEnumNames( schema: JSONSchema7 | undefined, placeholderToReal: Map @@ -3107,6 +3145,10 @@ async function generateRpc(schemaPath?: string, sessionEventsSchema?: JSONSchema }; const externalRefs = rewriteExternalRefsForPython(singleSchema as JSONSchema7 & { definitions?: Record }); const externalEnumNames = collectPythonExternalEnumNames(sessionEventsSchema, externalRefs.placeholderNames); + const externalDiscriminatedUnionNames = collectPythonExternalDiscriminatedUnionNames( + sessionEventsSchema, + externalRefs.placeholderNames + ); const externalUnionAliases = collectExternalUnionAliasesForPython( singleSchema.definitions as Record, externalRefs.placeholderNames @@ -3158,12 +3200,23 @@ async function generateRpc(schemaPath?: string, sessionEventsSchema?: JSONSchema const knownDefNames = new Set(Object.keys(allDefinitions).map((n) => n.toLowerCase())); typesCode = collapsePlaceholderPythonDataclasses(typesCode, knownDefNames); typesCode = postProcessExternalUnionAliasesForPython(typesCode, externalUnionAliases); - typesCode = postProcessExternalRefsForPython(typesCode, externalRefs.placeholderNames, externalEnumNames); + typesCode = postProcessExternalRefsForPython( + typesCode, + externalRefs.placeholderNames, + externalEnumNames, + externalDiscriminatedUnionNames + ); typesCode = removeShadowedSessionEventEnumsForPython( typesCode, externalRefs.imports.get(".session_events") ?? new Set(), sessionEventsSchema ); + const sessionEventImports = externalRefs.imports.get(".session_events"); + if (sessionEventImports) { + for (const unionName of externalDiscriminatedUnionNames) { + sessionEventImports.add(`_load_${unionName}`); + } + } const { code: typesCodeAfterUnions, unions: refBasedUnions } = postProcessRefBasedDiscriminatedUnionsForPython( typesCode, allDefinitions, diff --git a/test/harness/modelProtocolAdapters.test.ts b/test/harness/modelProtocolAdapters.test.ts index add59ae77f..6773edca80 100644 --- a/test/harness/modelProtocolAdapters.test.ts +++ b/test/harness/modelProtocolAdapters.test.ts @@ -549,6 +549,72 @@ describe("protocol-aware replay", () => { }, ); + test.each( + backends.flatMap((backend) => + [false, true].map((splitHistory) => ({ backend, splitHistory })), + ), + )( + "preserves the continuation response boundary for $backend with splitHistory=$splitHistory", + async ({ backend, splitHistory }) => { + const history: NormalizedData["conversations"][number]["messages"] = [ + { role: "system", content: "${system}" }, + { role: "user", content: "Hello" }, + ...(splitHistory + ? [ + { role: "assistant" as const, content: "CONTEXT_" }, + { role: "assistant" as const, content: "READY" }, + ] + : [{ role: "assistant" as const, content: "CONTEXT_READY" }]), + ]; + await writeFile( + cachePath, + yaml.stringify( + { + models: ["captured-capi-model"], + conversations: [ + { messages: history }, + { + messages: [ + ...history, + { role: "assistant", content: "CONTINUATION_DONE" }, + ], + }, + ], + } satisfies NormalizedData, + { aliasDuplicateObjects: false }, + ), + ); + const request = requestFor(backend, "Hello"); + if (backend === "openai-responses") { + (request.input as unknown[]).push({ + type: "message", + role: "assistant", + content: [{ type: "output_text", text: "CONTEXT_READY" }], + }); + } else { + (request.messages as unknown[]).push( + ...(backend === "anthropic-messages" + ? [{ role: "assistant", content: "CONTEXT_READY" }] + : history.slice(2)), + ); + } + await withProxy(backend, async (proxyUrl) => { + for (const stream of [false, true]) { + const response = await fetch(`${proxyUrl}${endpoints[backend]}`, { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify({ ...request, stream }), + signal: AbortSignal.timeout(2_000), + }); + expect(response.status).toBe(200); + const body = await response.text(); + expect(body).toContain("CONTINUATION_DONE"); + expect(body).not.toContain("CONTEXT_READY"); + } + }); + }, + ); + test("does not rewrite canonical snapshots after BYOK replay", async () => { const original = await readFile(cachePath, "utf8"); const proxy = new ReplayingCapiProxy( diff --git a/test/harness/replayingCapiProxy.test.ts b/test/harness/replayingCapiProxy.test.ts index 470638bbaa..833a941119 100644 --- a/test/harness/replayingCapiProxy.test.ts +++ b/test/harness/replayingCapiProxy.test.ts @@ -463,7 +463,8 @@ Always include PINEAPPLE_COCONUT_42. ]); const result = await readYamlOutput(outputPath); - expect(result.conversations[0].messages[0].content).toBe(` + expect(result.conversations[0].messages[0].content) + .toBe(` Base directory for this skill: ${workingDirPlaceholder}/.test_skills/test-skill # Test Skill Instructions @@ -803,6 +804,57 @@ Always include PINEAPPLE_COCONUT_42. }); } + test("replay-only mode rejects cache misses without contacting the upstream", async () => { + let upstreamRequests = 0; + const upstream = http.createServer((_request, response) => { + upstreamRequests++; + response.writeHead(200, { "content-type": "application/json" }); + response.end(JSON.stringify({ choices: [] })); + }); + await new Promise((resolve) => + upstream.listen(0, "127.0.0.1", resolve), + ); + const address = upstream.address(); + if (!address || typeof address === "string") { + throw new Error("Upstream test server did not expose a TCP port."); + } + + const cachePath = path.join(tempDir, "cache.yaml"); + await writeFile( + cachePath, + yaml.stringify({ + models: ["test-model"], + conversations: [], + } satisfies NormalizedData), + ); + const proxy = new ReplayingCapiProxy(`http://127.0.0.1:${address.port}`); + await proxy.updateConfig({ + filePath: cachePath, + workDir, + backend: "capi", + replayOnly: true, + }); + const proxyUrl = await proxy.start(); + + try { + const response = await makeRequest(proxyUrl, "/chat/completions", { + body: { + model: "test-model", + messages: [{ role: "user", content: "cache miss" }], + }, + }); + + expect(response.status).toBe(500); + expect(response.body).toBe("Proxy error"); + expect(upstreamRequests).toBe(0); + } finally { + await proxy.stop(true); + await new Promise((resolve, reject) => + upstream.close((error) => (error ? reject(error) : resolve())), + ); + } + }); + test.each([ ["should_accept_blob_attachments", "pixel.png"], ["vision_disabled_then_enabled_via_setmodel", "test.png"], @@ -1002,7 +1054,9 @@ Always include PINEAPPLE_COCONUT_42. test("matches shell tool results with shell ID completion markers", async () => { const originalShellConfig = - process.platform === "win32" ? ShellConfig.powerShell : ShellConfig.bash; + process.platform === "win32" + ? ShellConfig.powerShell + : ShellConfig.bash; const cachePath = path.join(tempDir, "cache.yaml"); const cacheContent = yaml.stringify({ models: ["test-model"], @@ -1761,7 +1815,9 @@ Always include PINEAPPLE_COCONUT_42. const parsed = JSON.parse(response.body) as { data: Array<{ id: string }>; }; - expect(parsed.data.map((model) => model.id)).toEqual(["claude-sonnet-5"]); + expect(parsed.data.map((model) => model.id)).toEqual([ + "claude-sonnet-5", + ]); } finally { await proxy.stop(); } diff --git a/test/harness/replayingCapiProxy.ts b/test/harness/replayingCapiProxy.ts index 4ecbcdc52d..2511ca9255 100644 --- a/test/harness/replayingCapiProxy.ts +++ b/test/harness/replayingCapiProxy.ts @@ -165,6 +165,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { workDir, testInfo, backend: "capi", + replayOnly: false, autoResponseIndex: 0, toolResultNormalizers: [...this.defaultToolResultNormalizers], }; @@ -190,6 +191,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { // would silently overwrite the file with that subset, breaking subsequent runs. if ( this.state?.backend === "capi" && + !this.state.replayOnly && process.env.GITHUB_ACTIONS !== "true" ) { await writeCapturesToDisk(this.exchanges, this.state); @@ -200,6 +202,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { workDir: config.workDir, testInfo: config.testInfo, backend: parseReplayBackend(config.backend), + replayOnly: config.replayOnly === true, autoResponseIndex: 0, toolResultNormalizers: [...this.defaultToolResultNormalizers], }; @@ -229,10 +232,6 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { normalizeToolResultOrder(this.state.storedData.conversations); normalizeStoredUserMessages(this.state.storedData.conversations); normalizeStoredToolMessages(this.state.storedData.conversations); - normalizeStoredMessagesForBackend( - this.state.storedData.conversations, - this.state.backend, - ); } async stop(skipWritingCache?: boolean): Promise { @@ -242,6 +241,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { // same canonical snapshots replay through each provider protocol. if ( this.state?.backend === "capi" && + !this.state.replayOnly && !skipWritingCache && process.env.GITHUB_ACTIONS !== "true" ) { @@ -323,8 +323,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { options.requestOptions.path === "/exchanges" && options.requestOptions.method === "GET" ) { - const protocol = - replayProtocols[this.state?.backend ?? "capi"]; + const protocol = replayProtocols[this.state?.backend ?? "capi"]; const parsedExchanges = await Promise.all( this.exchanges .filter((exchange) => exchange.request.url === protocol.endpoint) @@ -552,7 +551,8 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { : options.body; if (state.storedData && isModelRequest && normalizedBody) { const streamingIsRequested = - (JSON.parse(normalizedBody) as { stream?: boolean }).stream === true; + (JSON.parse(normalizedBody) as { stream?: boolean }).stream === + true; const savedError = await findSavedChatCompletionError( state.storedData, @@ -589,6 +589,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { normalizedBody, state.workDir, state.toolResultNormalizers, + state.backend, ); if (savedResponse) { @@ -611,6 +612,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { normalizedBody, state.workDir, state.toolResultNormalizers, + state.backend, ) ) { const headers = { @@ -645,7 +647,7 @@ export class ReplayingCapiProxy extends CapturingHttpProxy { // Fallback to normal proxying if no cached response found // This implicitly captures the new exchange too const isCI = process.env.GITHUB_ACTIONS === "true"; - if (isCI || state.backend !== "capi") { + if (isCI || state.replayOnly || state.backend !== "capi") { await exitWithNoMatchingRequestError( options, state.testInfo, @@ -743,7 +745,7 @@ async function writeCapturesToDisk( /** * Produces a human-readable explanation of why no stored conversation matched * a given request. For each stored conversation it reports the first reason - * matching failed, mirroring the logic in {@link findAssistantIndexAfterPrefix}. + * matching failed against the uncoalesced canonical messages. */ function diagnoseMatchFailure( requestMessages: NormalizedMessage[], @@ -763,7 +765,7 @@ function diagnoseMatchFailure( for (let c = 0; c < storedData.conversations.length; c++) { const saved = storedData.conversations[c].messages; - // Same check as findAssistantIndexAfterPrefix: request must be a strict prefix + // Coalescing can only reduce the number of saved request messages. if (requestMessages.length >= saved.length) { lines.push( `Conversation ${c} (${saved.length} messages): ` + @@ -865,6 +867,7 @@ async function findSavedChatCompletionResponse( requestBody: string | undefined, workDir: string, toolResultNormalizers: ToolResultNormalizer[], + backend: ReplayBackend, ): Promise { // Normalize the incoming request the same way we normalize for caching const normalized = await parseAndNormalizeRequest( @@ -884,6 +887,7 @@ async function findSavedChatCompletionResponse( const replyIndex = findAssistantIndexAfterPrefix( requestMessages, conversation.messages, + backend, ); if (replyIndex !== undefined) { return createOpenAIResponse( @@ -937,6 +941,7 @@ async function isRequestOnlySnapshot( requestBody: string | undefined, workDir: string, toolResultNormalizers: ToolResultNormalizer[], + backend: ReplayBackend, ): Promise { const normalized = await parseAndNormalizeRequest( requestBody, @@ -946,11 +951,14 @@ async function isRequestOnlySnapshot( const requestMessages = normalized.conversations[0]?.messages ?? []; for (const conversation of storedData.conversations) { + const messages = normalizeMessagesForBackend( + conversation.messages, + backend, + ); if ( - requestMessages.length === conversation.messages.length && + requestMessages.length === messages.length && requestMessages.every( - (msg, i) => - JSON.stringify(msg) === JSON.stringify(conversation.messages[i]), + (msg, i) => JSON.stringify(msg) === JSON.stringify(messages[i]), ) ) { return true; @@ -1047,10 +1055,7 @@ function coalesceAdjacentUserMessages(requestBody: string): string { return JSON.stringify(request); } -function openAIErrorBody( - code: string | undefined, - message: string, -): unknown { +function openAIErrorBody(code: string | undefined, message: string): unknown { const type = code ?? "rate_limited"; return { error: { message, type, code: type } }; } @@ -1132,9 +1137,7 @@ function normalizeToolCalls( } if (tc.function?.name === "task") { - const configuredName = getBackgroundAgentName( - tc.function.arguments, - ); + const configuredName = getBackgroundAgentName(tc.function.arguments); const fallbackName = unnamedBackgroundAgentCounter === 0 ? "background-agent" @@ -1406,18 +1409,16 @@ function normalizeStoredUserMessages(conversations: NormalizedConversation[]) { } } -function normalizeStoredMessagesForBackend( - conversations: NormalizedConversation[], +function normalizeMessagesForBackend( + messages: NormalizedMessage[], backend: ReplayBackend, -) { - if (backend === "capi") return; +): NormalizedMessage[] { + if (backend === "capi") return messages; - for (const conversation of conversations) { - conversation.messages = coalesceMessages( - conversation.messages, - backend !== "openai-completions", - ); - } + return coalesceMessages( + messages.map((message) => ({ ...message })), + backend !== "openai-completions", + ); } function coalesceMessages( @@ -1553,15 +1554,12 @@ function normalizeGh401AuthMessages(result: string): string { function normalizeReadAgentResult(result: string): string { const normalized = result + .replace(/^Agent is idle \(waiting for messages\)\./, "Agent completed.") .replace( - /^Agent is idle \(waiting for messages\)\./, - "Agent completed.", - ) - .replace(/^Agent completed\. (.*), status: idle,/, "Agent completed. $1, status: completed,") - .replace( - /, total_turns: \d+(?=\r?\n|$)/, - ", total_turns: 0, duration: 0s", + /^Agent completed\. (.*), status: idle,/, + "Agent completed. $1, status: completed,", ) + .replace(/, total_turns: \d+(?=\r?\n|$)/, ", total_turns: 0, duration: 0s") .replace(/\r?\n\r?\n\[Turn \d+\]\r?\n/, "\n\n"); return normalized @@ -1854,37 +1852,54 @@ async function parseOpenAIResponse( function findAssistantIndexAfterPrefix( requestMessages: NormalizedMessage[], savedMessages: NormalizedMessage[], + backend: ReplayBackend, ): number | undefined { const logFile = process.env.PROXY_DEBUG_LOG; - const log = (msg: string) => { if (logFile) try { appendFileSync(logFile, msg + "\n"); } catch {} }; + const log = (msg: string) => { + if (logFile) + try { + appendFileSync(logFile, msg + "\n"); + } catch {} + }; if (requestMessages.length >= savedMessages.length) { - log(`prefix check failed: request.length=${requestMessages.length} >= saved.length=${savedMessages.length}`); + log( + `prefix check failed: request.length=${requestMessages.length} >= saved.length=${savedMessages.length}`, + ); return undefined; } - for (let i = 0; i < requestMessages.length; i++) { - const reqMsg = JSON.stringify(requestMessages[i]); - const savedMsg = JSON.stringify(savedMessages[i]); - if (reqMsg !== savedMsg) { - log(`mismatch at index ${i}:`); - log(` REQ: ${reqMsg.substring(0, 1000)}`); - log(` SAVED: ${savedMsg.substring(0, 1000)}`); - return undefined; - } - } - - // The next message after the prefix should be an assistant message - const nextIndex = requestMessages.length; - if ( - nextIndex < savedMessages.length && - savedMessages[nextIndex].role === "assistant" + for ( + let nextIndex = requestMessages.length; + nextIndex < savedMessages.length; + nextIndex++ ) { - log(`MATCH found at index ${nextIndex}`); - return nextIndex; + if (savedMessages[nextIndex].role !== "assistant") continue; + + // A continuation can start after an assistant message. Never coalesce + // across this candidate request/response boundary. + const prefix = normalizeMessagesForBackend( + savedMessages.slice(0, nextIndex), + backend, + ); + if (prefix.length > requestMessages.length) break; + if (prefix.length !== requestMessages.length) continue; + + const mismatchIndex = requestMessages.findIndex( + (message, i) => JSON.stringify(message) !== JSON.stringify(prefix[i]), + ); + if (mismatchIndex === -1) { + log(`MATCH found at index ${nextIndex}`); + return nextIndex; + } + log(`mismatch at index ${mismatchIndex} for reply index ${nextIndex}:`); + log( + ` REQ: ${JSON.stringify(requestMessages[mismatchIndex]).substring(0, 1000)}`, + ); + log(` SAVED: ${JSON.stringify(prefix[mismatchIndex]).substring(0, 1000)}`); } - log(`no assistant at nextIndex=${nextIndex}, saved.length=${savedMessages.length}`); + log(`no matching assistant boundary, saved.length=${savedMessages.length}`); return undefined; } @@ -2120,6 +2135,7 @@ type ReplayingCapiProxyState = { workDir: string; testInfo?: { file: string; line?: number }; backend: ReplayBackend; + replayOnly: boolean; storedData?: NormalizedData | undefined; autoResponseIndex: number; toolResultNormalizers: ToolResultNormalizer[]; diff --git a/test/harness/test-mcp-app-server.mjs b/test/harness/test-mcp-app-server.mjs new file mode 100644 index 0000000000..ef361f27fd --- /dev/null +++ b/test/harness/test-mcp-app-server.mjs @@ -0,0 +1,65 @@ +#!/usr/bin/env node +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { Server } from "@modelcontextprotocol/sdk/server/index.js"; +import { StdioServerTransport } from "@modelcontextprotocol/sdk/server/stdio.js"; +import { + CallToolRequestSchema, + ListToolsRequestSchema, + ReadResourceRequestSchema, +} from "@modelcontextprotocol/sdk/types.js"; + +const resourceUri = "ui://scenario/app"; +const server = new Server( + { name: "scenario-mcp-app", version: "1.0.0" }, + { capabilities: { resources: {}, tools: {} } } +); + +server.setRequestHandler(ListToolsRequestSchema, async () => ({ + tools: [ + { + name: "app_visible", + description: "Visible to MCP App views.", + inputSchema: { + type: "object", + properties: { value: { type: "string" } }, + required: ["value"], + }, + _meta: { "ui.visibility": ["model", "app"] }, + }, + ], +})); + +server.setRequestHandler(CallToolRequestSchema, async request => ({ + content: [ + { + type: "text", + text: `APP_VISIBLE:${request.params.arguments?.value ?? ""}`, + }, + ], +})); + +server.setRequestHandler(ReadResourceRequestSchema, async request => { + if (request.params.uri !== resourceUri) { + throw new Error(`Unknown resource: ${request.params.uri}`); + } + + return { + contents: [ + { + uri: resourceUri, + mimeType: "text/html", + text: "SCENARIO_MCP_APP", + _meta: { + "ui.csp": { + connectDomains: ["https://api.example.test"], + }, + }, + }, + ], + }; +}); + +await server.connect(new StdioServerTransport()); diff --git a/test/snapshots/mode_handlers/should_invoke_auto_mode_switch_handler_when_rate_limited.yaml b/test/snapshots/mode_handlers/should_invoke_auto_mode_switch_handler_when_rate_limited.yaml index 30fee89306..9048e85a0f 100644 --- a/test/snapshots/mode_handlers/should_invoke_auto_mode_switch_handler_when_rate_limited.yaml +++ b/test/snapshots/mode_handlers/should_invoke_auto_mode_switch_handler_when_rate_limited.yaml @@ -19,4 +19,4 @@ conversations: - role: user content: Explain that auto mode recovered from a rate limit in one short sentence. - role: assistant - content: Auto mode recovered from the rate limit and the session can continue. + content: Auto mode recovered from the rate limit and continued automatically. diff --git a/test/snapshots/mode_handlers/should_invoke_exit_plan_mode_handler_when_model_uses_tool.yaml b/test/snapshots/mode_handlers/should_invoke_exit_plan_mode_handler_when_model_uses_tool.yaml index 9ee28ad083..ba9e2c5cea 100644 --- a/test/snapshots/mode_handlers/should_invoke_exit_plan_mode_handler_when_model_uses_tool.yaml +++ b/test/snapshots/mode_handlers/should_invoke_exit_plan_mode_handler_when_model_uses_tool.yaml @@ -12,8 +12,7 @@ conversations: type: function function: name: exit_plan_mode - arguments: '{"summary":"Greeting file implementation - plan","actions":["autopilot","interactive","exit_only"],"recommendedAction":"interactive"}' + arguments: '{"summary":"Greeting file implementation plan","actions":["autopilot","interactive","exit_only"],"recommendedAction":"interactive"}' - role: tool tool_call_id: toolcall_0 content: >- @@ -23,4 +22,4 @@ conversations: You are now in interactive mode. Start implementing the plan now, in this same response. Approving the plan is your go-signal, so do not stop to ask whether to proceed or wait for another message. - role: assistant - content: Plan approved; I will wait for the next instruction before making changes. + content: The greeting file implementation plan was approved. diff --git a/test/snapshots/scenario_testing_callbacks/should_approve_scenario_exit_plan_with_full_callback_and_event_state.yaml b/test/snapshots/scenario_testing_callbacks/should_approve_scenario_exit_plan_with_full_callback_and_event_state.yaml new file mode 100644 index 0000000000..11a5eb48c6 --- /dev/null +++ b/test/snapshots/scenario_testing_callbacks/should_approve_scenario_exit_plan_with_full_callback_and_event_state.yaml @@ -0,0 +1,26 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Create a scenario client plan, then request approval with exit_plan_mode. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: exit_plan_mode + arguments: '{"summary":"scenario client implementation + plan","actions":["autopilot","interactive","exit_only"],"recommendedAction":"interactive"}' + - role: tool + tool_call_id: toolcall_0 + content: >- + Plan approved! Exited plan mode. + + + You are now in interactive mode. Start implementing the plan now, in this same response. Approving the plan is + your go-signal, so do not stop to ask whether to proceed or wait for another message. + - role: assistant + content: The scenario client plan was approved. diff --git a/test/snapshots/scenario_testing_callbacks/should_auto_switch_scenario_mode_after_rate_limit.yaml b/test/snapshots/scenario_testing_callbacks/should_auto_switch_scenario_mode_after_rate_limit.yaml new file mode 100644 index 0000000000..fd6a6b369f --- /dev/null +++ b/test/snapshots/scenario_testing_callbacks/should_auto_switch_scenario_mode_after_rate_limit.yaml @@ -0,0 +1,22 @@ +models: + - claude-sonnet-5 + - auto +errors: + - model: claude-sonnet-5 + status: 429 + code: user_weekly_rate_limited + message: You've reached your weekly rate limit. + retryAfterSeconds: 1 + messages: + - role: system + content: ${system} + - role: user + content: Explain that the scenario client recovered from a rate limit in one short sentence. +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Explain that the scenario client recovered from a rate limit in one short sentence. + - role: assistant + content: The scenario client recovered from the rate limit and continued automatically. diff --git a/test/snapshots/scenario_testing_callbacks/should_cancel_scenario_host_callback_when_channel_disconnects.yaml b/test/snapshots/scenario_testing_callbacks/should_cancel_scenario_host_callback_when_channel_disconnects.yaml new file mode 100644 index 0000000000..8050ab601f --- /dev/null +++ b/test/snapshots/scenario_testing_callbacks/should_cancel_scenario_host_callback_when_channel_disconnects.yaml @@ -0,0 +1,15 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_host_callback with value 'disconnect' and wait for it. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_host_callback + arguments: '{"value":"disconnect"}' diff --git a/test/snapshots/scenario_testing_callbacks/should_run_scenario_prompt_and_tool_hooks_with_full_context_and_suppression.yaml b/test/snapshots/scenario_testing_callbacks/should_run_scenario_prompt_and_tool_hooks_with_full_context_and_suppression.yaml new file mode 100644 index 0000000000..5d5a1bb9d5 --- /dev/null +++ b/test/snapshots/scenario_testing_callbacks/should_run_scenario_prompt_and_tool_hooks_with_full_context_and_suppression.yaml @@ -0,0 +1,32 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_hook_tool with value 'original', then reply with exactly SCENARIO_POST_RESULT. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_hook_tool + arguments: '{"value":"original"}' + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_hook_tool with value 'original', then reply with exactly SCENARIO_POST_RESULT. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_hook_tool + arguments: '{"value":"pre-hook"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_POST_RESULT + - role: assistant + content: SCENARIO_POST_RESULT diff --git a/test/snapshots/scenario_testing_canvas/should_handle_structured_scenario_canvas_error.yaml b/test/snapshots/scenario_testing_canvas/should_handle_structured_scenario_canvas_error.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_canvas/should_handle_structured_scenario_canvas_error.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/scenario_testing_canvas/should_reattach_scenario_canvas_and_route_all_callbacks_after_resume.yaml b/test/snapshots/scenario_testing_canvas/should_reattach_scenario_canvas_and_route_all_callbacks_after_resume.yaml new file mode 100644 index 0000000000..fa06483278 --- /dev/null +++ b/test/snapshots/scenario_testing_canvas/should_reattach_scenario_canvas_and_route_all_callbacks_after_resume.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_CANVAS_READY. + - role: assistant + content: SCENARIO_CANVAS_READY diff --git a/test/snapshots/scenario_testing_canvas/should_run_ordered_scenario_canvas_lifecycle_with_exact_context_and_snapshot.yaml b/test/snapshots/scenario_testing_canvas/should_run_ordered_scenario_canvas_lifecycle_with_exact_context_and_snapshot.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_canvas/should_run_ordered_scenario_canvas_lifecycle_with_exact_context_and_snapshot.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/scenario_testing_cloud/should_notify_steerability_then_send_first_message_without_remote_enable.yaml b/test/snapshots/scenario_testing_cloud/should_notify_steerability_then_send_first_message_without_remote_enable.yaml new file mode 100644 index 0000000000..4a05363039 --- /dev/null +++ b/test/snapshots/scenario_testing_cloud/should_notify_steerability_then_send_first_message_without_remote_enable.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_STEERABLE_FIRST_SEND. + - role: assistant + content: SCENARIO_STEERABLE_FIRST_SEND diff --git a/test/snapshots/scenario_testing_composition/should_classify_queued_and_immediate_scenario_messages_while_busy.yaml b/test/snapshots/scenario_testing_composition/should_classify_queued_and_immediate_scenario_messages_while_busy.yaml new file mode 100644 index 0000000000..170480d129 --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_classify_queued_and_immediate_scenario_messages_while_busy.yaml @@ -0,0 +1,26 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call wait_for_scenario_release, then reply with its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: wait_for_scenario_release + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: ACTIVE_TURN_RELEASED + - role: user + content: Reply with STEERING_SCENARIO_MESSAGE instead. + - role: assistant + content: STEERING_SCENARIO_MESSAGE + - role: user + content: Reply with QUEUED_SCENARIO_MESSAGE after the active turn. + - role: assistant + content: QUEUED_SCENARIO_MESSAGE diff --git a/test/snapshots/scenario_testing_composition/should_not_emit_redundant_model_change_when_resuming_same_model.yaml b/test/snapshots/scenario_testing_composition/should_not_emit_redundant_model_change_when_resuming_same_model.yaml new file mode 100644 index 0000000000..140ff7e907 --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_not_emit_redundant_model_change_when_resuming_same_model.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_SAME_MODEL_HISTORY_READY. + - role: assistant + content: SCENARIO_SAME_MODEL_HISTORY_READY diff --git a/test/snapshots/scenario_testing_composition/should_read_persisted_scenario_events_without_resuming.yaml b/test/snapshots/scenario_testing_composition/should_read_persisted_scenario_events_without_resuming.yaml new file mode 100644 index 0000000000..5aae15f15c --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_read_persisted_scenario_events_without_resuming.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_PERSISTED_HISTORY. + - role: assistant + content: SCENARIO_PERSISTED_HISTORY diff --git a/test/snapshots/scenario_testing_composition/should_resume_with_reattached_scenario_host_state.yaml b/test/snapshots/scenario_testing_composition/should_resume_with_reattached_scenario_host_state.yaml new file mode 100644 index 0000000000..b0ffa70a09 --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_resume_with_reattached_scenario_host_state.yaml @@ -0,0 +1,24 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Remember SCENARIO_RESUME_MARKER and reply with exactly INITIALIZED. + - role: assistant + content: INITIALIZED + - role: user + content: Call scenario_host_lookup with key ALPHA, then reply with exactly its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_host_lookup + arguments: '{"key":"ALPHA"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_HOST_VALUE_ALPHA + - role: assistant + content: SCENARIO_HOST_VALUE_ALPHA diff --git a/test/snapshots/scenario_testing_composition/should_retry_resume_on_replacement_client_after_recoverable_setup_failure.yaml b/test/snapshots/scenario_testing_composition/should_retry_resume_on_replacement_client_after_recoverable_setup_failure.yaml new file mode 100644 index 0000000000..797f41ec16 --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_retry_resume_on_replacement_client_after_recoverable_setup_failure.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_RETRY_RESUME_READY. + - role: assistant + content: SCENARIO_RETRY_RESUME_READY diff --git a/test/snapshots/scenario_testing_composition/should_send_scenario_message_with_metadata_and_extension_context.yaml b/test/snapshots/scenario_testing_composition/should_send_scenario_message_with_metadata_and_extension_context.yaml new file mode 100644 index 0000000000..ca67d619a5 --- /dev/null +++ b/test/snapshots/scenario_testing_composition/should_send_scenario_message_with_metadata_and_extension_context.yaml @@ -0,0 +1,15 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: |- + Reply with exactly TRACE_SENTINEL from the attached extension context. + + + + {"selection":"TRACE_SENTINEL","line":42} + - role: assistant + content: TRACE_SENTINEL diff --git a/test/snapshots/scenario_testing_control_state/should_report_processing_while_scenario_tool_is_running.yaml b/test/snapshots/scenario_testing_control_state/should_report_processing_while_scenario_tool_is_running.yaml new file mode 100644 index 0000000000..b10f18df5e --- /dev/null +++ b/test/snapshots/scenario_testing_control_state/should_report_processing_while_scenario_tool_is_running.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call wait_for_scenario_control, then reply with exactly SCENARIO_CONTROL_DONE. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: wait_for_scenario_control + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_CONTROL_DONE + - role: assistant + content: SCENARIO_CONTROL_DONE diff --git a/test/snapshots/scenario_testing_empty_runtime/empty_mode_minimal_toolless_session_has_no_tools.yaml b/test/snapshots/scenario_testing_empty_runtime/empty_mode_minimal_toolless_session_has_no_tools.yaml new file mode 100644 index 0000000000..f8dcea99db --- /dev/null +++ b/test/snapshots/scenario_testing_empty_runtime/empty_mode_minimal_toolless_session_has_no_tools.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Start. + - role: assistant + content: EMPTY_SCENARIO_READY diff --git a/test/snapshots/scenario_testing_event_subscriptions/should_deliver_mixed_scenario_event_stream_in_order_after_handler_lag.yaml b/test/snapshots/scenario_testing_event_subscriptions/should_deliver_mixed_scenario_event_stream_in_order_after_handler_lag.yaml new file mode 100644 index 0000000000..cf626e6494 --- /dev/null +++ b/test/snapshots/scenario_testing_event_subscriptions/should_deliver_mixed_scenario_event_stream_in_order_after_handler_lag.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_event_lookup with key 'ordered', then reply with exactly its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_event_lookup + arguments: '{"key":"ordered"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_EVENT_ORDERED + - role: assistant + content: SCENARIO_EVENT_ORDERED diff --git a/test/snapshots/scenario_testing_event_subscriptions/should_stop_closed_and_replaced_scenario_event_sources.yaml b/test/snapshots/scenario_testing_event_subscriptions/should_stop_closed_and_replaced_scenario_event_sources.yaml new file mode 100644 index 0000000000..6d7b487220 --- /dev/null +++ b/test/snapshots/scenario_testing_event_subscriptions/should_stop_closed_and_replaced_scenario_event_sources.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_EVENT_SOURCE_ONE. + - role: assistant + content: SCENARIO_EVENT_SOURCE_ONE diff --git a/test/snapshots/scenario_testing_js_extension_bridge/should_bridge_js_extension_canvas_context_log_and_session_continuation.yaml b/test/snapshots/scenario_testing_js_extension_bridge/should_bridge_js_extension_canvas_context_log_and_session_continuation.yaml new file mode 100644 index 0000000000..9a236f1397 --- /dev/null +++ b/test/snapshots/scenario_testing_js_extension_bridge/should_bridge_js_extension_canvas_context_log_and_session_continuation.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly JS_EXTENSION_CONTINUATION. + - role: assistant + content: JS_EXTENSION_CONTINUATION diff --git a/test/snapshots/scenario_testing_js_extension_bridge/should_surface_structured_canvaserror_from_js_extension.yaml b/test/snapshots/scenario_testing_js_extension_bridge/should_surface_structured_canvaserror_from_js_extension.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_js_extension_bridge/should_surface_structured_canvaserror_from_js_extension.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/external_tool_cancellation/should_cancel_tool_handler_when_session_disposes.yaml b/test/snapshots/scenario_testing_lifecycle_recovery/should_abort_active_scenario_turn_and_remain_usable.yaml similarity index 53% rename from test/snapshots/external_tool_cancellation/should_cancel_tool_handler_when_session_disposes.yaml rename to test/snapshots/scenario_testing_lifecycle_recovery/should_abort_active_scenario_turn_and_remain_usable.yaml index 028b44e73f..74fee0966a 100644 --- a/test/snapshots/external_tool_cancellation/should_cancel_tool_handler_when_session_disposes.yaml +++ b/test/snapshots/scenario_testing_lifecycle_recovery/should_abort_active_scenario_turn_and_remain_usable.yaml @@ -5,14 +5,18 @@ conversations: - role: system content: ${system} - role: user - content: Use slow_analysis with value 'test_abort'. Wait for the result. + content: Call scenario_blocking_lookup with key 'abort', then reply with the result. - role: assistant tool_calls: - id: toolcall_0 type: function function: - name: slow_analysis - arguments: '{"value":"test_abort"}' + name: scenario_blocking_lookup + arguments: '{"key":"abort"}' - role: tool tool_call_id: toolcall_0 content: The execution of this tool, or a previous tool was interrupted. + - role: user + content: Reply with exactly SCENARIO_ABORT_RECOVERY_OK. + - role: assistant + content: SCENARIO_ABORT_RECOVERY_OK diff --git a/test/snapshots/scenario_testing_lifecycle_recovery/should_suspend_disconnect_and_resume_scenario_state_without_delete.yaml b/test/snapshots/scenario_testing_lifecycle_recovery/should_suspend_disconnect_and_resume_scenario_state_without_delete.yaml new file mode 100644 index 0000000000..409c22dd5d --- /dev/null +++ b/test/snapshots/scenario_testing_lifecycle_recovery/should_suspend_disconnect_and_resume_scenario_state_without_delete.yaml @@ -0,0 +1,14 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Remember SCENARIO_LIFECYCLE_MEMORY and reply with exactly SCENARIO_LIFECYCLE_INITIALIZED. + - role: assistant + content: SCENARIO_LIFECYCLE_INITIALIZED + - role: user + content: Reply with exactly the scenario lifecycle memory value from the earlier turn. + - role: assistant + content: SCENARIO_LIFECYCLE_MEMORY diff --git a/test/snapshots/scenario_testing_mcp/should_preserve_disabled_scenario_mcp_servers_across_reload_and_resume.yaml b/test/snapshots/scenario_testing_mcp/should_preserve_disabled_scenario_mcp_servers_across_reload_and_resume.yaml new file mode 100644 index 0000000000..f92459d730 --- /dev/null +++ b/test/snapshots/scenario_testing_mcp/should_preserve_disabled_scenario_mcp_servers_across_reload_and_resume.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_MCP_DISABLED_STATE. + - role: assistant + content: SCENARIO_MCP_DISABLED_STATE diff --git a/test/snapshots/scenario_testing_permissions/should_forward_exact_scenario_permission_callback_payload.yaml b/test/snapshots/scenario_testing_permissions/should_forward_exact_scenario_permission_callback_payload.yaml new file mode 100644 index 0000000000..5d6dc1c25a --- /dev/null +++ b/test/snapshots/scenario_testing_permissions/should_forward_exact_scenario_permission_callback_payload.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_permission_tool with key 'payload', then reply with exactly its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_permission_tool + arguments: '{"key":"payload"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_PERMISSION_PAYLOAD + - role: assistant + content: SCENARIO_PERMISSION_PAYLOAD diff --git a/test/snapshots/scenario_testing_persistence/should_page_persisted_events_backward_without_resuming.yaml b/test/snapshots/scenario_testing_persistence/should_page_persisted_events_backward_without_resuming.yaml new file mode 100644 index 0000000000..f4cfe2e07d --- /dev/null +++ b/test/snapshots/scenario_testing_persistence/should_page_persisted_events_backward_without_resuming.yaml @@ -0,0 +1,14 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly PERSISTED_SCENARIO_FIRST. + - role: assistant + content: PERSISTED_SCENARIO_FIRST + - role: user + content: Reply with exactly PERSISTED_SCENARIO_SECOND. + - role: assistant + content: PERSISTED_SCENARIO_SECOND diff --git a/test/snapshots/scenario_testing_persistence/should_retry_from_existing_history_with_empty_sendmessages.yaml b/test/snapshots/scenario_testing_persistence/should_retry_from_existing_history_with_empty_sendmessages.yaml new file mode 100644 index 0000000000..b5efa8ed65 --- /dev/null +++ b/test/snapshots/scenario_testing_persistence/should_retry_from_existing_history_with_empty_sendmessages.yaml @@ -0,0 +1,19 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly EMPTY_BATCH_CONTEXT_READY. + - role: assistant + content: EMPTY_BATCH_CONTEXT_READY + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly EMPTY_BATCH_CONTEXT_READY. + - role: assistant + content: EMPTY_BATCH_CONTEXT_READY + - role: assistant + content: EMPTY_BATCH_RETRY_DONE diff --git a/test/snapshots/scenario_testing_persistence/should_truncate_history_and_resend_from_boundary.yaml b/test/snapshots/scenario_testing_persistence/should_truncate_history_and_resend_from_boundary.yaml new file mode 100644 index 0000000000..b85628e963 --- /dev/null +++ b/test/snapshots/scenario_testing_persistence/should_truncate_history_and_resend_from_boundary.yaml @@ -0,0 +1,25 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly HISTORY_SCENARIO_FIRST. + - role: assistant + content: HISTORY_SCENARIO_FIRST + - role: user + content: Reply with exactly HISTORY_SCENARIO_DISCARDED. + - role: assistant + content: HISTORY_SCENARIO_DISCARDED + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly HISTORY_SCENARIO_FIRST. + - role: assistant + content: HISTORY_SCENARIO_FIRST + - role: user + content: Reply with exactly HISTORY_SCENARIO_REPLACEMENT. + - role: assistant + content: HISTORY_SCENARIO_REPLACEMENT diff --git a/test/snapshots/scenario_testing_providers/should_apply_reasoning_context_and_auto_atomically_without_implicit_reset.yaml b/test/snapshots/scenario_testing_providers/should_apply_reasoning_context_and_auto_atomically_without_implicit_reset.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_providers/should_apply_reasoning_context_and_auto_atomically_without_implicit_reset.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/scenario_testing_runtime/should_ping_then_reuse_client_across_two_sessions.yaml b/test/snapshots/scenario_testing_runtime/should_ping_then_reuse_client_across_two_sessions.yaml new file mode 100644 index 0000000000..4acb5280d7 --- /dev/null +++ b/test/snapshots/scenario_testing_runtime/should_ping_then_reuse_client_across_two_sessions.yaml @@ -0,0 +1,17 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly FIRST_SCENARIO_SESSION. + - role: assistant + content: FIRST_SCENARIO_SESSION + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SECOND_SCENARIO_SESSION. + - role: assistant + content: SECOND_SCENARIO_SESSION diff --git a/test/snapshots/scenario_testing_sends/should_order_idle_queued_and_immediate_scenario_delivery.yaml b/test/snapshots/scenario_testing_sends/should_order_idle_queued_and_immediate_scenario_delivery.yaml new file mode 100644 index 0000000000..a0291ce5db --- /dev/null +++ b/test/snapshots/scenario_testing_sends/should_order_idle_queued_and_immediate_scenario_delivery.yaml @@ -0,0 +1,46 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly IDLE_ENQUEUE. + - role: assistant + content: IDLE_ENQUEUE + - role: user + content: Reply with exactly IDLE_IMMEDIATE. + - role: assistant + content: IDLE_IMMEDIATE + - role: user + content: Call scenario_send_blocker, then reply with its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_send_blocker + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_SEND_BLOCKER_RELEASED + - role: user + content: Call scenario_send_blocker again, then reply with exactly FIRST_STEERING. + - role: assistant + tool_calls: + - id: toolcall_1 + type: function + function: + name: scenario_send_blocker + arguments: "{}" + - role: tool + tool_call_id: toolcall_1 + content: SCENARIO_SEND_BLOCKER_RELEASED_AGAIN + - role: user + content: Reply with exactly SECOND_IMMEDIATE. + - role: assistant + content: SECOND_IMMEDIATE + - role: user + content: Reply with exactly FINAL_QUEUED. + - role: assistant + content: FINAL_QUEUED diff --git a/test/snapshots/scenario_testing_skills_and_agents/should_classify_agent_method_not_found_as_remote_protocol_error.yaml b/test/snapshots/scenario_testing_skills_and_agents/should_classify_agent_method_not_found_as_remote_protocol_error.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_skills_and_agents/should_classify_agent_method_not_found_as_remote_protocol_error.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/scenario_testing_skills_and_agents/should_reload_atomically_replaced_skill_and_replay_it_on_resume.yaml b/test/snapshots/scenario_testing_skills_and_agents/should_reload_atomically_replaced_skill_and_replay_it_on_resume.yaml new file mode 100644 index 0000000000..0c6b353c19 --- /dev/null +++ b/test/snapshots/scenario_testing_skills_and_agents/should_reload_atomically_replaced_skill_and_replay_it_on_resume.yaml @@ -0,0 +1,3 @@ +models: + - claude-sonnet-5 +conversations: [] diff --git a/test/snapshots/scenario_testing_tools/should_advertise_scenario_tool_schema_override_and_availability.yaml b/test/snapshots/scenario_testing_tools/should_advertise_scenario_tool_schema_override_and_availability.yaml new file mode 100644 index 0000000000..12e5d6bcae --- /dev/null +++ b/test/snapshots/scenario_testing_tools/should_advertise_scenario_tool_schema_override_and_availability.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_lookup_issue for owner octo and issue number 42. Reply with its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_lookup_issue + arguments: '{"owner":"octo","number":42}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_ISSUE_octo_42 + - role: assistant + content: SCENARIO_ISSUE_octo_42 diff --git a/test/snapshots/scenario_testing_tools/should_cancel_scenario_tool_handler_when_session_disposes.yaml b/test/snapshots/scenario_testing_tools/should_cancel_scenario_tool_handler_when_session_disposes.yaml new file mode 100644 index 0000000000..a4a8f1be8d --- /dev/null +++ b/test/snapshots/scenario_testing_tools/should_cancel_scenario_tool_handler_when_session_disposes.yaml @@ -0,0 +1,15 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_wait_for_operation with operation sync-installation. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_wait_for_operation + arguments: '{"operation":"sync-installation"}' diff --git a/test/snapshots/scenario_testing_tools/should_deliver_expanded_scenario_tool_result_to_the_model.yaml b/test/snapshots/scenario_testing_tools/should_deliver_expanded_scenario_tool_result_to_the_model.yaml new file mode 100644 index 0000000000..048dd825ae --- /dev/null +++ b/test/snapshots/scenario_testing_tools/should_deliver_expanded_scenario_tool_result_to_the_model.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_get_deployment for environment production. Reply with its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_get_deployment + arguments: '{"environment":"production"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_DEPLOYMENT_READY:production + - role: assistant + content: SCENARIO_DEPLOYMENT_READY:production diff --git a/test/snapshots/scenario_testing_tools/should_isolate_scenario_tool_handler_error.yaml b/test/snapshots/scenario_testing_tools/should_isolate_scenario_tool_handler_error.yaml new file mode 100644 index 0000000000..a34fbf246f --- /dev/null +++ b/test/snapshots/scenario_testing_tools/should_isolate_scenario_tool_handler_error.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_failing_lookup. If it fails, reply with exactly SCENARIO_LOOKUP_UNAVAILABLE. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_failing_lookup + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: "Failed to execute `scenario_failing_lookup` tool with arguments: {} due to error: Error: Tool execution failed" + - role: assistant + content: SCENARIO_LOOKUP_UNAVAILABLE diff --git a/test/snapshots/scenario_testing_tools/should_preserve_scenario_tool_invocation_identity_arguments_and_text.yaml b/test/snapshots/scenario_testing_tools/should_preserve_scenario_tool_invocation_identity_arguments_and_text.yaml new file mode 100644 index 0000000000..16c7b8c296 --- /dev/null +++ b/test/snapshots/scenario_testing_tools/should_preserve_scenario_tool_invocation_identity_arguments_and_text.yaml @@ -0,0 +1,20 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call scenario_search_pull_requests with query is:open label:bug. Reply with its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: scenario_search_pull_requests + arguments: '{"query":"is:open label:bug"}' + - role: tool + tool_call_id: toolcall_0 + content: SCENARIO_SEARCH_TEXT:is:open label:bug + - role: assistant + content: SCENARIO_SEARCH_TEXT:is:open label:bug diff --git a/test/snapshots/scenario_testing_utility/should_send_wait_observe_idle_events_and_delete_suggestion_session.yaml b/test/snapshots/scenario_testing_utility/should_send_wait_observe_idle_events_and_delete_suggestion_session.yaml new file mode 100644 index 0000000000..0410b31b07 --- /dev/null +++ b/test/snapshots/scenario_testing_utility/should_send_wait_observe_idle_events_and_delete_suggestion_session.yaml @@ -0,0 +1,10 @@ +models: + - claude-sonnet-5 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Reply with exactly SCENARIO_SUGGESTION_ACCEPTED. + - role: assistant + content: SCENARIO_SUGGESTION_ACCEPTED