diff --git a/apps/sim/app/api/billing/update-cost/route.integration.ts b/apps/sim/app/api/billing/update-cost/route.integration.ts new file mode 100644 index 00000000000..082009c92dc --- /dev/null +++ b/apps/sim/app/api/billing/update-cost/route.integration.ts @@ -0,0 +1,121 @@ +/** + * 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. + */ +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 { NextRequest } from 'next/server' +import { afterAll, describe, expect, it, vi } from 'vitest' + +vi.mock('@/lib/core/config/env-flags', () => ({ + ...envFlagsMock, + isHosted: true, + isBillingEnabled: true, +})) +vi.mock('@/lib/mothership/request/http', async (importOriginal) => ({ + ...(await importOriginal()), + checkInternalApiKey: () => ({ success: true }), +})) + +import { + BILLING_ACCOUNT_DECISION_HEADER, + serializeAccountBillingDecisionHeader, +} from '@/lib/billing/core/billing-attribution' +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() + +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)) +}) + +function callback(requestKey: string, cost: number, decision: string): 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, + }, + body: JSON.stringify({ + userId, + cost, + model: 'test-model', + source: 'copilot', + idempotencyKey: requestKey, + }), + }) +} + +describe('direct-v1 cost callbacks in PostgreSQL', () => { + it('records spend after a Stripe rollover in the payer current period', async () => { + 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 requestKey = generateId() + + expect((await POST(callback(requestKey, 0.5, decision), {})).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) + + const rows = await db + .select({ + eventKey: usageLog.eventKey, + cost: usageLog.cost, + billingPeriodStart: usageLog.billingPeriodStart, + }) + .from(usageLog) + .where(eq(usageLog.userId, userId)) + 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) + expect(byKey.get(`update-cost:${requestKey}`)?.billingPeriodStart?.getTime()).toBe( + admitted.start.getTime() + ) + expect(Number(byKey.get(`update-cost:${requestKey}@1`)?.cost)).toBeCloseTo(0.3) + expect(byKey.get(`update-cost:${requestKey}@1`)?.billingPeriodStart?.getTime()).toBe( + rolled.start.getTime() + ) + }) +}) diff --git a/apps/sim/app/api/billing/update-cost/route.test.ts b/apps/sim/app/api/billing/update-cost/route.test.ts index dd89be51224..60407452e78 100644 --- a/apps/sim/app/api/billing/update-cost/route.test.ts +++ b/apps/sim/app/api/billing/update-cost/route.test.ts @@ -4,12 +4,19 @@ import { billingAttributionMock, billingAttributionMockFns, } from '@sim/testing/mocks/billing-attribution.mock' +import { billingCoreMock, billingCoreMockFns } from '@sim/testing/mocks/billing-core.mock' +import { billingPlanMock, billingPlanMockFns } from '@sim/testing/mocks/billing-plan.mock' import { billingUsageLogMock, billingUsageLogMockFns, } from '@sim/testing/mocks/billing-usage-log.mock' +import { + billingUsageMonitorMock, + billingUsageMonitorMockFns, +} from '@sim/testing/mocks/billing-usage-monitor.mock' import { copilotHttpMock, copilotHttpMockFns } from '@sim/testing/mocks/copilot-http.mock' import { mothershipOtelMock } from '@sim/testing/mocks/mothership-otel.mock' +import { sleep } from '@sim/utils/helpers' import { afterAll, beforeEach, describe, expect, it, vi } from 'vitest' const { @@ -41,12 +48,20 @@ vi.mock('@/lib/billing/core/usage-log', () => billingUsageLogMock) vi.mock('@/lib/billing/core/billing-attribution', () => billingAttributionMock) +vi.mock('@/lib/billing/core/billing', () => billingCoreMock) + +vi.mock('@/lib/billing/core/plan', () => billingPlanMock) + +vi.mock('@/lib/billing/calculations/usage-monitor', () => billingUsageMonitorMock) + vi.mock('@/lib/billing/threshold-billing', () => ({ checkAndBillOverageThreshold: mockCheckAndBillOverageThreshold, checkAndBillPayerOverageThreshold: mockCheckAndBillPayerOverageThreshold, ThresholdSettlementError: MockThresholdSettlementError, })) +import { resetMidRunUsageCaches } from '@/lib/billing/core/mid-run-usage' +import { resetUsageGateCache } from '@/lib/billing/core/usage-gate-cache' import { BillingCallbackBody, BillingCallbackHeaders, @@ -69,6 +84,8 @@ const mockRequireBillingAttributionHeader = const mockResolveLegacyV0BillingAttribution = billingAttributionMockFns.mockResolveLegacyV0BillingAttribution const mockToBillingContext = billingAttributionMockFns.mockToBillingContext +const mockCheckAttributedUsageLimits = billingAttributionMockFns.mockCheckAttributedUsageLimits +const mockRefreshAttributionPeriod = billingAttributionMockFns.mockRefreshAttributionPeriod afterAll(resetEnvFlagsMock) @@ -236,6 +253,18 @@ describe('POST /api/billing/update-cost — workspaceId attribution', () => { expect(mockRecordCumulativeUsage).not.toHaveBeenCalled() }) + it('rejects an idempotency key that could collide with a period row key', async () => { + const res = await POST( + createMockRequest( + 'POST', + { ...SELF_HOSTED_UPDATE_COST_BODY, idempotencyKey: 'old-go-key@1' }, + { 'x-api-key': 'internal' } + ) + ) + + expect(res.status).toBe(400) + }) + it('rejects billing-enabled callbacks without a stable idempotency key', async () => { const res = await POST( createMockRequest('POST', KEYLESS_UPDATE_COST_BODY, { 'x-api-key': 'internal' }) @@ -847,3 +876,430 @@ describe('POST /api/billing/update-cost — workspaceId attribution', () => { expect(mockRecordCumulativeUsage).not.toHaveBeenCalled() }) }) + +describe('POST /api/billing/update-cost — mid-run usage gate', () => { + let callbackSequence = 0 + /** A Stripe-period payer admitted in a period that has since ended. */ + const STRIPE_ATTRIBUTION = { + ...ATTRIBUTION, + billingPeriod: { ...ATTRIBUTION.billingPeriod, source: 'stripe' as const }, + } + const CURRENT_ATTRIBUTION = { + ...ATTRIBUTION, + billingPeriod: { + start: '2026-07-01T00:00:00.000Z', + end: '2099-01-01T00:00:00.000Z', + source: 'stripe' as const, + }, + } + + function attributedCallback() { + callbackSequence += 1 + const billingRequestId = `0190c03f-9f7d-4b79-8b58-${String(callbackSequence).padStart(12, '0')}` + return createMockRequest( + 'POST', + { + userId: 'user-1', + cost: 0.5 * callbackSequence, + model: 'claude-opus-4.8', + source: 'workspace-chat', + workspaceId: 'ws-1', + idempotencyKey: billingRequestId, + }, + { + 'x-api-key': 'internal', + 'x-sim-billing-protocol': 'attribution-v1', + 'x-sim-billing-request-id': billingRequestId, + 'x-sim-billing-attribution': 'serialized-attribution', + } + ) + } + + beforeEach(() => { + resetUsageGateCache() + resetMidRunUsageCaches() + setEnvFlags({ isBillingEnabled: true, isHosted: true }) + mockCheckInternalApiKey.mockReturnValue({ success: true }) + mockRecordCumulativeUsage.mockResolvedValue({ billed: true, delta: 0.5, total: 0.5 }) + mockCheckAndBillPayerOverageThreshold.mockResolvedValue(undefined) + mockRequireBillingAttributionHeader.mockReturnValue(CURRENT_ATTRIBUTION) + mockRefreshAttributionPeriod.mockImplementation(async (attribution: unknown) => attribution) + mockToBillingContext.mockReturnValue({ + billingEntity: { type: 'organization', id: 'org-1' }, + billingPeriod: { + start: new Date('2026-07-01T00:00:00.000Z'), + end: new Date('2026-08-01T00:00:00.000Z'), + }, + }) + }) + + it('tells the worker when the run payer has crossed its usage limit', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const res = await POST(attributedCallback()) + + expect(res.status).toBe(200) + await expect(res.json()).resolves.toMatchObject({ + success: true, + usageExceeded: true, + usageUpgrade: { + reason: 'usage_limit', + action: 'upgrade_plan', + message: expect.stringContaining('usage limit'), + }, + }) + }) + + it('offers a paid organization payer the increase-limit card', async () => { + mockRequireBillingAttributionHeader.mockReturnValue({ + ...CURRENT_ATTRIBUTION, + payerSubscription: { id: 'sub-1', plan: 'team', status: 'active', seats: 4 }, + }) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const res = await POST(attributedCallback()) + + const body = await res.json() + expect(body.usageUpgrade).toMatchObject({ + action: 'increase_limit', + message: expect.stringContaining('organization'), + }) + }) + + it('serves a cached admission to every step and re-reads a refusal', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: false }) + + expect((await (await POST(attributedCallback())).json()).usageExceeded).toBe(false) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + for (let step = 0; step < 4; step++) { + const body = await (await POST(attributedCallback())).json() + expect(body.usageExceeded).toBe(false) + expect(body).not.toHaveProperty('usageUpgrade') + } + + resetUsageGateCache() + expect((await (await POST(attributedCallback())).json()).usageExceeded).toBe(true) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: false }) + expect((await (await POST(attributedCallback())).json()).usageExceeded).toBe(false) + }) + + it('answers a duplicate retry with the verdict its lost first answer carried', async () => { + mockRecordCumulativeUsage.mockResolvedValue({ billed: false, delta: 0, total: 0.5 }) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const res = await POST(attributedCallback()) + + expect(res.status).toBe(409) + await expect(res.json()).resolves.toMatchObject({ + code: 'DUPLICATE_BILLING_EVENT', + usageExceeded: true, + usageUpgrade: { action: 'upgrade_plan' }, + }) + }) + + describe('a direct-v1 run', () => { + function directCallback() { + callbackSequence += 1 + const billingRequestId = `0190c03f-9f7d-4b79-8b58-${String(callbackSequence).padStart(12, '0')}` + return createMockRequest( + 'POST', + { + userId: 'user-1', + cost: 0.5 * callbackSequence, + model: 'claude-opus-4.8', + source: 'workspace-chat', + idempotencyKey: billingRequestId, + }, + { + 'x-api-key': 'internal', + 'x-sim-billing-protocol': 'direct-v1', + 'x-sim-billing-request-id': billingRequestId, + 'x-sim-billing-account-decision': 'serialized-account-decision', + } + ) + } + + beforeEach(() => { + mockRequireAccountBillingDecisionHeader.mockReturnValue(ACCOUNT_BILLING_DECISION) + billingCoreMockFns.mockGetOrganizationSubscription.mockResolvedValue(null) + billingPlanMockFns.mockGetHighestPrioritySubscription.mockResolvedValue(null) + billingAttributionMockFns.mockCheckAccountBillingBlocks.mockResolvedValue({ blocked: false }) + billingUsageMonitorMockFns.mockCheckUsageStatus.mockResolvedValue({ + isExceeded: true, + currentUsage: 12, + limit: 10, + }) + }) + + it('tells the worker when its admitted payer has crossed its usage limit', async () => { + const body = await (await POST(directCallback())).json() + + expect(body).toMatchObject({ + usageExceeded: true, + usageUpgrade: { reason: 'usage_limit' }, + }) + }) + + it('never pauses a blocked payer with the usage card', async () => { + billingAttributionMockFns.mockCheckAccountBillingBlocks.mockResolvedValue({ + blocked: true, + scope: 'payer', + }) + + const body = await (await POST(directCallback())).json() + + expect(body.usageExceeded).toBe(false) + }) + }) + + describe('a run that outlives its billing period', () => { + const PAYER_SUBSCRIPTION = { + id: 'sub-1', + plan: 'team', + status: 'active', + seats: 4, + } + const ADMITTED_PERIOD = { + start: new Date('2026-07-01T00:00:00.000Z'), + end: new Date('2026-08-01T00:00:00.000Z'), + } + const CURRENT_PERIOD = { + start: new Date('2026-08-01T00:00:00.000Z'), + end: new Date('2026-09-01T00:00:00.000Z'), + } + + function admittedWithSource(source: 'stripe' | 'reporting' | 'default') { + mockToBillingContext.mockReturnValue({ + billingEntity: { type: 'organization', id: 'org-1' }, + billingPeriod: { ...ADMITTED_PERIOD, source }, + }) + } + + /** Threshold settlement for a payer whose charges belong to `period` refuses any other. */ + function settlesOnlyAgainst(period: typeof ADMITTED_PERIOD) { + mockCheckAndBillPayerOverageThreshold.mockImplementation( + async (_payer: unknown, options: { expectedBillingPeriod: typeof ADMITTED_PERIOD }) => { + if (options.expectedBillingPeriod.start.getTime() !== period.start.getTime()) { + throw new Error('Settled against a period the charge did not land in') + } + } + ) + } + + beforeEach(() => { + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: false }) + mockRequireBillingAttributionHeader.mockReturnValue({ + ...CURRENT_ATTRIBUTION, + payerSubscription: PAYER_SUBSCRIPTION, + }) + // As the ledger behaves: a charge given the payer's subscription lands in its current + // period, any other stays in the period it was admitted in. + mockRecordCumulativeUsage.mockImplementation( + async (params: { + payerSubscriptionId?: string + billingPeriod: typeof ADMITTED_PERIOD + }) => ({ + billed: true, + delta: 0.5, + total: 1.5, + billingPeriod: params.payerSubscriptionId + ? CURRENT_PERIOD + : { start: params.billingPeriod.start, end: params.billingPeriod.end }, + }) + ) + }) + + it("records a Stripe payer's charge in its current period and settles it there", async () => { + admittedWithSource('stripe') + settlesOnlyAgainst(CURRENT_PERIOD) + + expect((await POST(attributedCallback())).status).toBe(200) + }) + + it('leaves a period that closed under a recorded charge to the cycle close', async () => { + admittedWithSource('stripe') + mockRecordCumulativeUsage.mockResolvedValue({ + billed: true, + delta: 0.5, + total: 1.5, + billingPeriod: ADMITTED_PERIOD, + }) + mockCheckAndBillPayerOverageThreshold.mockRejectedValue( + new MockThresholdSettlementError('billing_period_elapsed') + ) + + const res = await POST(attributedCallback()) + + expect(res.status).toBe(200) + }) + + it.each(['reporting', 'default'] as const)( + 'keeps a payer with a %s period on the period it was admitted in', + async (source) => { + admittedWithSource(source) + settlesOnlyAgainst(ADMITTED_PERIOD) + + expect((await POST(attributedCallback())).status).toBe(200) + } + ) + }) + + it.each([ + ['an unreadable ledger', { isExceeded: true, reason: 'usage_unavailable' }], + ['a blocked account', { isExceeded: true, reason: 'billing_blocked', scope: 'payer' }], + ])('does not pause a run for %s', async (_case, verdict) => { + mockCheckAttributedUsageLimits.mockResolvedValue(verdict) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(false) + expect(body).not.toHaveProperty('usageUpgrade') + }) + + it('tells a member over the cap their organization set who can raise it', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'member' }) + + const body = await (await POST(attributedCallback())).json() + + expect(body).toMatchObject({ + usageUpgrade: { message: expect.stringMatching(/limit your organization set for you/) }, + }) + }) + + /** The gate refuses only when it judges the payer's current period. */ + function refuseOnlyCurrentPeriod() { + mockCheckAttributedUsageLimits.mockImplementation( + async (attribution: typeof CURRENT_ATTRIBUTION) => ({ + isExceeded: attribution.billingPeriod.end === CURRENT_ATTRIBUTION.billingPeriod.end, + scope: 'payer', + }) + ) + } + + it('judges a run past its admitted period against the payer current period', async () => { + mockRequireBillingAttributionHeader.mockReturnValue(STRIPE_ATTRIBUTION) + mockRefreshAttributionPeriod.mockResolvedValue(CURRENT_ATTRIBUTION) + refuseOnlyCurrentPeriod() + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(true) + }) + + it('judges a run whose payer period moved early against the moved period', async () => { + const moved = { + ...CURRENT_ATTRIBUTION, + billingPeriod: { start: '2026-07-15T00:00:00.000Z', end: '2099-02-01T00:00:00.000Z' }, + } + mockRefreshAttributionPeriod.mockResolvedValue(moved) + mockCheckAttributedUsageLimits.mockImplementation(async (attribution: typeof moved) => ({ + isExceeded: attribution.billingPeriod.start === moved.billingPeriod.start, + scope: 'payer', + })) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(true) + }) + + it('judges a reporting-window run against its admitted window after that window ends', async () => { + const admitted = { + ...ATTRIBUTION, + billingPeriod: { ...ATTRIBUTION.billingPeriod, source: 'reporting' as const }, + } + mockRequireBillingAttributionHeader.mockReturnValue(admitted) + mockRefreshAttributionPeriod.mockResolvedValue({ + ...CURRENT_ATTRIBUTION, + billingPeriod: { ...CURRENT_ATTRIBUTION.billingPeriod, source: 'reporting' as const }, + }) + mockCheckAttributedUsageLimits.mockImplementation(async (attribution: typeof admitted) => ({ + isExceeded: attribution.billingPeriod.end === admitted.billingPeriod.end, + scope: 'payer', + })) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(true) + }) + + it('keeps a run going when its current period cannot be read', async () => { + mockRequireBillingAttributionHeader.mockReturnValue(STRIPE_ATTRIBUTION) + mockRefreshAttributionPeriod.mockRejectedValue(new Error('subscription read timed out')) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(false) + }) + + it('answers not exceeded when the standing read outlasts the callback budget', async () => { + mockCheckAttributedUsageLimits.mockImplementation(async () => { + await sleep(1500) + return { isExceeded: true, scope: 'payer' } + }) + const startedAt = Date.now() + + const res = await POST(attributedCallback()) + + expect(res.status).toBe(200) + await expect(res.json()).resolves.toMatchObject({ success: true, usageExceeded: false }) + expect(Date.now() - startedAt).toBeLessThan(1400) + }) + + it('does not pause a run on a verdict read across the end of its period', async () => { + const straddling = { + ...CURRENT_ATTRIBUTION, + billingPeriod: { + start: '2026-07-01T00:00:00.000Z', + end: new Date(Date.now() + 40).toISOString(), + source: 'stripe' as const, + }, + } + mockRefreshAttributionPeriod.mockResolvedValue(straddling) + mockCheckAttributedUsageLimits.mockImplementation(async () => { + await sleep(80) + return { isExceeded: true, scope: 'payer' } + }) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(false) + }) + + it('reloads a cached current period once it has ended', async () => { + const ending = { + ...CURRENT_ATTRIBUTION, + billingPeriod: { + start: '2026-07-01T00:00:00.000Z', + end: new Date(Date.now() + 50).toISOString(), + }, + } + mockRefreshAttributionPeriod + .mockResolvedValueOnce(ending) + .mockResolvedValue(CURRENT_ATTRIBUTION) + refuseOnlyCurrentPeriod() + await POST(attributedCallback()) + await sleep(100) + + const body = await (await POST(attributedCallback())).json() + + expect(body.usageExceeded).toBe(true) + }) + + it('keeps a recorded charge successful when the gate read fails', async () => { + mockCheckAttributedUsageLimits.mockRejectedValue(new Error('ledger read timed out')) + + const res = await POST(attributedCallback()) + + expect(res.status).toBe(200) + await expect(res.json()).resolves.toMatchObject({ success: true, usageExceeded: false }) + }) + + it('reports no exceeded usage when billing is disabled', async () => { + setEnvFlags({ isBillingEnabled: false, isHosted: true }) + + const res = await POST(attributedCallback()) + + await expect(res.json()).resolves.toMatchObject({ usageExceeded: false }) + }) +}) diff --git a/apps/sim/app/api/billing/update-cost/route.ts b/apps/sim/app/api/billing/update-cost/route.ts index 0f789f87c99..8f4be83be39 100644 --- a/apps/sim/app/api/billing/update-cost/route.ts +++ b/apps/sim/app/api/billing/update-cost/route.ts @@ -2,7 +2,11 @@ import type { Span } from '@opentelemetry/api' import { createLogger } from '@sim/logger' import { getPostgresConstraintName, getPostgresErrorCode, toError } from '@sim/utils/errors' import { type NextRequest, NextResponse } from 'next/server' -import { billingUpdateCostContract } from '@/lib/api/contracts/subscription' +import { + type BillingUpdateCostResponse, + type BillingUsageVerdict, + billingUpdateCostContract, +} from '@/lib/api/contracts/subscription' import { parseRequest } from '@/lib/api/server' import { type AccountBillingDecision, @@ -18,6 +22,11 @@ import { resolveLegacyV0BillingAttribution, toBillingContext, } from '@/lib/billing/core/billing-attribution' +import { + type MidRunUsageVerdict, + readMidRunAccountUsageVerdict, + readMidRunUsageVerdict, +} from '@/lib/billing/core/mid-run-usage' import { type CumulativeUsageContextField, CumulativeUsageContextMismatchError, @@ -28,7 +37,9 @@ import { checkAndBillPayerOverageThreshold, ThresholdSettlementError, } from '@/lib/billing/threshold-billing' +import { resolveUsageUpgradePayload } from '@/lib/billing/usage-upgrade' import { isBillingEnabled, isHosted } from '@/lib/core/config/env-flags' +import { withinDeadline } from '@/lib/core/utils/deadline' import { generateRequestId } from '@/lib/core/utils/request' import { withRouteHandler } from '@/lib/core/utils/with-route-handler' import { BILLING_CALLBACK_OUTCOME } from '@/lib/mothership/generated/billing-protocol-v1' @@ -39,6 +50,14 @@ import { checkInternalApiKey } from '@/lib/mothership/request/http' import { withIncomingGoSpan } from '@/lib/mothership/request/otel' const logger = createLogger('BillingUpdateCostAPI') +/** + * How long a cost callback waits on the payer's standing. The worker gives up on the whole + * callback after 5 s, and a cold gate read can wait on the ledger far longer; past this the + * callback answers not-exceeded. The abandoned read keeps running and caches its admission, and + * the next step or re-check reads a refusal again. + */ +const USAGE_STANDING_TIMEOUT_MS = 1000 + const RETRYABLE_SETTLEMENT_RESPONSE = { code: 'BILLING_SETTLEMENT_RETRYABLE', error: 'Billing settlement temporarily unavailable', @@ -57,6 +76,44 @@ function invalidBillingProtocolResponse(requestId: string, span: Span): NextResp ) } +/** + * Reads the run payer's standing after a cost callback, so a long run stops at its next step + * once it crosses the limit instead of at its next admission, with the card the worker writes + * to its log. The payer is the attributed run's, or the one a direct-v1 run was admitted with. + * A duplicate callback answers too: it is often a retry whose first answer was lost. An + * admission is cached per payer and actor for the gate TTL and a refusal is always re-read, so + * steady-state steps cost no ledger read. The charge is + * already recorded when this runs; a gate that cannot answer reports not-exceeded and leaves the + * refusal to the next step or re-check rather than ending a paying run on a database blip, + * and so does a read that outlasts {@link USAGE_STANDING_TIMEOUT_MS}. + */ +async function readUsageStanding( + userId: string, + billingAttribution: BillingAttributionSnapshot | undefined, + accountDecision: AccountBillingDecision | undefined +): Promise { + const readVerdict = billingAttribution + ? () => readMidRunUsageVerdict(billingAttribution) + : accountDecision + ? () => readMidRunAccountUsageVerdict(accountDecision) + : null + if (!isHosted || !readVerdict) return { usageExceeded: false } + let verdict: MidRunUsageVerdict + try { + verdict = await withinDeadline(readVerdict, Date.now() + USAGE_STANDING_TIMEOUT_MS) + } catch { + logger.warn('Usage standing read outlasted the callback budget; answering not exceeded') + return { usageExceeded: false } + } + // Only a spent limit pauses the run. A blocked account is refused at the run's next + // continuation or re-check, with blocked-account copy rather than the upgrade card. + if (verdict.status !== 'exceeded') return { usageExceeded: false } + return { + usageExceeded: true, + usageUpgrade: await resolveUsageUpgradePayload(userId, billingAttribution, verdict.scope), + } +} + function getBillingResolution( isMarkerlessLegacy: boolean, billingAttribution: BillingAttributionSnapshot | undefined @@ -112,9 +169,10 @@ async function updateCostInner(req: NextRequest, span: Span): Promise({ success: true, message: 'Billing disabled, cost update skipped', + usageExceeded: false, data: { billingEnabled: false, processedAt: new Date().toISOString(), @@ -202,6 +260,10 @@ async function updateCostInner(req: NextRequest, span: Span): Promise@`). + if (idempotencyKey?.includes('@')) { + return invalidBillingProtocolResponse(requestId, span) + } const isMcp = source === 'mcp_copilot' span.setAttributes({ @@ -312,6 +374,13 @@ async function updateCostInner(req: NextRequest, span: Span): Promise({ success: true, + ...usageVerdict, data: { userId, cost, diff --git a/apps/sim/app/api/copilot/api-keys/validate/route.test.ts b/apps/sim/app/api/copilot/api-keys/validate/route.test.ts index 48118441140..f9a88e61d3f 100644 --- a/apps/sim/app/api/copilot/api-keys/validate/route.test.ts +++ b/apps/sim/app/api/copilot/api-keys/validate/route.test.ts @@ -6,6 +6,7 @@ import { schemaMock, setEnvFlags, } from '@sim/testing' +import { billingCoreMock, billingCoreMockFns } from '@sim/testing/mocks/billing-core.mock' import { billingPlanMock, billingPlanMockFns } from '@sim/testing/mocks/billing-plan.mock' import { billingSubscriptionMock, @@ -63,7 +64,7 @@ const ATTRIBUTION = { billingEntity: { type: 'organization' as const, id: 'org-1' }, billingPeriod: { start: '2026-07-01T00:00:00.000Z', - end: '2026-08-01T00:00:00.000Z', + end: '2099-01-01T00:00:00.000Z', source: 'reporting' as const, }, payerSubscription: null, @@ -75,7 +76,7 @@ const ACCOUNT_BILLING_DECISION = { billingEntity: { type: 'organization' as const, id: 'account-org' }, billingPeriod: { start: '2026-07-01T00:00:00.000Z', - end: '2026-08-01T00:00:00.000Z', + end: '2099-01-01T00:00:00.000Z', source: 'reporting' as const, }, } @@ -118,6 +119,8 @@ vi.mock('@/lib/billing/calculations/usage-monitor', () => billingUsageMonitorMoc vi.mock('@/lib/billing/core/plan', () => billingPlanMock) +vi.mock('@/lib/billing/core/billing', () => billingCoreMock) + vi.mock('@/lib/billing/core/subscription', () => billingSubscriptionMock) vi.mock('@/lib/billing/core/usage-log', () => billingUsageLogMock) @@ -138,6 +141,8 @@ vi.mock('@/lib/workspaces/permissions/utils', () => permissionsMock) vi.mock('@/lib/workspaces/utils', () => workspacesUtilsMock) import { validateCopilotApiKeyBodySchema } from '@/lib/api/contracts/copilot' +import { resetMidRunUsageCaches } from '@/lib/billing/core/mid-run-usage' +import { resetUsageGateCache } from '@/lib/billing/core/usage-gate-cache' import { POST } from '@/app/api/copilot/api-keys/validate/route' const { mockGetWorkspaceBillingSettings } = workspacesUtilsMockFns @@ -146,7 +151,13 @@ const { mockAuthorizeOrganizationChatDelegation: mockAuthorizeOrganizationChat } mothershipOrganizationChatsMockFns const { mockDeriveBillingContext } = billingUsageLogMockFns const { mockGetHighestPrioritySubscription } = billingPlanMockFns -const { mockCheckServerSideUsageLimits } = billingUsageMonitorMockFns +const { mockGetOrganizationSubscription } = billingCoreMockFns +const { + mockCheckBillingBlocked, + mockCheckBillingEntityBlocked, + mockCheckServerSideUsageLimits, + mockCheckUsageStatus, +} = billingUsageMonitorMockFns const mockIsEnterprisePlan = billingSubscriptionMockFns.mockIsEnterprisePlan const mockGetUserEntityPermissions = permissionsMockFns.mockGetUserEntityPermissions @@ -389,6 +400,9 @@ describe('POST /api/copilot/api-keys/validate billing protocols', () => { }) it('admits a direct-v1 key without Redis while ignoring a local workspace ID', async () => { + mockSerializeAccountBillingDecisionHeader.mockImplementation((decision: object) => + encodeURIComponent(JSON.stringify(decision)) + ) mockGetUserEntityPermissions.mockResolvedValueOnce(null) mockGetWorkspaceBillingSettings.mockResolvedValueOnce({ billedAccountUserId: 'different-owner', @@ -415,8 +429,9 @@ describe('POST /api/copilot/api-keys/validate billing protocols', () => { expect(mockResolveBillingAttribution).not.toHaveBeenCalled() expect(mockGetUserEntityPermissions).not.toHaveBeenCalled() expect(mockGetWorkspaceBillingSettings).not.toHaveBeenCalled() - expect(mockSerializeAccountBillingDecisionHeader).toHaveBeenCalledWith(ACCOUNT_BILLING_DECISION) - expect(res.headers.get('x-sim-billing-account-decision')).toBe('serialized-account-decision') + expect( + JSON.parse(decodeURIComponent(res.headers.get('x-sim-billing-account-decision') ?? '')) + ).toEqual({ ...ACCOUNT_BILLING_DECISION, payerSubscriptionId: ACCOUNT_SUBSCRIPTION.id }) }) it('fails direct-v1 admission closed when its payer cannot be resolved', async () => { @@ -506,7 +521,21 @@ describe('validation lifecycle purposes', () => { mockCheckInternalApiKey.mockReturnValue({ success: true }) mockAuthorizeCallback.mockReset().mockResolvedValue(undefined) mockCheckContinuationBilling.mockReset().mockResolvedValue({ blocked: false }) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: false }) + mockCheckUsageStatus.mockResolvedValue({ isExceeded: false, currentUsage: 1, limit: 10 }) + mockCheckBillingBlocked.mockResolvedValue({ blocked: false }) + mockCheckBillingEntityBlocked.mockResolvedValue({ blocked: false }) mockIsEnterprisePlan.mockResolvedValue(false) + mockGetOrganizationSubscription.mockResolvedValue({ + id: 'sub-org-1', + referenceId: 'org-1', + plan: 'enterprise', + status: 'active', + periodStart: new Date(ATTRIBUTION.billingPeriod.start), + periodEnd: new Date(ATTRIBUTION.billingPeriod.end), + }) + resetUsageGateCache() + resetMidRunUsageCaches() }) it('defaults older callers to full admission and rejects unknown purposes', () => { @@ -516,7 +545,7 @@ describe('validation lifecycle purposes', () => { ).toBe(false) }) - it('checks original payer and current scope without repeating spend admission', async () => { + it('checks original payer, current scope, and the original payer spend', async () => { const response = await POST(request(body, attributedHeaders)) expect(response.status).toBe(200) expect(mockAuthorizeCallback).toHaveBeenCalledWith({ ...body, delegationId: requestId }) @@ -527,7 +556,6 @@ describe('validation lifecycle purposes', () => { expect(mockAuthorizeCallback.mock.invocationCallOrder[0]).toBeLessThan( mockCheckContinuationBilling.mock.invocationCallOrder[0] ) - expect(mockCheckAttributedUsageLimits).not.toHaveBeenCalled() expect(mockCheckServerSideUsageLimits).not.toHaveBeenCalled() expect(mockResolveLegacyV0BillingAttribution).not.toHaveBeenCalled() expect(mockGetHighestPrioritySubscription).not.toHaveBeenCalled() @@ -553,6 +581,23 @@ describe('validation lifecycle purposes', () => { expect(response.headers.get('x-sim-billing-account-decision')).toBeNull() }) + it('refuses a direct-v1 continuation whose account is over its usage limit', async () => { + mockCheckUsageStatus.mockResolvedValue({ isExceeded: true, currentUsage: 12, limit: 10 }) + + const response = await POST(request(body, directHeaders)) + + expect(response.status).toBe(402) + await expect(response.json()).resolves.toMatchObject({ + code: 'USAGE_LIMIT_EXCEEDED', + usageUpgrade: { reason: 'usage_limit' }, + }) + }) + + it('admits a direct-v1 continuation whose usage cannot be read', async () => { + mockCheckUsageStatus.mockResolvedValue({ isExceeded: true, unavailable: true }) + expect((await POST(request(body, directHeaders))).status).toBe(200) + }) + it.each([ ['missing attribution', { ...attributedHeaders, 'x-sim-billing-attribution': '' }], [ @@ -638,11 +683,186 @@ describe('validation lifecycle purposes', () => { }) it.each(['actor', 'payer'])('refuses a newly blocked %s on continuation', async (scope) => { - mockCheckContinuationBilling.mockResolvedValueOnce({ blocked: true, scope }) - expect((await POST(request(body, attributedHeaders))).status).toBe(402) + mockCheckContinuationBilling.mockResolvedValueOnce({ + blocked: true, + scope, + message: 'Billing account frozen.', + }) + const response = await POST(request(body, attributedHeaders)) + expect(response.status).toBe(402) + await expect(response.json()).resolves.toEqual({ + code: 'BILLING_BLOCKED', + error: 'Billing account frozen.', + }) expect(mockCheckAttributedUsageLimits).not.toHaveBeenCalled() }) + it('refuses a payer the usage gate finds blocked as blocked, without the usage card', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ + isExceeded: true, + reason: 'billing_blocked', + message: 'Organization billing issue.', + scope: 'payer', + }) + const response = await POST(request(body, attributedHeaders)) + expect(response.status).toBe(402) + await expect(response.json()).resolves.toEqual({ + code: 'BILLING_BLOCKED', + error: 'Organization billing issue.', + }) + }) + + it('admits a polled continuation whose spend cannot be read', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ + isExceeded: true, + reason: 'usage_unavailable', + }) + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + mockCheckAttributedUsageLimits.mockRejectedValue(new Error('ledger read timed out')) + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + }) + + it('refuses a continuation over its usage limit with the card the worker writes', async () => { + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const response = await POST(request(body, attributedHeaders)) + + expect(response.status).toBe(402) + await expect(response.json()).resolves.toEqual({ + code: 'USAGE_LIMIT_EXCEEDED', + error: expect.stringContaining('usage limit'), + usageUpgrade: { + reason: 'usage_limit', + action: 'upgrade_plan', + message: expect.stringContaining('usage limit'), + }, + }) + }) + + it('answers a polled re-check from the cached admission and always re-reads a refusal', async () => { + for (let call = 0; call < 2; call++) queueTableRows(schemaMock.user, [{ id: 'user-1' }]) + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + for (let poll = 0; poll < 2; poll++) { + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + } + + resetUsageGateCache() + expect((await POST(request(body, attributedHeaders))).status).toBe(402) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: false }) + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + }) + + it('checks the payer saved at admission for a direct-v1 run whose actor changed orgs', async () => { + const endedDecision = { + ...ACCOUNT_BILLING_DECISION, + billingPeriod: { start: '2026-06-01T00:00:00.000Z', end: '2026-07-01T00:00:00.000Z' }, + } + mockGetHighestPrioritySubscription.mockResolvedValue({ + id: 'sub-new-org', + referenceId: 'new-org', + plan: 'team', + status: 'active', + periodStart: new Date('2026-07-01T00:00:00.000Z'), + periodEnd: new Date('2099-01-01T00:00:00.000Z'), + }) + mockGetOrganizationSubscription.mockResolvedValue({ + id: 'sub-account-org', + referenceId: 'account-org', + plan: 'team', + status: 'active', + periodStart: new Date('2026-07-01T00:00:00.000Z'), + periodEnd: new Date('2099-01-01T00:00:00.000Z'), + }) + mockCheckUsageStatus.mockImplementation( + async ( + _userId: string, + _subscription: unknown, + context?: { billingEntity: { id: string } } + ) => ({ + isExceeded: context?.billingEntity.id === 'account-org', + currentUsage: 12, + limit: 10, + }) + ) + + const response = await POST( + request(body, { ...directHeaders, 'x-sim-billing-account-decision': encode(endedDecision) }) + ) + + expect(response.status).toBe(402) + }) + + it('answers repeated direct-v1 continuations from the cached admission and re-reads a refusal', async () => { + for (let call = 0; call < 2; call++) queueTableRows(schemaMock.user, [{ id: 'user-1' }]) + expect((await POST(request(body, directHeaders))).status).toBe(200) + mockCheckUsageStatus.mockResolvedValue({ isExceeded: true, currentUsage: 12, limit: 10 }) + for (let leg = 0; leg < 2; leg++) { + expect((await POST(request(body, directHeaders))).status).toBe(200) + } + + resetMidRunUsageCaches() + expect((await POST(request(body, directHeaders))).status).toBe(402) + mockCheckUsageStatus.mockResolvedValue({ isExceeded: false, currentUsage: 1, limit: 10 }) + expect((await POST(request(body, directHeaders))).status).toBe(200) + }) + + it('never reads the usage gate for an attributed continuation when billing is off', async () => { + setEnvFlags({ isHosted: false, isBillingEnabled: false }) + mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + expect((await POST(request(body, attributedHeaders))).status).toBe(200) + }) + + it('never reads the ledger for a direct-v1 continuation when billing is off', async () => { + setEnvFlags({ isHosted: false, isBillingEnabled: false }) + mockCheckUsageStatus.mockResolvedValue({ isExceeded: true, currentUsage: 12, limit: 10 }) + + expect((await POST(request(body, directHeaders))).status).toBe(200) + }) + + it('judges a direct-v1 reporting-window run against its admitted window after it ends', async () => { + const admittedWindow = { + ...ACCOUNT_BILLING_DECISION, + billingPeriod: { + start: '2026-06-01T00:00:00.000Z', + end: '2026-07-01T00:00:00.000Z', + source: 'reporting' as const, + }, + } + mockCheckUsageStatus.mockImplementation( + async ( + _userId: string, + _subscription: unknown, + context?: { billingPeriod: { start: Date } } + ) => ({ + isExceeded: + context?.billingPeriod.start.toISOString() === admittedWindow.billingPeriod.start, + currentUsage: 12, + limit: 10, + }) + ) + + const response = await POST( + request(body, { ...directHeaders, 'x-sim-billing-account-decision': encode(admittedWindow) }) + ) + + expect(response.status).toBe(402) + }) + + it('judges a direct-v1 organization payer without a subscription as that organization', async () => { + mockGetOrganizationSubscription.mockResolvedValue(null) + mockCheckUsageStatus.mockImplementation( + async (_userId: string, subscription: { referenceId?: string } | null) => ({ + isExceeded: subscription?.referenceId === 'account-org', + currentUsage: 12, + limit: 10, + }) + ) + + expect((await POST(request(body, directHeaders))).status).toBe(402) + }) + it('allows cancellation without billing material or spending/standing/plan checks', async () => { const response = await POST( request({ ...body, purpose: 'cancellation' }, { 'x-sim-billing-protocol': 'attribution-v1' }) @@ -752,7 +972,6 @@ describe('validation lifecycle purposes', () => { expect((await POST(request({ ...body, purpose: 'new-turn' }, attributedHeaders))).status).toBe( 402 ) - expect(mockCheckAttributedUsageLimits).toHaveBeenCalledTimes(1) expect((await POST(request({ ...body, purpose: 'new-turn' }, directHeaders))).status).toBe(400) expect(mockCheckServerSideUsageLimits).not.toHaveBeenCalled() }) diff --git a/apps/sim/app/api/copilot/api-keys/validate/route.ts b/apps/sim/app/api/copilot/api-keys/validate/route.ts index 0b431ed12d4..8723a9d5c90 100644 --- a/apps/sim/app/api/copilot/api-keys/validate/route.ts +++ b/apps/sim/app/api/copilot/api-keys/validate/route.ts @@ -4,7 +4,13 @@ import { createLogger } from '@sim/logger' import { generateId } from '@sim/utils/id' import { eq } from 'drizzle-orm' import { type NextRequest, NextResponse } from 'next/server' -import { validateCopilotApiKeyContract } from '@/lib/api/contracts/copilot' +import { + COPILOT_BILLING_BLOCKED_CODE, + COPILOT_USAGE_LIMIT_EXCEEDED_CODE, + type ValidateCopilotApiKeyBillingBlocked, + type ValidateCopilotApiKeyUsageExceeded, + validateCopilotApiKeyContract, +} from '@/lib/api/contracts/copilot' import { parseRequest, validationErrorResponse } from '@/lib/api/server' import { checkServerSideUsageLimits } from '@/lib/billing/calculations/usage-monitor' import { @@ -20,9 +26,14 @@ import { serializeAccountBillingDecisionHeader, serializeBillingAttributionHeader, } from '@/lib/billing/core/billing-attribution' +import { + readMidRunAccountUsageVerdict, + readMidRunUsageVerdict, +} from '@/lib/billing/core/mid-run-usage' import { getHighestPrioritySubscription } from '@/lib/billing/core/plan' import { isEnterprisePlan } from '@/lib/billing/core/subscription' import { deriveBillingContext } from '@/lib/billing/core/usage-log' +import { resolveUsageUpgradePayload } from '@/lib/billing/usage-upgrade' import { isBillingEnabled, isHosted } from '@/lib/core/config/env-flags' import { asOrchestrationError } from '@/lib/core/orchestration/types' import { withRouteHandler } from '@/lib/core/utils/with-route-handler' @@ -53,6 +64,8 @@ import { withIncomingGoSpan } from '@/lib/mothership/request/otel' const logger = createLogger('CopilotApiKeysValidate') +const CONTINUATION_BLOCKED_MESSAGE = 'Continuation billing account is blocked' + function invalidBillingProtocolResponse(): NextResponse { return NextResponse.json({ error: 'Invalid billing attribution protocol' }, { status: 400 }) } @@ -272,6 +285,7 @@ async function checkAdmissionUsage(admission: AdmissionBillingDecision): Promise ? { source: billingContext.billingPeriod.source } : {}), }, + ...(subscription ? { payerSubscriptionId: subscription.id } : {}), }, } } @@ -402,13 +416,54 @@ export const POST = withRouteHandler((req: NextRequest) => blocked: blocked?.blocked ?? false, elapsedMs: Math.round(performance.now() - startedAt), }) - if (blocked?.blocked) { + // A continuation, and a worker's periodic re-check of a long run, also reads the + // original payer's spend through the cached execution usage gate. A read that fails + // admits: the run is already under way, and the next re-check reads again. + const verdict = + !blocked?.blocked && purpose === COPILOT_VALIDATION_PURPOSE.continuation && billing + ? billing.kind === 'attributed' + ? await readMidRunUsageVerdict(billing.attribution) + : await readMidRunAccountUsageVerdict(billing.decision) + : null + if (blocked?.blocked || verdict?.status === 'blocked') { + span.setAttribute( + TraceAttr.CopilotValidateOutcome, + CopilotValidateOutcome.UsageExceeded + ) + span.setAttribute(TraceAttr.HttpStatusCode, 402) + return NextResponse.json( + { + code: COPILOT_BILLING_BLOCKED_CODE, + error: + (blocked?.blocked + ? blocked.message + : verdict?.status === 'blocked' + ? verdict.message + : undefined) ?? CONTINUATION_BLOCKED_MESSAGE, + }, + { status: 402 } + ) + } + if (verdict?.status === 'exceeded') { + logger.info('[API VALIDATION] Continuation usage exceeded', { userId }) span.setAttribute( TraceAttr.CopilotValidateOutcome, CopilotValidateOutcome.UsageExceeded ) span.setAttribute(TraceAttr.HttpStatusCode, 402) - return new NextResponse(null, { status: 402 }) + const usageUpgrade = await resolveUsageUpgradePayload( + userId, + billing?.kind === 'attributed' ? billing.attribution : undefined, + verdict.scope + ) + return NextResponse.json( + { + code: COPILOT_USAGE_LIMIT_EXCEEDED_CODE, + error: usageUpgrade.message, + usageUpgrade, + }, + { status: 402 } + ) } const isEnterprise = purpose === COPILOT_VALIDATION_PURPOSE.cancellation diff --git a/apps/sim/app/workspace/[workspaceId]/home/components/message-content/components/special-tags/special-tags.test.ts b/apps/sim/app/workspace/[workspaceId]/home/components/message-content/components/special-tags/special-tags.test.ts index 5e8d4273732..00894654cba 100644 --- a/apps/sim/app/workspace/[workspaceId]/home/components/message-content/components/special-tags/special-tags.test.ts +++ b/apps/sim/app/workspace/[workspaceId]/home/components/message-content/components/special-tags/special-tags.test.ts @@ -10,6 +10,7 @@ import { describe, expect, it, vi } from 'vitest' */ vi.mock('@/lib/auth/auth-client', () => authClientMock) +import { formatUsageUpgradeTag } from '@/lib/billing/usage-upgrade' import type { ContentSegment, CredentialItemData, @@ -875,3 +876,21 @@ describe('source tag', () => { } }) }) + +describe('usage card written to a worker log', () => { + it('renders the card Sim hands the worker when the text is replayed after a reload', () => { + const usageUpgrade = { + reason: 'usage_limit', + action: 'increase_limit', + message: "You've reached your usage limit for this billing period.", + } as const + const replayed = `Finished the first report.${formatUsageUpgradeTag(usageUpgrade)}` + + const { segments } = parseSpecialTags(replayed, false) + + expect(segments).toEqual([ + { type: 'text', content: 'Finished the first report.' }, + { type: 'usage_upgrade', data: usageUpgrade }, + ]) + }) +}) diff --git a/apps/sim/lib/api/contracts/copilot.ts b/apps/sim/lib/api/contracts/copilot.ts index d9fcd2ff2a5..d3682550460 100644 --- a/apps/sim/lib/api/contracts/copilot.ts +++ b/apps/sim/lib/api/contracts/copilot.ts @@ -3,6 +3,7 @@ import { persistedContentBlockSchema } from '@/lib/api/contracts/copilot-message import { workspaceSearchFiltersSchema } from '@/lib/api/contracts/knowledge/search' import { mothershipResourceSchema } from '@/lib/api/contracts/mothership-resources' import { requiredFieldSchema, workspaceIdSchema } from '@/lib/api/contracts/primitives' +import { usageUpgradePayloadSchema } from '@/lib/api/contracts/subscription' import { type ContractJsonResponse, defineRouteContract } from '@/lib/api/contracts/types' import { ASYNC_TOOL_CONFIRMATION_STATUS, @@ -270,6 +271,38 @@ export const validateCopilotApiKeyResponseSchema = z.object({ }) export type ValidateCopilotApiKeyResponse = z.output +export const COPILOT_USAGE_LIMIT_EXCEEDED_CODE = 'USAGE_LIMIT_EXCEEDED' +export const COPILOT_BILLING_BLOCKED_CODE = 'BILLING_BLOCKED' + +/** + * A continuation refused because the run's original payer is over its usage limit. A worker + * polling continuation validation mid-run writes `usageUpgrade` as a `` tag into + * its log and pauses the run for the limit. + */ +export const validateCopilotApiKeyUsageExceededSchema = z.object({ + code: z.literal(COPILOT_USAGE_LIMIT_EXCEEDED_CODE), + error: z.string(), + usageUpgrade: usageUpgradePayloadSchema, +}) +export type ValidateCopilotApiKeyUsageExceeded = z.output< + typeof validateCopilotApiKeyUsageExceededSchema +> + +/** A continuation refused because the actor or payer account is blocked (payment, dispute). */ +export const validateCopilotApiKeyBillingBlockedSchema = z.object({ + code: z.literal(COPILOT_BILLING_BLOCKED_CODE), + error: z.string(), +}) +export type ValidateCopilotApiKeyBillingBlocked = z.output< + typeof validateCopilotApiKeyBillingBlockedSchema +> + +/** A continuation 402 carries one of these bodies; a new turn's 402 is empty. */ +export const validateCopilotApiKeyRefusalSchema = z.union([ + validateCopilotApiKeyUsageExceededSchema, + validateCopilotApiKeyBillingBlockedSchema, +]) + export const listCopilotApiKeysContract = defineRouteContract({ method: 'GET', path: '/api/copilot/api-keys', @@ -392,7 +425,15 @@ export const validateCopilotApiKeyContract = defineRouteContract({ path: '/api/copilot/api-keys/validate', headers: validateCopilotApiKeyHeadersSchema, body: validateCopilotApiKeyBodySchema, - response: { mode: 'json', schema: validateCopilotApiKeyResponseSchema }, + response: { + mode: 'json', + schema: validateCopilotApiKeyResponseSchema, + status: [200, 402], + statusSchemas: { + 200: validateCopilotApiKeyResponseSchema, + 402: validateCopilotApiKeyRefusalSchema, + }, + }, error: validateCopilotApiKeyErrorSchema, }) diff --git a/apps/sim/lib/api/contracts/subscription.ts b/apps/sim/lib/api/contracts/subscription.ts index 4d1efa6dd63..424f17ceb10 100644 --- a/apps/sim/lib/api/contracts/subscription.ts +++ b/apps/sim/lib/api/contracts/subscription.ts @@ -325,17 +325,54 @@ export const billingSwitchPlanResponseSchema = z.object({ message: z.string().optional(), }) -export const billingUpdateCostResponseSchema = z.object({ - success: z.literal(true), - message: z.string().optional(), - data: z.object({ - userId: z.string().optional(), - cost: z.number().optional(), - billingEnabled: z.boolean().optional(), - processedAt: z.string(), - requestId: z.string(), - }), +/** + * The upgrade card's payload: the JSON body of a `` tag in assistant text, which + * the chat renders as the usage card wherever that text appears (live, replayed, or reloaded). + * Sim decides the action and copy from the payer's plan; a worker ending a run at the usage + * limit writes the tag with this payload verbatim into its durable log. + */ +export const usageUpgradePayloadSchema = z.object({ + reason: z.literal('usage_limit'), + action: z.enum(['upgrade_plan', 'increase_limit']), + message: z.string(), }) +export type UsageUpgradePayload = z.infer + +/** + * The payer's standing after a cost callback, read through the cached execution usage gate. It + * sits at the top level of the body, beside `success`, where the worker's shared + * `BillingCallbackResult` reads it, on a 200 and on a duplicate 409 alike. A worker that + * predates the fields ignores them. + */ +export const billingUsageVerdictSchema = z.discriminatedUnion('usageExceeded', [ + z.object({ + /** The payer is within its usage limit, or its standing could not be read. */ + usageExceeded: z.literal(false), + usageUpgrade: z.never().optional(), + }), + z.object({ + /** The payer is over its usage limit; the worker pauses the run at its next step boundary. */ + usageExceeded: z.literal(true), + /** The card the worker writes to its log. */ + usageUpgrade: usageUpgradePayloadSchema, + }), +]) +export type BillingUsageVerdict = z.infer + +export const billingUpdateCostResponseSchema = z + .object({ + success: z.literal(true), + message: z.string().optional(), + data: z.object({ + userId: z.string().optional(), + cost: z.number().optional(), + billingEnabled: z.boolean().optional(), + processedAt: z.string(), + requestId: z.string(), + }), + }) + .and(billingUsageVerdictSchema) +export type BillingUpdateCostResponse = z.infer export const billingSwitchPlanContract = defineRouteContract({ method: 'POST', diff --git a/apps/sim/lib/billing/calculations/usage-monitor.test.ts b/apps/sim/lib/billing/calculations/usage-monitor.test.ts index 7571686ee5b..5b587a36342 100644 --- a/apps/sim/lib/billing/calculations/usage-monitor.test.ts +++ b/apps/sim/lib/billing/calculations/usage-monitor.test.ts @@ -130,6 +130,21 @@ describe('checkUsageStatus', () => { }) }) + it('refuses on a ledger read failure but marks the answer unavailable', async () => { + mockGetBillingPeriodUsageCost.mockRejectedValueOnce(new Error('canceling statement')) + + await expect( + checkUsageStatus('user-1', { + referenceId: 'user-1', + plan: 'free', + status: 'active', + seats: 1, + periodStart: new Date('2026-06-01T00:00:00.000Z'), + periodEnd: new Date('2026-07-01T00:00:00.000Z'), + }) + ).resolves.toMatchObject({ isExceeded: true, unavailable: true }) + }) + it('preserves negative ledger-only personal usage', async () => { const periodStart = new Date('2026-06-01T00:00:00.000Z') const periodEnd = new Date('2026-07-01T00:00:00.000Z') @@ -198,6 +213,23 @@ describe('checkServerSideUsageLimits', () => { mockGetBillingPeriodUsageCost.mockResolvedValue(125) }) + it('does not describe an unreadable ledger as a spent limit', async () => { + dbChainMockFns.limit.mockResolvedValueOnce([{ blocked: false }]) + mockGetBillingPeriodUsageCost.mockRejectedValueOnce(new Error('canceling statement')) + + const result = await checkServerSideUsageLimits('user-1', { + referenceId: 'user-1', + plan: 'free', + status: 'active', + seats: 1, + periodStart: new Date('2026-06-01T00:00:00.000Z'), + periodEnd: new Date('2026-07-01T00:00:00.000Z'), + }) + + expect(result.isExceeded).toBe(true) + expect(result.message ?? '').not.toMatch(/\$/) + }) + it('keeps blocked accounts blocked while reporting their real ledger usage', async () => { dbChainMockFns.limit.mockResolvedValueOnce([{ blocked: true, blockedReason: 'payment_failed' }]) const subscription = { diff --git a/apps/sim/lib/billing/calculations/usage-monitor.ts b/apps/sim/lib/billing/calculations/usage-monitor.ts index f4ebe5bc66b..9ebb874e8b8 100644 --- a/apps/sim/lib/billing/calculations/usage-monitor.ts +++ b/apps/sim/lib/billing/calculations/usage-monitor.ts @@ -3,6 +3,7 @@ import { userStats } from '@sim/db/schema' import { createLogger } from '@sim/logger' import { toError } from '@sim/utils/errors' import { eq } from 'drizzle-orm' +import { USAGE_UNAVAILABLE_MESSAGE } from '@/lib/billing/constants' import { isOrganizationBillingBlocked } from '@/lib/billing/core/access' import { defaultBillingPeriod } from '@/lib/billing/core/billing-period' import { getHighestPrioritySubscription } from '@/lib/billing/core/plan' @@ -43,6 +44,11 @@ interface UsageData { scope: 'user' | 'organization' /** Present only when `scope === 'organization'`. */ organizationId: string | null + /** + * The ledger could not be read, so `isExceeded` is a fail-closed refusal rather than a + * measured one. Admission refuses on it; a run already under way treats it as unknown. + */ + unavailable?: true } /** @@ -183,6 +189,7 @@ export async function checkUsageStatus( limit: 0, scope: 'user', organizationId: null, + unavailable: true, } } } @@ -352,7 +359,11 @@ export async function checkServerSideUsageLimits( isExceeded: usageData.isExceeded, currentUsage: usageData.currentUsage, limit: usageData.limit, - message: usageData.isExceeded ? exceededMessage : undefined, + message: usageData.unavailable + ? USAGE_UNAVAILABLE_MESSAGE + : usageData.isExceeded + ? exceededMessage + : undefined, } } catch (error) { logger.error('Error in server-side usage limit check', { diff --git a/apps/sim/lib/billing/constants.ts b/apps/sim/lib/billing/constants.ts index c1fd884a293..a3db9fe9acf 100644 --- a/apps/sim/lib/billing/constants.ts +++ b/apps/sim/lib/billing/constants.ts @@ -103,3 +103,7 @@ export const ANNUAL_DISCOUNT_RATE = 0.15 * Effectively unlimited — any limit >= this threshold is treated as uncapped. */ export const ON_DEMAND_UNLIMITED = 999999 + +/** Shown when usage could not be read, instead of a limit the read never measured. */ +export const USAGE_UNAVAILABLE_MESSAGE = + 'Usage could not be verified right now. Please try again in a moment.' diff --git a/apps/sim/lib/billing/core/billing-attribution.test.ts b/apps/sim/lib/billing/core/billing-attribution.test.ts index 5cec0fa1918..8e40d3456da 100644 --- a/apps/sim/lib/billing/core/billing-attribution.test.ts +++ b/apps/sim/lib/billing/core/billing-attribution.test.ts @@ -17,8 +17,10 @@ import { assertBillingAttributionOwner, assertBillingAttributionSnapshot, billingAttributionsEqual, + checkAccountBillingBlocks, checkAttributedBillingBlocks, checkAttributedUsageLimits, + requireAccountBillingDecisionHeader, requireBillingAttributionHeader, requireBillingCallbackAttribution, requireBillingRequestIdHeader, @@ -195,6 +197,38 @@ describe('resolveBillingAttribution', () => { }) }) +describe('account billing decision header', () => { + const decision = { + userId: 'actor', + billingEntity: { type: 'user', id: 'actor' }, + billingPeriod: { + start: '2026-07-01T00:00:00.000Z', + end: '2026-08-01T00:00:00.000Z', + source: 'stripe', + }, + } + const header = (value: unknown) => + new Headers({ 'x-sim-billing-account-decision': encodeURIComponent(JSON.stringify(value)) }) + + it('restores the admitted payer subscription', () => { + expect( + requireAccountBillingDecisionHeader(header({ ...decision, payerSubscriptionId: 'sub-1' })) + ).toMatchObject({ payerSubscriptionId: 'sub-1' }) + expect(requireAccountBillingDecisionHeader(header(decision))).not.toHaveProperty( + 'payerSubscriptionId' + ) + }) + + it.each([42, '', ' ', null, { id: 'sub-1' }])( + 'refuses a payer subscription of %j', + (payerSubscriptionId) => { + expect(() => + requireAccountBillingDecisionHeader(header({ ...decision, payerSubscriptionId })) + ).toThrow('Account billing decision header is malformed') + } + ) +}) + describe('serialized attribution boundaries', () => { const attribution = { actorUserId: 'actor-a', @@ -333,6 +367,55 @@ describe('serialized attribution boundaries', () => { }) }) +describe('checkAccountBillingBlocks', () => { + const decision = { + userId: 'actor', + billingEntity: { type: 'organization' as const, id: 'original-payer' }, + billingPeriod: { start: '2026-07-01T00:00:00.000Z', end: '2026-08-01T00:00:00.000Z' }, + } + + beforeEach(() => { + mockCheckBillingBlocked.mockReset().mockResolvedValue({ blocked: false }) + mockCheckBillingEntityBlocked.mockReset().mockResolvedValue({ blocked: false }) + }) + + it('refuses the exact actor and original payer when either is blocked', async () => { + mockCheckBillingBlocked.mockImplementation(async (userId: string) => ({ + blocked: userId === 'actor', + })) + await expect(checkAccountBillingBlocks(decision)).resolves.toMatchObject({ + blocked: true, + scope: 'actor', + }) + + mockCheckBillingBlocked.mockResolvedValue({ blocked: false }) + mockCheckBillingEntityBlocked.mockImplementation(async (entity: { id: string }) => ({ + blocked: entity.id === 'original-payer', + })) + await expect(checkAccountBillingBlocks(decision)).resolves.toMatchObject({ + blocked: true, + scope: 'payer', + }) + }) + + it('reports an actor block ahead of a payer block', async () => { + mockCheckBillingBlocked.mockResolvedValue({ blocked: true, message: 'Actor frozen.' }) + mockCheckBillingEntityBlocked.mockResolvedValue({ blocked: true, message: 'Payer frozen.' }) + await expect(checkAccountBillingBlocks(decision)).resolves.toEqual({ + blocked: true, + message: 'Actor frozen.', + scope: 'actor', + }) + }) + + it('answers a personal payer from the actor standing alone', async () => { + mockCheckBillingEntityBlocked.mockResolvedValue({ blocked: true }) + await expect( + checkAccountBillingBlocks({ ...decision, billingEntity: { type: 'user', id: 'actor' } }) + ).resolves.toMatchObject({ blocked: false }) + }) +}) + describe('checkAttributedUsageLimits', () => { beforeEach(() => { resetDbChainMock() @@ -432,6 +515,28 @@ describe('checkAttributedUsageLimits', () => { expect(mockCheckOrganizationMemberUsageLimit).not.toHaveBeenCalled() }) + it('names a billing block and an unreadable ledger apart from a spent limit', async () => { + mockCheckBillingBlocked.mockResolvedValueOnce({ blocked: true, message: 'Frozen.' }) + await expect(checkAttributedUsageLimits(attribution)).resolves.toMatchObject({ + isExceeded: true, + reason: 'billing_blocked', + }) + + mockCheckUsageStatus.mockResolvedValueOnce({ + currentUsage: 0, + isExceeded: true, + limit: 0, + organizationId: null, + percentUsed: 100, + isWarning: false, + scope: 'user', + unavailable: true, + }) + const unavailable = await checkAttributedUsageLimits(attribution) + expect(unavailable).toMatchObject({ isExceeded: true, reason: 'usage_unavailable' }) + expect(unavailable.message ?? '').not.toMatch(/\$/) + }) + it('returns payer exhaustion before checking the actor member cap', async () => { mockCheckUsageStatus.mockResolvedValue({ currentUsage: 100, diff --git a/apps/sim/lib/billing/core/billing-attribution.ts b/apps/sim/lib/billing/core/billing-attribution.ts index 6c205ca33bd..9e734dc487c 100644 --- a/apps/sim/lib/billing/core/billing-attribution.ts +++ b/apps/sim/lib/billing/core/billing-attribution.ts @@ -10,6 +10,7 @@ import { checkUsageStatus, } from '@/lib/billing/calculations/usage-monitor' import { parseBillingConcurrencyLimit } from '@/lib/billing/concurrency-defaults' +import { USAGE_UNAVAILABLE_MESSAGE } from '@/lib/billing/constants' import { getOrganizationSubscription } from '@/lib/billing/core/billing' import { defaultBillingPeriod } from '@/lib/billing/core/billing-period' import { getHighestPriorityPersonalSubscription } from '@/lib/billing/core/plan' @@ -96,6 +97,14 @@ export interface AccountBillingDecision { readonly end: string readonly source?: UsagePeriodSource } + /** + * The payer's subscription at admission, so a run that outlives a Stripe period bills its + * later spend to the period it was spent in, as an attributed run's `payerSubscription` does. + * Absent for a payer without a subscription, and in decisions minted before it existed. + * If that subscription is replaced mid-run, spend stays in its period while the mid-run + * verdict judges the payer's current one, so the limit can only be under-enforced. + */ + readonly payerSubscriptionId?: string } export interface ResolveBillingAttributionParams { @@ -111,6 +120,11 @@ export interface AttributedUsageLimitsResult { isExceeded: boolean message?: string scope?: 'actor' | 'payer' | 'member' + /** + * Why an `isExceeded` refusal is not a spent limit: the account is blocked (payment failed, + * dispute), or the payer's usage could not be read and the gate failed closed. + */ + reason?: 'billing_blocked' | 'usage_unavailable' payerUsage?: { currentUsage: number limit: number @@ -540,6 +554,10 @@ function assertAccountBillingDecision(value: unknown): AccountBillingDecision { ) { throw new Error('Account billing decision must contain a valid billing period source') } + const payerSubscriptionId = value.payerSubscriptionId + if (payerSubscriptionId !== undefined && !isNonEmptyString(payerSubscriptionId)) { + throw new Error('Account billing decision must contain a valid payer subscription ID') + } return Object.freeze({ userId: value.userId, @@ -552,6 +570,7 @@ function assertAccountBillingDecision(value: unknown): AccountBillingDecision { end: end.toISOString(), ...(source !== undefined ? { source } : {}), }), + ...(payerSubscriptionId !== undefined ? { payerSubscriptionId } : {}), }) } @@ -720,6 +739,36 @@ function buildBillingAttributionSnapshot(params: { }) } +/** + * The same payer's attribution for its current usage period, for a run that outlived the period + * it was admitted in. The actor, workspace and payer are kept; only the payer's subscription, + * and so its period, is read again, and a subscription that no longer belongs to the payer is + * refused rather than re-selected. + */ +export async function refreshAttributionPeriod( + attribution: BillingAttributionSnapshot +): Promise { + const validated = assertBillingAttributionSnapshot(attribution) + const payerSubscription = validated.organizationId + ? await getOrganizationSubscription(validated.organizationId, { onError: 'throw' }) + : await getHighestPriorityPersonalSubscription(validated.billedAccountUserId, { + onError: 'throw', + }) + const expectedReferenceId = validated.organizationId ?? validated.billedAccountUserId + if (payerSubscription && payerSubscription.referenceId !== expectedReferenceId) { + throw new Error( + `Resolved subscription ${payerSubscription.id} does not belong to payer ${expectedReferenceId}` + ) + } + return buildBillingAttributionSnapshot({ + actorUserId: validated.actorUserId, + workspaceId: validated.workspaceId, + billedAccountUserId: validated.billedAccountUserId, + organizationId: validated.organizationId, + payerSubscription, + }) +} + /** * Resolves the payer from the workspace without consulting the actor's * subscriptions or organization memberships. @@ -901,6 +950,20 @@ export async function checkAttributedBillingBlocks( return { blocked: false } } +/** + * The same freeze checks for a direct-v1 run: the actor's own account, then the payer saved in + * its admission decision, never one re-selected from the actor's current memberships. + */ +export async function checkAccountBillingBlocks( + decision: AccountBillingDecision +): Promise { + const actorBlock = await checkBillingBlocked(decision.userId) + if (actorBlock.blocked) return { ...actorBlock, scope: 'actor' } + const payer = decision.billingEntity + if (payer.type === 'user' && payer.id === decision.userId) return actorBlock + return { ...(await checkBillingEntityBlocked(payer)), scope: 'payer' } +} + /** * Applies hosted billing gates in canonical order: actor account, workspace * payer pool, then `(organizationId, actorUserId)` member cap. @@ -919,6 +982,7 @@ export async function checkAttributedUsageLimits( isExceeded: true, message: billingBlock.message, scope: billingBlock.scope, + reason: 'billing_blocked', } } @@ -934,8 +998,9 @@ export async function checkAttributedUsageLimits( if (payerUsage.isExceeded) { const formattedUsage = payerUsage.currentUsage.toFixed(2) const formattedLimit = payerUsage.limit.toFixed(2) - const message = - validatedAttribution.billingEntity.type === 'organization' + const message = payerUsage.unavailable + ? USAGE_UNAVAILABLE_MESSAGE + : validatedAttribution.billingEntity.type === 'organization' ? `Organization usage limit exceeded: $${formattedUsage} pooled of $${formattedLimit} organization limit. Ask a team admin to raise the organization usage limit to continue.` : `Usage limit exceeded: $${formattedUsage} used of $${formattedLimit} limit. Please upgrade your plan or raise your usage limit to continue.` @@ -944,6 +1009,7 @@ export async function checkAttributedUsageLimits( message, scope: 'payer', payerUsage: payerSnapshot, + ...(payerUsage.unavailable ? { reason: 'usage_unavailable' as const } : {}), } } diff --git a/apps/sim/lib/billing/core/mid-run-usage.ts b/apps/sim/lib/billing/core/mid-run-usage.ts new file mode 100644 index 00000000000..36ea034f81e --- /dev/null +++ b/apps/sim/lib/billing/core/mid-run-usage.ts @@ -0,0 +1,232 @@ +import { createLogger } from '@sim/logger' +import { getErrorMessage } from '@sim/utils/errors' +import { LRUCache } from 'lru-cache' +import { checkUsageStatus } from '@/lib/billing/calculations/usage-monitor' +import { getOrganizationSubscription } from '@/lib/billing/core/billing' +import { + type AccountBillingDecision, + type AttributedUsageLimitsResult, + type BillingAttributionSnapshot, + checkAccountBillingBlocks, + refreshAttributionPeriod, +} from '@/lib/billing/core/billing-attribution' +import { defaultBillingPeriod } from '@/lib/billing/core/billing-period' +import { getHighestPriorityPersonalSubscription } from '@/lib/billing/core/plan' +import { resolveSubscriptionUsagePeriod } from '@/lib/billing/core/reporting-period' +import { + checkExecutionUsageLimits, + USAGE_GATE_SETTLE_TIMEOUT_MS, + USAGE_GATE_TTL_MS, +} from '@/lib/billing/core/usage-gate-cache' +import { coalesceLocally } from '@/lib/concurrency/singleflight' +import { isBillingEnabled, isHosted } from '@/lib/core/config/env-flags' + +const logger = createLogger('MidRunUsage') + +/** + * A run's standing while it is under way, read through the execution usage gate: + * - `exceeded`: the payer (or the actor's member cap) spent its limit; the run pauses with the + * upgrade card. + * - `blocked`: the account is blocked (payment failed, dispute); the run is refused as a blocked + * account, never with the upgrade card. + * - `unknown`: usage, or the payer's current period, could not be read. Admission fails closed + * on an unreadable ledger, but a run already under way continues: a database blip must not + * end a paying user's long run, and the next step or re-check reads again. + */ +export type MidRunUsageVerdict = + | { status: 'within' } + | { status: 'exceeded'; scope?: AttributedUsageLimitsResult['scope'] } + | { status: 'blocked'; message?: string } + | { status: 'unknown' } + +/** + * How long a payer's current period stays cached for mid-run checks. Every model step settles a + * cost callback that reads it; a subscription period change (rollover, anchor reset) reaches the + * check within this long. + */ +const CURRENT_PERIOD_TTL_MS = 60 * 1000 + +const currentPeriodCache = new LRUCache({ + max: 10_000, + ttl: CURRENT_PERIOD_TTL_MS, +}) + +function currentPeriodKey(attribution: BillingAttributionSnapshot): string { + return [ + attribution.actorUserId, + attribution.workspaceId ?? '', + attribution.organizationId ?? '', + attribution.billedAccountUserId, + ].join(':') +} + +/** The admitted payer's attribution for its current subscription period. */ +async function currentAttribution( + attribution: BillingAttributionSnapshot +): Promise { + const key = currentPeriodKey(attribution) + const cached = currentPeriodCache.get(key) + // A cached period that has since ended is stale: the payer may already be in the next one. + if (cached && !periodHasEnded(cached)) return cached + const current = await refreshAttributionPeriod(attribution) + currentPeriodCache.set(key, current) + return current +} + +/** Mirrors the cost callback's rollover gate: only a Stripe period rolls forward. */ +function rollsIntoCurrentPeriod(period: { source?: string }): boolean { + return period.source === 'stripe' +} + +function periodHasEnded(attribution: BillingAttributionSnapshot): boolean { + return Date.now() >= new Date(attribution.billingPeriod.end).getTime() +} + +async function readGateVerdict( + attribution: BillingAttributionSnapshot +): Promise { + let usage: AttributedUsageLimitsResult + try { + usage = await checkExecutionUsageLimits(attribution) + } catch (error) { + logger.warn('Mid-run usage read failed; continuing the run', { + error: getErrorMessage(error), + }) + return { status: 'unknown' } + } + if (!usage.isExceeded) return { status: 'within' } + if (usage.reason === 'billing_blocked') { + return { status: 'blocked', ...(usage.message ? { message: usage.message } : {}) } + } + if (usage.reason === 'usage_unavailable') { + logger.warn('Mid-run usage could not be read; continuing the run') + return { status: 'unknown' } + } + return { status: 'exceeded', ...(usage.scope ? { scope: usage.scope } : {}) } +} + +/** + * Judges a run against the period its charges land in. A Stripe-period payer's charges roll into + * whatever period the subscription is in now (a rollover or an early anchor reset included), so + * such a run is judged against the payer's CURRENT period. A current period that cannot be read, + * or that ends before its verdict is read, makes the verdict unknown, so the run continues and the + * next callback judges the next period. Any other payer's charges stay in the admitted + * period (a reporting window, or the open default one), so that period is judged, even after it + * ends. + */ +export async function readMidRunUsageVerdict( + attribution: BillingAttributionSnapshot +): Promise { + if (!isHosted || !isBillingEnabled) return { status: 'within' } + if (!rollsIntoCurrentPeriod(attribution.billingPeriod)) return readGateVerdict(attribution) + let judged: BillingAttributionSnapshot + try { + judged = await currentAttribution(attribution) + } catch (error) { + logger.warn('Current billing period could not be read; continuing the run', { + error: getErrorMessage(error), + }) + return { status: 'unknown' } + } + if (periodHasEnded(judged)) return { status: 'unknown' } + const verdict = await readGateVerdict(judged) + // A period that ended during the read is judged at the next callback, which reloads it. + return periodHasEnded(judged) ? { status: 'unknown' } : verdict +} + +/** + * Admitted direct-v1 verdicts, served for the execution gate's TTL like + * {@link checkExecutionUsageLimits} serves attributed ones: the worker re-validates a run on + * every resume leg, and each uncached read sums the payer's ledger for the period. Only a + * `within` verdict is stored, so a refusal or an unreadable ledger is always read again. + */ +const accountVerdictCache = new LRUCache({ + max: 10_000, + ttl: USAGE_GATE_TTL_MS, +}) + +/** + * The same verdict for a direct-v1 run billed to an account decision rather than an attributed + * payer, in the gate's order: a blocked actor or payer first, then the payer's spend. The payer + * is the one saved in the decision at admission, never re-selected from the actor's current + * memberships, and the period judged is the one its charges land in, as for attributed runs. + */ +export async function readMidRunAccountUsageVerdict( + decision: AccountBillingDecision +): Promise { + if (!isHosted || !isBillingEnabled) return { status: 'within' } + try { + const payer = decision.billingEntity + const subscription = + payer.type === 'organization' + ? await getOrganizationSubscription(payer.id, { onError: 'throw' }) + : await getHighestPriorityPersonalSubscription(payer.id, { onError: 'throw' }) + const billingPeriod = rollsIntoCurrentPeriod(decision.billingPeriod) + ? (resolveSubscriptionUsagePeriod(subscription) ?? { + ...defaultBillingPeriod(), + source: 'default' as const, + }) + : { + start: new Date(decision.billingPeriod.start), + end: new Date(decision.billingPeriod.end), + source: decision.billingPeriod.source ?? ('default' as const), + } + const key = [ + payer.type, + payer.id, + billingPeriod.start.toISOString(), + billingPeriod.end.toISOString(), + billingPeriod.source, + decision.userId, + subscription?.id ?? '', + subscription?.plan ?? '', + subscription?.status ?? '', + subscription?.seats ?? '', + ].join(':') + const cached = accountVerdictCache.get(key) + if (cached) return cached + const block = await checkAccountBillingBlocks(decision) + if (block.blocked) { + return { status: 'blocked', ...(block.message ? { message: block.message } : {}) } + } + // An organization payer without a subscription stays organization-scoped on the free plan, + // as `toUsageLimitSubscription` does for attributed runs, never the actor's personal ledger. + const usageSubscription = + subscription ?? + (payer.type === 'organization' + ? { + referenceId: payer.id, + plan: 'free', + status: null, + seats: null, + periodStart: billingPeriod.start, + periodEnd: billingPeriod.end, + } + : null) + const usage = await coalesceLocally( + `mid-run-account-usage:${key}`, + () => + checkUsageStatus(decision.userId, usageSubscription, { + billingEntity: payer, + billingPeriod, + }), + USAGE_GATE_SETTLE_TIMEOUT_MS + ) + if (usage.unavailable) return { status: 'unknown' } + if (usage.isExceeded) return { status: 'exceeded', scope: 'payer' } + const within: MidRunUsageVerdict = { status: 'within' } + accountVerdictCache.set(key, within) + return within + } catch (error) { + logger.warn('Mid-run account usage read failed; continuing the run', { + error: getErrorMessage(error), + }) + return { status: 'unknown' } + } +} + +/** Drops every cached current period and account verdict. Test seam; never called in production code. */ +export function resetMidRunUsageCaches(): void { + currentPeriodCache.clear() + accountVerdictCache.clear() +} diff --git a/apps/sim/lib/billing/core/usage-analytics.ts b/apps/sim/lib/billing/core/usage-analytics.ts index 92e6b50f9c7..60d5ff1a1e9 100644 --- a/apps/sim/lib/billing/core/usage-analytics.ts +++ b/apps/sim/lib/billing/core/usage-analytics.ts @@ -513,6 +513,10 @@ export function usageBucketTimestamps( * cost in place for as long as its stream runs — which {@link STREAM_TIMEOUT_MS} * caps — plus the retry flushes that follow it. Past the cap and this margin a day or * hour can no longer change and is treated as settled. + * + * Without a run deadline a Chat turn can top up its row for longer than that, so a + * settled hour's cached aggregate can under-report that turn's later spend. This is + * display only: invoices, threshold billing, and the usage gate read live ledger sums. */ export const USAGE_SETTLE_MS = STREAM_TIMEOUT_MS + 2 * 60 * 60 * 1000 diff --git a/apps/sim/lib/billing/core/usage-log.integration.ts b/apps/sim/lib/billing/core/usage-log.integration.ts index 7b8da8fb23f..e03401fb602 100644 --- a/apps/sim/lib/billing/core/usage-log.integration.ts +++ b/apps/sim/lib/billing/core/usage-log.integration.ts @@ -10,7 +10,7 @@ import * as schema from '@sim/db/schema' import { readTestDatabaseUrl } from '@sim/db/testing/test-infrastructure' import { getPostgresErrorCode } from '@sim/utils/errors' import { generateId } from '@sim/utils/id' -import { sql } from 'drizzle-orm' +import { eq, sql } from 'drizzle-orm' import { drizzle } from 'drizzle-orm/postgres-js' import postgres from 'postgres' import { afterAll, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest' @@ -26,6 +26,7 @@ import { CumulativeUsageContextMismatchError, getBillingPeriodUsageCost, getBillingPeriodUsageCostByUser, + getStampedPeriodRangeUsageCostByUser, type RecordCumulativeUsageParams, recordCumulativeUsage, } from '@/lib/billing/core/usage-log' @@ -119,7 +120,8 @@ 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 driver_probe (id text PRIMARY KEY); + CREATE TABLE subscription (id text PRIMARY KEY, period_start timestamp, period_end timestamp) `) transaction.mockImplementation(async (callback: (tx: Transaction) => Promise) => { const pause = nextPause @@ -142,6 +144,7 @@ describe('Cumulative billing with PostgreSQL', () => { beforeEach(async () => { nextPause = undefined await connection`truncate usage_log` + await connection`truncate subscription` }) it.each([ @@ -216,12 +219,12 @@ describe('Cumulative billing with PostgreSQL', () => { expect(recovered.billed).toBe(true) expect(recovered.delta).toBeCloseTo(0.8 - initial, 9) expect(recovered.total).toBe(0.8) - expect(await recordCumulativeUsage(usage(0.8))).toEqual({ + expect(await recordCumulativeUsage(usage(0.8))).toMatchObject({ billed: false, delta: 0, total: 0.8, }) - expect(await recordCumulativeUsage(usage(0.3))).toEqual({ + expect(await recordCumulativeUsage(usage(0.3))).toMatchObject({ billed: false, delta: 0, total: 0.8, @@ -245,7 +248,11 @@ describe('Cumulative billing with PostgreSQL', () => { ) ) expect(await ledgerRows()).toHaveLength(33) - expect(await recordCumulativeUsage(usage(0.8))).toEqual({ billed: false, delta: 0, total: 0.8 }) + expect(await recordCumulativeUsage(usage(0.8))).toMatchObject({ + billed: false, + delta: 0, + total: 0.8, + }) }) it('reads committed pooled and member charges freshly after concurrent executions', async () => { @@ -305,4 +312,217 @@ describe('Cumulative billing with PostgreSQL', () => { expect(await ledgerRows()).toEqual([{ event_key: usage(0).eventKey, cost: '0.8' }]) } ) + + it("counts a reporting run's top-ups after its window ends in that window, and a later run's charges in the next", async () => { + const payer = { type: 'organization', id: 'payer' } as const + const dayMs = 24 * 60 * 60 * 1000 + // A reporting window is summed by when each row was created; the stamped period only binds a + // request's rows to each other. + const stamp = { + start: new Date('2026-01-01'), + end: new Date('2027-01-01'), + source: 'reporting' as const, + } + await recordCumulativeUsage({ ...usage(0.4, 'update-cost:long-run'), billingPeriod: stamp }) + const [first] = await database + .select({ createdAt: schema.usageLog.createdAt }) + .from(schema.usageLog) + .where(eq(schema.usageLog.eventKey, 'update-cost:long-run')) + // The admitted window ends right after the run's first charge, and every later write starts + // once the database clock has passed that boundary. + const boundary = new Date(first.createdAt.getTime() + 1) + for (;;) { + const [{ passed }] = await connection<{ passed: boolean }[]>` + select clock_timestamp()::timestamp > created_at + interval '1 millisecond' as passed + from usage_log where event_key = 'update-cost:long-run' + ` + if (passed) break + } + + await recordCumulativeUsage({ ...usage(1, 'update-cost:long-run'), billingPeriod: stamp }) + await recordCumulativeUsage({ ...usage(0.25, 'update-cost:next-run'), billingPeriod: stamp }) + + const windowTotal = (start: Date, end: Date) => + getBillingPeriodUsageCost(payer, { start, end, source: 'reporting' }, undefined, database) + expect(await windowTotal(new Date(boundary.getTime() - 30 * dayMs), boundary)).toBeCloseTo(1, 9) + expect(await windowTotal(boundary, new Date(boundary.getTime() + 30 * dayMs))).toBeCloseTo( + 0.25, + 9 + ) + }) + + describe('a request that outlives its billing period', () => { + // Past periods: the old period's row is written under the subscription lock only once + // that period has ended. + const periods = [ + new Date('2025-09-01T00:00:00.000Z'), + new Date('2025-10-01T00:00:00.000Z'), + new Date('2025-11-01T00:00:00.000Z'), + new Date('2025-12-01T00:00:00.000Z'), + ] + const payer = { type: 'organization', id: 'payer' } as const + + 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') + on conflict (id) do update + set period_start = excluded.period_start, period_end = excluded.period_end + ` + } + + async function setSubscriptionPeriod(index: number) { + await setSubscriptionWindow(periods[index], periods[index + 1]) + } + + function charge(cost: number, frozen = { start: periods[0], end: periods[1] }) { + return recordCumulativeUsage({ + ...usage(cost), + billingPeriod: frozen, + payerSubscriptionId: 'sub-1', + }) + } + + /** What the cycle close invoices for one period: the ledger rows stamped with it. */ + async function stampedWindowTotal(from: Date, to: Date) { + const byUser = await getStampedPeriodRangeUsageCostByUser( + payer, + { from, to }, + undefined, + database + ) + return [...byUser.values()].reduce((total, cost) => total + cost, 0) + } + + function stampedTotal(index: number) { + return stampedWindowTotal(periods[index], periods[index + 1]) + } + + it('invoices a charge that spans a period close exactly once in total', async () => { + await setSubscriptionPeriod(0) + expect(await charge(0.4)).toMatchObject({ billed: true, total: 0.4 }) + + await setSubscriptionPeriod(1) + const closedTotal = await stampedTotal(0) + expect(closedTotal).toBeCloseTo(0.4, 9) + + const afterClose = await charge(1) + expect(afterClose).toMatchObject({ billed: true, total: 1 }) + expect(afterClose.billingPeriod).toEqual({ start: periods[1], end: periods[2] }) + expect(await charge(0.9)).toMatchObject({ billed: false, total: 1 }) + expect(await charge(1.3)).toMatchObject({ billed: true, total: 1.3 }) + expect(await charge(1.3)).toMatchObject({ billed: false, total: 1.3 }) + + await setSubscriptionPeriod(2) + expect(await charge(1.5)).toMatchObject({ billed: true, total: 1.5 }) + + expect(await stampedTotal(0)).toBeCloseTo(closedTotal, 9) + expect(await stampedTotal(1)).toBeCloseTo(0.9, 9) + expect(await stampedTotal(2)).toBeCloseTo(0.2, 9) + const invoiced = (await stampedTotal(0)) + (await stampedTotal(1)) + (await stampedTotal(2)) + expect(invoiced).toBeCloseTo(1.5, 9) + }) + + it('gives each period row only the tokens spent after the rows before it', async () => { + await setSubscriptionPeriod(0) + await recordCumulativeUsage({ + ...usage(0.4), + billingPeriod: { start: periods[0], end: periods[1] }, + payerSubscriptionId: 'sub-1', + }) + await setSubscriptionPeriod(1) + await recordCumulativeUsage({ + ...usage(1), + billingPeriod: { start: periods[0], end: periods[1] }, + payerSubscriptionId: 'sub-1', + metadata: { inputTokens: 25, outputTokens: 12 }, + }) + + const rows = await connection<{ event_key: string; metadata: Record }[]>` + select event_key, metadata from usage_log order by event_key + ` + expect(rows.map((row) => [row.event_key, row.metadata])).toEqual([ + ['update-cost:shared-request', { inputTokens: 10, outputTokens: 5 }], + ['update-cost:shared-request@1', { inputTokens: 15, outputTokens: 7 }], + ]) + }) + + it('never stamps a charge into a period earlier than its latest row', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + await setSubscriptionPeriod(1) + await charge(1) + await setSubscriptionPeriod(0) + expect(await charge(1.2)).toMatchObject({ billed: true, total: 1.2 }) + expect(await stampedTotal(0)).toBeCloseTo(0.4, 9) + expect(await stampedTotal(1)).toBeCloseTo(0.8, 9) + }) + + it('stamps a first charge that lands after the close into the current period', async () => { + await setSubscriptionPeriod(1) + expect(await charge(0.7)).toMatchObject({ billed: true, total: 0.7 }) + expect(await charge(0.9)).toMatchObject({ billed: true, total: 0.9 }) + expect(await stampedTotal(0)).toBe(0) + expect(await stampedTotal(1)).toBeCloseTo(0.9, 9) + }) + + it('holds the period advance until an in-flight top-up of the old period commits', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + const pause = pauseNextTransaction() + const inFlight = charge(0.6) + try { + await pause.reached.promise + const advance = await connection + .begin(async (tx) => { + await tx`select set_config('lock_timeout', '300ms', true)` + await tx`update subscription set period_start = ${periods[1].toISOString()}::timestamptz at time zone 'UTC' where id = 'sub-1'` + }) + .catch((error: unknown) => error) + expect(getPostgresErrorCode(advance)).toBe('55P03') + } finally { + pause.release.resolve() + await inFlight + } + expect(await stampedTotal(0)).toBeCloseTo(0.6, 9) + }) + + it('rolls into a period whose start moved forward before the old period ended', async () => { + await setSubscriptionPeriod(0) + await charge(0.4) + const resetStart = new Date('2025-09-15T00:00:00.000Z') + const resetEnd = new Date('2025-10-15T00:00:00.000Z') + await setSubscriptionWindow(resetStart, resetEnd) + + expect(await charge(1)).toMatchObject({ + billed: true, + billingPeriod: { start: resetStart, end: resetEnd }, + }) + expect(await stampedTotal(0)).toBeCloseTo(0.4, 9) + expect(await stampedWindowTotal(resetStart, resetEnd)).toBeCloseTo(0.6, 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) + await setSubscriptionWindow(start, end) + await charge(0.4, { start, end }) + const pause = pauseNextTransaction() + const inFlight = charge(0.6, { start, end }) + try { + await pause.reached.promise + const reset = await connection + .begin(async (tx) => { + await tx`select set_config('lock_timeout', '300ms', true)` + await tx`update subscription set period_start = now() at time zone 'UTC' where id = 'sub-1'` + }) + .catch((error: unknown) => error) + expect(getPostgresErrorCode(reset)).toBe('55P03') + } finally { + pause.release.resolve() + await inFlight + } + expect(await stampedWindowTotal(start, end)).toBeCloseTo(0.6, 9) + }) + }) }) diff --git a/apps/sim/lib/billing/core/usage-log.test.ts b/apps/sim/lib/billing/core/usage-log.test.ts index f3ccfba347e..3c87c10b9fa 100644 --- a/apps/sim/lib/billing/core/usage-log.test.ts +++ b/apps/sim/lib/billing/core/usage-log.test.ts @@ -280,7 +280,7 @@ describe('recordCumulativeUsage', () => { eventKey: 'update-cost:msg-1-billing', metadata: { inputTokens: 100, outputTokens: 5 }, }) - expect(result).toEqual({ billed: true, delta: 0.3474447, total: 0.3474447 }) + expect(result).toMatchObject({ billed: true, delta: 0.3474447, total: 0.3474447 }) expect(mockInsert).toHaveBeenCalledTimes(1) expect(mockUpdate).not.toHaveBeenCalled() expect(mockValues.mock.calls[0][0][0]).toMatchObject({ @@ -315,7 +315,7 @@ describe('recordCumulativeUsage', () => { cost: 0.4662453, eventKey: 'update-cost:msg-1-billing', }) - expect(result).toEqual({ billed: false, delta: 0, total: 0.4662453 }) + expect(result).toMatchObject({ billed: false, delta: 0, total: 0.4662453 }) expect(updateSet).not.toHaveBeenCalled() expect(mockInsert).not.toHaveBeenCalled() }) diff --git a/apps/sim/lib/billing/core/usage-log.ts b/apps/sim/lib/billing/core/usage-log.ts index 5a845acad67..40276626118 100644 --- a/apps/sim/lib/billing/core/usage-log.ts +++ b/apps/sim/lib/billing/core/usage-log.ts @@ -1,9 +1,11 @@ import { createHash } from 'node:crypto' import { db, dbReplica } from '@sim/db' -import { usageLog, workflow } from '@sim/db/schema' +import { subscription as subscriptionTable, usageLog, workflow } from '@sim/db/schema' import { createLogger } from '@sim/logger' +import { toNumberOrNull } from '@sim/utils/coerce' import { getPostgresErrorCode, toError } from '@sim/utils/errors' import { generateId } from '@sim/utils/id' +import { toRecordOrNull } from '@sim/utils/object' import { and, desc, eq, gte, inArray, lt, lte, notInArray, or, sql } from 'drizzle-orm' import { type CursorKey, @@ -583,6 +585,21 @@ export interface RecordCumulativeUsageParams { /** Stable per-request key; the single ledger row is keyed on this. */ eventKey: string metadata?: UsageLogMetadata + /** + * The Stripe-period subscription that pays for this request. When given, a top-up that + * 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. + * + * 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 + * has period rows, it re-adds those rows' amount to the first row. That only double-counts + * when it lands between the rollover and that period's close, which waits at least an hour, + * and only for runs spanning a rollover; the exposure is one run's post-rollover spend, cents + * to dollars. + */ + payerSubscriptionId?: string } export interface RecordCumulativeUsageResult { @@ -592,6 +609,57 @@ export interface RecordCumulativeUsageResult { delta: number /** The request's recorded cumulative cost after this flush. */ total: number + /** The billing period of the row this flush wrote to, or of the request's latest row. */ + billingPeriod: { start: Date; end: Date } +} + +/** + * The most period rows one request may span: its first row plus one per later billing period. + * A request still billing twelve periods after it started is refused rather than scanned. + */ +const MAX_CUMULATIVE_PERIOD_ROWS = 12 + +/** + * The ledger key of the `index`-th period a cumulative request rolled into; 0 is the request key. + * The cost callback refuses a request key containing `@`, so these never collide with another + * request's. + */ +function cumulativePeriodEventKey(eventKey: string, index: number): string { + return index === 0 ? eventKey : `${eventKey}@${index}` +} + +/** Decimal places kept when period row costs are summed or subtracted as floats. */ +const PERIOD_COST_DECIMALS = 12 + +function sumLedgerCost(rows: readonly { cost: string }[]): number { + if (rows.length <= 1) return rows[0] ? Number.parseFloat(rows[0].cost) : 0 + const total = rows.reduce((sum, row) => sum + Number.parseFloat(row.cost), 0) + return Number(total.toFixed(PERIOD_COST_DECIMALS)) +} + +const CUMULATIVE_TOKEN_FIELDS = ['inputTokens', 'outputTokens'] as const + +/** + * A period row's share of a cumulative callback's token counts: the cumulative counts minus what + * the request's other rows already hold, so summing the rows never counts a token twice. + */ +function periodUsageMetadata( + metadata: UsageLogMetadata | undefined, + otherRows: readonly { metadata: unknown }[] +): UsageLogMetadata | undefined { + const cumulative = toRecordOrNull(metadata) + if (!cumulative || otherRows.length === 0) return metadata + const share: Record = { ...cumulative } + for (const field of CUMULATIVE_TOKEN_FIELDS) { + const total = toNumberOrNull(cumulative[field]) + if (total === null) continue + const recorded = otherRows.reduce( + (sum, row) => sum + (toNumberOrNull(toRecordOrNull(row.metadata)?.[field]) ?? 0), + 0 + ) + share[field] = Math.max(0, total - recorded) + } + return share } export type CumulativeUsageContextField = @@ -628,6 +696,8 @@ function assertCumulativeUsageLedgerBinding( workspaceId?: string billingContext: BillingContext eventKey: string + /** A request whose first charge landed after its period closed is stamped with a later one. */ + allowLaterPeriod?: boolean } ): void { const mismatchedFields: CumulativeUsageContextField[] = [] @@ -643,11 +713,15 @@ function assertCumulativeUsageLedgerBinding( ) { mismatchedFields.push('billing entity') } - if ( - existing.billingPeriodStart?.getTime() !== - expected.billingContext.billingPeriod.start.getTime() || - existing.billingPeriodEnd?.getTime() !== expected.billingContext.billingPeriod.end.getTime() - ) { + const frozenPeriod = expected.billingContext.billingPeriod + const samePeriod = + existing.billingPeriodStart?.getTime() === frozenPeriod.start.getTime() && + existing.billingPeriodEnd?.getTime() === frozenPeriod.end.getTime() + const laterPeriod = + expected.allowLaterPeriod === true && + existing.billingPeriodStart !== null && + existing.billingPeriodStart.getTime() >= frozenPeriod.end.getTime() + if (!samePeriod && !laterPeriod) { mismatchedFields.push('billing period') } @@ -703,6 +777,7 @@ export async function recordCumulativeUsage( cost, eventKey, metadata, + payerSubscriptionId, } = params if (workspaceId && (!billingEntity || !billingPeriod)) { @@ -744,10 +819,12 @@ export async function recordCumulativeUsage( await acquireAdvisoryXactLock(tx, 'usage_log_event', eventKey) enterStage('read') - const [existing] = await tx + const rows = await tx .select({ id: usageLog.id, + eventKey: usageLog.eventKey, cost: usageLog.cost, + metadata: usageLog.metadata, userId: usageLog.userId, workspaceId: usageLog.workspaceId, billingEntityType: usageLog.billingEntityType, @@ -756,55 +833,125 @@ export async function recordCumulativeUsage( billingPeriodEnd: usageLog.billingPeriodEnd, }) .from(usageLog) - .where(eq(usageLog.eventKey, eventKey)) - .limit(1) - - if (existing) { - assertCumulativeUsageLedgerBinding(existing, { + .where( + payerSubscriptionId + ? inArray( + usageLog.eventKey, + Array.from({ length: MAX_CUMULATIVE_PERIOD_ROWS }, (_, index) => + cumulativePeriodEventKey(eventKey, index) + ) + ) + : eq(usageLog.eventKey, eventKey) + ) + .limit(MAX_CUMULATIVE_PERIOD_ROWS) + + // Period rows are written in order under this lock, so they are the keys 0..n-1. + const chain = payerSubscriptionId + ? Array.from({ length: rows.length }, (_, index) => + rows.find((row) => row.eventKey === cumulativePeriodEventKey(eventKey, index)) + ).filter((row) => row !== undefined) + : rows.slice(0, 1) + if (payerSubscriptionId && chain.length !== rows.length) { + throw new Error(`Cumulative usage event "${eventKey}" has a gap in its period rows`) + } + const [anchor] = chain + if (anchor) { + assertCumulativeUsageLedgerBinding(anchor, { userId, workspaceId, billingContext, eventKey, + allowLaterPeriod: Boolean(payerSubscriptionId), }) } - const recorded = existing ? Number.parseFloat(existing.cost) : 0 + const latest = chain.at(-1) + const latestPeriod = + latest?.billingPeriodStart && latest.billingPeriodEnd + ? { start: latest.billingPeriodStart, end: latest.billingPeriodEnd } + : billingContext.billingPeriod + const recorded = sumLedgerCost(chain) const { shouldBill, delta, newTotal } = resolveCumulativeTopUp(recorded, cost) if (!shouldBill) { enterStage('commit') - return { billed: false, delta: 0, total: recorded } + 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. + const [currentPeriod] = payerSubscriptionId + ? await tx + .select({ + start: subscriptionTable.periodStart, + end: subscriptionTable.periodEnd, + }) + .from(subscriptionTable) + .where(eq(subscriptionTable.id, payerSubscriptionId)) + .for('share') + .limit(1) + : [] + + // Only ever forward: a subscription period that does not start after the latest row's + // keeps topping up that row, whatever the wall clock or a replayed webhook says. A start + // that moved forward inside the old period (anchor reset, resync) still rolls, so the old + // period's close is never topped up after the fact. + const rolledPeriod = + currentPeriod?.start && + currentPeriod.end && + currentPeriod.start.getTime() > latestPeriod.start.getTime() + ? { start: currentPeriod.start, end: currentPeriod.end } + : null + if (rolledPeriod && latest && chain.length >= MAX_CUMULATIVE_PERIOD_ROWS) { + throw new Error(`Cumulative usage event "${eventKey}" spans too many billing periods`) } enterStage('write') - if (existing) { + if (latest && !rolledPeriod) { + const otherRows = chain.slice(0, -1) + const latestCost = + otherRows.length === 0 + ? newTotal + : Number((newTotal - sumLedgerCost(otherRows)).toFixed(PERIOD_COST_DECIMALS)) await tx .update(usageLog) - .set({ cost: newTotal.toString(), metadata: metadata ?? null }) - .where(eq(usageLog.id, existing.id)) - } else { - await recordUsage({ - userId, - workspaceId, - tx, - billingEntity: billingContext.billingEntity, - billingPeriod: billingContext.billingPeriod, - entries: [ - { - category: 'model', - source, - description: model, - cost: newTotal, - eventKey, - sourceReference: eventKey, - ...(metadata ? { metadata } : {}), - }, - ], - }) + .set({ + cost: latestCost.toString(), + metadata: periodUsageMetadata(metadata, otherRows) ?? null, + }) + .where(eq(usageLog.id, latest.id)) + enterStage('commit') + return { billed: true, delta, total: newTotal, billingPeriod: latestPeriod } } + const targetPeriod = rolledPeriod ?? billingContext.billingPeriod + const rowMetadata = periodUsageMetadata(metadata, chain) + await recordUsage({ + userId, + workspaceId, + tx, + billingEntity: billingContext.billingEntity, + billingPeriod: targetPeriod, + entries: [ + { + category: 'model', + source, + description: model, + cost: chain.length === 0 ? newTotal : Number(delta.toFixed(PERIOD_COST_DECIMALS)), + eventKey: cumulativePeriodEventKey(eventKey, chain.length), + sourceReference: eventKey, + ...(rowMetadata ? { metadata: rowMetadata } : {}), + }, + ], + }) enterStage('commit') - return { billed: true, delta, total: newTotal } + return { + billed: true, + delta, + total: newTotal, + billingPeriod: { start: targetPeriod.start, end: targetPeriod.end }, + } }) succeeded = true return result diff --git a/apps/sim/lib/billing/usage-upgrade.ts b/apps/sim/lib/billing/usage-upgrade.ts new file mode 100644 index 00000000000..e72713a8d8c --- /dev/null +++ b/apps/sim/lib/billing/usage-upgrade.ts @@ -0,0 +1,67 @@ +import { createLogger } from '@sim/logger' +import { getErrorMessage } from '@sim/utils/errors' +import type { UsageUpgradePayload } from '@/lib/api/contracts/subscription' +import type { + AttributedUsageLimitsResult, + BillingAttributionSnapshot, +} from '@/lib/billing/core/billing-attribution' +import { getHighestPrioritySubscription } from '@/lib/billing/core/plan' +import { isEnterprise, isPaid } from '@/lib/billing/plan-helpers' +import { isOrgScopedSubscription } from '@/lib/billing/subscriptions/utils' + +const logger = createLogger('UsageUpgrade') + +const UPGRADE_PLAN_MESSAGE = + "You've reached your usage limit. Please upgrade your plan to continue." + +const MEMBER_CAP_MESSAGE = + "You've reached the usage limit your organization set for you this billing period. Only an organization owner or admin can raise it — please ask them to continue." + +/** + * The upgrade card for a payer over its usage limit: a plan upgrade for a free payer, a limit + * increase for a paid one, with copy naming who can raise an organization's limit. A member + * over the cap their organization set gets copy naming who can raise that cap. An attributed + * run reads the plan from its admission snapshot without a query; otherwise the actor's current + * subscription decides, and a failed lookup falls back to the plan-upgrade card. + */ +export async function resolveUsageUpgradePayload( + userId: string, + billingAttribution?: BillingAttributionSnapshot, + scope?: AttributedUsageLimitsResult['scope'] +): Promise { + if (scope === 'member') { + return { reason: 'usage_limit', action: 'increase_limit', message: MEMBER_CAP_MESSAGE } + } + let plan: string | undefined + let orgScoped = false + try { + if (billingAttribution) { + plan = billingAttribution.payerSubscription?.plan + orgScoped = billingAttribution.billingEntity.type === 'organization' + } else { + const subscription = await getHighestPrioritySubscription(userId) + plan = subscription?.plan + orgScoped = isOrgScopedSubscription(subscription, userId) + } + } catch (error) { + logger.warn('Failed to determine subscription plan, defaulting to upgrade_plan', { + error: getErrorMessage(error), + }) + } + + if (!plan || !isPaid(plan)) { + return { reason: 'usage_limit', action: 'upgrade_plan', message: UPGRADE_PLAN_MESSAGE } + } + // Paid plans get `increase_limit`; the copy says who can raise it when the user cannot. + const message = !orgScoped + ? "You've reached your usage limit for this billing period. Please increase your usage limit from billing settings to continue." + : isEnterprise(plan) + ? "You've reached your organization's usage limit for this billing period. Only an organization admin or Sim support can raise an enterprise limit — reach out to them to continue." + : "You've reached your organization's usage limit for this billing period. Only an organization owner or admin can raise the limit — please ask them to update it from the team billing settings." + return { reason: 'usage_limit', action: 'increase_limit', message } +} + +/** The assistant text that renders {@link payload} as the usage card. */ +export function formatUsageUpgradeTag(payload: UsageUpgradePayload): string { + return `${JSON.stringify(payload)}` +} diff --git a/apps/sim/lib/mothership/application/authorize-chat-callback.test.ts b/apps/sim/lib/mothership/application/authorize-chat-callback.test.ts index 798b96f37e3..2b84d1e20ab 100644 --- a/apps/sim/lib/mothership/application/authorize-chat-callback.test.ts +++ b/apps/sim/lib/mothership/application/authorize-chat-callback.test.ts @@ -2,10 +2,6 @@ import { billingAttributionMock, billingAttributionMockFns, } from '@sim/testing/mocks/billing-attribution.mock' -import { - billingUsageMonitorMock, - billingUsageMonitorMockFns, -} from '@sim/testing/mocks/billing-usage-monitor.mock' import { mothershipOrganizationChatsMock, mothershipOrganizationChatsMockFns, @@ -37,16 +33,14 @@ vi.mock('@/lib/permission-groups/capability-assertions', async (importOriginal) })) vi.mock('@/lib/mothership/chat/organization-chats', () => mothershipOrganizationChatsMock) vi.mock('@/lib/billing/core/billing-attribution', () => billingAttributionMock) +const mockCheckAccountBillingBlocks = billingAttributionMockFns.mockCheckAccountBillingBlocks const mockCheckAttributedBillingBlocks = billingAttributionMockFns.mockCheckAttributedBillingBlocks const mocks = { ...hoisted, organization: mothershipOrganizationChatsMockFns.mockAuthorizeOrganizationChatDelegation, - actorBlock: billingUsageMonitorMockFns.mockCheckBillingBlocked, - payerBlock: billingUsageMonitorMockFns.mockCheckBillingEntityBlocked, loadWorkspace: workspaceContextMockFns.mockResolveActiveWorkspaceApplicationContext, permission: workspaceAuthzMockFns.mockResolveEffectiveWorkspacePermission, } -vi.mock('@/lib/billing/calculations/usage-monitor', () => billingUsageMonitorMock) const context = { userId: 'actor', @@ -80,9 +74,8 @@ beforeEach(() => { }) mocks.permission.mockResolvedValue('read') mocks.organization.mockResolvedValue(undefined) - mocks.actorBlock.mockResolvedValue({ blocked: false }) - mocks.payerBlock.mockResolvedValue({ blocked: false }) mockCheckAttributedBillingBlocks.mockResolvedValue({ blocked: false }) + mockCheckAccountBillingBlocks.mockResolvedValue({ blocked: false }) }) describe('fresh chat callback authorization', () => { @@ -193,40 +186,21 @@ describe('fresh chat callback authorization', () => { }) describe('continuation account standing', () => { - it('uses the existing attributed block policy with the original snapshot', async () => { - await checkCopilotContinuationBilling({ kind: 'attributed', attribution }) - expect(mockCheckAttributedBillingBlocks).toHaveBeenCalledWith(attribution) - expect(mocks.actorBlock).not.toHaveBeenCalled() - expect(mocks.payerBlock).not.toHaveBeenCalled() - }) + it('judges each run kind by its own block policy and original billing material', async () => { + mockCheckAttributedBillingBlocks.mockImplementation(async (value: unknown) => ({ + blocked: value === attribution, + scope: 'payer', + })) + mockCheckAccountBillingBlocks.mockImplementation(async (value: unknown) => ({ + blocked: value === account, + scope: 'actor', + })) - it('checks both actor and the exact original direct-account payer', async () => { - await checkCopilotContinuationBilling({ kind: 'account', decision: account }) - expect(mocks.actorBlock).toHaveBeenCalledWith('actor') - expect(mocks.payerBlock).toHaveBeenCalledWith({ type: 'organization', id: 'original-payer' }) - }) - - it('refuses an actor block before reading the payer', async () => { - mocks.actorBlock.mockResolvedValueOnce({ blocked: true }) await expect( - checkCopilotContinuationBilling({ kind: 'account', decision: account }) - ).resolves.toMatchObject({ blocked: true, scope: 'actor' }) - expect(mocks.payerBlock).not.toHaveBeenCalled() - }) - - it('refuses a payer block independently of actor standing', async () => { - mocks.payerBlock.mockResolvedValueOnce({ blocked: true }) + checkCopilotContinuationBilling({ kind: 'attributed', attribution }) + ).resolves.toEqual({ blocked: true, scope: 'payer' }) await expect( checkCopilotContinuationBilling({ kind: 'account', decision: account }) - ).resolves.toMatchObject({ blocked: true, scope: 'payer' }) - }) - - it('reads the same personal actor/payer only once', async () => { - await checkCopilotContinuationBilling({ - kind: 'account', - decision: { ...account, billingEntity: { type: 'user', id: 'actor' } }, - }) - expect(mocks.actorBlock).toHaveBeenCalledTimes(1) - expect(mocks.payerBlock).not.toHaveBeenCalled() + ).resolves.toEqual({ blocked: true, scope: 'actor' }) }) }) diff --git a/apps/sim/lib/mothership/application/authorize-chat-callback.ts b/apps/sim/lib/mothership/application/authorize-chat-callback.ts index c8401eeb01f..327a1cabca2 100644 --- a/apps/sim/lib/mothership/application/authorize-chat-callback.ts +++ b/apps/sim/lib/mothership/application/authorize-chat-callback.ts @@ -1,11 +1,8 @@ import type { DelegatedPrincipal } from '@sim/auth/principal' -import { - checkBillingBlocked, - checkBillingEntityBlocked, -} from '@/lib/billing/calculations/usage-monitor' import { type AccountBillingDecision, type BillingAttributionSnapshot, + checkAccountBillingBlocks, checkAttributedBillingBlocks, } from '@/lib/billing/core/billing-attribution' import { defineAuthorizedWorkspaceUseCase } from '@/lib/core/application/authorized-workspace-use-case' @@ -98,11 +95,7 @@ export type CopilotContinuationBilling = /** Checks account standing against the original admission; never reads spend or selects a new payer. */ export async function checkCopilotContinuationBilling(billing: CopilotContinuationBilling) { - if (billing.kind === 'attributed') return checkAttributedBillingBlocks(billing.attribution) - - const actor = await checkBillingBlocked(billing.decision.userId) - if (actor.blocked) return { ...actor, scope: 'actor' } - const payer = billing.decision.billingEntity - if (payer.type === 'user' && payer.id === billing.decision.userId) return actor - return { ...(await checkBillingEntityBlocked(payer)), scope: 'payer' } + return billing.kind === 'attributed' + ? checkAttributedBillingBlocks(billing.attribution) + : checkAccountBillingBlocks(billing.decision) } diff --git a/apps/sim/lib/mothership/generated/billing.ts b/apps/sim/lib/mothership/generated/billing.ts index 12d4ba14b82..abdec1159ad 100644 --- a/apps/sim/lib/mothership/generated/billing.ts +++ b/apps/sim/lib/mothership/generated/billing.ts @@ -64,9 +64,30 @@ export const BillingCallbackHeaders = z context.addIssue({ code: "custom", message: "Incomplete or conflicting billing protocol headers" }); }); +/** Sim's plan-aware usage card: the JSON body of the `` tag its chat renders. */ +export const UsageUpgrade = z.object({ + reason: z.literal("usage_limit"), + action: z.enum(["upgrade_plan", "increase_limit"]), + message: z.string().min(1).max(1_000), +}); +export type UsageUpgrade = z.infer; + export const BillingCallbackResult = z.object({ success: z.boolean(), code: z.string().optional(), + /** The payer is over its plan usage limit after this charge. Absent (older Sim) means not over. */ + usageExceeded: z.boolean().optional(), + /** The card for an over-limit payer; a malformed card never turns a settled charge into a retry. */ + usageUpgrade: UsageUpgrade.optional().catch(undefined), +}); + +/** + * A continuation refused for the usage limit. A body-less 402 is a blocked account instead. + * The code decides; a malformed card falls back to the default one. + */ +export const UsageLimitRefusal = z.object({ + code: z.literal("USAGE_LIMIT_EXCEEDED"), + usageUpgrade: UsageUpgrade.optional().catch(undefined), }); export const BillingDuplicateCode = "DUPLICATE_BILLING_EVENT"; diff --git a/apps/sim/lib/mothership/request/go/stream.ts b/apps/sim/lib/mothership/request/go/stream.ts index 336423cbdd7..cd16ae492ad 100644 --- a/apps/sim/lib/mothership/request/go/stream.ts +++ b/apps/sim/lib/mothership/request/go/stream.ts @@ -126,7 +126,11 @@ function backendErrorMessage(status: number, body: string): string { } export class BillingLimitError extends Error { - constructor(public readonly userId: string) { + /** `member` when the actor hit the cap their organization set, so the card names who can raise it. */ + constructor( + public readonly userId: string, + public readonly scope?: 'actor' | 'payer' | 'member' + ) { super('Usage limit reached') this.name = 'BillingLimitError' } diff --git a/apps/sim/lib/mothership/request/lifecycle/admission.test.ts b/apps/sim/lib/mothership/request/lifecycle/admission.test.ts index 1e21943fc8b..d49e468b854 100644 --- a/apps/sim/lib/mothership/request/lifecycle/admission.test.ts +++ b/apps/sim/lib/mothership/request/lifecycle/admission.test.ts @@ -1,6 +1,14 @@ import { resetEnvFlagsMock, setEnvFlags } from '@sim/testing' +import { billingCoreMock, billingCoreMockFns } from '@sim/testing/mocks/billing-core.mock' +import { + billingUsageGateCacheMock, + billingUsageGateCacheMockFns, +} from '@sim/testing/mocks/billing-usage-gate-cache.mock' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import { createAttributedBillingRequestEnvelope } from '@/lib/billing/core/billing-attribution' +import { resetMidRunUsageCaches } from '@/lib/billing/core/mid-run-usage' +import { OrchestrationError } from '@/lib/core/orchestration/types' +import { BillingLimitError } from '@/lib/mothership/request/go/stream' import { authorizeLifecycleContinuation, restoreBillingAdmission } from './admission' const mocks = vi.hoisted(() => ({ authorize: vi.fn(), standing: vi.fn() })) @@ -8,13 +16,17 @@ vi.mock('@/lib/mothership/application/authorize-chat-callback', () => ({ authorizeCopilotChatCallback: mocks.authorize, checkCopilotContinuationBilling: mocks.standing, })) +vi.mock('@/lib/billing/core/usage-gate-cache', () => billingUsageGateCacheMock) +vi.mock('@/lib/billing/core/billing', () => billingCoreMock) +const { mockGetOrganizationSubscription } = billingCoreMockFns +const { mockCheckExecutionUsageLimits } = billingUsageGateCacheMockFns const attribution = { actorUserId: 'actor', workspaceId: 'workspace', organizationId: 'original-org', billedAccountUserId: 'original-owner', billingEntity: { type: 'organization' as const, id: 'original-org' }, - billingPeriod: { start: '2026-09-01T00:00:00.000Z', end: '2026-10-01T00:00:00.000Z' }, + billingPeriod: { start: '2026-09-01T00:00:00.000Z', end: '2099-01-01T00:00:00.000Z' }, payerSubscription: null, } const context = { @@ -25,8 +37,19 @@ const context = { billingAttribution: attribution, } beforeEach(() => { - setEnvFlags({ isHosted: true }) + setEnvFlags({ isHosted: true, isBillingEnabled: true }) mocks.standing.mockResolvedValue({ blocked: false }) + mockCheckExecutionUsageLimits.mockResolvedValue({ isExceeded: false }) + mockGetOrganizationSubscription.mockResolvedValue({ + id: 'sub-org', + referenceId: 'original-org', + plan: 'team', + status: 'active', + seats: 4, + periodStart: new Date(attribution.billingPeriod.start), + periodEnd: new Date(attribution.billingPeriod.end), + }) + resetMidRunUsageCaches() }) afterEach(resetEnvFlagsMock) @@ -69,4 +92,86 @@ describe('continuation admission', () => { authorizeLifecycleContinuation({ ...context, billingAttribution: undefined }) ).rejects.toThrow('missing') }) + it('refuses a continuation whose original payer has crossed its usage limit', async () => { + mockCheckExecutionUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'payer' }) + + const refusal = authorizeLifecycleContinuation(context) + + await expect(refusal).rejects.toBeInstanceOf(BillingLimitError) + await expect(refusal).rejects.toMatchObject({ userId: 'actor' }) + }) + it('keeps a blocked account a forbidden refusal without reading spend', async () => { + mocks.standing.mockResolvedValue({ blocked: true }) + + await expect(authorizeLifecycleContinuation(context)).rejects.toThrow('blocked') + expect(mockCheckExecutionUsageLimits).not.toHaveBeenCalled() + }) + it('does not read spend for a self-hosted continuation', async () => { + setEnvFlags({ isHosted: false }) + + await authorizeLifecycleContinuation(context) + expect(mockCheckExecutionUsageLimits).not.toHaveBeenCalled() + }) + it('continues a leg when spend cannot be read', async () => { + mockCheckExecutionUsageLimits.mockRejectedValueOnce(new Error('ledger read timed out')) + await expect(authorizeLifecycleContinuation(context)).resolves.toBeUndefined() + + mockCheckExecutionUsageLimits.mockResolvedValueOnce({ + isExceeded: true, + reason: 'usage_unavailable', + }) + await expect(authorizeLifecycleContinuation(context)).resolves.toBeUndefined() + }) + it('refuses a payer the gate finds blocked as a blocked account, not with the usage card', async () => { + mockCheckExecutionUsageLimits.mockResolvedValueOnce({ + isExceeded: true, + reason: 'billing_blocked', + scope: 'payer', + }) + + const refusal = authorizeLifecycleContinuation(context) + + await expect(refusal).rejects.toBeInstanceOf(OrchestrationError) + await expect(refusal).rejects.toThrow('blocked') + }) + it('judges a leg past its admitted period against the payer current period', async () => { + const ended = { + ...attribution, + billingPeriod: { + start: '2026-07-01T00:00:00.000Z', + end: '2026-08-01T00:00:00.000Z', + source: 'stripe' as const, + }, + } + mockGetOrganizationSubscription.mockResolvedValue({ + id: 'sub-org', + referenceId: 'original-org', + plan: 'team', + status: 'active', + seats: 4, + periodStart: new Date(attribution.billingPeriod.start), + periodEnd: new Date(attribution.billingPeriod.end), + }) + mockCheckExecutionUsageLimits.mockImplementation(async (judged: typeof attribution) => ({ + isExceeded: judged.billingPeriod.end === attribution.billingPeriod.end, + scope: 'payer', + })) + + await expect( + authorizeLifecycleContinuation({ ...context, billingAttribution: ended }) + ).rejects.toBeInstanceOf(BillingLimitError) + + resetMidRunUsageCaches() + mockGetOrganizationSubscription.mockRejectedValue(new Error('subscription read failed')) + await expect( + authorizeLifecycleContinuation({ ...context, billingAttribution: ended }) + ).resolves.toBeUndefined() + }) + it('carries a member cap into the refusal so the card names who can raise it', async () => { + mockCheckExecutionUsageLimits.mockResolvedValue({ isExceeded: true, scope: 'member' }) + + await expect(authorizeLifecycleContinuation(context)).rejects.toMatchObject({ + scope: 'member', + }) + }) }) diff --git a/apps/sim/lib/mothership/request/lifecycle/admission.ts b/apps/sim/lib/mothership/request/lifecycle/admission.ts index b685b7bf4d8..efdcba71e38 100644 --- a/apps/sim/lib/mothership/request/lifecycle/admission.ts +++ b/apps/sim/lib/mothership/request/lifecycle/admission.ts @@ -6,12 +6,14 @@ import { COPILOT_BILLING_PROTOCOL_HEADER, requireBillingCallbackAttribution, } from '@/lib/billing/core/billing-attribution' +import { readMidRunUsageVerdict } from '@/lib/billing/core/mid-run-usage' import { isHosted } from '@/lib/core/config/env-flags' import { OrchestrationError } from '@/lib/core/orchestration/types' import { authorizeCopilotChatCallback, checkCopilotContinuationBilling, } from '@/lib/mothership/application/authorize-chat-callback' +import { BillingLimitError } from '@/lib/mothership/request/go/stream' import { BillingAdmissionSchema } from '@/lib/mothership/request/lifecycle/recovery-config' import type { ExecutionContext } from '@/lib/mothership/request/types' @@ -35,7 +37,13 @@ export function restoreBillingAdmission( return { attribution, envelope } } -/** Every resumed model leg rechecks authority and account standing without reading spend. */ +/** + * Every resumed model leg rechecks authority, account standing, and the original payer's spend. + * Spend is read through the execution usage gate, so a run under its limit pays no ledger read on + * most legs while an over-limit payer is re-read and refused. A spent limit is a + * {@link BillingLimitError}, which the lifecycle turns into the same upgrade card as a refused + * dispatch; a blocked account is refused as blocked; a usage read that fails lets the leg run. + */ export async function authorizeLifecycleContinuation( context: Pick< ExecutionContext, @@ -72,5 +80,9 @@ export async function authorizeLifecycleContinuation( }) if (standing.blocked) throw new OrchestrationError('forbidden', 'Continuation billing account is blocked') + const usage = await readMidRunUsageVerdict(context.billingAttribution) + if (usage.status === 'blocked') + throw new OrchestrationError('forbidden', 'Continuation billing account is blocked') + if (usage.status === 'exceeded') throw new BillingLimitError(context.userId, usage.scope) } } diff --git a/apps/sim/lib/mothership/request/lifecycle/run.test.ts b/apps/sim/lib/mothership/request/lifecycle/run.test.ts index 963aeb78ee0..d9872ecabf2 100644 --- a/apps/sim/lib/mothership/request/lifecycle/run.test.ts +++ b/apps/sim/lib/mothership/request/lifecycle/run.test.ts @@ -100,11 +100,13 @@ vi.mock('@/lib/mothership/request/go/stream', () => { class BillingLimitError extends Error { userId: string + scope?: string - constructor(userId: string) { + constructor(userId: string, scope?: string) { super('Usage limit reached') this.name = 'BillingLimitError' this.userId = userId + this.scope = scope } } @@ -199,6 +201,11 @@ vi.mock('@/lib/mothership/request/tools/billing', () => ({ handleBillingLimitResponse: vi.fn(), })) +const mockRequestExplicitStreamAbort = vi.hoisted(() => vi.fn()) +vi.mock('@/lib/mothership/request/session/explicit-abort', () => ({ + requestExplicitStreamAbort: mockRequestExplicitStreamAbort, +})) + vi.mock('@/lib/mothership/request/tools/executor', () => ({ executeToolAndReport: vi.fn(), failPendingToolCall: mockForceFailHungToolCall, @@ -212,12 +219,14 @@ vi.mock('@/lib/mothership/request/enterprise-byok', () => ({ resolveEnterpriseByokKey: mockResolveEnterpriseByokKey, })) +import { resetUsageGateCache } from '@/lib/billing/core/usage-gate-cache' import { buildPersistedAssistantMessage } from '@/lib/mothership/chat/persisted-message' import { MothershipStreamV1CompletionStatus, MothershipStreamV1ToolOutcome, } from '@/lib/mothership/generated/mothership-stream-v1' import { + BillingLimitError, CopilotBackendError, STREAM_ENDED_WITHOUT_TERMINAL_MESSAGE, StreamEndedWithoutTerminalError, @@ -2256,7 +2265,7 @@ describe('runCopilotLifecycle', () => { }, payerSubscription: null, } - setEnvFlags({ isHosted: true }) + setEnvFlags({ isHosted: true, isBillingEnabled: true }) mockEnv.COPILOT_API_KEY = 'sim-agent-key' mockRunStreamLoop.mockImplementationOnce( async ( @@ -2330,18 +2339,18 @@ describe('runCopilotLifecycle', () => { } }) - it('cold recovery preserves billing identity and does not read spend again', async () => { + it('cold recovery reads spend only against the original billing identity', async () => { const attribution = { actorUserId: 'user-1', workspaceId: 'ws-1', organizationId: 'org-1', billedAccountUserId: 'original-owner', billingEntity: { type: 'organization' as const, id: 'org-1' }, - billingPeriod: { start: '2026-07-01T00:00:00.000Z', end: '2026-08-01T00:00:00.000Z' }, + billingPeriod: { start: '2026-07-01T00:00:00.000Z', end: '2099-01-01T00:00:00.000Z' }, payerSubscription: null, } - setEnvFlags({ isHosted: true }) - mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true }) + setEnvFlags({ isHosted: true, isBillingEnabled: true }) + resetUsageGateCache() const billingRequestId = generateId() const onBillingAdmission = vi.fn() await runCopilotLifecycle( @@ -2363,7 +2372,13 @@ describe('runCopilotLifecycle', () => { onBillingAdmission, } ) - expect(mockCheckAttributedUsageLimits).not.toHaveBeenCalled() + expect(mockCheckAttributedUsageLimits).toHaveBeenCalledOnce() + expect(mockCheckAttributedUsageLimits).toHaveBeenCalledWith( + expect.objectContaining({ + billedAccountUserId: 'original-owner', + billingEntity: attribution.billingEntity, + }) + ) expect(onBillingAdmission).not.toHaveBeenCalled() expect(continuationAuth).toHaveBeenCalled() expect(mockRunStreamLoop).toHaveBeenCalledOnce() @@ -2390,11 +2405,11 @@ describe('runCopilotLifecycle', () => { }, payerSubscription: null, } - setEnvFlags({ isHosted: true }) + setEnvFlags({ isHosted: true, isBillingEnabled: true }) mockCheckAttributedUsageLimits.mockResolvedValue({ isExceeded: true, message: 'limit reached', - scope: 'payer', + scope: 'member', }) const result = await runCopilotLifecycle( @@ -2412,10 +2427,87 @@ describe('runCopilotLifecycle', () => { expect(mockCheckAttributedUsageLimits).toHaveBeenCalledWith(billingAttribution) expect(handleBillingLimitResponse).toHaveBeenCalledTimes(1) + expect(vi.mocked(handleBillingLimitResponse).mock.calls[0][4]).toBe('member') expect(mockRunStreamLoop).not.toHaveBeenCalled() expect(result.cancelled).not.toBe(true) }) + it('stops the worker run when the worker itself refuses a leg at the usage limit', async () => { + mockRequestExplicitStreamAbort.mockResolvedValue({ settled: true }) + mockRunStreamLoop.mockRejectedValueOnce(new BillingLimitError('user-1')) + + const result = await runCopilotLifecycle( + { message: 'hello', messageId: 'message-1' }, + { userId: 'user-1', workspaceId: 'ws-1', chatId: 'chat-1', runId: 'run-1' } + ) + + expect(handleBillingLimitResponse).toHaveBeenCalledOnce() + expect(mockRequestExplicitStreamAbort).toHaveBeenCalledWith( + expect.objectContaining({ streamId: 'message-1', chatId: 'chat-1' }) + ) + expect(result.error).toBeUndefined() + }) + + it('shows the usage card instead of resuming a run whose payer crossed its limit', async () => { + const billingAttribution = { + actorUserId: 'user-1', + workspaceId: 'ws-1', + organizationId: 'org-1', + billedAccountUserId: 'user-1', + billingEntity: { type: 'organization' as const, id: 'org-1' }, + billingPeriod: { start: '2026-07-01T00:00:00.000Z', end: '2099-01-01T00:00:00.000Z' }, + payerSubscription: null, + } + setEnvFlags({ isHosted: true, isBillingEnabled: true }) + resetUsageGateCache() + mockCheckAttributedUsageLimits + .mockResolvedValueOnce({ isExceeded: false }) + .mockResolvedValue({ isExceeded: true, scope: 'payer' }) + mockRunStreamLoop.mockImplementationOnce( + async (_url: string, _init: RequestInit, context: StreamingContext) => { + context.toolCalls.set('tool-1', { + id: 'tool-1', + name: 'read', + status: MothershipStreamV1ToolOutcome.success, + result: { success: true, output: { content: 'file contents' } }, + }) + context.awaitingAsyncContinuation = { + checkpointId: 'ckpt-1', + pendingToolCallIds: ['tool-1'], + } + } + ) + + mockRequestExplicitStreamAbort.mockResolvedValue({ settled: true }) + const result = await runCopilotLifecycle( + { message: 'hello', messageId: 'message-1' }, + { + userId: 'user-1', + workspaceId: 'ws-1', + chatId: 'chat-1', + executionId: 'execution-1', + runId: 'run-1', + billingAttribution, + } + ) + + expect(mockRunStreamLoop).toHaveBeenCalledOnce() + expect(mockCheckAttributedUsageLimits).toHaveBeenCalledTimes(2) + expect(handleBillingLimitResponse).toHaveBeenCalledOnce() + expect(handleBillingLimitResponse).toHaveBeenCalledWith( + 'user-1', + expect.anything(), + expect.anything(), + expect.anything(), + 'payer' + ) + expect(result.cancelled).not.toBe(true) + expect(result.error).toBeUndefined() + expect(mockRequestExplicitStreamAbort).toHaveBeenCalledWith( + expect.objectContaining({ streamId: 'message-1', userId: 'user-1', chatId: 'chat-1' }) + ) + }) + it('preserves a resume tool name that collides with a configured secret', async () => { const registry = new ResolvedSecretTraceRegistry([ { name: 'TOKEN', plaintext: 'unsafe-tool', encryptedValue: 'ciphertext' }, @@ -2456,7 +2548,7 @@ describe('runCopilotLifecycle', () => { }) it('rejects hosted work without immutable billing attribution before egress', async () => { - setEnvFlags({ isHosted: true }) + setEnvFlags({ isHosted: true, isBillingEnabled: true }) await expect( runCopilotLifecycle( diff --git a/apps/sim/lib/mothership/request/lifecycle/run.ts b/apps/sim/lib/mothership/request/lifecycle/run.ts index 0f2ca8e919b..b141f7d7e06 100644 --- a/apps/sim/lib/mothership/request/lifecycle/run.ts +++ b/apps/sim/lib/mothership/request/lifecycle/run.ts @@ -509,7 +509,13 @@ export async function runCopilotLifecycle( // The worker terminal was already delivered before the relay died. // Rebuild persistence from that receipt without charging its usage twice. } else if (admission.isExceeded) { - await handleBillingLimitResponse(execContext.userId, context, execContext, lifecycleOptions) + await handleBillingLimitResponse( + execContext.userId, + context, + execContext, + lifecycleOptions, + 'scope' in admission ? admission.scope : undefined + ) } else { if (!isContinuation && hostedBillingRequest) await lifecycleOptions.onBillingAdmission?.(hostedBillingRequest) @@ -518,14 +524,29 @@ export async function runCopilotLifecycle( requestPayload, lifecycleOptions.workspaceId ) - await runCheckpointLoop( - modelSafeRequestPayload, - context, - execContext, - lifecycleOptions, - goRoute, - hostedBillingRequest - ) + try { + await runCheckpointLoop( + modelSafeRequestPayload, + context, + execContext, + lifecycleOptions, + goRoute, + hostedBillingRequest + ) + } catch (error) { + // A continuation refused on spend, or a worker 402 on any leg, ends the turn with the + // same card as a refused dispatch and stops the worker run. + if (!(error instanceof BillingLimitError)) throw error + context.awaitingAsyncContinuation = undefined + await handleBillingLimitResponse( + error.userId, + context, + execContext, + lifecycleOptions, + error.scope + ) + await stopWorkerRunAfterUsageRefusal(context.messageId, execContext) + } } // The backend's terminal `complete` is the turn's verdict. A failure it @@ -1228,10 +1249,6 @@ async function runCheckpointLoop( } catch (streamError) { context.trace.endSpan(streamSpan, RequestTraceV1SpanStatus.error) context.trace.setActiveSpan(undefined) - if (streamError instanceof BillingLimitError) { - await handleBillingLimitResponse(streamError.userId, context, execContext, options) - break - } const backoff = retry?.nextDelay(streamError, options.abortSignal) ?? null if (backoff !== null) { /** A recovered connection must not finalize with an earlier transport failure. */ @@ -1691,6 +1708,32 @@ function isAborted(options: CopilotLifecycleOptions, context: StreamingContext): return !!(options.abortSignal?.aborted || context.wasAborted) } +/** + * A refused continuation leaves the worker run parked on its checkpoint, and a parked run holds + * the chat: the next message would be refused as busy until the sweeper expires it. Stopping it + * frees the chat, so the message sent after an upgrade continues the conversation. + */ +async function stopWorkerRunAfterUsageRefusal( + streamId: string, + execContext: Pick +): Promise { + try { + const { requestExplicitStreamAbort } = await import( + '@/lib/mothership/request/session/explicit-abort' + ) + await requestExplicitStreamAbort({ + streamId, + userId: execContext.userId, + chatId: execContext.chatId, + }) + } catch (error) { + logger.warn('Worker stop after a usage-limit refusal was not delivered', { + streamId, + error: getErrorMessage(error), + }) + } +} + function cancelPendingTools(context: StreamingContext): void { for (const [, toolCall] of context.toolCalls) { if ( diff --git a/apps/sim/lib/mothership/request/lifecycle/stream-retry.test.ts b/apps/sim/lib/mothership/request/lifecycle/stream-retry.test.ts index 558566dbe31..42544e57c6c 100644 --- a/apps/sim/lib/mothership/request/lifecycle/stream-retry.test.ts +++ b/apps/sim/lib/mothership/request/lifecycle/stream-retry.test.ts @@ -1,5 +1,6 @@ import { afterEach, describe, expect, it, vi } from 'vitest' import { + BillingLimitError, CopilotBackendError, StreamEndedWithoutTerminalError, WorkerStreamInterruptedError, @@ -142,6 +143,7 @@ describe('stream recovery budget', () => { expect(retry.nextDelay(new DOMException('Stopped', 'AbortError'))).toBeNull() expect(retry.nextDelay(new CopilotBackendError('Forbidden', { status: 403 }))).toBeNull() expect(retry.nextDelay(new Error('Invalid operation'))).toBeNull() + expect(retry.nextDelay(new BillingLimitError('user-1'))).toBeNull() expect(retry.attempt).toBe(0) }) diff --git a/apps/sim/lib/mothership/request/tools/billing.test.ts b/apps/sim/lib/mothership/request/tools/billing.test.ts index d1c74216697..2154b06e8b3 100644 --- a/apps/sim/lib/mothership/request/tools/billing.test.ts +++ b/apps/sim/lib/mothership/request/tools/billing.test.ts @@ -102,4 +102,31 @@ describe('handleBillingLimitResponse', () => { payload: { text: expect.stringContaining('"action":"increase_limit"') }, }) }) + + it('terminates the turn with the card after a leg that already ended at a checkpoint', async () => { + const onEvent = vi.fn() + const pausedContext = { streamComplete: true } as StreamingContext + + await handleBillingLimitResponse('actor-1', pausedContext, createExecutionContext(), { + onEvent, + } as OrchestratorOptions) + + expect(onEvent.mock.calls.map(([event]) => event.type)).toEqual(['text', 'complete']) + }) + + it('names who can raise the cap for a member over the limit their organization set', async () => { + const onEvent = vi.fn() + + await handleBillingLimitResponse( + 'actor-1', + { streamComplete: false } as StreamingContext, + createExecutionContext(), + { onEvent } as OrchestratorOptions, + 'member' + ) + + expect(onEvent.mock.calls[0]?.[0]).toMatchObject({ + payload: { text: expect.stringContaining('limit your organization set for you') }, + }) + }) }) diff --git a/apps/sim/lib/mothership/request/tools/billing.ts b/apps/sim/lib/mothership/request/tools/billing.ts index fed35fd4156..295bf69ddcc 100644 --- a/apps/sim/lib/mothership/request/tools/billing.ts +++ b/apps/sim/lib/mothership/request/tools/billing.ts @@ -1,8 +1,5 @@ import { createLogger } from '@sim/logger' -import { getErrorMessage } from '@sim/utils/errors' -import { getHighestPrioritySubscription } from '@/lib/billing/core/plan' -import { isEnterprise, isPaid } from '@/lib/billing/plan-helpers' -import { isOrgScopedSubscription } from '@/lib/billing/subscriptions/utils' +import { formatUsageUpgradeTag, resolveUsageUpgradePayload } from '@/lib/billing/usage-upgrade' import { MothershipStreamV1CompletionStatus, MothershipStreamV1EventType, @@ -19,59 +16,24 @@ import type { const logger = createLogger('CopilotBillingEffect') /** - * Handle a 402 billing-limit response from the Go backend. + * Ends the turn with the usage card: a refused dispatch or continuation, a worker 402, or a + * worker usage-limit terminal that arrived without a card of its own. * - * Determines whether the user needs a plan upgrade or a limit increase, - * then dispatches synthetic text + complete events through the handler chain - * so the client renders the upgrade prompt. + * Dispatches synthetic text + complete events through the handler chain so the client renders + * the upgrade prompt and the turn finishes as complete, so the next message after an upgrade + * starts normally. */ export async function handleBillingLimitResponse( userId: string, context: StreamingContext, execContext: ExecutionContext, - options: OrchestratorOptions + options: OrchestratorOptions, + scope?: 'actor' | 'payer' | 'member' ): Promise { - let action: 'upgrade_plan' | 'increase_limit' = 'upgrade_plan' - let message = "You've reached your usage limit. Please upgrade your plan to continue." - try { - let plan: string | undefined - let orgScoped = false - if (execContext.billingAttribution) { - plan = execContext.billingAttribution.payerSubscription?.plan - orgScoped = execContext.billingAttribution.billingEntity.type === 'organization' - } else { - const sub = await getHighestPrioritySubscription(userId) - plan = sub?.plan - orgScoped = isOrgScopedSubscription(sub, userId) - } - - if (plan && isPaid(plan)) { - // Paid subs use the existing `increase_limit` action so the UI - // (`UsageUpgradeDisplay`) renders its standard button. The message - // text does the work of clarifying the action when the user can't - // actually self-serve the limit change. - action = 'increase_limit' - if (orgScoped) { - message = isEnterprise(plan) - ? "You've reached your organization's usage limit for this billing period. Only an organization admin or Sim support can raise an enterprise limit — reach out to them to continue." - : "You've reached your organization's usage limit for this billing period. Only an organization owner or admin can raise the limit — please ask them to update it from the team billing settings." - } else { - message = - "You've reached your usage limit for this billing period. Please increase your usage limit from billing settings to continue." - } - } - } catch (error) { - logger.warn('Failed to determine subscription plan, defaulting to upgrade_plan', { - error: getErrorMessage(error), - }) - } - - const upgradePayload = JSON.stringify({ - reason: 'usage_limit', - action, - message, - }) - const syntheticContent = `${upgradePayload}` + const payload = await resolveUsageUpgradePayload(userId, execContext.billingAttribution, scope) + const syntheticContent = formatUsageUpgradeTag(payload) + // The card is this turn's terminal even when the refused leg follows one that already ended. + context.streamComplete = false const syntheticEvents: StreamEvent[] = [ { diff --git a/packages/testing/src/mocks/billing-attribution.mock.ts b/packages/testing/src/mocks/billing-attribution.mock.ts index 3ad64d75b51..a54c31a63e6 100644 --- a/packages/testing/src/mocks/billing-attribution.mock.ts +++ b/packages/testing/src/mocks/billing-attribution.mock.ts @@ -160,8 +160,10 @@ export const billingAttributionMockFns = { mockResolveLegacyV0BillingAttribution: vi.fn(), mockResolveSystemBillingAttribution: vi.fn(), mockToBillingContext: vi.fn(toBillingContext), + mockCheckAccountBillingBlocks: vi.fn(), mockCheckAttributedBillingBlocks: vi.fn(), mockCheckAttributedUsageLimits: vi.fn(), + mockRefreshAttributionPeriod: vi.fn(), } /** @@ -210,6 +212,8 @@ export const billingAttributionMock = { billingAttributionMockFns.mockResolveLegacyV0BillingAttribution, resolveSystemBillingAttribution: billingAttributionMockFns.mockResolveSystemBillingAttribution, toBillingContext: billingAttributionMockFns.mockToBillingContext, + checkAccountBillingBlocks: billingAttributionMockFns.mockCheckAccountBillingBlocks, checkAttributedBillingBlocks: billingAttributionMockFns.mockCheckAttributedBillingBlocks, checkAttributedUsageLimits: billingAttributionMockFns.mockCheckAttributedUsageLimits, + refreshAttributionPeriod: billingAttributionMockFns.mockRefreshAttributionPeriod, }