194 lines
6.3 KiB
JavaScript
194 lines
6.3 KiB
JavaScript
import test from 'node:test';
|
|
import assert from 'node:assert/strict';
|
|
import { createUserAuth } from './user-auth.mjs';
|
|
|
|
function createConcurrentBillingPool({ userId, sessionId, startingBalanceCents = 100_000 }) {
|
|
const billingState = new Map();
|
|
const usageRecords = [];
|
|
const ledger = [];
|
|
let wallet = { balance_cents: startingBalanceCents, tokens_used: 0 };
|
|
let sessionLock = Promise.resolve();
|
|
|
|
const createConnection = () => {
|
|
let holdsSessionLock = false;
|
|
let releaseSessionLock = null;
|
|
|
|
const acquireSessionLock = async () => {
|
|
if (holdsSessionLock) return;
|
|
let release;
|
|
const gate = new Promise((resolve) => {
|
|
release = resolve;
|
|
});
|
|
const prev = sessionLock;
|
|
sessionLock = gate;
|
|
await prev;
|
|
holdsSessionLock = true;
|
|
releaseSessionLock = release;
|
|
};
|
|
|
|
const releaseLock = () => {
|
|
if (!holdsSessionLock) return;
|
|
holdsSessionLock = false;
|
|
releaseSessionLock?.();
|
|
releaseSessionLock = null;
|
|
};
|
|
|
|
return {
|
|
async beginTransaction() {},
|
|
async commit() {
|
|
releaseLock();
|
|
},
|
|
async rollback() {
|
|
releaseLock();
|
|
},
|
|
release() {},
|
|
async query(sql, params = []) {
|
|
if (
|
|
sql.includes('INSERT INTO h5_session_billing_state') &&
|
|
sql.includes('ON DUPLICATE KEY UPDATE agent_session_id')
|
|
) {
|
|
if (!billingState.has(sessionId)) {
|
|
billingState.set(sessionId, {
|
|
last_accumulated_cost: null,
|
|
last_input_tokens: 0,
|
|
last_output_tokens: 0,
|
|
});
|
|
}
|
|
return [{ affectedRows: 1 }, []];
|
|
}
|
|
if (sql.includes('FROM h5_session_billing_state') && sql.includes('FOR UPDATE')) {
|
|
await acquireSessionLock();
|
|
const row = billingState.get(sessionId) ?? {
|
|
last_accumulated_cost: null,
|
|
last_input_tokens: 0,
|
|
last_output_tokens: 0,
|
|
};
|
|
return [[{ ...row }], []];
|
|
}
|
|
if (
|
|
sql.includes('INSERT INTO h5_session_billing_state') &&
|
|
sql.includes('last_output_tokens = VALUES(last_output_tokens)')
|
|
) {
|
|
const [, , , inputTokens, outputTokens] = params;
|
|
billingState.set(sessionId, {
|
|
last_accumulated_cost: params[2],
|
|
last_input_tokens: inputTokens,
|
|
last_output_tokens: outputTokens,
|
|
});
|
|
return [{ affectedRows: 1 }, []];
|
|
}
|
|
if (sql.includes('FROM h5_usage_records WHERE request_id')) {
|
|
const requestId = params[0];
|
|
const hit = usageRecords.find((row) => row.request_id === requestId);
|
|
return hit ? [[{ cost_cents: hit.cost_cents }], []] : [[], []];
|
|
}
|
|
if (sql.includes('FROM h5_user_wallets') && sql.includes('FOR UPDATE')) {
|
|
return [[{ ...wallet }], []];
|
|
}
|
|
if (sql.includes('FROM h5_user_wallets') && !sql.includes('FOR UPDATE')) {
|
|
return [[{ ...wallet }], []];
|
|
}
|
|
if (sql.includes('UPDATE h5_user_wallets')) {
|
|
const [nextBalance, deltaTokens] = params;
|
|
wallet = {
|
|
balance_cents: nextBalance,
|
|
tokens_used: Number(wallet.tokens_used) + Number(deltaTokens),
|
|
};
|
|
return [{ affectedRows: 1 }, []];
|
|
}
|
|
if (sql.includes('INSERT INTO h5_usage_records')) {
|
|
usageRecords.push({
|
|
request_id: params[2],
|
|
input_tokens: params[3],
|
|
output_tokens: params[4],
|
|
cost_cents: params[5],
|
|
});
|
|
return [{ insertId: usageRecords.length }, []];
|
|
}
|
|
if (sql.includes('INSERT INTO h5_billing_ledger')) {
|
|
ledger.push({ amount_cents: params[1], note: params[5] });
|
|
return [{ affectedRows: 1 }, []];
|
|
}
|
|
if (sql.includes('UPDATE h5_users SET status')) {
|
|
return [{ affectedRows: 1 }, []];
|
|
}
|
|
throw new Error(`unexpected query: ${sql}`);
|
|
},
|
|
};
|
|
};
|
|
|
|
const userRow = {
|
|
id: userId,
|
|
username: 'tester',
|
|
slug: 'tester',
|
|
email: 'tester@example.com',
|
|
display_name: 'Tester',
|
|
role: 'user',
|
|
status: 'active',
|
|
plan_type: 'free',
|
|
workspace_root: '/tmp/tester',
|
|
balance_cents: startingBalanceCents,
|
|
tokens_used: 0,
|
|
};
|
|
|
|
const pool = {
|
|
async query(sql, params = []) {
|
|
if (sql.includes('FROM h5_users u') && sql.includes('WHERE u.id = ?')) {
|
|
return [[{ ...userRow, balance_cents: wallet.balance_cents, tokens_used: wallet.tokens_used }], []];
|
|
}
|
|
throw new Error(`unexpected pool query: ${sql}`);
|
|
},
|
|
async getConnection() {
|
|
return createConnection();
|
|
},
|
|
};
|
|
|
|
return { pool, usageRecords, ledger, getWallet: () => wallet };
|
|
}
|
|
|
|
test('billSessionUsage charges once under concurrent Finish replays', async () => {
|
|
const userId = 'user-concurrent-1';
|
|
const sessionId = '20260629_test';
|
|
const { pool, usageRecords, ledger, getWallet } = createConcurrentBillingPool({
|
|
userId,
|
|
sessionId,
|
|
startingBalanceCents: 100_000,
|
|
});
|
|
const auth = createUserAuth(pool, { persistSessions: false });
|
|
const tokenState = {
|
|
accumulatedInputTokens: 4877927,
|
|
accumulatedOutputTokens: 23310,
|
|
};
|
|
|
|
const results = await Promise.all(
|
|
Array.from({ length: 11 }, () => auth.billSessionUsage(userId, sessionId, tokenState, null)),
|
|
);
|
|
|
|
const charged = results.filter((result) => result.costCents > 0);
|
|
assert.equal(charged.length, 1);
|
|
assert.equal(usageRecords.length, 1);
|
|
assert.equal(ledger.length, 1);
|
|
assert.equal(usageRecords[0].input_tokens, 4877927);
|
|
assert.equal(usageRecords[0].output_tokens, 23310);
|
|
assert.ok(getWallet().balance_cents < 100_000);
|
|
});
|
|
|
|
test('billSessionUsage is idempotent when request_id repeats', async () => {
|
|
const userId = 'user-concurrent-2';
|
|
const sessionId = '20260629_req';
|
|
const { pool, usageRecords } = createConcurrentBillingPool({
|
|
userId,
|
|
sessionId,
|
|
startingBalanceCents: 50_000,
|
|
});
|
|
const auth = createUserAuth(pool, { persistSessions: false });
|
|
const tokenState = { accumulatedInputTokens: 1000, accumulatedOutputTokens: 200 };
|
|
|
|
const first = await auth.billSessionUsage(userId, sessionId, tokenState, 'req-abc');
|
|
const second = await auth.billSessionUsage(userId, sessionId, tokenState, 'req-abc');
|
|
|
|
assert.ok(first.costCents > 0);
|
|
assert.equal(second.costCents, 0);
|
|
assert.equal(usageRecords.length, 1);
|
|
});
|