Skip to content
File

Blob: src/cloudflare/internal/test/ai/ai-api-test.js

javascript454 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 
5import * as assert from 'node:assert';
6 
7export const tests = {
8 async test(_, env) {
9 {
10 // Test ai run response is object
11 const resp = await env.ai.run('testModel', { prompt: 'test' });
12 assert.deepStrictEqual(resp, { response: 'model response' });
13 
14 // Test request id is present
15 assert.deepStrictEqual(
16 env.ai.lastRequestId,
17 '3a1983d7-1ddd-453a-ab75-c4358c91b582'
18 );
19 // Test request http status code is present
20 assert.deepStrictEqual(env.ai.lastRequestHttpStatusCode, 200);
21 }
22 
23 {
24 // Test ai blob model run response is a blob/stream
25 const resp = await env.ai.run('blobResponseModel', { prompt: 'test' });
26 assert.deepStrictEqual(resp instanceof ReadableStream, true);
27 }
28 
29 {
30 // Test legacy fetch
31 const resp = await env.ai.fetch(
32 'http://workers-binding.ai/run?version=2',
33 {
34 method: 'POST',
35 headers: { 'content-type': 'application/json' },
36 body: JSON.stringify({
37 inputs: { prompt: 'test' },
38 options: {},
39 }),
40 }
41 );
42 assert.deepStrictEqual(await resp.json(), { response: 'model response' });
43 }
44 
45 {
46 // Test error response
47 try {
48 await env.ai.run('inputErrorModel', { prompt: 'test' });
49 } catch (e) {
50 assert.deepEqual(
51 {
52 name: e.name,
53 message: e.message,
54 },
55 {
56 name: 'InvalidInput',
57 message: '1001: prompt and messages are mutually exclusive',
58 }
59 );
60 // Test request internal status code is present
61 // eslint-disable-next-line @typescript-eslint/no-unused-expressions
62 assert.deepEqual;
63 assert.deepStrictEqual(env.ai.lastRequestInternalStatusCode, 1001);
64 }
65 }
66 
67 {
68 // Test error properties
69 const err = await env.ai._parseError(
70 Response.json({
71 internalCode: 1001,
72 message: 'InvalidInput: prompt and messages are mutually exclusive',
73 name: 'InvalidInput',
74 description: 'prompt and messages are mutually exclusive',
75 })
76 );
77 assert.equal(err.name, 'InvalidInput');
78 assert.equal(
79 err.message,
80 '1001: prompt and messages are mutually exclusive'
81 );
82 }
83 
84 {
85 // Test error properties from non json response
86 const err = await env.ai._parseError(new Response('Unknown error'));
87 assert.equal(err.name, 'InferenceUpstreamError');
88 assert.equal(err.message, 'Unknown error');
89 }
90 
91 {
92 // Test raw input
93 const resp = await env.ai.run('rawInputs', { prompt: 'test' });
94 
95 assert.deepStrictEqual(resp, {
96 inputs: { prompt: 'test' },
97 options: {},
98 requestUrl: 'https://workers-binding.ai/run?version=3',
99 });
100 }
101 
102 {
103 // Test one readable stream input
104 const encoder = new TextEncoder();
105 const arr = [1, 2, 3];
106 const resp = await env.ai.run('readableStreamIputs', {
107 audio: {
108 body: new ReadableStream({
109 start(controller) {
110 for (const ele of arr) {
111 controller.enqueue(encoder.encode(ele));
112 }
113 controller.close();
114 },
115 }),
116 contentType: 'audio/wav',
117 },
118 });
119 
120 assert.deepStrictEqual(resp, {
121 inputs: {},
122 options: { userInputs: '{}', version: '3' },
123 requestUrl:
124 'https://workers-binding.ai/run?version=3&userInputs=%7B%7D',
125 });
126 }
127 
128 {
129 // Test one readable stream input with additional parameters
130 const encoder = new TextEncoder();
131 const arr = [1, 2, 3];
132 const resp = await env.ai.run('readableStreamIputs', {
133 audio: {
134 body: new ReadableStream({
135 start(controller) {
136 for (const ele of arr) {
137 controller.enqueue(encoder.encode(ele));
138 }
139 controller.close();
140 },
141 }),
142 contentType: 'audio/wav',
143 },
144 detect_language: true,
145 prompt: 'test prompt',
146 });
147 
148 assert.deepStrictEqual(resp, {
149 inputs: {},
150 options: {
151 userInputs: '{"detect_language":true,"prompt":"test prompt"}',
152 version: '3',
153 },
154 requestUrl:
155 'https://workers-binding.ai/run?version=3&userInputs=%7B%22detect_language%22%3Atrue%2C%22prompt%22%3A%22test+prompt%22%7D',
156 });
157 }
158 
159 {
160 // Test errors from one readable stream input without content-type
161 await assert.rejects(
162 async () => {
163 const arr = [1, 2, 3];
164 const encoder = new TextEncoder();
165 const _resp = await env.ai.run('readableStreamIputs', {
166 audio: {
167 body: new ReadableStream({
168 start(controller) {
169 for (const ele of arr) {
170 controller.enqueue(encoder.encode(ele));
171 }
172 controller.close();
173 },
174 }),
175 },
176 });
177 },
178 {
179 name: 'AiInternalError',
180 message: 'Content-Type is required with ReadableStream inputs',
181 }
182 );
183 }
184 
185 {
186 // Test errors from two readable stream inputs
187 await assert.rejects(
188 async () => {
189 const arr = [1, 2, 3];
190 const stream = new ReadableStream({
191 start(controller) {
192 const encoder = new TextEncoder();
193 for (const ele of arr) {
194 controller.enqueue(encoder.encode(ele));
195 }
196 controller.close();
197 },
198 });
199 const _resp = await env.ai.run('readableStreamIputs', {
200 audio: {
201 body: stream,
202 contentType: 'audio/wav',
203 },
204 image: {
205 body: stream,
206 contentType: 'image/png',
207 },
208 });
209 },
210 {
211 name: 'AiInternalError',
212 message:
213 'Multiple ReadableStreams are not supported. Found streams in keys: [audio, image]',
214 }
215 );
216 }
217 
218 {
219 // Test form data input
220 const form = new FormData();
221 form.append('prompt', 'cat');
222 const resp = await env.ai.run('formDataInputs', {
223 audio: {
224 body: form,
225 contentType: 'multipart/form-data',
226 },
227 });
228 
229 assert.deepStrictEqual(resp, {
230 inputs: {},
231 options: { userInputs: '{}', version: '3' },
232 requestUrl:
233 'https://workers-binding.ai/run?version=3&userInputs=%7B%7D',
234 });
235 }
236 
237 {
238 // Test gateway option
239 const resp = await env.ai.run(
240 'rawInputs',
241 { prompt: 'test' },
242 { gateway: { id: 'my-gateway', skipCache: true } }
243 );
244 
245 assert.deepStrictEqual(resp, {
246 inputs: { prompt: 'test' },
247 options: { gateway: { id: 'my-gateway', skipCache: true } },
248 requestUrl: 'https://workers-binding.ai/ai-gateway/run?version=3',
249 });
250 }
251 
252 {
253 // Test unwanted options not getting sent upstream
254 const resp = await env.ai.run(
255 'rawInputs',
256 { prompt: 'test' },
257 {
258 extraHeaders: 'test',
259 example: 123,
260 gateway: { id: 'my-gateway', metadata: { employee: 1233 } },
261 }
262 );
263 
264 assert.deepStrictEqual(resp, {
265 inputs: { prompt: 'test' },
266 options: {
267 example: 123,
268 gateway: { id: 'my-gateway', metadata: { employee: 1233 } },
269 },
270 requestUrl: 'https://workers-binding.ai/ai-gateway/run?version=3',
271 });
272 }
273 
274 {
275 // Test models
276 const resp = await env.ai.models();
277 
278 assert.deepStrictEqual(resp, [
279 {
280 id: 'f8703a00-ed54-4f98-bdc3-cd9a813286f3',
281 source: 1,
282 name: '@cf/qwen/qwen1.5-0.5b-chat',
283 description:
284 'Qwen1.5 is the improved version of Qwen, the large language model series developed by Alibaba Cloud.',
285 task: {
286 id: 'c329a1f9-323d-4e91-b2aa-582dd4188d34',
287 name: 'Text Generation',
288 description:
289 'Family of generative text models, such as large language models (LLM), that can be adapted for a variety of natural language tasks.',
290 },
291 tags: [],
292 properties: [
293 {
294 property_id: 'debug',
295 value: 'https://workers-binding.ai/ai-api/models/search',
296 },
297 ],
298 },
299 ]);
300 }
301 
302 {
303 // Test models with params
304 const resp = await env.ai.models({
305 search: 'test',
306 per_page: 3,
307 page: 1,
308 task: 'asd',
309 });
310 
311 assert.deepStrictEqual(resp, [
312 {
313 id: 'f8703a00-ed54-4f98-bdc3-cd9a813286f3',
314 source: 1,
315 name: '@cf/qwen/qwen1.5-0.5b-chat',
316 description:
317 'Qwen1.5 is the improved version of Qwen, the large language model series developed by Alibaba Cloud.',
318 task: {
319 id: 'c329a1f9-323d-4e91-b2aa-582dd4188d34',
320 name: 'Text Generation',
321 description:
322 'Family of generative text models, such as large language models (LLM), that can be adapted for a variety of natural language tasks.',
323 },
324 tags: [],
325 properties: [
326 {
327 property_id: 'debug',
328 value:
329 'https://workers-binding.ai/ai-api/models/search?search=test&per_page=3&page=1&task=asd',
330 },
331 ],
332 },
333 ]);
334 }
335 
336 {
337 // Test `returnRawResponse` option is returning a Response object
338 const resp = await env.ai.run(
339 'rawInputs',
340 { prompt: 'test' },
341 { returnRawResponse: true }
342 );
343 
344 assert.ok(resp instanceof Response);
345 }
346 
347 {
348 // Test websocket option with basic inputs
349 const resp = await env.ai.run(
350 '@cf/test/websocket',
351 { encoding: 'utf8' },
352 { websocket: true }
353 );
354 assert.deepStrictEqual(resp instanceof Response, true);
355 const respData = await resp.json();
356 assert.deepStrictEqual(respData, {
357 inputs: { encoding: 'utf8' },
358 options: { websocket: true },
359 requestUrl:
360 'https://workers-binding.ai/run?version=3&body=%7B%22inputs%22%3A%7B%22encoding%22%3A%22utf8%22%7D%2C%22options%22%3A%7B%22websocket%22%3Atrue%7D%7D',
361 headers: {
362 'cf-consn-sdk-version': '2.0.0',
363 'cf-consn-model-id': '@cf/test/websocket',
364 upgrade: 'websocket',
365 },
366 });
367 }
368 
369 {
370 // Test signal option is not included in request body
371 const controller = new AbortController();
372 const resp = await env.ai.run(
373 'rawInputs',
374 { prompt: 'test' },
375 { signal: controller.signal }
376 );
377 
378 assert.deepStrictEqual(resp, {
379 inputs: { prompt: 'test' },
380 options: {},
381 requestUrl: 'https://workers-binding.ai/run?version=3',
382 });
383 }
384 
385 {
386 // Test already-aborted signal throws AbortError for websocket requests
387 await assert.rejects(
388 async () => {
389 await env.ai.run(
390 '@cf/test/websocket',
391 { encoding: 'utf8' },
392 { websocket: true, signal: AbortSignal.abort() }
393 );
394 },
395 { name: 'AbortError' }
396 );
397 }
398 
399 {
400 // Test already-aborted signal throws AbortError for readable stream inputs
401 await assert.rejects(
402 async () => {
403 await env.ai.run(
404 'readableStreamIputs',
405 {
406 audio: {
407 body: new ReadableStream({
408 start(controller) {
409 controller.enqueue(new TextEncoder().encode('1'));
410 controller.close();
411 },
412 }),
413 contentType: 'audio/wav',
414 },
415 },
416 { signal: AbortSignal.abort() }
417 );
418 },
419 { name: 'AbortError' }
420 );
421 }
422 
423 {
424 // Test already-aborted signal throws AbortError
425 await assert.rejects(
426 async () => {
427 await env.ai.run(
428 'rawInputs',
429 { prompt: 'test' },
430 { signal: AbortSignal.abort() }
431 );
432 },
433 { name: 'AbortError' }
434 );
435 }
436 
437 {
438 // Test aborting an in-flight request
439 const controller = new AbortController();
440 let resolved = false;
441 const promise = env.ai
442 .run('hangingModel', { prompt: 'test' }, { signal: controller.signal })
443 .finally(() => {
444 resolved = true;
445 });
446 // Wait a moment for the request to start.
447 await scheduler.wait(10);
448 assert.deepStrictEqual(resolved, false);
449 controller.abort();
450 await assert.rejects(promise, { name: 'AbortError' });
451 }
452 },
453};