diff --git a/README.md b/README.md index 5290e32..d32016a 100644 --- a/README.md +++ b/README.md @@ -442,6 +442,12 @@ Put the namespace id returned by the first command into `kv_namespaces` in redeem authorization codes against Testomat.io server-to-server and is never committed. +Before exposing the Worker publicly, configure Cloudflare rate-limiting rules for +`/register`, `/authorize`, `/token`, and `/mcp/*`. These endpoints intentionally +support unauthenticated OAuth discovery and client registration, so rate limiting +belongs at the edge rather than in per-isolate memory. Keep separate rules and KV +namespaces for beta and production. + A beta worker is the same code deployed to the `beta` environment, which targets `https://beta.testomat.io` and keeps its own KV namespace so beta grants never reach the production one: diff --git a/src/api/testomatio-client.js b/src/api/testomatio-client.js index fe463e8..1f14699 100644 --- a/src/api/testomatio-client.js +++ b/src/api/testomatio-client.js @@ -1,4 +1,5 @@ import { HttpClient } from './http-client.js'; +import { encodePathParameter } from '../core/path-segment.js'; export class TestomatioApiClient { constructor({ baseUrl, projectId, token, logger }) { @@ -14,9 +15,12 @@ export class TestomatioApiClient { } buildPath(resource, id = '') { + const safeProjectId = encodePathParameter(this.projectId, 'Project ID'); const safeResource = String(resource).replace(/^\/+|\/+$/g, ''); - const safeId = id ? `/${String(id).replace(/^\/+|\/+$/g, '')}` : ''; - return `/api/v2/${this.projectId}/${safeResource}${safeId}`; + const resourceId = String(id ?? ''); + const safeId = resourceId ? `/${encodePathParameter(resourceId, 'Resource ID')}` : ''; + + return `/api/v2/${safeProjectId}/${safeResource}${safeId}`; } list(resource, query = {}) { diff --git a/src/core/path-segment.js b/src/core/path-segment.js new file mode 100644 index 0000000..4773bd4 --- /dev/null +++ b/src/core/path-segment.js @@ -0,0 +1,21 @@ +function pathParameter(value, label = 'Path parameter') { + const segment = String(value ?? ''); + + if (!segment || segment === '.' || segment === '..' || segment.includes('/') || segment.includes('\\')) { + throw new TypeError(`${label} must be a single URL path segment`); + } + + return segment; +} + +export function encodePathParameter(value, label) { + return encodeURIComponent(pathParameter(value, label)); +} + +export function decodePathParameter(value, label) { + try { + return pathParameter(decodeURIComponent(value), label); + } catch { + return ''; + } +} diff --git a/src/mcp/server.js b/src/mcp/server.js index 8cdf63e..3c7997d 100644 --- a/src/mcp/server.js +++ b/src/mcp/server.js @@ -28,7 +28,8 @@ export class TestomatioMCPServer { tools, ...registryOptions, }); - this.cleanupStarted = false; + this.closePromise = null; + this.sessionCleanupPromise = null; this.server = new Server( { @@ -70,29 +71,42 @@ export class TestomatioMCPServer { this.logger.info('Testomatio MCP server started'); } - installSessionCleanup() { - const cleanup = async () => { - if (this.cleanupStarted) { - return; - } + close() { + if (!this.closePromise) { + this.closePromise = (async () => { + try { + await this.server.close(); + } finally { + await this.#stopSession(); + } + })(); + } - this.cleanupStarted = true; - await this.apiClient?.stopSession?.(); - }; + return this.closePromise; + } + installSessionCleanup() { this.server.onclose = () => { - void cleanup(); + void this.#stopSession(); }; process.once('beforeExit', () => { - void cleanup(); + void this.#stopSession(); }); for (const signal of ['SIGINT', 'SIGTERM']) { process.once(signal, async () => { - await cleanup(); + await this.#stopSession(); process.exit(0); }); } } + + #stopSession() { + if (!this.sessionCleanupPromise) { + this.sessionCleanupPromise = Promise.resolve().then(() => this.apiClient?.stopSession?.()); + } + + return this.sessionCleanupPromise; + } } diff --git a/test/create-server.test.js b/test/create-server.test.js index 36b8c4c..3c4c8a4 100644 --- a/test/create-server.test.js +++ b/test/create-server.test.js @@ -1,4 +1,4 @@ -import { describe, expect, it } from 'vitest'; +import { describe, expect, it, vi } from 'vitest'; import { createMcpServer } from '../src/mcp/create-server.js'; describe('createMcpServer', () => { @@ -17,6 +17,7 @@ describe('createMcpServer', () => { expect(server.toolRegistry.apiClient.projectId).toBe('demo'); expect(server.toolRegistry.apiClient.http.baseUrl).toBe('https://beta.testomat.io'); expect(typeof server.connect).toBe('function'); + expect(typeof server.close).toBe('function'); }); it('isolates configuration per instance', () => { @@ -28,4 +29,17 @@ describe('createMcpServer', () => { expect(first.toolRegistry.apiClient.projectId).toBe('one'); expect(second.toolRegistry.apiClient.projectId).toBe('two'); }); + + it('closes the server and its API session only once', async () => { + const server = createMcpServer({ token: 'a', projectId: 'one', baseUrl: 'https://app.testomat.io' }); + const stopSession = vi.fn(async () => {}); + const closeProtocol = vi.fn(async () => {}); + server.apiClient.stopSession = stopSession; + server.server.close = closeProtocol; + + await Promise.all([server.close(), server.close()]); + + expect(closeProtocol).toHaveBeenCalledTimes(1); + expect(stopSession).toHaveBeenCalledTimes(1); + }); }); diff --git a/worker/src/mcp-handler.js b/worker/src/mcp-handler.js index d36e8b8..5f086c6 100644 --- a/worker/src/mcp-handler.js +++ b/worker/src/mcp-handler.js @@ -4,6 +4,7 @@ import { CfWorkerJsonSchemaValidator } from '@modelcontextprotocol/sdk/validatio import { loadServerConfig } from '../../src/config/load-config.js'; import { createMcpServer } from '../../src/mcp/create-server.js'; import { createLogger } from '../../src/core/logger.js'; +import { decodePathParameter } from '../../src/core/path-segment.js'; import pkg from '../../package.json'; const PROJECT_PATH = /^\/mcp\/([^/]+)\/?$/; @@ -23,11 +24,11 @@ export class McpHandler extends WorkerEntrypoint { static projectId(pathname) { const match = PROJECT_PATH.exec(pathname); - return match ? decodeURIComponent(match[1]) : ''; + return match ? decodePathParameter(match[1], 'Project ID') : ''; } async fetch(request) { - if (request.method === 'GET') { + if (request.method !== 'POST') { return new Response('Method Not Allowed', { status: 405, headers: { Allow: 'POST' } }); } @@ -61,9 +62,12 @@ export class McpHandler extends WorkerEntrypoint { }); const transport = new WebStandardStreamableHTTPServerTransport({ enableJsonResponse: true }); - await server.connect(transport); - - return transport.handleRequest(request); + try { + await server.connect(transport); + return await transport.handleRequest(request); + } finally { + await server.close(); + } } async #mapRevokedToken(request, response) { diff --git a/worker/src/testomatio-handler.js b/worker/src/testomatio-handler.js index 62dd618..94fc38a 100644 --- a/worker/src/testomatio-handler.js +++ b/worker/src/testomatio-handler.js @@ -1,8 +1,9 @@ import { WorkerEntrypoint } from 'cloudflare:workers'; import { loadServerConfig } from '../../src/config/load-config.js'; +import { decodePathParameter } from '../../src/core/path-segment.js'; const AUTH_REQUEST_TTL_SECONDS = 600; -const RESOURCE_PROJECT = /\/mcp\/([^/?#]+)/; +const RESOURCE_PATH = /^\/mcp\/([^/]+)\/?$/; export class TestomatioAuthHandler extends WorkerEntrypoint { async fetch(request) { @@ -32,7 +33,13 @@ export class TestomatioAuthHandler extends WorkerEntrypoint { }); const client = await this.env.OAUTH_PROVIDER.lookupClient(oauthReqInfo.clientId).catch(() => null); - const project = this.#projectFromResource(oauthReqInfo.resource); + const project = this.#projectFromResource(oauthReqInfo.resource, request); + if (!project) { + return new Response('Authorization resource must identify one project on this server', { + status: 400, + }); + } + const redirect = new URL('/mcp/authorize', `${this.#baseUrl()}/`); redirect.searchParams.set('state', state); @@ -41,9 +48,7 @@ export class TestomatioAuthHandler extends WorkerEntrypoint { redirect.searchParams.set('client_name', client.clientName); } - if (project) { - redirect.searchParams.set('project', project); - } + redirect.searchParams.set('project', project); return Response.redirect(redirect.toString(), 302); } @@ -115,13 +120,22 @@ export class TestomatioAuthHandler extends WorkerEntrypoint { return crypto.randomUUID().replace(/-/g, ''); } - #projectFromResource(resource) { - const value = Array.isArray(resource) ? resource[0] : resource; - if (!value) { + #projectFromResource(resource, request) { + if (!resource || Array.isArray(resource)) { return ''; } - const match = RESOURCE_PROJECT.exec(String(value)); - return match ? decodeURIComponent(match[1]) : ''; + try { + const resourceUrl = new URL(String(resource)); + const requestUrl = new URL(request.url); + if (resourceUrl.origin !== requestUrl.origin || resourceUrl.search || resourceUrl.hash) { + return ''; + } + + const match = RESOURCE_PATH.exec(resourceUrl.pathname); + return match ? decodePathParameter(match[1], 'Project ID') : ''; + } catch { + return ''; + } } } diff --git a/worker/test/mcp-endpoint.test.js b/worker/test/mcp-endpoint.test.js index 86b4b06..3a05da1 100644 --- a/worker/test/mcp-endpoint.test.js +++ b/worker/test/mcp-endpoint.test.js @@ -99,8 +99,60 @@ describe('POST /mcp/', () => { expect(calls[0]).toBe('/api/v2/other-project/suites'); }); - it('rejects GET with 405', async () => { + it('closes the Testomat.io API session after a mutating request', async () => { + stubApi((url, init) => { + const method = init?.method; + calls.push({ method, path: url.pathname, session: new Headers(init?.headers).get('X-Session-Hash') }); + + if (method === 'POST' && url.pathname.endsWith('/sessions')) { + return Response.json({ data: { hash: 'session-1' } }); + } + + if (method === 'POST' && url.pathname.endsWith('/tests')) { + return Response.json({ data: { id: 'test-1' } }); + } + + if (method === 'DELETE' && url.pathname.endsWith('/sessions/session-1')) { + return Response.json({}); + } + + return Response.json({ error: 'Unexpected request' }, { status: 500 }); + }); + + const { client, transport } = connect(); + await client.connect(transport); + await client.callTool({ + name: 'tests_create', + arguments: { title: 'Session cleanup', suite_id: 'suite-1' }, + }); + + expect(calls).toEqual([ + { method: 'POST', path: '/api/v2/demo-project/sessions', session: null }, + { method: 'POST', path: '/api/v2/demo-project/tests', session: 'session-1' }, + { method: 'DELETE', path: '/api/v2/demo-project/sessions/session-1', session: null }, + ]); + }); + + it.each(['demo%2F..%2F..%2Fadmin', '%'])( + 'rejects an unsafe project path: %s', + async (project) => { + const response = await SELF.fetch(`https://mcp.testomat.test/mcp/${project}`, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + Accept: 'application/json, text/event-stream', + Authorization: `Bearer ${STATIC_TOKEN}`, + }, + body: JSON.stringify({ jsonrpc: '2.0', id: 1, method: 'tools/list', params: {} }), + }); + + expect(response.status).toBe(404); + } + ); + + it.each(['GET', 'PUT', 'DELETE'])('rejects %s with 405', async (method) => { const response = await SELF.fetch(MCP_URL, { + method, headers: { Authorization: `Bearer ${STATIC_TOKEN}` }, }); diff --git a/worker/test/oauth-flow.test.js b/worker/test/oauth-flow.test.js index 914abfe..5117e1f 100644 --- a/worker/test/oauth-flow.test.js +++ b/worker/test/oauth-flow.test.js @@ -2,6 +2,7 @@ import { SELF, env } from 'cloudflare:test'; import { afterEach, describe, expect, it, vi } from 'vitest'; const CLIENT_REDIRECT = 'https://claude.test/api/mcp/callback'; +const MCP_RESOURCE = 'https://mcp.testomat.test/mcp/demo-project'; async function registerClient() { const response = await SELF.fetch('https://mcp.testomat.test/register', { @@ -18,7 +19,7 @@ async function registerClient() { return response.json(); } -async function startAuthorization(clientId) { +async function startAuthorization(clientId, resource = MCP_RESOURCE) { const url = new URL('https://mcp.testomat.test/authorize'); url.searchParams.set('response_type', 'code'); url.searchParams.set('client_id', clientId); @@ -26,6 +27,11 @@ async function startAuthorization(clientId) { url.searchParams.set('state', 'client-state'); url.searchParams.set('code_challenge', 'E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM'); url.searchParams.set('code_challenge_method', 'S256'); + for (const value of Array.isArray(resource) ? resource : [resource]) { + if (value) { + url.searchParams.append('resource', value); + } + } return SELF.fetch(url.toString(), { redirect: 'manual' }); } @@ -68,6 +74,18 @@ describe('oauth authorize', () => { expect(stored.redirectUri).toBe(CLIENT_REDIRECT); }); + it.each([ + ['a missing resource', ''], + ['an unsafe project id', 'https://mcp.testomat.test/mcp/demo%2F..%2F..%2Fadmin'], + ['a resource on another origin', 'https://attacker.test/mcp/demo-project'], + ['multiple resources', [MCP_RESOURCE, 'https://mcp.testomat.test/mcp/other-project']], + ])('rejects %s', async (_label, resource) => { + const client = await registerClient(); + const response = await startAuthorization(client.client_id, resource); + + expect(response.status).toBe(400); + }); + it('exchanges the opaque code server to server and completes the grant', async () => { const client = await registerClient(); const authorize = await startAuthorization(client.client_id);