diff --git a/cli-auth.ts b/cli-auth.ts index 9948c92..29a6e92 100644 --- a/cli-auth.ts +++ b/cli-auth.ts @@ -8,5 +8,5 @@ const baseUrl = process.env.BETTER_AUTH_URL ?? CLI_DEV_BASE_URL; export const auth = betterAuth({ baseURL: baseUrl, secret: process.env.BETTER_AUTH_SECRET ?? randomUUID(), - ...createBaseAuthConfig(baseUrl), + ...createBaseAuthConfig(), }); diff --git a/package.json b/package.json index b0ed436..56392ee 100644 --- a/package.json +++ b/package.json @@ -46,11 +46,14 @@ }, "DATAFORSEO_API_KEY": { "description": "Base64-encoded `login:password` for DataForSEO API access." + }, + "OAUTH_KV": { + "description": "KV namespace used by the Cloudflare Workers OAuth provider for MCP client registrations, grants, and tokens." } } }, "dependencies": { - "@better-auth/oauth-provider": "^1.5.5", + "@cloudflare/workers-oauth-provider": "^0.4.0", "@every-app/sdk": "^0.1.14", "@modelcontextprotocol/sdk": "1.29.0", "@tanstack/query-core": "^5.90.9", diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 6b86017..0b97ab0 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -8,9 +8,9 @@ importers: .: dependencies: - '@better-auth/oauth-provider': - specifier: ^1.5.5 - version: 1.5.5(@better-auth/core@1.5.5(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(@cloudflare/workers-types@4.20260302.0)(better-call@1.3.2(zod@4.3.6))(jose@6.1.3)(kysely@0.28.12)(nanostores@1.1.1))(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(better-auth@1.5.5(@cloudflare/workers-types@4.20260302.0)(@tanstack/react-start@1.167.16(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(drizzle-kit@0.31.9)(drizzle-orm@0.44.7(@cloudflare/workers-types@4.20260302.0)(@libsql/client@0.15.15)(@opentelemetry/api@1.9.1)(kysely@0.28.12))(mongodb@7.2.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(solid-js@1.9.11)(vitest@3.2.4(@types/debug@4.1.13)(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(better-call@1.3.2(zod@4.3.6)) + '@cloudflare/workers-oauth-provider': + specifier: ^0.4.0 + version: 0.4.0 '@every-app/sdk': specifier: ^0.1.14 version: 0.1.14(@tanstack/react-router@1.168.10(react-dom@19.2.4(react@19.2.4))(react@19.2.4))(@tanstack/react-start@1.167.16(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(jose@6.1.3)(react@19.2.4) @@ -40,7 +40,7 @@ importers: version: 8.21.3(react-dom@19.2.4(react@19.2.4))(react@19.2.4) agents: specifier: 0.12.3 - version: 0.12.3(@babel/core@7.29.0)(@babel/runtime@7.29.2)(@cloudflare/workers-types@4.20260302.0)(ai@6.0.176(zod@4.3.6))(react@19.2.4)(rolldown@1.0.0)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0))(zod@4.3.6) + version: 0.12.3(@babel/core@7.29.0)(@babel/runtime@7.29.2)(@cloudflare/workers-types@4.20260302.0)(ai@6.0.177(zod@4.3.6))(react@19.2.4)(rolldown@1.0.0)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0))(zod@4.3.6) autumn-js: specifier: ^1.1.7 version: 1.1.7(better-auth@1.5.5(@cloudflare/workers-types@4.20260302.0)(@tanstack/react-start@1.167.16(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(drizzle-kit@0.31.9)(drizzle-orm@0.44.7(@cloudflare/workers-types@4.20260302.0)(@libsql/client@0.15.15)(@opentelemetry/api@1.9.1)(kysely@0.28.12))(mongodb@7.2.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(solid-js@1.9.11)(vitest@3.2.4(@types/debug@4.1.13)(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(better-call@1.3.2(zod@4.3.6))(express@5.2.1)(hono@4.12.18)(react@19.2.4) @@ -186,8 +186,8 @@ importers: packages: - '@ai-sdk/gateway@3.0.111': - resolution: {integrity: sha512-gzdRuEH9Mqeuu8zG6j4of3EH3fFJUI0UIubyeaA8gep6KzhCJF7uaTfagSE7x2vLAf381g/NrxsXhhH7Hon9iA==} + '@ai-sdk/gateway@3.0.112': + resolution: {integrity: sha512-jiBao9pR4owWyjo0BnuNc7WSQBGOD0thysE4AFgZXaG+zMFbISQXUkJr7ePw/phBvePy7jE5FSA2Lf7lwqUiiQ==} engines: {node: '>=18'} peerDependencies: zod: ^3.25.76 || ^4.1.8 @@ -398,15 +398,6 @@ packages: '@better-auth/utils': ^0.3.0 mongodb: ^6.0.0 || ^7.0.0 - '@better-auth/oauth-provider@1.5.5': - resolution: {integrity: sha512-zH2uKtvd6406MysWCTBldPHTKCXEK8caMrNId03bh4ej4f2vU8+GfNGE+IyxARucHGI1T+Og7QrUgKAeA2jQUQ==} - peerDependencies: - '@better-auth/core': 1.5.5 - '@better-auth/utils': 0.3.1 - '@better-fetch/fetch': 1.1.21 - better-auth: 1.5.5 - better-call: 1.3.2 - '@better-auth/prisma-adapter@1.5.5': resolution: {integrity: sha512-CliDd78CXHzzwQIXhCdwGr5Ml53i6JdCHWV7PYwTIJz9EAm6qb2RVBdpP3nqEfNjINGM22A6gfleCgCdZkTIZg==} peerDependencies: @@ -483,6 +474,9 @@ packages: cpu: [x64] os: [win32] + '@cloudflare/workers-oauth-provider@0.4.0': + resolution: {integrity: sha512-UtbV8hjC2NloB+Ds6J6v/9HiG8rx8MbdeYGCyFwOACT5vANWzDL6SKo3W5UZymsXiameAgC7jAmtUx4cc+Qpaw==} + '@cloudflare/workers-types@4.20260302.0': resolution: {integrity: sha512-mbFRnlu1lNCScMpXZk/X/uBPufYx5OSbq+euGonGRcY+DgOwm2kczGdK401rUh52NB0fFMEcOy/zqwxv7CdDNA==} @@ -2474,8 +2468,8 @@ packages: vite: optional: true - ai@6.0.176: - resolution: {integrity: sha512-dhxDef3VCIxaFr6tKyG0BrkkCelmnporlen8nHajIwCk7S4PvIaSVI/iyJenhFOZ9KBoKjCAoUs6TzZ3yrSjxw==} + ai@6.0.177: + resolution: {integrity: sha512-1xQtbeWwNcLyyM86ixZhkKvT+WRXc1lvarIKqPVtsyn8F9NDikwUMBqYu+aQKDgMht50SMXh4qboYuU8MeHZZA==} engines: {node: '>=18'} peerDependencies: zod: ^3.25.76 || ^4.1.8 @@ -3246,8 +3240,8 @@ packages: resolution: {integrity: sha512-DyFP3BM/3YHTQOCUL/w0OZHR0lpKeGrxotcHWcqNEdnltqFwXVfhEBQ94eIo34AfQpo0rGki4cyIiftY06h2Fg==} engines: {node: 6.* || 8.* || >= 10.*} - get-east-asian-width@1.5.0: - resolution: {integrity: sha512-CQ+bEO+Tva/qlmw24dCejulK5pMzVnUOFOijVogd3KQs07HnRIgp8TGipvCCRT06xeYEbpbgwaCxglFyiuIcmA==} + get-east-asian-width@1.6.0: + resolution: {integrity: sha512-QRbvDIbx6YklUe6RxeTeleMR0yv3cYH6PsPZHcnVn7xv7zO1BHN8r0XETu8n6Ye3Q+ahtSarc3WgtNWmehIBfA==} engines: {node: '>=18'} get-intrinsic@1.3.0: @@ -4666,7 +4660,7 @@ packages: snapshots: - '@ai-sdk/gateway@3.0.111(zod@4.3.6)': + '@ai-sdk/gateway@3.0.112(zod@4.3.6)': dependencies: '@ai-sdk/provider': 3.0.10 '@ai-sdk/provider-utils': 4.0.27(zod@4.3.6) @@ -4917,16 +4911,6 @@ snapshots: '@better-auth/utils': 0.3.1 mongodb: 7.2.0 - '@better-auth/oauth-provider@1.5.5(@better-auth/core@1.5.5(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(@cloudflare/workers-types@4.20260302.0)(better-call@1.3.2(zod@4.3.6))(jose@6.1.3)(kysely@0.28.12)(nanostores@1.1.1))(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(better-auth@1.5.5(@cloudflare/workers-types@4.20260302.0)(@tanstack/react-start@1.167.16(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(drizzle-kit@0.31.9)(drizzle-orm@0.44.7(@cloudflare/workers-types@4.20260302.0)(@libsql/client@0.15.15)(@opentelemetry/api@1.9.1)(kysely@0.28.12))(mongodb@7.2.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(solid-js@1.9.11)(vitest@3.2.4(@types/debug@4.1.13)(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(better-call@1.3.2(zod@4.3.6))': - dependencies: - '@better-auth/core': 1.5.5(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(@cloudflare/workers-types@4.20260302.0)(better-call@1.3.2(zod@4.3.6))(jose@6.1.3)(kysely@0.28.12)(nanostores@1.1.1) - '@better-auth/utils': 0.3.1 - '@better-fetch/fetch': 1.1.21 - better-auth: 1.5.5(@cloudflare/workers-types@4.20260302.0)(@tanstack/react-start@1.167.16(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)))(drizzle-kit@0.31.9)(drizzle-orm@0.44.7(@cloudflare/workers-types@4.20260302.0)(@libsql/client@0.15.15)(@opentelemetry/api@1.9.1)(kysely@0.28.12))(mongodb@7.2.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)(solid-js@1.9.11)(vitest@3.2.4(@types/debug@4.1.13)(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)) - better-call: 1.3.2(zod@4.3.6) - jose: 6.1.3 - zod: 4.3.6 - '@better-auth/prisma-adapter@1.5.5(@better-auth/core@1.5.5(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(@cloudflare/workers-types@4.20260302.0)(better-call@1.3.2(zod@4.3.6))(jose@6.1.3)(kysely@0.28.12)(nanostores@1.1.1))(@better-auth/utils@0.3.1)': dependencies: '@better-auth/core': 1.5.5(@better-auth/utils@0.3.1)(@better-fetch/fetch@1.1.21)(@cloudflare/workers-types@4.20260302.0)(better-call@1.3.2(zod@4.3.6))(jose@6.1.3)(kysely@0.28.12)(nanostores@1.1.1) @@ -4980,6 +4964,8 @@ snapshots: '@cloudflare/workerd-windows-64@1.20260219.0': optional: true + '@cloudflare/workers-oauth-provider@0.4.0': {} + '@cloudflare/workers-types@4.20260302.0': {} '@cspotcode/source-map-support@0.8.1': @@ -6541,13 +6527,13 @@ snapshots: dependencies: humanize-ms: 1.2.1 - agents@0.12.3(@babel/core@7.29.0)(@babel/runtime@7.29.2)(@cloudflare/workers-types@4.20260302.0)(ai@6.0.176(zod@4.3.6))(react@19.2.4)(rolldown@1.0.0)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0))(zod@4.3.6): + agents@0.12.3(@babel/core@7.29.0)(@babel/runtime@7.29.2)(@cloudflare/workers-types@4.20260302.0)(ai@6.0.177(zod@4.3.6))(react@19.2.4)(rolldown@1.0.0)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0))(zod@4.3.6): dependencies: '@babel/plugin-proposal-decorators': 7.29.0(@babel/core@7.29.0) '@cfworker/json-schema': 4.1.1 '@modelcontextprotocol/sdk': 1.29.0(@cfworker/json-schema@4.1.1)(zod@4.3.6) '@rolldown/plugin-babel': 0.2.3(@babel/core@7.29.0)(@babel/runtime@7.29.2)(rolldown@1.0.0)(vite@7.3.1(@types/node@22.19.11)(jiti@2.6.1)(lightningcss@1.31.1)(tsx@4.21.0)) - ai: 6.0.176(zod@4.3.6) + ai: 6.0.177(zod@4.3.6) cron-schedule: 6.0.0 mimetext: 3.0.28 nanoid: 5.1.11 @@ -6566,9 +6552,9 @@ snapshots: - rolldown - supports-color - ai@6.0.176(zod@4.3.6): + ai@6.0.177(zod@4.3.6): dependencies: - '@ai-sdk/gateway': 3.0.111(zod@4.3.6) + '@ai-sdk/gateway': 3.0.112(zod@4.3.6) '@ai-sdk/provider': 3.0.10 '@ai-sdk/provider-utils': 4.0.27(zod@4.3.6) '@opentelemetry/api': 1.9.0 @@ -7269,7 +7255,7 @@ snapshots: get-caller-file@2.0.5: {} - get-east-asian-width@1.5.0: {} + get-east-asian-width@1.6.0: {} get-intrinsic@1.3.0: dependencies: @@ -8550,7 +8536,7 @@ snapshots: string-width@7.2.0: dependencies: emoji-regex: 10.6.0 - get-east-asian-width: 1.5.0 + get-east-asian-width: 1.6.0 strip-ansi: 7.2.0 stringify-entities@4.0.4: diff --git a/src/env.d.ts b/src/env.d.ts index 1e957f0..73c778b 100644 --- a/src/env.d.ts +++ b/src/env.d.ts @@ -4,6 +4,7 @@ declare namespace Cloudflare { interface Env { R2: R2Bucket; + OAUTH_KV: KVNamespace; AUTH_MODE?: "cloudflare_access" | "local_noauth" | "hosted"; TEAM_DOMAIN?: string; diff --git a/src/lib/auth-client.ts b/src/lib/auth-client.ts index c3be5cf..f30dfb1 100644 --- a/src/lib/auth-client.ts +++ b/src/lib/auth-client.ts @@ -3,7 +3,6 @@ import { inferAdditionalFields, organizationClient, } from "better-auth/client/plugins"; -import { oauthProviderClient } from "@better-auth/oauth-provider/client"; import { captureClientEvent, resetAnalyticsUser } from "@/client/lib/posthog"; import { userAdditionalFields } from "@/lib/auth-options"; import { getSignInHrefForLocation } from "@/lib/auth-redirect"; @@ -12,7 +11,6 @@ export const authClient = createAuthClient({ baseURL: typeof window !== "undefined" ? window.location.origin : "", plugins: [ organizationClient(), - oauthProviderClient(), inferAdditionalFields({ user: userAdditionalFields }), ], }); diff --git a/src/lib/auth-config.ts b/src/lib/auth-config.ts index c8d6910..52703bb 100644 --- a/src/lib/auth-config.ts +++ b/src/lib/auth-config.ts @@ -1,76 +1,9 @@ -import { oauthProvider } from "@better-auth/oauth-provider"; -import { jwt, organization } from "better-auth/plugins"; +import { organization } from "better-auth/plugins"; import { baseAuthOptions } from "@/lib/auth-options"; -import { getActiveOrganizationId } from "@/lib/auth-session"; -import { - getMcpOrganizationIdClaim, - getMcpResource, - MCP_OAUTH_SCOPES, - MCP_SCOPE, -} from "@/lib/oauth-resource"; - -function assertSingleMcpAudience(audiences: string[]) { - if (audiences.length !== 1) { - throw new Error( - "MCP OAuth resource injection requires exactly one valid audience", - ); - } -} - -export function createBaseAuthConfig(baseUrl: string) { - const mcpResource = getMcpResource(baseUrl); - const mcpOrganizationIdClaim = getMcpOrganizationIdClaim(baseUrl); - const validAudiences = [mcpResource]; - - assertSingleMcpAudience(validAudiences); +export function createBaseAuthConfig() { return { ...baseAuthOptions, - plugins: [ - organization(), - jwt(), - oauthProvider({ - loginPage: "/sign-in", - consentPage: "/oauth-consent", - signup: { - page: "/sign-up", - }, - scopes: MCP_OAUTH_SCOPES, - // We publish /.well-known/oauth-authorization-server/api/auth via - // TanStack routes, so silence Better Auth's metadata reminder. - silenceWarnings: { - oauthAuthServerConfig: true, - }, - allowDynamicClientRegistration: true, - clientRegistrationDefaultScopes: MCP_OAUTH_SCOPES, - clientRegistrationAllowedScopes: MCP_OAUTH_SCOPES, - // TODO: drop once the MCP spec settles on a replacement for - // unauthenticated DCR — better-auth has flagged this option for removal. - allowUnauthenticatedClientRegistration: true, - // Single allowed audience — see `routes/api/auth/$.ts`, which defaults - // missing `resource` on /oauth2/token to this value. Adding a second - // audience here would make that injection unsafe (we'd no longer know - // which to pick) and require scope-conditional logic in the route. - validAudiences, - postLogin: { - page: "/oauth-consent", - shouldRedirect: () => false, - consentReferenceId: ({ session, scopes }) => { - if (!scopes.includes(MCP_SCOPE)) { - return undefined; - } - - return getActiveOrganizationId({ session }) ?? undefined; - }, - }, - customAccessTokenClaims: ({ referenceId, scopes }) => { - if (!scopes.includes(MCP_SCOPE)) { - return {}; - } - - return referenceId ? { [mcpOrganizationIdClaim]: referenceId } : {}; - }, - }), - ], + plugins: [organization()], }; } diff --git a/src/lib/auth-redirect.test.ts b/src/lib/auth-redirect.test.ts index 9831270..a532d4f 100644 --- a/src/lib/auth-redirect.test.ts +++ b/src/lib/auth-redirect.test.ts @@ -1,8 +1,10 @@ import { describe, expect, it } from "vitest"; import { getAuthRedirectFromSearch, + getCurrentAuthRedirectFromHref, getOAuthAuthorizeRedirectFromSearch, getOAuthSignedQuery, + getSignInHref, normalizeAuthRedirect, } from "./auth-redirect"; @@ -26,6 +28,32 @@ describe("auth redirect helpers", () => { expect(normalizeAuthRedirect("//evil.example/app")).toBe("/"); }); + it("keeps same-origin relative redirects", () => { + expect( + normalizeAuthRedirect("/api/auth/oauth2/authorize?client_id=abc"), + ).toBe("/api/auth/oauth2/authorize?client_id=abc"); + }); + + it("rejects external and protocol-relative redirects", () => { + expect(normalizeAuthRedirect("https://evil.test")).toBe("/"); + expect(normalizeAuthRedirect("//evil.test")).toBe("/"); + }); + + it("builds sign-in links with the redirect query only when needed", () => { + expect(getSignInHref("/")).toBe("/sign-in"); + expect(getSignInHref("/oauth-consent?client_id=abc")).toBe( + "/sign-in?redirect=%2Foauth-consent%3Fclient_id%3Dabc", + ); + }); + + it("extracts the current path, query, and hash from hrefs", () => { + expect( + getCurrentAuthRedirectFromHref( + "https://open-seo.test/projects?tab=keywords#top", + ), + ).toBe("/projects?tab=keywords#top"); + }); + it("preserves safe internal redirects", () => { expect(getAuthRedirectFromSearch("", "/app")).toBe("/app"); }); diff --git a/src/lib/auth.ts b/src/lib/auth.ts index 77cda79..565cb37 100644 --- a/src/lib/auth.ts +++ b/src/lib/auth.ts @@ -25,14 +25,11 @@ const hostedBaseUrlSchema = z function createAuth() { const baseUrl = getHostedBaseUrl(); const bypassEmail = Reflect.get(env, "BYPASS_EMAIL_VERIFICATION") === "true"; - const baseAuthConfig = createBaseAuthConfig(baseUrl); + const baseAuthConfig = createBaseAuthConfig(); const auth = betterAuth({ baseURL: baseUrl, secret: getHostedSecret(), - // Disable Better Auth's generic /token endpoint so OAuth access tokens only - // flow through /oauth2/token, where the MCP resource shim can run. - disabledPaths: ["/token"], ...baseAuthConfig, emailAndPassword: { ...baseAuthConfig.emailAndPassword, diff --git a/src/lib/oauth-provider-resource-client.ts b/src/lib/oauth-provider-resource-client.ts deleted file mode 100644 index 8ab19ee..0000000 --- a/src/lib/oauth-provider-resource-client.ts +++ /dev/null @@ -1,14 +0,0 @@ -import { oauthProviderResourceClient } from "@better-auth/oauth-provider/resource-client"; -import { getAuth } from "@/lib/auth"; - -type ResourceClientAuth = Parameters[0]; - -export function getOAuthProviderResourceActions() { - // Better Auth documents passing the server auth instance here, but the - // resource-client package currently types the generic too narrowly for the - // concrete `betterAuth(...)` return type. - return oauthProviderResourceClient( - // oxlint-disable-next-line typescript-eslint/no-unsafe-type-assertion - getAuth() as unknown as ResourceClientAuth, - ).getActions(); -} diff --git a/src/lib/oauth-resource.ts b/src/lib/oauth-resource.ts index afb859b..59e13af 100644 --- a/src/lib/oauth-resource.ts +++ b/src/lib/oauth-resource.ts @@ -5,19 +5,3 @@ export const MCP_OAUTH_SCOPES = ["offline_access", MCP_SCOPE]; export function getMcpResource(baseUrl: string) { return new URL(MCP_RESOURCE_PATH, baseUrl).toString(); } - -export function getMcpOrganizationIdClaim(baseUrl: string) { - return new URL( - `${MCP_RESOURCE_PATH}/claims/organization-id`, - baseUrl, - ).toString(); -} - -export function getMcpProtectedResourceMetadataUrl(resource: string) { - const url = new URL(resource); - const pathname = url.pathname.endsWith("/") - ? url.pathname.slice(0, -1) - : url.pathname; - - return `${url.origin}/.well-known/oauth-protected-resource${pathname}`; -} diff --git a/src/routeTree.gen.ts b/src/routeTree.gen.ts index ca7c478..c3e7a76 100644 --- a/src/routeTree.gen.ts +++ b/src/routeTree.gen.ts @@ -25,13 +25,9 @@ import { Route as AppSupportRouteImport } from './routes/_app/support' import { Route as AppSettingsRouteImport } from './routes/_app/settings' import { Route as AppBillingRouteImport } from './routes/_app/billing' import { Route as AppAiRouteImport } from './routes/_app/ai' -import { Route as DotwellKnownOpenidConfigurationRouteImport } from './routes/[.]well-known/openid-configuration' -import { Route as DotwellKnownOauthProtectedResourceRouteImport } from './routes/[.]well-known/oauth-protected-resource' -import { Route as DotwellKnownOauthAuthorizationServerRouteImport } from './routes/[.]well-known/oauth-authorization-server' import { Route as ApiAutumnSplatRouteImport } from './routes/api/autumn/$' import { Route as ApiAuthSplatRouteImport } from './routes/api/auth/$' import { Route as AppHelpDataforseoApiKeyRouteImport } from './routes/_app/help/dataforseo-api-key' -import { Route as DotwellKnownOauthProtectedResourceMcpRouteImport } from './routes/[.]well-known/oauth-protected-resource/mcp' import { Route as ProjectPProjectIdRouteRouteImport } from './routes/_project/p/$projectId/route' import { Route as ProjectPProjectIdIndexRouteImport } from './routes/_project/p/$projectId/index' import { Route as ProjectPProjectIdSavedRouteImport } from './routes/_project/p/$projectId/saved' @@ -42,7 +38,6 @@ import { Route as ProjectPProjectIdDomainRouteImport } from './routes/_project/p import { Route as ProjectPProjectIdBrandLookupRouteImport } from './routes/_project/p/$projectId/brand-lookup' import { Route as ProjectPProjectIdBacklinksRouteImport } from './routes/_project/p/$projectId/backlinks' import { Route as ProjectPProjectIdAuditRouteImport } from './routes/_project/p/$projectId/audit' -import { Route as DotwellKnownOauthAuthorizationServerApiAuthRouteImport } from './routes/[.]well-known/oauth-authorization-server/api/auth' import { Route as ProjectPProjectIdRankTrackingIndexRouteImport } from './routes/_project/p/$projectId/rank-tracking/index' import { Route as ProjectPProjectIdAuditIndexRouteImport } from './routes/_project/p/$projectId/audit/index' import { Route as ProjectPProjectIdRankTrackingConfigIdRouteImport } from './routes/_project/p/$projectId/rank-tracking/$configId' @@ -125,24 +120,6 @@ const AppAiRoute = AppAiRouteImport.update({ path: '/ai', getParentRoute: () => AppRouteRoute, } as any) -const DotwellKnownOpenidConfigurationRoute = - DotwellKnownOpenidConfigurationRouteImport.update({ - id: '/.well-known/openid-configuration', - path: '/.well-known/openid-configuration', - getParentRoute: () => rootRouteImport, - } as any) -const DotwellKnownOauthProtectedResourceRoute = - DotwellKnownOauthProtectedResourceRouteImport.update({ - id: '/.well-known/oauth-protected-resource', - path: '/.well-known/oauth-protected-resource', - getParentRoute: () => rootRouteImport, - } as any) -const DotwellKnownOauthAuthorizationServerRoute = - DotwellKnownOauthAuthorizationServerRouteImport.update({ - id: '/.well-known/oauth-authorization-server', - path: '/.well-known/oauth-authorization-server', - getParentRoute: () => rootRouteImport, - } as any) const ApiAutumnSplatRoute = ApiAutumnSplatRouteImport.update({ id: '/api/autumn/$', path: '/api/autumn/$', @@ -158,12 +135,6 @@ const AppHelpDataforseoApiKeyRoute = AppHelpDataforseoApiKeyRouteImport.update({ path: '/help/dataforseo-api-key', getParentRoute: () => AppRouteRoute, } as any) -const DotwellKnownOauthProtectedResourceMcpRoute = - DotwellKnownOauthProtectedResourceMcpRouteImport.update({ - id: '/mcp', - path: '/mcp', - getParentRoute: () => DotwellKnownOauthProtectedResourceRoute, - } as any) const ProjectPProjectIdRouteRoute = ProjectPProjectIdRouteRouteImport.update({ id: '/p/$projectId', path: '/p/$projectId', @@ -219,12 +190,6 @@ const ProjectPProjectIdAuditRoute = ProjectPProjectIdAuditRouteImport.update({ path: '/audit', getParentRoute: () => ProjectPProjectIdRouteRoute, } as any) -const DotwellKnownOauthAuthorizationServerApiAuthRoute = - DotwellKnownOauthAuthorizationServerApiAuthRouteImport.update({ - id: '/api/auth', - path: '/api/auth', - getParentRoute: () => DotwellKnownOauthAuthorizationServerRoute, - } as any) const ProjectPProjectIdRankTrackingIndexRoute = ProjectPProjectIdRankTrackingIndexRouteImport.update({ id: '/', @@ -255,9 +220,6 @@ export interface FileRoutesByFullPath { '/forgot-password': typeof ForgotPasswordRoute '/reset-password': typeof ResetPasswordRoute '/verify-email': typeof VerifyEmailRoute - '/.well-known/oauth-authorization-server': typeof DotwellKnownOauthAuthorizationServerRouteWithChildren - '/.well-known/oauth-protected-resource': typeof DotwellKnownOauthProtectedResourceRouteWithChildren - '/.well-known/openid-configuration': typeof DotwellKnownOpenidConfigurationRoute '/ai': typeof AppAiRoute '/billing': typeof AppBillingRoute '/settings': typeof AppSettingsRoute @@ -267,11 +229,9 @@ export interface FileRoutesByFullPath { '/oauth-consent': typeof AuthenticatedOauthConsentRoute '/subscribe': typeof AuthenticatedSubscribeRoute '/p/$projectId': typeof ProjectPProjectIdRouteRouteWithChildren - '/.well-known/oauth-protected-resource/mcp': typeof DotwellKnownOauthProtectedResourceMcpRoute '/help/dataforseo-api-key': typeof AppHelpDataforseoApiKeyRoute '/api/auth/$': typeof ApiAuthSplatRoute '/api/autumn/$': typeof ApiAutumnSplatRoute - '/.well-known/oauth-authorization-server/api/auth': typeof DotwellKnownOauthAuthorizationServerApiAuthRoute '/p/$projectId/audit': typeof ProjectPProjectIdAuditRouteWithChildren '/p/$projectId/backlinks': typeof ProjectPProjectIdBacklinksRoute '/p/$projectId/brand-lookup': typeof ProjectPProjectIdBrandLookupRoute @@ -291,9 +251,6 @@ export interface FileRoutesByTo { '/forgot-password': typeof ForgotPasswordRoute '/reset-password': typeof ResetPasswordRoute '/verify-email': typeof VerifyEmailRoute - '/.well-known/oauth-authorization-server': typeof DotwellKnownOauthAuthorizationServerRouteWithChildren - '/.well-known/oauth-protected-resource': typeof DotwellKnownOauthProtectedResourceRouteWithChildren - '/.well-known/openid-configuration': typeof DotwellKnownOpenidConfigurationRoute '/ai': typeof AppAiRoute '/billing': typeof AppBillingRoute '/settings': typeof AppSettingsRoute @@ -302,11 +259,9 @@ export interface FileRoutesByTo { '/sign-up': typeof AuthSignUpRoute '/oauth-consent': typeof AuthenticatedOauthConsentRoute '/subscribe': typeof AuthenticatedSubscribeRoute - '/.well-known/oauth-protected-resource/mcp': typeof DotwellKnownOauthProtectedResourceMcpRoute '/help/dataforseo-api-key': typeof AppHelpDataforseoApiKeyRoute '/api/auth/$': typeof ApiAuthSplatRoute '/api/autumn/$': typeof ApiAutumnSplatRoute - '/.well-known/oauth-authorization-server/api/auth': typeof DotwellKnownOauthAuthorizationServerApiAuthRoute '/p/$projectId/backlinks': typeof ProjectPProjectIdBacklinksRoute '/p/$projectId/brand-lookup': typeof ProjectPProjectIdBrandLookupRoute '/p/$projectId/domain': typeof ProjectPProjectIdDomainRoute @@ -328,9 +283,6 @@ export interface FileRoutesById { '/forgot-password': typeof ForgotPasswordRoute '/reset-password': typeof ResetPasswordRoute '/verify-email': typeof VerifyEmailRoute - '/.well-known/oauth-authorization-server': typeof DotwellKnownOauthAuthorizationServerRouteWithChildren - '/.well-known/oauth-protected-resource': typeof DotwellKnownOauthProtectedResourceRouteWithChildren - '/.well-known/openid-configuration': typeof DotwellKnownOpenidConfigurationRoute '/_app/ai': typeof AppAiRoute '/_app/billing': typeof AppBillingRoute '/_app/settings': typeof AppSettingsRoute @@ -341,11 +293,9 @@ export interface FileRoutesById { '/_authenticated/subscribe': typeof AuthenticatedSubscribeRoute '/_app/': typeof AppIndexRoute '/_project/p/$projectId': typeof ProjectPProjectIdRouteRouteWithChildren - '/.well-known/oauth-protected-resource/mcp': typeof DotwellKnownOauthProtectedResourceMcpRoute '/_app/help/dataforseo-api-key': typeof AppHelpDataforseoApiKeyRoute '/api/auth/$': typeof ApiAuthSplatRoute '/api/autumn/$': typeof ApiAutumnSplatRoute - '/.well-known/oauth-authorization-server/api/auth': typeof DotwellKnownOauthAuthorizationServerApiAuthRoute '/_project/p/$projectId/audit': typeof ProjectPProjectIdAuditRouteWithChildren '/_project/p/$projectId/backlinks': typeof ProjectPProjectIdBacklinksRoute '/_project/p/$projectId/brand-lookup': typeof ProjectPProjectIdBrandLookupRoute @@ -367,9 +317,6 @@ export interface FileRouteTypes { | '/forgot-password' | '/reset-password' | '/verify-email' - | '/.well-known/oauth-authorization-server' - | '/.well-known/oauth-protected-resource' - | '/.well-known/openid-configuration' | '/ai' | '/billing' | '/settings' @@ -379,11 +326,9 @@ export interface FileRouteTypes { | '/oauth-consent' | '/subscribe' | '/p/$projectId' - | '/.well-known/oauth-protected-resource/mcp' | '/help/dataforseo-api-key' | '/api/auth/$' | '/api/autumn/$' - | '/.well-known/oauth-authorization-server/api/auth' | '/p/$projectId/audit' | '/p/$projectId/backlinks' | '/p/$projectId/brand-lookup' @@ -403,9 +348,6 @@ export interface FileRouteTypes { | '/forgot-password' | '/reset-password' | '/verify-email' - | '/.well-known/oauth-authorization-server' - | '/.well-known/oauth-protected-resource' - | '/.well-known/openid-configuration' | '/ai' | '/billing' | '/settings' @@ -414,11 +356,9 @@ export interface FileRouteTypes { | '/sign-up' | '/oauth-consent' | '/subscribe' - | '/.well-known/oauth-protected-resource/mcp' | '/help/dataforseo-api-key' | '/api/auth/$' | '/api/autumn/$' - | '/.well-known/oauth-authorization-server/api/auth' | '/p/$projectId/backlinks' | '/p/$projectId/brand-lookup' | '/p/$projectId/domain' @@ -439,9 +379,6 @@ export interface FileRouteTypes { | '/forgot-password' | '/reset-password' | '/verify-email' - | '/.well-known/oauth-authorization-server' - | '/.well-known/oauth-protected-resource' - | '/.well-known/openid-configuration' | '/_app/ai' | '/_app/billing' | '/_app/settings' @@ -452,11 +389,9 @@ export interface FileRouteTypes { | '/_authenticated/subscribe' | '/_app/' | '/_project/p/$projectId' - | '/.well-known/oauth-protected-resource/mcp' | '/_app/help/dataforseo-api-key' | '/api/auth/$' | '/api/autumn/$' - | '/.well-known/oauth-authorization-server/api/auth' | '/_project/p/$projectId/audit' | '/_project/p/$projectId/backlinks' | '/_project/p/$projectId/brand-lookup' @@ -480,9 +415,6 @@ export interface RootRouteChildren { ForgotPasswordRoute: typeof ForgotPasswordRoute ResetPasswordRoute: typeof ResetPasswordRoute VerifyEmailRoute: typeof VerifyEmailRoute - DotwellKnownOauthAuthorizationServerRoute: typeof DotwellKnownOauthAuthorizationServerRouteWithChildren - DotwellKnownOauthProtectedResourceRoute: typeof DotwellKnownOauthProtectedResourceRouteWithChildren - DotwellKnownOpenidConfigurationRoute: typeof DotwellKnownOpenidConfigurationRoute ApiAuthSplatRoute: typeof ApiAuthSplatRoute ApiAutumnSplatRoute: typeof ApiAutumnSplatRoute } @@ -601,27 +533,6 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof AppAiRouteImport parentRoute: typeof AppRouteRoute } - '/.well-known/openid-configuration': { - id: '/.well-known/openid-configuration' - path: '/.well-known/openid-configuration' - fullPath: '/.well-known/openid-configuration' - preLoaderRoute: typeof DotwellKnownOpenidConfigurationRouteImport - parentRoute: typeof rootRouteImport - } - '/.well-known/oauth-protected-resource': { - id: '/.well-known/oauth-protected-resource' - path: '/.well-known/oauth-protected-resource' - fullPath: '/.well-known/oauth-protected-resource' - preLoaderRoute: typeof DotwellKnownOauthProtectedResourceRouteImport - parentRoute: typeof rootRouteImport - } - '/.well-known/oauth-authorization-server': { - id: '/.well-known/oauth-authorization-server' - path: '/.well-known/oauth-authorization-server' - fullPath: '/.well-known/oauth-authorization-server' - preLoaderRoute: typeof DotwellKnownOauthAuthorizationServerRouteImport - parentRoute: typeof rootRouteImport - } '/api/autumn/$': { id: '/api/autumn/$' path: '/api/autumn/$' @@ -643,13 +554,6 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof AppHelpDataforseoApiKeyRouteImport parentRoute: typeof AppRouteRoute } - '/.well-known/oauth-protected-resource/mcp': { - id: '/.well-known/oauth-protected-resource/mcp' - path: '/mcp' - fullPath: '/.well-known/oauth-protected-resource/mcp' - preLoaderRoute: typeof DotwellKnownOauthProtectedResourceMcpRouteImport - parentRoute: typeof DotwellKnownOauthProtectedResourceRoute - } '/_project/p/$projectId': { id: '/_project/p/$projectId' path: '/p/$projectId' @@ -720,13 +624,6 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof ProjectPProjectIdAuditRouteImport parentRoute: typeof ProjectPProjectIdRouteRoute } - '/.well-known/oauth-authorization-server/api/auth': { - id: '/.well-known/oauth-authorization-server/api/auth' - path: '/api/auth' - fullPath: '/.well-known/oauth-authorization-server/api/auth' - preLoaderRoute: typeof DotwellKnownOauthAuthorizationServerApiAuthRouteImport - parentRoute: typeof DotwellKnownOauthAuthorizationServerRoute - } '/_project/p/$projectId/rank-tracking/': { id: '/_project/p/$projectId/rank-tracking/' path: '/' @@ -884,36 +781,6 @@ const AuthenticatedRouteWithChildren = AuthenticatedRoute._addFileChildren( AuthenticatedRouteChildren, ) -interface DotwellKnownOauthAuthorizationServerRouteChildren { - DotwellKnownOauthAuthorizationServerApiAuthRoute: typeof DotwellKnownOauthAuthorizationServerApiAuthRoute -} - -const DotwellKnownOauthAuthorizationServerRouteChildren: DotwellKnownOauthAuthorizationServerRouteChildren = - { - DotwellKnownOauthAuthorizationServerApiAuthRoute: - DotwellKnownOauthAuthorizationServerApiAuthRoute, - } - -const DotwellKnownOauthAuthorizationServerRouteWithChildren = - DotwellKnownOauthAuthorizationServerRoute._addFileChildren( - DotwellKnownOauthAuthorizationServerRouteChildren, - ) - -interface DotwellKnownOauthProtectedResourceRouteChildren { - DotwellKnownOauthProtectedResourceMcpRoute: typeof DotwellKnownOauthProtectedResourceMcpRoute -} - -const DotwellKnownOauthProtectedResourceRouteChildren: DotwellKnownOauthProtectedResourceRouteChildren = - { - DotwellKnownOauthProtectedResourceMcpRoute: - DotwellKnownOauthProtectedResourceMcpRoute, - } - -const DotwellKnownOauthProtectedResourceRouteWithChildren = - DotwellKnownOauthProtectedResourceRoute._addFileChildren( - DotwellKnownOauthProtectedResourceRouteChildren, - ) - const rootRouteChildren: RootRouteChildren = { AppRouteRoute: AppRouteRouteWithChildren, ProjectRouteRoute: ProjectRouteRouteWithChildren, @@ -922,11 +789,6 @@ const rootRouteChildren: RootRouteChildren = { ForgotPasswordRoute: ForgotPasswordRoute, ResetPasswordRoute: ResetPasswordRoute, VerifyEmailRoute: VerifyEmailRoute, - DotwellKnownOauthAuthorizationServerRoute: - DotwellKnownOauthAuthorizationServerRouteWithChildren, - DotwellKnownOauthProtectedResourceRoute: - DotwellKnownOauthProtectedResourceRouteWithChildren, - DotwellKnownOpenidConfigurationRoute: DotwellKnownOpenidConfigurationRoute, ApiAuthSplatRoute: ApiAuthSplatRoute, ApiAutumnSplatRoute: ApiAutumnSplatRoute, } diff --git a/src/routes/[.]well-known/oauth-authorization-server.ts b/src/routes/[.]well-known/oauth-authorization-server.ts deleted file mode 100644 index 5d9acd1..0000000 --- a/src/routes/[.]well-known/oauth-authorization-server.ts +++ /dev/null @@ -1,31 +0,0 @@ -import { oauthProviderAuthServerMetadata } from "@better-auth/oauth-provider"; -import { createFileRoute } from "@tanstack/react-router"; -import { env } from "cloudflare:workers"; -import { getAuth, hasHostedAuthConfig } from "@/lib/auth"; -import { isHostedAuthMode } from "@/lib/auth-mode"; - -function unavailableMetadataResponse() { - if (!isHostedAuthMode(env.AUTH_MODE)) { - return new Response("Not found", { status: 404 }); - } - - return new Response("Missing Better Auth hosted configuration", { - status: 500, - }); -} - -export const Route = createFileRoute("/.well-known/oauth-authorization-server")( - { - server: { - handlers: { - GET: async ({ request }: { request: Request }) => { - if (!isHostedAuthMode(env.AUTH_MODE) || !hasHostedAuthConfig()) { - return unavailableMetadataResponse(); - } - - return oauthProviderAuthServerMetadata(getAuth())(request); - }, - }, - }, - }, -); diff --git a/src/routes/[.]well-known/oauth-authorization-server/api/auth.ts b/src/routes/[.]well-known/oauth-authorization-server/api/auth.ts deleted file mode 100644 index 0b6dc2f..0000000 --- a/src/routes/[.]well-known/oauth-authorization-server/api/auth.ts +++ /dev/null @@ -1,31 +0,0 @@ -import { oauthProviderAuthServerMetadata } from "@better-auth/oauth-provider"; -import { createFileRoute } from "@tanstack/react-router"; -import { env } from "cloudflare:workers"; -import { getAuth, hasHostedAuthConfig } from "@/lib/auth"; -import { isHostedAuthMode } from "@/lib/auth-mode"; - -function unavailableMetadataResponse() { - if (!isHostedAuthMode(env.AUTH_MODE)) { - return new Response("Not found", { status: 404 }); - } - - return new Response("Missing Better Auth hosted configuration", { - status: 500, - }); -} - -export const Route = createFileRoute( - "/.well-known/oauth-authorization-server/api/auth", -)({ - server: { - handlers: { - GET: async ({ request }: { request: Request }) => { - if (!isHostedAuthMode(env.AUTH_MODE) || !hasHostedAuthConfig()) { - return unavailableMetadataResponse(); - } - - return oauthProviderAuthServerMetadata(getAuth())(request); - }, - }, - }, -}); diff --git a/src/routes/[.]well-known/oauth-protected-resource.ts b/src/routes/[.]well-known/oauth-protected-resource.ts deleted file mode 100644 index 8748123..0000000 --- a/src/routes/[.]well-known/oauth-protected-resource.ts +++ /dev/null @@ -1,11 +0,0 @@ -import { createFileRoute } from "@tanstack/react-router"; -import { env } from "cloudflare:workers"; -import { mcpProtectedResourceMetadataResponse } from "@/server/mcp/protected-resource-metadata"; - -export const Route = createFileRoute("/.well-known/oauth-protected-resource")({ - server: { - handlers: { - GET: async () => mcpProtectedResourceMetadataResponse(env.AUTH_MODE), - }, - }, -}); diff --git a/src/routes/[.]well-known/oauth-protected-resource/mcp.ts b/src/routes/[.]well-known/oauth-protected-resource/mcp.ts deleted file mode 100644 index d9b23e5..0000000 --- a/src/routes/[.]well-known/oauth-protected-resource/mcp.ts +++ /dev/null @@ -1,13 +0,0 @@ -import { createFileRoute } from "@tanstack/react-router"; -import { env } from "cloudflare:workers"; -import { mcpProtectedResourceMetadataResponse } from "@/server/mcp/protected-resource-metadata"; - -export const Route = createFileRoute( - "/.well-known/oauth-protected-resource/mcp", -)({ - server: { - handlers: { - GET: async () => mcpProtectedResourceMetadataResponse(env.AUTH_MODE), - }, - }, -}); diff --git a/src/routes/[.]well-known/openid-configuration.ts b/src/routes/[.]well-known/openid-configuration.ts deleted file mode 100644 index 84c10aa..0000000 --- a/src/routes/[.]well-known/openid-configuration.ts +++ /dev/null @@ -1,29 +0,0 @@ -import { oauthProviderOpenIdConfigMetadata } from "@better-auth/oauth-provider"; -import { createFileRoute } from "@tanstack/react-router"; -import { env } from "cloudflare:workers"; -import { getAuth, hasHostedAuthConfig } from "@/lib/auth"; -import { isHostedAuthMode } from "@/lib/auth-mode"; - -function unavailableMetadataResponse() { - if (!isHostedAuthMode(env.AUTH_MODE)) { - return new Response("Not found", { status: 404 }); - } - - return new Response("Missing Better Auth hosted configuration", { - status: 500, - }); -} - -export const Route = createFileRoute("/.well-known/openid-configuration")({ - server: { - handlers: { - GET: async ({ request }: { request: Request }) => { - if (!isHostedAuthMode(env.AUTH_MODE) || !hasHostedAuthConfig()) { - return unavailableMetadataResponse(); - } - - return oauthProviderOpenIdConfigMetadata(getAuth())(request); - }, - }, - }, -}); diff --git a/src/routes/_auth.sign-in.tsx b/src/routes/_auth.sign-in.tsx index 94e42ba..585ddac 100644 --- a/src/routes/_auth.sign-in.tsx +++ b/src/routes/_auth.sign-in.tsx @@ -29,6 +29,7 @@ function SignInPage() { const { redirectTo, oauthQuery, isHostedMode } = useAuthPageState( search.redirect, ); + const authCallbackURL = redirectTo; const [verificationEmail, setVerificationEmail] = useState( null, ); @@ -53,7 +54,7 @@ function SignInPage() { const result = await authClient.signIn.email({ email, password: value.password, - callbackURL: redirectTo, + callbackURL: authCallbackURL, ...(oauthQuery ? { oauth_query: oauthQuery } : {}), }); @@ -103,12 +104,16 @@ function SignInPage() { setIsSendingVerification(true); try { - const callbackURL = new URL("/verify-email", window.location.origin); - if (redirectTo !== "/") - callbackURL.searchParams.set("redirect", redirectTo); + const verificationCallbackURL = new URL( + "/verify-email", + window.location.origin, + ); + if (authCallbackURL !== "/") { + verificationCallbackURL.searchParams.set("redirect", authCallbackURL); + } const result = await authClient.sendVerificationEmail({ email: verificationEmail, - callbackURL: callbackURL.toString(), + callbackURL: verificationCallbackURL.toString(), }); if (result.error) { diff --git a/src/routes/_authenticated.oauth-consent.tsx b/src/routes/_authenticated.oauth-consent.tsx index 2c8e16b..f05b1c4 100644 --- a/src/routes/_authenticated.oauth-consent.tsx +++ b/src/routes/_authenticated.oauth-consent.tsx @@ -1,9 +1,7 @@ -import { useQuery } from "@tanstack/react-query"; import { createFileRoute } from "@tanstack/react-router"; import { Check, Database, KeyRound, User } from "lucide-react"; import { useState } from "react"; -import { authClient, useSession } from "@/lib/auth-client"; -import { getOAuthClientInfo } from "@/serverFunctions/oauth"; +import { useSession } from "@/lib/auth-client"; export const Route = createFileRoute("/_authenticated/oauth-consent")({ component: OAuthConsentPage, @@ -27,42 +25,35 @@ function OAuthConsentPage() { const [isSubmitting, setIsSubmitting] = useState(false); const [error, setError] = useState(null); - const clientId = - typeof window !== "undefined" - ? new URLSearchParams(window.location.search).get("client_id") - : null; - - const clientInfoQuery = useQuery({ - queryKey: ["oauth-client-info", clientId], - queryFn: () => - clientId - ? getOAuthClientInfo({ data: { clientId } }) - : Promise.resolve(null), - enabled: Boolean(clientId), - staleTime: 60_000, - }); - - const clientName = clientInfoQuery.data?.name ?? null; const userEmail = session?.user?.email ?? null; - const isLoadingClient = clientInfoQuery.isLoading; - const named = Boolean(clientName); async function respond(accept: boolean) { setError(null); setIsSubmitting(true); - const { data, error: consentError } = await authClient.oauth2.consent({ - accept, + const response = await fetch("/api/oauth/consent", { + method: "POST", + headers: { + "Content-Type": "application/json", + }, + body: JSON.stringify({ + accept, + query: window.location.search, + }), }); + const data: { + redirectTo?: string; + error?: string; + } = await response.json(); - if (consentError) { - setError(consentError.message ?? "Unable to complete authorization."); + if (!response.ok) { + setError(data.error ?? "Unable to complete authorization."); setIsSubmitting(false); return; } - if (data?.redirect && data.url) { - window.location.assign(data.url); + if (data.redirectTo) { + window.location.assign(data.redirectTo); return; } @@ -78,33 +69,12 @@ function OAuthConsentPage() { alt="OpenSEO" className="size-10 rounded-lg" /> - {isLoadingClient ? ( -
- ) : ( -

- {named ? ( - <> - Authorize {clientName} - - ) : ( - "Authorize MCP access" - )} -

- )} +

Authorize MCP access

- {named - ? `${clientName} is requesting access to your OpenSEO workspace.` - : "An MCP client is requesting access to your OpenSEO workspace."} + An MCP client is requesting access to your OpenSEO workspace.

- {!named && !isLoadingClient ? ( -
- This client did not provide a name during registration. Only continue - if you started this connection yourself. -
- ) : null} - {userEmail ? (
@@ -119,7 +89,7 @@ function OAuthConsentPage() {
- {named ? `This will allow ${clientName} to` : "This will allow it to"} + This will allow it to
    {SCOPES.map((scope) => ( diff --git a/src/routes/api/auth/$.ts b/src/routes/api/auth/$.ts index c49dd94..0e87161 100644 --- a/src/routes/api/auth/$.ts +++ b/src/routes/api/auth/$.ts @@ -1,106 +1,7 @@ import { createFileRoute } from "@tanstack/react-router"; import { env } from "cloudflare:workers"; -import { getAuth, getHostedBaseUrl, hasHostedAuthConfig } from "@/lib/auth"; +import { getAuth, hasHostedAuthConfig } from "@/lib/auth"; import { isHostedAuthMode } from "@/lib/auth-mode"; -import { getMcpResource } from "@/lib/oauth-resource"; - -const TOKEN_PATH = "/api/auth/oauth2/token"; -const REGISTER_PATH = "/api/auth/oauth2/register"; -const PUBLIC_CLIENT_AUTH_METHOD = "none"; - -function isJsonObject(value: unknown): value is Record { - return value != null && typeof value === "object" && !Array.isArray(value); -} - -function hasRequestAuthContext(request: Request) { - return Boolean( - request.headers.get("authorization") || request.headers.get("cookie"), - ); -} - -function requestWithReplacedBody(request: Request, body: BodyInit) { - const headers = new Headers(request.headers); - headers.delete("content-length"); - - return new Request(request.url, { - method: request.method, - headers, - body, - }); -} - -// Inject RFC 8707 `resource` into /oauth2/token requests when the client -// omitted it. Some MCP clients (notably codex as of 2026-05) skip the -// resource indicator, which makes better-auth issue an opaque access token -// (see `checkResource` in @better-auth/oauth-provider — audience comes from -// `ctx.body.resource` at token-issuance time, not from the stored authorize -// query). Without an audience to bind, no `aud` claim → opaque token → no -// local JWT verify on the resource side. -// -// We only have one valid audience (`validAudiences: [mcpResource]` in -// auth-config.ts), so it is safe to default missing resources to it. Remove -// this shim once MCP clients reliably pass `resource` per spec. -export async function maybeInjectMcpResource( - request: Request, -): Promise { - if (request.method !== "POST") return request; - - const url = new URL(request.url); - if (url.pathname !== TOKEN_PATH) return request; - - const contentType = request.headers.get("content-type") ?? ""; - if (!contentType.includes("application/x-www-form-urlencoded")) - return request; - - const body = await request.clone().text(); - const params = new URLSearchParams(body); - if (params.has("resource")) return request; - - params.set("resource", getMcpResource(getHostedBaseUrl())); - - return requestWithReplacedBody(request, params.toString()); -} - -// Some hosted MCP clients attempt unauthenticated DCR while sending a -// confidential-client auth method. Better Auth only permits unauthenticated DCR -// for public clients, so normalize that case to the compatible public shape. -// Claude Desktop hit this path during connector setup: it had no session or -// registration bearer token, but sent a non-`none` token endpoint auth method. -export async function maybeDefaultMcpClientRegistrationAuthMethod( - request: Request, -): Promise { - if (request.method !== "POST") return request; - - const url = new URL(request.url); - if (url.pathname !== REGISTER_PATH) return request; - - const contentType = request.headers.get("content-type") ?? ""; - if (!contentType.includes("application/json")) return request; - - let body: unknown; - try { - body = await request.clone().json(); - } catch { - return request; - } - - if (!isJsonObject(body)) return request; - - if ( - hasRequestAuthContext(request) || - body.token_endpoint_auth_method === PUBLIC_CLIENT_AUTH_METHOD - ) { - return request; - } - - return requestWithReplacedBody( - request, - JSON.stringify({ - ...body, - token_endpoint_auth_method: PUBLIC_CLIENT_AUTH_METHOD, - }), - ); -} async function handleAuthRequest(request: Request) { if (!isHostedAuthMode(env.AUTH_MODE)) { @@ -116,12 +17,7 @@ async function handleAuthRequest(request: Request) { } const auth = getAuth(); - const requestWithRegistrationDefaults = - await maybeDefaultMcpClientRegistrationAuthMethod(request); - const requestWithResource = await maybeInjectMcpResource( - requestWithRegistrationDefaults, - ); - return auth.handler(requestWithResource); + return auth.handler(request); } export const Route = createFileRoute("/api/auth/$")({ diff --git a/src/routes/api/auth/-$.test.ts b/src/routes/api/auth/-$.test.ts deleted file mode 100644 index f23c322..0000000 --- a/src/routes/api/auth/-$.test.ts +++ /dev/null @@ -1,181 +0,0 @@ -import { describe, expect, it, vi } from "vitest"; - -vi.mock("cloudflare:workers", () => ({ - env: { - AUTH_MODE: "hosted", - }, -})); - -vi.mock("@tanstack/react-router", () => ({ - createFileRoute: () => (routeConfig: unknown) => routeConfig, -})); - -vi.mock("@/lib/auth", () => ({ - getAuth: () => ({ handler: vi.fn() }), - getHostedBaseUrl: () => "https://open-seo.test", - hasHostedAuthConfig: () => true, -})); - -describe("maybeInjectMcpResource", () => { - it("injects the MCP resource into form token requests when missing", async () => { - const { maybeInjectMcpResource } = await import("@/routes/api/auth/$"); - const request = new Request("https://open-seo.test/api/auth/oauth2/token", { - method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - "Content-Length": "13", - }, - body: new URLSearchParams({ - grant_type: "authorization_code", - code: "code_123", - }), - }); - - const result = await maybeInjectMcpResource(request); - const params = new URLSearchParams(await result.text()); - - expect(params.get("resource")).toBe("https://open-seo.test/mcp"); - expect(params.get("grant_type")).toBe("authorization_code"); - expect(params.get("code")).toBe("code_123"); - expect(result.headers.has("content-length")).toBe(false); - }); - - it("leaves token requests alone when a resource is already present", async () => { - const { maybeInjectMcpResource } = await import("@/routes/api/auth/$"); - const request = new Request("https://open-seo.test/api/auth/oauth2/token", { - method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - }, - body: new URLSearchParams({ - grant_type: "authorization_code", - resource: "https://other-resource.test/mcp", - }), - }); - - await expect(maybeInjectMcpResource(request)).resolves.toBe(request); - }); - - it("skips requests that are not matching form POST token requests", async () => { - const { maybeInjectMcpResource } = await import("@/routes/api/auth/$"); - const requests = [ - new Request("https://open-seo.test/api/auth/oauth2/token", { - method: "GET", - }), - new Request("https://open-seo.test/api/auth/oauth2/authorize", { - method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - }, - body: new URLSearchParams({ grant_type: "authorization_code" }), - }), - new Request("https://open-seo.test/api/auth/oauth2/token", { - method: "POST", - headers: { - "Content-Type": "application/json", - "Content-Length": "13", - }, - body: JSON.stringify({ grant_type: "authorization_code" }), - }), - ]; - - for (const request of requests) { - await expect(maybeInjectMcpResource(request)).resolves.toBe(request); - } - }); -}); - -describe("maybeDefaultMcpClientRegistrationAuthMethod", () => { - it("defaults JSON dynamic client registration to a public client when omitted", async () => { - const { maybeDefaultMcpClientRegistrationAuthMethod } = - await import("@/routes/api/auth/$"); - const request = new Request( - "https://open-seo.test/api/auth/oauth2/register", - { - method: "POST", - headers: { - "Content-Type": "application/json", - }, - body: JSON.stringify({ - client_name: "Claude", - redirect_uris: ["https://claude.ai/api/mcp/auth_callback"], - }), - }, - ); - - const result = await maybeDefaultMcpClientRegistrationAuthMethod(request); - const body = await result.json(); - - expect(body).toMatchObject({ - client_name: "Claude", - redirect_uris: ["https://claude.ai/api/mcp/auth_callback"], - token_endpoint_auth_method: "none", - }); - expect(result.headers.has("content-length")).toBe(false); - }); - - it("forces unauthenticated dynamic client registration to a public client", async () => { - const { maybeDefaultMcpClientRegistrationAuthMethod } = - await import("@/routes/api/auth/$"); - const request = new Request( - "https://open-seo.test/api/auth/oauth2/register", - { - method: "POST", - headers: { - "Content-Type": "application/json", - }, - body: JSON.stringify({ - token_endpoint_auth_method: "client_secret_basic", - }), - }, - ); - - const result = await maybeDefaultMcpClientRegistrationAuthMethod(request); - const body = await result.json(); - - expect(body).toMatchObject({ - token_endpoint_auth_method: "none", - }); - }); - - it("leaves explicit registration auth methods alone when auth context exists", async () => { - const { maybeDefaultMcpClientRegistrationAuthMethod } = - await import("@/routes/api/auth/$"); - const request = new Request( - "https://open-seo.test/api/auth/oauth2/register", - { - method: "POST", - headers: { - "Content-Type": "application/json", - Cookie: "better-auth.session_token=session_123", - }, - body: JSON.stringify({ - token_endpoint_auth_method: "client_secret_basic", - }), - }, - ); - - await expect( - maybeDefaultMcpClientRegistrationAuthMethod(request), - ).resolves.toBe(request); - }); - - it("leaves malformed JSON registration requests for the auth handler", async () => { - const { maybeDefaultMcpClientRegistrationAuthMethod } = - await import("@/routes/api/auth/$"); - const request = new Request( - "https://open-seo.test/api/auth/oauth2/register", - { - method: "POST", - headers: { - "Content-Type": "application/json", - }, - body: "{", - }, - ); - - await expect( - maybeDefaultMcpClientRegistrationAuthMethod(request), - ).resolves.toBe(request); - }); -}); diff --git a/src/server.ts b/src/server.ts index 303ec8f..5bef2b8 100644 --- a/src/server.ts +++ b/src/server.ts @@ -6,21 +6,34 @@ import { RankTrackingRepository } from "@/server/features/rank-tracking/reposito import { beginRankCheckRun } from "@/server/features/rank-tracking/services/rankCheckRunGuards"; import { customerHasPaidPlan } from "@/server/billing/subscription"; import { isHostedServerAuthMode } from "@/server/lib/runtime-env"; +import { isHostedAuthMode } from "@/lib/auth-mode"; +import { + createOpenSeoOAuthProvider, + type OpenSeoOAuthEnv, +} from "@/server/mcp/oauth-provider"; +import { requestWithPublicOrigin } from "@/server/mcp/public-origin"; import { computeNextCheckAt } from "@/shared/rank-tracking"; -import { handleMcpRequest, MCP_ROUTE } from "@/server/mcp/handler"; const appFetch = createStartHandler(defaultStreamHandler); -const fetch = ( +const handleAppFetch = (request: Request): Response | Promise => + appFetch(request); +const openSeoOAuthProvider = createOpenSeoOAuthProvider(handleAppFetch); + +function fetch( request: Request, env: Env, ctx: ExecutionContext, -): Response | Promise => { - if (new URL(request.url).pathname === MCP_ROUTE) { - return handleMcpRequest(request, env, ctx); +): Response | Promise { + if (isHostedAuthMode(env.AUTH_MODE)) { + return openSeoOAuthProvider.fetch( + requestWithPublicOrigin(request), + env as OpenSeoOAuthEnv, + ctx, + ); } - return appFetch(request); -}; + return handleAppFetch(request); +} // Export Workflow classes as named exports export { SiteAuditWorkflow } from "./server/workflows/SiteAuditWorkflow"; diff --git a/src/server/mcp/context.test.ts b/src/server/mcp/context.test.ts new file mode 100644 index 0000000..676f480 --- /dev/null +++ b/src/server/mcp/context.test.ts @@ -0,0 +1,43 @@ +import { describe, expect, it } from "vitest"; +import { + createWorkersOAuthMcpProps, + MCP_AUTH_CONTEXT_PROP, + withWorkersOAuthMcpScopes, + workersOAuthMcpPropsSchema, +} from "@/server/mcp/context"; + +const mcpContext = { + userId: "user_123", + userEmail: "alice@example.com", + organizationId: "org_123", + clientId: "client_123", + scopes: ["offline_access", "mcp"], + audience: "https://open-seo.test/mcp", + subject: "user_123", + baseUrl: "https://open-seo.test", +}; + +describe("withWorkersOAuthMcpScopes", () => { + it("stores the OpenSEO MCP context in Workers OAuth props", () => { + const props = createWorkersOAuthMcpProps(mcpContext); + + expect(workersOAuthMcpPropsSchema.parse(props)).toEqual({ + [MCP_AUTH_CONTEXT_PROP]: mcpContext, + }); + }); + + it("updates access-token props with downscoped token scopes", () => { + const props = createWorkersOAuthMcpProps(mcpContext); + + expect(withWorkersOAuthMcpScopes(props, ["mcp"])).toEqual({ + [MCP_AUTH_CONTEXT_PROP]: { + ...mcpContext, + scopes: ["mcp"], + }, + }); + }); + + it("leaves unrecognized provider props alone", () => { + expect(withWorkersOAuthMcpScopes({}, ["mcp"])).toBeUndefined(); + }); +}); diff --git a/src/server/mcp/context.ts b/src/server/mcp/context.ts index 39ba0f4..ac595df 100644 --- a/src/server/mcp/context.ts +++ b/src/server/mcp/context.ts @@ -1,4 +1,9 @@ -import { getMcpAuthContext } from "agents/mcp"; +import type { + ServerNotification, + ServerRequest, +} from "@modelcontextprotocol/sdk/types.js"; +import type { RequestHandlerExtra } from "@modelcontextprotocol/sdk/shared/protocol.js"; +import { AsyncLocalStorage } from "node:async_hooks"; import { z } from "zod"; import type { BillingCustomerContext } from "@/server/billing/subscription"; import { buildDashboardUrl } from "@/server/mcp/urls"; @@ -14,6 +19,7 @@ type McpAuth = { }; export const MCP_AUTH_CONTEXT_PROP = "openSeoAuth"; +export const MCP_ROUTE = "/mcp"; const mcpToolAuthContextSchema = z.object({ userId: z.string().min(1), @@ -28,10 +34,48 @@ const mcpToolAuthContextSchema = z.object({ type McpToolAuthContext = z.infer; -export type ToolExtra = unknown; +export type ToolExtra = RequestHandlerExtra; -export function requireMcpToolAuthContext(): McpToolAuthContext { - const rawContext = getMcpAuthContext()?.props[MCP_AUTH_CONTEXT_PROP]; +export const workersOAuthMcpPropsSchema = z.object({ + [MCP_AUTH_CONTEXT_PROP]: mcpToolAuthContextSchema, +}); + +const mcpToolAuthContextStorage = new AsyncLocalStorage(); + +export function createWorkersOAuthMcpProps( + context: McpToolAuthContext, +): Record { + return { + [MCP_AUTH_CONTEXT_PROP]: context, + }; +} + +export function withWorkersOAuthMcpScopes( + props: unknown, + scopes: string[], +): Record | undefined { + const result = workersOAuthMcpPropsSchema.safeParse(props); + if (!result.success) return undefined; + + return createWorkersOAuthMcpProps({ + ...result.data[MCP_AUTH_CONTEXT_PROP], + scopes, + }); +} + +export function runWithMcpToolAuthContext( + context: McpToolAuthContext, + callback: () => T, +) { + return mcpToolAuthContextStorage.run(context, callback); +} + +export function requireMcpToolAuthContext( + extra: ToolExtra, +): McpToolAuthContext { + const rawContext = + mcpToolAuthContextStorage.getStore() ?? + extra.authInfo?.extra?.[MCP_AUTH_CONTEXT_PROP]; const result = mcpToolAuthContextSchema.safeParse(rawContext); if (!result.success) { @@ -41,13 +85,13 @@ export function requireMcpToolAuthContext(): McpToolAuthContext { return result.data; } -export function getAuth(_extra?: ToolExtra): McpAuth { - const { baseUrl: _baseUrl, ...auth } = requireMcpToolAuthContext(); +export function getAuth(extra: ToolExtra): McpAuth { + const { baseUrl: _baseUrl, ...auth } = requireMcpToolAuthContext(extra); return auth; } -export function getBaseUrl(_extra?: ToolExtra): string { - return requireMcpToolAuthContext().baseUrl; +export function getBaseUrl(extra: ToolExtra): string { + return requireMcpToolAuthContext(extra).baseUrl; } export function buildBillingCustomer( diff --git a/src/server/mcp/handler.test.ts b/src/server/mcp/handler.test.ts deleted file mode 100644 index 66be24a..0000000 --- a/src/server/mcp/handler.test.ts +++ /dev/null @@ -1,332 +0,0 @@ -import type { CreateMcpHandlerOptions } from "agents/mcp"; -import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { z } from "zod"; -import { MCP_AUTH_CONTEXT_PROP } from "@/server/mcp/context"; - -const verifyMocks = vi.hoisted(() => ({ - verifyJwsAccessToken: vi.fn(), -})); - -const userEmailMocks = vi.hoisted(() => ({ - getMcpUserEmail: vi.fn(), -})); - -const serverMocks = vi.hoisted(() => ({ - nextServerId: 0, - createdServerIds: [] as number[], - serverIds: new WeakMap(), -})); - -vi.mock("@/lib/auth", () => ({ - getAuth: () => ({ api: { getJwks: vi.fn() } }), - getHostedBaseUrl: () => "https://open-seo.test", - hasHostedAuthConfig: () => true, -})); - -vi.mock("better-auth/oauth2", () => ({ - verifyJwsAccessToken: verifyMocks.verifyJwsAccessToken, -})); - -vi.mock("@/server/mcp/server", () => ({ - createOpenSeoMcpServer: () => { - serverMocks.nextServerId += 1; - const server = new McpServer({ name: "Test MCP", version: "0.0.0" }); - serverMocks.createdServerIds.push(serverMocks.nextServerId); - serverMocks.serverIds.set(server, serverMocks.nextServerId); - return server; - }, -})); - -vi.mock("@/server/mcp/user-email", () => ({ - getMcpUserEmail: userEmailMocks.getMcpUserEmail, -})); - -vi.mock("agents/mcp", () => ({ - createMcpHandler: (_server: McpServer, options: CreateMcpHandlerOptions) => { - return async () => - new Response( - JSON.stringify({ - serverId: serverMocks.serverIds.get(_server), - options, - }), - { - status: 200, - headers: { "Content-Type": "application/json" }, - }, - ); - }, -})); - -const ctx: ExecutionContext = { - waitUntil() {}, - passThroughOnException() {}, - props: {}, -}; - -const transportOptionsSchema = z.object({ - serverId: z.number().optional(), - options: z.object({ - route: z.string().optional(), - enableJsonResponse: z.boolean().optional(), - authContext: z - .object({ - props: z.record(z.string(), z.unknown()), - }) - .optional(), - }), -}); - -function createMcpRequest({ - method = "tools/call", - token, -}: { - method?: string; - token?: string; -} = {}) { - const headers = new Headers({ - Accept: "application/json, text/event-stream", - "Content-Type": "application/json", - }); - - if (token) { - headers.set("Authorization", `Bearer ${token}`); - } - - return new Request("https://open-seo.test/mcp", { - method: "POST", - headers, - body: JSON.stringify({ - jsonrpc: "2.0", - id: 1, - method, - params: - method === "tools/call" - ? { - name: "whoami", - arguments: {}, - } - : undefined, - }), - }); -} - -const jwtShapedToken = "header.payload.signature"; -const organizationIdClaim = "https://open-seo.test/mcp/claims/organization-id"; - -function createAccessTokenPayload( - overrides: Record = {}, -): Record { - return { - sub: "user_123", - azp: "client_123", - scope: "offline_access mcp", - aud: "https://open-seo.test/mcp", - [organizationIdClaim]: "org_123", - ...overrides, - }; -} - -describe("handleMcpRequest", () => { - beforeEach(() => { - vi.clearAllMocks(); - serverMocks.nextServerId = 0; - serverMocks.createdServerIds = []; - serverMocks.serverIds = new WeakMap(); - verifyMocks.verifyJwsAccessToken.mockResolvedValue( - createAccessTokenPayload(), - ); - userEmailMocks.getMcpUserEmail.mockResolvedValue("alice@example.com"); - }); - - it("accepts access tokens verified by Better Auth", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - - const response = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const body = transportOptionsSchema.parse(await response.json()); - - expect(response.status).toBe(200); - expect( - body.options.authContext?.props[MCP_AUTH_CONTEXT_PROP], - ).toMatchObject({ - userId: "user_123", - userEmail: "alice@example.com", - organizationId: "org_123", - clientId: "client_123", - scopes: ["offline_access", "mcp"], - audience: "https://open-seo.test/mcp", - subject: "user_123", - baseUrl: "https://open-seo.test", - }); - expect(body.options.route).toBe("/mcp"); - expect(body.options.enableJsonResponse).toBe(true); - - const functionMatcher: unknown = expect.any(Function); - expect(verifyMocks.verifyJwsAccessToken).toHaveBeenCalledWith( - jwtShapedToken, - expect.objectContaining({ - verifyOptions: { - audience: "https://open-seo.test/mcp", - issuer: "https://open-seo.test/api/auth", - }, - jwksFetch: functionMatcher, - }), - ); - }); - - it("creates a fresh server for each request without persisted transport state", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - - const first = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const second = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const firstBody = transportOptionsSchema.parse(await first.json()); - const secondBody = transportOptionsSchema.parse(await second.json()); - - expect(serverMocks.createdServerIds).toEqual([1, 2]); - expect(firstBody.serverId).toBe(1); - expect(secondBody.serverId).toBe(2); - expect(firstBody.options).not.toHaveProperty("sessionIdGenerator"); - expect(firstBody.options).not.toHaveProperty("storage"); - expect(firstBody.options).not.toHaveProperty("transport"); - }); - - it("lets the MCP transport handle OPTIONS without token verification", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - - const response = await handleMcpRequest( - new Request("https://open-seo.test/mcp", { method: "OPTIONS" }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const body = transportOptionsSchema.parse(await response.json()); - - expect(response.status).toBe(200); - expect(verifyMocks.verifyJwsAccessToken).not.toHaveBeenCalled(); - expect(body.options.authContext).toBeUndefined(); - }); - - it("lets unauthenticated initialize reach the MCP transport", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - - const response = await handleMcpRequest( - createMcpRequest({ method: "initialize" }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const body = transportOptionsSchema.parse(await response.json()); - - expect(response.status).toBe(200); - expect(verifyMocks.verifyJwsAccessToken).not.toHaveBeenCalled(); - expect(body.options.authContext).toBeUndefined(); - }); - - it("lets unauthenticated tools/list reach the MCP transport", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - - const response = await handleMcpRequest( - createMcpRequest({ method: "tools/list" }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - const body = transportOptionsSchema.parse(await response.json()); - - expect(response.status).toBe(200); - expect(verifyMocks.verifyJwsAccessToken).not.toHaveBeenCalled(); - expect(body.options.authContext).toBeUndefined(); - }); - - it("returns 401 when Better Auth rejects the access token", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - verifyMocks.verifyJwsAccessToken.mockRejectedValue( - new Error("invalid audience"), - ); - - const response = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - - expect(response.status).toBe(401); - expect(response.headers.get("WWW-Authenticate")).toBe( - 'Bearer resource_metadata="https://open-seo.test/.well-known/oauth-protected-resource/mcp"', - ); - }); - - it("returns 403 when the verified token is missing MCP organization context", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - verifyMocks.verifyJwsAccessToken.mockResolvedValue( - createAccessTokenPayload({ - [organizationIdClaim]: undefined, - }), - ); - - const response = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - - expect(response.status).toBe(403); - }); - - it("returns 403 when the verified user is not found", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - userEmailMocks.getMcpUserEmail.mockResolvedValue(null); - - const response = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - - expect(response.status).toBe(403); - }); - - it("returns 401 when the token is missing the required mcp scope", async () => { - const { handleMcpRequest } = await import("@/server/mcp/handler"); - verifyMocks.verifyJwsAccessToken.mockResolvedValue( - createAccessTokenPayload({ scope: "offline_access" }), - ); - - const response = await handleMcpRequest( - createMcpRequest({ token: jwtShapedToken }), - { - AUTH_MODE: "hosted", - }, - ctx, - ); - - expect(response.status).toBe(401); - }); -}); diff --git a/src/server/mcp/handler.ts b/src/server/mcp/handler.ts deleted file mode 100644 index d72d958..0000000 --- a/src/server/mcp/handler.ts +++ /dev/null @@ -1,191 +0,0 @@ -import { createMcpHandler } from "agents/mcp"; -import { verifyJwsAccessToken } from "better-auth/oauth2"; -import type { JWTPayload } from "jose"; -import { getAuth, getHostedBaseUrl, hasHostedAuthConfig } from "@/lib/auth"; -import { isHostedAuthMode } from "@/lib/auth-mode"; -import { - getMcpOrganizationIdClaim, - getMcpProtectedResourceMetadataUrl, - getMcpResource, - MCP_SCOPE, -} from "@/lib/oauth-resource"; -import { MCP_AUTH_CONTEXT_PROP } from "@/server/mcp/context"; -import { createOpenSeoMcpServer } from "@/server/mcp/server"; -import { getMcpUserEmail } from "@/server/mcp/user-email"; - -// MCP request flow: -// 1. Resource (`resource=`) is injected into /oauth2/token requests by -// `routes/api/auth/$.ts` so Better Auth always issues audience-bound JWTs -// (some MCP clients skip RFC 8707; without it tokens would be opaque). -// 2. Here we verify the JWT in-process via `verifyJwsAccessToken`, reading -// the JWKS through `auth.api.getJwks()` rather than HTTP self-fetching -// `/api/auth/jwks` (which 500s under workerd dev's self-routing and is -// pointless in prod since the auth server and resource server are the -// same Worker). -// 3. We expect `iss = baseURL + basePath` (basePath defaults to `/api/auth`) -// and `aud = mcpResource`, both confirmed against the published -// /.well-known/oauth-authorization-server metadata. -export const MCP_ROUTE = "/mcp"; - -type McpAccessTokenPayload = JWTPayload & { - azp?: unknown; - client_id?: unknown; - scope?: unknown; -}; - -type McpJsonRpcRequest = { - method?: unknown; -}; - -const PUBLIC_MCP_METHODS = new Set([ - "initialize", - "notifications/initialized", - "ping", - "tools/list", -]); - -function getTokenScopes(payload: McpAccessTokenPayload) { - return typeof payload.scope === "string" - ? payload.scope.split(/\s+/).filter(Boolean) - : []; -} - -function getStringClaim(payload: Record, claim: string) { - const value = payload[claim]; - return typeof value === "string" && value.length > 0 ? value : null; -} - -function unauthorizedResponse(resource: string) { - return new Response("Unauthorized", { - status: 401, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Expose-Headers": "WWW-Authenticate", - "WWW-Authenticate": `Bearer resource_metadata="${getMcpProtectedResourceMetadataUrl( - resource, - )}"`, - }, - }); -} - -async function getMcpJsonRpcMethod(request: Request) { - if (request.method !== "POST") return null; - - try { - const body: McpJsonRpcRequest = await request.clone().json(); - return typeof body.method === "string" ? body.method : null; - } catch { - return null; - } -} - -function isPublicMcpRequest(request: Request, jsonRpcMethod: string | null) { - if (request.method === "OPTIONS") return true; - return jsonRpcMethod != null && PUBLIC_MCP_METHODS.has(jsonRpcMethod); -} - -export async function handleMcpRequest( - request: Request, - env: { AUTH_MODE?: unknown }, - ctx: ExecutionContext, -) { - const authMode = - typeof env.AUTH_MODE === "string" ? env.AUTH_MODE : undefined; - const jsonRpcMethod = await getMcpJsonRpcMethod(request); - - if (!isHostedAuthMode(authMode)) { - return new Response("Not found", { status: 404 }); - } - - if (!hasHostedAuthConfig()) { - return new Response("Missing Better Auth hosted configuration", { - status: 500, - }); - } - - const baseUrl = getHostedBaseUrl(); - const auth = getAuth(); - const mcpResource = getMcpResource(baseUrl); - const issuer = `${baseUrl}/api/auth`; - const organizationIdClaim = getMcpOrganizationIdClaim(baseUrl); - const server = createOpenSeoMcpServer(); - - if (isPublicMcpRequest(request, jsonRpcMethod)) { - return createMcpHandler(server, { - route: MCP_ROUTE, - enableJsonResponse: true, - })(request, env, ctx); - } - - const accessToken = - request.headers - .get("Authorization") - ?.replace(/^Bearer\s+/i, "") - .trim() || undefined; - - let payload: McpAccessTokenPayload; - try { - if (!accessToken) throw new Error("missing access token"); - payload = await verifyJwsAccessToken(accessToken, { - jwksFetch: () => auth.api.getJwks(), - verifyOptions: { audience: mcpResource, issuer }, - }); - } catch { - return unauthorizedResponse(mcpResource); - } - - const scopes = getTokenScopes(payload); - if (!scopes.includes(MCP_SCOPE)) { - return unauthorizedResponse(mcpResource); - } - - const userId = getStringClaim(payload, "sub"); - const organizationId = getStringClaim(payload, organizationIdClaim); - const clientId = - getStringClaim(payload, "azp") ?? getStringClaim(payload, "client_id"); - - if (!userId || !organizationId) { - return new Response( - userId - ? "MCP organization context required" - : "MCP user context required", - { - status: 403, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Expose-Headers": "WWW-Authenticate", - }, - }, - ); - } - - const userEmail = await getMcpUserEmail(userId); - if (!userEmail) { - return new Response("MCP user context required", { - status: 403, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Expose-Headers": "WWW-Authenticate", - }, - }); - } - - return createMcpHandler(server, { - route: MCP_ROUTE, - enableJsonResponse: true, - authContext: { - props: { - [MCP_AUTH_CONTEXT_PROP]: { - userId, - userEmail, - organizationId, - clientId, - scopes, - audience: mcpResource, - subject: userId, - baseUrl, - }, - }, - }, - })(request, env, ctx); -} diff --git a/src/server/mcp/oauth-provider.ts b/src/server/mcp/oauth-provider.ts new file mode 100644 index 0000000..df42ba3 --- /dev/null +++ b/src/server/mcp/oauth-provider.ts @@ -0,0 +1,407 @@ +import { + OAuthProvider, + type AuthRequest, + type OAuthHelpers, + type OAuthProviderOptions, +} from "@cloudflare/workers-oauth-provider"; +import { z } from "zod"; +import { getHostedBaseUrl } from "@/lib/auth"; +import { + getMcpResource, + MCP_OAUTH_SCOPES, + MCP_SCOPE, +} from "@/lib/oauth-resource"; +import { asAppError } from "@/server/lib/errors"; +import { + createWorkersOAuthMcpProps, + MCP_ROUTE, + withWorkersOAuthMcpScopes, +} from "@/server/mcp/context"; +import { getPublicOrigin } from "@/server/mcp/public-origin"; +import { handleAuthenticatedOpenSeoMcpRequest } from "@/server/mcp/transport"; +import { resolveHostedContext } from "@/middleware/ensure-user/hosted"; + +const OAUTH_AUTHORIZE_PATH = "/api/auth/oauth2/authorize"; +const OAUTH_TOKEN_PATH = "/api/auth/oauth2/token"; +const OAUTH_REGISTER_PATH = "/api/auth/oauth2/register"; + +const OAUTH_CONSENT_RESPONSE_PATH = "/api/oauth/consent"; +const WWW_AUTHENTICATE_HEADER = "WWW-Authenticate"; +const OAUTH_AUTHORIZATION_PARAM_NAMES = [ + "response_type", + "client_id", + "redirect_uri", + "scope", + "state", + "code_challenge", + "code_challenge_method", + "resource", +] as const; +const REFRESH_TOKEN_TTL_SECONDS = 60 * 60 * 24 * 30; + +export type OpenSeoOAuthEnv = Env & { + OAUTH_KV: KVNamespace; + OAUTH_PROVIDER?: OAuthHelpers; +}; + +type AppFetch = (request: Request) => Response | Promise; + +type OAuthExecutionContext = ExecutionContext & { + props?: unknown; +}; + +type ExportedHandlerWithFetch = ExportedHandler & { + fetch: NonNullable["fetch"]>; +}; + +const consentResponseSchema = z.object({ + accept: z.boolean(), + query: z.string(), +}); + +function getOAuthHelpers(env: OpenSeoOAuthEnv) { + if (!env.OAUTH_PROVIDER) { + throw new Error("OAuth provider helpers are unavailable"); + } + + return env.OAUTH_PROVIDER; +} + +function getMcpResourceForRequest(request: Request) { + return getMcpResource(getPublicOrigin(request)); +} + +function getRelativeRequestTarget(request: Request) { + const url = new URL(request.url); + return `${url.pathname}${url.search}`; +} + +function redirectToSignIn(request: Request) { + const signInUrl = new URL("/sign-in", request.url); + signInUrl.searchParams.set("redirect", getRelativeRequestTarget(request)); + return Response.redirect(signInUrl.toString(), 302); +} + +function invalidOAuthRequestResponse(error: unknown) { + return new Response( + error instanceof Error ? error.message : "Invalid OAuth request", + { status: 400 }, + ); +} + +function jsonResponse(body: unknown, init?: ResponseInit) { + const headers = new Headers(init?.headers); + headers.set("Content-Type", "application/json"); + + return new Response(JSON.stringify(body), { + ...init, + headers, + }); +} + +function oauthErrorResponse(error: { + code: string; + description: string; + status: number; + headers: Record; +}) { + console.warn(`[oauth] ${error.status} ${error.code}: ${error.description}`); + + const headers = new Headers(error.headers); + headers.set("Content-Type", "application/json"); + if (headers.has(WWW_AUTHENTICATE_HEADER)) { + headers.set("Access-Control-Expose-Headers", WWW_AUTHENTICATE_HEADER); + } + + return new Response( + JSON.stringify({ + error: error.code, + error_description: error.description, + }), + { + status: error.status, + headers, + }, + ); +} + +function csrfProtected(request: Request) { + const origin = request.headers.get("Origin"); + return origin === getPublicOrigin(request); +} + +async function getAuthorizeSessionBlocker(request: Request) { + try { + await resolveHostedContext(request.headers); + return null; + } catch (error) { + const appError = asAppError(error); + if (appError?.code === "UNAUTHENTICATED") { + return redirectToSignIn(request); + } + + if (appError?.code === "AUTH_CONFIG_MISSING") { + return new Response("Missing Better Auth hosted configuration", { + status: 500, + }); + } + + throw error; + } +} + +async function resolveContextForConsent(request: Request) { + try { + return await resolveHostedContext(request.headers); + } catch (error) { + const appError = asAppError(error); + if (appError?.code === "UNAUTHENTICATED") { + return null; + } + throw error; + } +} + +function buildConsentUrl(request: Request) { + const sourceUrl = new URL(request.url); + const consentUrl = new URL("/oauth-consent", request.url); + + for (const key of OAUTH_AUTHORIZATION_PARAM_NAMES) { + for (const value of sourceUrl.searchParams.getAll(key)) { + consentUrl.searchParams.append(key, value); + } + } + + return consentUrl; +} + +function buildAuthorizeRequestFromConsentQuery( + request: Request, + query: string, +) { + const authorizeUrl = new URL(OAUTH_AUTHORIZE_PATH, request.url); + const params = new URLSearchParams(query); + + for (const key of OAUTH_AUTHORIZATION_PARAM_NAMES) { + for (const value of params.getAll(key)) { + authorizeUrl.searchParams.append(key, value); + } + } + + return new Request(authorizeUrl.toString(), { + headers: request.headers, + }); +} + +function withDefaultMcpResource(authRequest: AuthRequest, request: Request) { + const mcpResource = getMcpResourceForRequest(request); + if (!authRequest.resource) { + return { + ...authRequest, + resource: mcpResource, + }; + } + + const requestedResources = Array.isArray(authRequest.resource) + ? authRequest.resource + : [authRequest.resource]; + + if (requestedResources.some((resource) => resource !== mcpResource)) { + throw new Error(`OAuth resource must be ${mcpResource}`); + } + + return { + ...authRequest, + resource: mcpResource, + }; +} + +function getGrantedMcpScopes(requestedScopes: string[]) { + if (requestedScopes.length === 0) { + return [...MCP_OAUTH_SCOPES]; + } + + const requested = new Set(requestedScopes); + const granted = MCP_OAUTH_SCOPES.filter((scope) => requested.has(scope)); + + if (!granted.includes(MCP_SCOPE)) { + throw new Error("The mcp scope is required"); + } + + return granted; +} + +function deniedRedirect(authRequest: AuthRequest) { + const redirectUrl = new URL(authRequest.redirectUri); + redirectUrl.searchParams.set("error", "access_denied"); + redirectUrl.searchParams.set("error_description", "The user denied access"); + if (authRequest.state) { + redirectUrl.searchParams.set("state", authRequest.state); + } + + return redirectUrl.toString(); +} + +async function handleOAuthAuthorizeRequest( + request: Request, + env: OpenSeoOAuthEnv, +) { + const oauth = getOAuthHelpers(env); + + try { + await oauth.parseAuthRequest(request); + } catch (error) { + return invalidOAuthRequestResponse(error); + } + + const sessionBlocker = await getAuthorizeSessionBlocker(request); + if (sessionBlocker) return sessionBlocker; + + return Response.redirect(buildConsentUrl(request).toString(), 302); +} + +async function handleOAuthConsentResponse( + request: Request, + env: OpenSeoOAuthEnv, +) { + if (request.method !== "POST") { + return new Response("Method not allowed", { status: 405 }); + } + + if (!csrfProtected(request)) { + return jsonResponse({ error: "Invalid request origin" }, { status: 403 }); + } + + let rawBody: unknown; + try { + rawBody = await request.json(); + } catch { + return jsonResponse({ error: "Invalid consent response" }, { status: 400 }); + } + + const body = consentResponseSchema.safeParse(rawBody); + if (!body.success) { + return jsonResponse({ error: "Invalid consent response" }, { status: 400 }); + } + + const oauth = getOAuthHelpers(env); + const authorizeRequest = buildAuthorizeRequestFromConsentQuery( + request, + body.data.query, + ); + + let authRequest: AuthRequest; + try { + authRequest = await oauth.parseAuthRequest(authorizeRequest); + authRequest = withDefaultMcpResource(authRequest, request); + } catch (error) { + return jsonResponse( + { + error: error instanceof Error ? error.message : "Invalid OAuth request", + }, + { status: 400 }, + ); + } + + if (!body.data.accept) { + return jsonResponse({ redirectTo: deniedRedirect(authRequest) }); + } + + const context = await resolveContextForConsent(request); + if (!context) { + return jsonResponse({ error: "Sign in required" }, { status: 401 }); + } + + let scopes: string[]; + try { + scopes = getGrantedMcpScopes(authRequest.scope); + } catch (error) { + return jsonResponse( + { + error: error instanceof Error ? error.message : "Invalid OAuth scopes", + }, + { status: 400 }, + ); + } + + const audience = getMcpResourceForRequest(request); + const props = createWorkersOAuthMcpProps({ + userId: context.userId, + userEmail: context.userEmail, + organizationId: context.organizationId, + clientId: authRequest.clientId, + scopes, + audience, + subject: context.userId, + baseUrl: getHostedBaseUrl(), + }); + + const { redirectTo } = await oauth.completeAuthorization({ + request: authRequest, + userId: context.userId, + metadata: { + clientId: authRequest.clientId, + organizationId: context.organizationId, + }, + scope: scopes, + props, + }); + + return jsonResponse({ redirectTo }); +} + +function createDefaultHandler( + appFetch: AppFetch, +): ExportedHandlerWithFetch { + return { + async fetch(request, env) { + const url = new URL(request.url); + + if (url.pathname === OAUTH_AUTHORIZE_PATH) { + return handleOAuthAuthorizeRequest(request, env); + } + + if (url.pathname === OAUTH_CONSENT_RESPONSE_PATH) { + return handleOAuthConsentResponse(request, env); + } + + return appFetch(request); + }, + }; +} + +const mcpApiHandler: ExportedHandlerWithFetch = { + async fetch(request, env, ctx) { + return handleAuthenticatedOpenSeoMcpRequest( + request, + (ctx as OAuthExecutionContext).props, + env, + ctx, + ); + }, +}; + +export function createOpenSeoOAuthProvider(appFetch: AppFetch) { + const options: OAuthProviderOptions = { + apiRoute: MCP_ROUTE, + apiHandler: mcpApiHandler, + defaultHandler: createDefaultHandler(appFetch), + authorizeEndpoint: OAUTH_AUTHORIZE_PATH, + tokenEndpoint: OAUTH_TOKEN_PATH, + clientRegistrationEndpoint: OAUTH_REGISTER_PATH, + scopesSupported: [...MCP_OAUTH_SCOPES], + refreshTokenTTL: REFRESH_TOKEN_TTL_SECONDS, + resourceMetadata: { + scopes_supported: [...MCP_OAUTH_SCOPES], + resource_name: "OpenSEO MCP", + }, + tokenExchangeCallback: ({ props, requestedScope }) => { + const accessTokenProps = withWorkersOAuthMcpScopes(props, requestedScope); + + return accessTokenProps ? { accessTokenProps } : undefined; + }, + onError: oauthErrorResponse, + }; + + return new OAuthProvider(options); +} diff --git a/src/server/mcp/project-auth.test.ts b/src/server/mcp/project-auth.test.ts index d432e9b..24cc25a 100644 --- a/src/server/mcp/project-auth.test.ts +++ b/src/server/mcp/project-auth.test.ts @@ -1,15 +1,12 @@ +import type { AuthInfo } from "@modelcontextprotocol/sdk/server/auth/types.js"; +import type { ToolExtra } from "@/server/mcp/context"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { MCP_AUTH_CONTEXT_PROP } from "@/server/mcp/context"; const mocks = vi.hoisted(() => ({ - getMcpAuthContext: vi.fn(), getProjectForOrganization: vi.fn(), })); -vi.mock("agents/mcp", () => ({ - getMcpAuthContext: mocks.getMcpAuthContext, -})); - vi.mock("@/server/features/projects/services/ProjectService", () => ({ ProjectService: { getProjectForOrganization: mocks.getProjectForOrganization, @@ -27,14 +24,24 @@ const authContext = { baseUrl: "https://open-seo.test", }; +const toolExtra: ToolExtra = { + signal: new AbortController().signal, + requestId: 1, + sendNotification: vi.fn(), + sendRequest: vi.fn(), + authInfo: { + token: "token", + clientId: "client_123", + scopes: ["mcp"], + resource: new URL("https://open-seo.test/mcp"), + extra: { [MCP_AUTH_CONTEXT_PROP]: authContext }, + } satisfies AuthInfo, +}; + describe("withMcpProjectAuth", () => { beforeEach(() => { vi.resetModules(); - mocks.getMcpAuthContext.mockReset(); mocks.getProjectForOrganization.mockReset(); - mocks.getMcpAuthContext.mockReturnValue({ - props: { [MCP_AUTH_CONTEXT_PROP]: authContext }, - }); }); it("checks project access for the authenticated organization", async () => { @@ -43,7 +50,7 @@ describe("withMcpProjectAuth", () => { const wrapped = withMcpProjectAuth(handler); await expect( - wrapped({ projectId: "project_123" }, undefined), + wrapped({ projectId: "project_123" }, toolExtra), ).resolves.toBe("ok"); expect(mocks.getProjectForOrganization).toHaveBeenCalledWith( @@ -57,7 +64,7 @@ describe("withMcpProjectAuth", () => { const handler = vi.fn().mockReturnValue("ok"); const wrapped = withMcpProjectAuth(handler); - await wrapped({ projectId: "project_123" }, undefined); + await wrapped({ projectId: "project_123" }, toolExtra); expect(handler).toHaveBeenCalledWith( { projectId: "project_123" }, @@ -89,7 +96,7 @@ describe("withMcpProjectAuth", () => { const handler = vi.fn(); const wrapped = withMcpProjectAuth(handler); - await expect(wrapped({ projectId: "project_123" }, undefined)).rejects.toBe( + await expect(wrapped({ projectId: "project_123" }, toolExtra)).rejects.toBe( error, ); diff --git a/src/server/mcp/project-auth.ts b/src/server/mcp/project-auth.ts index a224572..11a488a 100644 --- a/src/server/mcp/project-auth.ts +++ b/src/server/mcp/project-auth.ts @@ -10,7 +10,7 @@ type ProjectScopedArgs = { }; async function requireProjectAccess(_extra: ToolExtra, projectId: string) { - const { baseUrl, ...auth } = requireMcpToolAuthContext(); + const { baseUrl, ...auth } = requireMcpToolAuthContext(_extra); // This lookup enforces that the project belongs to the authenticated org. await ProjectService.getProjectForOrganization( diff --git a/src/server/mcp/protected-resource-metadata.ts b/src/server/mcp/protected-resource-metadata.ts deleted file mode 100644 index 3d0496e..0000000 --- a/src/server/mcp/protected-resource-metadata.ts +++ /dev/null @@ -1,44 +0,0 @@ -import { getHostedBaseUrl, hasHostedAuthConfig } from "@/lib/auth"; -import { isHostedAuthMode } from "@/lib/auth-mode"; -import { getOAuthProviderResourceActions } from "@/lib/oauth-provider-resource-client"; -import { getMcpResource, MCP_OAUTH_SCOPES } from "@/lib/oauth-resource"; - -function unavailableMcpProtectedResourceMetadataResponse( - authMode: string | null | undefined, -) { - if (!isHostedAuthMode(authMode)) { - return new Response("Not found", { status: 404 }); - } - - return new Response("Missing Better Auth hosted configuration", { - status: 500, - }); -} - -export async function mcpProtectedResourceMetadataResponse( - authMode: string | null | undefined, -) { - if (!isHostedAuthMode(authMode) || !hasHostedAuthConfig()) { - return unavailableMcpProtectedResourceMetadataResponse(authMode); - } - - const baseUrl = getHostedBaseUrl(); - const resource = getMcpResource(baseUrl); - const authorizationServer = `${baseUrl}/api/auth`; - - const metadata = - await getOAuthProviderResourceActions().getProtectedResourceMetadata({ - resource, - authorization_servers: [authorizationServer], - scopes_supported: [...MCP_OAUTH_SCOPES], - resource_name: "OpenSEO MCP", - }); - - return new Response(JSON.stringify(metadata), { - headers: { - "Cache-Control": - "public, max-age=15, stale-while-revalidate=15, stale-if-error=86400", - "Content-Type": "application/json", - }, - }); -} diff --git a/src/server/mcp/public-origin.test.ts b/src/server/mcp/public-origin.test.ts new file mode 100644 index 0000000..4af3202 --- /dev/null +++ b/src/server/mcp/public-origin.test.ts @@ -0,0 +1,58 @@ +import { describe, expect, it } from "vitest"; +import { getPublicOrigin, requestWithPublicOrigin } from "./public-origin"; + +describe("getPublicOrigin", () => { + it("uses the forwarded public protocol and host for tunneled requests", () => { + const request = new Request("http://localhost:3102/api/oauth/consent", { + headers: { + "x-forwarded-proto": "https", + "x-forwarded-host": "scenario-tools.trycloudflare.com", + }, + }); + + expect(getPublicOrigin(request)).toBe( + "https://scenario-tools.trycloudflare.com", + ); + }); + + it("falls back to the request origin without proxy headers", () => { + const request = new Request("http://localhost:3102/api/oauth/consent"); + + expect(getPublicOrigin(request)).toBe("http://localhost:3102"); + }); + + it("ignores forwarded hosts when the request is already public https", () => { + const request = new Request("https://app.openseo.so/api/oauth/consent", { + headers: { + "x-forwarded-proto": "https", + "x-forwarded-host": "evil.test", + }, + }); + + expect(getPublicOrigin(request)).toBe("https://app.openseo.so"); + }); +}); + +describe("requestWithPublicOrigin", () => { + it("rewrites the request URL origin while preserving path and query", async () => { + const request = new Request("http://localhost:3102/mcp?x=1", { + method: "POST", + headers: { + "content-type": "application/json", + "x-forwarded-proto": "https", + "x-forwarded-host": "scenario-tools.trycloudflare.com", + }, + body: JSON.stringify({ method: "initialize" }), + }); + + const publicRequest = requestWithPublicOrigin(request); + + expect(publicRequest.url).toBe( + "https://scenario-tools.trycloudflare.com/mcp?x=1", + ); + expect(publicRequest.method).toBe("POST"); + await expect(publicRequest.json()).resolves.toEqual({ + method: "initialize", + }); + }); +}); diff --git a/src/server/mcp/public-origin.ts b/src/server/mcp/public-origin.ts new file mode 100644 index 0000000..c4b21e6 --- /dev/null +++ b/src/server/mcp/public-origin.ts @@ -0,0 +1,43 @@ +function firstHeaderValue(value: string | null) { + return value?.split(",")[0]?.trim() || null; +} + +function getForwardedProtocol(request: Request) { + const protocol = firstHeaderValue(request.headers.get("x-forwarded-proto")); + return protocol === "http" || protocol === "https" ? protocol : null; +} + +export function getPublicOrigin(request: Request) { + const url = new URL(request.url); + if (url.protocol === "https:") { + return url.origin; + } + + const protocol = getForwardedProtocol(request); + const host = firstHeaderValue(request.headers.get("x-forwarded-host")); + + if (!protocol || !host) { + return url.origin; + } + + try { + return new URL(`${protocol}://${host}`).origin; + } catch { + return url.origin; + } +} + +export function requestWithPublicOrigin(request: Request) { + const url = new URL(request.url); + const publicOrigin = getPublicOrigin(request); + + if (url.origin === publicOrigin) { + return request; + } + + const publicUrl = new URL( + `${url.pathname}${url.search}${url.hash}`, + publicOrigin, + ); + return new Request(publicUrl, request); +} diff --git a/src/server/mcp/server.ts b/src/server/mcp/server.ts index ba24bc0..6677f72 100644 --- a/src/server/mcp/server.ts +++ b/src/server/mcp/server.ts @@ -1,4 +1,4 @@ -import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"; +import type { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"; import { getBacklinksOverviewTool } from "@/server/mcp/tools/get-backlinks-overview"; import { getDomainKeywordSuggestionsTool } from "@/server/mcp/tools/get-domain-keyword-suggestions"; import { getDomainOverviewTool } from "@/server/mcp/tools/get-domain-overview"; @@ -10,12 +10,7 @@ import { researchKeywordsTool } from "@/server/mcp/tools/research-keywords"; import { saveKeywordsTool } from "@/server/mcp/tools/save-keywords"; import { whoamiTool } from "@/server/mcp/tools/whoami"; -export function createOpenSeoMcpServer() { - const server = new McpServer({ - name: "OpenSEO MCP", - version: "0.0.10", - }); - +export function registerOpenSeoMcpTools(server: McpServer) { server.registerTool(whoamiTool.name, whoamiTool.config, whoamiTool.handler); server.registerTool( listProjectsTool.name, @@ -62,6 +57,4 @@ export function createOpenSeoMcpServer() { getRankTrackerTool.config, getRankTrackerTool.handler, ); - - return server; } diff --git a/src/server/mcp/transport.ts b/src/server/mcp/transport.ts new file mode 100644 index 0000000..88cfccc --- /dev/null +++ b/src/server/mcp/transport.ts @@ -0,0 +1,52 @@ +import { createMcpHandler } from "agents/mcp"; +import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"; +import { MCP_SCOPE } from "@/lib/oauth-resource"; +import { + MCP_AUTH_CONTEXT_PROP, + MCP_ROUTE, + runWithMcpToolAuthContext, + workersOAuthMcpPropsSchema, +} from "@/server/mcp/context"; +import { registerOpenSeoMcpTools } from "@/server/mcp/server"; + +function createOpenSeoMcpServer() { + const server = new McpServer({ + name: "OpenSEO MCP", + version: "0.0.10", + }); + registerOpenSeoMcpTools(server); + + return server; +} + +export async function handleAuthenticatedOpenSeoMcpRequest( + request: Request, + props: unknown, + env: Env, + ctx: ExecutionContext, +): Promise { + const result = workersOAuthMcpPropsSchema.safeParse(props); + const scopes = result.success + ? result.data[MCP_AUTH_CONTEXT_PROP].scopes + : []; + + if (!result.success || !scopes.includes(MCP_SCOPE)) { + return new Response("MCP auth context required", { status: 403 }); + } + + const server = createOpenSeoMcpServer(); + const handler = createMcpHandler(server, { + route: MCP_ROUTE, + enableJsonResponse: true, + authContext: { props: result.data }, + corsOptions: { + headers: + "Authorization, Content-Type, Last-Event-ID, mcp-protocol-version, mcp-session-id", + exposeHeaders: "mcp-protocol-version, mcp-session-id", + }, + }); + + return runWithMcpToolAuthContext(result.data[MCP_AUTH_CONTEXT_PROP], () => + handler(request, env, ctx), + ); +} diff --git a/src/server/mcp/user-email.ts b/src/server/mcp/user-email.ts deleted file mode 100644 index bf7e977..0000000 --- a/src/server/mcp/user-email.ts +++ /dev/null @@ -1,12 +0,0 @@ -import { eq } from "drizzle-orm"; -import { db } from "@/db"; -import { user } from "@/db/schema"; - -export async function getMcpUserEmail(userId: string) { - const authUser = await db.query.user.findFirst({ - columns: { email: true }, - where: eq(user.id, userId), - }); - - return authUser?.email ?? null; -} diff --git a/src/serverFunctions/oauth.ts b/src/serverFunctions/oauth.ts deleted file mode 100644 index f425794..0000000 --- a/src/serverFunctions/oauth.ts +++ /dev/null @@ -1,33 +0,0 @@ -import { createServerFn } from "@tanstack/react-start"; -import { eq } from "drizzle-orm"; -import { z } from "zod"; -import { db } from "@/db"; -import { oauthClient } from "@/db/better-auth-schema"; -import { requireAuthenticatedContext } from "@/serverFunctions/middleware"; - -const getOAuthClientInfoSchema = z.object({ - clientId: z.string().min(1), -}); - -export const getOAuthClientInfo = createServerFn({ method: "POST" }) - .middleware(requireAuthenticatedContext) - .inputValidator((data: unknown) => getOAuthClientInfoSchema.parse(data)) - .handler(async ({ data }) => { - const row = await db - .select({ - name: oauthClient.name, - icon: oauthClient.icon, - uri: oauthClient.uri, - }) - .from(oauthClient) - .where(eq(oauthClient.clientId, data.clientId)) - .get(); - - if (!row) return null; - - return { - name: row.name ?? null, - icon: row.icon ?? null, - uri: row.uri ?? null, - }; - }); diff --git a/vite.config.ts b/vite.config.ts index 18fe23b..c332b27 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -10,12 +10,16 @@ export default defineConfig(({ mode }) => { const env = loadEnv(mode, process.cwd(), ""); const port = env.PORT ? Number(env.PORT) : 3001; const showDevtools = env.VITE_SHOW_DEVTOOLS !== "false"; - const allowedHosts = env.ALLOWED_HOST ? [env.ALLOWED_HOST] : undefined; + const allowedHosts = [ + env.ALLOWED_HOST, + env.BETTER_AUTH_URL ? new URL(env.BETTER_AUTH_URL).hostname : undefined, + ].filter((host): host is string => Boolean(host)); const emitSourcemaps = env.POSTHOG_SOURCEMAPS === "true"; return { envPrefix: ["VITE_", "AUTH_MODE", "POSTHOG_PUBLIC_KEY", "POSTHOG_HOST"], server: { + allowedHosts, port, }, preview: { diff --git a/wrangler.jsonc b/wrangler.jsonc index 19e4f2d..143b88d 100644 --- a/wrangler.jsonc +++ b/wrangler.jsonc @@ -34,6 +34,10 @@ "binding": "KV", "id": "4abd52f3f2c549ac83cc2cb4ceec8620", }, + { + "binding": "OAUTH_KV", + "id": "bd1759494309474a9b423b029967b0db", + }, ], "d1_databases": [ {