Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
136 changes: 87 additions & 49 deletions apps/sim/app/api/billing/update-cost/route.integration.ts
Original file line number Diff line number Diff line change
@@ -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'

Expand All @@ -25,31 +26,86 @@ 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<Payer> {
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: {
'content-type': 'application/json',
'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',
Expand All @@ -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)
Expand All @@ -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), {})
Comment thread
waleedlatif1 marked this conversation as resolved.

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])
})
})
4 changes: 3 additions & 1 deletion apps/sim/app/api/billing/update-cost/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ import {
import {
type CumulativeUsageContextField,
CumulativeUsageContextMismatchError,
CumulativeUsagePeriodClosedError,
recordCumulativeUsage,
} from '@/lib/billing/core/usage-log'
import {
Expand Down Expand Up @@ -512,7 +513,8 @@ async function updateCostInner(req: NextRequest, span: Span): Promise<NextRespon
const pgCode = getPostgresErrorCode(error)
const pgConstraint = getPostgresConstraintName(error)
const reconciliationOutcome =
error instanceof ThresholdSettlementError && !error.retryable
(error instanceof ThresholdSettlementError && !error.retryable) ||
error instanceof CumulativeUsagePeriodClosedError
? BILLING_CALLBACK_OUTCOME.billingPeriodElapsed
: pgCode === '23503' && pgConstraint === 'usage_log_user_id_user_id_fk'
? BILLING_CALLBACK_OUTCOME.billingUserNotFound
Expand Down
100 changes: 95 additions & 5 deletions apps/sim/lib/billing/core/usage-log.integration.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import type { db } from '@sim/db'
import * as schema from '@sim/db/schema'
import { readTestDatabaseUrl } from '@sim/db/testing/test-infrastructure'
import { getPostgresErrorCode } from '@sim/utils/errors'
import { sleep } from '@sim/utils/helpers'
import { generateId } from '@sim/utils/id'
import { eq, sql } from 'drizzle-orm'
import { drizzle } from 'drizzle-orm/postgres-js'
Expand All @@ -20,16 +21,21 @@ const databaseUrl = readTestDatabaseUrl()

vi.mock('@sim/db', () => ({ 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<typeof import('@/lib/billing/subscriptions/utils')>()),
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
Expand Down Expand Up @@ -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<unknown>) => {
const pause = nextPause
Expand Down Expand Up @@ -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
)
`
}

Expand Down Expand Up @@ -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<unknown>) {
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<unknown> = 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<unknown> = 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)
Expand Down
Loading
Loading