diff --git a/apps/sim/app/api/billing/update-cost/route.integration.ts b/apps/sim/app/api/billing/update-cost/route.integration.ts index 082009c92dc..6263a693d6a 100644 --- a/apps/sim/app/api/billing/update-cost/route.integration.ts +++ b/apps/sim/app/api/billing/update-cost/route.integration.ts @@ -1,13 +1,14 @@ /** * Cost callbacks against real PostgreSQL: a direct-v1 run that outlives its admitted Stripe period * records its later spend in the payer's current period, so the closed period is never topped up - * after its invoice. Only the internal-key check is stubbed. + * after its invoice, and spend after the payer's terminal settlement is refused. Only the + * internal-key check is stubbed. */ import { db } from '@sim/db' import { subscription, usageLog, user, userStats } from '@sim/db/schema' import { envFlagsMock } from '@sim/testing/mocks/env-flags.mock' import { generateId } from '@sim/utils/id' -import { eq } from 'drizzle-orm' +import { eq, inArray } from 'drizzle-orm' import { NextRequest } from 'next/server' import { afterAll, describe, expect, it, vi } from 'vitest' @@ -25,20 +26,75 @@ import { BILLING_ACCOUNT_DECISION_HEADER, serializeAccountBillingDecisionHeader, } from '@/lib/billing/core/billing-attribution' +import { claimTerminalPeriod } from '@/lib/billing/cycle-close' import { POST } from '@/app/api/billing/update-cost/route' const DAY_MS = 24 * 60 * 60 * 1000 -const userId = `update-cost-user-${generateId()}` -const subscriptionId = generateId() +const userIds: string[] = [] afterAll(async () => { - await db.delete(usageLog).where(eq(usageLog.userId, userId)) - await db.delete(subscription).where(eq(subscription.id, subscriptionId)) - await db.delete(userStats).where(eq(userStats.userId, userId)) - await db.delete(user).where(eq(user.id, userId)) + if (userIds.length === 0) return + await db.delete(usageLog).where(inArray(usageLog.userId, userIds)) + await db.delete(subscription).where(inArray(subscription.referenceId, userIds)) + await db.delete(userStats).where(inArray(userStats.userId, userIds)) + await db.delete(user).where(inArray(user.id, userIds)) }) -function callback(requestKey: string, cost: number, decision: string): NextRequest { +interface Payer { + userId: string + subscriptionId: string + /** The direct-v1 decision of a run admitted in the subscription's period. */ + decision: string +} + +/** A user on a pro subscription for `period`, whose close marker has caught up to it. */ +async function createPayer(period: { start: Date; end: Date }): Promise { + const userId = `update-cost-user-${generateId()}` + const subscriptionId = generateId() + userIds.push(userId) + await db.insert(user).values({ + id: userId, + name: 'Update Cost Test', + email: `${userId}@update-cost.test`, + emailVerified: true, + createdAt: new Date(), + updatedAt: new Date(), + }) + await db.insert(userStats).values({ id: generateId(), userId }) + await db.insert(subscription).values({ + id: subscriptionId, + plan: 'pro', + referenceId: userId, + status: 'active', + periodStart: period.start, + periodEnd: period.end, + lastClosedPeriodStart: period.start, + }) + const decision = serializeAccountBillingDecisionHeader({ + userId, + billingEntity: { type: 'user', id: userId }, + billingPeriod: { + start: period.start.toISOString(), + end: period.end.toISOString(), + source: 'stripe', + }, + payerSubscriptionId: subscriptionId, + }) + return { userId, subscriptionId, decision } +} + +function requestRows(requestKey: string) { + return db + .select({ + eventKey: usageLog.eventKey, + cost: usageLog.cost, + billingPeriodStart: usageLog.billingPeriodStart, + }) + .from(usageLog) + .where(inArray(usageLog.eventKey, [`update-cost:${requestKey}`, `update-cost:${requestKey}@1`])) +} + +function callback(payer: Payer, requestKey: string, cost: number): NextRequest { return new NextRequest('http://localhost:3000/api/billing/update-cost', { method: 'POST', headers: { @@ -46,10 +102,10 @@ function callback(requestKey: string, cost: number, decision: string): NextReque 'x-api-key': 'internal', 'x-sim-billing-protocol': 'direct-v1', 'x-sim-billing-request-id': requestKey, - [BILLING_ACCOUNT_DECISION_HEADER]: decision, + [BILLING_ACCOUNT_DECISION_HEADER]: payer.decision, }, body: JSON.stringify({ - userId, + userId: payer.userId, cost, model: 'test-model', source: 'copilot', @@ -63,50 +119,17 @@ describe('direct-v1 cost callbacks in PostgreSQL', () => { const now = Date.now() const admitted = { start: new Date(now - 10 * DAY_MS), end: new Date(now + 20 * DAY_MS) } const rolled = { start: new Date(now - 60 * 60 * 1000), end: new Date(now + 30 * DAY_MS) } - await db.insert(user).values({ - id: userId, - name: 'Update Cost Test', - email: `${userId}@update-cost.test`, - emailVerified: true, - createdAt: new Date(now), - updatedAt: new Date(now), - }) - await db.insert(userStats).values({ id: generateId(), userId }) - await db.insert(subscription).values({ - id: subscriptionId, - plan: 'pro', - referenceId: userId, - status: 'active', - periodStart: admitted.start, - periodEnd: admitted.end, - }) - const decision = serializeAccountBillingDecisionHeader({ - userId, - billingEntity: { type: 'user', id: userId }, - billingPeriod: { - start: admitted.start.toISOString(), - end: admitted.end.toISOString(), - source: 'stripe', - }, - payerSubscriptionId: subscriptionId, - }) + const payer = await createPayer(admitted) const requestKey = generateId() - expect((await POST(callback(requestKey, 0.5, decision), {})).status).toBe(200) + expect((await POST(callback(payer, requestKey, 0.5), {})).status).toBe(200) await db .update(subscription) .set({ periodStart: rolled.start, periodEnd: rolled.end }) - .where(eq(subscription.id, subscriptionId)) - expect((await POST(callback(requestKey, 0.8, decision), {})).status).toBe(200) + .where(eq(subscription.id, payer.subscriptionId)) + expect((await POST(callback(payer, requestKey, 0.8), {})).status).toBe(200) - const rows = await db - .select({ - eventKey: usageLog.eventKey, - cost: usageLog.cost, - billingPeriodStart: usageLog.billingPeriodStart, - }) - .from(usageLog) - .where(eq(usageLog.userId, userId)) + const rows = await requestRows(requestKey) const byKey = new Map(rows.map((row) => [row.eventKey, row])) expect(rows).toHaveLength(2) expect(Number(byKey.get(`update-cost:${requestKey}`)?.cost)).toBeCloseTo(0.5) @@ -118,4 +141,19 @@ describe('direct-v1 cost callbacks in PostgreSQL', () => { rolled.start.getTime() ) }) + + it("refuses spend that lands after the payer's terminal settlement", async () => { + const now = Date.now() + const period = { start: new Date(now - 10 * DAY_MS), end: new Date(now + 20 * DAY_MS) } + const payer = await createPayer(period) + const requestKey = generateId() + + expect((await POST(callback(payer, requestKey, 0.5), {})).status).toBe(200) + await claimTerminalPeriod(payer.subscriptionId) + const late = await POST(callback(payer, requestKey, 0.8), {}) + + expect(late.status).toBe(409) + expect(await late.json()).toMatchObject({ code: 'BILLING_PERIOD_ELAPSED', retryable: false }) + expect((await requestRows(requestKey)).map((row) => Number(row.cost))).toEqual([0.5]) + }) }) diff --git a/apps/sim/app/api/billing/update-cost/route.ts b/apps/sim/app/api/billing/update-cost/route.ts index bec648abdc6..0c66f9fc346 100644 --- a/apps/sim/app/api/billing/update-cost/route.ts +++ b/apps/sim/app/api/billing/update-cost/route.ts @@ -30,6 +30,7 @@ import { import { type CumulativeUsageContextField, CumulativeUsageContextMismatchError, + CumulativeUsagePeriodClosedError, recordCumulativeUsage, } from '@/lib/billing/core/usage-log' import { @@ -512,7 +513,8 @@ async function updateCostInner(req: NextRequest, span: Span): Promise ({ db: { transaction }, dbReplica: {} })) vi.mock('@/lib/billing/core/plan', () => ({ getHighestPrioritySubscription: vi.fn() })) -vi.mock('@/lib/billing/subscriptions/utils', () => ({ isOrgScopedSubscription: vi.fn() })) +vi.mock('@/lib/billing/subscriptions/utils', async (importOriginal) => ({ + ...(await importOriginal()), + isOrgScopedSubscription: vi.fn(), +})) import { CumulativeUsageContextMismatchError, + CumulativeUsagePeriodClosedError, getBillingPeriodUsageCost, getBillingPeriodUsageCostByUser, getStampedPeriodRangeUsageCostByUser, type RecordCumulativeUsageParams, recordCumulativeUsage, } from '@/lib/billing/core/usage-log' +import { claimTerminalPeriod } from '@/lib/billing/cycle-close' const require = createRequire(import.meta.url) const commonJsPostgres = require('postgres') as typeof postgres @@ -121,7 +127,10 @@ describe('Cumulative billing with PostgreSQL', () => { CREATE UNIQUE INDEX usage_log_event_key_unique ON usage_log(event_key) WHERE event_key IS NOT NULL; CREATE TABLE driver_probe (id text PRIMARY KEY); - CREATE TABLE subscription (id text PRIMARY KEY, period_start timestamp, period_end timestamp) + CREATE TABLE subscription ( + id text PRIMARY KEY, period_start timestamp, period_end timestamp, + last_closed_period_start timestamp + ) `) transaction.mockImplementation(async (callback: (tx: Transaction) => Promise) => { const pause = nextPause @@ -362,12 +371,16 @@ describe('Cumulative billing with PostgreSQL', () => { ] const payer = { type: 'organization', id: 'payer' } as const + /** Moves the subscription to a window whose predecessor the cycle close has settled. */ async function setSubscriptionWindow(start: Date, end: Date) { await connection` - insert into subscription (id, period_start, period_end) - values ('sub-1', ${start.toISOString()}::timestamptz at time zone 'UTC', ${end.toISOString()}::timestamptz at time zone 'UTC') + insert into subscription (id, period_start, period_end, last_closed_period_start) + values ('sub-1', ${start.toISOString()}::timestamptz at time zone 'UTC', ${end.toISOString()}::timestamptz at time zone 'UTC', ${start.toISOString()}::timestamptz at time zone 'UTC') on conflict (id) do update - set period_start = excluded.period_start, period_end = excluded.period_end + set period_start = excluded.period_start, period_end = excluded.period_end, + last_closed_period_start = greatest( + subscription.last_closed_period_start, excluded.last_closed_period_start + ) ` } @@ -528,6 +541,83 @@ describe('Cumulative billing with PostgreSQL', () => { expect(await ledgerRows()).toEqual([{ event_key: usage(0).eventKey, cost: '0.4' }]) }) + it('refuses a charge that would roll into a period the terminal settlement already summed', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + await setSubscriptionPeriod(1) + await connection` + update subscription + set last_closed_period_start = ${periods[2].toISOString()}::timestamptz at time zone 'UTC' + ` + + await expect(charge(1)).rejects.toBeInstanceOf(CumulativeUsagePeriodClosedError) + expect(await ledgerRows()).toEqual([{ event_key: usage(0).eventKey, cost: '0.4' }]) + }) + + /** + * Resolves true once a session waits on a row lock of the subscription table, or false once + * `work` settles without anyone waiting, so a missing lock fails instead of hanging. + */ + async function waitsOnSubscriptionRow(work: Promise) { + let settled = false + work.then( + () => { + settled = true + }, + () => { + settled = true + } + ) + while (!settled) { + const [row] = await connection<{ waiting: boolean }[]>` + select exists ( + select 1 from pg_locks + where locktype = 'tuple' and relation = 'subscription'::regclass + ) as waiting + ` + if (row.waiting) return true + await sleep(10) + } + return false + } + + it('makes the terminal claim wait for an in-flight charge, so the final sum includes it', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + const pause = pauseNextTransaction() + const inFlight = charge(0.6) + let claim: Promise = Promise.resolve() + try { + await pause.reached.promise + claim = claimTerminalPeriod('sub-1') + expect(await waitsOnSubscriptionRow(claim)).toBe(true) + } finally { + pause.release.resolve() + await inFlight + await claim + } + expect(await stampedTotal(0)).toBeCloseTo(0.6, 9) + await expect(charge(0.8)).rejects.toBeInstanceOf(CumulativeUsagePeriodClosedError) + }) + + it('refuses a charge that waited on an in-flight terminal claim', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + const pause = pauseNextTransaction() + const claim = claimTerminalPeriod('sub-1') + let late: Promise = Promise.resolve() + try { + await pause.reached.promise + late = charge(0.6) + expect(await waitsOnSubscriptionRow(late)).toBe(true) + } finally { + pause.release.resolve() + await claim + } + await expect(late).rejects.toBeInstanceOf(CumulativeUsagePeriodClosedError) + expect(await stampedTotal(0)).toBeCloseTo(0.4, 9) + }) + it('holds an early period-start move until an in-flight top-up commits', async () => { const start = new Date(Date.now() - 24 * 60 * 60 * 1000) const end = new Date(Date.now() + 30 * 24 * 60 * 60 * 1000) diff --git a/apps/sim/lib/billing/core/usage-log.ts b/apps/sim/lib/billing/core/usage-log.ts index 1989f215322..204a10f4e24 100644 --- a/apps/sim/lib/billing/core/usage-log.ts +++ b/apps/sim/lib/billing/core/usage-log.ts @@ -590,7 +590,9 @@ export interface RecordCumulativeUsageParams { * arrives after that subscription has moved past the period of the request's latest row is * recorded in a new row stamped with the subscription's current period, so a request that * outlives its billing period is invoiced by the period it was spent in rather than topping up - * a period that has already been closed. Omit it for reporting-window and free payers. + * a period that has already been closed. A charge into a period the subscription's close marker + * has already passed (its terminal settlement) throws {@link CumulativeUsagePeriodClosedError}. + * Omit it for reporting-window and free payers. * * Mixed versions: code that predates period rows reads only the request key. If such code * (during a deploy, or after a rollback) handles a later callback for a request that already @@ -680,6 +682,23 @@ export class CumulativeUsageContextMismatchError extends Error { } } +/** + * A cumulative charge whose billing period the payer has already settled: the subscription ended + * and its final invoice summed that period. The charge is refused rather than recorded where no + * invoice will ever read it. + */ +export class CumulativeUsagePeriodClosedError extends Error { + constructor( + readonly eventKey: string, + readonly billingPeriod: { start: Date; end: Date } + ) { + super( + `Cumulative usage event "${eventKey}" targets a billing period that has already been settled` + ) + this.name = 'CumulativeUsagePeriodClosedError' + } +} + interface CumulativeUsageLedgerBinding { userId: string workspaceId: string | null @@ -882,14 +901,15 @@ export async function recordCumulativeUsage( return { billed: false, delta: 0, total: recorded, billingPeriod: latestPeriod } } - // The payer's current period, share-locked so a change to the subscription's period (a - // rollover, or an anchor reset inside the old period) waits for this write to commit, and - // whatever a close later sums for the old period is final. + // The payer's current period and close marker, share-locked so a change to either (a + // rollover, an anchor reset inside the old period, or a terminal settlement) waits for this + // write to commit, and whatever a close later sums for the old period is final. const [currentPeriod] = payerSubscriptionId ? await tx .select({ start: subscriptionTable.periodStart, end: subscriptionTable.periodEnd, + closedThrough: subscriptionTable.lastClosedPeriodStart, }) .from(subscriptionTable) .where(eq(subscriptionTable.id, payerSubscriptionId)) @@ -910,6 +930,16 @@ export async function recordCumulativeUsage( if (rolledPeriod && latest && chain.length >= MAX_CUMULATIVE_PERIOD_ROWS) { throw new Error(`Cumulative usage event "${eventKey}" spans too many billing periods`) } + // A marker at or past the target period's end means that period is already settled — a + // terminal settlement marks it whatever the subscription's bounds — so nothing would + // ever invoice this charge. + const targetPeriod = rolledPeriod ?? latestPeriod + if ( + currentPeriod?.closedThrough && + currentPeriod.closedThrough.getTime() >= targetPeriod.end.getTime() + ) { + throw new CumulativeUsagePeriodClosedError(eventKey, targetPeriod) + } enterStage('write') if (latest && !rolledPeriod) { @@ -929,7 +959,6 @@ export async function recordCumulativeUsage( return { billed: true, delta, total: newTotal, billingPeriod: latestPeriod } } - const targetPeriod = rolledPeriod ?? billingContext.billingPeriod const rowMetadata = periodUsageMetadata(metadata, chain) await recordUsage({ userId, diff --git a/apps/sim/lib/billing/cycle-close.ts b/apps/sim/lib/billing/cycle-close.ts index 9040e5084c1..952b3cd30a9 100644 --- a/apps/sim/lib/billing/cycle-close.ts +++ b/apps/sim/lib/billing/cycle-close.ts @@ -191,12 +191,16 @@ export async function closeElapsedPeriodBeforeDeletion(subscriptionId: string): * Claim the terminal period for a subscription that is being deleted, BEFORE * the deletion handler computes and charges final overage. Reads the * subscription row fresh (webhook payloads can be stale across a rollover) - * and advances the close marker to its current `periodStart` in one - * transaction, serializing with the sweep on the subscription row: an - * in-flight sweep close then fails its guarded marker claim and rolls back — - * including its outbox invoice — so deletion and sweep can never both bill - * the same period. Call `closeElapsedPeriodBeforeDeletion` first so a lagging - * elapsed period is settled rather than jumped. Returns the fresh period + * and, in one transaction, advances the close marker to the terminal period's end: + * the period is settled from here on, so a cost callback that commits after + * this claim is refused rather than topping up a period the final invoice has + * already summed (`recordCumulativeUsage` reads the marker under a share lock + * on the same row, so every charge either commits before this claim or sees + * the marker). This also serializes with the sweep on the subscription row: + * an in-flight sweep close then fails its guarded marker claim and rolls + * back — including its outbox invoice — so deletion and sweep can never both + * bill the same period. Call `closeElapsedPeriodBeforeDeletion` first so a + * lagging elapsed period is settled rather than jumped. Returns the period * bounds for the deletion flow to settle against, plus `markerWasCurrent`: * whether the close marker had already caught up to the terminal period. * The `billedOverageThisPeriod` tracker only ever holds collections for the @@ -241,7 +245,10 @@ export async function claimTerminalPeriod( const markerWasCurrent = !!row.lastClosedPeriodStart && row.lastClosedPeriodStart.getTime() >= row.periodStart.getTime() - if (!markerWasCurrent && options.sealLagging) { + if (!markerWasCurrent && !options.sealLagging) { + return { periodStart: row.periodStart, periodEnd: row.periodEnd, markerWasCurrent } + } + if (!markerWasCurrent) { logger.error( 'Sealing an unclosed elapsed period at terminal claim; residual overage forgiven', { @@ -250,8 +257,8 @@ export async function claimTerminalPeriod( periodStart: row.periodStart.toISOString(), } ) - await claimCloseMarker(tx, subscriptionId, row.periodStart) } + await claimCloseMarker(tx, subscriptionId, row.periodEnd ?? row.periodStart) return { periodStart: row.periodStart, periodEnd: row.periodEnd, markerWasCurrent } }) } diff --git a/apps/sim/lib/billing/webhooks/subscription.ts b/apps/sim/lib/billing/webhooks/subscription.ts index f6528ca3906..6f6badb33c9 100644 --- a/apps/sim/lib/billing/webhooks/subscription.ts +++ b/apps/sim/lib/billing/webhooks/subscription.ts @@ -294,7 +294,9 @@ export async function handleSubscriptionDeleted( // Then claim the terminal period BEFORE computing or charging: this // reads the row's fresh period (webhook payloads can be stale across - // a rollover) and serializes with the cycle-close sweep. A lagging + // a rollover), serializes with the cycle-close sweep, and marks the + // terminal period settled so a still-running request's later charge + // is refused instead of landing after the final invoice. A lagging // marker here means the close above deferred OR a rollover committed // in between — run the close once more (it settles a freshly elapsed // period; a deferred close defers again, loudly), then seal so the diff --git a/packages/db/schema.ts b/packages/db/schema.ts index 89a12e9887c..ac6a68ea216 100644 --- a/packages/db/schema.ts +++ b/packages/db/schema.ts @@ -1479,7 +1479,9 @@ export const subscription = pgTable( * closes the previous period whenever this lags the row's `periodStart`, * then advances it. Null = never initialized; the first sweep initializes * it to the current `periodStart` without billing so historical periods - * are never retroactively closed. + * are never retroactively closed. A deleted subscription's terminal + * settlement advances it to `periodEnd`: every period ending at or before + * the marker is settled, and a later charge into one is refused. */ lastClosedPeriodStart: timestamp('last_closed_period_start'), }, diff --git a/packages/testing/src/mocks/billing-usage-log.mock.ts b/packages/testing/src/mocks/billing-usage-log.mock.ts index ce906325920..9e1133828fa 100644 --- a/packages/testing/src/mocks/billing-usage-log.mock.ts +++ b/packages/testing/src/mocks/billing-usage-log.mock.ts @@ -23,6 +23,22 @@ export class MockCumulativeUsageContextMismatchError extends Error { } } +/** + * Stand-in for `CumulativeUsagePeriodClosedError` with the real `name`, constructor args, + * `eventKey`/`billingPeriod` fields, and message. + */ +export class MockCumulativeUsagePeriodClosedError extends Error { + constructor( + readonly eventKey: string, + readonly billingPeriod: { start: Date; end: Date } + ) { + super( + `Cumulative usage event "${eventKey}" targets a billing period that has already been settled` + ) + this.name = 'CumulativeUsagePeriodClosedError' + } +} + /** * Stand-in for `UnknownUsageCursorError` with the real `name`, message, and `statusCode` 400. * It is NOT a subclass of the real `HttpError`, and its `cause` is a plain `Error` carrying @@ -90,7 +106,8 @@ export const billingUsageLogMockFns = { /** * Static mock module for `@/lib/billing/core/usage-log`. Constants carry the real values; - * `CumulativeUsageContextMismatchError` is {@link MockCumulativeUsageContextMismatchError} and + * `CumulativeUsageContextMismatchError` is {@link MockCumulativeUsageContextMismatchError}, + * `CumulativeUsagePeriodClosedError` is {@link MockCumulativeUsagePeriodClosedError}, and * `UnknownUsageCursorError` is {@link MockUnknownUsageCursorError}. * * @example @@ -104,6 +121,7 @@ export const billingUsageLogMock = { CUMULATIVE_COST_EPSILON, UNKNOWN_CURSOR_MESSAGE, CumulativeUsageContextMismatchError: MockCumulativeUsageContextMismatchError, + CumulativeUsagePeriodClosedError: MockCumulativeUsagePeriodClosedError, UnknownUsageCursorError: MockUnknownUsageCursorError, isUnbilledUsageCategory: billingUsageLogMockFns.mockIsUnbilledUsageCategory, stableEventKey: billingUsageLogMockFns.mockStableEventKey, diff --git a/packages/testing/src/mocks/index.ts b/packages/testing/src/mocks/index.ts index ea63009b76c..fc4cc9a55f2 100644 --- a/packages/testing/src/mocks/index.ts +++ b/packages/testing/src/mocks/index.ts @@ -145,6 +145,7 @@ export { billingUsageLogMock, billingUsageLogMockFns, MockCumulativeUsageContextMismatchError, + MockCumulativeUsagePeriodClosedError, MockUnknownUsageCursorError, } from './billing-usage-log.mock' export {