Compare commits

...
Sign in to create a new pull request.

1 commit

Author SHA1 Message Date
Aiden Cline
d88b599b1d fix(mcp): persist OAuth discovery state 2026-07-28 10:42:48 -05:00
3 changed files with 113 additions and 3 deletions

View file

@ -5,6 +5,15 @@ import { Global } from "@opencode-ai/core/global"
import { Effect, Layer, Context, Option, Schema } from "effect" import { Effect, Layer, Context, Option, Schema } from "effect"
import { FSUtil } from "@opencode-ai/core/fs-util" import { FSUtil } from "@opencode-ai/core/fs-util"
import { EffectFlock } from "@opencode-ai/core/util/effect-flock" import { EffectFlock } from "@opencode-ai/core/util/effect-flock"
import type { OAuthDiscoveryState } from "@modelcontextprotocol/client"
const DiscoveryState = Schema.declare<OAuthDiscoveryState>(
(value): value is OAuthDiscoveryState =>
typeof value === "object" &&
value !== null &&
"authorizationServerUrl" in value &&
typeof value.authorizationServerUrl === "string",
)
export const Tokens = Schema.Struct({ export const Tokens = Schema.Struct({
accessToken: Schema.mutableKey(Schema.String), accessToken: Schema.mutableKey(Schema.String),
@ -31,6 +40,7 @@ export const Entry = Schema.Struct({
clientInfo: Schema.mutableKey(Schema.optional(ClientInfo)), clientInfo: Schema.mutableKey(Schema.optional(ClientInfo)),
codeVerifier: Schema.mutableKey(Schema.optional(Schema.String)), codeVerifier: Schema.mutableKey(Schema.optional(Schema.String)),
oauthState: Schema.mutableKey(Schema.optional(Schema.String)), oauthState: Schema.mutableKey(Schema.optional(Schema.String)),
discoveryState: Schema.mutableKey(Schema.optional(DiscoveryState)),
serverUrl: Schema.mutableKey(Schema.optional(Schema.String)), serverUrl: Schema.mutableKey(Schema.optional(Schema.String)),
}) })
export type Entry = Schema.Schema.Type<typeof Entry> export type Entry = Schema.Schema.Type<typeof Entry>
@ -54,6 +64,8 @@ export interface Interface {
readonly updateOAuthState: (mcpName: string, oauthState: string) => Effect.Effect<void> readonly updateOAuthState: (mcpName: string, oauthState: string) => Effect.Effect<void>
readonly getOAuthState: (mcpName: string) => Effect.Effect<string | undefined> readonly getOAuthState: (mcpName: string) => Effect.Effect<string | undefined>
readonly clearOAuthState: (mcpName: string) => Effect.Effect<void> readonly clearOAuthState: (mcpName: string) => Effect.Effect<void>
readonly updateDiscoveryState: (mcpName: string, discoveryState: OAuthDiscoveryState) => Effect.Effect<void>
readonly clearDiscoveryState: (mcpName: string) => Effect.Effect<void>
} }
export class Service extends Context.Service<Service, Interface>()("@opencode/McpAuth") {} export class Service extends Context.Service<Service, Interface>()("@opencode/McpAuth") {}
@ -137,8 +149,10 @@ const layer = Layer.effect(
const updateClientInfo = updateField("clientInfo", "updateClientInfo") const updateClientInfo = updateField("clientInfo", "updateClientInfo")
const updateCodeVerifier = updateField("codeVerifier", "updateCodeVerifier") const updateCodeVerifier = updateField("codeVerifier", "updateCodeVerifier")
const updateOAuthState = updateField("oauthState", "updateOAuthState") const updateOAuthState = updateField("oauthState", "updateOAuthState")
const updateDiscoveryState = updateField("discoveryState", "updateDiscoveryState")
const clearCodeVerifier = clearField("codeVerifier", "clearCodeVerifier") const clearCodeVerifier = clearField("codeVerifier", "clearCodeVerifier")
const clearOAuthState = clearField("oauthState", "clearOAuthState") const clearOAuthState = clearField("oauthState", "clearOAuthState")
const clearDiscoveryState = clearField("discoveryState", "clearDiscoveryState")
const getOAuthState = Effect.fn("McpAuth.getOAuthState")(function* (mcpName: string) { const getOAuthState = Effect.fn("McpAuth.getOAuthState")(function* (mcpName: string) {
const entry = yield* get(mcpName) const entry = yield* get(mcpName)
@ -158,6 +172,8 @@ const layer = Layer.effect(
updateOAuthState, updateOAuthState,
getOAuthState, getOAuthState,
clearOAuthState, clearOAuthState,
updateDiscoveryState,
clearDiscoveryState,
}) })
}), }),
) )

