Skip to content
File

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

typescript339 lines
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 
5interface Fetcher {
6 fetch: typeof fetch;
7}
8 
9type GatewayRetries = {
10 maxAttempts?: 1 | 2 | 3 | 4 | 5;
11 retryDelayMs?: number;
12 backoff?: 'constant' | 'linear' | 'exponential';
13};
14 
15export 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 
27export type UniversalGatewayOptions = Omit<GatewayOptions, 'id'> & {
28 /**
29 ** @deprecated
30 */
31 id?: string;
32};
33 
34export type AiGatewayPatchLog = {
35 score?: number | null;
36 feedback?: -1 | 1 | null;
37 metadata?: Record<string, number | string | boolean | null | bigint> | null;
38};
39 
40export 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 
68export 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 
90export 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 
112export type AIGatewayUniversalRequest = {
113 provider: AIGatewayProviders | string; // eslint-disable-line
114 endpoint: string;
115 headers: Partial<AIGatewayHeaders>;
116 query: unknown;
117};
118 
119export class AiGatewayInternalError extends Error {
120 constructor(message: string) {
121 super(message);
122 this.name = 'AiGatewayInternalError';
123 }
124}
125 
126export class AiGatewayLogNotFound extends Error {
127 constructor(message: string) {
128 super(message);
129 this.name = 'AiGatewayLogNotFound';
130 }
131}
132 
133async 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 
151export 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}