File
Blob: src/worker/routes/ai.ts
| 1 | import { Hono, type Context } from "hono"; |
| 2 | |
| 3 | import { requireAuth, optionalAuth } from "@/worker/middleware/auth"; |
| 4 | import { rateLimit } from "@/worker/middleware/rate-limit"; |
| 5 | import type { AppContext } from "@/worker/app-context"; |
| 6 | import { getPage } from "@/worker/lib/page-access"; |
| 7 | import { resolvePageAccessLevels, resolvePrincipal } from "@/worker/lib/permissions"; |
| 8 | import { parseBody } from "@/worker/lib/validate"; |
| 9 | import { errorContext } from "@/worker/lib/logger"; |
| 10 | import { createAiClient } from "@/worker/lib/ai"; |
| 11 | import type { AiChatMessage, AiClient, AiFrame } from "@/worker/lib/ai"; |
| 12 | import { clientAiErrorMessage } from "@/worker/lib/ai/client-error"; |
| 13 | import { buildAskMessages, buildGenerateMessages, buildRewriteMessages } from "@/worker/lib/ai/prompts"; |
| 14 | import { |
| 15 | aiLogger, |
| 16 | logAiDenied, |
| 17 | logAiRequest, |
| 18 | logAiResponse, |
| 19 | type AiAction, |
| 20 | type AiLogContext, |
| 21 | } from "@/worker/lib/ai/logging"; |
| 22 | import { |
| 23 | getPageAiEntitlements, |
| 24 | type EntitlementSurface, |
| 25 | type PageAccessLevel, |
| 26 | type PageAiEntitlements, |
| 27 | } from "@/shared/entitlements"; |
| 28 | import { AiAskRequest, AiGenerateRequest, AiRewriteRequest } from "@/shared/types"; |
| 29 | import { encodeAiSseChunk, encodeAiSseDone, encodeAiSseError, type AiErrorCode, type AiUsage } from "@/shared/ai"; |
| 30 | |
| 31 | const log = aiLogger(); |
| 32 | |
| 33 | const SSE_HEADERS: HeadersInit = { |
| 34 | "content-type": "text/event-stream", |
| 35 | "cache-control": "no-cache", |
| 36 | "x-accel-buffering": "no", |
| 37 | }; |
| 38 | |
| 39 | const aiRouter = new Hono<AppContext>(); |
| 40 | |
| 41 | aiRouter.post("/workspaces/:wid/pages/:id/rewrite", requireAuth, rateLimit("RL_AI"), async (c) => { |
| 42 | const startedAt = Date.now(); |
| 43 | const workspaceId = c.req.param("wid"); |
| 44 | const pageId = c.req.param("id"); |
| 45 | const data = await parseBody(c, AiRewriteRequest); |
| 46 | if (data instanceof Response) return data; |
| 47 | |
| 48 | const gate = await gateAiAction(c, workspaceId, pageId, "rewrite", (ent) => ent.useAiRewrite); |
| 49 | if (gate instanceof Response) return gate; |
| 50 | |
| 51 | const logCtx = buildLogContext(c, "rewrite", workspaceId, pageId, gate); |
| 52 | logAiRequest(logCtx); |
| 53 | |
| 54 | const client = resolveClient(c); |
| 55 | if (client instanceof Response) return client; |
| 56 | |
| 57 | const messages = buildRewriteMessages(data); |
| 58 | return streamChat(c, client, messages, `rewrite:${c.get("user")!.id}:${pageId}`, logCtx, startedAt); |
| 59 | }); |
| 60 | |
| 61 | aiRouter.post("/workspaces/:wid/pages/:id/generate", requireAuth, rateLimit("RL_AI"), async (c) => { |
| 62 | const startedAt = Date.now(); |
| 63 | const workspaceId = c.req.param("wid"); |
| 64 | const pageId = c.req.param("id"); |
| 65 | const data = await parseBody(c, AiGenerateRequest); |
| 66 | if (data instanceof Response) return data; |
| 67 | |
| 68 | const gate = await gateAiAction(c, workspaceId, pageId, "generate", (ent) => ent.useAiGenerate); |
| 69 | if (gate instanceof Response) return gate; |
| 70 | |
| 71 | const logCtx = buildLogContext(c, "generate", workspaceId, pageId, gate); |
| 72 | logAiRequest(logCtx); |
| 73 | |
| 74 | const client = resolveClient(c); |
| 75 | if (client instanceof Response) return client; |
| 76 | |
| 77 | const messages = buildGenerateMessages(data); |
| 78 | return streamChat(c, client, messages, `generate:${c.get("user")!.id}:${pageId}`, logCtx, startedAt); |
| 79 | }); |
| 80 | |
| 81 | aiRouter.post("/workspaces/:wid/pages/:id/summarize", optionalAuth, rateLimit("RL_AI"), async (c) => { |
| 82 | const startedAt = Date.now(); |
| 83 | const workspaceId = c.req.param("wid"); |
| 84 | const pageId = c.req.param("id"); |
| 85 | |
| 86 | const gate = await gateAiAction(c, workspaceId, pageId, "summarize", (ent) => ent.summarizePage); |
| 87 | if (gate instanceof Response) return gate; |
| 88 | |
| 89 | const logCtx = buildLogContext(c, "summarize", workspaceId, pageId, gate); |
| 90 | logAiRequest(logCtx); |
| 91 | |
| 92 | const body = await requireNonEmptyBody(c, pageId); |
| 93 | if (body instanceof Response) { |
| 94 | logAiResponse(logCtx, startedAt, "error", { errorCode: "page_empty" }); |
| 95 | return body; |
| 96 | } |
| 97 | |
| 98 | const client = resolveClient(c); |
| 99 | if (client instanceof Response) { |
| 100 | logAiResponse(logCtx, startedAt, "error", { errorCode: "ai_misconfigured" }); |
| 101 | return client; |
| 102 | } |
| 103 | |
| 104 | try { |
| 105 | const result = await client.summarize(body.bodyText); |
| 106 | logAiResponse(logCtx, startedAt, "ok", result.usage ? { usage: result.usage } : undefined); |
| 107 | return c.json(result); |
| 108 | } catch (err) { |
| 109 | log.error("summarize_failed", { ...errorContext(err), action: "summarize", pageId, workspaceId }); |
| 110 | const safe = clientAiErrorMessage(err); |
| 111 | logAiResponse(logCtx, startedAt, "error", { errorCode: safe.code }); |
| 112 | return c.json({ error: safe.code, message: safe.message }, 502); |
| 113 | } |
| 114 | }); |
| 115 | |
| 116 | aiRouter.post("/workspaces/:wid/pages/:id/ask", optionalAuth, rateLimit("RL_AI"), async (c) => { |
| 117 | const startedAt = Date.now(); |
| 118 | const workspaceId = c.req.param("wid"); |
| 119 | const pageId = c.req.param("id"); |
| 120 | |
| 121 | const data = await parseBody(c, AiAskRequest); |
| 122 | if (data instanceof Response) return data; |
| 123 | |
| 124 | const gate = await gateAiAction(c, workspaceId, pageId, "ask", (ent) => ent.askPage); |
| 125 | if (gate instanceof Response) return gate; |
| 126 | |
| 127 | const logCtx = buildLogContext(c, "ask", workspaceId, pageId, gate); |
| 128 | logAiRequest(logCtx); |
| 129 | |
| 130 | const body = await requireNonEmptyBody(c, pageId); |
| 131 | if (body instanceof Response) { |
| 132 | logAiResponse(logCtx, startedAt, "error", { errorCode: "page_empty" }); |
| 133 | return body; |
| 134 | } |
| 135 | |
| 136 | const client = resolveClient(c); |
| 137 | if (client instanceof Response) { |
| 138 | logAiResponse(logCtx, startedAt, "error", { errorCode: "ai_misconfigured" }); |
| 139 | return client; |
| 140 | } |
| 141 | |
| 142 | // Shared surface has no AI entitlements, so reaching this point guarantees an authenticated user. |
| 143 | const userId = c.get("user")!.id; |
| 144 | const messages = buildAskMessages(body.title, body.bodyText.slice(0, 6000), data.question, data.history ?? []); |
| 145 | return streamChat(c, client, messages, `ask:${userId}:${pageId}`, logCtx, startedAt); |
| 146 | }); |
| 147 | |
| 148 | export { aiRouter }; |
| 149 | |
| 150 | async function gateAiAction( |
| 151 | c: Context<AppContext>, |
| 152 | workspaceId: string, |
| 153 | pageId: string, |
| 154 | action: AiAction, |
| 155 | select: (ent: PageAiEntitlements) => boolean, |
| 156 | ): Promise<{ surface: EntitlementSurface; pageAccess: PageAccessLevel } | Response> { |
| 157 | const user = c.get("user"); |
| 158 | const db = c.get("db"); |
| 159 | const shareToken = c.req.query("share"); |
| 160 | // Surface is the route shape, not a principal heuristic: any `?share=` request |
| 161 | // is shared-scoped regardless of whether the caller also holds workspace |
| 162 | // membership. `getPageAiEntitlements("shared", ...)` denies all AI actions. |
| 163 | const surface: EntitlementSurface = shareToken ? "shared" : "canonical"; |
| 164 | |
| 165 | const resolved = await resolvePrincipal(db, user, workspaceId, { surface, shareToken }); |
| 166 | if (!resolved) { |
| 167 | return c.json({ error: "unauthorized", message: "Authentication required" }, 401); |
| 168 | } |
| 169 | |
| 170 | // Resolve access before loading page metadata so an inaccessible canvas page |
| 171 | // returns the same `not_found` as a missing page (no kind leak via response code). |
| 172 | const levels = await resolvePageAccessLevels(db, resolved.principal, [pageId], workspaceId); |
| 173 | const pageAccess = levels.get(pageId) ?? "none"; |
| 174 | if (pageAccess === "none") { |
| 175 | return c.json({ error: "not_found", message: "Page not found" }, 404); |
| 176 | } |
| 177 | |
| 178 | const page = await getPage(db, pageId, workspaceId); |
| 179 | if (!page) { |
| 180 | return c.json({ error: "not_found", message: "Page not found" }, 404); |
| 181 | } |
| 182 | |
| 183 | if (page.kind === "canvas") { |
| 184 | return c.json({ error: "page_empty", message: "AI is not available on canvas pages yet" }, 404); |
| 185 | } |
| 186 | |
| 187 | // Role axis enforces member-only on canonical surface even if the caller has |
| 188 | // a share-based grant. `resolved.workspaceRole` is null on shared surface and |
| 189 | // for non-members, both of which `getPageAiEntitlements` already treats as |
| 190 | // all-deny. |
| 191 | const workspaceRole = resolved.workspaceRole ?? "none"; |
| 192 | if (!select(getPageAiEntitlements(surface, pageAccess, workspaceRole))) { |
| 193 | logAiDenied({ |
| 194 | action, |
| 195 | userId: user?.id, |
| 196 | workspaceId, |
| 197 | pageId, |
| 198 | surface, |
| 199 | pageAccess, |
| 200 | }); |
| 201 | return c.json({ error: "ai_not_entitled", message: "AI action not permitted on this page" }, 403); |
| 202 | } |
| 203 | |
| 204 | return { surface, pageAccess }; |
| 205 | } |
| 206 | |
| 207 | function buildLogContext( |
| 208 | c: Context<AppContext>, |
| 209 | action: AiAction, |
| 210 | workspaceId: string, |
| 211 | pageId: string, |
| 212 | gate: { surface: EntitlementSurface; pageAccess: PageAccessLevel }, |
| 213 | ): AiLogContext { |
| 214 | const user = c.get("user"); |
| 215 | return { |
| 216 | action, |
| 217 | workspaceId, |
| 218 | pageId, |
| 219 | surface: gate.surface, |
| 220 | pageAccess: gate.pageAccess, |
| 221 | ...(user ? { userId: user.id } : {}), |
| 222 | }; |
| 223 | } |
| 224 | |
| 225 | async function requireNonEmptyBody( |
| 226 | c: Context<AppContext>, |
| 227 | pageId: string, |
| 228 | ): Promise<{ bodyText: string; title: string } | Response> { |
| 229 | const payload = await c.env.DocSync.getByName(pageId).getIndexPayload(pageId); |
| 230 | if (payload.kind === "missing" || payload.bodyText.trim().length === 0) { |
| 231 | return c.json({ error: "page_empty", message: "Page has no body text yet" }, 404); |
| 232 | } |
| 233 | return { bodyText: payload.bodyText, title: payload.title }; |
| 234 | } |
| 235 | |
| 236 | function resolveClient(c: Context<AppContext>): AiClient | Response { |
| 237 | try { |
| 238 | return createAiClient(c.env); |
| 239 | } catch (err) { |
| 240 | log.error("ai_client_misconfigured", errorContext(err)); |
| 241 | const safe = clientAiErrorMessage(err); |
| 242 | return c.json({ error: safe.code, message: safe.message }, 503); |
| 243 | } |
| 244 | } |
| 245 | |
| 246 | async function streamChat( |
| 247 | c: Context<AppContext>, |
| 248 | client: AiClient, |
| 249 | messages: AiChatMessage[], |
| 250 | sessionKey: string, |
| 251 | logCtx: AiLogContext, |
| 252 | startedAt: number, |
| 253 | ): Promise<Response> { |
| 254 | // Single AbortController fed by both the inbound request signal (client gave |
| 255 | // up before we dispatched) and the outer stream's cancel() (client gave up |
| 256 | // mid-stream). Either should propagate to upstream fetch / ai.run. |
| 257 | const upstream = new AbortController(); |
| 258 | const requestSignal = c.req.raw.signal; |
| 259 | if (requestSignal.aborted) upstream.abort(); |
| 260 | else requestSignal.addEventListener("abort", () => upstream.abort(), { once: true }); |
| 261 | |
| 262 | let iter: AsyncIterable<AiFrame>; |
| 263 | try { |
| 264 | iter = await client.chat(messages, { sessionKey, signal: upstream.signal }); |
| 265 | } catch (err) { |
| 266 | log.error("ai_chat_failed", { ...errorContext(err), action: logCtx.action, pageId: logCtx.pageId }); |
| 267 | const safe = clientAiErrorMessage(err); |
| 268 | logAiResponse(logCtx, startedAt, "error", { errorCode: safe.code }); |
| 269 | // 200 + SSE error frame so the client SSE parser sees it. A non-2xx here |
| 270 | // would be swallowed by sendApiRequest's JSON-only error path and the |
| 271 | // structured code/message would never reach the user. |
| 272 | const body = new ReadableStream<Uint8Array>({ |
| 273 | start(controller) { |
| 274 | controller.enqueue(encodeAiSseError(safe.message, safe.code)); |
| 275 | controller.enqueue(encodeAiSseDone()); |
| 276 | controller.close(); |
| 277 | }, |
| 278 | }); |
| 279 | return new Response(body, { headers: SSE_HEADERS, status: 200 }); |
| 280 | } |
| 281 | |
| 282 | const body = new ReadableStream<Uint8Array>({ |
| 283 | async start(controller) { |
| 284 | let usage: AiUsage | undefined; |
| 285 | let errorCode: AiErrorCode | undefined; |
| 286 | try { |
| 287 | for await (const frame of iter) { |
| 288 | if (upstream.signal.aborted) return; |
| 289 | if (frame.type === "chunk") { |
| 290 | controller.enqueue(encodeAiSseChunk(frame.text)); |
| 291 | } else if (frame.type === "usage") { |
| 292 | usage = frame.usage; |
| 293 | } else if (frame.type === "error") { |
| 294 | errorCode = frame.code; |
| 295 | controller.enqueue(encodeAiSseError(frame.message, frame.code)); |
| 296 | } |
| 297 | } |
| 298 | controller.enqueue(encodeAiSseDone(usage)); |
| 299 | } catch (err) { |
| 300 | if (upstream.signal.aborted) return; |
| 301 | log.error("ai_chat_stream_failed", { ...errorContext(err), action: logCtx.action, pageId: logCtx.pageId }); |
| 302 | const safe = clientAiErrorMessage(err); |
| 303 | errorCode = safe.code; |
| 304 | controller.enqueue(encodeAiSseError(safe.message, safe.code)); |
| 305 | controller.enqueue(encodeAiSseDone(usage)); |
| 306 | } finally { |
| 307 | if (errorCode) { |
| 308 | logAiResponse(logCtx, startedAt, "error", { errorCode }); |
| 309 | } else if (!upstream.signal.aborted) { |
| 310 | logAiResponse(logCtx, startedAt, "ok", usage ? { usage } : undefined); |
| 311 | } |
| 312 | try { |
| 313 | controller.close(); |
| 314 | } catch { |
| 315 | // already closed via cancel() |
| 316 | } |
| 317 | } |
| 318 | }, |
| 319 | cancel() { |
| 320 | // Client disconnected mid-stream; abort upstream so we stop spending |
| 321 | // tokens on output no one will read. |
| 322 | upstream.abort(); |
| 323 | }, |
| 324 | }); |
| 325 | return new Response(body, { headers: SSE_HEADERS }); |
| 326 | } |