View file

@ -3,6 +3,7 @@ import type {
OAuthClientMetadata, OAuthClientMetadata,
StoredOAuthTokens, StoredOAuthTokens,
StoredOAuthClientInformation, StoredOAuthClientInformation,
OAuthDiscoveryState,
} from "@modelcontextprotocol/client" } from "@modelcontextprotocol/client"
import { Effect } from "effect" import { Effect } from "effect"
import { McpAuth } from "./auth" import { McpAuth } from "./auth"
@ -195,7 +196,16 @@ export class McpOAuthProvider implements OAuthClientProvider {
return newState return newState
} }
async invalidateCredentials(type: "all" | "client" | "tokens"): Promise<void> { async saveDiscoveryState(state: OAuthDiscoveryState): Promise<void> {
await Effect.runPromise(this.auth.updateDiscoveryState(this.mcpName, state))
}
async discoveryState(): Promise<OAuthDiscoveryState | undefined> {
const entry = await Effect.runPromise(this.auth.get(this.mcpName))
return entry?.discoveryState
}
async invalidateCredentials(type: "all" | "client" | "tokens" | "verifier" | "discovery"): Promise<void> {
const entry = await Effect.runPromise(this.auth.get(this.mcpName)) const entry = await Effect.runPromise(this.auth.get(this.mcpName))
if (!entry) return if (!entry) return
switch (type) { switch (type) {
@ -210,6 +220,12 @@ export class McpOAuthProvider implements OAuthClientProvider {
delete entry.tokens delete entry.tokens
await Effect.runPromise(this.auth.set(this.mcpName, entry)) await Effect.runPromise(this.auth.set(this.mcpName, entry))
break break
case "verifier":
await Effect.runPromise(this.auth.clearCodeVerifier(this.mcpName))
break
case "discovery":
await Effect.runPromise(this.auth.clearDiscoveryState(this.mcpName))
break
} }
} }
} }
@ -217,6 +233,7 @@ export class McpOAuthProvider implements OAuthClientProvider {
export class McpOAuthPendingProvider extends McpOAuthProvider { export class McpOAuthPendingProvider extends McpOAuthProvider {
private pendingClientInfo?: StoredOAuthClientInformation private pendingClientInfo?: StoredOAuthClientInformation
private pendingTokens?: StoredOAuthTokens private pendingTokens?: StoredOAuthTokens
private pendingDiscoveryState?: OAuthDiscoveryState
override async clientInformation(): Promise<StoredOAuthClientInformation | undefined> { override async clientInformation(): Promise<StoredOAuthClientInformation | undefined> {
if (!this.config.clientId) return this.pendingClientInfo if (!this.config.clientId) return this.pendingClientInfo
@ -238,9 +255,21 @@ export class McpOAuthPendingProvider extends McpOAuthProvider {
this.pendingTokens = tokens this.pendingTokens = tokens
} }
override async invalidateCredentials(type: "all" | "client" | "tokens"): Promise<void> { override async saveDiscoveryState(state: OAuthDiscoveryState): Promise<void> {
this.pendingDiscoveryState = state
}
override async discoveryState(): Promise<OAuthDiscoveryState | undefined> {
return this.pendingDiscoveryState
}
override async invalidateCredentials(
type: "all" | "client" | "tokens" | "verifier" | "discovery",
): Promise<void> {
if (type === "all" || type === "client") this.pendingClientInfo = undefined if (type === "all" || type === "client") this.pendingClientInfo = undefined
if (type === "all" || type === "tokens") this.pendingTokens = undefined if (type === "all" || type === "tokens") this.pendingTokens = undefined
if (type === "all" || type === "discovery") this.pendingDiscoveryState = undefined
if (type === "verifier") await super.invalidateCredentials(type)
} }
async commit(): Promise<void> { async commit(): Promise<void> {
@ -271,6 +300,7 @@ export class McpOAuthPendingProvider extends McpOAuthProvider {
issuer: this.pendingClientInfo.issuer, issuer: this.pendingClientInfo.issuer,
} }
: undefined, : undefined,
discoveryState: this.pendingDiscoveryState,
}, },
this.serverUrl, this.serverUrl,
), ),

View file

@ -1,6 +1,13 @@
import { test, expect, describe } from "bun:test" import { test, expect, describe } from "bun:test"
import { McpOAuthProvider, OAUTH_CALLBACK_PORT, OAUTH_CALLBACK_PATH } from "../../src/mcp/oauth-provider" import {
McpOAuthPendingProvider,
McpOAuthProvider,
OAUTH_CALLBACK_PORT,
OAUTH_CALLBACK_PATH,
} from "../../src/mcp/oauth-provider"
import type { McpAuth } from "../../src/mcp/auth" import type { McpAuth } from "../../src/mcp/auth"
import type { OAuthDiscoveryState } from "@modelcontextprotocol/client"
import { Effect } from "effect"
// Stub auth — only synchronous getters are exercised in these tests // Stub auth — only synchronous getters are exercised in these tests
const stubAuth = {} as McpAuth.Interface const stubAuth = {} as McpAuth.Interface
@ -59,3 +66,60 @@ describe("McpOAuthProvider.clientMetadata", () => {
expect(provider.clientMetadata.token_endpoint_auth_method).toBe("none") expect(provider.clientMetadata.token_endpoint_auth_method).toBe("none")
}) })
}) })
describe("McpOAuthProvider.discoveryState", () => {
const discoveryState: OAuthDiscoveryState = {
authorizationServerUrl: "https://auth.example.com",
authorizationServerMetadata: {
issuer: "https://auth.example.com",
authorization_endpoint: "https://auth.example.com/authorize",
token_endpoint: "https://auth.example.com/token",
response_types_supported: ["code"],
},
resourceMetadataUrl: "https://mcp.example.com/.well-known/oauth-protected-resource",
}
test("persists discovery state through the auth store", async () => {
let saved: OAuthDiscoveryState | undefined
const auth = {
...stubAuth,
get: () => Effect.succeed(saved ? { discoveryState: saved } : undefined),
updateDiscoveryState: (_name: string, value: OAuthDiscoveryState) => Effect.sync(() => void (saved = value)),
clearDiscoveryState: () => Effect.sync(() => void (saved = undefined)),
} satisfies McpAuth.Interface
const provider = new McpOAuthProvider(
"test-server",
"https://mcp.example.com/mcp",
{},
{ onRedirect: async () => {} },
auth,
)
await provider.saveDiscoveryState(discoveryState)
expect(await provider.discoveryState()).toEqual(discoveryState)
await provider.invalidateCredentials("discovery")
expect(await provider.discoveryState()).toBeUndefined()
})
test("commits pending discovery state with OAuth credentials", async () => {
let entry: McpAuth.Entry | undefined
const auth = {
...stubAuth,
set: (_name: string, value: McpAuth.Entry) => Effect.sync(() => void (entry = value)),
} satisfies McpAuth.Interface
const provider = new McpOAuthPendingProvider(
"test-server",
"https://mcp.example.com/mcp",
{},
{ onRedirect: async () => {} },
auth,
)
await provider.saveDiscoveryState(discoveryState)
await provider.saveTokens({ access_token: "token", token_type: "Bearer" })
await provider.commit()
expect(entry?.discoveryState).toEqual(discoveryState)
})
})