@github/copilot-sdk 1.0.2 → 1.0.4

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/dist/client.js CHANGED
@@ -15,11 +15,13 @@ import {
15
15
  import {
16
16
  createServerRpc,
17
17
  createInternalServerRpc,
18
+ registerClientGlobalApiHandlers,
18
19
  registerClientSessionApiHandlers
19
20
  } from "./generated/rpc.js";
20
21
  import { getSdkProtocolVersion } from "./sdkProtocolVersion.js";
21
22
  import { CopilotSession } from "./session.js";
22
23
  import { createSessionFsAdapter } from "./sessionFsProvider.js";
24
+ import { createCopilotRequestAdapter } from "./copilotRequestHandler.js";
23
25
  import { getTraceContext } from "./telemetry.js";
24
26
  import { ToolSet } from "./toolSet.js";
25
27
  import { defaultJoinSessionPermissionHandler } from "./types.js";
@@ -79,6 +81,32 @@ function toJsonSchema(parameters) {
79
81
  }
80
82
  return parameters;
81
83
  }
84
+ const DEFAULT_PROVIDER_NAME = "default";
85
+ function extractBearerTokenProviders(provider, providers) {
86
+ const callbacks = /* @__PURE__ */ new Map();
87
+ let wireProvider = provider;
88
+ if (provider?.getBearerToken) {
89
+ const { getBearerToken, ...rest } = provider;
90
+ callbacks.set(DEFAULT_PROVIDER_NAME, getBearerToken);
91
+ wireProvider = {
92
+ ...rest,
93
+ hasBearerTokenProvider: true
94
+ };
95
+ }
96
+ let wireProviders = providers;
97
+ if (providers?.some((p) => p.getBearerToken)) {
98
+ wireProviders = providers.map((p) => {
99
+ if (!p.getBearerToken) return p;
100
+ const { getBearerToken, ...rest } = p;
101
+ callbacks.set(p.name, getBearerToken);
102
+ return {
103
+ ...rest,
104
+ hasBearerTokenProvider: true
105
+ };
106
+ });
107
+ }
108
+ return { wireProvider, wireProviders, callbacks };
109
+ }
82
110
  function toWireMcpServers(mcpServers) {
83
111
  if (!mcpServers) return void 0;
84
112
  return Object.fromEntries(
@@ -156,28 +184,57 @@ function getNodeExecPath() {
156
184
  }
157
185
  return process.execPath;
158
186
  }
187
+ function getCliPlatformPackageNames() {
188
+ const arch = process.arch;
189
+ const variants = process.platform === "linux" ? ["linux", "linuxmusl"] : [process.platform];
190
+ return variants.map((variant) => `@github/copilot-${variant}-${arch}`);
191
+ }
159
192
  function getBundledCliPath() {
193
+ const packageNames = getCliPlatformPackageNames();
160
194
  if (typeof import.meta.resolve === "function") {
161
- const sdkUrl = import.meta.resolve("@github/copilot/sdk");
162
- const sdkPath = fileURLToPath(sdkUrl);
163
- return join(dirname(dirname(sdkPath)), "index.js");
195
+ for (const packageName of packageNames) {
196
+ try {
197
+ const sdkUrl = import.meta.resolve(`${packageName}/sdk`);
198
+ const sdkPath = fileURLToPath(sdkUrl);
199
+ return join(dirname(dirname(sdkPath)), "index.js");
200
+ } catch {
201
+ }
202
+ }
203
+ throw new Error(
204
+ `Could not resolve a @github/copilot platform package (tried ${packageNames.join(", ")}). Ensure @github/copilot is installed, or pass cliPath/cliUrl to CopilotClient.`
205
+ );
164
206
  }
165
207
  const req = createRequire(__filename);
166
208
  const searchPaths = req.resolve.paths("@github/copilot") ?? [];
167
209
  for (const base of searchPaths) {
168
- const candidate = join(base, "@github", "copilot", "index.js");
169
- if (existsSync(candidate)) {
170
- return candidate;
210
+ for (const packageName of packageNames) {
211
+ const candidate = join(base, ...packageName.split("/"), "index.js");
212
+ if (existsSync(candidate)) {
213
+ return candidate;
214
+ }
171
215
  }
172
216
  }
173
217
  throw new Error(
174
- `Could not find @github/copilot package. Searched ${searchPaths.length} paths. Ensure it is installed, or pass cliPath/cliUrl to CopilotClient.`
218
+ `Could not find a @github/copilot platform package (tried ${packageNames.join(", ")}). Searched ${searchPaths.length} paths. Ensure @github/copilot is installed, or pass cliPath/cliUrl to CopilotClient.`
175
219
  );
176
220
  }
221
+ class TeardownResilientStreamMessageWriter extends StreamMessageWriter {
222
+ suppressWriteErrors = false;
223
+ async write(msg) {
224
+ try {
225
+ await super.write(msg);
226
+ } catch (error) {
227
+ if (!this.suppressWriteErrors) {
228
+ throw error;
229
+ }
230
+ }
231
+ }
232
+ }
177
233
  class CopilotClient {
178
234
  cliStartTimeout = null;
179
235
  cliProcess = null;
180
236
  connection = null;
237
+ messageWriter = null;
181
238
  socket = null;
182
239
  runtimePort = null;
183
240
  actualHost = "localhost";
@@ -209,6 +266,8 @@ class CopilotClient {
209
266
  negotiatedProtocolVersion = null;
210
267
  /** Connection-level session filesystem config, set via constructor option. */
211
268
  sessionFsConfig = null;
269
+ requestHandler = null;
270
+ llmInferenceHandlers = {};
212
271
  /**
213
272
  * Typed server-scoped RPC methods.
214
273
  * @throws Error if the client is not connected
@@ -239,6 +298,13 @@ class CopilotClient {
239
298
  const level = this.options.logLevel?.toLowerCase();
240
299
  if (level === "debug" || level === "all") {
241
300
  process.stderr.write(`[copilot-sdk] ${message}. Elapsed=${Date.now() - startMs}ms
301
+ `);
302
+ }
303
+ }
304
+ logDebug(message) {
305
+ const level = this.options.logLevel?.toLowerCase();
306
+ if (level === "debug" || level === "all") {
307
+ process.stderr.write(`[copilot-sdk] ${message}
242
308
  `);
243
309
  }
244
310
  }
@@ -301,6 +367,8 @@ class CopilotClient {
301
367
  this.onListModels = options.onListModels;
302
368
  this.onGetTraceContext = options.onGetTraceContext;
303
369
  this.sessionFsConfig = options.sessionFs ?? null;
370
+ this.requestHandler = options.requestHandler ?? null;
371
+ this.setupLlmInference();
304
372
  const effectiveEnv = options.env ?? process.env;
305
373
  this.resolvedEnv = effectiveEnv;
306
374
  this.resolvedCliPath = conn.kind === "stdio" || conn.kind === "tcp" ? conn.path ?? effectiveEnv.COPILOT_CLI_PATH ?? getBundledCliPath() : void 0;
@@ -380,6 +448,20 @@ class CopilotClient {
380
448
  }
381
449
  session.clientSessionApis.sessionFs = createSessionFsAdapter(provider);
382
450
  }
451
+ setupLlmInference() {
452
+ if (!this.requestHandler) {
453
+ return;
454
+ }
455
+ this.llmInferenceHandlers = {
456
+ llmInference: createCopilotRequestAdapter(this.requestHandler, () => {
457
+ if (!this.connection) {
458
+ return void 0;
459
+ }
460
+ this._rpc ??= createServerRpc(this.connection);
461
+ return this._rpc;
462
+ })
463
+ };
464
+ }
383
465
  /**
384
466
  * Starts the CLI server and establishes a connection.
385
467
  *
@@ -417,6 +499,9 @@ class CopilotClient {
417
499
  capabilities: this.sessionFsConfig.capabilities
418
500
  });
419
501
  }
502
+ if (this.requestHandler) {
503
+ await this.connection.sendRequest("llmInference.setProvider", {});
504
+ }
420
505
  this.state = "connected";
421
506
  } catch (error) {
422
507
  this.state = "error";
@@ -449,7 +534,8 @@ class CopilotClient {
449
534
  */
450
535
  async stop() {
451
536
  const errors = [];
452
- for (const session of this.sessions.values()) {
537
+ const activeSessions = [...this.sessions.values()];
538
+ for (const session of activeSessions) {
453
539
  const sessionId = session.sessionId;
454
540
  let lastError = null;
455
541
  for (let attempt = 1; attempt <= 3; attempt++) {
@@ -473,8 +559,10 @@ class CopilotClient {
473
559
  );
474
560
  }
475
561
  }
562
+ for (const session of activeSessions) {
563
+ session._markDisconnected();
564
+ }
476
565
  this.sessions.clear();
477
- let runtimeShutdownCompleted = false;
478
566
  if (this.connection && this.cliProcess && !this.isExternalServer) {
479
567
  const runtimeShutdownStart = Date.now();
480
568
  const shutdownPromise = this.rpc.runtime.shutdown();
@@ -485,7 +573,6 @@ class CopilotClient {
485
573
  RUNTIME_SHUTDOWN_TIMEOUT_MS,
486
574
  `runtime.shutdown timed out after ${RUNTIME_SHUTDOWN_TIMEOUT_MS}ms`
487
575
  );
488
- runtimeShutdownCompleted = true;
489
576
  this.logDebugTiming(
490
577
  "CopilotClient.stop runtime shutdown complete",
491
578
  runtimeShutdownStart
@@ -502,6 +589,9 @@ class CopilotClient {
502
589
  );
503
590
  }
504
591
  }
592
+ if (this.messageWriter) {
593
+ this.messageWriter.suppressWriteErrors = true;
594
+ }
505
595
  if (this.connection) {
506
596
  try {
507
597
  this.connection.dispose();
@@ -513,6 +603,7 @@ class CopilotClient {
513
603
  );
514
604
  }
515
605
  this.connection = null;
606
+ this.messageWriter = null;
516
607
  this._rpc = null;
517
608
  this._internalRpc = null;
518
609
  }
@@ -540,16 +631,13 @@ class CopilotClient {
540
631
  this.cliProcess = null;
541
632
  try {
542
633
  if (child.exitCode == null && child.signalCode == null) {
543
- const exitedGracefully = runtimeShutdownCompleted ? await waitForChildExit(child, RUNTIME_SHUTDOWN_TIMEOUT_MS) : false;
544
- if (!exitedGracefully) {
545
- child.kill();
546
- if (!await waitForChildExit(child, RUNTIME_SHUTDOWN_TIMEOUT_MS)) {
547
- errors.push(
548
- new Error(
549
- `Timed out waiting for CLI process to exit after kill: ${RUNTIME_SHUTDOWN_TIMEOUT_MS}ms`
550
- )
551
- );
552
- }
634
+ child.kill();
635
+ if (!await waitForChildExit(child, RUNTIME_SHUTDOWN_TIMEOUT_MS)) {
636
+ errors.push(
637
+ new Error(
638
+ `Timed out waiting for CLI process to exit after kill: ${RUNTIME_SHUTDOWN_TIMEOUT_MS}ms`
639
+ )
640
+ );
553
641
  }
554
642
  }
555
643
  } catch (error) {
@@ -612,13 +700,20 @@ class CopilotClient {
612
700
  */
613
701
  async forceStop() {
614
702
  this.forceStopping = true;
703
+ for (const session of this.sessions.values()) {
704
+ session._markDisconnected();
705
+ }
615
706
  this.sessions.clear();
707
+ if (this.messageWriter) {
708
+ this.messageWriter.suppressWriteErrors = true;
709
+ }
616
710
  if (this.connection) {
617
711
  try {
618
712
  this.connection.dispose();
619
713
  } catch {
620
714
  }
621
715
  this.connection = null;
716
+ this.messageWriter = null;
622
717
  this._rpc = null;
623
718
  this._internalRpc = null;
624
719
  }
@@ -804,6 +899,11 @@ class CopilotClient {
804
899
  const callerSessionId = config.sessionId;
805
900
  const useServerGeneratedId = config.cloud != null && callerSessionId == null;
806
901
  const localSessionId = useServerGeneratedId ? void 0 : callerSessionId ?? randomUUID();
902
+ const {
903
+ wireProvider: bearerWireProvider,
904
+ wireProviders: bearerWireProviders,
905
+ callbacks: bearerTokenCallbacks
906
+ } = extractBearerTokenProviders(config.provider, config.providers);
807
907
  const { wirePayload: wireSystemMessage, transformCallbacks } = extractTransformCallbacks(
808
908
  config.systemMessage
809
909
  );
@@ -817,6 +917,9 @@ class CopilotClient {
817
917
  s.registerTools(config.tools);
818
918
  s.registerCanvases(config.canvases);
819
919
  s.registerCommands(config.commands);
920
+ if (bearerTokenCallbacks.size > 0) {
921
+ s.registerBearerTokenProviders(bearerTokenCallbacks);
922
+ }
820
923
  s.registerPermissionHandler(config.onPermissionRequest);
821
924
  if (config.onUserInputRequest) {
822
925
  s.registerUserInputHandler(config.onUserInputRequest);
@@ -880,7 +983,10 @@ class CopilotClient {
880
983
  availableTools: toolFilterOptions.availableTools,
881
984
  excludedTools: toolFilterOptions.excludedTools,
882
985
  toolFilterPrecedence: toolFilterOptions.toolFilterPrecedence,
883
- provider: config.provider,
986
+ provider: bearerWireProvider,
987
+ capi: config.capi,
988
+ providers: bearerWireProviders,
989
+ models: config.models,
884
990
  enableSessionTelemetry: config.enableSessionTelemetry,
885
991
  modelCapabilities: config.modelCapabilities,
886
992
  largeOutput: toWireLargeOutput(config.largeOutput),
@@ -918,7 +1024,8 @@ class CopilotClient {
918
1024
  memory: config.memory,
919
1025
  gitHubToken: config.gitHubToken,
920
1026
  remoteSession: config.remoteSession,
921
- cloud: config.cloud
1027
+ cloud: config.cloud,
1028
+ expAssignments: config.expAssignments
922
1029
  });
923
1030
  const {
924
1031
  sessionId: returnedSessionId,
@@ -985,6 +1092,14 @@ class CopilotClient {
985
1092
  session.registerTools(config.tools);
986
1093
  session.registerCanvases(config.canvases);
987
1094
  session.registerCommands(config.commands);
1095
+ const {
1096
+ wireProvider: bearerWireProvider,
1097
+ wireProviders: bearerWireProviders,
1098
+ callbacks: bearerTokenCallbacks
1099
+ } = extractBearerTokenProviders(config.provider, config.providers);
1100
+ if (bearerTokenCallbacks.size > 0) {
1101
+ session.registerBearerTokenProviders(bearerTokenCallbacks);
1102
+ }
988
1103
  session.registerPermissionHandler(config.onPermissionRequest);
989
1104
  if (config.onUserInputRequest) {
990
1105
  session.registerUserInputHandler(config.onUserInputRequest);
@@ -1046,7 +1161,10 @@ class CopilotClient {
1046
1161
  name: cmd.name,
1047
1162
  description: cmd.description
1048
1163
  })),
1049
- provider: config.provider,
1164
+ provider: bearerWireProvider,
1165
+ capi: config.capi,
1166
+ providers: bearerWireProviders,
1167
+ models: config.models,
1050
1168
  modelCapabilities: config.modelCapabilities,
1051
1169
  largeOutput: toWireLargeOutput(config.largeOutput),
1052
1170
  requestPermission: config.onPermissionRequest !== defaultJoinSessionPermissionHandler,
@@ -1085,7 +1203,8 @@ class CopilotClient {
1085
1203
  continuePendingWork: config.continuePendingWork,
1086
1204
  gitHubToken: config.gitHubToken,
1087
1205
  remoteSession: config.remoteSession,
1088
- openCanvases: config.openCanvases
1206
+ openCanvases: config.openCanvases,
1207
+ expAssignments: config.expAssignments
1089
1208
  });
1090
1209
  const { workspacePath, capabilities, openCanvases } = response;
1091
1210
  session["_workspacePath"] = workspacePath;
@@ -1627,13 +1746,21 @@ stderr: ${stderrOutput}`
1627
1746
  throw new Error("CLI process not started");
1628
1747
  }
1629
1748
  this.cliProcess.stdin?.on("error", (err) => {
1630
- if (!this.forceStopping) {
1631
- throw err;
1749
+ if (this.forceStopping) {
1750
+ return;
1751
+ }
1752
+ this.state = "error";
1753
+ const reason = err instanceof Error ? err.stack ?? err.message : String(err);
1754
+ this.logDebug(`stdin pipe error: ${reason}`);
1755
+ try {
1756
+ this.connection?.dispose();
1757
+ } catch {
1632
1758
  }
1633
1759
  });
1760
+ this.messageWriter = new TeardownResilientStreamMessageWriter(this.cliProcess.stdin);
1634
1761
  this.connection = createMessageConnection(
1635
1762
  new StreamMessageReader(this.cliProcess.stdout),
1636
- new StreamMessageWriter(this.cliProcess.stdin)
1763
+ this.messageWriter
1637
1764
  );
1638
1765
  this.attachConnectionHandlers();
1639
1766
  this.connection.listen();
@@ -1645,9 +1772,10 @@ stderr: ${stderrOutput}`
1645
1772
  if (this.cliProcess) {
1646
1773
  throw new Error("CLI child process was unexpectedly started in parent process mode");
1647
1774
  }
1775
+ this.messageWriter = new TeardownResilientStreamMessageWriter(process.stdout);
1648
1776
  this.connection = createMessageConnection(
1649
1777
  new StreamMessageReader(process.stdin),
1650
- new StreamMessageWriter(process.stdout)
1778
+ this.messageWriter
1651
1779
  );
1652
1780
  this.attachConnectionHandlers();
1653
1781
  this.connection.listen();
@@ -1667,9 +1795,10 @@ stderr: ${stderrOutput}`
1667
1795
  }, 1e4);
1668
1796
  this.socket.connect(this.runtimePort, this.actualHost, () => {
1669
1797
  clearTimeout(connectionTimeout);
1798
+ this.messageWriter = new TeardownResilientStreamMessageWriter(this.socket);
1670
1799
  this.connection = createMessageConnection(
1671
1800
  new StreamMessageReader(this.socket),
1672
- new StreamMessageWriter(this.socket)
1801
+ this.messageWriter
1673
1802
  );
1674
1803
  this.attachConnectionHandlers();
1675
1804
  this.connection.listen();
@@ -1717,6 +1846,7 @@ stderr: ${stderrOutput}`
1717
1846
  if (!session) throw new Error(`No session found for sessionId: ${sessionId}`);
1718
1847
  return session.clientSessionApis;
1719
1848
  });
1849
+ registerClientGlobalApiHandlers(this.connection, this.llmInferenceHandlers);
1720
1850
  this.connection.onClose(() => {
1721
1851
  this.state = "disconnected";
1722
1852
  });
@@ -0,0 +1,82 @@
1
+ import type { LlmInferenceHeaders } from "./generated/rpc.js";
2
+ declare const kSuppressCloseOnDispose: unique symbol;
3
+ /**
4
+ * Per-request context handed to every {@link CopilotRequestHandler} hook.
5
+ *
6
+ * @experimental
7
+ */
8
+ export interface CopilotRequestContext {
9
+ readonly requestId: string;
10
+ readonly sessionId?: string;
11
+ readonly transport: "http" | "websocket";
12
+ url: string;
13
+ headers: LlmInferenceHeaders;
14
+ readonly signal: AbortSignal;
15
+ }
16
+ /**
17
+ * Terminal status for a callback-owned WebSocket connection.
18
+ *
19
+ * @experimental
20
+ */
21
+ export declare class CopilotWebSocketCloseStatus {
22
+ readonly description?: string | undefined;
23
+ readonly errorCode?: string | undefined;
24
+ readonly error?: Error | undefined;
25
+ static readonly normalClosure: CopilotWebSocketCloseStatus;
26
+ constructor(description?: string | undefined, errorCode?: string | undefined, error?: Error | undefined);
27
+ }
28
+ /**
29
+ * Lower-level WebSocket handler with no upstream connection.
30
+ *
31
+ * This is the abstract base shared by all WebSocket handlers. It does not open
32
+ * or forward to any upstream server on its own — subclass it directly only when
33
+ * you want to service a fully synthetic connection yourself (e.g. answer the
34
+ * runtime without any real backend). For the common case of mutating and
35
+ * forwarding traffic to the real upstream, subclass {@link CopilotWebSocketForwarder}
36
+ * instead, which connects upstream and forwards by default.
37
+ *
38
+ * @experimental
39
+ */
40
+ export declare abstract class CopilotWebSocketHandler implements AsyncDisposable {
41
+ #private;
42
+ [kSuppressCloseOnDispose]: boolean;
43
+ protected readonly context: CopilotRequestContext;
44
+ protected constructor(context: CopilotRequestContext);
45
+ sendResponseMessage(data: string | Uint8Array): Promise<void>;
46
+ close(status?: CopilotWebSocketCloseStatus): Promise<void>;
47
+ abstract sendRequestMessage(data: string | Uint8Array): Promise<void> | void;
48
+ [Symbol.asyncDispose](): Promise<void>;
49
+ }
50
+ /**
51
+ * WebSocket handler that connects to the real upstream and forwards traffic by
52
+ * default. This is the type returned by the default
53
+ * {@link CopilotRequestHandler.openWebSocket}.
54
+ *
55
+ * Override nothing to get full pass-through. To mutate traffic, subclass this
56
+ * type and override a message hook, then call `super` to keep forwarding to the
57
+ * upstream. (Subclassing {@link CopilotWebSocketHandler} instead would drop
58
+ * forwarding entirely.)
59
+ *
60
+ * @experimental
61
+ */
62
+ export declare class CopilotWebSocketForwarder extends CopilotWebSocketHandler {
63
+ #private;
64
+ constructor(context: CopilotRequestContext);
65
+ sendRequestMessage(data: string | Uint8Array): void;
66
+ close(status?: CopilotWebSocketCloseStatus): Promise<void>;
67
+ [Symbol.asyncDispose](): Promise<void>;
68
+ }
69
+ /**
70
+ * Base class for SDK consumers who want to observe or mutate the outbound
71
+ * model-layer requests the runtime issues (for both CAPI and BYOK providers).
72
+ * Subclass and override {@link sendRequest} or {@link openWebSocket}; an
73
+ * instance that overrides nothing is a transparent pass-through.
74
+ *
75
+ * @experimental
76
+ */
77
+ export declare class CopilotRequestHandler {
78
+ #private;
79
+ protected sendRequest(request: Request, ctx: CopilotRequestContext): Promise<Response>;
80
+ protected openWebSocket(ctx: CopilotRequestContext): Promise<CopilotWebSocketHandler>;
81
+ }
82
+ export {};