From 0be3bc55c28b125ec42ec85fd4b0c5fa7478a3a9 Mon Sep 17 00:00:00 2001 From: Daniel Salazar Date: Thu, 17 Sep 2026 15:18:38 -0700 Subject: [PATCH] feat(metering): AI cost multiplier hook for AI drivers (#3898) Emits ai.cost.multiplier..: before recording AI usage, so what a model costs to charge is policy an extension owns rather than a number in core. Nothing listening records the provider cost. MeteringService.withAiCostMultiplier(driver) returns a view of the service whose recording paths scale costOverride by the hook's answer; every AI driver hands that view to its providers, so all of them are covered without touching provider code. --- .../clients/database/MySQLDatabaseClient.ts | 28 +- src/backend/clients/database/SQLBatcher.js | 28 +- .../clients/database/SQLBatcher.test.ts | 39 +++ src/backend/clients/event/types.ts | 17 ++ .../drivers/ai-chat/ChatCompletionDriver.ts | 10 +- .../drivers/ai-image/ImageGenerationDriver.ts | 8 +- src/backend/drivers/ai-ocr/OCRDriver.ts | 18 +- .../ai-speech2speech/VoiceChangerDriver.ts | 11 +- .../ai-speech2txt/SpeechToTextDriver.ts | 8 +- src/backend/drivers/ai-tts/TTSDriver.ts | 32 ++- .../drivers/ai-video/VideoGenerationDriver.ts | 17 +- .../services/events/kv.integration.test.ts | 34 +-- .../services/metering/MeteringService.ts | 54 ++++ .../services/metering/aiCostFactor.test.ts | 266 ++++++++++++++++++ src/backend/services/metering/aiCostFactor.ts | 191 +++++++++++++ src/backend/stores/fs/FSEntryStore.ts | 143 +++++++--- .../stores/fs/FSEntryStore.writeBack.test.ts | 242 ++++++++++++++++ 17 files changed, 1059 insertions(+), 87 deletions(-) create mode 100644 src/backend/services/metering/aiCostFactor.test.ts create mode 100644 src/backend/services/metering/aiCostFactor.ts create mode 100644 src/backend/stores/fs/FSEntryStore.writeBack.test.ts diff --git a/src/backend/clients/database/MySQLDatabaseClient.ts b/src/backend/clients/database/MySQLDatabaseClient.ts index 472e9c428..b9eee25e5 100644 --- a/src/backend/clients/database/MySQLDatabaseClient.ts +++ b/src/backend/clients/database/MySQLDatabaseClient.ts @@ -20,7 +20,7 @@ import { readdirSync, readFileSync } from 'fs'; import { isAbsolute, resolve as resolvePath } from 'path'; import { metrics } from '@opentelemetry/api'; -import { createPool, type Pool } from 'mysql2'; +import { createPool, ExecuteValues, type Pool } from 'mysql2'; import { Span } from '../../util/span.js'; import { AbstractDatabaseClient, type WriteResult } from './DatabaseClient'; import { SQLBatcher } from './SQLBatcher.js'; @@ -54,6 +54,8 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { private replicaPool!: Pool; private db!: SQLBatcher; private dbReplica!: SQLBatcher; + /** Primary pool, SELECT-only: same rows as `db`, without the transaction. */ + private dbPrimaryRead!: SQLBatcher; private configuration = Configuration.SINGLE; private shutdownStarted = false; private shutdownTimer: ReturnType | null = null; @@ -79,6 +81,7 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { console.log('[mysql] connected to primary'); this.db = this.createPrimaryBatcher(this.primaryPool); + this.dbPrimaryRead = this.createPrimaryReadBatcher(this.primaryPool); if (dbConf.replica) { this.replicaPool = this.createPool(dbConf.replica); @@ -162,7 +165,7 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { query: string, params: unknown[] = [], ): Promise[]> { - const result = await this.db.execute(query, params); + const result = await this.dbPrimaryRead.execute(query, params); if (!result) return []; return (result[0] as Record[]) ?? []; } @@ -201,7 +204,7 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { await conn.beginTransaction(); try { for (const { statement, values } of entries) { - await conn.execute(statement, values); + await conn.execute(statement, values as ExecuteValues); } await conn.commit(); } catch (err) { @@ -224,7 +227,7 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { // Run both reads in parallel — prefer replica when it returns rows, // otherwise fall back to primary to handle replication lag. - const primaryPromise = this.db.execute(query, params); + const primaryPromise = this.dbPrimaryRead.execute(query, params); try { const replicaResult = await this.dbReplica.execute(query, params); if ( @@ -354,6 +357,22 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { }); } + /** + * Reads that must see the primary still only read, so they skip the + * batcher's transaction wrapper. Worth a separate batcher because BEGIN and + * COMMIT are round trips: against a primary in another region they cost + * more than the query. + */ + private createPrimaryReadBatcher(pool: Pool): SQLBatcher { + return new SQLBatcher(pool, { + maxTimeInQueue: 30, + maxBatchSize: 5, + poolLabel: 'primary', + readOnly: true, + acquireTimeoutMs: this.config.database?.acquireTimeoutMs, + }); + } + private createReplicaBatcher(pool: Pool): SQLBatcher { return new SQLBatcher(pool, { maxTimeInQueue: 10, @@ -378,6 +397,7 @@ export class MySQLDatabaseClient extends AbstractDatabaseClient { database: dbConf.database ?? 'puter', }); this.db = this.createPrimaryBatcher(this.primaryPool); + this.dbPrimaryRead = this.createPrimaryReadBatcher(this.primaryPool); if (this.configuration === Configuration.SINGLE) { this.replicaPool = this.primaryPool; diff --git a/src/backend/clients/database/SQLBatcher.js b/src/backend/clients/database/SQLBatcher.js index 2b313ea57..7377deb62 100644 --- a/src/backend/clients/database/SQLBatcher.js +++ b/src/backend/clients/database/SQLBatcher.js @@ -284,11 +284,16 @@ export class SQLBatcher { // individually below. Without this, MySQL would commit every // statement up to the failure point and a per-item retry would // misreport already-committed inserts as duplicate-key failures. + // + // A read-only batch has nothing to roll back, and the wrapper is not + // free: BEGIN and COMMIT are round trips of their own, which triples + // the cost of a read against a pool in another region. + const wrapInTransaction = !this.readOnly; let batchSucceeded = false; try { - await connection.beginTransaction(); + if (wrapInTransaction) await connection.beginTransaction(); const [results, fields] = await connection.query(query, values); - await connection.commit(); + if (wrapInTransaction) await connection.commit(); batchSucceeded = true; this.#consecutiveFailures = 0; for (let i = 0; i < batch.length; i++) { @@ -296,10 +301,12 @@ export class SQLBatcher { b.resolve([results[i], fields?.[i]]); } } catch (batchError) { - try { - await connection.rollback(); - } catch (rollbackError) { - console.warn('SQLBatcher rollback failed:', rollbackError); + if (wrapInTransaction) { + try { + await connection.rollback(); + } catch (rollbackError) { + console.warn('SQLBatcher rollback failed:', rollbackError); + } } console.warn( 'SQLBatcher batch failed; retrying items individually:', @@ -311,10 +318,11 @@ export class SQLBatcher { if (batchSucceeded) return; - // Per-item fallback. The transaction was rolled back so no statement - // committed; re-running each item independently produces clean - // success/failure outcomes for each caller. Concurrency is capped to - // avoid briefly saturating the pool when a large batch fails. + // Per-item fallback. Nothing from the batch is live — it was rolled + // back, or carried only SELECTs — so re-running each item + // independently produces clean success/failure outcomes for each + // caller. Concurrency is capped to avoid briefly saturating the pool + // when a large batch fails. flushFailureCounter.add(1, this.#metricAttrs); fallbackInvocationsCounter.add(1, this.#metricAttrs); diff --git a/src/backend/clients/database/SQLBatcher.test.ts b/src/backend/clients/database/SQLBatcher.test.ts index 14c416bd4..85056d4ff 100644 --- a/src/backend/clients/database/SQLBatcher.test.ts +++ b/src/backend/clients/database/SQLBatcher.test.ts @@ -86,6 +86,45 @@ describe('SQLBatcher', () => { expect(conn.release).toHaveBeenCalledTimes(1); }); + it('skips the transaction wrapper on a readOnly batcher', async () => { + const conn = makeConnection(happyBatch); + const { pool } = makePool(conn); + const batcher = new SQLBatcher(pool, { + maxTimeInQueue: 5, + readOnly: true, + }); + + const [a, b] = await Promise.all([ + batcher.query('SELECT a', []), + batcher.query('SELECT b', []), + ]); + expect(a[0]).toEqual([{ n: 0 }]); + expect(b[0]).toEqual([{ n: 1 }]); + // BEGIN and COMMIT are round trips; a SELECT-only batch has nothing to + // roll back, so it must not pay for them. + expect(conn.beginTransaction).not.toHaveBeenCalled(); + expect(conn.commit).not.toHaveBeenCalled(); + expect(conn.release).toHaveBeenCalledTimes(1); + }); + + it('does not roll back a failed readOnly batch', async () => { + const conn = makeConnection((sql: string) => { + if (isBatchQuery(sql)) throw makeError('ER_LOCK_DEADLOCK'); + return [[{ n: 0 }], undefined]; + }); + const { pool } = makePool(conn); + const batcher = new SQLBatcher(pool, { + maxTimeInQueue: 5, + readOnly: true, + }); + + await Promise.all([ + batcher.query('SELECT a', []), + batcher.query('SELECT b', []), + ]); + expect(conn.rollback).not.toHaveBeenCalled(); + }); + it('drops the oldest item with reason queueOverflow at the high-water mark', async () => { const conn = makeConnection(happyBatch); const { pool } = makePool(conn); diff --git a/src/backend/clients/event/types.ts b/src/backend/clients/event/types.ts index edf33daf8..310c9bc8a 100644 --- a/src/backend/clients/event/types.ts +++ b/src/backend/clients/event/types.ts @@ -693,6 +693,11 @@ export type EventMap = { // normalized path: `route...before|after|error|reject`. Same // wildcard + veto semantics as the driver lifecycle above. [K in `route.${string}`]: RouteLifecycleEvent; +} & { + // Cost factor for recorded AI usage, keyed by driver and model: + // `ai.cost.factor..:`. Emitted once per model + // per batch; the last listener to set `factor` wins. + [K in `ai.cost.factor.${string}`]: AiCostFactorEvent; } & { [K in `pubsub.login.${string}`]: { authtoken: string }; } & { @@ -713,6 +718,18 @@ export type EventMap = { 'outer.pubsub.metering.credits-changed': { userUuid: string }; } & IExtensionEventMap; +/** Payload for `ai.cost.factor..` events. */ +export type AiCostFactorEvent = { + /** Driver doing the pricing, e.g. `ai-chat`. */ + driver: string; + /** `:` the usage is recorded under. */ + model: string; + /** Who the usage is being charged to. */ + actor: Actor; + /** Applied to the cost, starting at 1. Values <= 0 are ignored. */ + factor: number; +}; + /** * Phase of a request/method lifecycle. `reject` is emitted when a `before` * listener vetoes the call (sets `allow = false`); the call never runs and no diff --git a/src/backend/drivers/ai-chat/ChatCompletionDriver.ts b/src/backend/drivers/ai-chat/ChatCompletionDriver.ts index 1972b148a..33053c5d0 100644 --- a/src/backend/drivers/ai-chat/ChatCompletionDriver.ts +++ b/src/backend/drivers/ai-chat/ChatCompletionDriver.ts @@ -26,6 +26,7 @@ import { HttpError, isHttpError } from '../../core/http/HttpError.js'; import { FREE_SUBSCRIPTION_IDS } from '../../services/metering/consts.js'; import type { CreditHold } from '../../services/metering/types.js'; import { NO_CREDIT_HOLD } from '../../services/metering/types.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import type { DriverStreamResult } from '../meta.js'; import { PuterDriver } from '../types.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; @@ -303,6 +304,11 @@ export class ChatCompletionDriver extends PuterDriver { #providers: Record = {}; #modelIdMap: Record = {}; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override onServerStart() { this.#registerProviders(); this.#buildModelMap(); @@ -1029,7 +1035,7 @@ export class ChatCompletionDriver extends PuterDriver { }; const cost = this.#computeCost(usage, model); - this.services.metering.utilRecordUsageObject( + this.#aiMetering.utilRecordUsageObject( { [`estimated_${inputKey}`]: inputTokens, [`estimated_${outputKey}`]: outputTokens, @@ -1127,7 +1133,7 @@ export class ChatCompletionDriver extends PuterDriver { #registerProviders() { const providers = this.config.providers ?? {}; - const metering = this.services.metering; + const metering = this.#aiMetering; const readKey = (cfg: Record | undefined) => (cfg?.apiKey as string | undefined) ?? diff --git a/src/backend/drivers/ai-image/ImageGenerationDriver.ts b/src/backend/drivers/ai-image/ImageGenerationDriver.ts index dfd5ba3a1..644787fd3 100644 --- a/src/backend/drivers/ai-image/ImageGenerationDriver.ts +++ b/src/backend/drivers/ai-image/ImageGenerationDriver.ts @@ -24,6 +24,7 @@ import { Readable } from 'node:stream'; import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; import type { Actor } from '../../core/actor.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import { PuterDriver } from '../types.js'; import { secureFetch } from '../../util/secureHttp.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; @@ -73,6 +74,11 @@ export class ImageGenerationDriver extends PuterDriver { #providers: Record = {}; #modelIdMap: Record = {}; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override onServerStart() { this.#registerProviders(); this.#buildModelMap(); @@ -253,7 +259,7 @@ export class ImageGenerationDriver extends PuterDriver { #registerProviders() { const providers = this.config.providers ?? {}; - const m = this.services.metering; + const m = this.#aiMetering; const readKey = ( ...cfgs: Array | undefined> diff --git a/src/backend/drivers/ai-ocr/OCRDriver.ts b/src/backend/drivers/ai-ocr/OCRDriver.ts index 52e0fa7c9..3aef5db5a 100644 --- a/src/backend/drivers/ai-ocr/OCRDriver.ts +++ b/src/backend/drivers/ai-ocr/OCRDriver.ts @@ -27,6 +27,7 @@ import { Actor } from '../../core/actor.js'; import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; import { mimeFromName } from '../../util/fileSigning.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import { PuterDriver } from '../types.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; import { loadFileInput, type LoadedFile } from '../util/fileInput.js'; @@ -102,6 +103,11 @@ export class OCRDriver extends PuterDriver { readonly rateLimit = AI_RATE_LIMIT; readonly concurrent = AI_CONCURRENT; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override getReportedCosts() { return Object.entries(OCR_COSTS).map(([usageType, ucentsPerUnit]) => ({ usageType, @@ -132,9 +138,11 @@ export class OCRDriver extends PuterDriver { const providers = this.config.providers ?? {}; const textract = providers['aws-textract'] as - Record | undefined; + | Record + | undefined; const textractAws = (textract?.aws ?? textract) as - Record | undefined; + | Record + | undefined; const textractAccessKey = textractAws?.access_key as string | undefined; const textractSecretKey = textractAws?.secret_key as string | undefined; const textractRegion = @@ -342,7 +350,7 @@ export class OCRDriver extends PuterDriver { } const pages = pageCount || 1; - this.services.metering.incrementUsage( + this.#aiMetering.incrementUsage( actor, usageType, pages, @@ -454,14 +462,14 @@ export class OCRDriver extends PuterDriver { const pagesProcessed = response?.usageInfo?.pagesProcessed ?? (Array.isArray(response?.pages) ? response.pages.length : 1); - this.services.metering.incrementUsage( + this.#aiMetering.incrementUsage( actor, 'mistral-ocr:ocr:page', pagesProcessed, OCR_COSTS['mistral-ocr:ocr:page'] * pagesProcessed, ); if (annotations) { - this.services.metering.incrementUsage( + this.#aiMetering.incrementUsage( actor, 'mistral-ocr:annotations:page', pagesProcessed, diff --git a/src/backend/drivers/ai-speech2speech/VoiceChangerDriver.ts b/src/backend/drivers/ai-speech2speech/VoiceChangerDriver.ts index 8621e98d2..d90f28f5a 100644 --- a/src/backend/drivers/ai-speech2speech/VoiceChangerDriver.ts +++ b/src/backend/drivers/ai-speech2speech/VoiceChangerDriver.ts @@ -20,6 +20,7 @@ import { Readable } from 'node:stream'; import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import type { DriverStreamResult } from '../meta.js'; import { PuterDriver } from '../types.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; @@ -69,6 +70,11 @@ export class VoiceChangerDriver extends PuterDriver { readonly rateLimit = AI_RATE_LIMIT; readonly concurrent = AI_CONCURRENT; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override getReportedCosts(): Record[] { return Object.entries(VOICE_CHANGER_COSTS).map( ([usageType, ucentsPerUnit]) => ({ @@ -87,7 +93,8 @@ export class VoiceChangerDriver extends PuterDriver { override onServerStart() { const elevenlabs = this.config.providers?.elevenlabs as - Record | undefined; + | Record + | undefined; this.#apiKey = (elevenlabs?.apiKey as string | undefined) ?? @@ -296,7 +303,7 @@ export class VoiceChangerDriver extends PuterDriver { const arrayBuffer = await response.arrayBuffer(); const stream = Readable.from(Buffer.from(arrayBuffer)); - this.services.metering.incrementUsage( + this.#aiMetering.incrementUsage( actor, usageKey, estimatedSeconds, diff --git a/src/backend/drivers/ai-speech2txt/SpeechToTextDriver.ts b/src/backend/drivers/ai-speech2txt/SpeechToTextDriver.ts index 21142aecf..0daebeda8 100644 --- a/src/backend/drivers/ai-speech2txt/SpeechToTextDriver.ts +++ b/src/backend/drivers/ai-speech2txt/SpeechToTextDriver.ts @@ -19,6 +19,7 @@ import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import { PuterDriver } from '../types.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; import { @@ -66,6 +67,11 @@ export class SpeechToTextDriver extends PuterDriver { #providers: Record = {}; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override onServerStart() { this.#registerProviders(); } @@ -195,7 +201,7 @@ export class SpeechToTextDriver extends PuterDriver { const deps: ISpeechToTextDeps = { stores: this.stores, fs: this.services.fs, - metering: this.services.metering, + metering: this.#aiMetering, }; this.#providers['openai'] = new OpenAISpeechToTextProvider(deps, { diff --git a/src/backend/drivers/ai-tts/TTSDriver.ts b/src/backend/drivers/ai-tts/TTSDriver.ts index bee62755d..ea0e43c90 100644 --- a/src/backend/drivers/ai-tts/TTSDriver.ts +++ b/src/backend/drivers/ai-tts/TTSDriver.ts @@ -19,6 +19,7 @@ import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import type { DriverStreamResult } from '../meta.js'; import { PuterDriver } from '../types.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; @@ -73,6 +74,11 @@ export class TTSDriver extends PuterDriver { #providers: Record = {}; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override onServerStart() { this.#registerProviders(); } @@ -244,7 +250,7 @@ export class TTSDriver extends PuterDriver { #registerProviders() { const providers = this.config.providers ?? {}; - const m = this.services.metering; + const m = this.#aiMetering; const openaiConfig = (providers['openai-tts'] as Record | undefined) ?? @@ -266,7 +272,8 @@ export class TTSDriver extends PuterDriver { } const elevenlabs = providers['elevenlabs'] as - Record | undefined; + | Record + | undefined; const elevenKey = (elevenlabs?.apiKey as string | undefined) ?? (elevenlabs?.api_key as string | undefined) ?? @@ -277,7 +284,8 @@ export class TTSDriver extends PuterDriver { apiKey: elevenKey, apiBaseUrl: elevenlabs?.apiBaseUrl as string | undefined, defaultVoiceId: elevenlabs?.defaultVoiceId as - string | undefined, + | string + | undefined, }); } catch (e) { console.warn( @@ -288,9 +296,11 @@ export class TTSDriver extends PuterDriver { } const polly = providers['aws-polly'] as - Record | undefined; + | Record + | undefined; const pollyAws = (polly?.aws ?? polly) as - Record | undefined; + | Record + | undefined; const pollyAccessKey = pollyAws?.access_key as string | undefined; const pollySecretKey = pollyAws?.secret_key as string | undefined; const pollyRegion = @@ -317,9 +327,10 @@ export class TTSDriver extends PuterDriver { } #registerGeminiProvider(providers: Record) { - const m = this.services.metering; + const m = this.#aiMetering; const gemini = (providers['gemini'] ?? providers['gemini-tts']) as - Record | undefined; + | Record + | undefined; const geminiKey = (gemini?.apiKey as string | undefined) ?? (gemini?.api_key as string | undefined) ?? @@ -339,9 +350,10 @@ export class TTSDriver extends PuterDriver { } #registerXAIProvider(providers: Record) { - const m = this.services.metering; + const m = this.#aiMetering; const xai = (providers['xai'] ?? providers['xai-tts']) as - Record | undefined; + | Record + | undefined; const xaiKey = (xai?.apiKey as string | undefined) ?? (xai?.api_key as string | undefined) ?? @@ -361,7 +373,7 @@ export class TTSDriver extends PuterDriver { } #registerSpeechifyProvider(providers: Record) { - const m = this.services.metering; + const m = this.#aiMetering; const speechify = (providers['speechify'] ?? providers['speechify-tts']) as Record | undefined; const speechifyKey = diff --git a/src/backend/drivers/ai-video/VideoGenerationDriver.ts b/src/backend/drivers/ai-video/VideoGenerationDriver.ts index 98ae15e93..2c21694ee 100644 --- a/src/backend/drivers/ai-video/VideoGenerationDriver.ts +++ b/src/backend/drivers/ai-video/VideoGenerationDriver.ts @@ -23,6 +23,7 @@ import { Readable } from 'node:stream'; import { Context } from '../../core/context.js'; import { HttpError } from '../../core/http/HttpError.js'; import type { Actor } from '../../core/actor.js'; +import type { MeteringService } from '../../services/metering/MeteringService.js'; import { PuterDriver } from '../types.js'; import { secureFetch } from '../../util/secureHttp.js'; import { AI_CONCURRENT, AI_RATE_LIMIT } from '../util/aiLimits.js'; @@ -94,6 +95,11 @@ export class VideoGenerationDriver extends PuterDriver { #providers: Record = {}; #modelIdMap: Record = {}; + /** Metering scoped to this driver. Lazy: services wire up after drivers. */ + get #aiMetering(): MeteringService { + return this.services.metering.withAiCostFactor(this.driverName); + } + override onServerStart() { this.#registerProviders(); this.#buildModelMap(); @@ -303,7 +309,7 @@ export class VideoGenerationDriver extends PuterDriver { #registerProviders() { const providers = this.config.providers ?? {}; - const m = this.services.metering; + const m = this.#aiMetering; // Same lenient reader as ImageGenerationDriver — accept // `apiKey || secret_key`, and fall back from the video-specific @@ -345,9 +351,11 @@ export class VideoGenerationDriver extends PuterDriver { // pair its missing apiBaseUrl with the shared block's key (or vice // versa) and point a region-scoped key at the wrong endpoint. const byteplusVideoCfg = providers['byteplus-video-generation'] as - Record | undefined; + | Record + | undefined; const byteplusSharedCfg = providers['byteplus'] as - Record | undefined; + | Record + | undefined; const byteplusKey = readKey(byteplusVideoCfg, byteplusSharedCfg); if (byteplusKey) { this.#providers['byteplus-video-generation'] = @@ -356,7 +364,8 @@ export class VideoGenerationDriver extends PuterDriver { apiKey: byteplusKey, apiBaseUrl: (byteplusVideoCfg?.apiBaseUrl ?? byteplusSharedCfg?.apiBaseUrl) as - string | undefined, + | string + | undefined, }, m, ); diff --git a/src/backend/services/events/kv.integration.test.ts b/src/backend/services/events/kv.integration.test.ts index e4bc3c4e4..fc50dfc07 100644 --- a/src/backend/services/events/kv.integration.test.ts +++ b/src/backend/services/events/kv.integration.test.ts @@ -55,8 +55,24 @@ const settle = () => interval: 25, }); +/** + * One row's own delivery. Session rows from earlier tests stay live on the + * shared socket, so `settle()` can return on someone else's envelope. + */ +const eventFor = (subId: string) => + vi.waitFor( + () => { + const found = delivered.find((one) => one.subId === subId); + expect(found).toBeDefined(); + return found!.event as Record; + }, + { timeout: EVENTS_COALESCE_WINDOW_MS * 12, interval: 25 }, + ); + const quiet = () => - new Promise((resolve) => setTimeout(resolve, EVENTS_COALESCE_WINDOW_MS * 3)); + new Promise((resolve) => + setTimeout(resolve, EVENTS_COALESCE_WINDOW_MS * 3), + ); /** An app owned by the test user, registered the way the app store sees one. */ const makeApp = async (metadata?: object): Promise => { @@ -214,18 +230,6 @@ describe('a kv write reaches its subscribers', () => { }); describe('a subscription that asked for the value', () => { - // Session rows from earlier tests stay live on the shared socket, so - // every assertion here reads its own row's delivery rather than the first. - const eventFor = (subId: string) => - vi.waitFor( - () => { - const found = delivered.find((one) => one.subId === subId); - expect(found).toBeDefined(); - return found!.event as Record; - }, - { timeout: EVENTS_COALESCE_WINDOW_MS * 12, interval: 25 }, - ); - it('is handed what the key now holds, and a row that did not ask is not', async () => { const asking = await subscribe(`kv:${ownAppUid}:value:*`, ownAppToken, { includeValue: true, @@ -375,10 +379,8 @@ describe('the cross-app gate against real grants', () => { await kvSet(env.users.user.token, 'cart:items', [7], { appUuid: otherAppUid, }); - await settle(); - const own = delivered.find((one) => one.subId === sub.subId); - expect(own?.event).toMatchObject({ value: [7] }); + expect(await eventFor(sub.subId)).toMatchObject({ value: [7] }); await revokeRead(otherAppUid); }); diff --git a/src/backend/services/metering/MeteringService.ts b/src/backend/services/metering/MeteringService.ts index c946a3789..fb06e9d4c 100644 --- a/src/backend/services/metering/MeteringService.ts +++ b/src/backend/services/metering/MeteringService.ts @@ -22,6 +22,7 @@ import type { Actor } from '../../core/actor'; import { isSystemActor } from '../../core/actor'; import { HttpError } from '../../core/http/HttpError.js'; import { PuterService } from '../types'; +import { MAX_AI_COST_FACTOR, withAiCostFactor } from './aiCostFactor.js'; import { DEFAULT_FREE_SUBSCRIPTION, DEFAULT_TEMP_SUBSCRIPTION, @@ -370,6 +371,59 @@ export class MeteringService extends PuterService { this.defaultSubscriptionResolvers.push(fn); } + // -- AI cost factor ------------------------------------------- + + /** + * This service as an AI driver should use it: recorded costs pass through + * the `ai.cost.factor..` hook first. + */ + withAiCostFactor(driver: string): MeteringService { + return withAiCostFactor(this, driver); + } + + /** + * Whether anything prices this model. Synchronous so an unhooked deployment + * records in the caller's own tick, not after the request ends. + */ + hasAiCostFactor(driver: string, model: string): boolean { + return this.clients.event.hasListeners( + `ai.cost.factor.${driver}.${model}`, + ); + } + + /** One model's cost factor. 1 when unhooked or the answer is unusable. */ + async resolveAiCostFactor( + actor: Actor, + driver: string, + model: string, + ): Promise { + const key = `ai.cost.factor.${driver}.${model}` as const; + try { + if (!this.hasAiCostFactor(driver, model)) return 1; + const event = { driver, model, actor, factor: 1 }; + await this.clients.event.emitAndWait(key, event, {}); + const factor = Number(event.factor); + if ( + !Number.isFinite(factor) || + factor <= 0 || + factor > MAX_AI_COST_FACTOR + ) { + if (factor !== 1) { + console.warn( + `[metering] ignoring AI cost factor ${event.factor} for ${key}`, + ); + } + return 1; + } + return factor; + } catch (e) { + console.warn( + `[metering] AI cost factor lookup failed for ${key}: ${(e as Error).message}`, + ); + return 1; + } + } + // -- Public API: increment usage ---------------------------------- utilRecordUsageObject>( diff --git a/src/backend/services/metering/aiCostFactor.test.ts b/src/backend/services/metering/aiCostFactor.test.ts new file mode 100644 index 000000000..69127d3da --- /dev/null +++ b/src/backend/services/metering/aiCostFactor.test.ts @@ -0,0 +1,266 @@ +import { + afterEach, + beforeAll, + beforeEach, + describe, + expect, + it, + vi, +} from 'vitest'; +import type { Actor } from '../../core/actor.ts'; +import type { AiCostFactorEvent } from '../../clients/event/types.ts'; +import { PuterServer } from '../../server.ts'; +import { setupTestServer } from '../../testUtil.ts'; +import { aiModelKey } from './aiCostFactor.ts'; +import type { MeteringService } from './MeteringService.ts'; + +type Listener = ( + key: `ai.cost.factor.${string}`, + data: AiCostFactorEvent, +) => void; + +describe('AI cost factor', () => { + let server: PuterServer; + let metering: MeteringService; + let scoped: MeteringService; + let actor: Actor; + let listeners: Listener[]; + + beforeAll(async () => { + server = await setupTestServer(); + metering = server.services.metering; + scoped = metering.withAiCostFactor('ai-chat'); + // Counters buffer before they're written onward; stop the drain loop + // so nothing fires mid-assertion. + await server.stores.meteringBuffer.onServerShutdown(); + }); + + beforeEach(() => { + listeners = []; + actor = { + user: { + uuid: `ai-mult-${Math.random().toString(36).slice(2)}`, + username: 'ai-mult', + email: 'ai-mult@test.com', + }, + } as Actor; + }); + + afterEach(() => { + for (const listener of listeners) { + server.clients.event.off('ai.cost.factor.*', listener); + } + }); + + /** Subscribe for the duration of one test, recording what it was asked. */ + const listen = (factor: number | undefined) => { + const seen: Array<{ key: string; event: AiCostFactorEvent }> = []; + const listener: Listener = (key, event) => { + seen.push({ key, event: { ...event } }); + if (factor !== undefined) event.factor = factor; + }; + server.clients.event.on('ai.cost.factor.*', listener); + listeners.push(listener); + return seen; + }; + + describe('aiModelKey', () => { + it.each([ + [ + 'claude:claude-sonnet-4-5:input_tokens', + 'claude:claude-sonnet-4-5', + ], + ['gemini:gemini-2.5-flash:output:audio', 'gemini:gemini-2.5-flash'], + ['xai:stt:second', 'xai:stt'], + ['mistral-ocr', 'mistral-ocr'], + ])('reads the provider and model out of %s', (usageType, expected) => { + expect(aiModelKey(usageType)).toBe(expected); + }); + }); + + it('records the provider cost when nothing is listening', async () => { + const result = await scoped.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 10, + 1000, + ); + expect(result.total).toBe(1000); + }); + + // Drivers fire metering off unawaited. With no hook in play the record + // must still be issued in the same call, not a tick later — a request that + // ends in between would otherwise leave it settling behind the response. + it('records without waiting on the hook when nothing is listening', async () => { + const recorded: unknown[] = []; + const spy = vi + .spyOn(metering, 'incrementUsage') + .mockImplementation(async (...args) => { + recorded.push(args); + return { total: 0 }; + }); + try { + void scoped.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 1, + 1000, + ); + expect(recorded).toHaveLength(1); + } finally { + spy.mockRestore(); + } + }); + + it('multiplies the recorded cost and reports the driver and model', async () => { + const seen = listen(1.04); + + const result = await scoped.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 10, + 1000, + ); + + expect(result.total).toBe(1040); + expect(seen).toHaveLength(1); + expect(seen[0].key).toBe('ai.cost.factor.ai-chat.claude:sonnet'); + expect(seen[0].event).toMatchObject({ + driver: 'ai-chat', + model: 'claude:sonnet', + factor: 1, + }); + expect(seen[0].event.actor.user.uuid).toBe(actor.user.uuid); + }); + + it('scales the cost overrides of a recorded usage object', async () => { + listen(1.04); + + const result = await scoped.utilRecordUsageObject( + { input_tokens: 100, output_tokens: 50 }, + actor, + 'claude:sonnet', + { input_tokens: 1000, output_tokens: 2000 }, + ); + + expect(result.total).toBe(1040 + 2080); + }); + + it('asks once per model in a batch', async () => { + const seen = listen(2); + + const result = await scoped.batchIncrementUsages(actor, [ + { + usageType: 'claude:sonnet:input_tokens', + usageAmount: 1, + costOverride: 100, + }, + { + usageType: 'claude:sonnet:output_tokens', + usageAmount: 1, + costOverride: 200, + }, + { + usageType: 'openai:gpt-5:input_tokens', + usageAmount: 1, + costOverride: 400, + }, + ]); + + expect(result.total).toBe(1400); + expect(seen.map((s) => s.key)).toEqual([ + 'ai.cost.factor.ai-chat.claude:sonnet', + 'ai.cost.factor.ai-chat.openai:gpt-5', + ]); + }); + + // An entry with no cost is recorded unpriced, and multiplying "unpriced" + // would invent a price of zero. + it('leaves an unpriced entry alone', async () => { + const seen = listen(1.04); + + const result = await scoped.batchIncrementUsages(actor, [ + { usageType: 'claude:sonnet:input_tokens', usageAmount: 3 }, + ]); + + expect(result.total).toBe(0); + expect(result['claude:sonnet:input_tokens']).toMatchObject({ + units: 3, + }); + expect(seen).toHaveLength(0); + }); + + it.each([ + ['zero', 0], + ['negative', -2], + ['past the ceiling', 1000], + ['not a number', Number.NaN], + ])('ignores a %s factor', async (_label, factor) => { + listen(factor); + + const result = await scoped.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 1, + 1000, + ); + + expect(result.total).toBe(1000); + }); + + it('does not multiply usage recorded through the unscoped service', async () => { + listen(1.04); + + const result = await metering.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 1, + 1000, + ); + + expect(result.total).toBe(1000); + }); + + // Providers hold the scoped service and call all of it, not just the + // recording methods. + it('passes everything else through to the service', async () => { + const [scopedSub, realSub] = await Promise.all([ + scoped.getActorSubscription(actor), + metering.getActorSubscription(actor), + ]); + expect(scopedSub).toEqual(realSub); + expect(scoped.getRegisteredPolicy(realSub.id)?.id).toBe(realSub.id); + }); + + it('re-scoping returns the same view rather than stacking factors', async () => { + listen(1.04); + const again = scoped.withAiCostFactor('ai-chat'); + expect(again).toBe(scoped); + + const result = await again.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 1, + 1000, + ); + expect(result.total).toBe(1040); + }); + + it('survives a listener that throws', async () => { + const boom = vi.fn(() => { + throw new Error('nope'); + }) as unknown as Listener; + server.clients.event.on('ai.cost.factor.*', boom); + listeners.push(boom); + + const result = await scoped.incrementUsage( + actor, + 'claude:sonnet:input_tokens', + 1, + 1000, + ); + + expect(boom).toHaveBeenCalled(); + expect(result.total).toBe(1000); + }); +}); diff --git a/src/backend/services/metering/aiCostFactor.ts b/src/backend/services/metering/aiCostFactor.ts new file mode 100644 index 000000000..601c95339 --- /dev/null +++ b/src/backend/services/metering/aiCostFactor.ts @@ -0,0 +1,191 @@ +/* + * Copyright (C) 2024-present Puter Technologies Inc. + * + * This file is part of Puter. + * + * Puter is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published + * by the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +import type { Actor } from '../../core/actor'; +import type { MeteringService } from './MeteringService.js'; +import type { UsageByType, UsageInput } from './types'; + +/** Anything past this is read as a mistake and the cost is left as-is. */ +export const MAX_AI_COST_FACTOR = 10; + +/** + * The `:` head of a usage type — usage types are written + * `::`, e.g. `xai:stt:second`. + */ +export const aiModelKey = (usageType: string): string => + usageType.split(':').slice(0, 2).join(':'); + +const scaleCost = (cost: number, factor: number): number => + Math.round(cost * factor); + +/** One facade per service + driver, so call sites can ask for theirs freely. */ +const facades = new WeakMap>(); + +/** + * A view of `metering` whose recorded AI costs pass through the + * `ai.cost.factor..` hook. Everything else is the service + * itself, untouched. + */ +export function withAiCostFactor( + metering: MeteringService, + driver: string, +): MeteringService { + const forService = facades.get(metering) ?? new Map(); + facades.set(metering, forService); + const existing = forService.get(driver); + if (existing) return existing; + /** Whether a hook prices this model. Synchronous. */ + const hooked = (usage: UsageInput): boolean => + !!usage?.usageType && + Number.isFinite(usage.costOverride) && + metering.hasAiCostFactor(driver, aiModelKey(usage.usageType)); + + /** Factors for one recorded batch, resolved once per model. */ + const scaleUsages = async ( + actor: Actor, + usages: UsageInput[], + ): Promise => { + const byModel = new Map(); + const scaled: UsageInput[] = []; + for (const usage of usages) { + if (!hooked(usage)) { + scaled.push(usage); + continue; + } + const model = aiModelKey(usage.usageType); + let factor = byModel.get(model); + if (factor === undefined) { + factor = await metering.resolveAiCostFactor( + actor, + driver, + model, + ); + byModel.set(model, factor); + } + scaled.push( + factor === 1 + ? usage + : { + ...usage, + costOverride: scaleCost( + usage.costOverride as number, + factor, + ), + }, + ); + } + return scaled; + }; + + // Unhooked calls pass straight through, unawaited: callers fire metering + // off, and an await here would outlive the request. + + const incrementUsage = ( + actor: Actor, + usageType: string, + usageAmount: number, + costOverride?: number, + ): Promise => { + const usage = { usageType, usageAmount, costOverride }; + if (!hooked(usage)) + return metering.incrementUsage( + actor, + usageType, + usageAmount, + costOverride, + ); + return scaleUsages(actor, [usage]).then(([scaled]) => + metering.incrementUsage( + actor, + usageType, + usageAmount, + scaled?.costOverride, + ), + ); + }; + + const batchIncrementUsages = ( + actor: Actor, + usages: UsageInput[], + ): Promise => { + if (!usages?.some(hooked)) + return metering.batchIncrementUsages(actor, usages); + return scaleUsages(actor, usages).then((scaled) => + metering.batchIncrementUsages(actor, scaled), + ); + }; + + const utilRecordUsageObject = >( + trackedUsageObject: T, + actor: Actor, + modelPrefix: string, + costsOverrides?: Partial>, + ): Promise => { + // The prefix is the model here, so one lookup covers every entry. + if (!costsOverrides || !metering.hasAiCostFactor(driver, modelPrefix)) + return metering.utilRecordUsageObject( + trackedUsageObject, + actor, + modelPrefix, + costsOverrides, + ); + const scaleOverrides = (factor: number) => + factor === 1 + ? costsOverrides + : (Object.fromEntries( + Object.entries(costsOverrides).map(([key, cost]) => [ + key, + Number.isFinite(cost) + ? scaleCost(cost as number, factor) + : cost, + ]), + ) as Partial>); + + return metering + .resolveAiCostFactor(actor, driver, modelPrefix) + .then((factor) => + metering.utilRecordUsageObject( + trackedUsageObject, + actor, + modelPrefix, + scaleOverrides(factor), + ), + ); + }; + + const overrides: Record = { + incrementUsage, + batchIncrementUsages, + utilRecordUsageObject, + }; + + // A proxy, not a wrapper: providers use far more of the service than the + // three recording methods. + const facade = new Proxy(metering, { + get(target, prop, receiver) { + if (prop in overrides) return overrides[prop as string]; + // Already scoped — re-scoping would multiply twice. + if (prop === 'withAiCostFactor') return () => receiver; + const value = Reflect.get(target, prop, target); + return typeof value === 'function' ? value.bind(target) : value; + }, + }); + forService.set(driver, facade); + return facade; +} diff --git a/src/backend/stores/fs/FSEntryStore.ts b/src/backend/stores/fs/FSEntryStore.ts index d9420297a..99d5f75c1 100644 --- a/src/backend/stores/fs/FSEntryStore.ts +++ b/src/backend/stores/fs/FSEntryStore.ts @@ -1226,6 +1226,35 @@ export class FSEntryStore extends PuterStore { return entriesByPath; } + /** + * The entry as it stands after a patch this store just wrote. The writer + * already knows every column it set, so all it needs is a base, and a + * cache-first read supplies that without the cross-region round trip a + * primary read costs. + * + * Only for patches that leave `id`, `uuid` and `path` alone — the cache + * keys derive from those. A writer changing other columns in the same + * instant can have one field served stale until the entry's TTL lapses; + * reading the primary on every mutation is the alternative. + */ + async #entryAfterPatch( + uuid: string, + patch: Partial, + ): Promise { + const base = await this.getEntryByUuid(uuid); + if (!base) return null; + const entry = { ...base, ...patch }; + // Broadcast the new value rather than a hole: peers would otherwise + // serve their own cached copy of the pre-patch row until its TTL. + await this.publishCacheKeys({ + keys: this.#entryCacheKeys(entry), + serializedData: JSON.stringify(entry), + ttlSeconds: ENTRY_CACHE_TTL_SECONDS, + broadcast: true, + }); + return entry; + } + async getEntryByUuid(id: string): Promise { const cacheKey = `prodfsv2:fsentry:uuid:${id}`; const cached = await this.#readEntryFromCache(cacheKey); @@ -1384,22 +1413,21 @@ export class FSEntryStore extends PuterStore { } } - const refreshedRows = (await this.clients.db.pread( - `SELECT ${this.#selectFsentriesColumns()} FROM fsentries WHERE uuid = ? AND user_id = ? LIMIT 1`, - [uuid, userId], - )) as unknown as FSEntryRow[]; - const refreshedRow = refreshedRows[0]; - if (!refreshedRow) { + // The UPDATE above already proved the row exists and belongs to this + // user, and named every column that moved — so the result is knowable + // without reading the row back. + const updatedEntry = await this.#entryAfterPatch(uuid, { + thumbnail, + modified: now, + accessed: now, + }); + if (!updatedEntry || updatedEntry.userId !== userId) { throw new HttpError( 404, 'File entry was not found for thumbnail update', { legacyCode: 'not_found' }, ); } - - const updatedEntry = this.#mapFSEntryRow(refreshedRow); - await this.#invalidateEntryCache(updatedEntry); - await this.#writeEntryToCache(updatedEntry); return updatedEntry; } @@ -2310,7 +2338,15 @@ export class FSEntryStore extends PuterStore { input.kind === 'symlink', ); - await this.clients.db.write( + const isPublic = + input.isPublic === undefined || input.isPublic === null + ? null + : this.clients.db.booleanValue(input.isPublic); + const immutable = this.clients.db.booleanValue( + Boolean(input.immutable), + ); + + const written = await this.clients.db.write( `INSERT INTO fsentries ( uuid, user_id, @@ -2348,10 +2384,8 @@ export class FSEntryStore extends PuterStore { input.associatedAppId ?? null, input.metadata ?? null, input.thumbnail ?? null, - this.clients.db.booleanValue(Boolean(input.immutable)), - input.isPublic === undefined || input.isPublic === null - ? null - : this.clients.db.booleanValue(input.isPublic), + immutable, + isPublic, now, now, now, @@ -2359,6 +2393,52 @@ export class FSEntryStore extends PuterStore { ], ); + // The insert supplied every column; the rest take their schema default + // and a row this new has no subdomains. Reading it back would only + // return what we just sent, at the price of a primary round trip on a + // path app launches wait for. + const insertId = Number(written.insertId); + const row: FSEntryRow = insertId + ? ({ + id: insertId, + uuid, + user_id: input.parent.userId, + parent_id: input.parent.id, + parent_uid: input.parent.uuid, + name: input.name, + path, + is_dir: isDir, + is_shortcut: isShortcut, + shortcut_to: input.shortcutTo ?? null, + is_symlink: isSymlink, + symlink_path: input.symlinkPath ?? null, + associated_app_id: input.associatedAppId ?? null, + metadata: input.metadata ?? null, + thumbnail: input.thumbnail ?? null, + immutable, + is_public: isPublic, + created: now, + modified: now, + accessed: now, + size: 0, + bucket: null, + bucket_region: null, + public_token: null, + file_request_token: null, + layout: null, + sort_by: null, + sort_order: null, + subdomains_agg: null, + } as unknown as FSEntryRow) + : await this.#readCreatedEntryRow(uuid); + + const entry = this.#mapFSEntryRow(row); + await this.#writeEntryToCache(entry); + return entry; + } + + /** Fallback for engines that report no insert id: the row we just wrote. */ + async #readCreatedEntryRow(uuid: string): Promise { const rows = (await this.clients.db.pread( `SELECT ${this.#selectFsentriesColumns()} FROM fsentries WHERE uuid = ? LIMIT 1`, [uuid], @@ -2369,9 +2449,7 @@ export class FSEntryStore extends PuterStore { legacyCode: 'internal_error', }); } - const entry = this.#mapFSEntryRow(row); - await this.#writeEntryToCache(entry); - return entry; + return row; } /** @@ -2389,42 +2467,41 @@ export class FSEntryStore extends PuterStore { const now = Math.floor(Date.now() / 1000); const assignments: string[] = []; const values: unknown[] = []; + const patch: Partial = {}; if (options.setAccessed) { assignments.push('accessed = ?'); values.push(now); + patch.accessed = now; } if (options.setModified) { assignments.push('modified = ?'); values.push(now); + patch.modified = now; } if (options.setCreated) { assignments.push('created = ?'); values.push(now); + patch.created = now; } if (assignments.length === 0) { // Default: touch all three. assignments.push('accessed = ?', 'modified = ?', 'created = ?'); values.push(now, now, now); + patch.accessed = now; + patch.modified = now; + patch.created = now; } await this.clients.db.write( `UPDATE fsentries SET ${assignments.join(', ')} WHERE uuid = ?`, [...values, uuid], ); - // Re-read the row itself rather than going through `getEntryByUuid`: - // that read is cache-first and would hand back the pre-touch - // timestamps (and then re-cache them for another TTL). - const refreshedRows = (await this.clients.db.pread( - `SELECT ${this.#selectFsentriesColumns()} FROM fsentries WHERE uuid = ? LIMIT 1`, - [uuid], - )) as unknown as FSEntryRow[]; - const refreshedRow = refreshedRows[0]; - if (!refreshedRow) + // The timestamps above are the only columns that moved, so the entry + // is knowable without reading the row back. + const entry = await this.#entryAfterPatch(uuid, patch); + if (!entry) throw new HttpError(404, 'Entry not found after touch', { legacyCode: 'not_found', }); - const entry = this.#mapFSEntryRow(refreshedRow); - await this.#invalidateEntryCache(entry); - await this.#writeEntryToCache(entry); return entry; } @@ -2489,7 +2566,8 @@ export class FSEntryStore extends PuterStore { } = {}, ): Promise<{ entries: FSEntry[]; cursor?: string }> { const payload = decodeCursor(options.cursor) as - { v: unknown; id: number; s?: string; o?: string } | undefined; + | { v: unknown; id: number; s?: string; o?: string } + | undefined; const requestedSort = options.sortBy ?? null; const requestedOrder = options.sortOrder ?? null; @@ -2675,7 +2753,8 @@ export class FSEntryStore extends PuterStore { const limit = normalizeLimit(options.limit, { cap: 10_000 }) ?? 1000; const payload = decodeCursor(options.cursor) as - { p: string } | undefined; + | { p: string } + | undefined; const seek = payload ? 'AND path > ?' : ''; const params: unknown[] = payload ? [userId, likePattern, maxSlashes, payload.p, limit + 1] diff --git a/src/backend/stores/fs/FSEntryStore.writeBack.test.ts b/src/backend/stores/fs/FSEntryStore.writeBack.test.ts new file mode 100644 index 000000000..c76d6bce4 --- /dev/null +++ b/src/backend/stores/fs/FSEntryStore.writeBack.test.ts @@ -0,0 +1,242 @@ +/* + * Copyright (C) 2024-present Puter Technologies Inc. + * + * This file is part of Puter. + * + * Puter is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published + * by the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +import { describe, expect, it, vi } from 'vitest'; +import type { IConfig } from '../../types.js'; +import { FSEntryStore } from './FSEntryStore.js'; + +const makeStore = ( + writeResult: { insertId: number }, + cachedEntry?: Record, +) => { + const pread = vi.fn(async () => []); + const write = vi.fn(async () => ({ + insertId: writeResult.insertId, + affectedRows: 1, + anyRowsAffected: true, + })); + const setex = vi.fn(async () => 'OK'); + const eventEmit = vi.fn(); + const pipelineSet = vi.fn(); + const pipeline = { + del: vi.fn(), + set: pipelineSet, + exec: vi.fn(async () => []), + }; + const clients = { + db: { + write, + pread, + read: vi.fn(async () => []), + tryHardRead: vi.fn(async () => []), + booleanValue: (value: boolean) => (value ? 1 : 0), + case: () => '', + insertIgnoreInto: () => '', + }, + redis: { + setex, + get: vi.fn(async () => + cachedEntry ? JSON.stringify(cachedEntry) : null, + ), + del: vi.fn(async () => 1), + pipeline: () => pipeline, + }, + event: { emit: eventEmit }, + }; + const store = new FSEntryStore( + {} as IConfig, + clients as never, + {} as never, + ); + return { store, clients, write, pread, setex, eventEmit, pipelineSet }; +}; + +const cachedFile = { + id: 5, + uuid: '33333333-3333-4333-8333-333333333333', + uid: '33333333-3333-4333-8333-333333333333', + userId: 3, + parentId: 7, + parentUid: '11111111-1111-4111-8111-111111111111', + path: '/alice/report.txt', + name: 'report.txt', + isDir: false, + size: 1234, + thumbnail: null, + accessed: 100, + modified: 100, + created: 100, + subdomains: [], + workers: [], + hasWebsite: false, + suggestedApps: [], +}; + +const parent = { + id: 7, + uuid: '11111111-1111-4111-8111-111111111111', + userId: 3, + path: '/alice', +}; + +describe('FSEntryStore.createNonFileEntry', () => { + it('returns the created entry without reading it back from the primary', async () => { + const { store, pread, write } = makeStore({ insertId: 42 }); + + const entry = await store.createNonFileEntry({ + parent, + name: 'AppData', + kind: 'directory', + thumbnail: 'https://example.invalid/icon.png', + } as never); + + // The read-back was a primary round trip on a path app launches wait + // for, and the insert already supplied every column. + expect(pread).not.toHaveBeenCalled(); + expect(write).toHaveBeenCalledTimes(1); + expect(entry).toMatchObject({ + id: 42, + userId: 3, + parentId: 7, + parentUid: parent.uuid, + name: 'AppData', + path: '/alice/AppData', + isDir: true, + isShortcut: false, + isSymlink: false, + immutable: false, + thumbnail: 'https://example.invalid/icon.png', + size: 0, + subdomains: [], + hasWebsite: false, + }); + expect(entry.uuid).toMatch(/^[0-9a-f-]{36}$/); + expect(entry.uid).toBe(entry.uuid); + }); + + it('falls back to reading the row when the engine reports no insert id', async () => { + const { store, pread } = makeStore({ insertId: 0 }); + pread.mockResolvedValueOnce([ + { + id: 99, + uuid: '22222222-2222-4222-8222-222222222222', + user_id: 3, + parent_id: 7, + parent_uid: parent.uuid, + name: 'AppData', + path: '/alice/AppData', + is_dir: 1, + is_shortcut: 0, + is_symlink: 0, + immutable: 0, + modified: 1, + created: 1, + accessed: 1, + size: 0, + }, + ] as never); + + const entry = await store.createNonFileEntry({ + parent, + name: 'AppData', + kind: 'directory', + } as never); + + expect(pread).toHaveBeenCalledTimes(1); + expect(entry.id).toBe(99); + }); +}); + +describe('FSEntryStore.touchEntryTimestamps', () => { + it('applies the touched timestamps without reading the primary', async () => { + const { store, pread, write, eventEmit, pipelineSet } = makeStore( + { insertId: 0 }, + cachedFile, + ); + + const entry = await store.touchEntryTimestamps(cachedFile.uuid, { + setModified: true, + }); + + expect(pread).not.toHaveBeenCalled(); + expect(write).toHaveBeenCalledTimes(1); + // Only `modified` was assigned, so the others keep the cached values. + expect(entry.modified).toBeGreaterThan(cachedFile.modified); + expect(entry.accessed).toBe(100); + expect(entry.created).toBe(100); + // Columns the update never named survive untouched. + expect(entry.name).toBe('report.txt'); + expect(entry.size).toBe(1234); + // Peers must get the patched row, not a hole they would refill from a + // lagging replica. + expect(pipelineSet).toHaveBeenCalled(); + expect(eventEmit).toHaveBeenCalledWith( + 'outer.cacheUpdate', + expect.objectContaining({ data: expect.any(String) }), + {}, + ); + }); + + it('touches all three timestamps when none is named', async () => { + const { store } = makeStore({ insertId: 0 }, cachedFile); + + const entry = await store.touchEntryTimestamps(cachedFile.uuid, {}); + + expect(entry.accessed).toBeGreaterThan(100); + expect(entry.modified).toBeGreaterThan(100); + expect(entry.created).toBeGreaterThan(100); + }); + + it('reports not found when the entry is gone', async () => { + const { store } = makeStore({ insertId: 0 }); + + await expect( + store.touchEntryTimestamps(cachedFile.uuid, { setModified: true }), + ).rejects.toMatchObject({ statusCode: 404 }); + }); +}); + +describe('FSEntryStore.updateEntryThumbnailByUuidForUser', () => { + it('applies the new thumbnail without reading the primary', async () => { + const { store, pread, write } = makeStore({ insertId: 0 }, cachedFile); + + const entry = await store.updateEntryThumbnailByUuidForUser( + 3, + cachedFile.uuid, + 'data:image/png;base64,AAAA', + ); + + expect(pread).not.toHaveBeenCalled(); + expect(write).toHaveBeenCalledTimes(1); + expect(entry.thumbnail).toBe('data:image/png;base64,AAAA'); + expect(entry.modified).toBeGreaterThan(100); + expect(entry.name).toBe('report.txt'); + }); + + it('refuses to hand back an entry owned by another user', async () => { + const { store } = makeStore( + { insertId: 0 }, + { ...cachedFile, userId: 99 }, + ); + + await expect( + store.updateEntryThumbnailByUuidForUser(3, cachedFile.uuid, null), + ).rejects.toMatchObject({ statusCode: 404 }); + }); +});