File
Blob: src/cloudflare/internal/ai-api.ts
| 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 | |
| 5 | import { AiGateway, type GatewayOptions } from 'cloudflare-internal:aig-api'; |
| 6 | import { AutoRAG } from 'cloudflare-internal:autorag-api'; |
| 7 | import { |
| 8 | ToMarkdownService, |
| 9 | type ConversionRequestOptions, |
| 10 | type ConversionResponse, |
| 11 | type MarkdownDocument, |
| 12 | } from 'cloudflare-internal:to-markdown-api'; |
| 13 | |
| 14 | type AiSearchService = object; |
| 15 | |
| 16 | const aiBindingExperimental = !!Cloudflare.compatibilityFlags['experimental']; |
| 17 | |
| 18 | interface Fetcher { |
| 19 | fetch: typeof fetch; |
| 20 | aiSearch: () => AiSearchService; |
| 21 | gateway: (gatewayId: string) => AiGateway; |
| 22 | autorag: (autoragId?: string) => AutoRAG; |
| 23 | toMarkdown: () => ToMarkdownService; |
| 24 | } |
| 25 | |
| 26 | interface AiError { |
| 27 | internalCode: number; |
| 28 | message: string; |
| 29 | name: string; |
| 30 | description: string; |
| 31 | errors?: Array<{ code: number; message: string }>; |
| 32 | } |
| 33 | |
| 34 | export type SessionOptions = { |
| 35 | // Deprecated, do not use this |
| 36 | extraHeaders?: object; |
| 37 | }; |
| 38 | |
| 39 | export 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 | |
| 53 | type CleanedAiOptions = Omit< |
| 54 | AiOptions, |
| 55 | 'prefix' | 'extraHeaders' | 'sessionOptions' | 'signal' |
| 56 | >; |
| 57 | |
| 58 | export type AiInputReadableStream = { |
| 59 | body: ReadableStream | FormData; |
| 60 | contentType: string; |
| 61 | }; |
| 62 | |
| 63 | export 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 | |
| 73 | export 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 | |
| 90 | export class InferenceUpstreamError extends Error { |
| 91 | constructor(message: string, name = 'InferenceUpstreamError') { |
| 92 | super(message); |
| 93 | this.name = name; |
| 94 | } |
| 95 | } |
| 96 | |
| 97 | export class AiInternalError extends Error { |
| 98 | constructor(message: string, name = 'AiInternalError') { |
| 99 | super(message); |
| 100 | this.name = name; |
| 101 | } |
| 102 | } |
| 103 | |
| 104 | function isReadableStream(obj: unknown): obj is ReadableStream { |
| 105 | return obj instanceof ReadableStream; |
| 106 | } |
| 107 | |
| 108 | function isFormData(obj: unknown): obj is FormData { |
| 109 | return obj instanceof FormData; |
| 110 | } |
| 111 | |
| 112 | /** |
| 113 | * Find keys in inputs that have a ReadableStream |
| 114 | * */ |
| 115 | function 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 | |
| 136 | export 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 | |
| 467 | export default function makeBinding(env: { fetcher: Fetcher }): Ai { |
| 468 | return new Ai(env.fetcher); |
| 469 | } |