mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-04 02:33:34 +00:00
refactor(McpHub): cleanup crosswindow auths
This commit is contained in:
parent
ba5345c084
commit
a127704207
3 changed files with 100 additions and 21 deletions
|
|
@ -171,8 +171,11 @@ export class McpHub {
|
|||
private reauthPromises: Map<string, Promise<void>> = new Map()
|
||||
private _oauthWatchers: Map<string, { unsubscribe: () => void; abortHandle: NodeJS.Timeout }> = new Map()
|
||||
|
||||
constructor(provider: ClineProvider) {
|
||||
constructor(provider: ClineProvider, secretStorage?: SecretStorageService) {
|
||||
this.providerRef = new WeakRef(provider)
|
||||
if (secretStorage) {
|
||||
this.secretStorage = secretStorage
|
||||
}
|
||||
this.watchMcpSettingsFile()
|
||||
this.watchProjectMcpFile().catch(console.error)
|
||||
this.setupWorkspaceFoldersWatcher()
|
||||
|
|
@ -190,16 +193,12 @@ export class McpHub {
|
|||
await this.initializationPromise
|
||||
}
|
||||
|
||||
public setSecretStorage(secretStorage: SecretStorageService): void {
|
||||
this.secretStorage = secretStorage
|
||||
}
|
||||
/**
|
||||
* Registers a client (e.g., ClineProvider) using this hub.
|
||||
* Increments the reference count.
|
||||
*/
|
||||
public registerClient(): void {
|
||||
this.refCount++
|
||||
// console.log(`McpHub: Client registered. Ref count: ${this.refCount}`)
|
||||
}
|
||||
|
||||
/**
|
||||
|
|
@ -209,8 +208,6 @@ export class McpHub {
|
|||
public async unregisterClient(): Promise<void> {
|
||||
this.refCount--
|
||||
|
||||
// console.log(`McpHub: Client unregistered. Ref count: ${this.refCount}`)
|
||||
|
||||
if (this.refCount <= 0) {
|
||||
console.log("McpHub: Last client unregistered. Disposing hub.")
|
||||
await this.dispose()
|
||||
|
|
@ -1054,9 +1051,31 @@ export class McpHub {
|
|||
return
|
||||
}
|
||||
|
||||
// Check if another window already saved valid tokens
|
||||
// Register the cross-window token watcher BEFORE the initial token read so we
|
||||
// don't miss a write that lands in the gap between the read and the subscription.
|
||||
// After the read we immediately unsubscribe; the watcher inside _runOAuthFlow
|
||||
// will set up its own long-lived subscription for the duration of the flow.
|
||||
// Note: onDidChange fires on both writes AND deletes — the re-read below is
|
||||
// authoritative; we only reconnect if it returns a valid (non-expired) token.
|
||||
let tokenChangedDuringRead = false
|
||||
const unsubscribeCrossWindow = this.secretStorage.onDidChange(serverUrl, () => {
|
||||
tokenChangedDuringRead = true
|
||||
})
|
||||
|
||||
// Check if another window already saved valid tokens.
|
||||
// Re-read after registering the watcher in case a write landed in the gap.
|
||||
const existing = await this.secretStorage.getOAuthData(serverUrl)
|
||||
if (existing && Date.now() < existing.expires_at - TOKEN_EXPIRY_BUFFER_MS) {
|
||||
if (this.isDisposed) {
|
||||
unsubscribeCrossWindow()
|
||||
return
|
||||
}
|
||||
// If the first read missed but the watcher fired during it, re-read — a write
|
||||
// may have landed between subscription and read completion.
|
||||
const tokenToUse =
|
||||
existing ?? (tokenChangedDuringRead ? await this.secretStorage.getOAuthData(serverUrl) : undefined)
|
||||
unsubscribeCrossWindow()
|
||||
|
||||
if (!this.isDisposed && tokenToUse && Date.now() < tokenToUse.expires_at - TOKEN_EXPIRY_BUFFER_MS) {
|
||||
await authProvider.close()
|
||||
await this.deleteConnection(name, source)
|
||||
await this.connectToServer(name, config, source)
|
||||
|
|
@ -1185,8 +1204,8 @@ export class McpHub {
|
|||
// Register in _oauthWatchers so deleteConnection() and dispose() can clean up
|
||||
this._oauthWatchers.set(watcherKey, { unsubscribe, abortHandle: timeoutHandle })
|
||||
|
||||
// Non-modal toast — fires and forgets. When it auto-dismisses, update the
|
||||
// persistent progress bar to tell the user how to re-trigger the flow.
|
||||
// Non-modal toast — fires once. If dismissed without clicking Authenticate,
|
||||
// the flow stays alive via the persistent progress bar (Cancel to abort).
|
||||
const authenticateLabel = t("mcp:oauth.flow.authenticateButton")
|
||||
void (async () => {
|
||||
const choice = await vscode.window.showInformationMessage(
|
||||
|
|
@ -1197,8 +1216,6 @@ export class McpHub {
|
|||
if (disposed) return
|
||||
|
||||
if (choice !== authenticateLabel) {
|
||||
// Toast was dismissed (timed out or closed) without clicking Authenticate.
|
||||
// Update the progress bar so the user knows how to re-trigger.
|
||||
progress.report({ message: t("mcp:oauth.flow.dismissedHint") })
|
||||
return
|
||||
}
|
||||
|
|
|
|||
|
|
@ -37,10 +37,8 @@ export class McpServerManager {
|
|||
try {
|
||||
// Double-check instance in case it was created while we were waiting
|
||||
if (!this.instance) {
|
||||
const hub = new McpHub(provider)
|
||||
// Set the secret storage service for OAuth operations
|
||||
const secretStorage = new SecretStorageService(context)
|
||||
hub.setSecretStorage(secretStorage)
|
||||
const hub = new McpHub(provider, secretStorage)
|
||||
// Wait for all MCP servers to finish connecting (or timing out)
|
||||
await hub.waitUntilReady()
|
||||
this.instance = hub
|
||||
|
|
|
|||
|
|
@ -2463,8 +2463,6 @@ describe("McpHub", () => {
|
|||
})
|
||||
|
||||
it("should update progress bar hint when toast is dismissed without clicking Authenticate", async () => {
|
||||
// Toast dismissed (undefined) without clicking Authenticate — flow stays alive
|
||||
// via the progress bar. The progress bar message should update to the dismissedHint.
|
||||
let capturedProgress: any
|
||||
vsc.window.withProgress.mockImplementationOnce((_options: any, task: any) => {
|
||||
capturedProgress = { report: vi.fn() }
|
||||
|
|
@ -2476,7 +2474,6 @@ describe("McpHub", () => {
|
|||
return task(capturedProgress, cancellationToken)
|
||||
})
|
||||
|
||||
// Toast dismissed (no button clicked), then flow stays open forever (cancel via timeout)
|
||||
vsc.window.showInformationMessage.mockResolvedValueOnce(undefined as any)
|
||||
vi.useFakeTimers()
|
||||
|
||||
|
|
@ -2546,8 +2543,6 @@ describe("McpHub", () => {
|
|||
})
|
||||
|
||||
it("should disconnect and flag error when user cancels the OAuth flow", async () => {
|
||||
vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {}))
|
||||
|
||||
// Override withProgress for this test to capture the cancellation token
|
||||
let capturedCancellationToken: any
|
||||
vsc.window.withProgress.mockImplementationOnce((_options: any, task: any) => {
|
||||
|
|
@ -2738,6 +2733,75 @@ describe("McpHub", () => {
|
|||
expect((mcpHub as any).connectToServer).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("should not reconnect when cross-window watcher fires but token is missing or expired", async () => {
|
||||
vi.useFakeTimers()
|
||||
vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {}))
|
||||
|
||||
// onDidChange fires immediately (simulates a storage event during the read gap)
|
||||
// but getOAuthData returns undefined — no valid token written.
|
||||
mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => {
|
||||
cb() // fires synchronously — sets crossWindowTokenWritten flag
|
||||
return vi.fn()
|
||||
})
|
||||
mockSecretStorage.getOAuthData.mockResolvedValue(undefined)
|
||||
|
||||
const flowPromise = (mcpHub as any)._initiateOAuthFlow(
|
||||
serverName,
|
||||
source,
|
||||
config,
|
||||
mockAuthProvider,
|
||||
mockTransport,
|
||||
mockConnection,
|
||||
)
|
||||
|
||||
// Advance past the timeout so the flow can settle
|
||||
await vi.advanceTimersByTimeAsync(OAUTH_FLOW_TIMEOUT_MS)
|
||||
await flowPromise
|
||||
|
||||
// Should NOT have tried to reconnect — the watcher fired but no valid token exists
|
||||
expect((mcpHub as any).connectToServer).not.toHaveBeenCalledWith(serverName, expect.anything(), source)
|
||||
})
|
||||
|
||||
it("should reconnect when cross-window token is written while getOAuthData is in-flight", async () => {
|
||||
// This covers the race: onDidChange fires BEFORE getOAuthData resolves.
|
||||
// The flag causes a second read which finds the now-valid token.
|
||||
let resolveGetOAuthData!: (value: any) => void
|
||||
mockSecretStorage.getOAuthData
|
||||
// First call (during the race window) — delayed, returns undefined
|
||||
.mockImplementationOnce(
|
||||
() =>
|
||||
new Promise((r) => {
|
||||
resolveGetOAuthData = r
|
||||
}),
|
||||
)
|
||||
// Second call (after watcher fires) — valid token available
|
||||
.mockResolvedValueOnce({ expires_at: Date.now() + 10 * 60 * 1000 })
|
||||
|
||||
// Watcher fires synchronously before getOAuthData resolves
|
||||
mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => {
|
||||
cb()
|
||||
return vi.fn()
|
||||
})
|
||||
|
||||
const flowPromise = (mcpHub as any)._initiateOAuthFlow(
|
||||
serverName,
|
||||
source,
|
||||
config,
|
||||
mockAuthProvider,
|
||||
mockTransport,
|
||||
mockConnection,
|
||||
)
|
||||
|
||||
// Now let the first getOAuthData resolve with undefined
|
||||
resolveGetOAuthData(undefined)
|
||||
await flowPromise
|
||||
|
||||
expect(mockAuthProvider.close).toHaveBeenCalled()
|
||||
expect((mcpHub as any).deleteConnection).toHaveBeenCalledWith(serverName, source)
|
||||
expect((mcpHub as any).connectToServer).toHaveBeenCalled()
|
||||
expect(vsc.window.withProgress).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("should cancel the previous watcher when called again for the same server", async () => {
|
||||
vi.useFakeTimers()
|
||||
vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {}))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue