feat(McpOAuthClientProvider): adding callback to cancel callback server

This commit is contained in:
Elliott de Launay 2026-04-24 09:56:29 -04:00
parent dab4e2cece
commit f0e1121cc3
No known key found for this signature in database
GPG key ID: BB899BED766D1806
2 changed files with 21 additions and 5 deletions

View file

@ -78,6 +78,7 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
private _server: http.Server | null, private _server: http.Server | null,
private _port: number, private _port: number,
private _authCodePromise: Promise<string> | null, private _authCodePromise: Promise<string> | null,
private _cancelCallbackServer: (() => void) | null,
private readonly _tokenEndpointAuthMethod: string, private readonly _tokenEndpointAuthMethod: string,
private readonly _grantTypes: string[], private readonly _grantTypes: string[],
private readonly _scopes: string[], private readonly _scopes: string[],
@ -122,9 +123,9 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
} }
// Extract auth-method preferences. // Extract auth-method preferences.
// Prefer "none" → first supported → "client_secret_post" // Only pick methods we actually implement: "none" or "client_secret_post".
const authMethods: string[] = authServerMeta?.token_endpoint_auth_methods_supported ?? [] const authMethods: string[] = authServerMeta?.token_endpoint_auth_methods_supported ?? []
const tokenEndpointAuthMethod = authMethods.includes("none") ? "none" : (authMethods[0] ?? "client_secret_post") const tokenEndpointAuthMethod = authMethods.includes("none") ? "none" : "client_secret_post"
const grantTypes: string[] = authServerMeta?.grant_types_supported ?? ["authorization_code", "refresh_token"] const grantTypes: string[] = authServerMeta?.grant_types_supported ?? ["authorization_code", "refresh_token"]
const scopes: string[] = authServerMeta?.scopes_supported ?? [] const scopes: string[] = authServerMeta?.scopes_supported ?? []
@ -142,6 +143,7 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
null, null,
0, 0,
null, null,
null,
tokenEndpointAuthMethod, tokenEndpointAuthMethod,
grantTypes, grantTypes,
scopes, scopes,
@ -178,9 +180,10 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
private async _doStartCallbackServer(): Promise<void> { private async _doStartCallbackServer(): Promise<void> {
this._closed = false this._closed = false
const { server, port, result } = await startCallbackServer(this._port, this._state) const { server, port, result, cancel } = await startCallbackServer(this._port, this._state)
this._server = server this._server = server
this._port = port this._port = port
this._cancelCallbackServer = cancel
this._authCodePromise = result.then((r) => { this._authCodePromise = result.then((r) => {
if (r.error) throw new Error(`OAuth authorization failed: ${r.error}`) if (r.error) throw new Error(`OAuth authorization failed: ${r.error}`)
if (!r.code) throw new Error("No authorization code received in callback") if (!r.code) throw new Error("No authorization code received in callback")
@ -442,6 +445,12 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
code_verifier: codeVerifier, code_verifier: codeVerifier,
} }
// RFC 8707: include resource indicator so servers that bind token requests
// to a specific resource can validate the exchange.
if (this._resourceIndicator) {
params.resource = this._resourceIndicator
}
// Include client_secret in the body when the auth method is client_secret_post. // Include client_secret in the body when the auth method is client_secret_post.
if (this._tokenEndpointAuthMethod === "client_secret_post" && this._clientInfo.client_secret) { if (this._tokenEndpointAuthMethod === "client_secret_post" && this._clientInfo.client_secret) {
params.client_secret = this._clientInfo.client_secret params.client_secret = this._clientInfo.client_secret
@ -489,6 +498,11 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
client_id: clientId, client_id: clientId,
} }
// RFC 8707: include resource indicator in refresh requests too.
if (this._resourceIndicator) {
params.resource = this._resourceIndicator
}
if (this._tokenEndpointAuthMethod === "client_secret_post" && this._clientInfo?.client_secret) { if (this._tokenEndpointAuthMethod === "client_secret_post" && this._clientInfo?.client_secret) {
params.client_secret = this._clientInfo.client_secret params.client_secret = this._clientInfo.client_secret
} }
@ -521,8 +535,9 @@ export class McpOAuthClientProvider implements OAuthClientProvider {
} }
if (!this._closed && this._server) { if (!this._closed && this._server) {
this._closed = true this._closed = true
await stopCallbackServer(this._server).catch(() => {}) await stopCallbackServer(this._server, this._cancelCallbackServer ?? (() => {})).catch(() => {})
this._server = null this._server = null
this._cancelCallbackServer = null
this._authCodePromise = null this._authCodePromise = null
} }
} }

View file

@ -75,6 +75,7 @@ function setupCallbackServerMock(code = "test-auth-code", state?: string) {
server: mockServer, server: mockServer,
port: 12345, port: 12345,
result: resultPromise, result: resultPromise,
cancel: vi.fn(),
}) })
return { mockServer, resultPromise } return { mockServer, resultPromise }
} }
@ -675,7 +676,7 @@ describe("McpOAuthClientProvider", () => {
await provider.close() await provider.close()
expect(stopCallbackServer).toHaveBeenCalledWith(mockServer) expect(stopCallbackServer).toHaveBeenCalledWith(mockServer, expect.any(Function))
}) })
it("should be idempotent", async () => { it("should be idempotent", async () => {