diff --git a/MCP_SETUP.md b/MCP_SETUP.md index 6da9b6b..95a9b79 100644 --- a/MCP_SETUP.md +++ b/MCP_SETUP.md @@ -40,8 +40,11 @@ For production, add these Worker variables in the Cloudflare dashboard: ```text MCP_ALLOWED_HOSTNAMES= MCP_ALLOWED_ORIGIN_HOSTNAMES=chatgpt.com,claude.ai +MCP_TOKEN_CACHE_TTL_SECONDS=300 ``` +`MCP_TOKEN_CACHE_TTL_SECONDS` controls how long a verified token stays cached in memory per isolate. It defaults to 300 seconds; set it to `0` to verify every request against Ed. + ## Choose the right connection | Client | URL | Authentication | diff --git a/README.md b/README.md index 985f6a1..4b2e94a 100644 --- a/README.md +++ b/README.md @@ -115,7 +115,7 @@ The read-only `list_lesson_files` tool returns compact file metadata plus MCP re [![Deploy to Cloudflare](https://deploy.workers.cloudflare.com/button)](https://deploy.workers.cloudflare.com/?url=https://github.com/bunizao/edstem-cli) -The button copies this repository to your Git provider and deploys `src/worker.ts`. The Worker uses no database, storage binding, or protocol session. Clients send an Ed access token or API key with each request; the Worker validates the credential without storing it. The hosted OAuth service stores the Ed token encrypted with AES-256-GCM instead. +The button copies this repository to your Git provider and deploys `src/worker.ts`. The Worker uses no database, storage binding, or protocol session. Clients send an Ed access token or API key with each request; the Worker validates the credential without storing it. The hosted OAuth service stores the Ed token encrypted with AES-256-GCM instead. Verified tokens are cached in memory for five minutes per isolate, so repeated calls skip the extra Ed verification round-trip; set `MCP_TOKEN_CACHE_TTL_SECONDS=0` to disable it. After deployment, the endpoints are: diff --git a/src/worker.ts b/src/worker.ts index 592da43..9968267 100644 --- a/src/worker.ts +++ b/src/worker.ts @@ -16,12 +16,19 @@ import { const DEFAULT_ED_API_BASE_URL = "https://edstem.org/api/"; const READ_SCOPE = "mcp:tools.read"; const WRITE_SCOPE = "mcp:tools.write"; +const TOKEN_CACHE_TTL_MS = 5 * 60 * 1000; +const TOKEN_CACHE_MAX_ENTRIES = 1000; const handlers = new WeakMap(); +// Verified tokens, keyed by SHA-256 so the raw token never lives in the map. +// The cache is per-isolate and best-effort: a cold isolate just verifies again. +const verifiedTokens = new Map(); + export interface WorkerEnv { ED_API_BASE_URL?: string; MCP_ALLOWED_HOSTNAMES?: string; MCP_ALLOWED_ORIGIN_HOSTNAMES?: string; + MCP_TOKEN_CACHE_TTL_SECONDS?: string; } type Credential = { @@ -29,6 +36,11 @@ type Credential = { token: string; }; +type CachedIdentity = { + edUserId: number; + expiresAt: number; +}; + const worker = { async fetch( request: Request, @@ -80,10 +92,14 @@ function getHandler(env: WorkerEnv): StatelessMcpHandler { throw new Error("Verified MCP auth context is missing a token."); } - const client = new EdClient({ apiBaseUrl, token: authInfo.token }); + const token = authInfo.token; + const client = new EdClient({ apiBaseUrl, token }); return createEdMcpServer({ canWrite: () => true, - getClient: () => client + getClient: () => client, + onAuthExpired: async () => { + verifiedTokens.delete(await hashToken(token)); + } }); }, { @@ -140,20 +156,24 @@ async function verifyCredential( credential: Credential, env: WorkerEnv ): Promise { + const ttlMs = tokenCacheTtlMs(env); + const cacheKey = ttlMs > 0 ? await hashToken(credential.token) : undefined; + if (cacheKey) { + const cached = readCachedIdentity(cacheKey); + if (cached) { + return buildAuthInfo(credential, cached.edUserId); + } + } + try { const identity = await verifyEdToken( credential.token, normalizeApiBaseUrl(env.ED_API_BASE_URL) ); - return { - clientId: `ed:${identity.edUserId}`, - extra: { - authMethod: credential.source, - edUserId: identity.edUserId - }, - scopes: [READ_SCOPE, WRITE_SCOPE], - token: credential.token - }; + if (cacheKey) { + cacheIdentity(cacheKey, identity.edUserId, ttlMs); + } + return buildAuthInfo(credential, identity.edUserId); } catch (error) { if (error instanceof EdTokenInvalidError) { return unauthorizedResponse( @@ -175,6 +195,55 @@ async function verifyCredential( } } +function buildAuthInfo(credential: Credential, edUserId: number): AuthInfo { + return { + clientId: `ed:${edUserId}`, + extra: { + authMethod: credential.source, + edUserId + }, + scopes: [READ_SCOPE, WRITE_SCOPE], + token: credential.token + }; +} + +function readCachedIdentity(cacheKey: string): CachedIdentity | undefined { + const cached = verifiedTokens.get(cacheKey); + if (!cached) { + return undefined; + } + if (cached.expiresAt <= Date.now()) { + verifiedTokens.delete(cacheKey); + return undefined; + } + return cached; +} + +function cacheIdentity(cacheKey: string, edUserId: number, ttlMs: number): void { + verifiedTokens.delete(cacheKey); + if (verifiedTokens.size >= TOKEN_CACHE_MAX_ENTRIES) { + const oldest = verifiedTokens.keys().next(); + if (!oldest.done) { + verifiedTokens.delete(oldest.value); + } + } + verifiedTokens.set(cacheKey, { edUserId, expiresAt: Date.now() + ttlMs }); +} + +function tokenCacheTtlMs(env: WorkerEnv): number { + const raw = env.MCP_TOKEN_CACHE_TTL_SECONDS?.trim(); + if (!raw) { + return TOKEN_CACHE_TTL_MS; + } + const seconds = Number(raw); + return Number.isFinite(seconds) && seconds >= 0 ? seconds * 1000 : TOKEN_CACHE_TTL_MS; +} + +async function hashToken(token: string): Promise { + const digest = await crypto.subtle.digest("SHA-256", new TextEncoder().encode(token)); + return Array.from(new Uint8Array(digest), (byte) => byte.toString(16).padStart(2, "0")).join(""); +} + function normalizeApiBaseUrl(value: string | undefined): string { const normalized = value?.trim() || DEFAULT_ED_API_BASE_URL; return `${normalized.replace(/\/+$/, "")}/`; diff --git a/tests/worker/worker.test.ts b/tests/worker/worker.test.ts index a61b390..66e350d 100644 --- a/tests/worker/worker.test.ts +++ b/tests/worker/worker.test.ts @@ -213,6 +213,159 @@ describe("Cloudflare Worker MCP", () => { }); }); +describe("Cloudflare Worker token verification cache", () => { + const cleanups: Array<() => Promise> = []; + + afterEach(async () => { + while (cleanups.length > 0) { + await cleanups.pop()?.(); + } + }); + + it("verifies a token once and reuses the cached identity", async () => { + const fakeEd = await startFakeEdServer([edUser("cache-hit-token")]); + cleanups.push(fakeEd.close); + const verifications = countUserFetches(cleanups); + const env = { ED_API_BASE_URL: fakeEd.baseUrl }; + + const first = await fetchWorker(listToolsRequest("cache-hit-token"), env); + expect(first.status).toBe(200); + expect(verifications.count).toBe(1); + + const second = await fetchWorker(listToolsRequest("cache-hit-token"), env); + expect(second.status).toBe(200); + expect(verifications.count).toBe(1); + }); + + it("verifies each distinct token", async () => { + const fakeEd = await startFakeEdServer([ + edUser("distinct-token-a"), + edUser("distinct-token-b", 102) + ]); + cleanups.push(fakeEd.close); + const verifications = countUserFetches(cleanups); + const env = { ED_API_BASE_URL: fakeEd.baseUrl }; + + expect((await fetchWorker(listToolsRequest("distinct-token-a"), env)).status).toBe(200); + expect((await fetchWorker(listToolsRequest("distinct-token-b"), env)).status).toBe(200); + expect(verifications.count).toBe(2); + }); + + it("re-verifies once the cached entry expires", async () => { + const fakeEd = await startFakeEdServer([edUser("expiring-token")]); + cleanups.push(fakeEd.close); + const verifications = countUserFetches(cleanups); + const env = { + ED_API_BASE_URL: fakeEd.baseUrl, + MCP_TOKEN_CACHE_TTL_SECONDS: "0.05" + }; + + expect((await fetchWorker(listToolsRequest("expiring-token"), env)).status).toBe(200); + expect(verifications.count).toBe(1); + + await Bun.sleep(80); + + expect((await fetchWorker(listToolsRequest("expiring-token"), env)).status).toBe(200); + expect(verifications.count).toBe(2); + }); + + it("disables caching when the TTL is zero", async () => { + const fakeEd = await startFakeEdServer([edUser("uncached-token")]); + cleanups.push(fakeEd.close); + const verifications = countUserFetches(cleanups); + const env = { + ED_API_BASE_URL: fakeEd.baseUrl, + MCP_TOKEN_CACHE_TTL_SECONDS: "0" + }; + + expect((await fetchWorker(listToolsRequest("uncached-token"), env)).status).toBe(200); + expect((await fetchWorker(listToolsRequest("uncached-token"), env)).status).toBe(200); + expect(verifications.count).toBe(2); + }); + + it("never caches a rejected token", async () => { + const fakeEd = await startFakeEdServer([edUser("known-token")]); + cleanups.push(fakeEd.close); + const verifications = countUserFetches(cleanups); + const env = { ED_API_BASE_URL: fakeEd.baseUrl }; + + expect((await fetchWorker(listToolsRequest("bogus-token"), env)).status).toBe(401); + expect((await fetchWorker(listToolsRequest("bogus-token"), env)).status).toBe(401); + expect(verifications.count).toBe(2); + }); + + it("drops the cached entry when a tool call hits an expired token", async () => { + const fakeEd = await startFakeEdServer([edUser("revoked-token")]); + cleanups.push(fakeEd.close); + const env = { ED_API_BASE_URL: fakeEd.baseUrl }; + + expect((await fetchWorker(listToolsRequest("revoked-token"), env)).status).toBe(200); + + fakeEd.revokeToken("revoked-token"); + + const toolCall = await fetchWorker( + modernRequest( + "tools/call", + { + _meta: CLIENT_META, + arguments: { includeArchived: false }, + name: "list_courses" + }, + { + Authorization: "Bearer revoked-token", + "Mcp-Name": "list_courses" + } + ), + env + ); + expect(toolCall.status).toBe(200); + const toolPayload = await toolCall.json() as any; + expect(toolPayload.result.content[0].text).toContain("EDSTEM_REAUTH_REQUIRED"); + + const afterEviction = await fetchWorker(listToolsRequest("revoked-token"), env); + expect(afterEviction.status).toBe(401); + }); +}); + +function edUser(token: string, id = 101) { + return { + courses: [], + token, + user: { + avatar: "", + course_role: "student", + email: `user-${id}@example.com`, + id, + name: "Ada", + role: "student" + } + }; +} + +function listToolsRequest(token: string): Request { + return modernRequest( + "tools/list", + { _meta: CLIENT_META }, + { Authorization: `Bearer ${token}` } + ); +} + +function countUserFetches(cleanups: Array<() => Promise>): { count: number } { + const counter = { count: 0 }; + const original = globalThis.fetch; + globalThis.fetch = ((input: any, init?: any) => { + const url = typeof input === "string" ? input : input?.url ?? String(input); + if (url.endsWith("/api/user")) { + counter.count += 1; + } + return original(input, init); + }) as typeof fetch; + cleanups.push(async () => { + globalThis.fetch = original; + }); + return counter; +} + function modernRequest( method: string, params: Record,