Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions package.json
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,6 @@
"sharp@<0.35.0": "0.35.3",
"shell-quote": "1.9.0",
"svgo@>=4.0.0 <4.0.2": "4.0.2",
"undici@>=7.0.0 <7.29.0": "7.29.0",
"ws@>=8.0.0 <8.21.0": "8.21.0",
"browserslist@<4.28.9": ">=4.28.9",
"react-pdf>pdfjs-dist": "5.4.624",
Expand All @@ -137,7 +136,8 @@
"brace-expansion@>=1.0.0 <1.1.21": "1.1.21",
"brace-expansion@>=2.0.0 <2.1.7": "2.1.7",
"brace-expansion@>=5.0.0 <5.0.12": "5.0.12",
"undici@>=7.0.0 <7.29.1": "7.29.1"
"undici@>=7.0.0 <7.29.1": "7.29.1",
"devalue@<=5.9.2": "5.9.3"
}
}
}
10 changes: 5 additions & 5 deletions pnpm-lock.yaml

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

109 changes: 109 additions & 0 deletions src/lib/ai-cloudflare.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
import { generateText } from 'ai';
import { describe, expect, it, vi } from 'vitest';
import { getLanguageModel } from './ai-cloudflare';
import { findSharedAiBudgetDenied, SharedAiBudgetDenied } from './shared-ai-budget';

const MODEL = '@cf/meta/llama-3.3-70b-instruct-fp8-fast';
const DAILY_CAP = 9_500;

function makeNamespace(reply?: unknown) {
let used = 0;
const fetch = vi.fn(async (_url: string, init?: RequestInit) => {
const request = JSON.parse(String(init?.body)) as { neurons: number };
used += request.neurons;
if (reply !== undefined) return Response.json(reply);
return Response.json({
allowed: true,
used,
remaining: DAILY_CAP - used,
retryAfter: 0,
dayKey: new Date().toISOString().slice(0, 10),
});
});
const namespace = {
idFromName: (name: string) => name,
get: () => ({ fetch }),
} as unknown as DurableObjectNamespace;
return { namespace, fetch };
}

function makeBinding(run: ReturnType<typeof vi.fn>) {
return { run, marker: 'preserve-this-receiver' } as unknown as Ai;
}

describe('Workers AI language model budget integration', () => {
it('reserves each actual SDK retry before calling the binding', async () => {
const { namespace, fetch } = makeNamespace();
let receiver: unknown;
const run = vi.fn(function (this: unknown) {
receiver = this;
if (run.mock.calls.length === 1) {
throw Object.assign(new Error('synthetic retryable rate-limit response'), { code: 3036 });
}
return Promise.resolve({
choices: [
{ message: { role: 'assistant', content: 'Synthetic answer.' }, finish_reason: 'stop' },
],
});
});
const model = getLanguageModel({
binding: makeBinding(run),
budgetNamespace: namespace,
endpointUrl: '',
apiKey: '',
model: MODEL,
});

const result = await generateText({
model,
prompt: 'Synthetic UTF-8 probe: café ⚓',
maxOutputTokens: 512,
maxRetries: 1,
});

expect(result.text).toBe('Synthetic answer.');
expect(run).toHaveBeenCalledTimes(2);
expect(fetch).toHaveBeenCalledTimes(2);
expect(receiver).toBe(run.mock.contexts[1]);
expect(run.mock.contexts[0]).toBe(receiver);
expect(run.mock.calls.every(([, input]) => input.max_tokens === 512)).toBe(true);
});

it('returns a typed budget denial without binding inference when the receipt is malformed', async () => {
const { namespace, fetch } = makeNamespace(null);
const run = vi.fn();
const model = getLanguageModel({
binding: makeBinding(run),
budgetNamespace: namespace,
endpointUrl: '',
apiKey: '',
model: MODEL,
});

let caught: unknown;
try {
await generateText({ model, prompt: 'Synthetic only.', maxRetries: 0 });
} catch (error) {
caught = error;
}
expect(findSharedAiBudgetDenied(caught)).toBeInstanceOf(SharedAiBudgetDenied);
expect(fetch).toHaveBeenCalledOnce();
expect(run).not.toHaveBeenCalled();
});

it('keeps explicit BYOK ahead of the Cloudflare binding and shared budget', () => {
const { namespace, fetch } = makeNamespace();
const run = vi.fn();
const model = getLanguageModel({
binding: makeBinding(run),
budgetNamespace: namespace,
endpointUrl: 'https://provider.example/v1',
apiKey: 'synthetic-test-key',
model: 'provider-model',
});

expect(model.modelId).toBe('provider-model');
expect(fetch).not.toHaveBeenCalled();
expect(run).not.toHaveBeenCalled();
});
});
13 changes: 9 additions & 4 deletions src/lib/ai-cloudflare.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ import type { LanguageModel } from 'ai';
import { createWorkersAI, type WorkersAISettings } from 'workers-ai-provider';

import type { AIConfig } from './ai-vendor';
import { createBudgetedWorkersAiBinding, type SharedBudgetNamespace } from './shared-ai-budget';

type WorkersAiBinding = Extract<WorkersAISettings, { binding: unknown }>['binding'];

