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
13 changes: 6 additions & 7 deletions apps/sim/lib/billing/core/limit-notifications.test.ts
Original file line number Diff line number Diff line change
@@ -1,25 +1,24 @@
import { dbChainMockFns, queueTableRows, resetDbChainMock } from '@sim/testing/mocks/database.mock'
import { emailMailerMock, emailMailerMockFns } from '@sim/testing/mocks/email-mailer.mock'
import { emailTemplatesMock, emailTemplatesMockFns } from '@sim/testing/mocks/email-templates.mock'
import {
emailUnsubscribeMock,
emailUnsubscribeMockFns,
} from '@sim/testing/mocks/email-unsubscribe.mock'
import { resetEnvFlagsMock, setEnvFlags } from '@sim/testing/mocks/env-flags.mock'
import { schemaMock } from '@sim/testing/mocks/schema.mock'
import { resetUrlsMock, urlsMockFns } from '@sim/testing/mocks/urls.mock'
import { workspaceAuthzMock } from '@sim/testing/mocks/workspace-authz.mock'
import { afterAll, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest'

const { getEmailPreferencesMock } = vi.hoisted(() => ({
getEmailPreferencesMock: vi.fn(() => Promise.resolve(null as unknown)),
}))

vi.mock('@/lib/messaging/email/mailer', () => emailMailerMock)
vi.mock('@/lib/messaging/email/unsubscribe', () => ({
getEmailPreferences: getEmailPreferencesMock,
}))
vi.mock('@/lib/messaging/email/unsubscribe', () => emailUnsubscribeMock)
vi.mock('@/components/emails', () => emailTemplatesMock)
vi.mock('@sim/platform-authz/workspace', () => workspaceAuthzMock)

import { maybeSendLimitThresholdEmail } from '@/lib/billing/core/limit-notifications'

const getEmailPreferencesMock = emailUnsubscribeMockFns.mockGetEmailPreferences
const sendEmailSpy = emailMailerMockFns.mockSendEmail
const renderMock = emailTemplatesMockFns.mockRenderLimitThresholdEmail
const subjectMock = emailTemplatesMockFns.mockGetLimitEmailSubject
Expand Down
103 changes: 89 additions & 14 deletions apps/sim/lib/billing/core/limit-notifications.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ import { db } from '@sim/db'
import { member, organization, settings, user, userStats } from '@sim/db/schema'
import { createLogger } from '@sim/logger'
import { isOrgAdminRole } from '@sim/platform-authz/workspace'
import { and, eq, sql } from 'drizzle-orm'
import { and, eq, type SQL, sql } from 'drizzle-orm'
import type { HighestPrioritySubscription } from '@/lib/billing/core/plan'
import { getHighestPrioritySubscription } from '@/lib/billing/core/subscription'
import type { BillingEntity } from '@/lib/billing/core/usage-log'
Expand Down Expand Up @@ -32,6 +32,32 @@ function thresholdFor(percent: number): 0 | 80 | 100 {
return 0
}

/**
* Replace the account's `limitNotifications` with `next` when `condition` holds, returning
* whether the row was updated.
*/
async function writeLimitNotifications(
scope: 'user' | 'organization',
id: string,
next: SQL,
condition: SQL
): Promise<boolean> {
const written =
scope === 'user'
? await db
.update(userStats)
.set({ limitNotifications: next })
.where(and(eq(userStats.userId, id), condition))
.returning({ id: userStats.userId })
: await db
.update(organization)
.set({ limitNotifications: next })
.where(and(eq(organization.id, id), condition))
.returning({ id: organization.id })

return written.length > 0
}

/**
* Atomically claim a threshold for a category: advance the stored value to
* `threshold` only if it is currently lower, returning whether THIS call won the
Expand All @@ -50,20 +76,69 @@ async function claimThreshold(
? sql`coalesce((${userStats.limitNotifications} ->> ${category})::int, 0) < ${threshold}`
: sql`coalesce((${organization.limitNotifications} ->> ${category})::int, 0) < ${threshold}`

const claimed =
scope === 'user'
return writeLimitNotifications(scope, id, setExpr, onlyIfLower)
}

const DAY_MS = 24 * 60 * 60 * 1000

/** One account's credits threshold, keyed on its billing period and limit. */
export interface CreditsThresholdClaim {
scope: 'user' | 'organization'
id: string
periodStart: Date
limit: number
threshold: 80 | 100
}

/**
* The stored claims for a credits threshold, and the condition under which it is still unclaimed:
* `credits` holds the highest threshold emailed while `creditsPeriod` (the period's start day) and
* `creditsLimit` (the limit in cents) still match, so a new period or a changed limit — in either
* direction — re-arms both thresholds with no reset write, while within one a claim of 100 also
* retires 80, and never the reverse.
*/
function creditsThresholdSql(claim: CreditsThresholdClaim) {
const periodDay = Math.floor(claim.periodStart.getTime() / DAY_MS)
const limitCents = Math.round(claim.limit * 100)
const column =
claim.scope === 'user' ? userStats.limitNotifications : organization.limitNotifications
return {
next: sql`coalesce(${column}, '{}'::jsonb) || jsonb_build_object('credits', ${claim.threshold}::int, 'creditsPeriod', ${periodDay}::bigint, 'creditsLimit', ${limitCents}::bigint)`,
unclaimed: sql<boolean>`not (
(${column} ->> 'creditsPeriod')::bigint is not distinct from ${periodDay}::bigint
and (${column} ->> 'creditsLimit')::bigint is not distinct from ${limitCents}::bigint
and coalesce((${column} ->> 'credits')::int, 0) >= ${claim.threshold}::int
)`,
}
}

/**
* Whether a credits threshold is still unclaimed, read with one indexed lookup of the account
* row. A cheap pre-check only: callers run it on every completion above a threshold, so once
* the threshold is claimed they stop there instead of resolving recipients. The atomic
* {@link claimCreditsThreshold} remains the real dedup.
*/
export async function isCreditsThresholdUnclaimed(claim: CreditsThresholdClaim): Promise<boolean> {
const { unclaimed } = creditsThresholdSql(claim)
const [row] =
claim.scope === 'user'
? await db
.update(userStats)
.set({ limitNotifications: setExpr })
.where(and(eq(userStats.userId, id), onlyIfLower))
.returning({ id: userStats.userId })
.select({ unclaimed })
.from(userStats)
.where(eq(userStats.userId, claim.id))
.limit(1)
: await db
.update(organization)
.set({ limitNotifications: setExpr })
.where(and(eq(organization.id, id), onlyIfLower))
.returning({ id: organization.id })
.select({ unclaimed })
.from(organization)
.where(eq(organization.id, claim.id))
.limit(1)
return row?.unclaimed === true
}

return claimed.length > 0
/** Claim a credits threshold, returning whether THIS call won it. */
export function claimCreditsThreshold(claim: CreditsThresholdClaim): Promise<boolean> {
const { next, unclaimed } = creditsThresholdSql(claim)
return writeLimitNotifications(claim.scope, claim.id, next, unclaimed)
}

/** Re-arm a category (reset its stored threshold to 0) once usage falls back into the low band. */
Expand Down Expand Up @@ -108,7 +183,7 @@ async function isUnsubscribed(email: string): Promise<boolean> {
* Returning an empty list means "nobody to notify" — the caller then skips the
* claim so the dedup state isn't burned without an email going out.
*/
async function resolveRecipients(
export async function resolveLimitEmailRecipients(
scope: 'user' | 'organization',
params: { userId?: string; userEmail?: string; userName?: string; organizationId?: string }
): Promise<LimitEmailRecipient[]> {
Expand Down Expand Up @@ -207,7 +282,7 @@ export async function maybeSendLimitThresholdEmail(params: {

if (params.rearmOnly || desired === 0) return

const recipients = await resolveRecipients(scope, params)
const recipients = await resolveLimitEmailRecipients(scope, params)
if (recipients.length === 0) return

if (!(await claimThreshold(scope, stateId, category, desired))) return
Expand Down
164 changes: 164 additions & 0 deletions apps/sim/lib/billing/core/reporting-usage-cache.integration.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,164 @@
/**
* The shared reporting-usage read against a real ledger in a disposable PostgreSQL schema and a
* real Redis. Skipped without `TEST_REDIS_URL`. Each test uses a fresh payer, so the in-process
* cache is always cold and every read models a new process.
*/

import { type AddressInfo, createServer, type Socket } from 'node:net'
import type { db } from '@sim/db'
import * as schema from '@sim/db/schema'
import { readTestDatabaseUrl, readTestRedisUrl } from '@sim/db/testing/test-infrastructure'
import { redisConfigMock, redisConfigMockFns } from '@sim/testing/mocks/redis-config.mock'
import { generateId } from '@sim/utils/id'
import { drizzle } from 'drizzle-orm/postgres-js'
import Redis from 'ioredis'
import postgres from 'postgres'
import { afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest'

const { transaction } = vi.hoisted(() => ({ transaction: vi.fn() }))
const databaseUrl = readTestDatabaseUrl()
const redisUrl = readTestRedisUrl()

vi.mock('@sim/db', () => ({ db: { transaction }, dbReplica: {} }))
vi.mock('@/lib/core/config/redis', () => redisConfigMock)

import { readSoftGateUsageCost } from '@/lib/billing/core/reporting-usage-cache'
import type { BillingEntity, UsageQueryPeriod } from '@/lib/billing/core/usage-log'

const schemaName = `reporting_usage_${generateId().replaceAll('-', '')}`
const connection = postgres(databaseUrl, {
max: 2,
prepare: false,
connection: { search_path: schemaName },
onnotice: () => undefined,
})
const database = drizzle(connection, { schema }) as typeof db

const REPORTING: UsageQueryPeriod = {
start: new Date('2026-01-01T00:00:00.000Z'),
end: new Date('2027-01-01T00:00:00.000Z'),
source: 'reporting',
}

function sharedKey(payer: BillingEntity): string {
return `usage:reporting:v1:${payer.type}:${payer.id}:reporting:${REPORTING.start.toISOString()}:${REPORTING.end.toISOString()}`
}

describe.runIf(Boolean(redisUrl))('shared reporting usage read', () => {
let redis: Redis
let payer: BillingEntity

beforeAll(async () => {
redis = new Redis(redisUrl!, { lazyConnect: true, maxRetriesPerRequest: 0 })
await redis.connect()
await connection.unsafe(`CREATE SCHEMA "${schemaName}"`)
await connection.unsafe(`CREATE TABLE usage_log (
id text PRIMARY KEY, cost numeric NOT NULL, billing_entity_type text,
billing_entity_id text, billing_period_start timestamp, billing_period_end timestamp,
created_at timestamp NOT NULL
)`)
transaction.mockImplementation((callback) => database.transaction(callback))
})

beforeEach(async () => {
transaction.mockClear()
redisConfigMockFns.mockGetRedisClient.mockReturnValue(redis)
payer = { type: 'organization', id: generateId() }
await connection`INSERT INTO usage_log (id, cost, billing_entity_type, billing_entity_id, created_at)
VALUES (${generateId()}, 4.25, 'organization', ${payer.id}, '2026-03-01'),
(${generateId()}, 1.5, 'organization', ${payer.id}, '2026-06-01'),
(${generateId()}, 99, 'organization', ${payer.id}, '2025-12-31')`
})

afterEach(async () => {
await redis.del(sharedKey(payer))
})

afterAll(async () => {
await redis?.quit()
await connection.unsafe(`DROP SCHEMA IF EXISTS "${schemaName}" CASCADE`)
await connection.end()
})

it('serves a sum another process stored instead of the ledger', async () => {
await redis.set(sharedKey(payer), '12.5', 'PX', 30_000)

await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(12.5)
})

it('sums the ledger exactly on a miss and stores the sum for other processes', async () => {
await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)

await vi.waitFor(async () => expect(await redis.get(sharedKey(payer))).toBe('5.75'))
const ttl = await redis.pttl(sharedKey(payer))
expect(ttl).toBeGreaterThan(25_000)
expect(ttl).toBeLessThanOrEqual(35_000)
})

it('treats an unreadable stored sum as a miss', async () => {
await redis.set(sharedKey(payer), 'not-a-number', 'PX', 30_000)

await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)
})

it('never lets a slower, older sum replace one stored while it ran', async () => {
transaction.mockImplementationOnce(async (callback) => {
const result = await database.transaction(callback)
await redis.set(sharedKey(payer), '9.99', 'PX', 30_000)
return result
})

await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)
expect(await redis.get(sharedKey(payer))).toBe('9.99')
})

it('sums the ledger when the Redis client cannot be built', async () => {
redisConfigMockFns.mockGetRedisClient.mockImplementation(() => {
throw new Error('Invalid Redis configuration')
})

await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)
})

it('sums the ledger without issuing or queuing a command while Redis is not ready', async () => {
const notReady = new Redis(redisUrl!, { lazyConnect: true, maxRetriesPerRequest: 0 })
redisConfigMockFns.mockGetRedisClient.mockReturnValue(notReady)
try {
await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)
expect(notReady.status).toBe('wait')

await notReady.connect()
await notReady.ping()
expect(await redis.get(sharedKey(payer))).toBeNull()
} finally {
notReady.disconnect()
}
})

it('sums the ledger promptly when a connected Redis stops answering', async () => {
const sockets = new Set<Socket>()
/** Completes the client's handshake, then never answers a read. */
const silent = createServer((socket) => {
sockets.add(socket)
socket.on('data', (data) => {
const text = data.toString()
if (/\bGET\b/i.test(text)) return
socket.write('+OK\r\n'.repeat(text.match(/^\*\d+\r\n/gm)?.length ?? 0))
})
})
await new Promise<void>((resolve) => silent.listen(0, '127.0.0.1', resolve))
const { port } = silent.address() as AddressInfo
const hung = new Redis({ host: '127.0.0.1', port, lazyConnect: true, enableReadyCheck: false })
redisConfigMockFns.mockGetRedisClient.mockReturnValue(hung)
try {
await hung.connect()
const startedAt = Date.now()
await expect(readSoftGateUsageCost(payer, REPORTING)).resolves.toBe(5.75)
expect(Date.now() - startedAt).toBeLessThan(2_000)
} finally {
hung.disconnect()
for (const socket of sockets) socket.destroy()
silent.close()
}
})
})
Loading
Loading