Skip to content
File

Blob: src/worker/routes/ai.ts

typescript327 lines
1import { Hono, type Context } from "hono";
2 
3import { requireAuth, optionalAuth } from "@/worker/middleware/auth";
4import { rateLimit } from "@/worker/middleware/rate-limit";
5import type { AppContext } from "@/worker/app-context";
6import { getPage } from "@/worker/lib/page-access";
7import { resolvePageAccessLevels, resolvePrincipal } from "@/worker/lib/permissions";
8import { parseBody } from "@/worker/lib/validate";
9import { errorContext } from "@/worker/lib/logger";
10import { createAiClient } from "@/worker/lib/ai";
11import type { AiChatMessage, AiClient, AiFrame } from "@/worker/lib/ai";
12import { clientAiErrorMessage } from "@/worker/lib/ai/client-error";
13import { buildAskMessages, buildGenerateMessages, buildRewriteMessages } from "@/worker/lib/ai/prompts";
14import {
15 aiLogger,
16 logAiDenied,
17 logAiRequest,
18 logAiResponse,
19 type AiAction,
20 type AiLogContext,
21} from "@/worker/lib/ai/logging";
22import {
23 getPageAiEntitlements,
24 type EntitlementSurface,
25 type PageAccessLevel,
26 type PageAiEntitlements,
27} from "@/shared/entitlements";
28import { AiAskRequest, AiGenerateRequest, AiRewriteRequest } from "@/shared/types";
29import { encodeAiSseChunk, encodeAiSseDone, encodeAiSseError, type AiErrorCode, type AiUsage } from "@/shared/ai";
30 
31const log = aiLogger();
32 
33const SSE_HEADERS: HeadersInit = {
34 "content-type": "text/event-stream",
35 "cache-control": "no-cache",
36 "x-accel-buffering": "no",
37};
38 
39const aiRouter = new Hono<AppContext>();
40 
41aiRouter.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 
61aiRouter.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 
81aiRouter.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 
116aiRouter.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 
148export { aiRouter };
149 
150async 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 
207function 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 
225async 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 
236function 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 
246async 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}