diff --git a/src/services/mcp/SecretStorageService.ts b/src/services/mcp/SecretStorageService.ts index ac57207734..9a66c4c270 100644 --- a/src/services/mcp/SecretStorageService.ts +++ b/src/services/mcp/SecretStorageService.ts @@ -29,10 +29,9 @@ export class SecretStorageService { private _key(serverUrl: string): string { const url = new URL(serverUrl) - // Use host + pathname for stricter isolation between different MCP servers on the same host. - // We sanitize the pathname to ensure it's a valid key component. - const sanitizedPath = url.pathname.replace(/[^a-zA-Z0-9]/g, "_").replace(/^_+|_+$/g, "") - const pathSuffix = sanitizedPath ? `.${sanitizedPath}` : "" + const normalizedPath = url.pathname.replace(/\/$/, "") + // Use base64url encoding to avoid collisions between paths like /a-b, /a_b, /a/b. + const pathSuffix = normalizedPath ? `.${Buffer.from(normalizedPath).toString("base64url")}` : "" return `${this._namespace}${url.host}${pathSuffix}.data` } diff --git a/src/services/mcp/__tests__/SecretStorageService.spec.ts b/src/services/mcp/__tests__/SecretStorageService.spec.ts index 63d83ec764..a0d4adfebb 100644 --- a/src/services/mcp/__tests__/SecretStorageService.spec.ts +++ b/src/services/mcp/__tests__/SecretStorageService.spec.ts @@ -59,8 +59,8 @@ describe("SecretStorageService", () => { }) it("should return undefined for malformed JSON", async () => { - // Manually store garbage via the underlying mock - context.secrets.store("mcp.oauth.example.com.mcp.data", "not-json") + // Manually store garbage via the underlying mock (key uses base64url-encoded path) + context.secrets.store("mcp.oauth.example.com.L21jcA.data", "not-json") const result = await service.getOAuthData("https://example.com/mcp") expect(result).toBeUndefined() @@ -75,7 +75,10 @@ describe("SecretStorageService", () => { } await service.saveOAuthData("https://example.com/mcp", data) - expect(context.secrets.store).toHaveBeenCalledWith("mcp.oauth.example.com.mcp.data", JSON.stringify(data)) + expect(context.secrets.store).toHaveBeenCalledWith( + "mcp.oauth.example.com.L21jcA.data", + JSON.stringify(data), + ) }) it("should handle root path correctly", async () => { @@ -124,7 +127,7 @@ describe("SecretStorageService", () => { await service.deleteOAuthData("https://example.com/mcp") - expect(context.secrets.delete).toHaveBeenCalledWith("mcp.oauth.example.com.mcp.data") + expect(context.secrets.delete).toHaveBeenCalledWith("mcp.oauth.example.com.L21jcA.data") const result = await service.getOAuthData("https://example.com/mcp") expect(result).toBeUndefined() }) @@ -134,7 +137,7 @@ describe("SecretStorageService", () => { const cb = vi.fn() service.onDidChange("https://example.com/mcp", cb) - context.secrets._emit("mcp.oauth.example.com.mcp.data") + context.secrets._emit("mcp.oauth.example.com.L21jcA.data") expect(cb).toHaveBeenCalledTimes(1) }) @@ -143,7 +146,7 @@ describe("SecretStorageService", () => { const cb = vi.fn() service.onDidChange("https://example.com/mcp", cb) - context.secrets._emit("mcp.oauth.other.com.mcp.data") + context.secrets._emit("mcp.oauth.other.com.L21jcA.data") expect(cb).not.toHaveBeenCalled() }) @@ -153,7 +156,7 @@ describe("SecretStorageService", () => { const unsubscribe = service.onDidChange("https://example.com/mcp", cb) unsubscribe() - context.secrets._emit("mcp.oauth.example.com.mcp.data") + context.secrets._emit("mcp.oauth.example.com.L21jcA.data") expect(cb).not.toHaveBeenCalled() }) @@ -217,5 +220,19 @@ describe("SecretStorageService", () => { expect((await service.getOAuthData("https://example.com/service1"))?.tokens.access_token).toBe("path1") expect((await service.getOAuthData("https://example.com/service2"))?.tokens.access_token).toBe("path2") }) + + it("should not collide between paths that differ only in separators (/a-b, /a_b, /a/b)", async () => { + const urls = ["https://example.com/a-b", "https://example.com/a_b", "https://example.com/a/b"] + for (const [i, url] of urls.entries()) { + await service.saveOAuthData(url, { + tokens: { access_token: `tok-${i}`, token_type: "Bearer" }, + expires_at: i, + }) + } + + for (const [i, url] of urls.entries()) { + expect((await service.getOAuthData(url))?.tokens.access_token).toBe(`tok-${i}`) + } + }) }) })