diff --git a/README.md b/README.md index de77080..f2148f0 100644 --- a/README.md +++ b/README.md @@ -24,7 +24,7 @@ The Kernel MCP Server bridges AI assistants (like Claude, Cursor, or other MCP-c **Open-source & fully-managed** — the complete codebase is available here, and we run the production instance so you don't need to deploy anything. -The server uses OAuth 2.0 authentication via [Clerk](https://clerk.com) to ensure secure access to your Kernel resources. +The server uses OAuth 2.0 authentication via [Clerk](https://clerk.com) to ensure secure access to your Kernel resources. During authorization, users can grant organization-wide access or restrict the resulting access and refresh tokens to one Kernel project. Project-scoped tokens cannot switch projects; organization-wide authorization remains available for existing workflows. For a deeper dive into why and how we built this server, see our blog post: [Introducing Kernel MCP Server](https://blog.onkernel.com/p/introducing-kernel-mcp-server). diff --git a/src/app/authorize/route.test.ts b/src/app/authorize/route.test.ts new file mode 100644 index 0000000..6208570 --- /dev/null +++ b/src/app/authorize/route.test.ts @@ -0,0 +1,193 @@ +import { describe, expect, test } from "bun:test"; +import { NextRequest } from "next/server"; +import type { AuthorizeDependencies } from "./route"; + +process.env.KERNEL_CLI_PROD_CLIENT_ID ??= "cli_prod"; +process.env.KERNEL_CLI_STAGING_CLIENT_ID ??= "cli_staging"; +process.env.KERNEL_CLI_DEV_CLIENT_ID ??= "cli_dev"; +process.env.NEXT_PUBLIC_CLERK_DOMAIN ??= "clerk.example.test"; + +const { authorizeRequest } = await import("./route"); + +function request(query: string) { + return new NextRequest(`https://auth.example.test/authorize?${query}`); +} + +function dependencies({ + userId = "user_1", + orgId = "org_1", +}: { userId?: string | null; orgId?: string | null } = {}) { + const requestContexts: Parameters< + AuthorizeDependencies["setRequestContext"] + >[0][] = []; + const clientContexts: Parameters< + AuthorizeDependencies["setClientContext"] + >[0][] = []; + const projects: Parameters[0][] = []; + return { + requestContexts, + clientContexts, + projects, + value: { + getAuth: async () => ({ + userId, + orgId, + getToken: async () => "session-token", + }), + setRequestContext: async (value) => { + requestContexts.push(value); + }, + setClientContext: async (value) => { + clientContexts.push(value); + }, + requireProject: async (value) => { + projects.push(value); + return { id: value.projectId, name: "project", status: "active" }; + }, + } satisfies AuthorizeDependencies, + }; +} + +describe("GET /authorize", () => { + test("redirects to selection and preserves OAuth parameters", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request( + "client_id=client_1&state=opaque&resource=https%3A%2F%2Fmcp.example.test%2Fmcp&code_challenge=challenge&code_challenge_method=S256", + ), + deps.value, + ); + + expect(response.status).toBe(307); + const location = new URL(response.headers.get("location")!); + expect(location.pathname).toBe("/select-org"); + expect(location.searchParams.get("state")).toBe("opaque"); + expect(location.searchParams.get("resource")).toBe( + "https://mcp.example.test/mcp", + ); + }); + + test("stores PKCE-bound organization scope and strips internal parameters", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request( + "client_id=client_1&org_id=org_1&access_scope=organization&state=opaque&code_challenge=challenge&code_challenge_method=S256", + ), + deps.value, + ); + + expect(response.status).toBe(307); + expect(deps.requestContexts).toHaveLength(1); + expect(deps.requestContexts[0]).toMatchObject({ + clientId: "client_1", + codeChallenge: "challenge", + authorizationContext: { + version: 1, + clerk_user_id: "user_1", + clerk_org_id: "org_1", + access_scope: "organization", + }, + }); + const location = new URL(response.headers.get("location")!); + expect(location.host).toBe("clerk.example.test"); + expect(location.searchParams.get("org_id")).toBeNull(); + expect(location.searchParams.get("access_scope")).toBeNull(); + expect(location.searchParams.get("project_id")).toBeNull(); + expect(location.searchParams.get("state")).toBe("opaque"); + }); + + test("validates and stores project scope", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request( + "client_id=client_1&org_id=org_1&access_scope=project&project_id=proj_1&code_challenge=challenge&code_challenge_method=S256", + ), + deps.value, + ); + + expect(response.status).toBe(307); + expect(deps.projects).toEqual([ + { clerkSessionToken: "session-token", projectId: "proj_1" }, + ]); + expect(deps.requestContexts[0].authorizationContext).toMatchObject({ + access_scope: "project", + project_id: "proj_1", + }); + }); + + test("rejects project scope without S256 PKCE", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request( + "client_id=client_1&org_id=org_1&access_scope=project&project_id=proj_1", + ), + deps.value, + ); + + expect(response.status).toBe(400); + expect(await response.json()).toMatchObject({ error: "invalid_request" }); + expect(deps.requestContexts).toHaveLength(0); + }); + + test("rejects incomplete or unsupported PKCE parameters", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request( + "client_id=client_1&org_id=org_1&access_scope=organization&code_challenge=challenge&code_challenge_method=plain", + ), + deps.value, + ); + + expect(response.status).toBe(400); + expect(await response.json()).toMatchObject({ error: "invalid_request" }); + expect(deps.clientContexts).toHaveLength(0); + expect(deps.requestContexts).toHaveLength(0); + }); + + test("requires PKCE for shared clients", async () => { + const deps = dependencies(); + const response = await authorizeRequest( + request("client_id=cli_prod&org_id=org_1&access_scope=organization"), + deps.value, + ); + + expect(response.status).toBe(400); + expect(await response.json()).toMatchObject({ error: "invalid_request" }); + expect(deps.clientContexts).toHaveLength(0); + }); + + test("rejects an organization that is not active for the user", async () => { + const deps = dependencies({ orgId: "org_2" }); + const response = await authorizeRequest( + request("client_id=client_1&org_id=org_1&access_scope=organization"), + deps.value, + ); + + expect(response.status).toBe(403); + expect(await response.json()).toMatchObject({ error: "access_denied" }); + }); + + test("keeps shared-client state compatible and includes scope", async () => { + const deps = dependencies(); + const originalState = Buffer.from( + JSON.stringify({ csrf: "csrf_1" }), + ).toString("base64"); + const response = await authorizeRequest( + request( + `client_id=cli_prod&org_id=org_1&access_scope=project&project_id=proj_1&state=${encodeURIComponent(originalState)}&code_challenge=challenge&code_challenge_method=S256`, + ), + deps.value, + ); + + const location = new URL(response.headers.get("location")!); + const state = JSON.parse( + Buffer.from(location.searchParams.get("state")!, "base64").toString(), + ); + expect(state).toEqual({ + csrf: "csrf_1", + org_id: "org_1", + access_scope: "project", + project_id: "proj_1", + }); + }); +}); diff --git a/src/app/authorize/route.ts b/src/app/authorize/route.ts index fc1738c..2341207 100644 --- a/src/app/authorize/route.ts +++ b/src/app/authorize/route.ts @@ -1,196 +1,260 @@ +import { auth } from "@clerk/nextjs/server"; import { NextRequest, NextResponse } from "next/server"; -import { setOrgIdForClientId } from "../../lib/redis"; -import { SHARED_CLIENT_IDS } from "../../lib/const"; +import { + setAuthorizationContextForClientId, + setAuthorizationContextForRequest, +} from "@/lib/redis"; +import { SHARED_CLIENT_IDS } from "@/lib/const"; +import { + authorizationContextFromSelection, + type OAuthAuthorizationContext, + PROJECT_ACCESS_SCOPE, +} from "@/lib/oauth-context"; +import { + OAuthProjectsError, + requireActiveOAuthProject, +} from "@/lib/oauth-projects"; + +const CORS_HEADERS = { + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Methods": "GET, POST, OPTIONS", + "Access-Control-Allow-Headers": "Content-Type, Authorization", +}; + +const INTERNAL_AUTHORIZATION_PARAMS = new Set([ + "org_id", + "access_scope", + "project_id", +]); export async function OPTIONS(): Promise { - return new NextResponse(null, { - status: 204, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "GET, POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }); + return new NextResponse(null, { status: 204, headers: CORS_HEADERS }); } -export async function GET(request: NextRequest): Promise { - const searchParams = request.nextUrl.searchParams; +function errorResponse( + error: string, + errorDescription: string, + status = 400, +): NextResponse { + return NextResponse.json( + { error, error_description: errorDescription }, + { status, headers: CORS_HEADERS }, + ); +} + +function sharedClientState({ + originalState, + orgId, + accessScope, + projectId, +}: { + originalState: string | null; + orgId: string; + accessScope: string; + projectId?: string; +}): string { + let csrf = originalState || ""; + if (originalState) { + try { + const parsed = JSON.parse( + Buffer.from(originalState, "base64").toString(), + ) as { csrf?: string }; + if (parsed.csrf) csrf = parsed.csrf; + } catch { + // Older clients may send a plain CSRF value. + } + } + + return Buffer.from( + JSON.stringify({ + csrf, + org_id: orgId, + access_scope: accessScope, + ...(projectId ? { project_id: projectId } : {}), + }), + ).toString("base64"); +} + +export interface AuthorizeDependencies { + getAuth: () => Promise<{ + userId: string | null | undefined; + orgId: string | null | undefined; + getToken: () => Promise; + }>; + setRequestContext: (input: { + clientId: string; + codeChallenge: string; + authorizationContext: OAuthAuthorizationContext; + ttlSeconds: number; + }) => Promise; + setClientContext: (input: { + clientId: string; + authorizationContext: OAuthAuthorizationContext; + ttlSeconds: number; + }) => Promise; + requireProject: typeof requireActiveOAuthProject; +} + +const authorizeDependencies: AuthorizeDependencies = { + getAuth: async () => auth(), + setRequestContext: setAuthorizationContextForRequest, + setClientContext: setAuthorizationContextForClientId, + requireProject: requireActiveOAuthProject, +}; - // Step 1: Extract and validate required OAuth parameters +export async function authorizeRequest( + request: NextRequest, + dependencies: AuthorizeDependencies = authorizeDependencies, +): Promise { + const searchParams = request.nextUrl.searchParams; const clientId = searchParams.get("client_id"); const selectedOrgId = searchParams.get("org_id"); const originalState = searchParams.get("state"); + const accessScope = searchParams.get("access_scope"); + const projectId = searchParams.get("project_id"); + const codeChallenge = searchParams.get("code_challenge"); + const codeChallengeMethod = searchParams.get("code_challenge_method"); - console.debug("[authorize] start", { - hasClientId: Boolean(clientId), - hasSelectedOrgId: Boolean(selectedOrgId), - hasState: Boolean(originalState), - }); - - // Step 2: Validate minimum required parameters if (!clientId) { - return NextResponse.json( - { - error: "invalid_request", - error_description: "Missing required parameter: client_id", - }, - { - status: 400, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "GET, POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }, + return errorResponse( + "invalid_request", + "Missing required parameter: client_id", ); } - // Step 3: Redirect to organization selector if no org chosen yet if (!selectedOrgId) { - console.debug( - "[authorize] no org selected yet, redirecting to /select-org", - ); const selectOrgUrl = new URL("/select-org", request.nextUrl.origin); - - // Pass all OAuth parameters to the org selector searchParams.forEach((value, key) => { selectOrgUrl.searchParams.set(key, value); }); - - return NextResponse.redirect(selectOrgUrl.toString()); + return NextResponse.redirect(selectOrgUrl); } - // Step 4: Validate server configuration const clerkDomain = process.env.NEXT_PUBLIC_CLERK_DOMAIN; - if (!clerkDomain) { - return NextResponse.json( - { - error: "server_error", - error_description: "Server configuration error", - }, - { - status: 500, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "GET, POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }, + return errorResponse("server_error", "Server configuration error", 500); + } + + const { userId, orgId, getToken } = await dependencies.getAuth(); + if (!userId || orgId !== selectedOrgId) { + return errorResponse( + "access_denied", + "The selected organization is not active for this user", + 403, + ); + } + + const hasPKCEParameters = Boolean(codeChallenge || codeChallengeMethod); + if (hasPKCEParameters && (!codeChallenge || codeChallengeMethod !== "S256")) { + return errorResponse( + "invalid_request", + "PKCE requires code_challenge and code_challenge_method=S256", + ); + } + if (SHARED_CLIENT_IDS.includes(clientId) && !codeChallenge) { + return errorResponse( + "invalid_request", + "Shared OAuth clients require PKCE with S256", ); } - // Step 5: Store organization context for ephemeral clients - // Skip Redis storage for shared clients to avoid cross-user overwrites - if (!SHARED_CLIENT_IDS.includes(clientId)) { + let authorizationContext; + try { + authorizationContext = authorizationContextFromSelection({ + clerkUserId: userId, + clerkOrgId: selectedOrgId, + accessScope, + projectId, + }); + } catch (error) { + return errorResponse( + "invalid_request", + error instanceof Error ? error.message : "Invalid access scope", + ); + } + + if (authorizationContext.access_scope === PROJECT_ACCESS_SCOPE) { + if (!codeChallenge) { + return errorResponse( + "invalid_request", + "Project-scoped authorization requires PKCE with S256", + ); + } + const clerkSessionToken = await getToken(); + if (!clerkSessionToken) { + return errorResponse("access_denied", "Authentication required", 401); + } try { - // TTL only needs to last through the OAuth flow (authorization_code exchange) - await setOrgIdForClientId({ - clientId, - orgId: selectedOrgId, - ttlSeconds: 60 * 60, // 1 hour - }); - console.debug("[authorize] stored org_id for ephemeral client", { - clientIdMasked: clientId?.slice(0, 4) + "...", + await dependencies.requireProject({ + clerkSessionToken, + projectId: authorizationContext.project_id, }); } catch (error) { - console.error("[authorize] failed to store org_id in Redis", { error }); - return NextResponse.json( - { - error: "server_error", - error_description: "Failed to store organization context", - }, - { - status: 500, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "GET, POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }, + console.warn("[authorize] project validation failed", { error }); + if (error instanceof OAuthProjectsError && error.status >= 500) { + return errorResponse( + "server_error", + "Project validation is temporarily unavailable", + 503, + ); + } + return errorResponse( + "access_denied", + "Project not found or inactive", + 403, ); } - } else { - console.debug("[authorize] shared client, skipping Redis store", { - clientIdMasked: clientId?.slice(0, 4) + "...", - }); } - // Step 6: Handle state parameter based on client type - let modifiedState = originalState; - - // Only modify state for shared clients (CLI), not ephemeral clients (MCP) - if (selectedOrgId && SHARED_CLIENT_IDS.includes(clientId)) { - try { - // Extract CSRF token from original state if it's base64-encoded JSON - let csrfToken = originalState || ""; - - if (originalState) { - try { - // Try to decode originalState as base64-encoded JSON (from CLI) - const decodedState = Buffer.from(originalState, "base64").toString(); - const parsedState = JSON.parse(decodedState); - if (parsedState.csrf) { - csrfToken = parsedState.csrf; - console.debug("[authorize] extracted CSRF from CLI state param"); - } - } catch (decodeError) { - // If decoding fails, treat originalState as plain CSRF token - csrfToken = originalState; - console.debug("[authorize] using original state as plain CSRF token"); - } - } - - const stateData = { - csrf: csrfToken, - org_id: selectedOrgId, - }; - modifiedState = Buffer.from(JSON.stringify(stateData)).toString("base64"); - console.debug("[authorize] encoded org_id into state for shared client", { - clientIdMasked: clientId?.slice(0, 4) + "...", + try { + if (codeChallenge) { + await dependencies.setRequestContext({ + clientId, + codeChallenge, + authorizationContext, + ttlSeconds: 60 * 60, }); - } catch (error) { - console.error("[authorize] failed to encode org_id into state", { - error, + } else { + // Compatibility for existing non-PKCE organization-wide clients. New + // project-scoped grants always require the PKCE-bound request mapping. + await dependencies.setClientContext({ + clientId, + authorizationContext, + ttlSeconds: 60 * 60, }); - return NextResponse.json( - { - error: "server_error", - error_description: "Failed to encode organization context", - }, - { - status: 500, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "GET, POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }, - ); } - } else if (selectedOrgId) { - // For ephemeral clients, don't modify state - rely on Redis storage - console.debug("[authorize] ephemeral client: preserving original state"); + } catch (error) { + console.error("[authorize] failed to store authorization context", { + error, + }); + return errorResponse( + "server_error", + "Failed to store authorization context", + 500, + ); } - // Step 7: Build Clerk authorization URL with OAuth parameters - const clerkAuthUrl = new URL(`https://${clerkDomain}/oauth/authorize`); + let state = originalState; + if (SHARED_CLIENT_IDS.includes(clientId)) { + state = sharedClientState({ + originalState, + orgId: selectedOrgId, + accessScope: authorizationContext.access_scope, + projectId: authorizationContext.project_id, + }); + } - // Pass through all original parameters except our custom org_id + const clerkAuthUrl = new URL(`https://${clerkDomain}/oauth/authorize`); searchParams.forEach((value, key) => { - if (key !== "org_id") { + if (!INTERNAL_AUTHORIZATION_PARAMS.has(key)) { clerkAuthUrl.searchParams.set(key, value); } }); + if (state) clerkAuthUrl.searchParams.set("state", state); - // Use the modified state parameter that includes org_id - if (modifiedState) { - clerkAuthUrl.searchParams.set("state", modifiedState); - } + return NextResponse.redirect(clerkAuthUrl); +} - // Step 8: Redirect to Clerk for actual OAuth authentication - console.debug("[authorize] redirecting to clerk /oauth/authorize", { - hasModifiedState: Boolean(modifiedState), - }); - return NextResponse.redirect(clerkAuthUrl.toString()); +export async function GET(request: NextRequest): Promise { + return authorizeRequest(request); } diff --git a/src/app/oauth/projects/route.test.ts b/src/app/oauth/projects/route.test.ts new file mode 100644 index 0000000..fca1459 --- /dev/null +++ b/src/app/oauth/projects/route.test.ts @@ -0,0 +1,78 @@ +import { describe, expect, test } from "bun:test"; +import { NextRequest } from "next/server"; +import { + oauthProjectsRequest, + type OAuthProjectsRouteDependencies, +} from "./route"; + +function request(query: string) { + return new NextRequest(`https://mcp.example.test/oauth/projects?${query}`); +} + +function dependencies() { + const listCalls: Parameters< + OAuthProjectsRouteDependencies["listProjects"] + >[0][] = []; + return { + listCalls, + value: { + getAuth: async () => ({ + userId: "user_1", + orgId: "org_1", + getToken: async () => "session-token", + }), + listProjects: async (input) => { + listCalls.push(input); + return { + projects: [{ id: "proj_1", name: "production", status: "active" }], + hasMore: true, + nextOffset: 40, + }; + }, + } satisfies OAuthProjectsRouteDependencies, + }; +} + +describe("GET /oauth/projects", () => { + test("forwards server-side search and pagination", async () => { + const deps = dependencies(); + const response = await oauthProjectsRequest( + request("org_id=org_1&query=prod&limit=20&offset=20"), + deps.value, + ); + + expect(response.status).toBe(200); + expect(deps.listCalls).toEqual([ + { + clerkSessionToken: "session-token", + query: "prod", + limit: 20, + offset: 20, + }, + ]); + expect(await response.json()).toEqual({ + projects: [{ id: "proj_1", name: "production", status: "active" }], + has_more: true, + next_offset: 40, + }); + }); + + test("rejects organization and pagination mismatches", async () => { + const deps = dependencies(); + const wrongOrg = await oauthProjectsRequest( + request("org_id=org_2"), + deps.value, + ); + expect(wrongOrg.status).toBe(403); + + for (const query of [ + "org_id=org_1&limit=21", + "org_id=org_1&limit=0", + "org_id=org_1&offset=-1", + ]) { + const response = await oauthProjectsRequest(request(query), deps.value); + expect(response.status).toBe(400); + } + expect(deps.listCalls).toHaveLength(0); + }); +}); diff --git a/src/app/oauth/projects/route.ts b/src/app/oauth/projects/route.ts new file mode 100644 index 0000000..124e200 --- /dev/null +++ b/src/app/oauth/projects/route.ts @@ -0,0 +1,87 @@ +import { auth } from "@clerk/nextjs/server"; +import { NextRequest, NextResponse } from "next/server"; +import { + listOAuthProjectsPage, + OAuthProjectsError, +} from "@/lib/oauth-projects"; + +export interface OAuthProjectsRouteDependencies { + getAuth: () => Promise<{ + userId: string | null | undefined; + orgId: string | null | undefined; + getToken: () => Promise; + }>; + listProjects: typeof listOAuthProjectsPage; +} + +const routeDependencies: OAuthProjectsRouteDependencies = { + getAuth: async () => auth(), + listProjects: listOAuthProjectsPage, +}; + +function integerParameter( + value: string | null, + fallback: number, +): number | null { + if (value === null) return fallback; + const parsed = Number(value); + return Number.isInteger(parsed) && parsed >= 0 ? parsed : null; +} + +export async function oauthProjectsRequest( + request: NextRequest, + dependencies: OAuthProjectsRouteDependencies = routeDependencies, +): Promise { + const { userId, orgId, getToken } = await dependencies.getAuth(); + if (!userId || !orgId) { + return NextResponse.json({ error: "unauthorized" }, { status: 401 }); + } + + const requestedOrgId = request.nextUrl.searchParams.get("org_id"); + if (!requestedOrgId || requestedOrgId !== orgId) { + return NextResponse.json( + { error: "organization_mismatch" }, + { status: 403 }, + ); + } + + const limit = integerParameter(request.nextUrl.searchParams.get("limit"), 20); + const offset = integerParameter( + request.nextUrl.searchParams.get("offset"), + 0, + ); + if (limit === null || limit < 1 || limit > 20 || offset === null) { + return NextResponse.json({ error: "invalid_pagination" }, { status: 400 }); + } + const query = request.nextUrl.searchParams.get("query")?.trim() || undefined; + if (query && query.length > 255) { + return NextResponse.json({ error: "invalid_query" }, { status: 400 }); + } + + const token = await getToken(); + if (!token) { + return NextResponse.json({ error: "unauthorized" }, { status: 401 }); + } + + try { + const page = await dependencies.listProjects({ + clerkSessionToken: token, + query, + limit, + offset, + }); + return NextResponse.json({ + projects: page.projects, + has_more: page.hasMore, + next_offset: page.nextOffset, + }); + } catch (error) { + const status = error instanceof OAuthProjectsError ? error.status : 500; + console.error("[oauth/projects] failed to list projects", { error }); + return NextResponse.json({ error: "projects_unavailable" }, { status }); + } +} + +export async function GET(request: NextRequest): Promise { + return oauthProjectsRequest(request); +} diff --git a/src/app/register/route.test.ts b/src/app/register/route.test.ts new file mode 100644 index 0000000..0337cd0 --- /dev/null +++ b/src/app/register/route.test.ts @@ -0,0 +1,81 @@ +import { describe, expect, test } from "bun:test"; +import { NextRequest } from "next/server"; +import { registerRequest, type RegisterDependencies } from "./route"; + +function request(body: unknown, contentType = "application/json") { + return new NextRequest("https://auth.example.test/register", { + method: "POST", + headers: { "Content-Type": contentType }, + body: typeof body === "string" ? body : JSON.stringify(body), + }); +} + +describe("POST /register", () => { + test("registers a public client and expands localhost redirects", async () => { + const createCalls: Parameters< + RegisterDependencies["createOAuthApplication"] + >[0][] = []; + const response = await registerRequest( + request({ + client_name: "Test Client", + redirect_uris: ["http://localhost:58432/callback"], + token_endpoint_auth_method: "none", + grant_types: ["authorization_code", "refresh_token"], + response_types: ["code"], + scope: "openid", + }), + { + createOAuthApplication: async (value) => { + createCalls.push(value); + return { + id: "oauth_app_1", + clientId: "client_1", + clientSecret: null, + }; + }, + }, + ); + + expect(response.status).toBe(200); + expect(createCalls).toEqual([ + { + name: "Test Client", + redirectUris: [ + "http://localhost:58432/callback", + "http://127.0.0.1:58432/callback", + ], + scopes: "openid", + public: true, + }, + ]); + expect(await response.json()).toMatchObject({ + client_id: "client_1", + redirect_uris: ["http://localhost:58432/callback"], + token_endpoint_auth_method: "none", + grant_types: ["authorization_code", "refresh_token"], + }); + }); + + test("rejects malformed and unsafe registrations before Clerk", async () => { + let called = false; + const deps: RegisterDependencies = { + createOAuthApplication: async () => { + called = true; + return { id: "unexpected", clientId: "unexpected" }; + }, + }; + + const contentType = await registerRequest(request({}, "text/plain"), deps); + expect(contentType.status).toBe(400); + + const missingRedirect = await registerRequest(request({}), deps); + expect(missingRedirect.status).toBe(400); + + const insecureRedirect = await registerRequest( + request({ redirect_uris: ["http://example.com/callback"] }), + deps, + ); + expect(insecureRedirect.status).toBe(400); + expect(called).toBe(false); + }); +}); diff --git a/src/app/register/route.ts b/src/app/register/route.ts index 2554529..df19d4f 100644 --- a/src/app/register/route.ts +++ b/src/app/register/route.ts @@ -15,7 +15,32 @@ export async function OPTIONS(): Promise { }); } -export async function POST(request: NextRequest): Promise { +interface OAuthApplicationInput { + name: string; + redirectUris: string[]; + scopes: string; + public: boolean; +} + +export interface RegisterDependencies { + createOAuthApplication: (input: OAuthApplicationInput) => Promise<{ + id: string; + clientId: string; + clientSecret?: string | null; + }>; +} + +const registerDependencies: RegisterDependencies = { + createOAuthApplication: async (input) => { + const clerk = await clerkClient(); + return clerk.oauthApplications.create(input); + }, +}; + +export async function registerRequest( + request: NextRequest, + dependencies: RegisterDependencies = registerDependencies, +): Promise { const contentType = request.headers.get("content-type"); if (!contentType?.includes("application/json")) { return NextResponse.json( @@ -112,8 +137,7 @@ export async function POST(request: NextRequest): Promise { const expandedRedirectUris = expandLocalhostUris(redirect_uris); // Register the OAuth application with Clerk - const clerk = await clerkClient(); - const oauthApp = await clerk.oauthApplications.create({ + const oauthApp = await dependencies.createOAuthApplication({ name: client_name || "MCP Client", redirectUris: expandedRedirectUris, scopes: scope ? scope : "openid", @@ -161,3 +185,7 @@ export async function POST(request: NextRequest): Promise { ); } } + +export async function POST(request: NextRequest): Promise { + return registerRequest(request); +} diff --git a/src/app/select-org/page.tsx b/src/app/select-org/page.tsx index 4e0c924..26c9d92 100644 --- a/src/app/select-org/page.tsx +++ b/src/app/select-org/page.tsx @@ -1,126 +1,260 @@ -'use client'; +"use client"; -import { useAuth, useOrganizationList, useUser, CreateOrganization, UserButton } from '@clerk/nextjs'; -import { useSearchParams, useRouter } from 'next/navigation'; -import { useState, Suspense, useEffect, useRef } from 'react'; -import { Col } from '@/components/col' -import { Row } from '@/components/row' -import { LoadingState } from '@/components/spinner/loading-state'; -import { KernelWordmark } from '@/components/icons'; +import { + CreateOrganization, + UserButton, + useAuth, + useOrganizationList, + useUser, +} from "@clerk/nextjs"; +import { useSearchParams, useRouter } from "next/navigation"; +import { useState, Suspense, useCallback, useEffect, useRef } from "react"; +import { Col } from "@/components/col"; +import { Row } from "@/components/row"; +import { LoadingState } from "@/components/spinner/loading-state"; +import { KernelWordmark } from "@/components/icons"; + +interface OAuthProject { + id: string; + name: string; +} + +type SelectionStage = "organization" | "scope"; function SelectOrgContent(): React.ReactElement { const { isLoaded, setActive, userMemberships } = useOrganizationList({ - userMemberships: { - infinite: true, - pageSize: 100, - }, + userMemberships: { infinite: true, pageSize: 100 }, }); - - useEffect(() => { - if (userMemberships?.hasNextPage && !userMemberships.isFetching) { - userMemberships.fetchNext?.(); - } - }, [userMemberships?.hasNextPage, userMemberships?.isFetching]); const { orgId } = useAuth(); const { user } = useUser(); const searchParams = useSearchParams(); const router = useRouter(); + const [stage, setStage] = useState("organization"); const [isSelecting, setIsSelecting] = useState(false); - const [selectedOrgId, setSelectedOrgId] = useState(orgId || null); + const [selectedOrgId, setSelectedOrgId] = useState( + orgId || null, + ); + const [projects, setProjects] = useState([]); + const [projectQuery, setProjectQuery] = useState(""); + const [hasMoreProjects, setHasMoreProjects] = useState(false); + const [nextProjectOffset, setNextProjectOffset] = useState(); + const [isLoadingProjects, setIsLoadingProjects] = useState(false); + const [selectedScope, setSelectedScope] = useState("organization"); + const [projectsError, setProjectsError] = useState(false); + const [selectionError, setSelectionError] = useState(null); const [canScrollUp, setCanScrollUp] = useState(false); const [canScrollDown, setCanScrollDown] = useState(false); const scrollContainerRef = useRef(null); + const projectRequestRef = useRef(0); + const lastLoadedProjectQueryRef = useRef(null); + const supportsProjectScope = + searchParams.get("code_challenge_method") === "S256" && + Boolean(searchParams.get("code_challenge")); - // Check if we just returned from org creation and reload to get fresh data useEffect(() => { - if (searchParams.get('org_created') === 'true') { - // Remove the flag from URL and reload to get fresh organization data + if (userMemberships?.hasNextPage && !userMemberships.isFetching) { + userMemberships.fetchNext?.(); + } + }, [userMemberships?.hasNextPage, userMemberships?.isFetching]); + + useEffect(() => { + if (searchParams.get("org_created") === "true") { const newUrl = new URL(window.location.href); - newUrl.searchParams.delete('org_created'); + newUrl.searchParams.delete("org_created"); window.location.href = newUrl.toString(); } }, [searchParams]); - // Check scroll state const updateScrollState = (): void => { const container = scrollContainerRef.current; if (!container) return; - const { scrollTop, scrollHeight, clientHeight } = container; setCanScrollUp(scrollTop > 0); setCanScrollDown(scrollTop < scrollHeight - clientHeight); }; - // Update scroll state when content changes useEffect(() => { updateScrollState(); - }, [userMemberships?.data]); - - // Get the original OAuth parameters from the URL - const originalParams = { - client_id: searchParams.get('client_id'), - redirect_uri: searchParams.get('redirect_uri'), - response_type: searchParams.get('response_type'), - scope: searchParams.get('scope'), - state: searchParams.get('state'), - code_challenge: searchParams.get('code_challenge'), - code_challenge_method: searchParams.get('code_challenge_method'), - }; + }, [userMemberships?.data, projects, stage]); - const handleOrgSelect = (organizationId: string): void => { - setSelectedOrgId(organizationId); - }; + const loadProjectsPage = useCallback( + async ({ + organizationId, + query, + offset, + append, + }: { + organizationId: string; + query: string; + offset: number; + append: boolean; + }): Promise => { + const requestId = ++projectRequestRef.current; + setIsLoadingProjects(true); + setProjectsError(false); - const handleConfirm = async (): Promise => { - if (!setActive || isSelecting || !selectedOrgId) return; + try { + const params = new URLSearchParams({ + org_id: organizationId, + limit: "20", + offset: String(offset), + }); + if (query) params.set("query", query); + const response = await fetch(`/oauth/projects?${params.toString()}`, { + cache: "no-store", + }); + if (!response.ok) throw new Error("failed to load projects"); + const body = (await response.json()) as { + projects: OAuthProject[]; + has_more: boolean; + next_offset?: number; + }; + if (requestId !== projectRequestRef.current) return; + + setProjects((current) => + append ? [...current, ...body.projects] : body.projects, + ); + setHasMoreProjects(body.has_more); + setNextProjectOffset(body.next_offset); + lastLoadedProjectQueryRef.current = query; + } catch (error) { + if (requestId !== projectRequestRef.current) return; + console.error("Failed to load projects:", error); + if (!append) { + setProjects([]); + setHasMoreProjects(false); + setNextProjectOffset(undefined); + lastLoadedProjectQueryRef.current = null; + } + setProjectsError(true); + } finally { + if (requestId === projectRequestRef.current) { + setIsLoadingProjects(false); + } + } + }, + [], + ); + + useEffect(() => { + if ( + stage !== "scope" || + !supportsProjectScope || + !selectedOrgId || + projectQuery === lastLoadedProjectQueryRef.current + ) { + return; + } + const timeout = setTimeout(() => { + setSelectedScope("organization"); + void loadProjectsPage({ + organizationId: selectedOrgId, + query: projectQuery, + offset: 0, + append: false, + }); + }, 250); + return () => clearTimeout(timeout); + }, [ + loadProjectsPage, + projectQuery, + selectedOrgId, + stage, + supportsProjectScope, + ]); + + const handleOrgConfirm = async (): Promise => { + if (!setActive || isSelecting || !selectedOrgId) return; setIsSelecting(true); + setProjectsError(false); + setSelectionError(null); try { await setActive({ organization: selectedOrgId }); + } catch (error) { + console.error("Failed to select organization:", error); + setSelectionError("organization selection failed. please try again."); + setIsSelecting(false); + return; + } - // After setting active org, redirect back to authorize - const authorizeUrl = new URL('/authorize', window.location.origin); - - // Add all original OAuth parameters - Object.entries(originalParams).forEach(([key, value]) => { - if (value) authorizeUrl.searchParams.set(key, value); + setProjectQuery(""); + lastLoadedProjectQueryRef.current = null; + if (supportsProjectScope) { + await loadProjectsPage({ + organizationId: selectedOrgId, + query: "", + offset: 0, + append: false, }); + } else { + setProjects([]); + setHasMoreProjects(false); + setNextProjectOffset(undefined); + } - // Add the selected orgId as a parameter - authorizeUrl.searchParams.set('org_id', selectedOrgId); + setSelectedScope("organization"); + setStage("scope"); + setIsSelecting(false); + }; - router.push(authorizeUrl.toString()); - } catch (error) { - console.error('Failed to set active organization:', error); - setIsSelecting(false); + const handleAuthorize = (): void => { + if (!selectedOrgId || isSelecting) return; + setIsSelecting(true); + + const authorizeUrl = new URL("/authorize", window.location.origin); + const params = new URLSearchParams(searchParams.toString()); + params.delete("org_created"); + params.delete("org_id"); + params.delete("access_scope"); + params.delete("project_id"); + params.forEach((value, key) => authorizeUrl.searchParams.set(key, value)); + authorizeUrl.searchParams.set("org_id", selectedOrgId); + + if (selectedScope.startsWith("project:")) { + authorizeUrl.searchParams.set("access_scope", "project"); + authorizeUrl.searchParams.set( + "project_id", + selectedScope.slice("project:".length), + ); + } else { + authorizeUrl.searchParams.set("access_scope", "organization"); } + + router.push(authorizeUrl.toString()); }; if (!isLoaded || userMemberships?.isLoading) { return ( -

loading your organizations...

+

+ loading your organizations... +

); } - // Check if user has any organizations (only after loaded) - if (isLoaded && !userMemberships?.isLoading && (!userMemberships?.data || userMemberships.data.length === 0)) { + const memberships = + userMemberships?.data || user?.organizationMemberships || []; + + if (!memberships.length) { return ( - {/* User button in top-right corner */}
- - +

you need to be a member of at least one organization to continue.

@@ -128,7 +262,7 @@ function SelectOrgContent(): React.ReactElement { { const params = new URLSearchParams(searchParams.toString()); - params.set('org_created', 'true'); + params.set("org_created", "true"); return `/select-org?${params.toString()}`; })()} skipInvitationScreen={true} @@ -138,9 +272,12 @@ function SelectOrgContent(): React.ReactElement { ); } + const selectedOrg = memberships.find( + (membership) => membership.organization.id === selectedOrgId, + )?.organization; + return ( - {/* User button in top-right corner */}

- select an organization to authorize access. + {stage === "organization" + ? "select an organization to authorize access." + : `choose access for ${selectedOrg?.name || "this organization"}.`}

-
-
- {(userMemberships?.data || user?.organizationMemberships) - ?.sort((a, b) => { - // Put the currently active org first - if (a.organization.id === orgId) return -1; - if (b.organization.id === orgId) return 1; - return 0; - }) - ?.map((membership, index, arr) => { - const isSelected = membership.organization.id === selectedOrgId; - const isCurrentlyActive = membership.organization.id === orgId; - const isLast = index === arr.length - 1; - return ( - - ); - })} + + ); + })} +
+ {canScrollUp && ( +
+ )} + {canScrollDown && ( +
+ )}
+ ) : ( + +
+ setSelectedScope("organization")} + /> +
+ {supportsProjectScope && ( + <> + +
+ + or restrict access + +
+ +
+
+
+ + setProjectQuery(event.target.value) + } + placeholder="search projects" + aria-label="search projects" + className="w-full bg-transparent border-[0.5px] border-[#e1dccf] px-3 py-2 text-sm text-foreground placeholder:text-muted-foreground focus:outline-none focus:border-foreground" + /> +
+ {projects.map((project) => ( + + setSelectedScope(`project:${project.id}`) + } + /> + ))} + {isLoadingProjects && projects.length === 0 && ( +

+ loading projects... +

+ )} + {!isLoadingProjects && + projects.length === 0 && + !hasMoreProjects && + !projectsError && ( +

+ no projects found. +

+ )} + {hasMoreProjects && ( + + )} +
+ {canScrollUp && ( +
+ )} + {canScrollDown && ( +
+ )} +
+ + )} + + )} - {/* Fade overlays to indicate scrollability */} - {canScrollUp && ( -
- )} - {canScrollDown && ( -
- )} -
+ {selectionError && ( +

{selectionError}

+ )} + {stage === "scope" && projectsError && ( +

+ projects could not be loaded. organization-wide access is still + available. +

+ )} + {stage === "scope" && !supportsProjectScope && ( +

+ this client does not support PKCE, so project-scoped access is + unavailable. +

+ )} - + + {stage === "scope" && ( + + )} + + ); } +function ScopeButton({ + label, + detail, + selected, + onClick, +}: { + label: string; + detail: string; + selected: boolean; + onClick: () => void; +}): React.ReactElement { + return ( + + ); +} + function LoadingFallback(): React.ReactElement { return ( diff --git a/src/app/token/route.test.ts b/src/app/token/route.test.ts new file mode 100644 index 0000000..49c5799 --- /dev/null +++ b/src/app/token/route.test.ts @@ -0,0 +1,253 @@ +import { describe, expect, test } from "bun:test"; +import { NextRequest } from "next/server"; +import type { TokenDependencies } from "./route"; +import { + organizationAuthorizationContext, + projectAuthorizationContext, +} from "@/lib/oauth-context"; + +process.env.KERNEL_CLI_PROD_CLIENT_ID ??= "cli_prod"; +process.env.KERNEL_CLI_STAGING_CLIENT_ID ??= "cli_staging"; +process.env.KERNEL_CLI_DEV_CLIENT_ID ??= "cli_dev"; +process.env.NEXT_PUBLIC_CLERK_DOMAIN ??= "clerk.example.test"; +process.env.CLERK_SECRET_KEY ??= "clerk-secret"; + +const { tokenRequest } = await import("./route"); + +function request(values: Record) { + return new NextRequest("https://auth.example.test/token", { + method: "POST", + headers: { "Content-Type": "application/x-www-form-urlencoded" }, + body: new URLSearchParams(values), + }); +} + +function dependencies({ + authorizationContext = organizationAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + }), + subject = "user_1", + member = true, + clerkStatus = 200, + clerkTokens = { + access_token: "clerk-access", + id_token: "header.payload.signature", + refresh_token: "refresh-new", + expires_in: 3600, + token_type: "Bearer", + }, +}: { + authorizationContext?: + | ReturnType + | ReturnType; + subject?: string; + member?: boolean; + clerkStatus?: number; + clerkTokens?: Record; +} = {}) { + const calls = { + resolve: [] as Parameters[0][], + exchanges: [] as URLSearchParams[], + jwt: [] as Parameters[0][], + refresh: [] as Parameters[0][], + rotate: [] as Parameters[0][], + deleted: [] as Parameters[0][], + }; + + return { + calls, + value: { + resolveContext: async (value) => { + calls.resolve.push(value); + return { + authorizationContext, + ...(value.grantType === "authorization_code" + ? { requestCodeChallenge: "derived-challenge" } + : {}), + }; + }, + exchange: async (_input, init) => { + calls.exchanges.push( + new URLSearchParams(init?.body as URLSearchParams), + ); + return Response.json(clerkTokens, { status: clerkStatus }); + }, + verify: async () => ({ sub: subject }), + hasMembership: async () => member, + setJwtContext: async (value) => { + calls.jwt.push(value); + }, + setRefreshContext: async (value) => { + calls.refresh.push(value); + }, + rotateRefreshContext: async (value) => { + calls.rotate.push(value); + }, + deleteRequestContext: async (value) => { + calls.deleted.push(value); + }, + } satisfies TokenDependencies, + }; +} + +describe("POST /token", () => { + test("issues an organization-wide token and persists both contexts", async () => { + const deps = dependencies(); + const response = await tokenRequest( + request({ + grant_type: "authorization_code", + client_id: "client_1", + code: "code_1", + code_verifier: "verifier_1", + org_id: "forged-org", + access_scope: "project", + project_id: "forged-project", + }), + deps.value, + ); + + expect(response.status).toBe(200); + expect(await response.json()).toMatchObject({ + access_token: "header.payload.signature", + refresh_token: "refresh-new", + org_id: "org_1", + access_scope: "organization", + }); + expect(deps.calls.jwt).toHaveLength(1); + expect(deps.calls.refresh).toHaveLength(1); + expect(deps.calls.deleted).toEqual([ + { clientId: "client_1", codeChallenge: "derived-challenge" }, + ]); + expect(deps.calls.exchanges[0].has("org_id")).toBe(false); + expect(deps.calls.exchanges[0].has("project_id")).toBe(false); + expect(deps.calls.exchanges[0].has("access_scope")).toBe(false); + }); + + test("returns the server-bound project scope", async () => { + const deps = dependencies({ + authorizationContext: projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }), + }); + const response = await tokenRequest( + request({ + grant_type: "authorization_code", + client_id: "client_1", + code: "code_1", + code_verifier: "verifier_1", + }), + deps.value, + ); + + expect(await response.json()).toMatchObject({ + org_id: "org_1", + access_scope: "project", + project_id: "proj_1", + }); + }); + + test("refresh preserves stored scope and ignores body escalation", async () => { + const deps = dependencies({ + authorizationContext: projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }), + }); + const response = await tokenRequest( + request({ + grant_type: "refresh_token", + client_id: "cli_prod", + refresh_token: "refresh-old", + org_id: "org_other", + access_scope: "organization", + project_id: "proj_other", + }), + deps.value, + ); + + expect(response.status).toBe(200); + expect(deps.calls.resolve[0]).toMatchObject({ + grantType: "refresh_token", + clientId: "cli_prod", + refreshToken: "refresh-old", + }); + expect(deps.calls.rotate).toEqual([ + expect.objectContaining({ + oldRefreshToken: "refresh-old", + newRefreshToken: "refresh-new", + authorizationContext: expect.objectContaining({ + access_scope: "project", + project_id: "proj_1", + }), + }), + ]); + expect(deps.calls.exchanges[0].has("org_id")).toBe(false); + expect(deps.calls.exchanges[0].has("project_id")).toBe(false); + }); + + test("fails closed on subject and membership mismatches", async () => { + const wrongSubject = dependencies({ subject: "user_2" }); + const subjectResponse = await tokenRequest( + request({ + grant_type: "authorization_code", + client_id: "client_1", + code: "code_1", + code_verifier: "verifier_1", + }), + wrongSubject.value, + ); + expect(subjectResponse.status).toBe(400); + expect(wrongSubject.calls.jwt).toHaveLength(0); + + const removedMember = dependencies({ member: false }); + const memberResponse = await tokenRequest( + request({ + grant_type: "refresh_token", + client_id: "client_1", + refresh_token: "refresh-old", + }), + removedMember.value, + ); + expect(memberResponse.status).toBe(400); + expect(removedMember.calls.rotate).toHaveLength(0); + }); + + test("does not persist partial context when provider response is invalid", async () => { + const providerFailure = dependencies({ clerkStatus: 401 }); + const failedResponse = await tokenRequest( + request({ + grant_type: "authorization_code", + client_id: "client_1", + code: "bad", + code_verifier: "verifier_1", + }), + providerFailure.value, + ); + expect(failedResponse.status).toBe(400); + expect(providerFailure.calls.jwt).toHaveLength(0); + + const missingRefresh = dependencies({ + clerkTokens: { + access_token: "clerk-access", + id_token: "header.payload.signature", + expires_in: 3600, + token_type: "Bearer", + }, + }); + const missingRefreshResponse = await tokenRequest( + request({ + grant_type: "authorization_code", + client_id: "client_1", + code: "code_1", + code_verifier: "verifier_1", + }), + missingRefresh.value, + ); + expect(missingRefreshResponse.status).toBe(400); + expect(missingRefresh.calls.jwt).toHaveLength(0); + }); +}); diff --git a/src/app/token/route.ts b/src/app/token/route.ts index c1b3034..9781891 100644 --- a/src/app/token/route.ts +++ b/src/app/token/route.ts @@ -1,185 +1,203 @@ +import { clerkClient, verifyToken } from "@clerk/nextjs/server"; import { NextRequest, NextResponse } from "next/server"; import { - setOrgIdForJwt, - setOrgIdForRefreshToken, - deleteOrgIdForRefreshToken, + deleteAuthorizationContextForRequest, + rotateAuthorizationContextForRefreshToken, + setAuthorizationContextForJwt, + setAuthorizationContextForRefreshToken, } from "@/lib/redis"; -import { resolveOrgId } from "@/lib/org-utils"; +import { resolveAuthorizationContext } from "@/lib/org-utils"; import { REFRESH_TOKEN_ORG_TTL_SECONDS } from "@/lib/const"; import { normalizeLocalhostUri } from "@/lib/auth-utils"; +const CORS_HEADERS = { + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Methods": "POST, OPTIONS", + "Access-Control-Allow-Headers": "Content-Type, Authorization", +}; + +interface ClerkTokenResponse { + access_token: string; + expires_in: number; + refresh_token?: string; + token_type: string; + id_token?: string; +} + export async function OPTIONS(): Promise { - return new NextResponse(null, { - status: 204, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }); + return new NextResponse(null, { status: 204, headers: CORS_HEADERS }); } function createErrorResponse( error: string, errorDescription: string, - status: number = 400, + status = 400, ) { return NextResponse.json( - { - error, - error_description: errorDescription, - }, - { - status, - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", - }, - }, + { error, error_description: errorDescription }, + { status, headers: CORS_HEADERS }, ); } -interface ClerkTokenResponse { - access_token: string; - expires_in: number; - refresh_token: string; - token_type: string; - id_token?: string; // Optional, only present for authorization_code grants +function clientCredentials( + request: NextRequest, + body: FormData, + params: URLSearchParams, +): { clientId: string | null } { + let clientId = body.get("client_id")?.toString() || null; + let clientSecret = body.get("client_secret")?.toString() || null; + + if (!clientId) { + const authHeader = request.headers.get("authorization"); + if (authHeader?.startsWith("Basic ")) { + try { + const decoded = atob(authHeader.slice(6)); + const colonIndex = decoded.indexOf(":"); + if (colonIndex !== -1) { + clientId = decodeURIComponent(decoded.slice(0, colonIndex)); + clientSecret = decodeURIComponent(decoded.slice(colonIndex + 1)); + params.set("client_id", clientId); + if (clientSecret) params.set("client_secret", clientSecret); + } + } catch { + // Missing client_id is returned below. + } + } + } + + return { clientId }; } -export async function POST(request: NextRequest): Promise { - // Step 1: Validate request format +async function hasOrganizationMembership( + clerkUserId: string, + clerkOrgId: string, +): Promise { + const clerk = await clerkClient(); + let offset = 0; + const limit = 100; + + for (;;) { + const memberships = await clerk.users.getOrganizationMembershipList({ + userId: clerkUserId, + limit, + offset, + }); + if ( + memberships.data.some( + (membership) => membership.organization.id === clerkOrgId, + ) + ) { + return true; + } + offset += memberships.data.length; + if (memberships.data.length === 0 || offset >= memberships.totalCount) { + return false; + } + } +} + +type Fetcher = ( + input: URL | RequestInfo, + init?: RequestInit, +) => Promise; + +export interface TokenDependencies { + exchange: Fetcher; + resolveContext: typeof resolveAuthorizationContext; + verify: ( + token: string, + options: { secretKey?: string }, + ) => Promise<{ sub?: string }>; + hasMembership: typeof hasOrganizationMembership; + setJwtContext: typeof setAuthorizationContextForJwt; + setRefreshContext: typeof setAuthorizationContextForRefreshToken; + rotateRefreshContext: typeof rotateAuthorizationContextForRefreshToken; + deleteRequestContext: typeof deleteAuthorizationContextForRequest; +} + +const tokenDependencies: TokenDependencies = { + exchange: fetch, + resolveContext: resolveAuthorizationContext, + verify: verifyToken, + hasMembership: hasOrganizationMembership, + setJwtContext: setAuthorizationContextForJwt, + setRefreshContext: setAuthorizationContextForRefreshToken, + rotateRefreshContext: rotateAuthorizationContextForRefreshToken, + deleteRequestContext: deleteAuthorizationContextForRequest, +}; + +export async function tokenRequest( + request: NextRequest, + dependencies: TokenDependencies = tokenDependencies, +): Promise { const contentType = request.headers.get("content-type"); if (!contentType?.includes("application/x-www-form-urlencoded")) { - console.debug("[token] invalid content-type", { contentType }); return createErrorResponse( "invalid_request", "Content-Type must be application/x-www-form-urlencoded", ); } - const body = await request.formData(); - - // Step 2: Validate server configuration const clerkDomain = process.env.NEXT_PUBLIC_CLERK_DOMAIN; - if (!clerkDomain) { - console.error("NEXT_PUBLIC_CLERK_DOMAIN environment variable is not set"); return createErrorResponse( "server_error", - "Server configuration error - clerk domain not found", + "Server configuration error", 500, ); } - try { - // Step 3: Prepare parameters for Clerk token exchange - // Normalize redirect_uri to match Vercel's query param normalization (127.0.0.1 → localhost) - const params = new URLSearchParams(); - for (const [key, value] of body.entries()) { - if (key === "redirect_uri") { - params.append(key, normalizeLocalhostUri(value.toString())); - } else { - params.append(key, value.toString()); - } - } - - const grantType = body.get("grant_type") as string; - console.debug("[token] start", { grantType }); - - // Extract client_id from body or Authorization: Basic header - let clientId = body.get("client_id") as string | null; - let clientSecret = body.get("client_secret") as string | null; - - if (!clientId) { - const authHeader = request.headers.get("authorization"); - if (authHeader?.startsWith("Basic ")) { - try { - const decoded = atob(authHeader.slice(6)); - const colonIdx = decoded.indexOf(":"); - if (colonIdx !== -1) { - clientId = decodeURIComponent(decoded.slice(0, colonIdx)); - clientSecret = decodeURIComponent(decoded.slice(colonIdx + 1)); - params.set("client_id", clientId); - if (clientSecret) { - params.set("client_secret", clientSecret); - } - console.debug("[token] extracted client_id from Basic auth header"); - } - } catch { - console.debug("[token] failed to decode Basic auth header"); - } - } - } + const body = await request.formData(); + const params = new URLSearchParams(); + for (const [key, value] of body.entries()) { + params.append( + key, + key === "redirect_uri" + ? normalizeLocalhostUri(value.toString()) + : value.toString(), + ); + } - if (!clientId) { - console.debug("[token] missing client_id"); - return createErrorResponse( - "invalid_request", - "Missing required parameter: client_id", - ); - } + const grantType = body.get("grant_type")?.toString() || ""; + const { clientId } = clientCredentials(request, body, params); + if (!clientId) { + return createErrorResponse( + "invalid_request", + "Missing required parameter: client_id", + ); + } - // Extract direct org_id if provided (shared clients) - let directOrgId: string | undefined; - const directOrgIdParam = body.get("org_id"); - if (directOrgIdParam) { - const orgIdParam = directOrgIdParam.toString(); - directOrgId = orgIdParam; - const maskedOrgId = orgIdParam.slice(0, 4) + "..." + orgIdParam.slice(-4); - console.debug("[token] using org_id from request body", { maskedOrgId }); - } + const refreshToken = body.get("refresh_token")?.toString(); + const contextResult = await dependencies.resolveContext({ + grantType, + clientId, + codeVerifier: body.get("code_verifier")?.toString(), + refreshToken, + }); + if (contextResult.error) return contextResult.error; + const authorizationContext = contextResult.authorizationContext; + if (!authorizationContext) { + return createErrorResponse( + "invalid_grant", + "Authorization context not found. Please re-authorize.", + ); + } - // For refresh_token flow, resolve org before calling Clerk - let resolvedOrgId: string | null = null; - let refreshTokenFromBody: string | null = null; - if (grantType === "refresh_token") { - refreshTokenFromBody = body.get("refresh_token") as string | null; - if (!refreshTokenFromBody) { - console.debug("[token] missing refresh_token in refresh flow"); - return createErrorResponse( - "invalid_request", - "Missing required parameter: refresh_token", - ); - } - const orgResultPre = await resolveOrgId({ - grantType, - clientId, - directOrgId, - refreshToken: refreshTokenFromBody, - }); - if (orgResultPre.error) { - console.debug( - "[token] resolveOrgId (pre) returned error for refresh flow", - ); - return orgResultPre.error; - } - resolvedOrgId = orgResultPre.orgId; - if (!resolvedOrgId) { - console.debug("[token] no org_id resolved for refresh_token flow"); - return createErrorResponse( - "invalid_grant", - "Organization context not found for refresh token. Please re-authorize.", - ); - } - console.debug("[token] resolved org via refresh_token mapping"); - } + // Internal context parameters are not part of Clerk's token endpoint. + params.delete("org_id"); + params.delete("project_id"); + params.delete("access_scope"); - // Step 4: Exchange with Clerk - const clerkTokenResponse = await fetch( + try { + const clerkResponse = await dependencies.exchange( `https://${clerkDomain}/oauth/token`, { method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - }, + headers: { "Content-Type": "application/x-www-form-urlencoded" }, body: params, }, ); - - if (!clerkTokenResponse.ok) { - console.error("[token] clerk token exchange failed"); + if (!clerkResponse.ok) { return createErrorResponse( "invalid_grant", grantType === "refresh_token" @@ -188,61 +206,88 @@ export async function POST(request: NextRequest): Promise { ); } - const clerkTokens: ClerkTokenResponse = await clerkTokenResponse.json(); - - // Step 5: Resolve organization context for authorization_code flows (or confirm for refresh flows) - let orgId = resolvedOrgId; - if (!orgId) { - const orgResult = await resolveOrgId({ - grantType, - clientId, - directOrgId, - }); - if (orgResult.error) { - console.debug("[token] resolveOrgId returned error for auth_code flow"); - return orgResult.error; - } - orgId = orgResult.orgId; + const clerkTokens = (await clerkResponse.json()) as ClerkTokenResponse; + if (!clerkTokens.id_token) { + return createErrorResponse( + "invalid_grant", + "Failed to retrieve id_token from OAuth provider", + ); + } + if ( + !Number.isFinite(clerkTokens.expires_in) || + clerkTokens.expires_in <= 0 + ) { + return createErrorResponse( + "invalid_grant", + "OAuth provider returned an invalid token lifetime", + ); } - // Step 6: Validate organization context - if (!orgId || orgId === "") { - console.warn("[token] no org_id resolved for client", { - clientIdMasked: clientId.slice(0, 4) + "...", - }); + const payload = await dependencies.verify(clerkTokens.id_token, { + secretKey: process.env.CLERK_SECRET_KEY, + }); + if (!payload.sub) { + return createErrorResponse( + "invalid_grant", + "OAuth provider returned a token without a subject", + ); + } + if ( + authorizationContext.clerk_user_id && + authorizationContext.clerk_user_id !== payload.sub + ) { + return createErrorResponse( + "invalid_grant", + "OAuth token subject does not match authorization context", + ); + } + if ( + !(await dependencies.hasMembership( + payload.sub, + authorizationContext.clerk_org_id, + )) + ) { return createErrorResponse( "invalid_grant", - "Unable to resolve organization context. Please re-authorize.", + "Organization membership is no longer active", ); } - // Step 7: Validate grant type and extract JWT - let finalJwt: string; - let expiresIn: number; + const issuedRefreshToken = clerkTokens.refresh_token; + if (!issuedRefreshToken) { + return createErrorResponse( + "invalid_grant", + grantType === "refresh_token" + ? "OAuth provider did not rotate the refresh token" + : "OAuth provider did not return a refresh token", + ); + } + const finalJwt = clerkTokens.id_token; + await dependencies.setJwtContext({ + jwt: finalJwt, + authorizationContext, + ttlSeconds: clerkTokens.expires_in, + }); if (grantType === "authorization_code") { - // For authorization_code: Use id_token directly (already has proper structure) - if (!clerkTokens.id_token) { - console.debug("[token] missing id_token in auth_code response"); - return createErrorResponse( - "invalid_grant", - "Failed to retrieve id_token from Clerk authorization code", - ); - } - - finalJwt = clerkTokens.id_token; - expiresIn = clerkTokens.expires_in; + await dependencies.setRefreshContext({ + refreshToken: issuedRefreshToken, + authorizationContext, + ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, + }); } else if (grantType === "refresh_token") { - if (!clerkTokens.id_token) { - console.debug("[token] missing id_token in refresh response"); + if (!refreshToken) { return createErrorResponse( "invalid_grant", - "Failed to retrieve id_token from Clerk refresh token", + "Missing required parameter: refresh_token", ); } - - finalJwt = clerkTokens.id_token; - expiresIn = clerkTokens.expires_in; + await dependencies.rotateRefreshContext({ + oldRefreshToken: refreshToken, + newRefreshToken: issuedRefreshToken, + authorizationContext, + ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, + }); } else { return createErrorResponse( "unsupported_grant_type", @@ -250,113 +295,43 @@ export async function POST(request: NextRequest): Promise { ); } - // Step 8: Store refresh_token → org_id mapping (where applicable) - try { - if (grantType === "authorization_code" && clerkTokens.refresh_token) { - await setOrgIdForRefreshToken({ - refreshToken: clerkTokens.refresh_token, - orgId, - ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, + if (contextResult.requestCodeChallenge) { + try { + await dependencies.deleteRequestContext({ + clientId, + codeChallenge: contextResult.requestCodeChallenge, }); - console.debug( - "[token] stored refresh_token→org_id mapping (auth_code)", + } catch (error) { + // The provider has already consumed the authorization code and both + // token mappings are durable. Let the short pending-context TTL clean up. + console.warn( + "[token] failed to delete consumed authorization context", + { + error, + }, ); } - if (grantType === "refresh_token") { - // Update mapping for rotated refresh token if provided - if (clerkTokens.refresh_token) { - // Clean up old mapping before storing the new one - if (refreshTokenFromBody) { - try { - await deleteOrgIdForRefreshToken({ - refreshToken: refreshTokenFromBody, - }); - console.debug( - "[token] deleted old refresh_token→org_id mapping (refresh)", - ); - } catch (e) { - console.warn( - "[token] failed to delete old refresh_token mapping", - { error: e }, - ); - } - } - await setOrgIdForRefreshToken({ - refreshToken: clerkTokens.refresh_token, - orgId, - ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, - }); - console.debug( - "[token] updated refresh_token→org_id mapping (refresh)", - ); - } - } - } catch (error) { - console.error("[token] failed to store refresh_token→org_id mapping", { - error, - }); - return createErrorResponse( - "server_error", - "Failed to store refresh token context", - 500, - ); - } - - // Step 9: Store JWT to org_id mapping for verifyjwt.go - try { - // Store JWT to org_id mapping with JWT expiration time - await setOrgIdForJwt({ - jwt: finalJwt, - orgId, - ttlSeconds: expiresIn, - }); - console.debug("[token] stored jwt→org_id mapping", { - ttlSeconds: expiresIn, - }); - } catch (error) { - console.error("[token] failed to store jwt→org_id mapping", { error }); - return createErrorResponse( - "server_error", - "Failed to store authentication context", - 500, - ); } - // Step 10: Build final token response - const mcpTokenResponse = { - ...clerkTokens, - access_token: finalJwt, - expires_in: expiresIn, - }; - - console.debug("[token] success", { grantType }); - return NextResponse.json(mcpTokenResponse, { - headers: { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", + return NextResponse.json( + { + ...clerkTokens, + access_token: finalJwt, + expires_in: clerkTokens.expires_in, + org_id: authorizationContext.clerk_org_id, + access_scope: authorizationContext.access_scope, + ...(authorizationContext.project_id + ? { project_id: authorizationContext.project_id } + : {}), }, - }); + { headers: CORS_HEADERS }, + ); } catch (error) { - console.error("[token] unhandled error", { error }); - - // If it's a Clerk error, log the detailed error information - if (error && typeof error === "object" && "clerkError" in error) { - const clerkError = error as any; - console.error("Clerk error details:"); - console.error(" Status:", clerkError.status); - console.error(" Clerk Trace ID:", clerkError.clerkTraceId); - if (clerkError.errors && Array.isArray(clerkError.errors)) { - console.error(" Specific errors:"); - clerkError.errors.forEach((err: any, index: number) => { - console.error( - ` Error ${index + 1}:`, - JSON.stringify(err, null, 2), - ); - }); - } - } - + console.error("[token] token exchange failed", { error }); return createErrorResponse("server_error", "Internal server error", 500); } } + +export async function POST(request: NextRequest): Promise { + return tokenRequest(request); +} diff --git a/src/lib/oauth-context.test.ts b/src/lib/oauth-context.test.ts new file mode 100644 index 0000000..65d368b --- /dev/null +++ b/src/lib/oauth-context.test.ts @@ -0,0 +1,100 @@ +import { describe, expect, test } from "bun:test"; +import { + authorizationContextFromSelection, + deriveS256CodeChallenge, + organizationAuthorizationContext, + parseAuthorizationContext, + projectAuthorizationContext, + serializeAuthorizationContext, +} from "./oauth-context"; + +describe("OAuth authorization context", () => { + test("decodes legacy org mappings as organization-wide", () => { + expect(parseAuthorizationContext("org_legacy")).toEqual( + organizationAuthorizationContext({ clerkOrgId: "org_legacy" }), + ); + }); + + test("round-trips organization and project contexts", () => { + const contexts = [ + organizationAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + }), + projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }), + ]; + + for (const context of contexts) { + expect( + parseAuthorizationContext(serializeAuthorizationContext(context)), + ).toEqual(context); + } + }); + + test("rejects malformed authorization boundaries", () => { + for (const value of [ + "", + "{", + "not-an-org", + "null", + "[]", + '{"version":2,"clerk_org_id":"org_1","access_scope":"organization"}', + '{"version":1,"access_scope":"organization"}', + '{"version":1,"clerk_org_id":"org_1","access_scope":"project"}', + '{"version":1,"clerk_org_id":"org_1","access_scope":"organization","project_id":"proj_1"}', + '{"version":1,"clerk_org_id":"org_1","access_scope":"account"}', + ]) { + expect(() => parseAuthorizationContext(value)).toThrow(); + } + }); + + test("validates authorization selections", () => { + expect( + authorizationContextFromSelection({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + accessScope: null, + projectId: null, + }), + ).toEqual( + organizationAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + }), + ); + + expect( + authorizationContextFromSelection({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + accessScope: "project", + projectId: "proj_1", + }), + ).toEqual( + projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }), + ); + + expect(() => + authorizationContextFromSelection({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + accessScope: "project", + projectId: null, + }), + ).toThrow("Project access requires a project"); + }); +}); + +test("derives RFC 7636 S256 challenge", () => { + expect( + deriveS256CodeChallenge("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"), + ).toBe("E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"); +}); diff --git a/src/lib/oauth-context.ts b/src/lib/oauth-context.ts new file mode 100644 index 0000000..d87f5be --- /dev/null +++ b/src/lib/oauth-context.ts @@ -0,0 +1,159 @@ +import { createHash } from "crypto"; + +export const OAUTH_CONTEXT_VERSION = 1; +export const ORGANIZATION_ACCESS_SCOPE = "organization"; +export const PROJECT_ACCESS_SCOPE = "project"; + +export type OAuthAccessScope = + | typeof ORGANIZATION_ACCESS_SCOPE + | typeof PROJECT_ACCESS_SCOPE; + +interface OAuthAuthorizationContextBase { + version: typeof OAUTH_CONTEXT_VERSION; + clerk_user_id?: string; + clerk_org_id: string; +} + +export type OAuthAuthorizationContext = OAuthAuthorizationContextBase & + ( + | { + access_scope: typeof ORGANIZATION_ACCESS_SCOPE; + project_id?: never; + } + | { + access_scope: typeof PROJECT_ACCESS_SCOPE; + project_id: string; + } + ); + +export function organizationAuthorizationContext({ + clerkUserId, + clerkOrgId, +}: { + clerkUserId?: string; + clerkOrgId: string; +}): OAuthAuthorizationContext { + return { + version: OAUTH_CONTEXT_VERSION, + ...(clerkUserId ? { clerk_user_id: clerkUserId } : {}), + clerk_org_id: clerkOrgId, + access_scope: ORGANIZATION_ACCESS_SCOPE, + }; +} + +export function projectAuthorizationContext({ + clerkUserId, + clerkOrgId, + projectId, +}: { + clerkUserId?: string; + clerkOrgId: string; + projectId: string; +}): OAuthAuthorizationContext { + return { + version: OAUTH_CONTEXT_VERSION, + ...(clerkUserId ? { clerk_user_id: clerkUserId } : {}), + clerk_org_id: clerkOrgId, + access_scope: PROJECT_ACCESS_SCOPE, + project_id: projectId, + }; +} + +export function parseAuthorizationContext( + value: string, +): OAuthAuthorizationContext { + const trimmed = value.trim(); + if (!trimmed) throw new Error("OAuth authorization context is empty"); + + // Existing access and refresh token mappings contain only a Clerk org ID. + if (trimmed.startsWith("org_")) { + return organizationAuthorizationContext({ clerkOrgId: trimmed }); + } + + const decoded: unknown = JSON.parse(trimmed); + if (!decoded || typeof decoded !== "object" || Array.isArray(decoded)) { + throw new Error("OAuth authorization context must be an object"); + } + const parsed = decoded as Partial; + if (parsed.version !== OAUTH_CONTEXT_VERSION) { + throw new Error( + `Unsupported OAuth authorization context version: ${String(parsed.version)}`, + ); + } + if (!parsed.clerk_org_id) { + throw new Error("OAuth authorization context is missing clerk_org_id"); + } + if ( + parsed.clerk_user_id !== undefined && + typeof parsed.clerk_user_id !== "string" + ) { + throw new Error("OAuth authorization context has invalid clerk_user_id"); + } + + if (parsed.access_scope === ORGANIZATION_ACCESS_SCOPE) { + if (parsed.project_id) { + throw new Error( + "Organization-scoped OAuth context cannot include project_id", + ); + } + return organizationAuthorizationContext({ + clerkUserId: parsed.clerk_user_id, + clerkOrgId: parsed.clerk_org_id, + }); + } + + if (parsed.access_scope === PROJECT_ACCESS_SCOPE) { + if (!parsed.project_id) { + throw new Error("Project-scoped OAuth context is missing project_id"); + } + return projectAuthorizationContext({ + clerkUserId: parsed.clerk_user_id, + clerkOrgId: parsed.clerk_org_id, + projectId: parsed.project_id, + }); + } + + throw new Error( + `Unsupported OAuth access scope: ${String(parsed.access_scope)}`, + ); +} + +export function serializeAuthorizationContext( + context: OAuthAuthorizationContext, +): string { + // Validate before writing so malformed boundaries never enter Redis. + return JSON.stringify(parseAuthorizationContext(JSON.stringify(context))); +} + +export function deriveS256CodeChallenge(codeVerifier: string): string { + return createHash("sha256").update(codeVerifier).digest("base64url"); +} + +export function authorizationContextFromSelection({ + clerkUserId, + clerkOrgId, + accessScope, + projectId, +}: { + clerkUserId: string; + clerkOrgId: string; + accessScope: string | null; + projectId: string | null; +}): OAuthAuthorizationContext { + const normalizedScope = accessScope || ORGANIZATION_ACCESS_SCOPE; + if (normalizedScope === ORGANIZATION_ACCESS_SCOPE) { + if (projectId) { + throw new Error("Organization-wide access cannot include a project"); + } + return organizationAuthorizationContext({ clerkUserId, clerkOrgId }); + } + if (normalizedScope === PROJECT_ACCESS_SCOPE) { + if (!projectId) throw new Error("Project access requires a project"); + return projectAuthorizationContext({ + clerkUserId, + clerkOrgId, + projectId, + }); + } + throw new Error(`Unsupported OAuth access scope: ${normalizedScope}`); +} diff --git a/src/lib/oauth-projects.test.ts b/src/lib/oauth-projects.test.ts new file mode 100644 index 0000000..3fc8dae --- /dev/null +++ b/src/lib/oauth-projects.test.ts @@ -0,0 +1,124 @@ +import { afterEach, beforeEach, describe, expect, mock, test } from "bun:test"; +import { + listOAuthProjectsPage, + OAuthProjectsError, + requireActiveOAuthProject, +} from "./oauth-projects"; + +const originalApiBaseUrl = process.env.API_BASE_URL; + +beforeEach(() => { + process.env.API_BASE_URL = "https://api.example.test"; +}); + +afterEach(() => { + if (originalApiBaseUrl === undefined) delete process.env.API_BASE_URL; + else process.env.API_BASE_URL = originalApiBaseUrl; +}); + +describe("OAuth project lookup", () => { + test("forwards search and pagination to the API", async () => { + let requestedUrl: URL | undefined; + const fetcher = mock(async (input: RequestInfo | URL) => { + requestedUrl = new URL(input.toString()); + return Response.json( + [ + { id: "proj_1", name: "one", status: "active" }, + { id: "proj_old", name: "old", status: "archived" }, + ], + { headers: { "X-Has-More": "true", "X-Next-Offset": "40" } }, + ); + }); + + const page = await listOAuthProjectsPage({ + clerkSessionToken: "session-token", + query: "prod", + limit: 20, + offset: 20, + fetcher, + }); + + expect(requestedUrl?.pathname).toBe("/org/projects"); + expect(requestedUrl?.searchParams.get("query")).toBe("prod"); + expect(requestedUrl?.searchParams.get("limit")).toBe("20"); + expect(requestedUrl?.searchParams.get("offset")).toBe("20"); + expect(page).toEqual({ + projects: [{ id: "proj_1", name: "one", status: "active" }], + hasMore: true, + nextOffset: 40, + }); + }); + + test("validates the selected project directly by ID", async () => { + let requestedUrl: URL | undefined; + const fetcher = mock(async (input: RequestInfo | URL) => { + requestedUrl = new URL(input.toString()); + return Response.json({ id: "proj_1", name: "one", status: "active" }); + }); + + expect( + await requireActiveOAuthProject({ + clerkSessionToken: "session-token", + projectId: "proj_1", + fetcher, + }), + ).toEqual({ id: "proj_1", name: "one", status: "active" }); + expect(requestedUrl?.pathname).toBe("/org/projects/proj_1"); + }); + + test("rejects missing, inactive, and mismatched projects", async () => { + const missing = mock(async () => new Response(null, { status: 404 })); + await expect( + requireActiveOAuthProject({ + clerkSessionToken: "session-token", + projectId: "proj_1", + fetcher: missing, + }), + ).rejects.toMatchObject({ status: 404 }); + + const forbidden = mock(async () => new Response(null, { status: 403 })); + await expect( + requireActiveOAuthProject({ + clerkSessionToken: "session-token", + projectId: "proj_1", + fetcher: forbidden, + }), + ).rejects.toMatchObject({ status: 403 }); + + for (const project of [ + { id: "proj_1", name: "one", status: "archived" }, + { id: "proj_2", name: "two", status: "active" }, + ]) { + const fetcher = mock(async () => Response.json(project)); + await expect( + requireActiveOAuthProject({ + clerkSessionToken: "session-token", + projectId: "proj_1", + fetcher, + }), + ).rejects.toBeInstanceOf(OAuthProjectsError); + } + }); + + test("fails on API and pagination errors", async () => { + const unavailable = mock(async () => new Response(null, { status: 503 })); + await expect( + listOAuthProjectsPage({ + clerkSessionToken: "session-token", + fetcher: unavailable, + }), + ).rejects.toMatchObject({ status: 502 }); + + const invalidPagination = mock(async () => + Response.json([], { + headers: { "X-Has-More": "true", "X-Next-Offset": "0" }, + }), + ); + await expect( + listOAuthProjectsPage({ + clerkSessionToken: "session-token", + fetcher: invalidPagination, + }), + ).rejects.toThrow("Invalid project pagination response"); + }); +}); diff --git a/src/lib/oauth-projects.ts b/src/lib/oauth-projects.ts new file mode 100644 index 0000000..fa80789 --- /dev/null +++ b/src/lib/oauth-projects.ts @@ -0,0 +1,126 @@ +export type OAuthProjectsFetcher = ( + input: URL | RequestInfo, + init?: RequestInit, +) => Promise; + +export interface OAuthProject { + id: string; + name: string; + status: "active" | "archived"; +} + +export interface OAuthProjectsPage { + projects: OAuthProject[]; + hasMore: boolean; + nextOffset?: number; +} + +export class OAuthProjectsError extends Error { + constructor( + message: string, + readonly status: number, + ) { + super(message); + } +} + +function apiBaseUrl(): string { + const value = process.env.API_BASE_URL; + if (!value) { + throw new OAuthProjectsError("API_BASE_URL is not configured", 500); + } + return value; +} + +function projectHeaders(clerkSessionToken: string): HeadersInit { + return { + Authorization: `Bearer ${clerkSessionToken}`, + "X-Source": "oauth-server", + }; +} + +export async function listOAuthProjectsPage({ + clerkSessionToken, + query, + limit = 20, + offset = 0, + fetcher = fetch, +}: { + clerkSessionToken: string; + query?: string; + limit?: number; + offset?: number; + fetcher?: OAuthProjectsFetcher; +}): Promise { + const url = new URL("/org/projects", apiBaseUrl()); + url.searchParams.set("limit", String(limit)); + url.searchParams.set("offset", String(offset)); + if (query) url.searchParams.set("query", query); + + const response = await fetcher(url, { + headers: projectHeaders(clerkSessionToken), + cache: "no-store", + }); + if (!response.ok) { + throw new OAuthProjectsError( + `Failed to load projects (${response.status})`, + response.status === 401 || response.status === 403 + ? response.status + : 502, + ); + } + + const page = (await response.json()) as OAuthProject[]; + const hasMore = response.headers.get("X-Has-More") === "true"; + const nextOffsetValue = response.headers.get("X-Next-Offset"); + const nextOffset = nextOffsetValue ? Number(nextOffsetValue) : undefined; + if ( + hasMore && + (nextOffset === undefined || + !Number.isInteger(nextOffset) || + nextOffset <= offset) + ) { + throw new OAuthProjectsError("Invalid project pagination response", 502); + } + + return { + projects: page.filter((project) => project.status === "active"), + hasMore, + ...(hasMore ? { nextOffset } : {}), + }; +} + +export async function requireActiveOAuthProject({ + clerkSessionToken, + projectId, + fetcher = fetch, +}: { + clerkSessionToken: string; + projectId: string; + fetcher?: OAuthProjectsFetcher; +}): Promise { + const url = new URL( + `/org/projects/${encodeURIComponent(projectId)}`, + apiBaseUrl(), + ); + const response = await fetcher(url, { + headers: projectHeaders(clerkSessionToken), + cache: "no-store", + }); + if (!response.ok) { + throw new OAuthProjectsError( + "Project not found or inactive", + response.status === 401 || + response.status === 403 || + response.status === 404 + ? response.status + : 502, + ); + } + + const project = (await response.json()) as OAuthProject; + if (project.status !== "active" || project.id !== projectId) { + throw new OAuthProjectsError("Project not found or inactive", 404); + } + return project; +} diff --git a/src/lib/org-utils.test.ts b/src/lib/org-utils.test.ts new file mode 100644 index 0000000..e18b4b9 --- /dev/null +++ b/src/lib/org-utils.test.ts @@ -0,0 +1,123 @@ +import { beforeEach, describe, expect, test } from "bun:test"; +import { + type OAuthAuthorizationContext, + projectAuthorizationContext, +} from "./oauth-context"; +import type { AuthorizationContextDependencies } from "./org-utils"; + +process.env.KERNEL_CLI_PROD_CLIENT_ID ??= "cli_prod"; +process.env.KERNEL_CLI_STAGING_CLIENT_ID ??= "cli_staging"; +process.env.KERNEL_CLI_DEV_CLIENT_ID ??= "cli_dev"; + +const { resolveAuthorizationContext } = await import("./org-utils"); + +let requestContext: OAuthAuthorizationContext | null = null; +let clientContext: OAuthAuthorizationContext | null = null; +let refreshContext: OAuthAuthorizationContext | null = null; +let requestLookup: { clientId: string; codeChallenge: string } | null = null; + +const dependencies: AuthorizationContextDependencies = { + getRequestContext: async (value) => { + requestLookup = value; + return requestContext; + }, + getClientContext: async () => clientContext, + getRefreshContext: async () => refreshContext, +}; + +function resolveContext( + input: Parameters[0], +) { + return resolveAuthorizationContext(input, dependencies); +} + +beforeEach(() => { + requestContext = null; + clientContext = null; + refreshContext = null; + requestLookup = null; +}); + +describe("resolveAuthorizationContext", () => { + test("uses the PKCE-bound request context", async () => { + requestContext = projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }); + + const result = await resolveContext({ + grantType: "authorization_code", + clientId: "client_1", + codeVerifier: "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk", + }); + + expect(requestLookup).toEqual({ + clientId: "client_1", + codeChallenge: "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM", + }); + expect(result.authorizationContext).toEqual(requestContext); + expect(result.requestCodeChallenge).toBe( + "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM", + ); + }); + + test("refresh uses only the refresh-token mapping", async () => { + refreshContext = projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_1", + projectId: "proj_1", + }); + + const result = await resolveContext({ + grantType: "refresh_token", + clientId: "cli_prod", + refreshToken: "refresh-token", + }); + + expect(result.authorizationContext).toEqual(refreshContext); + }); + + test("shared clients cannot supply authorization context directly", async () => { + const result = await resolveContext({ + grantType: "authorization_code", + clientId: "cli_prod", + }); + + expect(result.authorizationContext).toBeNull(); + expect(result.error?.status).toBe(400); + }); + + test("does not fall back to client context when PKCE context is missing", async () => { + clientContext = projectAuthorizationContext({ + clerkUserId: "user_1", + clerkOrgId: "org_wrong", + projectId: "proj_wrong", + }); + + const result = await resolveContext({ + grantType: "authorization_code", + clientId: "client_1", + codeVerifier: "verifier", + }); + + expect(result.authorizationContext).toBeNull(); + expect(result.error?.status).toBe(400); + }); + + test("fails when request and refresh context have expired", async () => { + const authCode = await resolveContext({ + grantType: "authorization_code", + clientId: "client_1", + codeVerifier: "verifier", + }); + expect(authCode.error?.status).toBe(400); + + const refresh = await resolveContext({ + grantType: "refresh_token", + clientId: "client_1", + refreshToken: "refresh-token", + }); + expect(refresh.error?.status).toBe(400); + }); +}); diff --git a/src/lib/org-utils.ts b/src/lib/org-utils.ts index 92c6646..5587aae 100644 --- a/src/lib/org-utils.ts +++ b/src/lib/org-utils.ts @@ -1,17 +1,22 @@ import { NextResponse } from "next/server"; -import { getOrgIdForClientId, getOrgIdForRefreshTokenSliding } from "./redis"; -import { REFRESH_TOKEN_ORG_TTL_SECONDS, SHARED_CLIENT_IDS } from "./const"; +import { + getAuthorizationContextForClientId, + getAuthorizationContextForRefreshTokenSliding, + getAuthorizationContextForRequest, +} from "./redis"; +import { REFRESH_TOKEN_ORG_TTL_SECONDS } from "./const"; +import { + deriveS256CodeChallenge, + type OAuthAuthorizationContext, +} from "./oauth-context"; function createErrorResponse( error: string, errorDescription: string, - status: number = 400, + status = 400, ) { return NextResponse.json( - { - error, - error_description: errorDescription, - }, + { error, error_description: errorDescription }, { status, headers: { @@ -23,142 +28,106 @@ function createErrorResponse( ); } -/** - * Resolves organization ID based on client type and grant type - * @param grantType - OAuth grant type ("authorization_code" or "refresh_token") - * @param clientId - The OAuth client ID - * @param directOrgId - The org_id from OAuth state parameter (shared clients only) - * @param refreshToken - The refresh token from OAuth state parameter (shared clients only) - * @returns { orgId: string | null, error?: NextResponse } - org_id or error response - */ -export async function resolveOrgId({ - grantType, - clientId, - directOrgId, - refreshToken, -}: { - grantType: string; - clientId: string; - directOrgId?: string; - refreshToken?: string; -}): Promise<{ orgId: string | null; error?: NextResponse }> { - const isSharedClient = SHARED_CLIENT_IDS.includes(clientId); - const clientIdMasked = clientId ? clientId.slice(0, 4) + "..." : ""; - console.debug("[org-utils] resolveOrgId", { - grantType, - isSharedClient, - hasDirectOrgId: Boolean(directOrgId), - hasRefreshToken: Boolean(refreshToken), - }); +export interface AuthorizationContextDependencies { + getRequestContext: typeof getAuthorizationContextForRequest; + getClientContext: typeof getAuthorizationContextForClientId; + getRefreshContext: typeof getAuthorizationContextForRefreshTokenSliding; +} + +const authorizationContextDependencies: AuthorizationContextDependencies = { + getRequestContext: getAuthorizationContextForRequest, + getClientContext: getAuthorizationContextForClientId, + getRefreshContext: getAuthorizationContextForRefreshTokenSliding, +}; + +export interface ResolvedAuthorizationContext { + authorizationContext: OAuthAuthorizationContext | null; + requestCodeChallenge?: string; + error?: NextResponse; +} - if (isSharedClient) { - // Shared clients (CLI): Use org_id from OAuth state parameter - if (directOrgId) { - console.debug("[org-utils] shared client: using direct org_id from body"); - return { orgId: directOrgId }; - } else if (grantType === "refresh_token") { - if (!refreshToken) { - console.warn( - "[org-utils] shared client: missing refresh_token in request body", - ); +export async function resolveAuthorizationContext( + { + grantType, + clientId, + codeVerifier, + refreshToken, + }: { + grantType: string; + clientId: string; + codeVerifier?: string; + refreshToken?: string; + }, + dependencies: AuthorizationContextDependencies = authorizationContextDependencies, +): Promise { + if (grantType === "authorization_code") { + if (codeVerifier) { + const codeChallenge = deriveS256CodeChallenge(codeVerifier); + try { + const authorizationContext = await dependencies.getRequestContext({ + clientId, + codeChallenge, + }); + if (authorizationContext) { + return { + authorizationContext, + requestCodeChallenge: codeChallenge, + }; + } return { - orgId: null, + authorizationContext: null, error: createErrorResponse( - "invalid_request", - "Missing required parameter: refresh_token", + "invalid_grant", + "Authorization context expired. Please re-authorize.", ), }; - } - try { - const orgIdFromRefresh = await getOrgIdForRefreshTokenSliding({ - refreshToken, - ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, + } catch (error) { + console.error("[org-utils] failed to read PKCE authorization context", { + error, }); - if (orgIdFromRefresh) { - console.debug( - "[org-utils] shared client: resolved org via refresh_token mapping", - ); - return { orgId: orgIdFromRefresh }; - } - } catch (e) { - console.error( - "[org-utils] shared client: error reading refresh_token mapping", - { error: e }, - ); return { - orgId: null, + authorizationContext: null, error: createErrorResponse( "server_error", - "Failed to retrieve organization context for refresh token.", + "Failed to retrieve authorization context", + 500, ), }; } - console.warn( - "[org-utils] shared client: missing org_id and no refresh mapping", - ); - return { - orgId: null, - error: createErrorResponse( - "invalid_grant", - "Missing organization context in refresh request. Please re-authorize.", - ), - }; - } else { - console.warn("[org-utils] shared client: missing org_id in request body"); - return { - orgId: null, - error: createErrorResponse( - "invalid_grant", - "Missing organization context in OAuth request body. Please re-authorize.", - ), - }; } - } - // Ephemeral clients (MCP) - if (grantType === "authorization_code") { - // Use client_id mapping just to bridge the initial hop try { - const orgId = await getOrgIdForClientId({ clientId }); - if (orgId) { - console.debug( - "[org-utils] ephemeral client: resolved org via client_id mapping", - ); - return { orgId }; - } else { - console.warn( - "[org-utils] ephemeral client: missing org context for authorization_code", - { clientIdMasked }, - ); - return { - orgId: null, - error: createErrorResponse( - "invalid_grant", - "Organization context expired for client: " + - clientId + - ". Please re-authorize to select your organization.", - ), - }; - } + const authorizationContext = await dependencies.getClientContext({ + clientId, + }); + if (authorizationContext) return { authorizationContext }; } catch (error) { - console.error( - "[org-utils] ephemeral client: error reading client_id mapping", - { error, clientIdMasked }, - ); + console.error("[org-utils] failed to read client authorization context", { + error, + }); return { - orgId: null, + authorizationContext: null, error: createErrorResponse( "server_error", - "Failed to retrieve organization context for client: " + clientId, + "Failed to retrieve authorization context", + 500, ), }; } - } else if (grantType === "refresh_token") { - // Use refresh_token mapping for ongoing refresh flows + + return { + authorizationContext: null, + error: createErrorResponse( + "invalid_grant", + "Authorization context expired. Please re-authorize.", + ), + }; + } + + if (grantType === "refresh_token") { if (!refreshToken) { - console.debug("[org-utils] refresh flow without refresh_token param"); return { - orgId: null, + authorizationContext: null, error: createErrorResponse( "invalid_request", "Missing required parameter: refresh_token", @@ -166,40 +135,39 @@ export async function resolveOrgId({ }; } try { - const orgId = await getOrgIdForRefreshTokenSliding({ + const authorizationContext = await dependencies.getRefreshContext({ refreshToken, ttlSeconds: REFRESH_TOKEN_ORG_TTL_SECONDS, }); - if (orgId) { - console.debug("[org-utils] resolved org via refresh_token mapping"); - return { orgId }; - } - console.warn( - "[org-utils] no org mapping for refresh_token (expired or unknown)", - ); - return { - orgId: null, - error: createErrorResponse( - "invalid_grant", - "Organization context expired for this refresh token. Please re-authorize.", - ), - }; + if (authorizationContext) return { authorizationContext }; } catch (error) { - console.error("[org-utils] error reading refresh_token mapping", { - error, - }); + console.error( + "[org-utils] failed to read refresh authorization context", + { + error, + }, + ); return { - orgId: null, + authorizationContext: null, error: createErrorResponse( "server_error", - "Failed to retrieve organization context for refresh token.", + "Failed to retrieve authorization context for refresh token", + 500, ), }; } + + return { + authorizationContext: null, + error: createErrorResponse( + "invalid_grant", + "Authorization context expired for this refresh token. Please re-authorize.", + ), + }; } return { - orgId: null, + authorizationContext: null, error: createErrorResponse( "unsupported_grant_type", `Grant type '${grantType}' is not supported`, diff --git a/src/lib/redis.ts b/src/lib/redis.ts index a85ac27..b149483 100644 --- a/src/lib/redis.ts +++ b/src/lib/redis.ts @@ -1,6 +1,11 @@ import { createClient } from "redis"; import { createHmac } from "crypto"; import { mcpAppsMarkerKey } from "@/lib/mcp-apps-marker"; +import { + type OAuthAuthorizationContext, + parseAuthorizationContext, + serializeAuthorizationContext, +} from "@/lib/oauth-context"; const redisUrl = process.env.REDIS_URL; const redisTlsServerName = process.env.REDIS_TLS_SERVER_NAME; @@ -108,43 +113,110 @@ function hashOpaqueToken(token: string): string { return createHmac("sha256", secretKey).update(token).digest("hex"); } -export async function setOrgIdForClientId({ +function authorizationRequestKey( + clientId: string, + codeChallenge: string, +): string { + return `oauth-request:${hashOpaqueToken(`${clientId}:${codeChallenge}`)}`; +} + +export async function setAuthorizationContextForRequest({ clientId, - orgId, + codeChallenge, + authorizationContext, ttlSeconds, }: { clientId: string; - orgId: string; + codeChallenge: string; + authorizationContext: OAuthAuthorizationContext; ttlSeconds: number; }): Promise { await ensureConnected(); - const key = `client:${clientId}`; - await withReconnect(() => client.setEx(key, ttlSeconds, orgId)); + await withReconnect(() => + client.setEx( + authorizationRequestKey(clientId, codeChallenge), + ttlSeconds, + serializeAuthorizationContext(authorizationContext), + ), + ); +} + +export async function getAuthorizationContextForRequest({ + clientId, + codeChallenge, +}: { + clientId: string; + codeChallenge: string; +}): Promise { + await ensureConnected(); + const value = await withReconnect(() => + client.get(authorizationRequestKey(clientId, codeChallenge)), + ); + return value ? parseAuthorizationContext(value) : null; +} + +export async function deleteAuthorizationContextForRequest({ + clientId, + codeChallenge, +}: { + clientId: string; + codeChallenge: string; +}): Promise { + await ensureConnected(); + await withReconnect(() => + client.del(authorizationRequestKey(clientId, codeChallenge)), + ); } -export async function getOrgIdForClientId({ +export async function setAuthorizationContextForClientId({ clientId, + authorizationContext, + ttlSeconds, }: { clientId: string; -}): Promise { + authorizationContext: OAuthAuthorizationContext; + ttlSeconds: number; +}): Promise { await ensureConnected(); const key = `client:${clientId}`; - return await withReconnect(() => client.get(key)); + await withReconnect(() => + client.setEx( + key, + ttlSeconds, + serializeAuthorizationContext(authorizationContext), + ), + ); +} + +export async function getAuthorizationContextForClientId({ + clientId, +}: { + clientId: string; +}): Promise { + await ensureConnected(); + const value = await withReconnect(() => client.get(`client:${clientId}`)); + return value ? parseAuthorizationContext(value) : null; } -export async function setOrgIdForJwt({ +export async function setAuthorizationContextForJwt({ jwt, - orgId, + authorizationContext, ttlSeconds, }: { jwt: string; - orgId: string; + authorizationContext: OAuthAuthorizationContext; ttlSeconds: number; }): Promise { await ensureConnected(); const hashedJwt = hashJwt(jwt); const key = `jwt:${hashedJwt}`; - await withReconnect(() => client.setEx(key, ttlSeconds, orgId)); + await withReconnect(() => + client.setEx( + key, + ttlSeconds, + serializeAuthorizationContext(authorizationContext), + ), + ); } export { client as redisClient }; @@ -205,48 +277,65 @@ export async function hasMcpAppsClient({ return value !== null; } -export async function setOrgIdForRefreshToken({ +export async function setAuthorizationContextForRefreshToken({ refreshToken, - orgId, + authorizationContext, ttlSeconds, }: { refreshToken: string; - orgId: string; + authorizationContext: OAuthAuthorizationContext; ttlSeconds: number; }): Promise { await ensureConnected(); - const hashed = hashOpaqueToken(refreshToken); - const key = `refresh:${hashed}`; - await withReconnect(() => client.setEx(key, ttlSeconds, orgId)); + const key = `refresh:${hashOpaqueToken(refreshToken)}`; + await withReconnect(() => + client.setEx( + key, + ttlSeconds, + serializeAuthorizationContext(authorizationContext), + ), + ); } -export async function getOrgIdForRefreshTokenSliding({ +export async function getAuthorizationContextForRefreshTokenSliding({ refreshToken, ttlSeconds, }: { refreshToken: string; ttlSeconds: number; -}): Promise { +}): Promise { await ensureConnected(); - const hashed = hashOpaqueToken(refreshToken); - const key = `refresh:${hashed}`; - const orgId = await withReconnect(() => client.get(key)); - if (orgId) { - // Refresh TTL to implement sliding expiration on active tokens - await withReconnect(() => client.expire(key, ttlSeconds)); - } - return orgId; + const key = `refresh:${hashOpaqueToken(refreshToken)}`; + const value = await withReconnect(() => + client.getEx(key, { type: "EX", value: ttlSeconds }), + ); + return value ? parseAuthorizationContext(value) : null; } -export async function deleteOrgIdForRefreshToken({ - refreshToken, +export async function rotateAuthorizationContextForRefreshToken({ + oldRefreshToken, + newRefreshToken, + authorizationContext, + ttlSeconds, }: { - refreshToken: string; + oldRefreshToken: string; + newRefreshToken: string; + authorizationContext: OAuthAuthorizationContext; + ttlSeconds: number; }): Promise { await ensureConnected(); - const hashed = hashOpaqueToken(refreshToken); - const key = `refresh:${hashed}`; - await withReconnect(() => client.del(key)); + const oldKey = `refresh:${hashOpaqueToken(oldRefreshToken)}`; + const newKey = `refresh:${hashOpaqueToken(newRefreshToken)}`; + const value = serializeAuthorizationContext(authorizationContext); + + if (oldKey === newKey) { + await withReconnect(() => client.setEx(newKey, ttlSeconds, value)); + return; + } + + await withReconnect(async () => { + await client.multi().setEx(newKey, ttlSeconds, value).del(oldKey).exec(); + }); } function isTransientSocketError(error: unknown): boolean {