@tanstack/ai-cloudflare 0.0.0 → 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +21 -0
- package/README.md +53 -2
- package/dist/esm/adapters/embedding.d.ts +19 -0
- package/dist/esm/adapters/embedding.js +60 -0
- package/dist/esm/adapters/embedding.js.map +1 -0
- package/dist/esm/adapters/image.d.ts +26 -0
- package/dist/esm/adapters/image.js +62 -0
- package/dist/esm/adapters/image.js.map +1 -0
- package/dist/esm/adapters/summarize.d.ts +15 -0
- package/dist/esm/adapters/summarize.js +22 -0
- package/dist/esm/adapters/summarize.js.map +1 -0
- package/dist/esm/adapters/text.d.ts +71 -0
- package/dist/esm/adapters/text.js +88 -0
- package/dist/esm/adapters/text.js.map +1 -0
- package/dist/esm/adapters/transcription.d.ts +21 -0
- package/dist/esm/adapters/transcription.js +134 -0
- package/dist/esm/adapters/transcription.js.map +1 -0
- package/dist/esm/adapters/tts.d.ts +26 -0
- package/dist/esm/adapters/tts.js +68 -0
- package/dist/esm/adapters/tts.js.map +1 -0
- package/dist/esm/byok.d.ts +9 -0
- package/dist/esm/byok.js +24 -0
- package/dist/esm/byok.js.map +1 -0
- package/dist/esm/gateway.d.ts +30 -0
- package/dist/esm/gateway.js +45 -0
- package/dist/esm/gateway.js.map +1 -0
- package/dist/esm/index.d.ts +16 -0
- package/dist/esm/index.js +8 -0
- package/dist/esm/utils/config.d.ts +60 -0
- package/dist/esm/utils/config.js +48 -0
- package/dist/esm/utils/config.js.map +1 -0
- package/dist/esm/utils/fetch.d.ts +27 -0
- package/dist/esm/utils/fetch.js +94 -0
- package/dist/esm/utils/fetch.js.map +1 -0
- package/dist/esm/utils/models.d.ts +16 -0
- package/dist/esm/utils/run.d.ts +21 -0
- package/dist/esm/utils/run.js +62 -0
- package/dist/esm/utils/run.js.map +1 -0
- package/package.json +73 -4
- package/src/adapters/embedding.ts +84 -0
- package/src/adapters/image.ts +97 -0
- package/src/adapters/summarize.ts +46 -0
- package/src/adapters/text.ts +150 -0
- package/src/adapters/transcription.ts +227 -0
- package/src/adapters/tts.ts +100 -0
- package/src/byok.ts +21 -0
- package/src/gateway.ts +57 -0
- package/src/index.ts +68 -0
- package/src/utils/config.ts +129 -0
- package/src/utils/fetch.ts +131 -0
- package/src/utils/models.ts +41 -0
- package/src/utils/run.ts +116 -0
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
import type { Ai } from '@cloudflare/workers-types'
|
|
2
|
+
import type { CloudflareGatewayOptions, FetchLike } from './config'
|
|
3
|
+
|
|
4
|
+
/**
|
|
5
|
+
* Workers AI streams end with a usage-only trailer (`{"response":"","usage":
|
|
6
|
+
* {...}}`) that has no `choices` field. The OpenAI Chat Completions stream
|
|
7
|
+
* reader indexes `chunk.choices[0]` on every event, so give such events an
|
|
8
|
+
* empty `choices` array (OpenAI's own usage-only trailer shape) and keep the
|
|
9
|
+
* usage totals they carry.
|
|
10
|
+
*/
|
|
11
|
+
export function normalizeSseResponse(response: Response): Response {
|
|
12
|
+
const contentType = response.headers.get('content-type') ?? ''
|
|
13
|
+
if (!response.body || !contentType.includes('text/event-stream')) {
|
|
14
|
+
return response
|
|
15
|
+
}
|
|
16
|
+
let buffer = ''
|
|
17
|
+
const fixLine = (line: string): string => {
|
|
18
|
+
if (!line.startsWith('data: ') || line === 'data: [DONE]') return line
|
|
19
|
+
try {
|
|
20
|
+
const event = JSON.parse(line.slice(6)) as Record<string, unknown>
|
|
21
|
+
if (event && typeof event === 'object' && !('choices' in event)) {
|
|
22
|
+
return `data: ${JSON.stringify({ ...event, choices: [] })}`
|
|
23
|
+
}
|
|
24
|
+
} catch {
|
|
25
|
+
// Not JSON: forward untouched.
|
|
26
|
+
}
|
|
27
|
+
return line
|
|
28
|
+
}
|
|
29
|
+
const body = response.body
|
|
30
|
+
.pipeThrough(new TextDecoderStream())
|
|
31
|
+
.pipeThrough(
|
|
32
|
+
new TransformStream<string, string>({
|
|
33
|
+
transform(chunk, controller) {
|
|
34
|
+
buffer += chunk
|
|
35
|
+
const lines = buffer.split('\n')
|
|
36
|
+
buffer = lines.pop() ?? ''
|
|
37
|
+
for (const line of lines) controller.enqueue(`${fixLine(line)}\n`)
|
|
38
|
+
},
|
|
39
|
+
flush(controller) {
|
|
40
|
+
if (buffer) controller.enqueue(fixLine(buffer))
|
|
41
|
+
},
|
|
42
|
+
}),
|
|
43
|
+
)
|
|
44
|
+
.pipeThrough(new TextEncoderStream())
|
|
45
|
+
return new Response(body, {
|
|
46
|
+
status: response.status,
|
|
47
|
+
statusText: response.statusText,
|
|
48
|
+
headers: response.headers,
|
|
49
|
+
})
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
/**
|
|
53
|
+
* Cloudflare error bodies look like `{ name, message, internalCode }` or
|
|
54
|
+
* `{ errors: [{ code, message }] }`. The OpenAI SDK only reads
|
|
55
|
+
* `body.error.message`, so rewrap them or every failure reads as
|
|
56
|
+
* "status code (no body)".
|
|
57
|
+
*/
|
|
58
|
+
export async function normalizeErrorResponse(
|
|
59
|
+
response: Response,
|
|
60
|
+
): Promise<Response> {
|
|
61
|
+
const text = await response.text()
|
|
62
|
+
let body = text
|
|
63
|
+
try {
|
|
64
|
+
const json = JSON.parse(text) as {
|
|
65
|
+
error?: unknown
|
|
66
|
+
errors?: Array<{ code?: number; message?: string }>
|
|
67
|
+
message?: string
|
|
68
|
+
name?: string
|
|
69
|
+
internalCode?: number
|
|
70
|
+
}
|
|
71
|
+
if (json && typeof json === 'object' && !('error' in json)) {
|
|
72
|
+
const first = json.errors?.[0]
|
|
73
|
+
body = JSON.stringify({
|
|
74
|
+
error: {
|
|
75
|
+
message: json.message ?? first?.message ?? text,
|
|
76
|
+
type: json.name ?? 'cloudflare_error',
|
|
77
|
+
code: json.internalCode ?? first?.code ?? null,
|
|
78
|
+
},
|
|
79
|
+
})
|
|
80
|
+
}
|
|
81
|
+
} catch {
|
|
82
|
+
// Not JSON: forward the text as-is.
|
|
83
|
+
}
|
|
84
|
+
return new Response(body, {
|
|
85
|
+
status: response.status,
|
|
86
|
+
statusText: response.statusText,
|
|
87
|
+
headers: response.headers,
|
|
88
|
+
})
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
/** Applies the error and SSE normalizations a raw Cloudflare response needs. */
|
|
92
|
+
export async function normalizeResponse(response: Response): Promise<Response> {
|
|
93
|
+
return response.ok
|
|
94
|
+
? normalizeSseResponse(response)
|
|
95
|
+
: await normalizeErrorResponse(response)
|
|
96
|
+
}
|
|
97
|
+
|
|
98
|
+
/**
|
|
99
|
+
* Makes `env.AI` look like an OpenAI-compatible HTTP endpoint to the OpenAI
|
|
100
|
+
* SDK: the JSON request body becomes `binding.run(model, inputs)` and the
|
|
101
|
+
* raw inference `Response` (OpenAI-format JSON or SSE) is handed back.
|
|
102
|
+
*/
|
|
103
|
+
export function createBindingFetch(
|
|
104
|
+
binding: Ai,
|
|
105
|
+
gateway?: CloudflareGatewayOptions,
|
|
106
|
+
): FetchLike {
|
|
107
|
+
// `Ai` is typed against the bundled model catalog; the adapter accepts any
|
|
108
|
+
// model id, so widen the binding to the open catalog shape for this call.
|
|
109
|
+
const run = binding.run.bind(binding) as (
|
|
110
|
+
model: string,
|
|
111
|
+
inputs: Record<string, unknown>,
|
|
112
|
+
options: Record<string, unknown>,
|
|
113
|
+
) => Promise<unknown>
|
|
114
|
+
return async (_input, init) => {
|
|
115
|
+
const { model, ...inputs } = JSON.parse(
|
|
116
|
+
typeof init?.body === 'string' ? init.body : '{}',
|
|
117
|
+
) as { model: string } & Record<string, unknown>
|
|
118
|
+
const response = (await run(model, inputs, {
|
|
119
|
+
returnRawResponse: true,
|
|
120
|
+
...(gateway && { gateway }),
|
|
121
|
+
})) as Response
|
|
122
|
+
return await normalizeResponse(response)
|
|
123
|
+
}
|
|
124
|
+
}
|
|
125
|
+
|
|
126
|
+
/** Wraps a REST fetch so responses get the same error and trailer fixes. */
|
|
127
|
+
export function createRestFetch(baseFetch: FetchLike | undefined): FetchLike {
|
|
128
|
+
const fetchImpl = baseFetch ?? fetch
|
|
129
|
+
return async (input, init) =>
|
|
130
|
+
await normalizeResponse(await fetchImpl(input, init))
|
|
131
|
+
}
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
import type {
|
|
2
|
+
AiModels,
|
|
3
|
+
BaseAiAutomaticSpeechRecognition,
|
|
4
|
+
BaseAiTextEmbeddings,
|
|
5
|
+
BaseAiTextGeneration,
|
|
6
|
+
BaseAiTextToImage,
|
|
7
|
+
BaseAiTextToSpeech,
|
|
8
|
+
} from '@cloudflare/workers-types'
|
|
9
|
+
|
|
10
|
+
/** Model ids from the Workers AI catalog whose task shape matches `TTask`. */
|
|
11
|
+
type ModelsFor<TTask> = {
|
|
12
|
+
[K in keyof AiModels]: AiModels[K] extends TTask ? K : never
|
|
13
|
+
}[keyof AiModels]
|
|
14
|
+
|
|
15
|
+
/**
|
|
16
|
+
* Chat model id. Catalog ids get autocomplete; any other id works too,
|
|
17
|
+
* including third-party `provider/model` ids routed through AI Gateway
|
|
18
|
+
* (for example `openai/gpt-5.5`).
|
|
19
|
+
*/
|
|
20
|
+
export type CloudflareTextModel =
|
|
21
|
+
| ModelsFor<BaseAiTextGeneration>
|
|
22
|
+
| (string & {})
|
|
23
|
+
|
|
24
|
+
export type CloudflareEmbeddingModel =
|
|
25
|
+
| ModelsFor<BaseAiTextEmbeddings>
|
|
26
|
+
| (string & {})
|
|
27
|
+
|
|
28
|
+
export type CloudflareImageModel = ModelsFor<BaseAiTextToImage> | (string & {})
|
|
29
|
+
|
|
30
|
+
export type CloudflareTTSModel =
|
|
31
|
+
| ModelsFor<BaseAiTextToSpeech>
|
|
32
|
+
| '@cf/deepgram/aura-1'
|
|
33
|
+
| '@cf/deepgram/aura-2-en'
|
|
34
|
+
| '@cf/deepgram/aura-2-es'
|
|
35
|
+
| (string & {})
|
|
36
|
+
|
|
37
|
+
export type CloudflareTranscriptionModel =
|
|
38
|
+
| ModelsFor<BaseAiAutomaticSpeechRecognition>
|
|
39
|
+
| '@cf/openai/whisper-large-v3-turbo'
|
|
40
|
+
| '@cf/deepgram/nova-3'
|
|
41
|
+
| (string & {})
|
package/src/utils/run.ts
ADDED
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
import { arrayBufferToBase64 } from '@tanstack/ai-utils'
|
|
2
|
+
import { CLOUDFLARE_API_BASE, gatewayHeaders, isBindingConfig } from './config'
|
|
3
|
+
import type { CloudflareConfig } from './config'
|
|
4
|
+
|
|
5
|
+
export type RunInputs = Record<string, unknown>
|
|
6
|
+
|
|
7
|
+
export interface RunBinary {
|
|
8
|
+
/** Input field that carries the bytes on the binding path. */
|
|
9
|
+
field: string
|
|
10
|
+
body: Uint8Array | ArrayBuffer | Blob
|
|
11
|
+
contentType: string
|
|
12
|
+
}
|
|
13
|
+
|
|
14
|
+
/**
|
|
15
|
+
* Runs a Workers AI model with its native task inputs (embeddings, image,
|
|
16
|
+
* speech, transcription) through the binding or the REST `/ai/run` endpoint.
|
|
17
|
+
*
|
|
18
|
+
* Returns the model's decoded output: an object for JSON tasks, or bytes
|
|
19
|
+
* (`Uint8Array` / `ReadableStream`) for binary media outputs.
|
|
20
|
+
*/
|
|
21
|
+
export async function runModel(
|
|
22
|
+
config: CloudflareConfig,
|
|
23
|
+
model: string,
|
|
24
|
+
inputs: RunInputs,
|
|
25
|
+
options?: { signal?: AbortSignal; binary?: RunBinary },
|
|
26
|
+
): Promise<unknown> {
|
|
27
|
+
if (isBindingConfig(config)) {
|
|
28
|
+
const run = config.binding.run.bind(config.binding) as (
|
|
29
|
+
model: string,
|
|
30
|
+
inputs: RunInputs,
|
|
31
|
+
options?: Record<string, unknown>,
|
|
32
|
+
) => Promise<unknown>
|
|
33
|
+
// The binding serializes inputs as JSON; binary bodies must travel as a
|
|
34
|
+
// ReadableStream, which it forwards as the raw request body.
|
|
35
|
+
const bindingInputs = options?.binary
|
|
36
|
+
? {
|
|
37
|
+
...inputs,
|
|
38
|
+
[options.binary.field]: {
|
|
39
|
+
body: new Response(options.binary.body as BodyInit).body,
|
|
40
|
+
contentType: options.binary.contentType,
|
|
41
|
+
},
|
|
42
|
+
}
|
|
43
|
+
: inputs
|
|
44
|
+
return await run(
|
|
45
|
+
model,
|
|
46
|
+
bindingInputs,
|
|
47
|
+
config.gateway ? { gateway: config.gateway } : undefined,
|
|
48
|
+
)
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
const url = new URL(
|
|
52
|
+
`${CLOUDFLARE_API_BASE}/accounts/${config.accountId}/ai/run/${model}`,
|
|
53
|
+
)
|
|
54
|
+
const headers: Record<string, string> = {
|
|
55
|
+
Authorization: `Bearer ${config.apiKey}`,
|
|
56
|
+
...gatewayHeaders(config.gateway),
|
|
57
|
+
}
|
|
58
|
+
let body: BodyInit
|
|
59
|
+
if (options?.binary) {
|
|
60
|
+
// Binary tasks take the bytes as the body and the other inputs as query.
|
|
61
|
+
for (const [key, value] of Object.entries(inputs)) {
|
|
62
|
+
if (value !== undefined) url.searchParams.set(key, String(value))
|
|
63
|
+
}
|
|
64
|
+
headers['Content-Type'] = options.binary.contentType
|
|
65
|
+
body = options.binary.body as BodyInit
|
|
66
|
+
} else {
|
|
67
|
+
headers['Content-Type'] = 'application/json'
|
|
68
|
+
body = JSON.stringify(inputs)
|
|
69
|
+
}
|
|
70
|
+
const fetchImpl = config.fetch ?? fetch
|
|
71
|
+
const response = await fetchImpl(url, {
|
|
72
|
+
method: 'POST',
|
|
73
|
+
headers,
|
|
74
|
+
body,
|
|
75
|
+
signal: options?.signal,
|
|
76
|
+
})
|
|
77
|
+
if (!response.ok) {
|
|
78
|
+
throw new Error(
|
|
79
|
+
`Workers AI request for ${model} failed (${response.status}): ${await response.text()}`,
|
|
80
|
+
)
|
|
81
|
+
}
|
|
82
|
+
if (response.headers.get('content-type')?.includes('application/json')) {
|
|
83
|
+
const json = (await response.json()) as {
|
|
84
|
+
success?: boolean
|
|
85
|
+
result?: unknown
|
|
86
|
+
errors?: Array<{ message?: string }>
|
|
87
|
+
}
|
|
88
|
+
if (json.success === false) {
|
|
89
|
+
throw new Error(
|
|
90
|
+
`Workers AI request for ${model} failed: ${json.errors?.map((e) => e.message).join('; ')}`,
|
|
91
|
+
)
|
|
92
|
+
}
|
|
93
|
+
return 'result' in json ? json.result : json
|
|
94
|
+
}
|
|
95
|
+
return new Uint8Array(await response.arrayBuffer())
|
|
96
|
+
}
|
|
97
|
+
|
|
98
|
+
/** Base64-encodes a binary model output, whatever shape it arrived in. */
|
|
99
|
+
export async function outputToBase64(output: unknown): Promise<string> {
|
|
100
|
+
if (typeof output === 'string') return output
|
|
101
|
+
let bytes: Uint8Array
|
|
102
|
+
if (output instanceof Uint8Array) {
|
|
103
|
+
bytes = output
|
|
104
|
+
} else if (output instanceof ArrayBuffer) {
|
|
105
|
+
bytes = new Uint8Array(output)
|
|
106
|
+
} else if (output instanceof ReadableStream) {
|
|
107
|
+
bytes = new Uint8Array(await new Response(output).arrayBuffer())
|
|
108
|
+
} else {
|
|
109
|
+
throw new Error(
|
|
110
|
+
`Unexpected Workers AI output type: ${Object.prototype.toString.call(output)}`,
|
|
111
|
+
)
|
|
112
|
+
}
|
|
113
|
+
const copy = new Uint8Array(bytes.byteLength)
|
|
114
|
+
copy.set(bytes)
|
|
115
|
+
return arrayBufferToBase64(copy.buffer)
|
|
116
|
+
}
|