Expand All @@ -26,12 +27,13 @@ function createAIModel(
/** Default model when the project's direct endpoint is Workers AI. */
const DEFAULT_WORKERS_AI_MODEL = '@cf/meta/llama-3.3-70b-instruct-fp8-fast';

interface CreateLanguageModelArgs {
interface CreateLanguageModelArgs<Id = never> {
binding?: WorkersAiBinding;
endpointUrl: string;
apiKey: string;
model: string;
headers?: Record<string, string>;
budgetNamespace?: SharedBudgetNamespace<Id>;
}

function getDirectBaseUrl(): string {
Expand All @@ -50,20 +52,23 @@ function getDirectApiKey(): string {
* Returns a model for an explicit BYOK endpoint or the project's own direct
* free-provider/local endpoint. No shared gateway fallback exists.
*/
export function getLanguageModel({
export function getLanguageModel<Id = never>({
binding,
endpointUrl,
apiKey,
model,
headers,
}: CreateLanguageModelArgs): LanguageModel {
budgetNamespace,
}: CreateLanguageModelArgs<Id>): LanguageModel {
// Honour explicit BYO config first (settings UI etc.).
if (endpointUrl && apiKey) {
return createAIModel({ endpointUrl, apiKey, model } as AIConfig, { headers });
}

if (binding) {
return createWorkersAI({ binding })(model || DEFAULT_WORKERS_AI_MODEL);
return createWorkersAI({
binding: createBudgetedWorkersAiBinding(binding, budgetNamespace),
})(model || DEFAULT_WORKERS_AI_MODEL);
}

const resolvedModel = model || DEFAULT_WORKERS_AI_MODEL;
Expand Down
111 changes: 111 additions & 0 deletions src/lib/shared-ai-budget.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
import { describe, expect, it, vi } from 'vitest';
import { createBudgetedWorkersAiBinding, SharedAiBudgetDenied } from './shared-ai-budget';

const MODEL = '@cf/meta/llama-3.3-70b-instruct-fp8-fast';
const DAILY_CAP = 9_500;

function makeBudget(reply?: unknown) {
const fetch = vi.fn(async (_url: string, init?: RequestInit) => {
if (reply !== undefined) return Response.json(reply);
const request = JSON.parse(String(init?.body)) as { neurons: number };
return Response.json({
allowed: true,
used: request.neurons,
remaining: DAILY_CAP - request.neurons,
retryAfter: 0,
dayKey: new Date().toISOString().slice(0, 10),
});
});
const namespace = {
idFromName: vi.fn((name: string) => name),
get: vi.fn(() => ({ fetch })),
} as unknown as DurableObjectNamespace;
return { namespace, fetch };
}

function makeBinding() {
let receiver: unknown;
const binding = {
marker: 'binding-receiver',
run: vi.fn(function (
this: { marker: string },
_model: string,
_input: Record<string, unknown>
) {
receiver = this;
return Promise.resolve({ response: 'synthetic result' });
}),
};
return { binding, run: binding.run, getReceiver: () => receiver };
}

describe('shared Workers AI budget guard', () => {
it('reserves UTF-8 serialized input, applies the 512 default, and preserves the binding receiver', async () => {
const { namespace, fetch } = makeBudget();
const { binding, run, getReceiver } = makeBinding();
const guarded = createBudgetedWorkersAiBinding(binding as unknown as Ai, namespace);
const input = { messages: [{ role: 'user', content: 'café ⚓' }] };
const boundedInput = { ...input, max_tokens: 512 };

await guarded.run(MODEL, input);

const body = JSON.parse(String(fetch.mock.calls[0][1]?.body)) as { neurons: number };
const bytes = new TextEncoder().encode(JSON.stringify(boundedInput)).byteLength;
const inputTokens = Math.ceil(bytes * 1.2);
const expected = Math.ceil((inputTokens * 26_668 + 512 * 204_805) / 1_000_000);
expect(bytes).toBeGreaterThan(JSON.stringify(boundedInput).length);
expect(body.neurons).toBe(expected);
expect(run).toHaveBeenCalledOnce();
expect(run.mock.calls[0][1]).toMatchObject(boundedInput);
expect(getReceiver()).toBe(binding);
expect(namespace.idFromName).toHaveBeenCalledWith('global-budget');
});

it.each([0, -1, 8_193, '512', Number.NaN])(
'rejects invalid output token bound %j before debit or inference',
async (max_tokens) => {
const { namespace, fetch } = makeBudget();
const { binding, run } = makeBinding();
const guarded = createBudgetedWorkersAiBinding(binding as unknown as Ai, namespace);

await expect(
guarded.run(MODEL, { messages: [], max_tokens } as Record<string, unknown>)
).rejects.toBeInstanceOf(SharedAiBudgetDenied);
expect(fetch).not.toHaveBeenCalled();
expect(run).not.toHaveBeenCalled();
}
);

it.each([null, [], 'allowed', 0, {}])(
'fails closed for malformed debit receipt %j without inference',
async (reply) => {
const { namespace, fetch } = makeBudget(reply);
const { binding, run } = makeBinding();
const guarded = createBudgetedWorkersAiBinding(binding as unknown as Ai, namespace);

await expect(
guarded.run(MODEL, { messages: [], max_tokens: 512 } as Record<string, unknown>)
).rejects.toBeInstanceOf(SharedAiBudgetDenied);
expect(fetch).toHaveBeenCalledOnce();
expect(run).not.toHaveBeenCalled();
}
);

it('fails closed on exhausted daily budget before inference', async () => {
const { namespace, fetch } = makeBudget({
allowed: false,
used: DAILY_CAP,
remaining: 0,
retryAfter: 1,
dayKey: new Date().toISOString().slice(0, 10),
});
const { binding, run } = makeBinding();
const guarded = createBudgetedWorkersAiBinding(binding as unknown as Ai, namespace);

await expect(
guarded.run(MODEL, { messages: [], max_tokens: 512 } as Record<string, unknown>)
).rejects.toBeInstanceOf(SharedAiBudgetDenied);
expect(fetch).toHaveBeenCalledOnce();
expect(run).not.toHaveBeenCalled();
});
});
Loading
Loading