From fe9b59e111f50b5fbf55b11e7f5ccdca9d786f2f Mon Sep 17 00:00:00 2001 From: Amruth Pillai Date: Sat, 5 Sep 2026 09:33:15 -0700 Subject: [PATCH] fix: restore MCP OAuth registration and authorization (#3421) * fix: align MCP OAuth provider schema and authorization flow * test: isolate OpenAPI generation from OAuth initialization * fix: accept auth routes without a callback query * fix: require explicit OAuth consent and preserve signed requests * test: verify OAuth audiences through real MCP initialization * test(e2e): isolate OAuth token audience validation --- apps/server/src/http/app.test.ts | 14 + apps/server/src/http/app.ts | 8 + apps/server/src/http/auth.test.ts | 133 +- apps/server/src/http/auth.ts | 165 +- .../src/http/oauth-flow.integration.test.ts | 254 + apps/server/src/index.test.ts | 35 + apps/server/src/index.ts | 9 +- apps/server/src/openapi/generator.test.ts | 7 +- apps/server/turbo.json | 6 +- apps/web/locales/en-US.po | 49 + .../features/auth/components/social-auth.tsx | 50 +- .../src/features/auth/login-redirect.test.tsx | 209 + .../src/features/auth/pages/consent.test.tsx | 96 + apps/web/src/features/auth/pages/consent.tsx | 123 + apps/web/src/features/auth/pages/login.tsx | 33 +- apps/web/src/features/auth/pages/register.tsx | 27 +- .../src/features/auth/pages/verify-2fa.tsx | 17 +- apps/web/src/features/auth/redirect.test.ts | 58 + apps/web/src/features/auth/redirect.ts | 63 + apps/web/src/libs/auth/client.ts | 10 +- apps/web/src/routeTree.gen.ts | 21 + apps/web/src/routes/auth/consent.tsx | 26 + apps/web/src/routes/auth/index.tsx | 7 +- apps/web/src/routes/auth/login.tsx | 5 +- apps/web/src/routes/auth/register.tsx | 7 +- apps/web/src/routes/auth/route.tsx | 2 + .../web/src/routes/auth/verify-2fa-backup.tsx | 5 +- apps/web/src/routes/auth/verify-2fa.tsx | 5 +- docs/guides/using-the-mcp-server.mdx | 12 +- .../migration.sql | 58 + .../snapshot.json | 6235 +++++++++++++++++ packages/auth/src/config.ts | 12 +- packages/auth/src/oauth-schema.test.ts | 18 + packages/db/src/schema/auth.ts | 83 + tests/e2e/specs/oauth-consent.spec.ts | 119 + 35 files changed, 7837 insertions(+), 144 deletions(-) create mode 100644 apps/server/src/http/oauth-flow.integration.test.ts create mode 100644 apps/server/src/index.test.ts create mode 100644 apps/web/src/features/auth/login-redirect.test.tsx create mode 100644 apps/web/src/features/auth/pages/consent.test.tsx create mode 100644 apps/web/src/features/auth/pages/consent.tsx create mode 100644 apps/web/src/features/auth/redirect.test.ts create mode 100644 apps/web/src/features/auth/redirect.ts create mode 100644 apps/web/src/routes/auth/consent.tsx create mode 100644 migrations/20260905113135_oauth_provider_schema/migration.sql create mode 100644 migrations/20260905113135_oauth_provider_schema/snapshot.json create mode 100644 packages/auth/src/oauth-schema.test.ts create mode 100644 tests/e2e/specs/oauth-consent.spec.ts diff --git a/apps/server/src/http/app.test.ts b/apps/server/src/http/app.test.ts index 902644e24..e626d254d 100644 --- a/apps/server/src/http/app.test.ts +++ b/apps/server/src/http/app.test.ts @@ -205,3 +205,17 @@ describe("createApp", () => { expect(mocks.serveWebDistStatic).not.toHaveBeenCalled(); }); }); + +it.each(["/auth/consent", "/auth/consent/", "/auth/login"])("prevents framing or caching %s", async (path) => { + const { createApp } = await import("./app"); + mocks.serveWebDistStatic.mockImplementationOnce(async (_context: unknown, next: () => Promise) => { + await next(); + }); + const response = await createApp().request(`http://localhost:3000${path}?sig=signed`); + expect(response.status).toBe(200); + expect(await response.text()).toBe("web"); + expect(response.headers.get("content-security-policy")).toBe("frame-ancestors 'none'"); + expect(response.headers.get("x-frame-options")).toBe("DENY"); + expect(response.headers.get("referrer-policy")).toBe("no-referrer"); + expect(response.headers.get("cache-control")).toBe("no-store"); +}); diff --git a/apps/server/src/http/app.ts b/apps/server/src/http/app.ts index 677f92d5a..a53e56763 100644 --- a/apps/server/src/http/app.ts +++ b/apps/server/src/http/app.ts @@ -36,6 +36,14 @@ const getTrustedClient = (context: Context): string => { export function createApp() { const app = new Hono(); + app.use("/auth/*", async (c, next) => { + await next(); + c.header("Content-Security-Policy", "frame-ancestors 'none'"); + c.header("X-Frame-Options", "DENY"); + c.header("Referrer-Policy", "no-referrer"); + c.header("Cache-Control", "no-store"); + }); + app.all("/api/rpc", (c) => handleRpc(c.req.raw, getTrustedClient(c))); app.all("/api/rpc/*", (c) => handleRpc(c.req.raw, getTrustedClient(c))); app.all("/api/openapi", (c) => handleOpenApi(c.req.raw, getTrustedClient(c))); diff --git a/apps/server/src/http/auth.test.ts b/apps/server/src/http/auth.test.ts index 8017ece9b..4807795e8 100644 --- a/apps/server/src/http/auth.test.ts +++ b/apps/server/src/http/auth.test.ts @@ -2,6 +2,8 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; const mocks = vi.hoisted(() => ({ getSession: vi.fn(), + consent: vi.fn(), + continueOAuth: vi.fn(), handler: vi.fn(), env: { SERVER_PORT: 3001, @@ -14,6 +16,8 @@ vi.mock("@reactive-resume/auth/config", () => ({ auth: { api: { getSession: mocks.getSession, + oauth2Consent: mocks.consent, + oauth2Continue: mocks.continueOAuth, }, handler: mocks.handler, }, @@ -32,6 +36,23 @@ beforeEach(() => { }); describe("handleAuth", () => { + it.for([null, false, 42, "client", [], [{ redirect_uris: [] }]])( + "rejects non-object registration payload %j", + async (body) => { + const { handleAuth } = await import("./auth"); + const response = await handleAuth( + new Request("http://localhost:3000/api/auth/oauth2/register", { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify(body), + }), + ); + expect(response.status).toBe(400); + await expect(response.json()).resolves.toEqual({ message: "Invalid registration payload" }); + expect(mocks.handler).not.toHaveBeenCalled(); + }, + ); + it("rejects untrusted dynamic OAuth redirect URIs in safe mode", async () => { const { handleAuth } = await import("./auth"); @@ -66,6 +87,57 @@ describe("handleAuth", () => { expect(response.status).toBe(200); expect(mocks.handler).toHaveBeenCalledOnce(); }); + it.each(["localhost", "127.0.0.1", "[::1]"])( + "infers native application type for exact %s loopback callbacks", + async (host) => { + const { handleAuth } = await import("./auth"); + await handleAuth( + new Request("http://localhost:3000/api/auth/oauth2/register", { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify({ redirect_uris: [`http://${host}:3210/callback`] }), + }), + ); + const forwarded = mocks.handler.mock.calls[0]?.[0] as Request; + await expect(forwarded.json()).resolves.toMatchObject({ + application_type: "native", + token_endpoint_auth_method: "none", + }); + }, + ); + + it.each([ + { redirect_uris: ["https://example.com/callback"] }, + { redirect_uris: ["http://localhost.evil.example/callback"] }, + { redirect_uris: ["http://localhost:3210/callback"], application_type: "web" }, + { redirect_uris: ["http://localhost:3210/callback", "https://example.com/callback"] }, + ])("does not infer native for explicit web or non-loopback clients: %j", async (body) => { + const { handleAuth } = await import("./auth"); + mocks.env.FLAG_ALLOW_UNSAFE_OAUTH_REDIRECT_URI = true; + await handleAuth( + new Request("http://localhost:3000/api/auth/oauth2/register", { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify(body), + }), + ); + const forwarded = mocks.handler.mock.calls[0]?.[0] as Request; + expect((await forwarded.json()).application_type).not.toBe("native"); + }); + + it("preserves repeated resource indicators during authorization sanitization", async () => { + const { handleAuth } = await import("./auth"); + await handleAuth( + new Request( + "http://localhost:3000/api/auth/oauth2/authorize?resource=http%3A%2F%2Flocalhost%3A3000&resource=http%3A%2F%2Flocalhost%3A3000%2Fmcp", + ), + ); + const forwarded = mocks.handler.mock.calls[0]?.[0] as Request; + expect(new URL(forwarded.url).searchParams.getAll("resource")).toEqual([ + "http://localhost:3000", + "http://localhost:3000/mcp", + ]); + }); }); describe("handleOAuth", () => { @@ -91,7 +163,64 @@ describe("handleOAuth", () => { expect(callbackUrl.searchParams.get("client_id")).toBe("test-client"); expect(callbackUrl.searchParams.get("redirect_uri")).toBe("https://example.com/callback"); expect(callbackUrl.searchParams.get("state")).toBe("abc"); - expect(callbackUrl.searchParams.has("exp")).toBe(false); - expect(callbackUrl.searchParams.has("sig")).toBe(false); + expect(callbackUrl.searchParams.get("exp")).toBe("123"); + expect(callbackUrl.searchParams.get("sig")).toBe("456"); + }); + it("continues signed authorization without approving consent on GET", async () => { + const { handleOAuth } = await import("./auth"); + mocks.getSession.mockResolvedValueOnce({ user: { id: "owner" } }); + mocks.continueOAuth.mockResolvedValueOnce( + Response.json({ redirect: true, url: "/auth/consent?client_id=client&sig=signed" }), + ); + const query = "client_id=client&resource=one&resource=two&exp=123&sig=456"; + const response = await handleOAuth(new Request(`http://localhost:3000/api/auth/oauth?${query}`)); + expect(mocks.continueOAuth).toHaveBeenCalledWith( + expect.objectContaining({ body: { postLogin: true, oauth_query: query } }), + ); + expect(mocks.consent).not.toHaveBeenCalled(); + expect(response.status).toBe(302); + expect(response.headers.get("location")).toBe("/auth/consent?client_id=client&sig=signed"); + }); + + it("preserves provider failures instead of issuing an authorization code", async () => { + const { handleOAuth } = await import("./auth"); + mocks.getSession.mockResolvedValueOnce({ user: { id: "owner" } }); + mocks.continueOAuth.mockResolvedValueOnce(Response.json({ error: "invalid_signature" }, { status: 400 })); + const response = await handleOAuth(new Request("http://localhost:3000/api/auth/oauth?sig=invalid")); + expect(response.status).toBe(400); + expect(response.headers.get("location")).toBeNull(); + await expect(response.json()).resolves.toEqual({ error: "invalid_signature" }); + }); + it("preserves provider cookies and cache headers on forced reauthentication", async () => { + const { handleOAuth } = await import("./auth"); + mocks.getSession.mockResolvedValueOnce({ user: { id: "owner" } }); + const headers = new Headers({ "cache-control": "no-store", "content-length": "123" }); + headers.append("set-cookie", "oauth_state=state; Path=/; HttpOnly"); + headers.append("set-cookie", "session=refreshed; Path=/; HttpOnly"); + mocks.continueOAuth.mockResolvedValueOnce( + Response.json({ redirect: true, url: "/api/auth/oauth?prompt=login&sig=signed" }, { headers }), + ); + const response = await handleOAuth(new Request("http://localhost:3000/api/auth/oauth?sig=original")); + expect(response.status).toBe(302); + expect(response.headers.get("location")).toMatch(/^\/auth\/login\?reauthenticate=true&/); + expect(response.headers.getSetCookie()).toEqual(headers.getSetCookie()); + expect(response.headers.get("cache-control")).toBe("no-store"); + expect(response.headers.get("content-type")).toBeNull(); + expect(response.headers.get("content-length")).toBeNull(); }); }); + +describe("OAuth provider response validation", () => { + it.for([{}, { url: null }, { url: 7 }, { url: "" }, { url: "undefined" }, { url: "javascript:alert(1)" }])( + "fails closed for malformed provider response %j", + async (body) => { + const { handleOAuth } = await import("./auth"); + mocks.getSession.mockResolvedValueOnce({ user: { id: "owner" } }); + mocks.continueOAuth.mockResolvedValueOnce(Response.json(body)); + const response = await handleOAuth(new Request("http://localhost:3000/api/auth/oauth?sig=signed")); + expect(response.status).toBe(502); + expect(response.headers.get("location")).toBeNull(); + expect(mocks.consent).not.toHaveBeenCalled(); + }, + ); +}); diff --git a/apps/server/src/http/auth.ts b/apps/server/src/http/auth.ts index ea2736e51..80b6ce27e 100644 --- a/apps/server/src/http/auth.ts +++ b/apps/server/src/http/auth.ts @@ -1,10 +1,6 @@ -import crypto from "node:crypto"; -import { eq } from "drizzle-orm"; +import { APIError } from "better-auth/api"; import { auth } from "@reactive-resume/auth/config"; -import { db } from "@reactive-resume/db/client"; -import { oauthClient, verification } from "@reactive-resume/db/schema"; import { env } from "@reactive-resume/env/server"; -import { generateId } from "@reactive-resume/utils/string"; import { isAllowedOAuthRedirectUri } from "@reactive-resume/utils/url-security.node"; const oauthAuthorizeSanitizedParams = [ @@ -19,8 +15,6 @@ const oauthAuthorizeSanitizedParams = [ "resource", ] as const; -const oauthCallbackPassthroughExcludedParams = new Set(["exp", "sig"]); - function sanitizeOAuthAuthorizeRequest(request: Request): Request { if (request.method !== "GET") return request; @@ -33,9 +27,10 @@ function sanitizeOAuthAuthorizeRequest(request: Request): Request { .replace(/\s+/g, " ") .trim(); const sanitizeParam = (key: string) => { - const value = url.searchParams.get(key); - if (!value) return; - url.searchParams.set(key, sanitizeValue(value)); + const values = url.searchParams.getAll(key); + if (!values.length) return; + url.searchParams.delete(key); + for (const value of values) url.searchParams.append(key, sanitizeValue(value)); }; for (const key of oauthAuthorizeSanitizedParams) sanitizeParam(key); @@ -56,6 +51,10 @@ function sanitizeOAuthAuthorizeRequest(request: Request): Request { return new Request(url.toString(), request); } +function isRegistrationPayload(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + async function defaultPublicClientRegistration(request: Request): Promise { if (request.method !== "POST") return request; @@ -66,11 +65,23 @@ async function defaultPublicClientRegistration(request: Request): Promise; try { - body = await cloned.json(); + const payload: unknown = await cloned.json(); + if (!isRegistrationPayload(payload)) return request; + body = payload; } catch { return request; } + // MCP native clients often omit OIDC application_type. Infer it only for + // exact HTTP loopback callbacks; the provider still validates every URI. + if (body.application_type === undefined && Array.isArray(body.redirect_uris) && body.redirect_uris.length > 0) { + const allLoopback = body.redirect_uris.every( + (uri: unknown) => + typeof uri === "string" && /^http:\/\/(?:localhost|127\.0\.0\.1|\[::1\])(?::[0-9]+)?(?:[/?]|$)/i.test(uri), + ); + if (allLoopback) body.application_type = "native"; + } + if (!request.headers.get("authorization")) { body.token_endpoint_auth_method = "none"; } @@ -92,7 +103,11 @@ async function validateDynamicClientRegistrationRequest(request: Request): Promi let body: Record; try { - body = await cloned.json(); + const payload: unknown = await cloned.json(); + if (!isRegistrationPayload(payload)) { + return Response.json({ message: "Invalid registration payload" }, { status: 400 }); + } + body = payload; } catch { return Response.json({ message: "Invalid registration payload" }, { status: 400 }); } @@ -125,90 +140,68 @@ export async function handleAuth(request: Request) { return auth.handler(finalRequest); } -function generateCode() { - return crypto.randomBytes(32).toString("base64url"); -} - -function hashCode(code: string) { - return crypto.createHash("sha256").update(code).digest("base64url"); -} - export async function handleOAuth(request: Request) { + try { + return await resumeOAuth(request); + } catch (error) { + // Before-hooks can throw even when the provider is called with asResponse. + if (error instanceof APIError) return Response.json(error.body, { status: error.statusCode }); + throw error; + } +} + +async function resumeOAuth(request: Request) { const session = await auth.api.getSession({ headers: request.headers }); const url = new URL(request.url); if (session?.user) { - const clientId = url.searchParams.get("client_id"); - const redirectUri = url.searchParams.get("redirect_uri"); - const state = url.searchParams.get("state"); - const scope = url.searchParams.get("scope"); - const codeChallenge = url.searchParams.get("code_challenge"); - const codeChallengeMethod = url.searchParams.get("code_challenge_method"); - - if (!clientId || !redirectUri) { - return Response.json({ error: "missing client_id or redirect_uri" }, { status: 400 }); - } - - const [client] = await db.select().from(oauthClient).where(eq(oauthClient.clientId, clientId)).limit(1); - - if (!client) { - return Response.json({ error: "invalid client" }, { status: 400 }); - } - - if (!client.redirectUris.includes(redirectUri)) { - return Response.json({ error: "invalid redirect_uri" }, { status: 400 }); - } - - const code = generateCode(); - const hashedCode = hashCode(code); - const now = new Date(); - const expiresAt = new Date(now.getTime() + 600_000); - - await db.insert(verification).values({ - id: generateId(), - identifier: hashedCode, - value: JSON.stringify({ - type: "authorization_code", - query: { - response_type: "code", - client_id: clientId, - redirect_uri: redirectUri, - scope, - state, - code_challenge: codeChallenge, - code_challenge_method: codeChallengeMethod, - }, - userId: session.user.id, - sessionId: session.session.id, - authTime: new Date(session.session.createdAt).getTime(), - }), - expiresAt, - createdAt: now, - updatedAt: now, - }); - - const callbackUrl = new URL(redirectUri); - callbackUrl.searchParams.set("code", code); - if (state) callbackUrl.searchParams.set("state", state); - callbackUrl.searchParams.set("iss", `${env.APP_URL}/api/auth`); - - return new Response(null, { - status: 302, - headers: { Location: callbackUrl.toString() }, + // Resume authorization without granting consent. The provider decides whether + // the user must sign in, explicitly approve a client, or reuse an existing grant. + // Its signed query must survive the login round trip byte-for-byte. + const response = await auth.api.oauth2Continue({ + asResponse: true, + request, + headers: request.headers, + body: { postLogin: true, oauth_query: url.search.slice(1) }, }); + if (!(response instanceof Response)) throw new Error("OAuth provider did not return a response"); + if (!response.ok) return response; + const result: unknown = await response.json().catch(() => null); + if ( + !result || + typeof result !== "object" || + !("url" in result) || + typeof result.url !== "string" || + !result.url || + !(result.url.startsWith("/") || URL.canParse(result.url)) || + !URL.canParse(result.url, env.APP_URL) + ) + return Response.json({ error: "invalid_provider_response" }, { status: 502 }); + const headers = new Headers(response.headers); + headers.delete("content-type"); + headers.delete("content-length"); + const target = new URL(result.url, env.APP_URL); + if (["javascript:", "data:", "vbscript:", "file:", "blob:"].includes(target.protocol)) { + return Response.json({ error: "invalid_provider_response" }, { status: 502 }); + } + if (target.origin === new URL(env.APP_URL).origin && target.pathname === "/api/auth/oauth") { + return redirectToOAuthLogin(target, true, headers); + } + headers.set("Location", result.url); + return new Response(null, { status: 302, headers }); } - const loginUrl = new URL("/auth/login", env.APP_URL); - const oauthParams = new URLSearchParams(); - for (const [key, value] of url.searchParams) { - if (!oauthCallbackPassthroughExcludedParams.has(key)) { - oauthParams.set(key, value); - } - } - loginUrl.searchParams.set("callbackURL", `/api/auth/oauth?${oauthParams.toString()}`); + return redirectToOAuthLogin(url); +} +function redirectToOAuthLogin(url: URL, reauthenticate = false, headers = new Headers()) { + const prompt = new Set(url.searchParams.get("prompt")?.split(" ") ?? []); + const loginUrl = new URL(prompt.has("create") ? "/auth/register" : "/auth/login", env.APP_URL); + if (reauthenticate) loginUrl.searchParams.set("reauthenticate", "true"); + loginUrl.searchParams.set("callbackURL", `/api/auth/oauth${url.search}`); + headers.set("Location", `${loginUrl.pathname}${loginUrl.search}`); return new Response(null, { status: 302, - headers: { Location: `${loginUrl.pathname}${loginUrl.search}` }, + headers, }); } diff --git a/apps/server/src/http/oauth-flow.integration.test.ts b/apps/server/src/http/oauth-flow.integration.test.ts new file mode 100644 index 000000000..94a07d9b6 --- /dev/null +++ b/apps/server/src/http/oauth-flow.integration.test.ts @@ -0,0 +1,254 @@ +import { createHash, randomBytes } from "node:crypto"; +import { describe, expect, it, vi } from "vitest"; + +vi.mock("@reactive-resume/email/transport", () => ({ sendEmail: vi.fn() })); + +// Run only against an explicitly supplied disposable database, after applying migrations. +const databaseURL = process.env.OAUTH_TEST_DATABASE_URL; + +describe.skipIf(!databaseURL)("MCP OAuth flow with PostgreSQL", () => { + it("registers public clients, resumes login, and exchanges a resource-bound PKCE code", async () => { + if (!databaseURL) return; + process.env.DATABASE_URL = databaseURL; + process.env.APP_URL = "http://localhost:33920"; + process.env.AUTH_SECRET = "oauth-integration-test-secret-only"; + const { handleAuth, handleOAuth } = await import("./auth"); + // Better Auth disables origin checks by default in test mode; exercise production behavior. + const { auth } = await import("@reactive-resume/auth/config"); + (await auth.$context).skipOriginCheck = false; + const origin = process.env.APP_URL; + const redirectURI = "http://127.0.0.1:33921/callback"; + const request = (path: string, body: object, cookie = "") => + new Request(`${origin}/api/auth/${path}`, { + method: "POST", + headers: { "content-type": "application/json", origin, cookie }, + body: JSON.stringify(body), + }); + const registration = await handleAuth( + request("oauth2/register", { client_name: "OAuth integration", redirect_uris: [redirectURI] }), + ); + expect(registration.status, await registration.clone().text()).toBe(201); + const client = await registration.json(); + expect(client.token_endpoint_auth_method).toBe("none"); + + const deniedRegistration = await handleAuth( + request("oauth2/register", { + client_name: "Denied resource", + redirect_uris: [redirectURI], + resources: ["https://untrusted.example/mcp"], + }), + ); + expect(deniedRegistration.status).toBe(400); + await expect(deniedRegistration.json()).resolves.toMatchObject({ error: "invalid_target" }); + + const verifier = randomBytes(32).toString("base64url"); + const query = new URLSearchParams({ + client_id: client.client_id, + redirect_uri: redirectURI, + response_type: "code", + scope: "openid profile offline_access", + code_challenge: createHash("sha256").update(verifier).digest("base64url"), + code_challenge_method: "S256", + resource: `${origin}/mcp`, + state: "opaque-state", + }); + const authorize = await handleAuth(new Request(`${origin}/api/auth/oauth2/authorize?${query}`)); + expect(authorize.status, await authorize.clone().text()).toBe(302); + const bridgeURL = authorize.headers.get("location"); + expect(bridgeURL).toBeTruthy(); + const login = await handleOAuth(new Request(new URL(bridgeURL ?? "", origin))); + const loginURL = new URL(login.headers.get("location") ?? "", origin); + const callbackURL = loginURL.searchParams.get("callbackURL"); + expect(callbackURL).toContain("sig="); + expect(callbackURL).toContain("resource="); + + const unique = randomBytes(6).toString("hex"); + const signup = await handleAuth( + request("sign-up/email", { + name: "OAuth Test", + email: `oauth-${unique}@example.com`, + username: `oauth-${unique}`, + password: "password123", + }), + ); + expect(signup.status, await signup.clone().text()).toBe(200); + const cookie = signup.headers + .getSetCookie() + .map((value) => value.split(";", 1)[0]) + .join("; "); + const tamperedURL = new URL(`${origin}${callbackURL}`); + tamperedURL.searchParams.set("state", "tampered"); + const tampered = await handleOAuth(new Request(tamperedURL, { headers: { cookie } })); + expect(tampered.status).toBe(400); + await expect(tampered.json()).resolves.toMatchObject({ error: "invalid_signature" }); + const callback = await handleOAuth(new Request(`${origin}${callbackURL}`, { headers: { cookie } })); + expect(callback.status, await callback.clone().text()).toBe(302); + const consentURL = new URL(callback.headers.get("location") ?? "", origin); + expect(consentURL.pathname).toBe("/auth/consent"); + expect(consentURL.searchParams.has("code")).toBe(false); + const oauth_query = consentURL.search.slice(1); + const consents = async () => { + const response = await handleAuth(new Request(`${origin}/api/auth/oauth2/get-consents`, { headers: { cookie } })); + expect(response.status).toBe(200); + return response.json(); + }; + expect(await consents()).toEqual([]); + const silent = await handleAuth( + new Request(`${origin}/api/auth/oauth2/authorize?${query}&prompt=none`, { headers: { cookie } }), + ); + expect(new URL(silent.headers.get("location") ?? "").searchParams.get("error")).toBe("consent_required"); + const tamperedConsentQuery = new URLSearchParams(oauth_query); + tamperedConsentQuery.set("scope", "openid profile email offline_access"); + const tamperedConsent = await handleAuth( + request( + "oauth2/consent", + { + accept: true, + oauth_query: tamperedConsentQuery.toString(), + }, + cookie, + ), + ); + expect(tamperedConsent.status).toBe(400); + expect(await consents()).toEqual([]); + const csrf = await handleAuth( + new Request(`${origin}/api/auth/oauth2/consent`, { + method: "POST", + headers: { cookie, origin: "https://untrusted.example", "content-type": "application/json" }, + body: JSON.stringify({ accept: true, oauth_query }), + }), + ); + expect(csrf.status).toBe(403); + const denied = await handleAuth(request("oauth2/consent", { accept: false, oauth_query }, cookie)); + expect(denied.status, await denied.clone().text()).toBe(200); + const deniedURL = new URL((await denied.json()).url); + expect(deniedURL.searchParams.get("error")).toBe("access_denied"); + expect(deniedURL.searchParams.get("state")).toBe("opaque-state"); + expect(deniedURL.searchParams.has("code")).toBe(false); + expect(await consents()).toEqual([]); + const accepted = await handleAuth(request("oauth2/consent", { accept: true, oauth_query }, cookie)); + expect(accepted.status, await accepted.clone().text()).toBe(200); + expect(await consents()).toHaveLength(1); + const codeURL = new URL((await accepted.json()).url); + expect(codeURL.origin).toBe(new URL(redirectURI).origin); + expect(codeURL.searchParams.get("state")).toBe("opaque-state"); + const code = codeURL.searchParams.get("code"); + expect(code).toBeTruthy(); + const tokenRequest = () => + new Request(`${origin}/api/auth/oauth2/token`, { + method: "POST", + headers: { "content-type": "application/x-www-form-urlencoded" }, + body: new URLSearchParams({ + grant_type: "authorization_code", + client_id: client.client_id, + code: code ?? "", + redirect_uri: redirectURI, + code_verifier: verifier, + resource: `${origin}/mcp`, + }), + }); + const tokenResponse = await handleAuth(tokenRequest()); + expect(tokenResponse.status, await tokenResponse.clone().text()).toBe(200); + const token = await tokenResponse.json(); + expect(token.access_token).toBeTruthy(); + expect(token.refresh_token).toBeTruthy(); + const claims = JSON.parse(Buffer.from(token.access_token.split(".")[1], "base64url").toString()); + expect([claims.aud].flat()).toContain(`${origin}/mcp`); + expect((await handleAuth(tokenRequest())).status).toBe(400); + }, 30_000); + it.each(["login", "max-age", "create"])( + "requires fresh authentication for %s without looping", + async (mode) => { + if (!databaseURL) return; + process.env.DATABASE_URL = databaseURL; + process.env.APP_URL = "http://localhost:33920"; + process.env.AUTH_SECRET = "oauth-integration-test-secret-only"; + const { handleAuth, handleOAuth } = await import("./auth"); + const origin = process.env.APP_URL; + const cookieOf = (response: Response) => + response.headers + .getSetCookie() + .map((value) => value.split(";", 1)[0]) + .join("; "); + const post = (path: string, body: object, cookie = "") => + handleAuth( + new Request(`${origin}/api/auth/${path}`, { + method: "POST", + headers: { "content-type": "application/json", origin, cookie }, + body: JSON.stringify(body), + }), + ); + const unique = randomBytes(6).toString("hex"); + const credentials = { + name: "Reauth Test", + email: `reauth-${unique}@example.com`, + username: `reauth-${unique}`, + password: "password123", + }; + const existingSignup = await post("sign-up/email", credentials); + expect(existingSignup.status).toBe(200); + const oldCookie = cookieOf(existingSignup); + const registration = await post("oauth2/register", { + client_name: "Reauth integration", + redirect_uris: ["http://127.0.0.1:33921/callback"], + }); + expect(registration.status).toBe(201); + const client = await registration.json(); + const query = new URLSearchParams({ + client_id: client.client_id, + redirect_uri: "http://127.0.0.1:33921/callback", + response_type: "code", + scope: "openid profile", + resource: `${origin}/mcp`, + code_challenge: createHash("sha256").update(randomBytes(32)).digest("base64url"), + code_challenge_method: "S256", + ...(mode === "max-age" ? { max_age: "0" } : { prompt: mode }), + }); + const authorization = await handleAuth( + new Request(`${origin}/api/auth/oauth2/authorize?${query}`, { headers: { cookie: oldCookie } }), + ); + expect(authorization.status).toBe(302); + const bridge = await handleOAuth( + new Request(new URL(authorization.headers.get("location") ?? "", origin), { headers: { cookie: oldCookie } }), + ); + expect(bridge.status).toBe(302); + const loginURL = new URL(bridge.headers.get("location") ?? "", origin); + expect(loginURL.pathname).toBe(mode === "create" ? "/auth/register" : "/auth/login"); + expect(loginURL.searchParams.get("reauthenticate")).toBe("true"); + const callbackURL = new URL(loginURL.searchParams.get("callbackURL") ?? "", origin); + const oauth_query = callbackURL.search.slice(1); + const authenticated = + mode === "create" + ? await post( + "sign-up/email", + { ...credentials, email: `new-${unique}@example.com`, username: `new-${unique}` }, + oldCookie, + ) + : await post( + "sign-in/email", + { email: credentials.email, password: credentials.password, oauth_query }, + oldCookie, + ); + expect(authenticated.status, await authenticated.clone().text()).toBe(200); + const newCookie = cookieOf(authenticated); + expect(newCookie).not.toBe(oldCookie); + const continuation = + mode === "create" ? await post("oauth2/continue", { created: true, oauth_query }, newCookie) : authenticated; + expect(continuation.status, await continuation.clone().text()).toBe(200); + const result = await continuation.json(); + let target = new URL(result.url, origin); + if (target.pathname === "/api/auth/oauth") { + const response = await handleOAuth(new Request(target, { headers: { cookie: newCookie } })); + expect(response.status, await response.clone().text()).toBe(302); + target = new URL(response.headers.get("location") ?? "", origin); + } + expect(target.pathname).toBe("/auth/consent"); + const accepted = await post("oauth2/consent", { accept: true, oauth_query: target.search.slice(1) }, newCookie); + expect(accepted.status, await accepted.clone().text()).toBe(200); + target = new URL((await accepted.json()).url, origin); + expect(target.origin).toBe("http://127.0.0.1:33921"); + expect(target.searchParams.get("code")).toBeTruthy(); + }, + 30_000, + ); +}); diff --git a/apps/server/src/index.test.ts b/apps/server/src/index.test.ts new file mode 100644 index 000000000..094ce6158 --- /dev/null +++ b/apps/server/src/index.test.ts @@ -0,0 +1,35 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; + +const events = vi.hoisted(() => [] as string[]); +vi.mock("./startup/checks", () => ({ + runStartupChecks: async () => { + await Promise.resolve(); + events.push("migrations complete"); + }, +})); +vi.mock("./http/app", () => { + events.push("auth imported"); + return { + createApp: () => { + events.push("app created"); + return { fetch: vi.fn() }; + }, + }; +}); +vi.mock("@hono/node-server", () => ({ + serve: () => { + events.push("server listening"); + }, +})); +vi.mock("@reactive-resume/env/server", () => ({ env: { SERVER_PORT: 3001 } })); +afterEach(() => vi.restoreAllMocks()); + +describe("server startup", () => { + it("finishes migrations before importing auth and seeding OAuth resources", async () => { + vi.spyOn(process, "on").mockReturnValue(process); + const entry = await import("./index"); + expect(events).toEqual([]); + await entry.main(); + expect(events).toEqual(["migrations complete", "auth imported", "app created", "server listening"]); + }); +}); diff --git a/apps/server/src/index.ts b/apps/server/src/index.ts index dc0dfc48c..b13631919 100644 --- a/apps/server/src/index.ts +++ b/apps/server/src/index.ts @@ -1,14 +1,15 @@ import { pathToFileURL } from "node:url"; import { serve } from "@hono/node-server"; import { env } from "@reactive-resume/env/server"; -import { createApp } from "./http/app"; import { runStartupChecks } from "./startup/checks"; -export { createApp } from "./http/app"; - -async function main() { +export async function main() { await runStartupChecks(); + // OAuth resource seeding starts when auth is imported, so load the app only + // after migrations have created the provider tables. + const { createApp } = await import("./http/app"); + // Safety net: Node 24 crashes the whole process on an unhandled rejection. One request's // stray promise must not take the server down for everyone, so log and keep serving. // Registered after startup checks so a broken startup still fails loudly. (Left uncaught diff --git a/apps/server/src/openapi/generator.test.ts b/apps/server/src/openapi/generator.test.ts index 491132cb4..a4cb26dd8 100644 --- a/apps/server/src/openapi/generator.test.ts +++ b/apps/server/src/openapi/generator.test.ts @@ -1,9 +1,14 @@ -import { describe, expect, it } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import z from "zod"; import { defaultResumeData } from "@reactive-resume/schema/resume/default"; import { createResumeDataJsonSchema } from "@reactive-resume/schema/resume/json-schema"; import { writableResumeDataSchema } from "@reactive-resume/schema/resume/write"; +// Spec generation reads procedure contracts without executing authentication. Keep the +// provider's resource seeding out of this unit test; real OAuth initialization is covered +// by the opt-in PostgreSQL integration suite after migrations run. +vi.mock("@reactive-resume/auth/config", () => ({ auth: {}, verifyOAuthToken: vi.fn() })); + type GeneratedSpecView = { components?: { schemas?: Record }; paths?: Record< diff --git a/apps/server/turbo.json b/apps/server/turbo.json index 3cbfc7ade..6b0b91d68 100644 --- a/apps/server/turbo.json +++ b/apps/server/turbo.json @@ -2,8 +2,12 @@ "extends": ["//"], "tags": ["app:server", "runtime:server", "role:adapter"], "tasks": { + "test": { "env": ["OAUTH_TEST_DATABASE_URL"] }, + "test:coverage": { "env": ["OAUTH_TEST_DATABASE_URL"] }, + "test:agent": { "env": ["OAUTH_TEST_DATABASE_URL"] }, "test:ci": { - "cache": false + "cache": false, + "env": ["OAUTH_TEST_DATABASE_URL"] } } } diff --git a/apps/web/locales/en-US.po b/apps/web/locales/en-US.po index 7f29c4beb..de80be782 100644 --- a/apps/web/locales/en-US.po +++ b/apps/web/locales/en-US.po @@ -234,6 +234,10 @@ msgstr "A web address is written without https://." msgid "A4" msgstr "A4" +#: src/features/auth/pages/consent.tsx +msgid "Access your account through the API, including reading and changing your resumes and job applications." +msgstr "Access your account through the API, including reading and changing your resumes and job applications." + #: src/routes/_home/-sections/features.tsx msgid "Access your resumes and data programmatically using the API." msgstr "Access your resumes and data programmatically using the API." @@ -523,6 +527,10 @@ msgstr "Albanian" msgid "All applications (including archived)" msgstr "All applications (including archived)" +#: src/features/auth/pages/consent.tsx +msgid "Allow access" +msgstr "Allow access" + #: src/routes/builder/$resumeId/-sidebar/right/sections/sharing.tsx msgid "Allow Public Access" msgstr "Allow Public Access" @@ -1093,6 +1101,10 @@ msgstr "Clear selection" msgid "Click here to select a file to import" msgstr "Click here to select a file to import" +#: src/features/auth/pages/consent.tsx +msgid "Client ID" +msgstr "Client ID" + #: src/routes/agent/-components/agent-chat.tsx msgid "Close AI assistant" msgstr "Close AI assistant" @@ -1172,6 +1184,10 @@ msgstr "Confirm Password" msgid "Connect" msgstr "Connect" +#: src/features/auth/pages/consent.tsx +msgid "Connect an application" +msgstr "Connect an application" + #: src/features/ats-checker/ai-review/ai-review-card.tsx msgid "Connect your own AI provider to get a review of the writing. The checks above need no provider and no account." msgstr "Connect your own AI provider to get a review of the writing. The checks above need no provider and no account." @@ -1295,6 +1311,10 @@ msgstr "Correct the year, or use \"Present\" for ongoing work." msgid "Correct the year, or write \"Present\" for ongoing work." msgstr "Correct the year, or write \"Present\" for ongoing work." +#: src/features/auth/pages/consent.tsx +msgid "Could not complete this connection. Restart the connection from your client and try again." +msgstr "Could not complete this connection. Restart the connection from your client and try again." + #: src/features/resume/export/use-resume-export.ts msgid "Could not generate the DOCX. Please try again." msgstr "Could not generate the DOCX. Please try again." @@ -1722,6 +1742,7 @@ msgstr "Deleting your resume..." msgid "Denied, waiting for the agent…" msgstr "Denied, waiting for the agent…" +#: src/features/auth/pages/consent.tsx #: src/routes/agent/-components/patch-approval-card.tsx msgid "Deny" msgstr "Deny" @@ -2925,6 +2946,10 @@ msgstr "Kannada" msgid "Keep a single address, so software does not have to guess which one to use." msgstr "Keep a single address, so software does not have to guess which one to use." +#: src/features/auth/pages/consent.tsx +msgid "Keep access when you are not using the application." +msgstr "Keep access when you are not using the application." + #: src/features/ats-checker/messages.ts msgid "Keep it under 2.5 MB, usually by removing or shrinking images." msgstr "Keep it under 2.5 MB, usually by removing or shrinking images." @@ -3158,6 +3183,10 @@ msgstr "Loading AI providers. Please try again in a moment." msgid "Loading applications…" msgstr "Loading applications…" +#: src/features/auth/pages/consent.tsx +msgid "Loading connection request..." +msgstr "Loading connection request..." + #: src/features/settings/integrations/components/ai-provider-picker.tsx #: src/features/settings/integrations/components/ai-section.tsx msgid "Loading providers…" @@ -3630,6 +3659,10 @@ msgstr "One-Click Sign-In" msgid "Ongoing Maintenance" msgstr "Ongoing Maintenance" +#: src/features/auth/pages/consent.tsx +msgid "Only allow applications you trust. This application will be able to:" +msgstr "Only allow applications you trust. This application will be able to:" + #. Helper note explaining the keep-together limitation #: src/routes/builder/$resumeId/-sidebar/right/sections/layout/pages.tsx msgid "Only applies when the section fits on a single page." @@ -4173,6 +4206,14 @@ msgstr "Read the Applying Custom Styles guide." msgid "Read the resume" msgstr "Read the resume" +#: src/features/auth/pages/consent.tsx +msgid "Read your email address." +msgstr "Read your email address." + +#: src/features/auth/pages/consent.tsx +msgid "Read your profile information." +msgstr "Read your profile information." + #: src/features/ats-checker/messages.ts msgid "Readability" msgstr "Readability" @@ -4881,6 +4922,10 @@ msgstr "Sign up" msgid "Sign-in didn't complete" msgstr "Sign-in didn't complete" +#: src/features/auth/pages/consent.tsx +msgid "Signed in as {email}" +msgstr "Signed in as {email}" + #: src/features/auth/components/social-auth.tsx #: src/features/auth/pages/login.tsx msgid "Signing in..." @@ -5495,6 +5540,10 @@ msgstr "This check does not apply to this file." msgid "This check only runs on resumes written in English." msgstr "This check only runs on resumes written in English." +#: src/features/auth/pages/consent.tsx +msgid "This connection request is invalid or has expired." +msgstr "This connection request is invalid or has expired." + #: src/dialogs/api-key/create.tsx msgid "This creates a new API key, which lets other programs read and change your resume data through the Reactive Resume API." msgstr "This creates a new API key, which lets other programs read and change your resume data through the Reactive Resume API." diff --git a/apps/web/src/features/auth/components/social-auth.tsx b/apps/web/src/features/auth/components/social-auth.tsx index eb03c9af8..c66acc933 100644 --- a/apps/web/src/features/auth/components/social-auth.tsx +++ b/apps/web/src/features/auth/components/social-auth.tsx @@ -3,13 +3,14 @@ import { t } from "@lingui/core/macro"; import { Trans } from "@lingui/react/macro"; import { FingerprintIcon, GithubLogoIcon, GoogleLogoIcon, LinkedinLogoIcon, VaultIcon } from "@phosphor-icons/react"; import { useQuery } from "@tanstack/react-query"; -import { useRouter } from "@tanstack/react-router"; +import { useRouter, useSearch } from "@tanstack/react-router"; import { Button } from "@reactive-resume/ui/components/button"; import { Skeleton } from "@reactive-resume/ui/components/skeleton"; import { toast } from "@reactive-resume/ui/components/toast"; import { cn } from "@reactive-resume/utils/style"; import { authClient } from "@/libs/auth/client"; import { orpc } from "@/libs/orpc/client"; +import { getAuthRedirectOptions, getOAuthPasskeyOptions, getOAuthSignInOptions, isOAuthRedirect } from "../redirect"; export function SocialAuth() { const { data: providers = {}, isLoading } = useQuery(orpc.auth.providers.list.queryOptions()); @@ -48,10 +49,14 @@ type SocialAuthButtonsProps = { function SocialAuthButtons({ providers }: SocialAuthButtonsProps) { const router = useRouter(); + const { callbackURL } = useSearch({ from: "/auth" }); - const runSignIn = async (fn: () => Promise<{ error: { message?: string } | null }>) => { + const runSignIn = async ( + fn: () => Promise<{ data?: unknown; error: { message?: string } | null }>, + isPasskey = false, + ) => { const toastId = toast.add({ type: "loading", description: t`Signing in...` }); - const { error } = await fn(); + const { data, error } = await fn(); if (error) { toast.add({ type: "error", @@ -66,7 +71,9 @@ function SocialAuthButtons({ providers }: SocialAuthButtonsProps) { return; } toast.close(toastId); + if (isOAuthRedirect(data)) return; await router.invalidate(); + if (isPasskey) void router.navigate(getAuthRedirectOptions(callbackURL)); }; return ( @@ -77,7 +84,8 @@ function SocialAuthButtons({ providers }: SocialAuthButtonsProps) { runSignIn(() => authClient.signIn.social({ provider: "custom", - callbackURL: "/dashboard", + callbackURL: callbackURL ?? "/dashboard", + ...getOAuthSignInOptions(callbackURL), }), ) } @@ -89,7 +97,9 @@ function SocialAuthButtons({ providers }: SocialAuthButtonsProps) { + + + + ) : null} + + ); +} diff --git a/apps/web/src/features/auth/pages/login.tsx b/apps/web/src/features/auth/pages/login.tsx index 128817377..9ad9ae84c 100644 --- a/apps/web/src/features/auth/pages/login.tsx +++ b/apps/web/src/features/auth/pages/login.tsx @@ -2,7 +2,7 @@ import { t } from "@lingui/core/macro"; import { Trans } from "@lingui/react/macro"; import { ArrowRightIcon, EyeIcon, EyeSlashIcon } from "@phosphor-icons/react"; import { useQuery } from "@tanstack/react-query"; -import { Link, useNavigate, useRouter } from "@tanstack/react-router"; +import { Link, useNavigate, useRouter, useSearch } from "@tanstack/react-router"; import { useEffect, useRef } from "react"; import { useToggle } from "usehooks-ts"; import z from "zod"; @@ -14,6 +14,7 @@ import { authClient } from "@/libs/auth/client"; import { orpc } from "@/libs/orpc/client"; import { useAppForm } from "@/libs/tanstack-form"; import { SocialAuth } from "../components/social-auth"; +import { getAuthRedirectOptions, getOAuthPasskeyOptions, getOAuthSignInOptions, isOAuthRedirect } from "../redirect"; const formSchema = z.object({ identifier: z.string().trim().toLowerCase(), @@ -27,6 +28,7 @@ type Props = { export function LoginPage({ disableEmailAuth, disableSignups }: Props) { const router = useRouter(); + const { callbackURL, reauthenticate } = useSearch({ from: "/auth" }); const navigate = useNavigate(); const hasStartedConditionalPasskeyRef = useRef(false); @@ -44,8 +46,16 @@ export function LoginPage({ disableEmailAuth, disableSignups }: Props) { const isEmail = value.identifier.includes("@"); const result = isEmail - ? await authClient.signIn.email({ email: value.identifier, password: value.password }) - : await authClient.signIn.username({ username: value.identifier, password: value.password }); + ? await authClient.signIn.email({ + email: value.identifier, + password: value.password, + ...getOAuthSignInOptions(callbackURL), + }) + : await authClient.signIn.username({ + username: value.identifier, + password: value.password, + ...getOAuthSignInOptions(callbackURL), + }); if (result.error) { toast.add({ @@ -69,13 +79,14 @@ export function LoginPage({ disableEmailAuth, disableSignups }: Props) { if (requiresTwoFactor) { toast.close(toastId); - void navigate({ to: "/auth/verify-2fa", replace: true }); + void navigate({ to: "/auth/verify-2fa", search: { callbackURL, reauthenticate }, replace: true }); return; } toast.close(toastId); + if (isOAuthRedirect(result.data)) return; await router.invalidate(); - void navigate({ to: "/dashboard", replace: true }); + void navigate(getAuthRedirectOptions(callbackURL)); } catch { toast.add({ type: "error", description: t`Failed to sign in. Please try again.`, id: toastId }); } @@ -94,12 +105,16 @@ export function LoginPage({ disableEmailAuth, disableSignups }: Props) { void PublicKeyCredential.isConditionalMediationAvailable().then(async (isAvailable) => { if (!isAvailable) return; - const { error } = await authClient.signIn.passkey({ autoFill: true }); - if (error) return; + const { data, error } = await authClient.signIn.passkey({ + autoFill: true, + ...getOAuthPasskeyOptions(callbackURL), + }); + if (error || isOAuthRedirect(data)) return; await router.invalidate(); + void navigate(getAuthRedirectOptions(callbackURL)); }); - }, [providers, router]); + }, [providers, router, navigate, callbackURL]); return ( <> @@ -117,7 +132,7 @@ export function LoginPage({ disableEmailAuth, disableSignups }: Props) { nativeButton={false} className="h-auto gap-1.5 px-1! py-0" render={ - + Create one now {" "} diff --git a/apps/web/src/features/auth/pages/register.tsx b/apps/web/src/features/auth/pages/register.tsx index 96205f403..eadaa9a96 100644 --- a/apps/web/src/features/auth/pages/register.tsx +++ b/apps/web/src/features/auth/pages/register.tsx @@ -1,7 +1,7 @@ import { t } from "@lingui/core/macro"; import { Trans } from "@lingui/react/macro"; import { ArrowRightIcon, EyeIcon, EyeSlashIcon } from "@phosphor-icons/react"; -import { Link } from "@tanstack/react-router"; +import { Link, useSearch } from "@tanstack/react-router"; import { useState } from "react"; import { useToggle } from "usehooks-ts"; import z from "zod"; @@ -13,6 +13,7 @@ import { toast } from "@reactive-resume/ui/components/toast"; import { authClient } from "@/libs/auth/client"; import { useAppForm } from "@/libs/tanstack-form"; import { SocialAuth } from "../components/social-auth"; +import { getOAuthSignInOptions, isOAuthRedirect } from "../redirect"; const formSchema = z.object({ name: z.string().min(3).max(64), @@ -34,6 +35,7 @@ type Props = { }; export function RegisterPage({ disableEmailAuth }: Props) { + const { callbackURL, reauthenticate } = useSearch({ from: "/auth" }); const [submitted, setSubmitted] = useState(false); const [showPassword, toggleShowPassword] = useToggle(false); @@ -43,13 +45,16 @@ export function RegisterPage({ disableEmailAuth }: Props) { onSubmit: async ({ value }) => { const toastId = toast.add({ type: "loading", description: t`Signing up...` }); - const { error } = await authClient.signUp.email({ + const oauthOptions = getOAuthSignInOptions(callbackURL); + const createPrompt = new URLSearchParams(oauthOptions.oauth_query).get("prompt")?.split(" ").includes("create"); + const { data, error } = await authClient.signUp.email({ name: value.name, email: value.email, password: value.password, username: value.username, displayUsername: value.username, - callbackURL: "/dashboard", + callbackURL: callbackURL ?? "/dashboard", + ...(!createPrompt ? oauthOptions : {}), }); if (error) { @@ -66,6 +71,15 @@ export function RegisterPage({ disableEmailAuth }: Props) { return; } + if (isOAuthRedirect(data)) return; + if (createPrompt && oauthOptions.oauth_query) { + const continuation = await authClient.oauth2.continue({ created: true, oauth_query: oauthOptions.oauth_query }); + if (continuation.error) { + toast.add({ type: "error", description: continuation.error.message, id: toastId }); + return; + } + if (isOAuthRedirect(continuation.data)) return; + } setSubmitted(true); toast.close(toastId); }, @@ -88,7 +102,7 @@ export function RegisterPage({ disableEmailAuth }: Props) { nativeButton={false} className="h-auto gap-1.5 px-1! py-0" render={ - + Sign in now{" "} @@ -250,6 +264,7 @@ export function RegisterPage({ disableEmailAuth }: Props) { } function PostSignupScreen() { + const { callbackURL } = useSearch({ from: "/auth" }); return ( <>
@@ -273,10 +288,10 @@ function PostSignupScreen() {