mirror of
https://github.com/HeyPuter/puter.git
synced 2026-09-29 08:38:06 +00:00
feat(metering): AI cost multiplier hook for AI drivers (#3898)
Emits ai.cost.multiplier.<driver>.<provider>:<model> 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.
This commit is contained in:
@@ -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<typeof setTimeout> | 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<Record<string, unknown>[]> {
|
||||
const result = await this.db.execute(query, params);
|
||||
const result = await this.dbPrimaryRead.execute(query, params);
|
||||
if (!result) return [];
|
||||
return (result[0] as Record<string, unknown>[]) ?? [];
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -693,6 +693,11 @@ export type EventMap = {
|
||||
// normalized path: `route.<method>.<path>.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.<driverName>.<provider>:<model>`. 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.<driver>.<model>` events. */
|
||||
export type AiCostFactorEvent = {
|
||||
/** Driver doing the pricing, e.g. `ai-chat`. */
|
||||
driver: string;
|
||||
/** `<provider>:<model>` 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
|
||||
|
||||
@@ -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<string, IChatProvider> = {};
|
||||
#modelIdMap: Record<string, IChatModel[]> = {};
|
||||
|
||||
/** 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<string, unknown> | undefined) =>
|
||||
(cfg?.apiKey as string | undefined) ??
|
||||
|
||||
@@ -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<string, IImageProvider> = {};
|
||||
#modelIdMap: Record<string, IImageModel[]> = {};
|
||||
|
||||
/** 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<Record<string, unknown> | undefined>
|
||||
|
||||
@@ -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<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| undefined;
|
||||
const textractAws = (textract?.aws ?? textract) as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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,
|
||||
|
||||
@@ -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<string, unknown>[] {
|
||||
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<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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,
|
||||
|
||||
@@ -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<string, ISpeechToTextProvider> = {};
|
||||
|
||||
/** 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, {
|
||||
|
||||
@@ -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<string, ITTSProvider> = {};
|
||||
|
||||
/** 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<string, unknown> | undefined) ??
|
||||
@@ -266,7 +272,8 @@ export class TTSDriver extends PuterDriver {
|
||||
}
|
||||
|
||||
const elevenlabs = providers['elevenlabs'] as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| undefined;
|
||||
const pollyAws = (polly?.aws ?? polly) as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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<string, unknown>) {
|
||||
const m = this.services.metering;
|
||||
const m = this.#aiMetering;
|
||||
const gemini = (providers['gemini'] ?? providers['gemini-tts']) as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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<string, unknown>) {
|
||||
const m = this.services.metering;
|
||||
const m = this.#aiMetering;
|
||||
const xai = (providers['xai'] ?? providers['xai-tts']) as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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<string, unknown>) {
|
||||
const m = this.services.metering;
|
||||
const m = this.#aiMetering;
|
||||
const speechify = (providers['speechify'] ??
|
||||
providers['speechify-tts']) as Record<string, unknown> | undefined;
|
||||
const speechifyKey =
|
||||
|
||||
@@ -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<string, IVideoProvider> = {};
|
||||
#modelIdMap: Record<string, IVideoModel[]> = {};
|
||||
|
||||
/** 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<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| undefined;
|
||||
const byteplusSharedCfg = providers['byteplus'] as
|
||||
Record<string, unknown> | undefined;
|
||||
| Record<string, unknown>
|
||||
| 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,
|
||||
);
|
||||
|
||||
@@ -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<string, unknown>;
|
||||
},
|
||||
{ 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<string> => {
|
||||
@@ -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<string, unknown>;
|
||||
},
|
||||
{ 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);
|
||||
});
|
||||
|
||||
|
||||
@@ -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.<driver>.<model>` 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<number> {
|
||||
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<T extends Record<string, number>>(
|
||||
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
@@ -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 <https://www.gnu.org/licenses/>.
|
||||
*/
|
||||
|
||||
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 `<provider>:<model>` head of a usage type — usage types are written
|
||||
* `<provider>:<model>:<what>`, 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<MeteringService, Map<string, MeteringService>>();
|
||||
|
||||
/**
|
||||
* A view of `metering` whose recorded AI costs pass through the
|
||||
* `ai.cost.factor.<driver>.<model>` 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<UsageInput[]> => {
|
||||
const byModel = new Map<string, number>();
|
||||
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<UsageByType> => {
|
||||
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<UsageByType> => {
|
||||
if (!usages?.some(hooked))
|
||||
return metering.batchIncrementUsages(actor, usages);
|
||||
return scaleUsages(actor, usages).then((scaled) =>
|
||||
metering.batchIncrementUsages(actor, scaled),
|
||||
);
|
||||
};
|
||||
|
||||
const utilRecordUsageObject = <T extends Record<string, number>>(
|
||||
trackedUsageObject: T,
|
||||
actor: Actor,
|
||||
modelPrefix: string,
|
||||
costsOverrides?: Partial<Record<keyof T, number>>,
|
||||
): Promise<UsageByType> => {
|
||||
// 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<Record<keyof T, number>>);
|
||||
|
||||
return metering
|
||||
.resolveAiCostFactor(actor, driver, modelPrefix)
|
||||
.then((factor) =>
|
||||
metering.utilRecordUsageObject(
|
||||
trackedUsageObject,
|
||||
actor,
|
||||
modelPrefix,
|
||||
scaleOverrides(factor),
|
||||
),
|
||||
);
|
||||
};
|
||||
|
||||
const overrides: Record<string, unknown> = {
|
||||
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;
|
||||
}
|
||||
@@ -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<FSEntry>,
|
||||
): Promise<FSEntry | null> {
|
||||
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<FSEntry | null> {
|
||||
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<FSEntryRow> {
|
||||
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<FSEntry> = {};
|
||||
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]
|
||||
|
||||
@@ -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 <https://www.gnu.org/licenses/>.
|
||||
*/
|
||||
|
||||
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<string, unknown>,
|
||||
) => {
|
||||
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 });
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user