diff --git a/.github/workflows/java-sdk-tests.yml b/.github/workflows/java-sdk-tests.yml index d1dd9c63e5..8fb6c0b5d9 100644 --- a/.github/workflows/java-sdk-tests.yml +++ b/.github/workflows/java-sdk-tests.yml @@ -72,6 +72,90 @@ jobs: - name: Verify CLI works run: node ../nodejs/node_modules/@github/copilot/npm-loader.js --version + - name: Prime Spotless Eclipse formatter cache + if: matrix.test-jdk == '25' + run: | + set -euo pipefail + + p2_data="$HOME/.m2/repository/dev/equo/p2-data" + bundle_dir="$p2_data/bundle-pool/https-download.eclipse.org-eclipse-updates-4.33-R-4.33-202409030240-" + query_dir="$p2_data/queries/1.8.1-1468712279" + jna_jar="$bundle_dir/com.sun.jna_5.14.0.v20231211-1200.jar" + + mkdir -p "$bundle_dir" "$query_dir" + printf '1' > "$p2_data/queries/version" + printf 'https://download.eclipse.org/eclipse/updates/4.33/R-4.33-202409030240/' > "$bundle_dir/.url" + + mvn -q dependency:get -Dartifact=net.java.dev.jna:jna:5.14.0 -Dtransitive=false + cp "$HOME/.m2/repository/net/java/dev/jna/jna/5.14.0/jna-5.14.0.jar" "$jna_jar" + + helper_dir="$(mktemp -d)" + trap 'rm -rf "$helper_dir"' EXIT + mkdir -p "$helper_dir/dev/equo/solstice/p2" + cat > "$helper_dir/dev/equo/solstice/p2/P2QueryResult.java" <<'JAVA' + package dev.equo.solstice.p2; + + import java.io.File; + import java.io.FileOutputStream; + import java.io.ObjectOutputStream; + import java.io.Serializable; + import java.util.ArrayList; + import java.util.Collections; + import java.util.List; + + public class P2QueryResult implements Serializable { + private static final long serialVersionUID = 1L; + + private final List mavenCoordinates; + private final List downloadedP2Jars; + + private P2QueryResult(List mavenCoordinates, List downloadedP2Jars) { + this.mavenCoordinates = mavenCoordinates; + this.downloadedP2Jars = downloadedP2Jars; + } + + public static void main(String[] args) throws Exception { + var coordinates = new ArrayList(); + Collections.addAll( + coordinates, + "net.java.dev.jna:jna-platform:5.14.0", + "org.apache.felix:org.apache.felix.scr:2.2.12", + "org.eclipse.platform:org.eclipse.core.commands:3.12.200", + "org.eclipse.platform:org.eclipse.core.contenttype:3.9.500", + "org.eclipse.platform:org.eclipse.core.expressions:3.9.400", + "org.eclipse.platform:org.eclipse.core.filesystem:1.11.0", + "org.eclipse.platform:org.eclipse.core.jobs:3.15.400", + "org.eclipse.platform:org.eclipse.core.resources:3.21.0", + "org.eclipse.platform:org.eclipse.core.runtime:3.31.100", + "org.eclipse.platform:org.eclipse.equinox.app:1.7.200", + "org.eclipse.platform:org.eclipse.equinox.common:3.19.100", + "org.eclipse.platform:org.eclipse.equinox.event:1.7.100", + "org.eclipse.platform:org.eclipse.equinox.preferences:3.11.100", + "org.eclipse.platform:org.eclipse.equinox.registry:3.12.100", + "org.eclipse.platform:org.eclipse.equinox.supplement:1.11.0", + "org.eclipse.jdt:org.eclipse.jdt.core:3.39.0", + "org.eclipse.jdt:ecj:3.39.0", + "org.eclipse.platform:org.eclipse.osgi:3.21.0", + "org.eclipse.platform:org.eclipse.text:3.14.100", + "org.osgi:org.osgi.service.cm:1.6.1", + "org.osgi:org.osgi.service.component:1.5.1", + "org.osgi:org.osgi.service.event:1.4.1", + "org.osgi:org.osgi.service.metatype:1.4.1", + "org.osgi:org.osgi.service.prefs:1.1.2", + "org.osgi:org.osgi.util.function:1.2.0", + "org.osgi:org.osgi.util.promise:1.3.0"); + + var result = new P2QueryResult(coordinates, List.of(new File(args[1]))); + try (var out = new ObjectOutputStream(new FileOutputStream(args[0]))) { + out.writeObject(result); + } + } + } + JAVA + + javac "$helper_dir/dev/equo/solstice/p2/P2QueryResult.java" + java -cp "$helper_dir" dev.equo.solstice.p2.P2QueryResult "$query_dir/content" "$jna_jar" + - name: Run spotless check if: matrix.test-jdk == '25' run: | diff --git a/dotnet/src/Client.cs b/dotnet/src/Client.cs index 6041fe2391..72408e5c22 100644 --- a/dotnet/src/Client.cs +++ b/dotnet/src/Client.cs @@ -1802,7 +1802,17 @@ private async Task VerifyProtocolVersionAsync(Connection connection, Cancellatio _ => null, }; var connectResponse = await InvokeRpcAsync( - connection.Rpc, "connect", [new ConnectRequest { Token = token }], connection.StderrBuffer, cancellationToken); + connection.Rpc, + "connect", + [new ConnectHandshakeRequest( + token, + // Opt in to GitHub telemetry forwarding at the connection level when a + // handler is registered (mirrors the runtime, which reads this flag on the + // `connect` handshake so the first session's un-replayable `session.start` + // event is forwarded). Also sent on session.create/resume for older CLIs. + _options.OnGitHubTelemetry != null ? true : null)], + connection.StderrBuffer, + cancellationToken); serverVersion = (int)connectResponse.ProtocolVersion; } catch (IOException ex) when (ex.InnerException is RemoteRpcException remoteEx && IsUnsupportedConnectMethod(remoteEx)) @@ -2639,6 +2649,10 @@ internal record GetSessionMetadataRequest( internal record GetSessionMetadataResponse( SessionMetadata? Session); + internal record ConnectHandshakeRequest( + string? Token, + [property: JsonPropertyName("enableGitHubTelemetryForwarding")] bool? EnableGitHubTelemetryForwarding = null); + internal record SetForegroundSessionRequest( string SessionId); @@ -2673,6 +2687,7 @@ internal record HooksInvokeResponse( [JsonSerializable(typeof(ListSessionsResponse))] [JsonSerializable(typeof(GetSessionMetadataRequest))] [JsonSerializable(typeof(GetSessionMetadataResponse))] + [JsonSerializable(typeof(ConnectHandshakeRequest))] [JsonSerializable(typeof(McpOAuthTokenStorageMode))] [JsonSerializable(typeof(EmbeddingCacheStorageMode))] [JsonSerializable(typeof(ModelCapabilitiesOverride))] diff --git a/dotnet/test/E2E/GitHubTelemetryForwardingE2ETests.cs b/dotnet/test/E2E/GitHubTelemetryForwardingE2ETests.cs index 85f3706e25..343c4815b7 100644 --- a/dotnet/test/E2E/GitHubTelemetryForwardingE2ETests.cs +++ b/dotnet/test/E2E/GitHubTelemetryForwardingE2ETests.cs @@ -43,7 +43,7 @@ await TestHelper.WaitForConditionAsync( timeoutMessage: "Timed out waiting for GitHub telemetry notification."); Assert.True(notifications.TryPeek(out var notification)); - Assert.NotEmpty(notification.SessionId); + Assert.False(string.IsNullOrEmpty(notification.SessionId)); Assert.NotNull(notification.Event); Assert.NotEmpty(notification.Event.Kind); Assert.IsType(notification.Restricted); diff --git a/dotnet/test/E2E/McpOAuthE2ETests.cs b/dotnet/test/E2E/McpOAuthE2ETests.cs index 417b7ad1bd..2bea715b7f 100644 --- a/dotnet/test/E2E/McpOAuthE2ETests.cs +++ b/dotnet/test/E2E/McpOAuthE2ETests.cs @@ -172,7 +172,7 @@ public async Task Should_Cancel_Pending_MCP_OAuth_Request() } }); - await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.Failed); + await WaitForMcpServerStatusAsync(session, serverName, McpServerStatus.NeedsAuth); Assert.NotNull(observedRequest); Assert.NotEmpty(observedRequest!.RequestId); diff --git a/dotnet/test/E2E/RpcSessionStateExtrasE2ETests.cs b/dotnet/test/E2E/RpcSessionStateExtrasE2ETests.cs index 5b1ce14843..e969df068c 100644 --- a/dotnet/test/E2E/RpcSessionStateExtrasE2ETests.cs +++ b/dotnet/test/E2E/RpcSessionStateExtrasE2ETests.cs @@ -64,19 +64,19 @@ public async Task Should_Get_And_Set_AllowAll_Permissions() var initial = await session.Rpc.Permissions.GetAllowAllAsync(); Assert.False(initial.Enabled, "Allow-all should be disabled on a fresh session."); - var enable = await session.Rpc.Permissions.SetAllowAllAsync(true); + var enable = await session.Rpc.Permissions.SetAllowAllAsync(enabled: true); Assert.True(enable.Success); Assert.True(enable.Enabled); Assert.True((await session.Rpc.Permissions.GetAllowAllAsync()).Enabled); - var disable = await session.Rpc.Permissions.SetAllowAllAsync(false); + var disable = await session.Rpc.Permissions.SetAllowAllAsync(enabled: false); Assert.True(disable.Success); Assert.False(disable.Enabled); Assert.False((await session.Rpc.Permissions.GetAllowAllAsync()).Enabled); } finally { - await session.Rpc.Permissions.SetAllowAllAsync(false); + await session.Rpc.Permissions.SetAllowAllAsync(enabled: false); } } diff --git a/dotnet/test/Unit/GitHubTelemetryTests.cs b/dotnet/test/Unit/GitHubTelemetryTests.cs index 4a41c1cb83..f82e0db6e0 100644 --- a/dotnet/test/Unit/GitHubTelemetryTests.cs +++ b/dotnet/test/Unit/GitHubTelemetryTests.cs @@ -53,6 +53,39 @@ public async Task ResumeSession_Opts_Into_Forwarding_When_Handler_Provided() Assert.True(flag.GetBoolean()); } + [Fact] + public async Task Connect_Opts_Into_Forwarding_When_Handler_Provided() + { + await using var server = await FakeTelemetryServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(server.Url), + OnGitHubTelemetry = _ => Task.CompletedTask, + }); + await client.StartAsync(); + + var connectParams = server.LastConnectParams ?? throw new InvalidOperationException("connect was not captured."); + Assert.True(connectParams.TryGetProperty("enableGitHubTelemetryForwarding", out var flag)); + Assert.True(flag.GetBoolean()); + } + + [Fact] + public async Task Connect_Does_Not_Opt_In_Without_Handler() + { + await using var server = await FakeTelemetryServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions + { + Connection = RuntimeConnection.ForUri(server.Url), + }); + await client.StartAsync(); + + var connectParams = server.LastConnectParams ?? throw new InvalidOperationException("connect was not captured."); + var present = connectParams.TryGetProperty("enableGitHubTelemetryForwarding", out var flag); + Assert.True( + !present || flag.ValueKind == JsonValueKind.Null, + "connect request should omit enableGitHubTelemetryForwarding (or send null) when no handler is registered"); + } + [Fact] public async Task CreateSession_Does_Not_Opt_In_Without_Handler() { @@ -187,6 +220,8 @@ public string Url public JsonElement? LastResumeParams { get; private set; } + public JsonElement? LastConnectParams { get; private set; } + public static Task StartAsync() { var listener = new TcpListener(IPAddress.Loopback, 0); @@ -267,12 +302,7 @@ private async Task HandleRequestAsync(Stream stream, JsonElement request, Cancel object? result = method switch { - "connect" => new Dictionary - { - ["ok"] = true, - ["protocolVersion"] = 3, - ["version"] = "test", - }, + "connect" => CaptureConnect(request), "session.create" => CaptureCreate(request), "session.resume" => CaptureResume(request), "session.send" => new Dictionary { ["messageId"] = "message-1" }, @@ -289,6 +319,17 @@ private async Task HandleRequestAsync(Stream stream, JsonElement request, Cancel }, cancellationToken); } + private Dictionary CaptureConnect(JsonElement request) + { + LastConnectParams = request.TryGetProperty("params", out var p) ? p.Clone() : null; + return new Dictionary + { + ["ok"] = true, + ["protocolVersion"] = 3, + ["version"] = "test", + }; + } + private Dictionary CaptureCreate(JsonElement request) { LastCreateParams = request.TryGetProperty("params", out var p) ? p.Clone() : null; diff --git a/go/client.go b/go/client.go index 1bbf615d7b..5357f28585 100644 --- a/go/client.go +++ b/go/client.go @@ -1685,7 +1685,15 @@ func (c *Client) verifyProtocolVersion(ctx context.Context) error { t := c.effectiveConnectionToken tokenPtr = &t } - connectResult, err := c.internalRPC.Connect(ctx, &rpc.ConnectRequest{Token: tokenPtr}) + connectReq := &connectHandshakeRequest{Token: tokenPtr} + // Opt in to GitHub telemetry forwarding at the connection level when a handler is + // registered (mirrors the runtime, which reads this flag on the `connect` handshake + // so the first session's un-replayable `session.start` event is forwarded). Also + // sent on session.create/resume for older CLIs. + if c.options.OnGitHubTelemetry != nil { + connectReq.EnableGitHubTelemetryForwarding = Bool(true) + } + rawConnectResult, err := c.client.Request(ctx, "connect", connectReq) if err != nil { var rpcErr *jsonrpc2.Error if errors.As(err, &rpcErr) && (rpcErr.Code == jsonrpc2.ErrMethodNotFound.Code || rpcErr.Message == "Unhandled method connect") { @@ -1700,6 +1708,10 @@ func (c *Client) verifyProtocolVersion(ctx context.Context) error { return err } } else { + var connectResult rpc.ConnectResult + if err := json.Unmarshal(rawConnectResult, &connectResult); err != nil { + return err + } v := int(connectResult.ProtocolVersion) serverVersion = &v } @@ -1716,6 +1728,11 @@ func (c *Client) verifyProtocolVersion(ctx context.Context) error { return nil } +type connectHandshakeRequest struct { + Token *string `json:"token,omitempty"` + EnableGitHubTelemetryForwarding *bool `json:"enableGitHubTelemetryForwarding,omitempty"` +} + // stderrBufferSize is the maximum number of bytes kept from the CLI process's // stderr. Only the tail is retained so that memory stays bounded even when the // process produces a large amount of diagnostic output. diff --git a/go/client_test.go b/go/client_test.go index f7d5f50c6c..bd48baacd7 100644 --- a/go/client_test.go +++ b/go/client_test.go @@ -2487,6 +2487,52 @@ func assertForwardingFlagAbsent(t *testing.T, params json.RawMessage) { } } +func TestClient_ForwardsGitHubTelemetryForwardingOnConnect(t *testing.T) { + rpcClient, server, _ := newRuntimeShutdownRpcPair(t) + t.Cleanup(server.Stop) + client := &Client{ + client: rpcClient, + RPC: rpc.NewServerRPC(rpcClient), + internalRPC: rpc.NewInternalServerRPC(rpcClient), + sessions: make(map[string]*Session), + options: ClientOptions{OnGitHubTelemetry: func(*rpc.GitHubTelemetryNotification) {}}, + } + + connectParams := make(chan json.RawMessage, 1) + server.SetRequestHandler("connect", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + connectParams <- append(json.RawMessage(nil), params...) + return []byte(`{"ok":true,"protocolVersion":3,"version":"test"}`), nil + }) + + if err := client.verifyProtocolVersion(t.Context()); err != nil { + t.Fatalf("verifyProtocolVersion failed: %v", err) + } + assertForwardingFlagTrue(t, <-connectParams) +} + +func TestClient_OmitsGitHubTelemetryForwardingOnConnectWhenNoHandler(t *testing.T) { + rpcClient, server, _ := newRuntimeShutdownRpcPair(t) + t.Cleanup(server.Stop) + client := &Client{ + client: rpcClient, + RPC: rpc.NewServerRPC(rpcClient), + internalRPC: rpc.NewInternalServerRPC(rpcClient), + sessions: make(map[string]*Session), + options: ClientOptions{}, + } + + connectParams := make(chan json.RawMessage, 1) + server.SetRequestHandler("connect", func(params json.RawMessage) (json.RawMessage, *jsonrpc2.Error) { + connectParams <- append(json.RawMessage(nil), params...) + return []byte(`{"ok":true,"protocolVersion":3,"version":"test"}`), nil + }) + + if err := client.verifyProtocolVersion(t.Context()); err != nil { + t.Fatalf("verifyProtocolVersion failed: %v", err) + } + assertForwardingFlagAbsent(t, <-connectParams) +} + func TestGitHubTelemetryNotificationRoutesToCallback(t *testing.T) { // The runtime forwards telemetry via a JSON-RPC *notification* (no id). // Drive a real Content-Length-framed notification through the transport and @@ -2547,8 +2593,12 @@ func TestGitHubTelemetryNotificationRoutesToCallback(t *testing.T) { select { case n := <-received: - if n.SessionID != "sess-telemetry" { - t.Errorf("session id = %q, want sess-telemetry", n.SessionID) + sessionID := "" + if n.SessionID != nil { + sessionID = *n.SessionID + } + if sessionID != "sess-telemetry" { + t.Errorf("session id = %q, want sess-telemetry", sessionID) } if !n.Restricted { t.Error("expected restricted to be true") diff --git a/go/internal/e2e/github_telemetry_e2e_test.go b/go/internal/e2e/github_telemetry_e2e_test.go index 666817451e..aa26ba31f2 100644 --- a/go/internal/e2e/github_telemetry_e2e_test.go +++ b/go/internal/e2e/github_telemetry_e2e_test.go @@ -35,7 +35,7 @@ func TestGitHubTelemetryE2E(t *testing.T) { t.Cleanup(func() { session.Disconnect() }) notification := waitForGitHubTelemetryNotification(t, &mu, ¬ifications, 30*time.Second) - if notification.SessionID == "" { + if notification.SessionID == nil || *notification.SessionID == "" { t.Fatal("Expected a non-empty SessionID") } if notification.Event.Kind == "" { diff --git a/go/internal/e2e/mcp_oauth_e2e_test.go b/go/internal/e2e/mcp_oauth_e2e_test.go index e423f12d12..1655e3bd1a 100644 --- a/go/internal/e2e/mcp_oauth_e2e_test.go +++ b/go/internal/e2e/mcp_oauth_e2e_test.go @@ -218,7 +218,7 @@ func TestMCPOAuthE2E(t *testing.T) { } t.Cleanup(func() { session.Disconnect() }) - waitForMCPServerStatus(t, session, serverName, rpc.MCPServerStatusFailed) + waitForMCPServerStatus(t, session, serverName, rpc.MCPServerStatusNeedsAuth) if observedRequest.ServerName != serverName { t.Fatalf("Expected serverName %q, got %q", serverName, observedRequest.ServerName) } diff --git a/go/internal/e2e/rpc_session_state_extras_e2e_test.go b/go/internal/e2e/rpc_session_state_extras_e2e_test.go index f4de1c1867..36f12ac548 100644 --- a/go/internal/e2e/rpc_session_state_extras_e2e_test.go +++ b/go/internal/e2e/rpc_session_state_extras_e2e_test.go @@ -70,7 +70,7 @@ func TestRpcSessionStateExtras(t *testing.T) { session := createPortedSession(t, client, nil) defer session.Disconnect() defer func() { - _, _ = session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: false}) + _, _ = session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: copilot.Bool(false)}) }() initial, err := session.RPC.Permissions.GetAllowAll(t.Context()) @@ -81,7 +81,7 @@ func TestRpcSessionStateExtras(t *testing.T) { t.Fatal("Allow-all should be disabled on a fresh session") } - enable, err := session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: true}) + enable, err := session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: copilot.Bool(true)}) if err != nil { t.Fatalf("Permissions.SetAllowAll(true) failed: %v", err) } @@ -96,7 +96,7 @@ func TestRpcSessionStateExtras(t *testing.T) { t.Fatal("Expected allow-all to be enabled") } - disable, err := session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: false}) + disable, err := session.RPC.Permissions.SetAllowAll(t.Context(), &rpc.PermissionsSetAllowAllRequest{Enabled: copilot.Bool(false)}) if err != nil { t.Fatalf("Permissions.SetAllowAll(false) failed: %v", err) } diff --git a/java/src/main/java/com/github/copilot/CopilotClient.java b/java/src/main/java/com/github/copilot/CopilotClient.java index 31a8929142..b90ccd545a 100644 --- a/java/src/main/java/com/github/copilot/CopilotClient.java +++ b/java/src/main/java/com/github/copilot/CopilotClient.java @@ -26,7 +26,7 @@ import com.github.copilot.rpc.CreateSessionResponse; import com.github.copilot.generated.rpc.SessionOptionsUpdateParams; import com.github.copilot.generated.rpc.SessionInstalledPlugin; -import com.github.copilot.generated.rpc.ConnectParams; +import com.github.copilot.generated.rpc.ConnectResult; import com.github.copilot.generated.rpc.GitHubTelemetryNotification; import com.github.copilot.generated.rpc.ServerRpc; import com.github.copilot.generated.rpc.SessionEventLogRegisterInterestParams; @@ -306,11 +306,20 @@ private void verifyProtocolVersion(Connection connection) throws Exception { Integer serverVersion; try { - // Try the new 'connect' RPC which supports connection tokens - var connectParams = new ConnectParams(effectiveConnectionToken); - var connectResponse = connection.rpc - .invoke("connect", connectParams, com.github.copilot.generated.rpc.ConnectResult.class) - .get(30, TimeUnit.SECONDS); + // Try the new 'connect' RPC which supports connection tokens. + var connectParams = new HashMap(); + if (effectiveConnectionToken != null) { + connectParams.put("token", effectiveConnectionToken); + } + // Opt into GitHub telemetry forwarding at the connection level when a handler + // is registered, so the runtime can forward the first session's un-replayable + // start event. Also sent on session create/resume for backward compatibility + // with servers that read the flag there instead. + if (this.options.getOnGitHubTelemetry() != null) { + connectParams.put("enableGitHubTelemetryForwarding", true); + } + var connectResponse = connection.rpc.invoke("connect", connectParams, ConnectResult.class).get(30, + TimeUnit.SECONDS); serverVersion = connectResponse.protocolVersion() != null ? connectResponse.protocolVersion().intValue() : null; diff --git a/java/src/test/java/com/github/copilot/GitHubTelemetryTest.java b/java/src/test/java/com/github/copilot/GitHubTelemetryTest.java index ad950b8233..8e35bd9a92 100644 --- a/java/src/test/java/com/github/copilot/GitHubTelemetryTest.java +++ b/java/src/test/java/com/github/copilot/GitHubTelemetryTest.java @@ -32,8 +32,9 @@ /** * Exercises the hand-written GitHub telemetry forwarding surface: the * {@code gitHubTelemetry.event} notification adapter, the - * {@code enableGitHubTelemetryForwarding} capability flag on the create/resume - * requests, and the {@code onGitHubTelemetry} client option. + * {@code enableGitHubTelemetryForwarding} capability flag on the connect + * handshake and the create/resume requests, and the {@code onGitHubTelemetry} + * client option. */ @AllowCopilotExperimental class GitHubTelemetryTest { @@ -146,6 +147,12 @@ void clientOptsSessionsIntoForwardingAndReceivesEvents() throws Exception { client.start().get(15, TimeUnit.SECONDS); + // Connecting must opt into telemetry forwarding at the connection level so + // the runtime can forward the first session's un-replayable start event. + JsonNode connectParams = server.awaitConnect(); + assertTrue(connectParams.path("enableGitHubTelemetryForwarding").asBoolean(), + "connect request should carry enableGitHubTelemetryForwarding=true"); + // Creating a session must opt it into telemetry forwarding. client.createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)).get(15, TimeUnit.SECONDS); @@ -178,6 +185,10 @@ void clientOmitsForwardingWhenNoHandler() throws Exception { client.start().get(15, TimeUnit.SECONDS); + JsonNode connectParams = server.awaitConnect(); + assertFalse(connectParams.has("enableGitHubTelemetryForwarding"), + "connect request should omit the flag when no handler is registered"); + client.createSession(new SessionConfig().setOnPermissionRequest(PermissionHandler.APPROVE_ALL)).get(15, TimeUnit.SECONDS); JsonNode createParams = server.awaitCreate(); @@ -214,6 +225,7 @@ private static final class FakeRuntimeServer implements AutoCloseable { private final ServerSocket serverSocket; private final Thread acceptThread; private final CompletableFuture ready = new CompletableFuture<>(); + private final CompletableFuture connectParams = new CompletableFuture<>(); private final CompletableFuture createParams = new CompletableFuture<>(); private final CompletableFuture resumeParams = new CompletableFuture<>(); @@ -228,6 +240,10 @@ String url() { return "127.0.0.1:" + serverSocket.getLocalPort(); } + JsonNode awaitConnect() throws Exception { + return connectParams.get(15, TimeUnit.SECONDS); + } + JsonNode awaitCreate() throws Exception { return createParams.get(15, TimeUnit.SECONDS); } @@ -244,8 +260,10 @@ private void acceptLoop() { try { Socket socket = serverSocket.accept(); JsonRpcClient server = JsonRpcClient.fromSocket(socket); - server.registerMethodHandler("connect", - (id, params) -> respond(server, id, Map.of("protocolVersion", 2))); + server.registerMethodHandler("connect", (id, params) -> { + connectParams.complete(params); + respond(server, id, Map.of("protocolVersion", 2)); + }); server.registerMethodHandler("session.create", (id, params) -> { createParams.complete(params); respond(server, id, Map.of("sessionId", params.path("sessionId").asText("created"), "workspacePath", @@ -261,6 +279,7 @@ private void acceptLoop() { ready.complete(server); } catch (IOException e) { ready.completeExceptionally(e); + connectParams.completeExceptionally(e); createParams.completeExceptionally(e); resumeParams.completeExceptionally(e); } diff --git a/java/src/test/java/com/github/copilot/McpOAuthE2ETest.java b/java/src/test/java/com/github/copilot/McpOAuthE2ETest.java index 8da3b08b47..1e602caaa1 100644 --- a/java/src/test/java/com/github/copilot/McpOAuthE2ETest.java +++ b/java/src/test/java/com/github/copilot/McpOAuthE2ETest.java @@ -182,7 +182,7 @@ void testShouldCancelPendingMcpOauthRequest() throws Exception { }).setMcpServers(Map.of(serverName, new McpHttpServerConfig() .setUrl(oauthServer.url() + "/mcp").setTools(List.of("*"))))) .get()) { - waitForMcpServerStatus(session, serverName, McpServerStatus.FAILED, observedRequest); + waitForMcpServerStatus(session, serverName, McpServerStatus.NEEDS_AUTH, observedRequest); } var request = observedRequest.get(); diff --git a/nodejs/src/client.ts b/nodejs/src/client.ts index 160a12d480..9f430600ca 100644 --- a/nodejs/src/client.ts +++ b/nodejs/src/client.ts @@ -1846,9 +1846,18 @@ export class CopilotClient { let serverVersion: number | undefined; try { - const result = await raceAgainstExit( - this.internalRpc.connect({ token: this.effectiveConnectionToken }) - ); + const connectParams: { + token?: string; + enableGitHubTelemetryForwarding?: boolean; + } = { token: this.effectiveConnectionToken }; + // Opt in to GitHub telemetry forwarding at the connection level when a + // handler is registered (mirrors the runtime, which reads this flag on the + // `connect` handshake so the first session's un-replayable `session.start` + // event is forwarded). Also sent on session.create/resume for older CLIs. + if (this.onGitHubTelemetry != null) { + connectParams.enableGitHubTelemetryForwarding = true; + } + const result = await raceAgainstExit(this.internalRpc.connect(connectParams)); serverVersion = result.protocolVersion; } catch (err) { if ( diff --git a/nodejs/test/client.test.ts b/nodejs/test/client.test.ts index b174494548..96c32a5951 100644 --- a/nodejs/test/client.test.ts +++ b/nodejs/test/client.test.ts @@ -488,6 +488,40 @@ describe("CopilotClient", () => { expect(resumePayload.enableGitHubTelemetryForwarding).toBe(true); }); + it("opts into GitHub telemetry forwarding on the connect handshake when a handler is provided", async () => { + const client = new CopilotClient({ onGitHubTelemetry: () => {} }); + onTestFinished(() => client.forceStop()); + + const sendRequest = vi.fn(async (method: string) => { + if (method === "connect") return { ok: true, protocolVersion: 3, version: "test" }; + throw new Error(`Unexpected method: ${method}`); + }); + (client as any).connection = { sendRequest }; + + await (client as any).verifyProtocolVersion(); + + const connectCall = sendRequest.mock.calls.find(([method]) => method === "connect"); + expect(connectCall).toBeDefined(); + expect((connectCall![1] as any).enableGitHubTelemetryForwarding).toBe(true); + }); + + it("does not opt into GitHub telemetry forwarding on the connect handshake without a handler", async () => { + const client = new CopilotClient(); + onTestFinished(() => client.forceStop()); + + const sendRequest = vi.fn(async (method: string) => { + if (method === "connect") return { ok: true, protocolVersion: 3, version: "test" }; + throw new Error(`Unexpected method: ${method}`); + }); + (client as any).connection = { sendRequest }; + + await (client as any).verifyProtocolVersion(); + + const connectCall = sendRequest.mock.calls.find(([method]) => method === "connect"); + expect(connectCall).toBeDefined(); + expect((connectCall![1] as any).enableGitHubTelemetryForwarding).toBeUndefined(); + }); + it("does not opt into GitHub telemetry forwarding without a handler", async () => { const client = new CopilotClient(); await client.start(); diff --git a/nodejs/test/e2e/mcp_oauth.e2e.test.ts b/nodejs/test/e2e/mcp_oauth.e2e.test.ts index 29ed089edb..509932f971 100644 --- a/nodejs/test/e2e/mcp_oauth.e2e.test.ts +++ b/nodejs/test/e2e/mcp_oauth.e2e.test.ts @@ -193,7 +193,7 @@ describe("MCP OAuth host auth", async () => { }); onTestFinished(() => disconnectSession(session)); - await waitForMcpServerStatus(session, serverName, "failed"); + await waitForMcpServerStatus(session, serverName, "needs-auth"); expect(authRequest).toMatchObject({ serverName, diff --git a/python/copilot/client.py b/python/copilot/client.py index 269aaf96ce..55d01c5b57 100644 --- a/python/copilot/client.py +++ b/python/copilot/client.py @@ -70,8 +70,7 @@ OpenCanvasInstance, RemoteSessionMode, ServerRpc, - _ConnectRequest, - _InternalServerRpc, + _ConnectResult, from_datetime, register_client_global_api_handlers, register_client_session_api_handlers, @@ -3303,8 +3302,17 @@ async def _verify_protocol_version(self) -> None: server_version: int | None try: - connect_result = await _InternalServerRpc(self._client)._connect( - _ConnectRequest(token=self._effective_connection_token) + connect_params: dict[str, Any] = {} + if self._effective_connection_token is not None: + connect_params["token"] = self._effective_connection_token + # Opt in to GitHub telemetry forwarding at the connection level when a + # handler is registered (mirrors the runtime, which reads this flag on the + # `connect` handshake so the first session's un-replayable `session.start` + # event is forwarded). Also sent on session.create/resume for older CLIs. + if self._on_github_telemetry is not None: + connect_params["enableGitHubTelemetryForwarding"] = True + connect_result = _ConnectResult.from_dict( + await self._client.request("connect", connect_params) ) server_version = connect_result.protocol_version except JsonRpcError as err: diff --git a/python/e2e/test_mcp_oauth_e2e.py b/python/e2e/test_mcp_oauth_e2e.py index 47897f69c0..6c33165000 100644 --- a/python/e2e/test_mcp_oauth_e2e.py +++ b/python/e2e/test_mcp_oauth_e2e.py @@ -249,7 +249,7 @@ def on_mcp_auth_request(request, _invocation): on_mcp_auth_request=on_mcp_auth_request, mcp_servers=mcp_servers, ) as session: - await _wait_for_mcp_server_status(session, server_name, McpServerStatus.FAILED) + await _wait_for_mcp_server_status(session, server_name, McpServerStatus.NEEDS_AUTH) assert observed_request is not None assert observed_request["serverName"] == server_name diff --git a/python/test_client.py b/python/test_client.py index 13fc50e73f..5e1b8be634 100644 --- a/python/test_client.py +++ b/python/test_client.py @@ -2381,6 +2381,37 @@ async def mock_request(method, params, **kwargs): finally: await client.force_stop() + @pytest.mark.asyncio + async def test_connect_enables_forwarding_when_handler_registered(self): + client = CopilotClient( + connection=RuntimeConnection.for_stdio(path=CLI_PATH), + on_github_telemetry=lambda _notification: None, + ) + captured = {} + + class _FakeClient: + async def request(self, method, params, **kwargs): + captured[method] = params + return {"ok": True, "protocolVersion": 3, "version": "test"} + + client._client = _FakeClient() + await client._verify_protocol_version() + assert captured["connect"]["enableGitHubTelemetryForwarding"] is True + + @pytest.mark.asyncio + async def test_connect_omits_forwarding_without_handler(self): + client = CopilotClient(connection=RuntimeConnection.for_stdio(path=CLI_PATH)) + captured = {} + + class _FakeClient: + async def request(self, method, params, **kwargs): + captured[method] = params + return {"ok": True, "protocolVersion": 3, "version": "test"} + + client._client = _FakeClient() + await client._verify_protocol_version() + assert "enableGitHubTelemetryForwarding" not in captured["connect"] + @pytest.mark.asyncio async def test_event_routes_to_handler(self): from copilot.generated.rpc import GitHubTelemetryNotification diff --git a/rust/src/errors.rs b/rust/src/errors.rs index 5690f6412c..6e05bbfae1 100644 --- a/rust/src/errors.rs +++ b/rust/src/errors.rs @@ -63,6 +63,12 @@ pub enum ProtocolErrorKind { max: u32, }, + /// The CLI server reported a protocol version that can't be represented by the SDK. + InvalidProtocolVersion { + /// Version reported by the server. + server: i64, + }, + /// The CLI server's protocol version changed between calls. VersionChanged { /// Previously negotiated version. @@ -94,6 +100,9 @@ impl fmt::Display for ProtocolErrorKind { "version mismatch: server={server}, supported={min}\u{2013}{max}" ) } + ProtocolErrorKind::InvalidProtocolVersion { server } => { + write!(f, "invalid protocol version: server={server}") + } ProtocolErrorKind::VersionChanged { previous, current } => { write!(f, "version changed: was {previous}, now {current}") } diff --git a/rust/src/lib.rs b/rust/src/lib.rs index c31e80dc52..0333281f59 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -1818,13 +1818,26 @@ impl Client { /// param. Server-side, the token is required when the server was /// started with `COPILOT_CONNECTION_TOKEN`. async fn connect_handshake(&self) -> Result> { - let result = self - .rpc() - .connect(crate::generated::api_types::ConnectRequest { - token: self.inner.effective_connection_token.clone(), - }) + let params = crate::generated::api_types::ConnectRequest { + token: self.inner.effective_connection_token.clone(), + enable_git_hub_telemetry_forwarding: self + .inner + .on_github_telemetry + .is_some() + .then_some(true), + }; + let value = self + .call( + crate::generated::api_types::rpc_methods::CONNECT, + Some(serde_json::to_value(params)?), + ) .await?; - Ok(u32::try_from(result.protocol_version).ok()) + let result: crate::generated::api_types::ConnectResult = serde_json::from_value(value)?; + Ok(Some(u32::try_from(result.protocol_version).map_err( + |_| ProtocolErrorKind::InvalidProtocolVersion { + server: result.protocol_version, + }, + )?)) } /// Send a `ping` RPC and return the typed [`PingResponse`]. diff --git a/rust/tests/e2e/github_telemetry.rs b/rust/tests/e2e/github_telemetry.rs index 26e2a3f94b..2047ee34ff 100644 --- a/rust/tests/e2e/github_telemetry.rs +++ b/rust/tests/e2e/github_telemetry.rs @@ -47,7 +47,12 @@ async fn should_forward_github_telemetry_on_session_create() { let first = notifications .first() .expect("github telemetry notification"); - assert!(!first.session_id.is_empty()); + assert!( + first + .session_id + .as_deref() + .is_some_and(|session_id| !session_id.is_empty()) + ); let _: bool = first.restricted; assert!(!first.event.kind.is_empty()); } diff --git a/rust/tests/e2e/mcp_oauth.rs b/rust/tests/e2e/mcp_oauth.rs index b1d932372c..98d4cf0309 100644 --- a/rust/tests/e2e/mcp_oauth.rs +++ b/rust/tests/e2e/mcp_oauth.rs @@ -217,7 +217,7 @@ async fn should_cancel_pending_mcp_oauth_request() { .await .expect("create session"); - wait_for_mcp_server_status(&session, server_name, McpServerStatus::Failed).await; + wait_for_mcp_server_status(&session, server_name, McpServerStatus::NeedsAuth).await; let request = handler .request diff --git a/rust/tests/e2e/rpc_session_state_extras.rs b/rust/tests/e2e/rpc_session_state_extras.rs index 10954b4e27..b8b8073a38 100644 --- a/rust/tests/e2e/rpc_session_state_extras.rs +++ b/rust/tests/e2e/rpc_session_state_extras.rs @@ -102,7 +102,9 @@ async fn should_get_and_set_allowall_permissions() { .rpc() .permissions() .set_allow_all(PermissionsSetAllowAllRequest { - enabled: true, + enabled: Some(true), + mode: None, + model: None, source: None, }) .await @@ -123,7 +125,9 @@ async fn should_get_and_set_allowall_permissions() { .rpc() .permissions() .set_allow_all(PermissionsSetAllowAllRequest { - enabled: false, + enabled: Some(false), + mode: None, + model: None, source: None, }) .await diff --git a/rust/tests/session_test.rs b/rust/tests/session_test.rs index 08f8a7653e..2599ea6d3a 100644 --- a/rust/tests/session_test.rs +++ b/rust/tests/session_test.rs @@ -25,7 +25,7 @@ use github_copilot_sdk::types::{ MessageOptions, RequestId, SessionConfig, SessionId, SetModelOptions, Tool, ToolInvocation, ToolResult, }; -use github_copilot_sdk::{Client, ContextTier, tool}; +use github_copilot_sdk::{Client, ContextTier, ErrorKind, ProtocolErrorKind, tool}; use serde_json::Value; use tokio::io::{AsyncWrite, AsyncWriteExt, duplex}; use tokio::time::timeout; @@ -911,6 +911,94 @@ async fn resume_session_omits_github_telemetry_forwarding_without_callback() { timeout(TIMEOUT, resume_handle).await.unwrap().unwrap(); } +#[tokio::test] +async fn connect_sends_github_telemetry_forwarding_when_callback_registered() { + let callback: github_copilot_sdk::github_telemetry::GitHubTelemetryCallback = + Arc::new(|_notification| {}); + let (client, mut server_read, mut server_write) = make_client_with_telemetry(callback); + + let handle = tokio::spawn({ + let client = client.clone(); + async move { client.verify_protocol_version().await.unwrap() } + }); + + let request = read_framed(&mut server_read).await; + assert_eq!(request["method"], "connect"); + assert_eq!(request["params"]["enableGitHubTelemetryForwarding"], true); + + let id = request["id"].as_u64().unwrap(); + let response = serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "result": { "ok": true, "protocolVersion": 3, "version": "test" }, + }); + write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await; + timeout(TIMEOUT, handle).await.unwrap().unwrap(); +} + +#[tokio::test] +async fn connect_omits_github_telemetry_forwarding_without_callback() { + let (client, mut server_read, mut server_write) = make_client(); + + let handle = tokio::spawn({ + let client = client.clone(); + async move { client.verify_protocol_version().await.unwrap() } + }); + + let request = read_framed(&mut server_read).await; + assert_eq!(request["method"], "connect"); + assert!( + request["params"] + .get("enableGitHubTelemetryForwarding") + .is_none_or(Value::is_null), + "forwarding flag should be omitted when no callback is registered" + ); + + let id = request["id"].as_u64().unwrap(); + let response = serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "result": { "ok": true, "protocolVersion": 3, "version": "test" }, + }); + write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await; + timeout(TIMEOUT, handle).await.unwrap().unwrap(); +} + +#[tokio::test] +async fn connect_rejects_invalid_protocol_version_values() { + for protocol_version in [-1, i64::from(u32::MAX) + 1] { + let (client, mut server_read, mut server_write) = make_client(); + + let handle = tokio::spawn({ + let client = client.clone(); + async move { client.verify_protocol_version().await } + }); + + let request = read_framed(&mut server_read).await; + assert_eq!(request["method"], "connect"); + + let id = request["id"].as_u64().unwrap(); + let response = serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "result": { "ok": true, "protocolVersion": protocol_version, "version": "test" }, + }); + write_framed(&mut server_write, &serde_json::to_vec(&response).unwrap()).await; + + let err = timeout(TIMEOUT, handle) + .await + .unwrap() + .unwrap() + .unwrap_err(); + match err.kind() { + ErrorKind::Protocol(ProtocolErrorKind::InvalidProtocolVersion { server }) => { + assert_eq!(*server, protocol_version); + } + other => panic!("unexpected error kind: {other:?}"), + } + } +} + #[tokio::test] async fn github_telemetry_event_dispatches_to_callback() { use github_copilot_sdk::github_telemetry::GitHubTelemetryNotification; @@ -965,7 +1053,7 @@ async fn github_telemetry_event_dispatches_to_callback() { .await; let received = timeout(TIMEOUT, rx.recv()).await.unwrap().unwrap(); - assert_eq!(received.session_id, session_id); + assert_eq!(received.session_id.as_deref(), Some(session_id.as_str())); assert!(!received.restricted); assert_eq!(received.event.kind, "tool_call_executed"); assert_eq!(