Skip to content
File

Blob: src/worker/api/public/oidc.ts

typescript367 lines
1import { eq } from "drizzle-orm";
2import { DrizzleQueryError } from "drizzle-orm/errors";
3import { deleteCookie, getSignedCookie, setSignedCookie } from "hono/cookie";
4import * as oidc from "openid-client";
5 
6import { OwnerSlug } from "@/contracts";
7import { createSession } from "@/worker/auth/sessions";
8import {
9 OIDC_SCOPE,
10 OIDC_TX_COOKIE_NAME,
11 OIDC_TX_COOKIE_MAX_AGE_SECONDS,
12 appendOidcMarker,
13 decodeTxCookiePayload,
14 encodeTxCookiePayload,
15 exchangeAuthorizationCode,
16 getOidcConfig,
17 getTxCookieSecret,
18 oidcErrorContext,
19 sanitizeReturnTo,
20 validateClaims,
21 type ResolvedClaims,
22 type TxCookiePayload,
23} from "@/worker/auth/oidc";
24import { setSessionCookie } from "@/worker/auth/cookies";
25import type { AppContext } from "@/worker/hono";
26import { type D1DbExecutor, tesseraIdentities, users } from "@/worker/db/d1";
27import { findIdentityByUserId, findUserByEmail, findUserIdBySub, insertIdentity } from "@/worker/db/d1/repositories";
28import { createLogger, generateDurableEntityId } from "@/worker/services";
29import { normalizeDisplayName } from "@/worker/validation";
30 
31const logger = createLogger("worker.oidc");
32const TX_COOKIE_PREFIX = "host" as const;
33 
34const TX_COOKIE_OPTIONS = {
35 path: "/" as const,
36 httpOnly: true,
37 secure: true,
38 sameSite: "Lax" as const,
39 maxAge: OIDC_TX_COOKIE_MAX_AGE_SECONDS,
40 prefix: TX_COOKIE_PREFIX,
41};
42 
43const TX_COOKIE_CLEAR_OPTIONS = {
44 path: "/" as const,
45 secure: true,
46 prefix: TX_COOKIE_PREFIX,
47};
48 
49type BindSuccessKind = "signed_in" | "email_updated" | "bound_legacy" | "created" | "raced_to_existing";
50type BindFailureCode = "tessera_email_conflict" | "identity_conflict" | "user_disabled";
51 
52type BindOutcome = { ok: true; kind: BindSuccessKind; userId: string } | { ok: false; code: BindFailureCode };
53 
54interface ExistingIdentityUser {
55 userId: string;
56 userEmail: string;
57 disabledAt: number | null;
58}
59 
60const UNIQUE_CONSTRAINT_PATTERN = /UNIQUE|SQLITE_CONSTRAINT/iu;
61const UNIQUE_SLUG_PATTERN = /users\.slug|idx_users_slug/iu;
62 
63const isUniqueViolation = (error: unknown): boolean => {
64 const cause = error instanceof DrizzleQueryError ? error.cause : error;
65 return cause instanceof Error && UNIQUE_CONSTRAINT_PATTERN.test(cause.message);
66};
67 
68const isSlugUniqueViolation = (error: unknown): boolean => {
69 const cause = error instanceof DrizzleQueryError ? error.cause : error;
70 return cause instanceof Error && UNIQUE_SLUG_PATTERN.test(cause.message);
71};
72 
73const clearTxCookie = (c: AppContext): void => {
74 deleteCookie(c, OIDC_TX_COOKIE_NAME, TX_COOKIE_CLEAR_OPTIONS);
75};
76 
77const loginErrorRedirect = (c: AppContext, code: string): Response => {
78 clearTxCookie(c);
79 return c.redirect(`/app/login?error=${encodeURIComponent(code)}`, 302);
80};
81 
82const buildRedirectUri = (requestUrl: string): string => `${new URL(requestUrl).origin}/api/public/oidc/callback`;
83 
84const findIdentityUserBySub = async (db: D1DbExecutor, sub: string): Promise<ExistingIdentityUser | undefined> => {
85 const rows = await db
86 .select({
87 userId: users.id,
88 userEmail: users.email,
89 disabledAt: users.disabledAt,
90 })
91 .from(tesseraIdentities)
92 .innerJoin(users, eq(tesseraIdentities.userId, users.id))
93 .where(eq(tesseraIdentities.sub, sub))
94 .limit(1);
95 
96 return rows[0];
97};
98 
99const localPartFromEmail = (email: string): string => email.split("@", 1)[0] ?? "";
100 
101const displayNameFromClaims = (claims: ResolvedClaims): string => {
102 const candidate = claims.name ?? claims.preferredUsername ?? localPartFromEmail(claims.email);
103 return normalizeDisplayName(candidate.trim().length > 0 ? candidate : "anvil user");
104};
105 
106const normalizeSlugCandidate = (value: string): string =>
107 value
108 .toLowerCase()
109 .replace(/[^a-z0-9_-]+/gu, "-")
110 .replace(/-+/gu, "-")
111 .replace(/^-+|-+$/gu, "");
112 
113export const generateUserSlug = (claims: ResolvedClaims, userId: string, attempt: number): OwnerSlug => {
114 const suffix = userId.slice(-6);
115 const sourceCandidates = [claims.preferredUsername, localPartFromEmail(claims.email), claims.name];
116 const source = sourceCandidates.find((value) => value && value.trim().length > 0) ?? "";
117 const base = normalizeSlugCandidate(source) || `usr-${suffix}`;
118 const slug = attempt === 0 || base === `usr-${suffix}` ? base : `${base}-${suffix}`;
119 
120 return OwnerSlug.assertDecode(slug);
121};
122 
123const updateExistingIdentity = async (
124 db: D1DbExecutor,
125 claims: ResolvedClaims,
126 existing: ExistingIdentityUser,
127 now: number,
128): Promise<BindOutcome> => {
129 if (existing.disabledAt !== null) {
130 return { ok: false, code: "user_disabled" };
131 }
132 
133 if (existing.userEmail === claims.email) {
134 await db.update(tesseraIdentities).set({ lastSeenAt: now }).where(eq(tesseraIdentities.sub, claims.sub));
135 return { ok: true, kind: "signed_in", userId: existing.userId };
136 }
137 
138 const emailOwner = await findUserByEmail(db, claims.email);
139 if (emailOwner && emailOwner.id !== existing.userId) {
140 return { ok: false, code: "tessera_email_conflict" };
141 }
142 
143 try {
144 await db.batch([
145 db.update(users).set({ email: claims.email }).where(eq(users.id, existing.userId)),
146 db.update(tesseraIdentities).set({ lastSeenAt: now }).where(eq(tesseraIdentities.sub, claims.sub)),
147 ]);
148 } catch (error) {
149 if (isUniqueViolation(error)) {
150 return { ok: false, code: "tessera_email_conflict" };
151 }
152 
153 throw error;
154 }
155 
156 return { ok: true, kind: "email_updated", userId: existing.userId };
157};
158 
159const bindLegacyUser = async (
160 db: D1DbExecutor,
161 claims: ResolvedClaims,
162 userId: string,
163 now: number,
164): Promise<BindOutcome> => {
165 const existingUserIdentity = await findIdentityByUserId(db, userId);
166 if (existingUserIdentity) {
167 return { ok: false, code: "identity_conflict" };
168 }
169 
170 try {
171 const inserted = await insertIdentity(db, {
172 sub: claims.sub,
173 userId,
174 createdAt: now,
175 lastSeenAt: now,
176 });
177 
178 if (inserted) {
179 return { ok: true, kind: "bound_legacy", userId };
180 }
181 } catch (error) {
182 if (!isUniqueViolation(error)) {
183 throw error;
184 }
185 }
186 
187 const recheckUserId = await findUserIdBySub(db, claims.sub);
188 if (recheckUserId === userId) {
189 return { ok: true, kind: "bound_legacy", userId };
190 }
191 
192 return { ok: false, code: "identity_conflict" };
193};
194 
195const createUserForIdentity = async (db: D1DbExecutor, claims: ResolvedClaims, now: number): Promise<BindOutcome> => {
196 for (let attempt = 0; attempt < 3; attempt += 1) {
197 const userId = generateDurableEntityId("usr", now + attempt);
198 const slug = generateUserSlug(claims, userId, attempt);
199 
200 try {
201 await db.batch([
202 db.insert(users).values({
203 id: userId,
204 slug,
205 email: claims.email,
206 displayName: displayNameFromClaims(claims),
207 createdAt: now,
208 disabledAt: null,
209 }),
210 db.insert(tesseraIdentities).values({
211 sub: claims.sub,
212 userId,
213 createdAt: now,
214 lastSeenAt: now,
215 }),
216 ]);
217 
218 return { ok: true, kind: "created", userId };
219 } catch (error) {
220 if (!isUniqueViolation(error)) {
221 throw error;
222 }
223 
224 const racedUserId = await findUserIdBySub(db, claims.sub);
225 if (racedUserId) {
226 return { ok: true, kind: "raced_to_existing", userId: racedUserId };
227 }
228 
229 if (isSlugUniqueViolation(error) && attempt < 2) {
230 continue;
231 }
232 
233 return { ok: false, code: "identity_conflict" };
234 }
235 }
236 
237 return { ok: false, code: "identity_conflict" };
238};
239 
240export const bindIdentity = async (db: D1DbExecutor, claims: ResolvedClaims): Promise<BindOutcome> => {
241 const now = Date.now();
242 const existing = await findIdentityUserBySub(db, claims.sub);
243 
244 if (existing) {
245 return await updateExistingIdentity(db, claims, existing, now);
246 }
247 
248 const legacyUser = await findUserByEmail(db, claims.email);
249 if (legacyUser) {
250 if (legacyUser.disabledAt !== null) {
251 return { ok: false, code: "user_disabled" };
252 }
253 
254 return await bindLegacyUser(db, claims, legacyUser.id, now);
255 }
256 
257 return await createUserForIdentity(db, claims, now);
258};
259 
260export const handleOidcStart = async (c: AppContext): Promise<Response> => {
261 const returnTo = sanitizeReturnTo(c.req.query("return_to"));
262 const redirectUri = buildRedirectUri(c.req.url);
263 
264 let config: oidc.Configuration;
265 let txCookieSecret: ArrayBuffer;
266 try {
267 txCookieSecret = await getTxCookieSecret(c.env);
268 config = await getOidcConfig(c.env);
269 } catch (error) {
270 logger.error("oidc_discovery_failed", oidcErrorContext(error, c.env));
271 return loginErrorRedirect(c, "oidc_provider_error");
272 }
273 
274 const codeVerifier = oidc.randomPKCECodeVerifier();
275 const codeChallenge = await oidc.calculatePKCECodeChallenge(codeVerifier);
276 const state = oidc.randomState();
277 const nonce = oidc.randomNonce();
278 const txPayload: TxCookiePayload = {
279 state,
280 nonce,
281 codeVerifier,
282 redirectUri,
283 returnTo,
284 createdAt: Date.now(),
285 };
286 
287 await setSignedCookie(c, OIDC_TX_COOKIE_NAME, encodeTxCookiePayload(txPayload), txCookieSecret, TX_COOKIE_OPTIONS);
288 
289 const authorizationUrl = oidc.buildAuthorizationUrl(config, {
290 redirect_uri: redirectUri,
291 code_challenge: codeChallenge,
292 code_challenge_method: "S256",
293 state,
294 nonce,
295 scope: OIDC_SCOPE,
296 });
297 
298 logger.info("oidc_start", { returnTo });
299 return c.redirect(authorizationUrl.toString(), 302);
300};
301 
302export const handleOidcCallback = async (c: AppContext): Promise<Response> => {
303 let txCookieValue: string | false | undefined;
304 try {
305 const txCookieSecret = await getTxCookieSecret(c.env);
306 txCookieValue = await getSignedCookie(c, txCookieSecret, OIDC_TX_COOKIE_NAME, TX_COOKIE_PREFIX);
307 } catch (error) {
308 logger.error("oidc_cookie_read_failed", oidcErrorContext(error, c.env));
309 return loginErrorRedirect(c, "oidc_provider_error");
310 }
311 
312 const txPayload = decodeTxCookiePayload(txCookieValue);
313 if (!txPayload) {
314 return loginErrorRedirect(c, "oidc_session_expired");
315 }
316 
317 const callbackUrl = new URL(c.req.url);
318 const queryState = callbackUrl.searchParams.get("state");
319 if (!queryState || queryState !== txPayload.state) {
320 return loginErrorRedirect(c, "oidc_session_expired");
321 }
322 
323 const expectedRedirectUri = new URL(txPayload.redirectUri);
324 if (expectedRedirectUri.origin !== callbackUrl.origin || expectedRedirectUri.pathname !== callbackUrl.pathname) {
325 return loginErrorRedirect(c, "oidc_session_expired");
326 }
327 
328 let config: oidc.Configuration;
329 try {
330 config = await getOidcConfig(c.env);
331 } catch (error) {
332 logger.error("oidc_discovery_failed", oidcErrorContext(error, c.env));
333 return loginErrorRedirect(c, "oidc_provider_error");
334 }
335 
336 let tokens: { claims(): oidc.IDToken | undefined };
337 try {
338 tokens = await exchangeAuthorizationCode(config, callbackUrl, {
339 pkceCodeVerifier: txPayload.codeVerifier,
340 expectedNonce: txPayload.nonce,
341 expectedState: txPayload.state,
342 });
343 } catch (error) {
344 logger.warn("oidc_token_exchange_failed", oidcErrorContext(error, c.env));
345 return loginErrorRedirect(c, "oidc_provider_error");
346 }
347 
348 const claimsResult = validateClaims(tokens.claims());
349 if (!claimsResult.ok) {
350 logger.warn("oidc_callback_failed", { code: claimsResult.code });
351 return loginErrorRedirect(c, claimsResult.code);
352 }
353 
354 const outcome = await bindIdentity(c.get("db"), claimsResult.claims);
355 if (!outcome.ok) {
356 logger.warn("oidc_callback_failed", { code: outcome.code });
357 return loginErrorRedirect(c, outcome.code);
358 }
359 
360 const { sessionId } = await createSession(c.env, outcome.userId);
361 clearTxCookie(c);
362 setSessionCookie(c, sessionId);
363 
364 logger.info("oidc_callback_success", { userId: outcome.userId, outcome: outcome.kind });
365 return c.redirect(appendOidcMarker(sanitizeReturnTo(txPayload.returnTo)), 302);
366};