diff --git a/package.json b/package.json index 7f62613..87fb661 100644 --- a/package.json +++ b/package.json @@ -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", @@ -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" } } } diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index f63ce77..d4aa3b0 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -14,7 +14,6 @@ overrides: 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 @@ -26,6 +25,7 @@ overrides: 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 + devalue@<=5.9.2: 5.9.3 importers: @@ -3368,8 +3368,8 @@ packages: detect-node-es@1.1.0: resolution: {integrity: sha512-ypdmJU/TbBby2Dxibuv7ZLW3Bs1QEmM7nHjEANfohJLvE0XVujisn1qPJcZxg+qDucsr+bP6fLD1rPS3AhJ7EQ==} - devalue@5.8.1: - resolution: {integrity: sha512-4CXDYRBGqN+57wVJkuXBYmpAVUSg3L6JAQa/DFqm238G73E1wuyc/JhGQJzN7vUf/CMphYau2zXbfWzDR5aTEw==} + devalue@5.9.3: + resolution: {integrity: sha512-xRumYOCUZN/EesqHEU3WOXanOZNvfZFZ/o1AHVFDX1yI0UAkZkOgDXt341CzKoBVIkgQba55/+DjGBKrIoKcHw==} devlop@1.1.0: resolution: {integrity: sha512-RWmIqhcFf1lRYBvNmr7qTNuyCt/7/ns2jbpp1+PalgE/rDQcBT0fioSMUpJ93irlUhC5hrg4cYqe6U+0ImW0rA==} @@ -7461,7 +7461,7 @@ snapshots: clsx: 2.1.1 common-ancestor-path: 2.0.0 cookie: 2.0.1 - devalue: 5.8.1 + devalue: 5.9.3 diff: 9.0.0 dset: 3.1.4 es-module-lexer: 2.3.1 @@ -7820,7 +7820,7 @@ snapshots: detect-node-es@1.1.0: {} - devalue@5.8.1: {} + devalue@5.9.3: {} devlop@1.1.0: dependencies: diff --git a/src/lib/ai-cloudflare.test.ts b/src/lib/ai-cloudflare.test.ts new file mode 100644 index 0000000..03ce99f --- /dev/null +++ b/src/lib/ai-cloudflare.test.ts @@ -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) { + 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(); + }); +}); diff --git a/src/lib/ai-cloudflare.ts b/src/lib/ai-cloudflare.ts index 853606e..f005393 100644 --- a/src/lib/ai-cloudflare.ts +++ b/src/lib/ai-cloudflare.ts @@ -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['binding']; @@ -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 { binding?: WorkersAiBinding; endpointUrl: string; apiKey: string; model: string; headers?: Record; + budgetNamespace?: SharedBudgetNamespace; } function getDirectBaseUrl(): string { @@ -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({ binding, endpointUrl, apiKey, model, headers, -}: CreateLanguageModelArgs): LanguageModel { + budgetNamespace, +}: CreateLanguageModelArgs): 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; diff --git a/src/lib/shared-ai-budget.test.ts b/src/lib/shared-ai-budget.test.ts new file mode 100644 index 0000000..39b1ebe --- /dev/null +++ b/src/lib/shared-ai-budget.test.ts @@ -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 + ) { + 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) + ).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) + ).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) + ).rejects.toBeInstanceOf(SharedAiBudgetDenied); + expect(fetch).toHaveBeenCalledOnce(); + expect(run).not.toHaveBeenCalled(); + }); +}); diff --git a/src/lib/shared-ai-budget.ts b/src/lib/shared-ai-budget.ts new file mode 100644 index 0000000..c96c63d --- /dev/null +++ b/src/lib/shared-ai-budget.ts @@ -0,0 +1,136 @@ +import type { WorkersAISettings } from 'workers-ai-provider'; + +type WorkersAiBinding = Extract['binding']; +type BindingRunOptions = Parameters[2]; +export type SharedBudgetNamespace = { + idFromName(name: string): Id; + get(id: Id): { fetch(input: string, init?: RequestInit): Promise }; +}; + +const DAILY_CAP = 9_500; +const DEFAULT_OUTPUT_TOKENS = 512; +const MAX_OUTPUT_TOKENS = 8_192; +const PRICED_MODEL_RATES: Record = { + '@cf/meta/llama-3.3-70b-instruct-fp8-fast': { input: 26_668, output: 204_805 }, +}; + +export class SharedAiBudgetDenied extends Error { + constructor() { + super('The shared daily Workers AI budget is unavailable or exhausted.'); + this.name = 'SharedAiBudgetDenied'; + } +} + +export function findSharedAiBudgetDenied(error: unknown): SharedAiBudgetDenied | undefined { + const seen = new Set(); + let current: unknown = error; + while (current && typeof current === 'object' && !seen.has(current)) { + if (current instanceof SharedAiBudgetDenied) return current; + seen.add(current); + current = (current as { cause?: unknown }).cause; + } + return undefined; +} + +function deny(): never { + throw new SharedAiBudgetDenied(); +} + +function isRecord(value: unknown): value is Record { + return value !== null && typeof value === 'object' && !Array.isArray(value); +} + +function boundedOutputTokens(input: Record): number { + const value = input.max_tokens === undefined ? DEFAULT_OUTPUT_TOKENS : input.max_tokens; + if ( + !Number.isSafeInteger(value) || + (value as number) <= 0 || + (value as number) > MAX_OUTPUT_TOKENS + ) { + return deny(); + } + return value as number; +} + +async function reserveWorkersAiCall( + namespace: SharedBudgetNamespace | undefined, + model: string, + input: Record, + outputTokens: number +): Promise { + const rates = PRICED_MODEL_RATES[model]; + if (!rates || !namespace) return deny(); + + let serialized: string; + try { + const json = JSON.stringify(input); + if (typeof json !== 'string') return deny(); + serialized = json; + } catch { + return deny(); + } + const inputBytes = new TextEncoder().encode(serialized).byteLength; + const estimatedInputTokens = Math.ceil(inputBytes * 1.2); + const neurons = Math.ceil( + (estimatedInputTokens * rates.input + outputTokens * rates.output) / 1_000_000 + ); + if (!Number.isSafeInteger(neurons) || neurons <= 0 || neurons > DAILY_CAP) return deny(); + + let response: Response; + try { + const stub = namespace.get(namespace.idFromName('global-budget')); + response = await stub.fetch('https://internal.local/try-debit', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ neurons }), + }); + } catch { + return deny(); + } + if (response.status !== 200) return deny(); + + let receipt: unknown; + try { + receipt = await response.json(); + } catch { + return deny(); + } + if (!isRecord(receipt)) return deny(); + if ( + receipt.allowed !== true || + receipt.dayKey !== new Date().toISOString().slice(0, 10) || + receipt.retryAfter !== 0 || + !Number.isSafeInteger(receipt.used) || + (receipt.used as number) < neurons || + !Number.isSafeInteger(receipt.remaining) || + (receipt.remaining as number) < 0 || + (receipt.used as number) + (receipt.remaining as number) !== DAILY_CAP + ) { + return deny(); + } +} + +export function createBudgetedWorkersAiBinding( + binding: WorkersAiBinding, + namespace: SharedBudgetNamespace | undefined +): WorkersAiBinding { + return new Proxy(binding, { + get(target, property) { + if (property === 'run') { + return async ( + model: string, + input: Record, + options?: BindingRunOptions + ) => { + if (!isRecord(input)) return deny(); + const outputTokens = boundedOutputTokens(input); + const boundedInput = { ...input, max_tokens: outputTokens }; + await reserveWorkersAiCall(namespace, model, boundedInput, outputTokens); + // Invoke through the target object: Ai.run depends on the binding receiver. + return target.run(model, boundedInput, options); + }; + } + return Reflect.get(target, property, target); + }, + }); +} diff --git a/src/lib/worker-env.ts b/src/lib/worker-env.ts index 985705a..bda8e8c 100644 --- a/src/lib/worker-env.ts +++ b/src/lib/worker-env.ts @@ -11,6 +11,7 @@ export type WorkerEnv = { AI_API_KEY?: string; AI_BASE_URL?: string; AI?: Ai; + NEURON_BUDGET?: DurableObjectNamespace; LOCAL_AI_URL?: string; CLI_BRIDGE_URL?: string; NODE_ENV?: string; diff --git a/src/worker/routes/ai.ts b/src/worker/routes/ai.ts index 3e57b64..9df7c69 100644 --- a/src/worker/routes/ai.ts +++ b/src/worker/routes/ai.ts @@ -91,6 +91,7 @@ ai.post('/chat', async (c) => { const result = streamText({ model: getLanguageModel({ binding: c.env.AI, + budgetNamespace: c.env.NEURON_BUDGET, endpointUrl, apiKey, model, @@ -168,6 +169,7 @@ Remember to respond with valid JSON in the exact format specified.`; const result = await generateText({ model: getLanguageModel({ binding: c.env.AI, + budgetNamespace: c.env.NEURON_BUDGET, endpointUrl, apiKey, model, diff --git a/src/worker/routes/articles.ts b/src/worker/routes/articles.ts index ad7fd49..28102f1 100644 --- a/src/worker/routes/articles.ts +++ b/src/worker/routes/articles.ts @@ -293,6 +293,7 @@ articles.post('/:id/session-review', async (c) => { const result = await generateText({ model: getLanguageModel({ binding: c.env.AI, + budgetNamespace: c.env.NEURON_BUDGET, endpointUrl, apiKey, model, diff --git a/src/worker/routes/misc.ts b/src/worker/routes/misc.ts index 7ad9c2c..1c350ea 100644 --- a/src/worker/routes/misc.ts +++ b/src/worker/routes/misc.ts @@ -373,6 +373,7 @@ misc.post('/ext/chat', async (c) => { const result = streamText({ model: getLanguageModel({ binding: c.env.AI, + budgetNamespace: c.env.NEURON_BUDGET, endpointUrl: '', apiKey: '', model: '', diff --git a/wrangler.toml b/wrangler.toml index e02bdcf..837ca04 100644 --- a/wrangler.toml +++ b/wrangler.toml @@ -23,6 +23,11 @@ cpu_ms = 30000 [ai] binding = "AI" +[[durable_objects.bindings]] +name = "NEURON_BUDGET" +class_name = "NeuronBudgetDO" +script_name = "free-ai-gateway" + [vars] NODE_ENV = "production" BETTER_AUTH_URL = "https://read.significanthobbies.com"