diff --git a/app/(authenticated)/link/_tests/page.test.tsx b/app/(authenticated)/link/_tests/page.test.tsx new file mode 100644 index 00000000..e05cab14 --- /dev/null +++ b/app/(authenticated)/link/_tests/page.test.tsx @@ -0,0 +1,77 @@ +import { renderToStaticMarkup } from "react-dom/server"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import type { getLinkAccount as readLinkAccount } from "@db/services/auth/link"; +import Page from "../page"; + +const { getLinkAccount } = vi.hoisted(() => ({ + getLinkAccount: vi.fn(), +})); + +vi.mock("@db/services/auth/link", () => ({ + getLinkAccount, + linkConfigured: () => true, +})); +vi.mock("@web/auth/request-scope", () => ({ + requireRequestScope: () => ({ userId: "better-auth:phone-user" }), +})); + +beforeEach(() => { + getLinkAccount.mockResolvedValue({ id: "wallet-account" }); +}); + +describe("Link wallet recovery", () => { + it("offers reauthentication and returns to the pending wallet request", async () => { + const attempt = "6738a4d4-c973-42f4-9e4a-e8339c09687f"; + const html = renderToStaticMarkup( + await Page({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ + error: "reauthentication_required", + attempt, + }), + }) + ); + + expect(html).toContain("Sign in again before disconnecting your wallet."); + expect(html).toContain( + `href="/sign-in?reauthenticate=true&callbackUrl=%2Flink%3Fattempt%3D${attempt}"` + ); + expect(html).not.toContain("Disconnect wallet"); + }); + + it("offers a retry for failed revocation without requesting a new sign-in", async () => { + const html = renderToStaticMarkup( + await Page({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ error: "disconnection_failed" }), + }) + ); + + expect(html).toContain("Your wallet could not be disconnected."); + expect(html).toContain("It is still connected. Try again shortly."); + expect(html).toContain("Disconnect wallet"); + expect(html).not.toContain("Sign in again"); + }); + + it.each(["disconnection_failed", "reauthentication_required"])( + "shows the current disconnected state for a stale %s error", + async (code) => { + getLinkAccount.mockResolvedValue(undefined); + const html = renderToStaticMarkup( + await Page({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ error: code }), + }) + ); + + expect(html).toContain("Connect Link"); + expect(html).not.toContain("Disconnect wallet"); + expect(html).toContain("Your wallet is no longer connected."); + expect(html).not.toContain("It is still connected."); + expect(html).not.toContain("Sign in again"); + expect(html).not.toContain( + "The wallet connection could not be completed" + ); + } + ); +}); diff --git a/app/(authenticated)/link/page.tsx b/app/(authenticated)/link/page.tsx index fc877767..0501957d 100644 --- a/app/(authenticated)/link/page.tsx +++ b/app/(authenticated)/link/page.tsx @@ -21,6 +21,12 @@ export default async function Page({ searchParams }: PageProps<"/link">) { const connected = configured && Boolean(await getLinkAccount(scope.userId.slice("better-auth:".length))); + const reauthenticationRequired = + connected && params.error === "reauthentication_required"; + const disconnectError = + params.error === "reauthentication_required" || + params.error === "disconnection_failed"; + const returnUrl = attempt.success ? `/link?attempt=${attempt.data}` : "/link"; return (
@@ -93,9 +99,13 @@ export default async function Page({ searchParams }: PageProps<"/link">) { )} {params.error && (

- The wallet connection could not be completed. Try again, or sign in - again before disconnecting. To use a different wallet, disconnect the - current wallet first. + {disconnectError && !connected + ? "Your wallet is no longer connected." + : reauthenticationRequired + ? "Sign in again before disconnecting your wallet. Your wallet is still connected." + : disconnectError + ? "Your wallet could not be disconnected. It is still connected. Try again shortly." + : "The wallet connection could not be completed. Try again. To use a different wallet, disconnect the current wallet first."}

)} {configured && ( @@ -111,13 +121,27 @@ export default async function Page({ searchParams }: PageProps<"/link">) { )} - {connected && ( -
- - -
+ {reauthenticationRequired ? ( + + ) : ( + connected && ( +
+ + +
+ ) )}
)} diff --git a/app/api/link/route.ts b/app/api/link/route.ts index c260fa5f..1999475c 100644 --- a/app/api/link/route.ts +++ b/app/api/link/route.ts @@ -1,4 +1,5 @@ import { z } from "zod"; +import { APIError } from "better-auth/api"; import { getAuth } from "@db/services/auth"; import { getAuthSession } from "@db/services/auth/session"; import { @@ -77,8 +78,14 @@ export async function POST(request: Request) { status: 303, headers: { ...privateHeaders, location: "/link" }, }); - } catch { - return connectionFailed(input.data.attempt); + } catch (error) { + return connectionFailed( + input.data.attempt, + error instanceof APIError && + (error.body?.code === "SESSION_NOT_FRESH" || error.statusCode === 401) + ? "reauthentication_required" + : "disconnection_failed" + ); } } return connectLink(session.user.id, request.headers, input.data.attempt); @@ -114,9 +121,9 @@ async function connectLink( } } -function connectionFailed(attempt?: string) { +function connectionFailed(attempt?: string, error = "connection_failed") { const destination = new URL("/link", applicationOrigin()); - destination.searchParams.set("error", "connection_failed"); + destination.searchParams.set("error", error); if (attempt) destination.searchParams.set("attempt", attempt); return new Response(null, { status: 303, diff --git a/app/sign-in/_tests/page.test.tsx b/app/sign-in/_tests/page.test.tsx new file mode 100644 index 00000000..c8493a4b --- /dev/null +++ b/app/sign-in/_tests/page.test.tsx @@ -0,0 +1,63 @@ +import { createElement } from "react"; +import { renderToStaticMarkup } from "react-dom/server"; +import { describe, expect, it, vi } from "vitest"; +import SignInPage from "../page"; + +vi.mock("next/headers", () => ({ headers: () => new Headers() })); +vi.mock("next/navigation", () => ({ + redirect: (url: string) => { + throw new Error(`Redirect: ${url}`); + }, +})); +vi.mock("@db/services/auth/session", () => ({ + getAuthSession: () => ({ user: { id: "phone-user" } }), +})); +vi.mock("@shared/environment", () => ({ + env: {}, + localPhoneAuthBypassEnabled: true, +})); +vi.mock("@app/sign-in/_components/local-form", () => ({ + LocalPhoneAuthForm: ({ callbackUrl }: { readonly callbackUrl: string }) => + createElement("form", { action: callbackUrl }), +})); + +describe("phone sign-in reauthentication", () => { + it("allows an already signed-in user to verify again and return to Link", async () => { + const html = renderToStaticMarkup( + await SignInPage({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ + reauthenticate: "true", + callbackUrl: "/link", + }), + }) + ); + + expect(html).toContain("Sign in again"); + expect(html).toContain('action="/link"'); + }); + + it("still redirects signed-in users outside the reauthentication flow", async () => { + await expect( + SignInPage({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ callbackUrl: "/link" }), + }) + ).rejects.toThrow("Redirect: /"); + }); + + it("keeps the reauthentication callback on this app", async () => { + const html = renderToStaticMarkup( + await SignInPage({ + params: Promise.resolve({}), + searchParams: Promise.resolve({ + reauthenticate: "true", + callbackUrl: "//attacker.example", + }), + }) + ); + + expect(html).toContain('action="/"'); + expect(html).not.toContain("attacker.example"); + }); +}); diff --git a/app/sign-in/page.tsx b/app/sign-in/page.tsx index 46d19f33..f783789e 100644 --- a/app/sign-in/page.tsx +++ b/app/sign-in/page.tsx @@ -9,9 +9,11 @@ import { readLinqOnboardingPhoneNumber } from "@db/services/auth/linq"; export default async function SignInPage({ searchParams, }: PageProps<"/sign-in">) { - if (await getAuthSession(await headers())) redirect("/"); + const params = await searchParams; + const reauthenticate = params.reauthenticate === "true"; + if (!reauthenticate && (await getAuthSession(await headers()))) redirect("/"); - const callbackValue = (await searchParams).callbackUrl; + const callbackValue = params.callbackUrl; const requestedCallback = Array.isArray(callbackValue) ? callbackValue[0] : callbackValue; @@ -30,7 +32,9 @@ export default async function SignInPage({
-

Sign In

+

+ {reauthenticate ? "Sign in again" : "Sign In"} +

Enter your phone number to sign in.

diff --git a/tests/integration/link.test.ts b/tests/integration/link.test.ts index 98dc2810..716cd93d 100644 --- a/tests/integration/link.test.ts +++ b/tests/integration/link.test.ts @@ -333,12 +333,52 @@ describe("Link wallet integration", () => { }) ).rejects.toMatchObject({ cause: { code: "23505" } }); - await expect(link.disconnectLink(userId, headers)).rejects.toThrow( - "Unable to revoke Link access" + // A valid session older than the freshness window needs a new phone sign-in. + await database + .update(schema.session) + .set({ createdAt: new Date(Date.now() - 25 * 60 * 60_000) }) + .where(eq(schema.session.userId, userId)); + expect((await auth.api.getSession({ headers }))?.user.id).toBe(userId); + const disconnectRequest = () => + new Request("http://localhost:3000/api/link", { + method: "POST", + headers, + body: new URLSearchParams({ operation: "disconnect" }), + }); + const staleDisconnect = await POST(disconnectRequest()); + expect(staleDisconnect.status).toBe(303); + expect( + new URL(staleDisconnect.headers.get("location") ?? "").searchParams.get( + "error" + ) + ).toBe("reauthentication_required"); + expect(revokedTokens).toEqual([]); + expect(await link.getLinkAccount(userId)).toBeDefined(); + + const reauthenticated = await auth.api.verifyPhoneNumber({ + body: { phoneNumber: "+12025550123", code: "123456" }, + headers, + returnHeaders: true, + }); + expect(reauthenticated.response.user.id).toBe(userId); + expect(reauthenticated.response.token).not.toBe(signedIn.response.token); + updateCookies(headers, reauthenticated.headers); + const failedDisconnect = await POST(disconnectRequest()); + expect(failedDisconnect.status).toBe(303); + expect( + new URL( + failedDisconnect.headers.get("location") ?? "" + ).searchParams.get("error") + ).toBe("disconnection_failed"); + expect(failedDisconnect.headers.get("cache-control")).toBe("no-store"); + expect(failedDisconnect.headers.get("referrer-policy")).toBe( + "no-referrer" ); expect(await link.getLinkAccount(userId)).toBeDefined(); revokeSucceeds = true; - await link.disconnectLink(userId, headers); + const disconnected = await POST(disconnectRequest()); + expect(disconnected.status).toBe(303); + expect(disconnected.headers.get("location")).toBe("/link"); expect(await link.getLinkAccount(userId)).toBeUndefined(); expect(revokedTokens).toEqual(["refresh-rotated", "refresh-rotated"]); expect((await auth.api.getSession({ headers }))?.user.id).toBe(userId);