mirror of
https://github.com/simstudioai/sim.git
synced 2026-09-24 15:45:35 +08:00
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:
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user