Skip to content
File

Blob: src/worker/lib/ai/workers-ai.ts

typescript230 lines
1import type { AiUsage } from "@/shared/ai";
2import {
3 AiBackendError,
4 type AiChatMessage,
5 type AiChatOptions,
6 type AiClient,
7 type AiFrame,
8 type AiSummarizeOptions,
9 type AiSummarizeResult,
10} from "@/worker/lib/ai/types";
11import { parseProviderSseFrames, parseProviderTokenUsage } from "@/worker/lib/ai/sse-transform";
12import { aiLogger } from "@/worker/lib/ai/logging";
13 
14const log = aiLogger();
15const DEBUG_SAMPLE_LIMIT = 3;
16 
17export const DEFAULT_WORKERS_CHAT_MODEL = "@cf/google/gemma-4-26b-a4b-it";
18export const DEFAULT_WORKERS_SUMMARIZE_MODEL = "@cf/google/gemma-4-26b-a4b-it";
19 
20export interface WorkersAiConfig {
21 chatModel: string;
22 summarizeModel: string;
23}
24 
25export function createWorkersAiClient(ai: Ai, config: WorkersAiConfig): AiClient {
26 return {
27 async chat(messages: AiChatMessage[], opts?: AiChatOptions): Promise<AsyncIterable<AiFrame>> {
28 const runOpts = buildRunOptions(opts);
29 const result = await ai.run(
30 config.chatModel,
31 {
32 messages,
33 stream: true,
34 ...tokenBudget(config.chatModel, opts?.maxTokens ?? 1024),
35 temperature: opts?.temperature ?? 0.7,
36 ...thinkingOverride(config.chatModel),
37 },
38 runOpts,
39 );
40 
41 if (!(result instanceof ReadableStream)) {
42 throw new AiBackendError("Workers AI did not return a stream for chat", "ai_chat_no_stream");
43 }
44 return instrumentedWorkersStream(result as ReadableStream<Uint8Array>, config.chatModel);
45 },
46 
47 async summarize(text: string, _opts?: AiSummarizeOptions): Promise<AiSummarizeResult> {
48 return isDedicatedSummarizer(config.summarizeModel)
49 ? runDedicatedSummarize(ai, config.summarizeModel, text)
50 : runChatSummarize(ai, config.summarizeModel, text);
51 },
52 };
53}
54 
55function buildRunOptions(opts?: AiChatOptions): AiOptions | undefined {
56 const runOpts: AiOptions = {};
57 if (opts?.sessionKey) runOpts.extraHeaders = { "x-session-affinity": opts.sessionKey };
58 if (opts?.signal) runOpts.signal = opts.signal;
59 return runOpts.extraHeaders || runOpts.signal ? runOpts : undefined;
60}
61 
62// ---- Model quirks ----
63 
64// Models bound to Cloudflare's `ChatCompletionsInput` schema in
65// `worker-configuration.d.ts` — these four are the only ones that accept the
66// full OpenAI-compat option set (`max_completion_tokens`, `reasoning_effort`,
67// `chat_template_kwargs`, …). Every other Workers AI chat/instruct model uses
68// a bespoke per-model input type with only the legacy `max_tokens` field and
69// ignores the modern knobs.
70//
71// When Cloudflare migrates more models onto `ChatCompletionsInput` (new
72// bindings near `worker-configuration.d.ts:9409-9423`), add them here.
73const OPENAI_COMPAT_SCHEMA_MODELS = new Set<string>([
74 "@cf/google/gemma-4-26b-a4b-it",
75 "@cf/zai-org/glm-4.7-flash",
76 "@cf/moonshotai/kimi-k2.5",
77 "@cf/nvidia/nemotron-3-120b-a12b",
78]);
79 
80// bart-large-cnn and siblings take { input_text } and return { summary }; all
81// chat/instruct models take { messages } and return chat-completion shape.
82// Match the narrow set that needs the legacy summarization path.
83function isDedicatedSummarizer(model: string): boolean {
84 return model.includes("/bart-");
85}
86 
87// `chat_template_kwargs.enable_thinking: false` suppresses the <think> block on
88// the four OpenAI-compat-schema reasoning models. Bespoke-schema reasoning
89// models (Qwen 3, QwQ, DeepSeek R1 distill, GPT-OSS) don't expose the kwarg —
90// they'd either ignore it or error, so we skip it.
91function thinkingOverride(model: string): { chat_template_kwargs?: ChatTemplateKwargs } {
92 return OPENAI_COMPAT_SCHEMA_MODELS.has(model) ? { chat_template_kwargs: { enable_thinking: false } } : {};
93}
94 
95// Workers AI models disagree on the token-budget field name. The
96// ChatCompletionsInput models (Gemma 4, GLM, Kimi, Nemotron) honor
97// `max_completion_tokens`; every other model uses the legacy per-model schema
98// and honors `max_tokens`. Sending the wrong field is silently ignored, so the
99// model runs to its server-side default — we want explicit control.
100function tokenBudget(model: string, budget: number): { max_tokens: number } | { max_completion_tokens: number } {
101 return OPENAI_COMPAT_SCHEMA_MODELS.has(model) ? { max_completion_tokens: budget } : { max_tokens: budget };
102}
103 
104// ---- Summarize paths ----
105 
106async function runDedicatedSummarize(ai: Ai, model: string, text: string): Promise<AiSummarizeResult> {
107 const result = (await ai.run(model, { input_text: text })) as AiSummarizationOutput;
108 const summary = typeof result.summary === "string" ? result.summary.trim() : "";
109 if (!summary) {
110 throw new AiBackendError("Workers AI summarize returned empty summary", "ai_summarize_empty");
111 }
112 const usage = parseProviderTokenUsage(result.usage);
113 return usage ? { summary, usage } : { summary };
114}
115 
116async function runChatSummarize(ai: Ai, model: string, text: string): Promise<AiSummarizeResult> {
117 const result = await ai.run(model, {
118 messages: [
119 { role: "system", content: "Summarize the user's document in 3–5 sentences. Be concise and faithful." },
120 { role: "user", content: text },
121 ],
122 ...tokenBudget(model, 768),
123 temperature: 0.2,
124 ...thinkingOverride(model),
125 });
126 const summary = extractMessageContent(result);
127 if (!summary) {
128 throw new AiBackendError(
129 `Workers AI summarize returned empty summary (shape: ${describeShape(result)})`,
130 "ai_summarize_empty",
131 );
132 }
133 const usage = parseProviderTokenUsage((result as { usage?: unknown } | null)?.usage);
134 return usage ? { summary, usage } : { summary };
135}
136 
137// ---- Stream instrumentation (debug canary) ----
138 
139// Logs the first few payload shapes plus a summary line once the stream ends.
140// Useful when swapping models — reasoning frames on an instruct-only model, or
141// unexpected payload shape, both surface here at debug level.
142async function* instrumentedWorkersStream(source: ReadableStream<Uint8Array>, model: string): AsyncGenerator<AiFrame> {
143 let chunks = 0;
144 let sampled = 0;
145 let reasoningFrames = 0;
146 
147 const frames = parseProviderSseFrames(source, {
148 extractChunkText,
149 extractUsage,
150 errorLabel: "workers-ai stream error",
151 onPayload: (payload, index) => {
152 if (hasReasoningDelta(payload)) reasoningFrames++;
153 if (index < DEBUG_SAMPLE_LIMIT) {
154 sampled++;
155 log.debug("ai_stream_sample", { model, index, shape: describeShape(payload) });
156 }
157 },
158 });
159 
160 for await (const frame of frames) {
161 if (frame.type === "chunk") chunks++;
162 yield frame;
163 }
164 log.debug("ai_stream_summary", { model, sampled, chunks, reasoningFrames });
165}
166 
167// ---- Shape-aware extractors (streaming + non-streaming) ----
168 
169// Stream frames come in two shapes: legacy `{response: "…"}` (Llama, Mistral)
170// and OpenAI-compat `{choices: [{delta: {content: "…"}}]}` (Gemma, Qwen 3,
171// Scout, …). Try both.
172function extractChunkText(payload: unknown): string | null {
173 if (typeof payload !== "object" || payload === null) return null;
174 const legacy = (payload as AiTextGenerationOutput).response;
175 if (typeof legacy === "string" && legacy.length > 0) return legacy;
176 const completion = payload as { choices?: Array<{ delta?: { content?: unknown } }> };
177 const content = completion.choices?.[0]?.delta?.content;
178 return typeof content === "string" && content.length > 0 ? content : null;
179}
180 
181function extractUsage(payload: unknown): AiUsage | null {
182 if (typeof payload !== "object" || payload === null) return null;
183 return parseProviderTokenUsage((payload as { usage?: unknown }).usage);
184}
185 
186// Non-streaming chat response: `{choices: [{message: {content}}]}` (Gemma,
187// Scout) or legacy `{response}` (older Llama/Mistral). Trim on the way out.
188function extractMessageContent(result: unknown): string {
189 if (typeof result !== "object" || result === null) return "";
190 const legacy = (result as AiTextGenerationOutput).response;
191 if (typeof legacy === "string" && legacy.trim().length > 0) return legacy.trim();
192 const completion = result as Partial<ChatCompletionsOutput>;
193 const content = completion.choices?.[0]?.message?.content;
194 return typeof content === "string" && content.trim().length > 0 ? content.trim() : "";
195}
196 
197function hasReasoningDelta(payload: unknown): boolean {
198 if (typeof payload !== "object" || payload === null) return false;
199 const delta = (payload as { choices?: Array<{ delta?: Record<string, unknown> }> }).choices?.[0]?.delta;
200 if (!delta) return false;
201 return typeof delta.reasoning === "string" || typeof delta.reasoning_content === "string";
202}
203 
204// Structural fingerprint used in debug logs and error messages. Leaks no content,
205// just the keys at each relevant nesting level so we can diagnose shape drift.
206function describeShape(value: unknown): string {
207 if (value === null || typeof value !== "object") return typeof value;
208 const top = Object.keys(value).sort();
209 const choices = (value as { choices?: unknown }).choices;
210 if (!Array.isArray(choices) || choices.length === 0) return `top=[${top.join(",")}]`;
211 
212 const first = choices[0] as Record<string, unknown> | null;
213 const choiceKeys = first ? Object.keys(first).sort() : [];
214 const deltaKeys =
215 first && typeof first.delta === "object" && first.delta !== null ? Object.keys(first.delta).sort() : null;
216 const messageKeys =
217 first && typeof first.message === "object" && first.message !== null ? Object.keys(first.message).sort() : null;
218 return `top=[${top.join(",")}] choices[0]=[${choiceKeys.join(",")}]${
219 deltaKeys ? ` delta=[${deltaKeys.join(",")}]` : ""
220 }${messageKeys ? ` message=[${messageKeys.join(",")}]` : ""}`;
221}
222 
223// `@cf/facebook/bart-large-cnn` returns `{ summary }` instead of the chat-completion
224// shape; Workers AI's generated `AiSummarizationOutput` type is named but it isn't
225// exported as a value, so we redeclare the minimum surface we read.
226interface AiSummarizationOutput {
227 summary?: string;
228 usage?: unknown;
229}