File
Blob: src/cloudflare/internal/aig-api.ts
| 1 | // Copyright (c) 2024 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 | interface Fetcher { |
| 6 | fetch: typeof fetch; |
| 7 | } |
| 8 | |
| 9 | type GatewayRetries = { |
| 10 | maxAttempts?: 1 | 2 | 3 | 4 | 5; |
| 11 | retryDelayMs?: number; |
| 12 | backoff?: 'constant' | 'linear' | 'exponential'; |
| 13 | }; |
| 14 | |
| 15 | export type GatewayOptions = { |
| 16 | id: string; |
| 17 | cacheKey?: string; |
| 18 | cacheTtl?: number; |
| 19 | skipCache?: boolean; |
| 20 | metadata?: Record<string, number | string | boolean | null | bigint>; |
| 21 | collectLog?: boolean; |
| 22 | eventId?: string; |
| 23 | requestTimeoutMs?: number; |
| 24 | retries?: GatewayRetries; |
| 25 | }; |
| 26 | |
| 27 | export type UniversalGatewayOptions = Omit<GatewayOptions, 'id'> & { |
| 28 | /** |
| 29 | ** @deprecated |
| 30 | */ |
| 31 | id?: string; |
| 32 | }; |
| 33 | |
| 34 | export type AiGatewayPatchLog = { |
| 35 | score?: number | null; |
| 36 | feedback?: -1 | 1 | null; |
| 37 | metadata?: Record<string, number | string | boolean | null | bigint> | null; |
| 38 | }; |
| 39 | |
| 40 | export type AiGatewayLog = { |
| 41 | id: string; |
| 42 | provider: string; |
| 43 | model: string; |
| 44 | model_type?: string; |
| 45 | path: string; |
| 46 | duration: number; |
| 47 | request_type?: string; |
| 48 | request_content_type?: string; |
| 49 | status_code: number; |
| 50 | response_content_type?: string; |
| 51 | success: boolean; |
| 52 | cached: boolean; |
| 53 | tokens_in?: number; |
| 54 | tokens_out?: number; |
| 55 | metadata?: Record<string, number | string | boolean | null | bigint>; |
| 56 | step?: number; |
| 57 | cost?: number; |
| 58 | custom_cost?: boolean; |
| 59 | request_size: number; |
| 60 | request_head?: string; |
| 61 | request_head_complete: boolean; |
| 62 | response_size: number; |
| 63 | response_head?: string; |
| 64 | response_head_complete: boolean; |
| 65 | created_at: Date; |
| 66 | }; |
| 67 | |
| 68 | export type AIGatewayProviders = |
| 69 | | 'workers-ai' |
| 70 | | 'anthropic' |
| 71 | | 'aws-bedrock' |
| 72 | | 'azure-openai' |
| 73 | | 'google-vertex-ai' |
| 74 | | 'huggingface' |
| 75 | | 'openai' |
| 76 | | 'perplexity-ai' |
| 77 | | 'replicate' |
| 78 | | 'groq' |
| 79 | | 'cohere' |
| 80 | | 'google-ai-studio' |
| 81 | | 'mistral' |
| 82 | | 'grok' |
| 83 | | 'openrouter' |
| 84 | | 'deepseek' |
| 85 | | 'cerebras' |
| 86 | | 'cartesia' |
| 87 | | 'elevenlabs' |
| 88 | | 'adobe-firefly'; |
| 89 | |
| 90 | export type AIGatewayHeaders = { |
| 91 | 'cf-aig-metadata': |
| 92 | | Record<string, number | string | boolean | null | bigint> |
| 93 | | string; |
| 94 | 'cf-aig-custom-cost': |
| 95 | | { per_token_in?: number; per_token_out?: number } |
| 96 | | { total_cost?: number } |
| 97 | | string; |
| 98 | 'cf-aig-cache-ttl': number | string; |
| 99 | 'cf-aig-skip-cache': boolean | string; |
| 100 | 'cf-aig-cache-key': string; |
| 101 | 'cf-aig-event-id': string; |
| 102 | 'cf-aig-request-timeout': number | string; |
| 103 | 'cf-aig-max-attempts': number | string; |
| 104 | 'cf-aig-retry-delay': number | string; |
| 105 | 'cf-aig-backoff': string; |
| 106 | 'cf-aig-collect-log': boolean | string; |
| 107 | Authorization: string; |
| 108 | 'Content-Type': string; |
| 109 | [key: string]: string | number | boolean | object; |
| 110 | }; |
| 111 | |
| 112 | export type AIGatewayUniversalRequest = { |
| 113 | provider: AIGatewayProviders | string; // eslint-disable-line |
| 114 | endpoint: string; |
| 115 | headers: Partial<AIGatewayHeaders>; |
| 116 | query: unknown; |
| 117 | }; |
| 118 | |
| 119 | export class AiGatewayInternalError extends Error { |
| 120 | constructor(message: string) { |
| 121 | super(message); |
| 122 | this.name = 'AiGatewayInternalError'; |
| 123 | } |
| 124 | } |
| 125 | |
| 126 | export class AiGatewayLogNotFound extends Error { |
| 127 | constructor(message: string) { |
| 128 | super(message); |
| 129 | this.name = 'AiGatewayLogNotFound'; |
| 130 | } |
| 131 | } |
| 132 | |
| 133 | async function parseError( |
| 134 | res: Response, |
| 135 | defaultMsg = 'Internal Error', |
| 136 | errorCls = AiGatewayInternalError |
| 137 | ): Promise<Error> { |
| 138 | const content = await res.text(); |
| 139 | |
| 140 | try { |
| 141 | const parsedContent = JSON.parse(content) as { |
| 142 | errors: { message: string }[]; |
| 143 | }; |
| 144 | |
| 145 | return new errorCls(parsedContent.errors.at(0)?.message || defaultMsg); |
| 146 | } catch { |
| 147 | return new AiGatewayInternalError(content); |
| 148 | } |
| 149 | } |
| 150 | |
| 151 | export class AiGateway { |
| 152 | readonly #fetcher: Fetcher; |
| 153 | readonly #gatewayId: string; |
| 154 | |
| 155 | constructor(fetcher: Fetcher, gatewayId: string) { |
| 156 | this.#fetcher = fetcher; |
| 157 | this.#gatewayId = gatewayId; |
| 158 | } |
| 159 | |
| 160 | // eslint-disable-next-line |
| 161 | async getUrl(provider?: AIGatewayProviders | string): Promise<string> { |
| 162 | const res = await this.#fetcher.fetch( |
| 163 | `https://workers-binding.ai/ai-gateway/gateways/${this.#gatewayId}/url/${provider ?? 'universal'}`, |
| 164 | { method: 'GET' } |
| 165 | ); |
| 166 | |
| 167 | if (!res.ok) { |
| 168 | throw await parseError(res); |
| 169 | } |
| 170 | |
| 171 | const data = (await res.json()) as { result: { url: string } }; |
| 172 | |
| 173 | return data.result.url; |
| 174 | } |
| 175 | |
| 176 | async getLog(logId: string): Promise<AiGatewayLog> { |
| 177 | const res = await this.#fetcher.fetch( |
| 178 | `https://workers-binding.ai/ai-gateway/gateways/${this.#gatewayId}/logs/${logId}`, |
| 179 | { |
| 180 | method: 'GET', |
| 181 | } |
| 182 | ); |
| 183 | |
| 184 | switch (res.status) { |
| 185 | case 200: { |
| 186 | const data = (await res.json()) as { result: AiGatewayLog }; |
| 187 | |
| 188 | return { |
| 189 | ...data.result, |
| 190 | created_at: new Date(data.result.created_at), |
| 191 | }; |
| 192 | } |
| 193 | case 404: { |
| 194 | throw await parseError(res, 'Log Not Found', AiGatewayLogNotFound); |
| 195 | } |
| 196 | default: { |
| 197 | throw await parseError(res); |
| 198 | } |
| 199 | } |
| 200 | } |
| 201 | |
| 202 | async patchLog(logId: string, data: AiGatewayPatchLog): Promise<void> { |
| 203 | const res = await this.#fetcher.fetch( |
| 204 | `https://workers-binding.ai/ai-gateway/gateways/${this.#gatewayId}/logs/${logId}`, |
| 205 | { |
| 206 | method: 'PATCH', |
| 207 | body: JSON.stringify(data), |
| 208 | headers: { |
| 209 | 'content-type': 'application/json', |
| 210 | }, |
| 211 | } |
| 212 | ); |
| 213 | |
| 214 | switch (res.status) { |
| 215 | case 200: { |
| 216 | return; |
| 217 | } |
| 218 | case 404: { |
| 219 | throw await parseError(res, 'Log Not Found', AiGatewayLogNotFound); |
| 220 | } |
| 221 | default: { |
| 222 | throw await parseError(res); |
| 223 | } |
| 224 | } |
| 225 | } |
| 226 | |
| 227 | run( |
| 228 | data: AIGatewayUniversalRequest | AIGatewayUniversalRequest[], |
| 229 | options?: { |
| 230 | gateway?: UniversalGatewayOptions; |
| 231 | extraHeaders?: Record<string, string>; |
| 232 | signal?: AbortSignal; |
| 233 | } |
| 234 | ): Promise<Response> { |
| 235 | const input = Array.isArray(data) ? data : [data]; |
| 236 | |
| 237 | const headers = this.#getHeadersFromOptions( |
| 238 | options?.gateway, |
| 239 | options?.extraHeaders |
| 240 | ); |
| 241 | |
| 242 | // Convert header values to string |
| 243 | for (const req of input) { |
| 244 | for (const [k, v] of Object.entries(req.headers)) { |
| 245 | if (typeof v === 'number' || typeof v === 'boolean') { |
| 246 | req.headers[k] = v.toString(); |
| 247 | // eslint-disable-next-line |
| 248 | } else if (typeof v === 'object' && v != null) { |
| 249 | req.headers[k] = JSON.stringify(v); |
| 250 | } |
| 251 | } |
| 252 | } |
| 253 | |
| 254 | const fetchOptions: RequestInit = { |
| 255 | method: 'POST', |
| 256 | body: JSON.stringify(input), |
| 257 | headers: headers, |
| 258 | }; |
| 259 | if (options?.signal) { |
| 260 | fetchOptions.signal = options.signal; |
| 261 | } |
| 262 | |
| 263 | return this.#fetcher.fetch( |
| 264 | `https://workers-binding.ai/ai-gateway/universal/run/${this.#gatewayId}`, |
| 265 | fetchOptions |
| 266 | ); |
| 267 | } |
| 268 | |
| 269 | #getHeadersFromOptions( |
| 270 | options?: UniversalGatewayOptions, |
| 271 | extraHeaders?: Record<string, string> |
| 272 | ): Headers { |
| 273 | const headers = new Headers(); |
| 274 | headers.set('content-type', 'application/json'); |
| 275 | |
| 276 | if (options) { |
| 277 | if (options.skipCache !== undefined) { |
| 278 | headers.set('cf-aig-skip-cache', options.skipCache ? 'true' : 'false'); |
| 279 | } |
| 280 | |
| 281 | if (options.cacheTtl) { |
| 282 | headers.set('cf-aig-cache-ttl', options.cacheTtl.toString()); |
| 283 | } |
| 284 | |
| 285 | if (options.metadata) { |
| 286 | headers.set('cf-aig-metadata', JSON.stringify(options.metadata)); |
| 287 | } |
| 288 | |
| 289 | if (options.cacheKey) { |
| 290 | headers.set('cf-aig-cache-key', options.cacheKey); |
| 291 | } |
| 292 | |
| 293 | if (options.collectLog !== undefined) { |
| 294 | headers.set( |
| 295 | 'cf-aig-collect-log', |
| 296 | options.collectLog ? 'true' : 'false' |
| 297 | ); |
| 298 | } |
| 299 | |
| 300 | if (options.eventId !== undefined) { |
| 301 | headers.set('cf-aig-event-id', options.eventId); |
| 302 | } |
| 303 | |
| 304 | if (options.requestTimeoutMs !== undefined) { |
| 305 | headers.set( |
| 306 | 'cf-aig-request-timeout', |
| 307 | options.requestTimeoutMs.toString() |
| 308 | ); |
| 309 | } |
| 310 | |
| 311 | if (options.retries !== undefined) { |
| 312 | if (options.retries.maxAttempts !== undefined) { |
| 313 | headers.set( |
| 314 | 'cf-aig-max-attempts', |
| 315 | options.retries.maxAttempts.toString() |
| 316 | ); |
| 317 | } |
| 318 | if (options.retries.retryDelayMs !== undefined) { |
| 319 | headers.set( |
| 320 | 'cf-aig-retry-delay', |
| 321 | options.retries.retryDelayMs.toString() |
| 322 | ); |
| 323 | } |
| 324 | if (options.retries.backoff !== undefined) { |
| 325 | headers.set('cf-aig-backoff', options.retries.backoff); |
| 326 | } |
| 327 | } |
| 328 | } |
| 329 | |
| 330 | if (extraHeaders) { |
| 331 | for (const [key, value] of Object.entries(extraHeaders)) { |
| 332 | headers.set(key, value); |
| 333 | } |
| 334 | } |
| 335 | |
| 336 | return headers; |
| 337 | } |
| 338 | } |