@github/copilot-sdk 1.0.3 → 1.0.5-preview.0

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?.bearerTokenProvider) {
89
+ const { bearerTokenProvider, ...rest } = provider;
90
+ callbacks.set(DEFAULT_PROVIDER_NAME, bearerTokenProvider);
91
+ wireProvider = {
92
+ ...rest,
93
+ hasBearerTokenProvider: true
94
+ };
95
+ }
96
+ let wireProviders = providers;
97
+ if (providers?.some((p) => p.bearerTokenProvider)) {
98
+ wireProviders = providers.map((p) => {
99
+ if (!p.bearerTokenProvider) return p;
100
+ const { bearerTokenProvider, ...rest } = p;
101
+ callbacks.set(p.name, bearerTokenProvider);
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(
@@ -190,10 +218,23 @@ function getBundledCliPath() {
190
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.`
191
219
  );
192
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
+ }
193
233
  class CopilotClient {
194
234
  cliStartTimeout = null;
195
235
  cliProcess = null;
196
236
  connection = null;
237
+ messageWriter = null;
197
238
  socket = null;
198
239
  runtimePort = null;
199
240
  actualHost = "localhost";
@@ -225,6 +266,8 @@ class CopilotClient {
225
266
  negotiatedProtocolVersion = null;
226
267
  /** Connection-level session filesystem config, set via constructor option. */
227
268
  sessionFsConfig = null;
269
+ requestHandler = null;
270
+ llmInferenceHandlers = {};
228
271
  /**
229
272
  * Typed server-scoped RPC methods.
230
273
  * @throws Error if the client is not connected
@@ -255,6 +298,13 @@ class CopilotClient {
255
298
  const level = this.options.logLevel?.toLowerCase();
256
299
  if (level === "debug" || level === "all") {
257
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}
258
308
  `);
259
309
  }
260
310
  }
@@ -317,6 +367,8 @@ class CopilotClient {
317
367
  this.onListModels = options.onListModels;
318
368
  this.onGetTraceContext = options.onGetTraceContext;
319
369
  this.sessionFsConfig = options.sessionFs ?? null;
370
+ this.requestHandler = options.requestHandler ?? null;
371
+ this.setupLlmInference();
320
372
  const effectiveEnv = options.env ?? process.env;
321
373
  this.resolvedEnv = effectiveEnv;
322
374
  this.resolvedCliPath = conn.kind === "stdio" || conn.kind === "tcp" ? conn.path ?? effectiveEnv.COPILOT_CLI_PATH ?? getBundledCliPath() : void 0;
@@ -396,6 +448,20 @@ class CopilotClient {
396
448
  }
397
449
  session.clientSessionApis.sessionFs = createSessionFsAdapter(provider);
398
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
+ }
399
465
  /**
400
466
  * Starts the CLI server and establishes a connection.
401
467
  *
@@ -433,6 +499,9 @@ class CopilotClient {
433
499
  capabilities: this.sessionFsConfig.capabilities
434
500
  });
435
501
  }
502
+ if (this.requestHandler) {
503
+ await this.connection.sendRequest("llmInference.setProvider", {});
504
+ }
436
505
  this.state = "connected";
437
506
  } catch (error) {
438
507
  this.state = "error";
@@ -465,7 +534,8 @@ class CopilotClient {
465
534
  */
466
535
  async stop() {
467
536
  const errors = [];
468
- for (const session of this.sessions.values()) {
537
+ const activeSessions = [...this.sessions.values()];
538
+ for (const session of activeSessions) {
469
539
  const sessionId = session.sessionId;
470
540
  let lastError = null;
471
541
  for (let attempt = 1; attempt <= 3; attempt++) {
@@ -489,8 +559,10 @@ class CopilotClient {
489
559
  );
490
560
  }
491
561
  }
562
+ for (const session of activeSessions) {
563
+ session._markDisconnected();
564
+ }
492
565
  this.sessions.clear();
493
- let runtimeShutdownCompleted = false;
494
566
  if (this.connection && this.cliProcess && !this.isExternalServer) {
495
567
  const runtimeShutdownStart = Date.now();
496
568
  const shutdownPromise = this.rpc.runtime.shutdown();
@@ -501,7 +573,6 @@ class CopilotClient {
501
573
  RUNTIME_SHUTDOWN_TIMEOUT_MS,
502
574
  `runtime.shutdown timed out after ${RUNTIME_SHUTDOWN_TIMEOUT_MS}ms`
503
575
  );
504
- runtimeShutdownCompleted = true;
505
576
  this.logDebugTiming(
506
577
  "CopilotClient.stop runtime shutdown complete",
507
578
  runtimeShutdownStart
@@ -518,6 +589,9 @@ class CopilotClient {
518
589
  );
519
590
  }
520
591
  }
592
+ if (this.messageWriter) {
593
+ this.messageWriter.suppressWriteErrors = true;
594
+ }
521
595
  if (this.connection) {
522
596
  try {
523
597
  this.connection.dispose();
@@ -529,6 +603,7 @@ class CopilotClient {
529
603
  );
530
604
  }
531
605
  this.connection = null;
606
+ this.messageWriter = null;
532
607
  this._rpc = null;
533
608
  this._internalRpc = null;
534
609
  }
@@ -556,16 +631,13 @@ class CopilotClient {
556
631
  this.cliProcess = null;
557
632
  try {
558
633
  if (child.exitCode == null && child.signalCode == null) {
559
- const exitedGracefully = runtimeShutdownCompleted ? await waitForChildExit(child, RUNTIME_SHUTDOWN_TIMEOUT_MS) : false;
560
- if (!exitedGracefully) {
561
- child.kill();
562
- if (!await waitForChildExit(child, RUNTIME_SHUTDOWN_TIMEOUT_MS)) {
563
- errors.push(
564
- new Error(
565
- `Timed out waiting for CLI process to exit after kill: ${RUNTIME_SHUTDOWN_TIMEOUT_MS}ms`
566
- )
567
- );
568
- }
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
+ );
569
641
  }
570
642
  }
571
643
  } catch (error) {
@@ -628,13 +700,20 @@ class CopilotClient {
628
700
  */
629
701
  async forceStop() {
630
702
  this.forceStopping = true;
703
+ for (const session of this.sessions.values()) {
704
+ session._markDisconnected();
705
+ }
631
706
  this.sessions.clear();
707
+ if (this.messageWriter) {
708
+ this.messageWriter.suppressWriteErrors = true;
709
+ }
632
710
  if (this.connection) {
633
711
  try {
634
712
  this.connection.dispose();
635
713
  } catch {
636
714
  }
637
715
  this.connection = null;
716
+ this.messageWriter = null;
638
717
  this._rpc = null;
639
718
  this._internalRpc = null;
640
719
  }
@@ -820,6 +899,11 @@ class CopilotClient {
820
899
  const callerSessionId = config.sessionId;
821
900
  const useServerGeneratedId = config.cloud != null && callerSessionId == null;
822
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);
823
907
  const { wirePayload: wireSystemMessage, transformCallbacks } = extractTransformCallbacks(
824
908
  config.systemMessage
825
909
  );
@@ -833,6 +917,9 @@ class CopilotClient {
833
917
  s.registerTools(config.tools);
834
918
  s.registerCanvases(config.canvases);
835
919
  s.registerCommands(config.commands);
920
+ if (bearerTokenCallbacks.size > 0) {
921
+ s.registerBearerTokenProviders(bearerTokenCallbacks);
922
+ }
836
923
  s.registerPermissionHandler(config.onPermissionRequest);
837
924
  if (config.onUserInputRequest) {
838
925
  s.registerUserInputHandler(config.onUserInputRequest);
@@ -896,8 +983,9 @@ class CopilotClient {
896
983
  availableTools: toolFilterOptions.availableTools,
897
984
  excludedTools: toolFilterOptions.excludedTools,
898
985
  toolFilterPrecedence: toolFilterOptions.toolFilterPrecedence,
899
- provider: config.provider,
900
- providers: config.providers,
986
+ provider: bearerWireProvider,
987
+ capi: config.capi,
988
+ providers: bearerWireProviders,
901
989
  models: config.models,
902
990
  enableSessionTelemetry: config.enableSessionTelemetry,
903
991
  modelCapabilities: config.modelCapabilities,
@@ -936,7 +1024,8 @@ class CopilotClient {
936
1024
  memory: config.memory,
937
1025
  gitHubToken: config.gitHubToken,
938
1026
  remoteSession: config.remoteSession,
939
- cloud: config.cloud
1027
+ cloud: config.cloud,
1028
+ expAssignments: config.expAssignments
940
1029
  });
941
1030
  const {
942
1031
  sessionId: returnedSessionId,
@@ -1003,6 +1092,14 @@ class CopilotClient {
1003
1092
  session.registerTools(config.tools);
1004
1093
  session.registerCanvases(config.canvases);
1005
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
+ }
1006
1103
  session.registerPermissionHandler(config.onPermissionRequest);
1007
1104
  if (config.onUserInputRequest) {
1008
1105
  session.registerUserInputHandler(config.onUserInputRequest);
@@ -1064,8 +1161,9 @@ class CopilotClient {
1064
1161
  name: cmd.name,
1065
1162
  description: cmd.description
1066
1163
  })),
1067
- provider: config.provider,
1068
- providers: config.providers,
1164
+ provider: bearerWireProvider,
1165
+ capi: config.capi,
1166
+ providers: bearerWireProviders,
1069
1167
  models: config.models,
1070
1168
  modelCapabilities: config.modelCapabilities,
1071
1169
  largeOutput: toWireLargeOutput(config.largeOutput),
@@ -1105,7 +1203,8 @@ class CopilotClient {
1105
1203
  continuePendingWork: config.continuePendingWork,
1106
1204
  gitHubToken: config.gitHubToken,
1107
1205
  remoteSession: config.remoteSession,
1108
- openCanvases: config.openCanvases
1206
+ openCanvases: config.openCanvases,
1207
+ expAssignments: config.expAssignments
1109
1208
  });
1110
1209
  const { workspacePath, capabilities, openCanvases } = response;
1111
1210
  session["_workspacePath"] = workspacePath;
@@ -1647,13 +1746,21 @@ stderr: ${stderrOutput}`
1647
1746
  throw new Error("CLI process not started");
1648
1747
  }
1649
1748
  this.cliProcess.stdin?.on("error", (err) => {
1650
- if (!this.forceStopping) {
1651
- 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 {
1652
1758
  }
1653
1759
  });
1760
+ this.messageWriter = new TeardownResilientStreamMessageWriter(this.cliProcess.stdin);
1654
1761
  this.connection = createMessageConnection(
1655
1762
  new StreamMessageReader(this.cliProcess.stdout),
1656
- new StreamMessageWriter(this.cliProcess.stdin)
1763
+ this.messageWriter
1657
1764
  );
1658
1765
  this.attachConnectionHandlers();
1659
1766
  this.connection.listen();
@@ -1665,9 +1772,10 @@ stderr: ${stderrOutput}`
1665
1772
  if (this.cliProcess) {
1666
1773
  throw new Error("CLI child process was unexpectedly started in parent process mode");
1667
1774
  }
1775
+ this.messageWriter = new TeardownResilientStreamMessageWriter(process.stdout);
1668
1776
  this.connection = createMessageConnection(
1669
1777
  new StreamMessageReader(process.stdin),
1670
- new StreamMessageWriter(process.stdout)
1778
+ this.messageWriter
1671
1779
  );
1672
1780
  this.attachConnectionHandlers();
1673
1781
  this.connection.listen();
@@ -1687,9 +1795,10 @@ stderr: ${stderrOutput}`
1687
1795
  }, 1e4);
1688
1796
  this.socket.connect(this.runtimePort, this.actualHost, () => {
1689
1797
  clearTimeout(connectionTimeout);
1798
+ this.messageWriter = new TeardownResilientStreamMessageWriter(this.socket);
1690
1799
  this.connection = createMessageConnection(
1691
1800
  new StreamMessageReader(this.socket),
1692
- new StreamMessageWriter(this.socket)
1801
+ this.messageWriter
1693
1802
  );
1694
1803
  this.attachConnectionHandlers();
1695
1804
  this.connection.listen();
@@ -1737,6 +1846,7 @@ stderr: ${stderrOutput}`
1737
1846
  if (!session) throw new Error(`No session found for sessionId: ${sessionId}`);
1738
1847
  return session.clientSessionApis;
1739
1848
  });
1849
+ registerClientGlobalApiHandlers(this.connection, this.llmInferenceHandlers);
1740
1850
  this.connection.onClose(() => {
1741
1851
  this.state = "disconnected";
1742
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 {};