fix(billing): close TOCTOU race in subscription transfer, centralize stripe test mocks (#4239)

* fix(billing): close TOCTOU race in subscription transfer, centralize stripe test mocks

* more mocks

* fix(testing): provide complete Stripe.Event defaults in createMockStripeEvent

* fix(testing): make dbChainMock .for('update') chainable with .limit()

* fix(billing): gate subscription transfer noop behind membership check

Previously the 'already belongs to this organization' early return fired
before the org/member lookups, letting any authenticated caller probe
sub-to-org pairings without being a member of the target org. Move the
noop check after the admin/owner verification so unauthorized callers
hit the 403 first.
This commit is contained in:
Waleed
2026-04-20 16:46:28 -07:00
committed by GitHub
parent 0cd14f4ac9
commit ac4ccfcac8
14 changed files with 482 additions and 473 deletions
+23 -5
View File
@@ -64,13 +64,19 @@ export function createMockSqlOperators() {
* are wired at module load time:
*
* - `select().from().where()` → returns a builder with `.limit` / `.orderBy` /
* `.returning` / `.groupBy` terminals
* `.returning` / `.groupBy` / `.for` terminals
* - `select().from().innerJoin()|leftJoin()` → returns the same where-builder
* - `insert().values().returning()` / `update().set().where()` / `delete().where()`
*
* Terminals (`limit`, `orderBy`, `returning`, `groupBy`, `values`) default to
* resolving `[]` (or `undefined` for `values`). Override per-test with
* `dbChainMockFns.limit.mockResolvedValueOnce([...])`.
* Terminals (`limit`, `orderBy`, `returning`, `groupBy`, `for`, `values`)
* default to resolving `[]` (or `undefined` for `values`). Override per-test
* with `dbChainMockFns.limit.mockResolvedValueOnce([...])`. `for` mirrors
* drizzle's `.for('update')` — it returns a Promise with `.limit` / `.orderBy`
* / `.returning` / `.groupBy` attached, so both `await .where().for('update')`
* (terminal) and `await .where().for('update').limit(1)` (chained) work.
* Override the terminal result with `dbChainMockFns.for.mockResolvedValueOnce(
* [...])`; override the chained result by mocking the downstream terminal
* (e.g. `dbChainMockFns.limit.mockResolvedValueOnce([...])`).
*
* `vi.clearAllMocks()` clears call history but preserves default wiring. Tests
* that replace a wiring with `mockReturnValue(...)` (not `...Once`) must re-wire
@@ -94,10 +100,20 @@ const returning = vi.fn(() => Promise.resolve([] as unknown[]))
const groupBy = vi.fn(() => Promise.resolve([] as unknown[]))
const execute = vi.fn(() => Promise.resolve([] as unknown[]))
const forBuilder = () => {
const thenable: any = Promise.resolve([] as unknown[])
thenable.limit = limit
thenable.orderBy = orderBy
thenable.returning = returning
thenable.groupBy = groupBy
return thenable
}
const forClause = vi.fn(forBuilder)
const onConflictDoUpdate = vi.fn(() => ({ returning }) as unknown as Promise<void>)
const onConflictDoNothing = vi.fn(() => ({ returning }) as unknown as Promise<void>)
const whereBuilder = () => ({ limit, orderBy, returning, groupBy })
const whereBuilder = () => ({ limit, orderBy, returning, groupBy, for: forClause })
const where = vi.fn(whereBuilder)
const joinBuilder = (): { where: typeof where; innerJoin: any; leftJoin: any } => ({
@@ -134,6 +150,7 @@ export const dbChainMockFns = {
leftJoin,
groupBy,
execute,
for: forClause,
insert,
values,
onConflictDoUpdate,
@@ -173,6 +190,7 @@ export function resetDbChainMock(): void {
returning.mockImplementation(() => Promise.resolve([] as unknown[]))
groupBy.mockImplementation(() => Promise.resolve([] as unknown[]))
execute.mockImplementation(() => Promise.resolve([] as unknown[]))
forClause.mockImplementation(forBuilder)
transaction.mockImplementation(async (cb: (tx: typeof dbChainMock.db) => unknown) =>
cb(dbChainMock.db)
)
+8
View File
@@ -108,6 +108,14 @@ export {
} from './socket.mock'
// Storage mocks
export { clearStorageMocks, createMockStorage, setupGlobalStorageMocks } from './storage.mock'
// Stripe mocks
export {
createMockStripeEvent,
stripeClientMock,
stripeClientMockFns,
stripePaymentMethodMock,
stripePaymentMethodMockFns,
} from './stripe.mock'
// Telemetry mocks
export { telemetryMock } from './telemetry.mock'
// URL mocks
+71
View File
@@ -0,0 +1,71 @@
import type Stripe from 'stripe'
import { vi } from 'vitest'
/**
* Mock for `@/lib/billing/stripe-client`.
*
* @example
* ```ts
* import { stripeClientMock, stripeClientMockFns } from '@sim/testing'
* vi.mock('@/lib/billing/stripe-client', () => stripeClientMock)
*
* stripeClientMockFns.mockRequireStripeClient.mockReturnValue(fakeStripe)
* ```
*/
export const stripeClientMockFns = {
mockRequireStripeClient: vi.fn(),
mockGetStripeClient: vi.fn(),
mockHasValidStripeCredentials: vi.fn(() => true),
}
export const stripeClientMock = {
requireStripeClient: stripeClientMockFns.mockRequireStripeClient,
getStripeClient: stripeClientMockFns.mockGetStripeClient,
hasValidStripeCredentials: stripeClientMockFns.mockHasValidStripeCredentials,
}
/**
* Mock for `@/lib/billing/stripe-payment-method`.
*
* @example
* ```ts
* import { stripePaymentMethodMock, stripePaymentMethodMockFns } from '@sim/testing'
* vi.mock('@/lib/billing/stripe-payment-method', () => stripePaymentMethodMock)
* ```
*/
export const stripePaymentMethodMockFns = {
mockResolveDefaultPaymentMethod: vi.fn(async () => ({
paymentMethodId: undefined as string | undefined,
collectionMethod: 'charge_automatically' as 'charge_automatically' | 'send_invoice' | null,
})),
mockGetCustomerId: vi.fn(),
}
export const stripePaymentMethodMock = {
resolveDefaultPaymentMethod: stripePaymentMethodMockFns.mockResolveDefaultPaymentMethod,
getCustomerId: stripePaymentMethodMockFns.mockGetCustomerId,
}
/**
* Build a minimal `Stripe.Event` with the given type and object payload.
* Fills in a deterministic `id` (`evt_${type}`) and nests `object` under
* `data.object` as Stripe does.
*/
export function createMockStripeEvent<T = unknown>(
type: string,
object: T,
overrides: Partial<Stripe.Event> = {}
): Stripe.Event {
return {
id: `evt_${type}`,
object: 'event',
api_version: '2024-06-20',
created: Math.floor(Date.now() / 1000),
livemode: false,
pending_webhooks: 0,
request: null,
type,
data: { object: object as unknown as Stripe.Event.Data.Object },
...overrides,
} as Stripe.Event
}