Skip to content
File

Blob: src/cloudflare/internal/ai-api.ts

typescript470 lines
1// Copyright (c) 2025 Cloudflare, Inc.
2// Licensed under the Apache 2.0 license found in the LICENSE file or at:
3// https://opensource.org/licenses/Apache-2.0
4 
5import { AiGateway, type GatewayOptions } from 'cloudflare-internal:aig-api';
6import { AutoRAG } from 'cloudflare-internal:autorag-api';
7import {
8 ToMarkdownService,
9 type ConversionRequestOptions,
10 type ConversionResponse,
11 type MarkdownDocument,
12} from 'cloudflare-internal:to-markdown-api';
13 
14type AiSearchService = object;
15 
16const aiBindingExperimental = !!Cloudflare.compatibilityFlags['experimental'];
17 
18interface Fetcher {
19 fetch: typeof fetch;
20 aiSearch: () => AiSearchService;
21 gateway: (gatewayId: string) => AiGateway;
22 autorag: (autoragId?: string) => AutoRAG;
23 toMarkdown: () => ToMarkdownService;
24}
25 
26interface AiError {
27 internalCode: number;
28 message: string;
29 name: string;
30 description: string;
31 errors?: Array<{ code: number; message: string }>;
32}
33 
34export type SessionOptions = {
35 // Deprecated, do not use this
36 extraHeaders?: object;
37};
38 
39export type AiOptions = {
40 gateway?: GatewayOptions;
41 websocket?: boolean;
42 /** If true it will return a Response object */
43 returnRawResponse?: boolean;
44 prefix?: string;
45 extraHeaders?: object;
46 /*
47 * @deprecated this option is deprecated, do not use this
48 */
49 sessionOptions?: SessionOptions;
50 signal?: AbortSignal;
51};
52 
53type CleanedAiOptions = Omit<
54 AiOptions,
55 'prefix' | 'extraHeaders' | 'sessionOptions' | 'signal'
56>;
57 
58export type AiInputReadableStream = {
59 body: ReadableStream | FormData;
60 contentType: string;
61};
62 
63export type AiModelsSearchParams = {
64 author?: string;
65 hide_experimental?: boolean;
66 page?: number;
67 per_page?: number;
68 search?: string;
69 source?: number;
70 task?: string;
71};
72 
73export type AiModelsSearchObject = {
74 id: string;
75 source: number;
76 name: string;
77 description: string;
78 task: {
79 id: string;
80 name: string;
81 description: string;
82 };
83 tags: string[];
84 properties: {
85 property_id: string;
86 value: string;
87 }[];
88};
89 
90export class InferenceUpstreamError extends Error {
91 constructor(message: string, name = 'InferenceUpstreamError') {
92 super(message);
93 this.name = name;
94 }
95}
96 
97export class AiInternalError extends Error {
98 constructor(message: string, name = 'AiInternalError') {
99 super(message);
100 this.name = name;
101 }
102}
103 
104function isReadableStream(obj: unknown): obj is ReadableStream {
105 return obj instanceof ReadableStream;
106}
107 
108function isFormData(obj: unknown): obj is FormData {
109 return obj instanceof FormData;
110}
111 
112/**
113 * Find keys in inputs that have a ReadableStream
114 * */
115function findReadableStreamKeys(
116 inputs: Record<string, unknown>
117): Array<string> {
118 const readableStreamKeys: Array<string> = [];
119 
120 for (const [key, value] of Object.entries(inputs)) {
121 // Check if value has a body property that's a ReadableStream
122 const hasReadableStreamBody =
123 value &&
124 typeof value === 'object' &&
125 'body' in value &&
126 (isReadableStream(value.body) || isFormData(value.body));
127 
128 if (hasReadableStreamBody || isReadableStream(value) || isFormData(value)) {
129 readableStreamKeys.push(key);
130 }
131 }
132 
133 return readableStreamKeys;
134}
135 
136export class Ai {
137 #fetcher: Fetcher;
138 
139 /*
140 * @deprecated this option is deprecated, do not use this
141 */
142 // @ts-expect-error: deprecated var
143 // eslint-disable-next-line no-unused-private-class-members
144 #logs: Array<string> = [];
145 #options: AiOptions = {};
146 #endpointURL = 'https://workers-binding.ai';
147 lastRequestId: string | null = null;
148 aiGatewayLogId: string | null = null;
149 lastRequestHttpStatusCode: number | null = null;
150 lastRequestInternalStatusCode: number | null = null;
151 
152 constructor(fetcher: Fetcher) {
153 this.#fetcher = fetcher;
154 }
155 
156 async fetch(input: RequestInfo | URL, init?: RequestInit): Promise<Response> {
157 return this.#fetcher.fetch(input, init);
158 }
159 
160 /**
161 * Generate fetch call for JSON inputs
162 * */
163 async #generateFetch(
164 inputs: object,
165 cleanedOptions: CleanedAiOptions,
166 model: string
167 ): Promise<Response> {
168 // Treat inputs as regular JS objects
169 const body = JSON.stringify({
170 inputs,
171 options: cleanedOptions,
172 });
173 
174 const fetchOptions: RequestInit = {
175 method: 'POST',
176 body: body,
177 headers: {
178 ...this.#options.sessionOptions?.extraHeaders,
179 ...this.#options.extraHeaders,
180 'content-type': 'application/json',
181 'cf-consn-sdk-version': '2.0.0',
182 'cf-consn-model-id': `${this.#options.prefix ? `${this.#options.prefix}:` : ''}${model}`,
183 },
184 };
185 if (this.#options.signal) {
186 fetchOptions.signal = this.#options.signal;
187 }
188 
189 let endpointUrl = `${this.#endpointURL}/run?version=3`;
190 if (cleanedOptions.gateway?.id) {
191 endpointUrl = `${this.#endpointURL}/ai-gateway/run?version=3`;
192 }
193 
194 return await this.#fetcher.fetch(endpointUrl, fetchOptions);
195 }
196 
197 /**
198 * Generate fetch call for inputs with ReadableStream
199 * */
200 async #generateStreamFetch(
201 inputs: Record<string, string | AiInputReadableStream>,
202 cleanedOptions: CleanedAiOptions,
203 model: string,
204 streamKeys: string[]
205 ): Promise<Response> {
206 const streamKey = streamKeys[0] ?? '';
207 const stream = streamKey ? inputs[streamKey] : null;
208 const body = (stream as AiInputReadableStream).body;
209 const contentType = (stream as AiInputReadableStream).contentType;
210 
211 if (cleanedOptions.gateway?.id) {
212 throw new AiInternalError(
213 'AI Gateway does not support ReadableStreams yet.'
214 );
215 }
216 
217 // Make sure user has supplied the Content-Type
218 // This allows AI binding to treat the ReadableStream correctly
219 if (!contentType) {
220 throw new AiInternalError(
221 'Content-Type is required with ReadableStream inputs'
222 );
223 }
224 
225 // Pass single ReadableStream in request body
226 const fetchOptions: RequestInit = {
227 method: 'POST',
228 body: body,
229 headers: {
230 ...this.#options.sessionOptions?.extraHeaders,
231 ...this.#options.extraHeaders,
232 'content-type': contentType,
233 'cf-consn-sdk-version': '2.0.0',
234 'cf-consn-model-id': `${this.#options.prefix ? `${this.#options.prefix}:` : ''}${model}`,
235 },
236 };
237 if (this.#options.signal) {
238 fetchOptions.signal = this.#options.signal;
239 }
240 
241 // Fetch the additional input params
242 const { [streamKey]: streamInput, ...userInputs } = inputs;
243 
244 // Construct query params
245 // Append inputs with ai.run options that are passed to the inference request
246 const query = {
247 ...cleanedOptions,
248 version: '3',
249 userInputs: JSON.stringify({ ...userInputs }),
250 };
251 const aiEndpoint = new URL(`${this.#endpointURL}/run`);
252 for (const [key, value] of Object.entries(query)) {
253 aiEndpoint.searchParams.set(key, value as string);
254 }
255 
256 return await this.#fetcher.fetch(aiEndpoint, fetchOptions);
257 }
258 
259 /**
260 * Generate call to open a websocket connection
261 * */
262 async #generateWebsocketFetch(
263 inputs: object,
264 cleanedOptions: CleanedAiOptions,
265 model: string
266 ): Promise<Response> {
267 // Treat inputs as regular JS objects
268 const body = JSON.stringify({
269 inputs,
270 options: cleanedOptions,
271 });
272 
273 const fetchOptions: RequestInit = {
274 headers: {
275 ...this.#options.sessionOptions?.extraHeaders,
276 ...this.#options.extraHeaders,
277 'cf-consn-sdk-version': '2.0.0',
278 'cf-consn-model-id': `${this.#options.prefix ? `${this.#options.prefix}:` : ''}${model}`,
279 Upgrade: 'websocket',
280 },
281 };
282 if (this.#options.signal) {
283 fetchOptions.signal = this.#options.signal;
284 }
285 
286 const aiEndpoint = new URL(`${this.#endpointURL}/run`);
287 aiEndpoint.searchParams.set('version', '3');
288 aiEndpoint.searchParams.set('body', body);
289 
290 return await this.#fetcher.fetch(aiEndpoint, fetchOptions);
291 }
292 
293 async run(
294 model: string,
295 inputs: Record<string, string | AiInputReadableStream>,
296 options: AiOptions = {}
297 ): Promise<Response | ReadableStream<Uint8Array> | object | null> {
298 this.#options = options;
299 this.lastRequestId = '';
300 
301 // This removes some unwanted options from getting sent in the body
302 const cleanedOptions = (({
303 prefix,
304 extraHeaders,
305 sessionOptions,
306 signal,
307 ...object
308 }): CleanedAiOptions => object)(this.#options);
309 
310 let res: Response;
311 
312 if (this.#options.websocket) {
313 res = await this.#generateWebsocketFetch(inputs, cleanedOptions, model);
314 } else {
315 /**
316 * Inputs that contain a ReadableStream which will be sent directly to
317 * the fetcher object along with other keys parsed as a query parameters
318 * */
319 const streamKeys = findReadableStreamKeys(inputs);
320 
321 if (streamKeys.length === 0) {
322 res = await this.#generateFetch(inputs, cleanedOptions, model);
323 } else if (streamKeys.length > 1) {
324 throw new AiInternalError(
325 `Multiple ReadableStreams are not supported. Found streams in keys: [${streamKeys.join(', ')}]`
326 );
327 } else {
328 res = await this.#generateStreamFetch(
329 inputs,
330 cleanedOptions,
331 model,
332 streamKeys
333 );
334 }
335 }
336 
337 this.lastRequestId = res.headers.get('cf-ai-req-id');
338 this.aiGatewayLogId = res.headers.get('cf-aig-log-id');
339 this.lastRequestHttpStatusCode = res.status;
340 
341 if (this.#options.returnRawResponse || this.#options.websocket) {
342 return res;
343 }
344 
345 if (!res.ok || !res.body) {
346 throw await this._parseError(res);
347 }
348 
349 const contentType = res.headers.get('content-type');
350 if (contentType === 'application/json') {
351 return (await res.json()) as object;
352 }
353 
354 return res.body;
355 }
356 
357 /*
358 * @deprecated this method is deprecated, do not use this
359 */
360 getLogs(): string[] {
361 return [];
362 }
363 
364 // TODO(soon): Can we use the # syntax here?
365 // eslint-disable-next-line no-restricted-syntax
366 private async _parseError(res: Response): Promise<InferenceUpstreamError> {
367 const content = await res.text();
368 
369 try {
370 const parsedContent = JSON.parse(content) as AiError;
371 if (parsedContent.internalCode) {
372 this.lastRequestInternalStatusCode = parsedContent.internalCode;
373 return new InferenceUpstreamError(
374 `${parsedContent.internalCode}: ${parsedContent.description}`,
375 parsedContent.name
376 );
377 } else if (
378 parsedContent.errors &&
379 parsedContent.errors.length > 0 &&
380 parsedContent.errors[0]
381 ) {
382 return new InferenceUpstreamError(
383 `${parsedContent.errors[0].code}: ${parsedContent.errors[0].message}`
384 );
385 } else {
386 return new InferenceUpstreamError(content);
387 }
388 } catch {
389 return new InferenceUpstreamError(content);
390 }
391 }
392 
393 async models(
394 params: AiModelsSearchParams = {}
395 ): Promise<AiModelsSearchObject[]> {
396 const url = new URL(`${this.#endpointURL}/ai-api/models/search`);
397 
398 for (const [key, value] of Object.entries(params)) {
399 url.searchParams.set(key, value.toString());
400 }
401 
402 const res = await this.#fetcher.fetch(url, { method: 'GET' });
403 
404 switch (res.status) {
405 case 200: {
406 const data = (await res.json()) as { result: AiModelsSearchObject[] };
407 return data.result;
408 }
409 default: {
410 const data = (await res.json()) as { errors: { message: string }[] };
411 
412 throw new AiInternalError(data.errors[0]?.message || 'Internal Error');
413 }
414 }
415 }
416 
417 toMarkdown(): ToMarkdownService;
418 async toMarkdown(
419 files: MarkdownDocument[],
420 options?: ConversionRequestOptions
421 ): Promise<ConversionResponse[]>;
422 async toMarkdown(
423 files: MarkdownDocument,
424 options?: ConversionRequestOptions
425 ): Promise<ConversionResponse>;
426 toMarkdown(
427 files?: MarkdownDocument | MarkdownDocument[],
428 options?: ConversionRequestOptions
429 ): ToMarkdownService | Promise<ConversionResponse | ConversionResponse[]> {
430 const service = aiBindingExperimental
431 ? this.#fetcher.toMarkdown()
432 : new ToMarkdownService(this.#fetcher);
433 
434 if (arguments.length < 1 || !files) return service;
435 
436 // NOTE(nunopereira): assuming type A = { name: string; blob: Blob }, 'files' here can be of type A | A[].
437 // However, 'service.transform' has no overload that accepts that union, rather it has one overload for each variant.
438 // We know the type of 'files' satisfies whatever type 'service.transform' expects and
439 // instead it's Typescript that is failing to narrow the type, so just ignore this error.
440 // @ts-expect-error unable to narrow type of file to either { name: string; blob: Blob } or { name: string; blob: Blob }[]
441 return service.transform(files, options);
442 }
443 
444 gateway(gatewayId: string, options?: { beta?: boolean }): AiGateway {
445 // Use RPC if the `beta` flag is set or if the "experimental" compat flag is in use. Note that
446 // this is explicitly structured to allow opting out with `{beta: false}`, even if
447 // "experimental" is set, because currently `AbortSignal` does not work over RPC, and Gadgets
448 // needs it (and sets "experimental" for other reasons).
449 if (options?.beta ?? aiBindingExperimental) {
450 return this.#fetcher.gateway(gatewayId);
451 }
452 return new AiGateway(this.#fetcher, gatewayId);
453 }
454 
455 autorag(autoragId?: string): AutoRAG {
456 if (aiBindingExperimental) {
457 return this.#fetcher.autorag(autoragId);
458 }
459 return new AutoRAG(this.#fetcher, autoragId);
460 }
461 
462 aiSearch(): AiSearchService {
463 return this.#fetcher.aiSearch();
464 }
465}
466 
467export default function makeBinding(env: { fetcher: Fetcher }): Ai {
468 return new Ai(env.fetcher);
469}