diff --git a/.github/scripts/start-cloudflare-workers-with-retry.sh b/.github/scripts/start-cloudflare-workers-with-retry.sh new file mode 100755 index 0000000000..99afe12efd --- /dev/null +++ b/.github/scripts/start-cloudflare-workers-with-retry.sh @@ -0,0 +1,57 @@ +#!/usr/bin/env bash + +set -euo pipefail + +max_attempts="${CLOUDFLARE_WORKERS_START_MAX_ATTEMPTS:-2}" +worker_port_offset="${CLOUDFLARE_WORKER_PORT_OFFSET:-0}" +log_path="${CLOUDFLARE_WORKERS_LOG_PATH:-${RUNNER_TEMP:-/tmp}/cloudflare-workers.log}" +wait_timeout_ms="${CLOUDFLARE_WORKERS_WAIT_TIMEOUT_MS:-180000}" +worker_pids=() + +stop_worker_pid() { + local pid="$1" + if [[ -z "${pid}" ]]; then + return + fi + kill "${pid}" 2>/dev/null || true + pkill -P "${pid}" 2>/dev/null || true +} + +stop_all_worker_pids() { + local pid + for pid in "${worker_pids[@]}"; do + stop_worker_pid "${pid}" + done +} + +export BACKGROUND_SERVICE_NAME="${BACKGROUND_SERVICE_NAME:-Cloudflare Workers}" +export BACKGROUND_RUN_COMMAND="${BACKGROUND_RUN_COMMAND:-$'chmod +x scripts/start-cloudflare-workers.sh\nexec ./scripts/start-cloudflare-workers.sh'}" +export BACKGROUND_LOG_PATH="${log_path}" +export BACKGROUND_WAIT_TIMEOUT_MS="${wait_timeout_ms}" +export BACKGROUND_TAIL_LINES="${BACKGROUND_TAIL_LINES:-400}" +export BACKGROUND_WAIT_ON="http-get://127.0.0.1:$((8787 + worker_port_offset))/ok +http-get://127.0.0.1:$((8788 + worker_port_offset))/ok +http-get://127.0.0.1:$((8789 + worker_port_offset))/ok" + +for attempt in $(seq 1 "${max_attempts}"); do + if bash .github/scripts/start-background-service.sh; then + exit 0 + else + exit_code=$? + if [[ -n "${GITHUB_OUTPUT:-}" ]] && [[ -f "${GITHUB_OUTPUT}" ]]; then + latest_pid="$(grep '^pid=' "${GITHUB_OUTPUT}" | tail -1 | cut -d= -f2- || true)" + if [[ -n "${latest_pid}" ]]; then + worker_pids+=("${latest_pid}") + stop_worker_pid "${latest_pid}" + fi + fi + echo "Cloudflare Workers failed to become ready with exit ${exit_code} (attempt ${attempt}/${max_attempts})" >&2 + if [ "${attempt}" -eq "${max_attempts}" ]; then + stop_all_worker_pids + exit "${exit_code}" + fi + sleep_seconds=$((attempt * 10)) + echo "Retrying Cloudflare Workers startup in ${sleep_seconds}s..." >&2 + sleep "${sleep_seconds}" + fi +done diff --git a/.github/scripts/start-supabase-worktree-with-retry.sh b/.github/scripts/start-supabase-worktree-with-retry.sh new file mode 100755 index 0000000000..04081c8473 --- /dev/null +++ b/.github/scripts/start-supabase-worktree-with-retry.sh @@ -0,0 +1,22 @@ +#!/usr/bin/env bash + +set -euo pipefail + +max_attempts="${SUPABASE_START_MAX_ATTEMPTS:-3}" +exclude_services="${SUPABASE_START_EXCLUDE:-imgproxy,studio,mailpit,realtime,postgres-meta,supavisor,logflare,vector}" + +for attempt in $(seq 1 "${max_attempts}"); do + bun scripts/supabase-worktree.ts stop --no-backup || true + if bun scripts/supabase-worktree.ts start -x "${exclude_services}"; then + exit 0 + else + exit_code=$? + echo "Supabase start failed with exit ${exit_code} (attempt ${attempt}/${max_attempts})" >&2 + if [ "${attempt}" -eq "${max_attempts}" ]; then + exit "${exit_code}" + fi + sleep_seconds=$((attempt * 5)) + echo "Retrying Supabase start in ${sleep_seconds}s..." >&2 + sleep "${sleep_seconds}" + fi +done diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index d2b541cd14..a5867d23c5 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -418,7 +418,7 @@ jobs: warm_get() { local path="$1" for attempt in 1 2 3 4 5; do - status=$(curl -s -o /tmp/edge-warm.body -w '%{http_code}' "${base_url}${path}" || true) + status=$(curl --max-time 10 -s -o /tmp/edge-warm.body -w '%{http_code}' "${base_url}${path}" || true) echo "edge warm GET path=${path} attempt=${attempt} status=${status}" if [ "${status}" != "502" ] && [ "${status}" != "503" ] && [ "${status}" != "000" ]; then return 0 @@ -430,7 +430,7 @@ jobs: warm_post() { local path="$1" for attempt in 1 2 3 4 5; do - status=$(curl -s -o /tmp/edge-warm.body -w '%{http_code}' -X POST "${base_url}${path}" \ + status=$(curl --max-time 10 -s -o /tmp/edge-warm.body -w '%{http_code}' -X POST "${base_url}${path}" \ -H 'Content-Type: application/json' \ -H 'apisecret: testsecret' \ -d '{}' || true) @@ -742,7 +742,7 @@ jobs: needs: [changes, lint_typecheck, dead_code] if: needs.changes.outputs.run_capgo == 'true' runs-on: ubuntu-latest - timeout-minutes: 5 + timeout-minutes: 7 name: Run Cloudflare Workers tests (shard ${{ matrix.shard }}) permissions: contents: read @@ -813,7 +813,7 @@ jobs: - name: Install dependencies run: bun install - name: Run Supabase Start - run: bun scripts/supabase-worktree.ts start -x imgproxy,studio,mailpit,realtime,postgres-meta,supavisor,logflare,vector + run: bash .github/scripts/start-supabase-worktree-with-retry.sh - name: Export isolated test endpoints run: | { @@ -835,20 +835,7 @@ jobs: bash .github/scripts/start-background-service.sh - id: start_cloudflare_workers name: Start Cloudflare Workers for testing - env: - BACKGROUND_SERVICE_NAME: Cloudflare Workers - BACKGROUND_RUN_COMMAND: | - chmod +x scripts/start-cloudflare-workers.sh - exec ./scripts/start-cloudflare-workers.sh - BACKGROUND_LOG_PATH: ${{ runner.temp }}/cloudflare-workers.log - BACKGROUND_WAIT_TIMEOUT_MS: 120000 - BACKGROUND_TAIL_LINES: 400 - run: | - worker_port_offset="${CLOUDFLARE_WORKER_PORT_OFFSET:-0}" - export BACKGROUND_WAIT_ON="http-get://127.0.0.1:$((8787 + worker_port_offset))/ok - http-get://127.0.0.1:$((8788 + worker_port_offset))/ok - http-get://127.0.0.1:$((8789 + worker_port_offset))/ok" - bash .github/scripts/start-background-service.sh + run: bash .github/scripts/start-cloudflare-workers-with-retry.sh - name: Run Cloudflare Workers integration tests env: VITEST_SHARD: ${{ matrix.shard }} @@ -885,7 +872,7 @@ jobs: needs: [changes, lint_typecheck, dead_code] if: needs.changes.outputs.run_capgo == 'true' runs-on: ubuntu-latest - timeout-minutes: 5 + timeout-minutes: 7 name: Run Cloudflare plugin tests (serial) permissions: contents: read @@ -919,7 +906,7 @@ jobs: - name: Install dependencies run: bun install - name: Run Supabase Start - run: bun scripts/supabase-worktree.ts start -x imgproxy,studio,mailpit,realtime,postgres-meta,supavisor,logflare,vector + run: bash .github/scripts/start-supabase-worktree-with-retry.sh - name: Export isolated test endpoints run: | { @@ -941,19 +928,8 @@ jobs: - id: start_cloudflare_workers name: Start Cloudflare Workers for testing env: - BACKGROUND_SERVICE_NAME: Cloudflare Workers - BACKGROUND_RUN_COMMAND: | - chmod +x scripts/start-cloudflare-workers.sh - exec ./scripts/start-cloudflare-workers.sh - BACKGROUND_LOG_PATH: ${{ runner.temp }}/cloudflare-plugin-workers.log - BACKGROUND_WAIT_TIMEOUT_MS: 120000 - BACKGROUND_TAIL_LINES: 400 - run: | - worker_port_offset="${CLOUDFLARE_WORKER_PORT_OFFSET:-0}" - export BACKGROUND_WAIT_ON="http-get://127.0.0.1:$((8787 + worker_port_offset))/ok - http-get://127.0.0.1:$((8788 + worker_port_offset))/ok - http-get://127.0.0.1:$((8789 + worker_port_offset))/ok" - bash .github/scripts/start-background-service.sh + CLOUDFLARE_WORKERS_LOG_PATH: ${{ runner.temp }}/cloudflare-plugin-workers.log + run: bash .github/scripts/start-cloudflare-workers-with-retry.sh - name: Warm plugin /updates endpoint run: | plugin_url="http://127.0.0.1:$((8788 + CLOUDFLARE_WORKER_PORT_OFFSET))/updates" @@ -1503,6 +1479,45 @@ jobs: run: | export BACKGROUND_WAIT_ON="${SUPABASE_FUNCTIONS_HEALTH_URL:?Missing isolated Supabase health URL}" bash .github/scripts/start-background-service.sh + - name: Warm edge functions before CLI integration tests + run: | + # GET /ok alone is not enough — cold first requests under file parallelism 502. + health_url="${SUPABASE_FUNCTIONS_HEALTH_URL#http-get://}" + health_url="${health_url#https-get://}" + base_url="http://${health_url%/ok}" + warm_get() { + local path="$1" + for attempt in 1 2 3 4 5; do + status=$(curl --max-time 10 -s -o /tmp/edge-warm.body -w '%{http_code}' "${base_url}${path}" || true) + echo "edge warm GET path=${path} attempt=${attempt} status=${status}" + if [ "${status}" != "502" ] && [ "${status}" != "503" ] && [ "${status}" != "000" ]; then + return 0 + fi + sleep "${attempt}" + done + return 1 + } + warm_post() { + local path="$1" + for attempt in 1 2 3 4 5; do + status=$(curl --max-time 10 -s -o /tmp/edge-warm.body -w '%{http_code}' -X POST "${base_url}${path}" \ + -H 'Content-Type: application/json' \ + -H 'apisecret: testsecret' \ + -d '{}' || true) + echo "edge warm POST path=${path} attempt=${attempt} status=${status}" + if [ "${status}" != "502" ] && [ "${status}" != "503" ] && [ "${status}" != "000" ]; then + return 0 + fi + sleep "${attempt}" + done + return 1 + } + warm_get /ok + warm_post /triggers/cron_email + warm_post /triggers/cron_stat_org + warm_post /triggers/cron_stat_app + warm_post /apikey + warm_post /app - name: Run Capgo CLI integration tests run: | bun run supabase:with-env -- bunx vitest run tests/cli* --exclude=tests/cli-min-version.test.ts --exclude=tests/cli-new-encryption.test.ts diff --git a/BOUNTY.md b/BOUNTY.md index ab08e12408..f88f4a5c09 100644 --- a/BOUNTY.md +++ b/BOUNTY.md @@ -25,9 +25,6 @@ Anyone from the community can review the pull request and leave comments. -Review are rewarded with a tip of $20 when requested on merged pull request. -AI review does not qualify. - ## What is a good review? Check code pattern repetition, and things that can be done better. diff --git a/cli/src/types/supabase.types.ts b/cli/src/types/supabase.types.ts index b958946e65..d3497f914b 100644 --- a/cli/src/types/supabase.types.ts +++ b/cli/src/types/supabase.types.ts @@ -2267,6 +2267,7 @@ export type Database = { build_time_unit: number created_at: string credit_id: string + credit_id_us: string | null description: string id: string market_desc: string | null @@ -2277,8 +2278,11 @@ export type Database = { price_m_id: string price_y: number price_y_id: string + price_y_id_us: string | null + price_m_id_us: string | null storage: number stripe_id: string + stripe_id_us: string | null updated_at: string } Insert: { @@ -2286,6 +2290,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2296,8 +2301,11 @@ export type Database = { price_m_id: string price_y?: number price_y_id: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Update: { @@ -2305,6 +2313,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id?: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2315,8 +2324,11 @@ export type Database = { price_m_id?: string price_y?: number price_y_id?: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage?: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Relationships: [] @@ -2642,6 +2654,7 @@ export type Database = { stripe_info: { Row: { bandwidth_exceeded: boolean | null + billing_account: string build_time_exceeded: boolean | null canceled_at: string | null churn_reason: string | null @@ -2669,6 +2682,7 @@ export type Database = { } Insert: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -2696,6 +2710,7 @@ export type Database = { } Update: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -2721,15 +2736,7 @@ export type Database = { updated_at?: string upgraded_at?: string | null } - Relationships: [ - { - foreignKeyName: "stripe_info_product_id_fkey" - columns: ["product_id"] - isOneToOne: false - referencedRelation: "plans" - referencedColumns: ["stripe_id"] - }, - ] + Relationships: [] } trial_extension_events: { Row: { diff --git a/cloudflare_workers/api/index.ts b/cloudflare_workers/api/index.ts index c35fa12eef..1dd258fb92 100644 --- a/cloudflare_workers/api/index.ts +++ b/cloudflare_workers/api/index.ts @@ -95,6 +95,7 @@ import { app as pluginNotifications } from '../../supabase/functions/_backend/tr import { app as queue_consumer } from '../../supabase/functions/_backend/triggers/queue_consumer.ts' import { app as send_email } from './triggers/send_email.ts' import { app as stripe_event } from '../../supabase/functions/_backend/triggers/stripe_event.ts' +import { app as stripe_event_us } from '../../supabase/functions/_backend/triggers/stripe_event_us.ts' import { app as webhook_delivery } from '../../supabase/functions/_backend/triggers/webhook_delivery.ts' import { app as webhook_dispatcher } from '../../supabase/functions/_backend/triggers/webhook_dispatcher.ts' import { BRES, createAllCatch, createHono } from '../../supabase/functions/_backend/utils/hono.ts' @@ -220,6 +221,7 @@ appTriggers.route('/on_version_delete', on_version_delete) appTriggers.route('/on_manifest_create', on_manifest_create) appTriggers.route('/on_deploy_history_create', on_deploy_history_create) appTriggers.route('/stripe_event', stripe_event) +appTriggers.route('/stripe_event_us', stripe_event_us) appTriggers.route('/on_organization_create', on_organization_create) appTriggers.route('/cron_stat_app', cron_stat_app) appTriggers.route('/cron_stat_org', cron_stat_org) diff --git a/read_replicate/schema_replicate.catalog.json b/read_replicate/schema_replicate.catalog.json index fce31cbb4e..bc69191dfb 100644 --- a/read_replicate/schema_replicate.catalog.json +++ b/read_replicate/schema_replicate.catalog.json @@ -1819,6 +1819,16 @@ "position": 26, "table": "stripe_info", "type": "boolean" + }, + { + "default": "'ee'::text", + "generated": "", + "identity": "", + "name": "billing_account", + "notNull": true, + "position": 27, + "table": "stripe_info", + "type": "text" } ], "constraints": [ @@ -2067,6 +2077,13 @@ "type": "u", "valid": true }, + { + "definition": "CHECK (billing_account = ANY (ARRAY['ee'::text, 'us'::text]))", + "name": "stripe_info_billing_account_check", + "table": "stripe_info", + "type": "c", + "valid": true + }, { "definition": "PRIMARY KEY (customer_id)", "name": "stripe_info_pkey", diff --git a/read_replicate/schema_replicate.sql b/read_replicate/schema_replicate.sql index 77ffb0b6ae..e0e958bc07 100644 --- a/read_replicate/schema_replicate.sql +++ b/read_replicate/schema_replicate.sql @@ -446,7 +446,9 @@ CREATE TABLE public.stripe_info ( last_stripe_event_at timestamp with time zone, past_due_at timestamp with time zone, churn_reason text, - is_above_plan boolean + is_above_plan boolean, + billing_account text DEFAULT 'ee'::text NOT NULL, + CONSTRAINT stripe_info_billing_account_check CHECK ((billing_account = ANY (ARRAY['ee'::text, 'us'::text]))) ); ALTER TABLE ONLY public.stripe_info REPLICA IDENTITY FULL; diff --git a/scripts/serve-backend-playwright.ts b/scripts/serve-backend-playwright.ts index cfb0ebe612..5fb74c0cbd 100644 --- a/scripts/serve-backend-playwright.ts +++ b/scripts/serve-backend-playwright.ts @@ -30,10 +30,13 @@ function upsertEnvValue(content: string, key: string, value: string): string { const baseEnv = existsSync(sourceEnvPath) ? readFileSync(sourceEnvPath, 'utf8') : '' const overriddenEnv = [ + ['ENV', env.ENV || 'local'], ['S3_ENDPOINT', `127.0.0.1:${supabaseConfig.ports.api}/storage/v1/s3`], ['STRIPE_SECRET_KEY', env.STRIPE_SECRET_KEY || 'sk_test_emulator'], + ['STRIPE_SECRET_KEY_US', env.STRIPE_SECRET_KEY_US || env.STRIPE_SECRET_KEY || 'sk_test_emulator'], ['STRIPE_API_BASE_URL', stripeApiBaseUrl], ['STRIPE_WEBHOOK_SECRET', env.STRIPE_WEBHOOK_SECRET || 'testsecret'], + ['STRIPE_WEBHOOK_SECRET_US', env.STRIPE_WEBHOOK_SECRET_US || env.STRIPE_WEBHOOK_SECRET || 'testsecret'], ['WEBAPP_URL', webAppUrl], ] as const diff --git a/scripts/supabase-worktree.ts b/scripts/supabase-worktree.ts index f050466cc2..81d3838237 100644 --- a/scripts/supabase-worktree.ts +++ b/scripts/supabase-worktree.ts @@ -279,7 +279,7 @@ function parseInlineEnvAssignments(args: string[]): { env: Record 'us' AND p.stripe_id = si.product_id)` + export async function getAdminPayingOrgBreakdown(c: Context): Promise { const emptyResult: AdminPayingOrgBreakdown = { paying_orgs_subscription: 0, @@ -2128,7 +2130,7 @@ export async function getAdminPayingOrgBreakdown(c: Context): Promise= INTERVAL '330 days' @@ -2641,7 +2643,7 @@ export async function getAdminOrganizationInsights( ${billingTypeExpression} AS billing_type FROM orgs o LEFT JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE true ${planFilter} ${billingFilter} @@ -2837,7 +2839,7 @@ export async function getAdminOrganizationInsights( SELECT COUNT(*)::int AS total FROM orgs o LEFT JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE true ${planFilter} ${billingFilter} @@ -2912,6 +2914,7 @@ export interface AdminCancelledOrganizationRow { plan_name: string | null billing_type: 'monthly' | 'yearly' | null subscription_or_signup_date: string + billing_account: 'ee' | 'us' } export interface AdminCancelledOrganizationsResult { @@ -2949,10 +2952,11 @@ export async function getAdminCancelledOrganizations( si.customer_id, si.churn_reason, si.subscription_id, + COALESCE(si.billing_account, 'ee') AS billing_account, p.name AS plan_name, CASE - WHEN si.price_id = p.price_y_id THEN 'yearly' - WHEN si.price_id = p.price_m_id THEN 'monthly' + WHEN si.price_id IN (p.price_y_id, p.price_y_id_us) THEN 'yearly' + WHEN si.price_id IN (p.price_m_id, p.price_m_id_us) THEN 'monthly' WHEN si.subscription_anchor_start IS NOT NULL AND si.subscription_anchor_end IS NOT NULL AND si.subscription_anchor_end::timestamp - si.subscription_anchor_start::timestamp >= INTERVAL '330 days' @@ -2965,7 +2969,7 @@ export async function getAdminCancelledOrganizations( COALESCE(si.paid_at, u.created_at, o.created_at) AS subscription_or_signup_date FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} LEFT JOIN users u ON u.id = o.created_by WHERE si.canceled_at IS NOT NULL ${dateFilter} @@ -2998,6 +3002,7 @@ export async function getAdminCancelledOrganizations( plan_name: row.plan_name ?? null, billing_type: row.billing_type ?? null, subscription_or_signup_date: normalizeTimestamp(row.subscription_or_signup_date) ?? '', + billing_account: row.billing_account === 'us' ? 'us' : 'ee', })) const total = Number((countResult.rows[0] as any)?.total) || 0 @@ -3080,7 +3085,7 @@ export async function getAdminTrialOrganizations( lbu.last_bundle_upload_at FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} LEFT JOIN latest_bundle_uploads lbu ON lbu.owner_org = o.id WHERE si.trial_at::date >= CURRENT_DATE AND (si.status IS NULL OR si.status != 'succeeded') @@ -3175,7 +3180,7 @@ export async function getAdminTrialPlanBreakdown( COUNT(DISTINCT o.id)::int AS trials FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE o.created_at >= ${startDay.toISOString()}::timestamptz AND o.created_at < ${endExclusive.toISOString()}::timestamptz AND si.trial_at IS NOT NULL diff --git a/supabase/functions/_backend/plugin_runtime/utils/supabase.types.ts b/supabase/functions/_backend/plugin_runtime/utils/supabase.types.ts index b06b835891..b3ba355dea 100644 --- a/supabase/functions/_backend/plugin_runtime/utils/supabase.types.ts +++ b/supabase/functions/_backend/plugin_runtime/utils/supabase.types.ts @@ -2507,6 +2507,7 @@ export type Database = { build_time_unit: number created_at: string credit_id: string + credit_id_us: string | null description: string id: string market_desc: string | null @@ -2517,8 +2518,11 @@ export type Database = { price_m_id: string price_y: number price_y_id: string + price_y_id_us: string | null + price_m_id_us: string | null storage: number stripe_id: string + stripe_id_us: string | null updated_at: string } Insert: { @@ -2526,6 +2530,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2536,8 +2541,11 @@ export type Database = { price_m_id: string price_y?: number price_y_id: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Update: { @@ -2545,6 +2553,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id?: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2555,8 +2564,11 @@ export type Database = { price_m_id?: string price_y?: number price_y_id?: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage?: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Relationships: [] @@ -2882,6 +2894,7 @@ export type Database = { stripe_info: { Row: { bandwidth_exceeded: boolean | null + billing_account: string build_time_exceeded: boolean | null canceled_at: string | null churn_reason: string | null @@ -2910,6 +2923,7 @@ export type Database = { } Insert: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -2938,6 +2952,7 @@ export type Database = { } Update: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -2964,15 +2979,7 @@ export type Database = { updated_at?: string upgraded_at?: string | null } - Relationships: [ - { - foreignKeyName: "stripe_info_product_id_fkey" - columns: ["product_id"] - isOneToOne: false - referencedRelation: "plans" - referencedColumns: ["stripe_id"] - }, - ] + Relationships: [] } tmp_users: { Row: { diff --git a/supabase/functions/_backend/private/admin_stats.ts b/supabase/functions/_backend/private/admin_stats.ts index 03a1a99815..07be89b852 100644 --- a/supabase/functions/_backend/private/admin_stats.ts +++ b/supabase/functions/_backend/private/admin_stats.ts @@ -317,7 +317,7 @@ app.post('/', middlewareAuth, async (c) => { details = detailsCache.get(org.subscription_id) ?? null } else { - details = await getCancellationDetails(c, org.subscription_id) + details = await getCancellationDetails(c, org.subscription_id, org.billing_account) detailsCache.set(org.subscription_id, details) } } diff --git a/supabase/functions/_backend/private/credits.ts b/supabase/functions/_backend/private/credits.ts index 2c2ea5aae7..7c3b9b2797 100644 --- a/supabase/functions/_backend/private/credits.ts +++ b/supabase/functions/_backend/private/credits.ts @@ -13,7 +13,7 @@ import { parseBody, simpleError, useCors } from '../utils/hono.ts' import { getClaimsFromJWT, middlewareAuth } from '../utils/hono_jwt.ts' import { cloudlog, cloudlogErr } from '../utils/logging.ts' import { checkPermission } from '../utils/rbac.ts' -import { createOneTimeCheckout, getCreditCheckoutDetails, getStripe, isStripeEmulatorEnabled } from '../utils/stripe.ts' +import { createOneTimeCheckout, getBillingAccountForCustomer, getCreditCheckoutDetails, getStripe, isStripeEmulatorEnabled, planProductIdOrFilter, resolvePlanCreditProductId } from '../utils/stripe.ts' import { supabaseAdmin, supabaseClient } from '../utils/supabase.ts' import { getEnv } from '../utils/utils.ts' @@ -231,6 +231,7 @@ async function getScopedCreditSteps(c: AppContext, orgId?: string): Promise { + const billingAccount = await getBillingAccountForCustomer(c, customerId) const supabase = supabaseClient(c, token) const { data: stripeInfo, error: stripeInfoError } = await supabase .from('stripe_info') @@ -249,23 +250,23 @@ async function getCreditTopUpProductId(c: AppContext, customerId: string, token: const productId = await getFallbackCreditProductId(c, customerId, async () => { const { data, error } = await supabase .from('plans') - .select('credit_id') + .select('*') .eq('name', 'Solo') .single() if (error) throw error - return data ?? null + return data ? { credit_id: resolvePlanCreditProductId(data, billingAccount) } : null }) return { productId } } const { data: plan, error: planError } = await supabase .from('plans') - .select('credit_id, name') - .eq('stripe_id', stripeInfo.product_id) + .select('*') + .or(planProductIdOrFilter(stripeInfo.product_id)) .single() - if (planError || !plan?.credit_id) { + if (planError || !plan) { cloudlogErr({ requestId: c.get('requestId'), message: 'credit_top_up_product_missing', @@ -276,17 +277,32 @@ async function getCreditTopUpProductId(c: AppContext, customerId: string, token: const productId = await getFallbackCreditProductId(c, customerId, async () => { const { data, error } = await supabase .from('plans') - .select('credit_id') + .select('*') .eq('name', 'Solo') .single() if (error) throw error - return data ?? null + return data ? { credit_id: resolvePlanCreditProductId(data, billingAccount) } : null }) return { productId } } - return { productId: plan.credit_id } + const productId = resolvePlanCreditProductId(plan, billingAccount) + if (!productId) { + const fallbackProductId = await getFallbackCreditProductId(c, customerId, async () => { + const { data, error } = await supabase + .from('plans') + .select('*') + .eq('name', 'Solo') + .single() + if (error) + throw error + return data ? { credit_id: resolvePlanCreditProductId(data, billingAccount) } : null + }) + return { productId: fallbackProductId } + } + + return { productId } } async function resolveOrgStripeContext(c: AppContext, orgId: string) { @@ -593,7 +609,8 @@ app.post('/complete-top-up', middlewareAuth, async (c) => { const { customerId, token } = await resolveOrgStripeContext(c, body.orgId) const supabase = supabaseClient(c, token) - const stripe = getStripe(c) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + const stripe = getStripe(c, billingAccount) const session = await resolveCheckoutSession(c, stripe, supabase, body.orgId, customerId, body.sessionId) const resolvedSessionId = session.id @@ -609,7 +626,7 @@ app.post('/complete-top-up', middlewareAuth, async (c) => { const { productId } = await getCreditTopUpProductId(c, customerId, token) const paymentIntentId = getCheckoutSessionPaymentIntentId(session) - const { creditQuantity, itemsSummary } = await getCreditCheckoutDetails(c, session, productId) + const { creditQuantity, itemsSummary } = await getCreditCheckoutDetails(c, session, productId, billingAccount) if (creditQuantity <= 0) throw simpleError('credit_product_not_found', 'Checkout session does not include the credit product') diff --git a/supabase/functions/_backend/public/organization/post.ts b/supabase/functions/_backend/public/organization/post.ts index 7a9eaa3e9c..3f7a024656 100644 --- a/supabase/functions/_backend/public/organization/post.ts +++ b/supabase/functions/_backend/public/organization/post.ts @@ -8,6 +8,7 @@ import { closeClient, getPgClient } from '../../utils/pg.ts' import { assertJwtMfaAssurance } from '../../utils/jwt_mfa_assurance.ts' import { supabaseAdmin, supabaseWithAuth } from '../../utils/supabase.ts' import { parseOrgOnboardingDevelopmentEnvironment, parseOrgOnboardingIntent } from '../../utils/org_onboarding_intent.ts' +import { getNewCustomersBillingAccount, getPlanProductId } from '../../utils/stripe.ts' import { normalizeWebsiteUrl } from './website.ts' const MAX_ESTIMATED_MAU = 1_000_000 @@ -34,23 +35,32 @@ interface PgTransactionClient { } async function getInitialPlanForMau(c: Context, estimatedMau: number) { + const billingAccount = getNewCustomersBillingAccount(c) const adminClient = supabaseAdmin(c) const { data: plan, error } = await adminClient .from('plans') - .select('name, stripe_id, mau') + .select('name, stripe_id, stripe_id_us, mau') .gte('mau', estimatedMau) .order('mau', { ascending: true }) .limit(1) .single() - if (error || !plan?.stripe_id) { - throw simpleError('cannot_get_plan', 'Cannot get plan', { error: error?.message, estimatedMau }) + if (error || !plan) { + throw simpleError('cannot_get_plan', 'Cannot get plan', { error: error?.message, estimatedMau, billingAccount }) + } + + try { + getPlanProductId(plan, billingAccount) + } + catch { + throw simpleError('cannot_get_plan', 'Cannot get plan', { estimatedMau, billingAccount, plan: plan.name }) } return plan } async function createPendingStripeInfo(c: Context, orgId: string, estimatedMau: number) { + const billingAccount = getNewCustomersBillingAccount(c) const plan = await getInitialPlanForMau(c, estimatedMau) const pendingCustomerId = `pending_${orgId}` const trialAt = new Date() @@ -60,7 +70,8 @@ async function createPendingStripeInfo(c: Context, orgId .from('stripe_info') .insert({ customer_id: pendingCustomerId, - product_id: plan.stripe_id, + product_id: getPlanProductId(plan, billingAccount), + billing_account: billingAccount, trial_at: trialAt.toISOString(), status: null, is_good_plan: true, diff --git a/supabase/functions/_backend/triggers/stripe_event.ts b/supabase/functions/_backend/triggers/stripe_event.ts index 929479282d..2edcc25339 100644 --- a/supabase/functions/_backend/triggers/stripe_event.ts +++ b/supabase/functions/_backend/triggers/stripe_event.ts @@ -10,6 +10,7 @@ import { isBentoConfigured, syncBentoSubscriberTags, trackBentoEvent } from '../ import { purgeOnPremCacheForOrg, purgePlanCacheForOrg } from '../utils/cloudflare_cache_purge.ts' import { handleAutoTopUpPaymentIntent } from '../utils/credit_auto_top_up.ts' import { getFallbackCreditProductId } from '../utils/credits.ts' +import { getRetryablePostgrestStatus, isRetryablePostgrestError } from '../utils/retry.ts' import { BRES, quickError, simpleError } from '../utils/hono.ts' import { middlewareStripeWebhook } from '../utils/hono_middleware_stripe.ts' import { cloudlog, cloudlogErr } from '../utils/logging.ts' @@ -17,13 +18,46 @@ import { getOrgAdminMemberEmailsForTags } from '../utils/org_email_notifications import { closeClient, getDrizzleClient, getPgClient } from '../utils/pg.ts' import * as schema from '../utils/postgres_schema.ts' import { groupIdentifyPosthog } from '../utils/posthog.ts' -import { ensureCustomerMetadata, getCreditCheckoutDetails, getStripe, syncStripeCustomerCountry } from '../utils/stripe.ts' +import type { BillingAccount } from '../utils/stripe_billing.ts' +import { ensureCustomerMetadata, getBillingAccountForCustomer, getCreditCheckoutDetails, getStripe, isStripeConfiguredForAccount, normalizeBillingAccount, planProductIdOrFilter, resolvePlanCreditProductId, syncStripeCustomerCountry } from '../utils/stripe.ts' import { buildTransferInvoiceFooter, getTransferInvoiceFooterUpdate, isTransferInvoice, normalizeBillingEmail, shouldStampTransferInvoiceFooter, TRANSFER_INVOICE_FOOTER, TRANSFER_INVOICE_FOOTER_MAX_LENGTH } from '../utils/stripe_event.ts' import { customerToSegmentOrg, supabaseAdmin } from '../utils/supabase.ts' import { sendEventToTracking } from '../utils/tracking.ts' -import { backgroundTask, isStripeConfigured } from '../utils/utils.ts' +import { backgroundTask } from '../utils/utils.ts' -export const app = new Hono() +function getWebhookBillingAccount(c: Context): BillingAccount { + return c.get('stripeBillingAccount') ?? 'ee' +} + +async function assertStripeBillingAccount( + c: Context, + customer: Pick, +) { + const expected = getWebhookBillingAccount(c) + const actual = normalizeBillingAccount(customer.billing_account) + if (actual !== expected) { + cloudlogErr({ + requestId: c.get('requestId'), + message: 'Stripe webhook billing_account mismatch', + customerId: customer.customer_id, + expected, + actual, + }) + throw simpleError('webhook_billing_account_mismatch', 'Stripe webhook billing account mismatch', { + customerId: customer.customer_id, + expected, + actual, + }) + } +} + +export function createStripeEventApp(webhookBillingAccount: BillingAccount = 'ee') { + const app = new Hono() + app.post('/', middlewareStripeWebhook(webhookBillingAccount), stripeEventHandler) + return app +} + +export const app = createStripeEventApp('ee') interface Org { id: string @@ -43,7 +77,7 @@ type StripeInfoRevenueState = { product_id?: string | null status?: Database['public']['Enums']['stripe_status'] | null } | null | undefined -type RevenuePlanRow = Pick +type RevenuePlanRow = Pick type RevenuePlanKey = 'solo' | 'maker' | 'team' | 'enterprise' type RevenuePlanBreakdown = Record type ChurnReason = 'past_due_unresolved' @@ -200,15 +234,17 @@ function compactMetadata(metadata: Record) { ) as Record } +type PlanPriceIdsRow = Pick + function getPlanType( - plan: Pick, + plan: PlanPriceIdsRow, priceId: string | null | undefined, ) { if (!priceId) return undefined - if (plan.price_m_id === priceId) + if (plan.price_m_id === priceId || plan.price_m_id_us === priceId) return 'monthly' - if (plan.price_y_id === priceId) + if (plan.price_y_id === priceId || plan.price_y_id_us === priceId) return 'yearly' return undefined } @@ -229,11 +265,14 @@ function getSubscriptionTrackingState( function buildSubscriptionEventMetadata( stripeData: Pick, - currentPlan: Pick, - previousPlan?: Pick | null, + currentPlan: Pick, + previousPlan?: Pick | null, ) { const currentPlanType = getPlanType(currentPlan, stripeData.data.price_id) - const fallbackPreviousPlan = stripeData.previousProductId === currentPlan.stripe_id ? currentPlan : previousPlan + const fallbackPreviousPlan = stripeData.previousProductId === currentPlan.stripe_id + || stripeData.previousProductId === currentPlan.stripe_id_us + ? currentPlan + : previousPlan const previousPlanType = fallbackPreviousPlan ? getPlanType(fallbackPreviousPlan, stripeData.previousPriceId) : undefined @@ -281,10 +320,10 @@ function getPlanMrr(plan: RevenuePlanRow | null | undefined, priceId: string | n if (!plan || !priceId) return 0 - if (plan.price_m_id === priceId) + if (plan.price_m_id === priceId || plan.price_m_id_us === priceId) return Number(plan.price_m) || 0 - if (plan.price_y_id === priceId) + if (plan.price_y_id === priceId || plan.price_y_id_us === priceId) return (Number(plan.price_y) || 0) / 12 return 0 @@ -294,7 +333,7 @@ function getPlanByProductId(plans: RevenuePlanRow[], productId: string | null | if (!productId) return null - return plans.find(plan => plan.stripe_id === productId) ?? null + return plans.find(plan => plan.stripe_id === productId || plan.stripe_id_us === productId) ?? null } async function lookupOrgCreatorEmail( @@ -318,10 +357,13 @@ async function lookupOrgCreatorEmail( } async function retrieveStripeCustomerBillingEmail(c: Context, customerId: string): Promise { - if (!customerId || !isStripeConfigured(c)) + if (!customerId) + return null + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return null - const customer = await getStripe(c).customers.retrieve(customerId) + const customer = await getStripe(c, billingAccount).customers.retrieve(customerId) if ('deleted' in customer && customer.deleted) return null @@ -465,7 +507,7 @@ export async function syncBillingBentoTagsFromStoredStripeInfo(c: Context, org: return const plans = await getBillingPlans(c) - const plan = plans.find(candidate => candidate.stripe_id === stripeInfo.product_id) ?? null + const plan = plans.find(candidate => candidate.stripe_id === stripeInfo.product_id || candidate.stripe_id_us === stripeInfo.product_id) ?? null const trialPlanNames = plans.map(candidate => candidate.name) const segment = await customerToSegmentOrg(c, org.id, stripeInfo.price_id, plan, trialPlanNames) await syncBillingBentoTags(c, org, customerId, segment) @@ -600,7 +642,7 @@ function isStaleStripeEvent( async function getRevenuePlans(c: Context): Promise { const { data: plans, error } = await supabaseAdmin(c) .from('plans') - .select('name, stripe_id, price_m, price_y, price_m_id, price_y_id') + .select('name, stripe_id, stripe_id_us, price_m, price_y, price_m_id, price_y_id, price_m_id_us, price_y_id_us') .in('name', ['Solo', 'Maker', 'Team', 'Enterprise']) if (error) { @@ -812,76 +854,80 @@ async function writePaidAtAtomically(c: Context, customerId: string, eventOccurr } async function getCreditTopUpProductIdFromCustomer(c: Context, customerId: string): Promise { - const pgClient = getPgClient(c, true) - const drizzleClient = getDrizzleClient(pgClient) - - try { - let stripeInfoError: unknown | null = null - let stripeInfo: { product_id: string | null } | undefined - try { - [stripeInfo] = await drizzleClient - .select({ product_id: schema.stripe_info.product_id }) - .from(schema.stripe_info) - .where(eq(schema.stripe_info.customer_id, customerId)) - .limit(1) - } - catch (error) { - stripeInfoError = error - } + const billingAccount = await getBillingAccountForCustomer(c, customerId) + const fetchFallbackPlan = async () => { + const { data, error } = await supabaseAdmin(c) + .from('plans') + .select('*') + .eq('name', 'Solo') + .single() + if (error) + throw error + return data ? { credit_id: resolvePlanCreditProductId(data, billingAccount) } : null + } + const { data: stripeInfo, error: stripeInfoError } = await supabaseAdmin(c) + .from('stripe_info') + .select('product_id') + .eq('customer_id', customerId) + .maybeSingle() - if (stripeInfoError || !stripeInfo?.product_id) { - cloudlog({ - requestId: c.get('requestId'), - message: 'credit_plan_missing', + if (stripeInfoError) { + if (isRetryablePostgrestError(stripeInfoError)) { + const retryStatus = getRetryablePostgrestStatus(stripeInfoError) ?? 503 + throw quickError(retryStatus, 'stripe_info_lookup_failed', 'Temporary stripe_info lookup failure', { customerId, - error: stripeInfoError, - }) - return await getFallbackCreditProductId(c, customerId, async () => { - const [fallbackPlan] = await drizzleClient - .select({ credit_id: schema.plans.credit_id }) - .from(schema.plans) - .where(eq(schema.plans.name, 'Solo')) - .limit(1) - return fallbackPlan ?? null + stripeInfoError, }) } + throw simpleError('stripe_info_lookup_failed', 'stripe_info lookup failed', { customerId, stripeInfoError }) + } - let planError: unknown | null = null - let plan: { credit_id: string | null } | undefined - try { - [plan] = await drizzleClient - .select({ credit_id: schema.plans.credit_id }) - .from(schema.plans) - .where(eq(schema.plans.stripe_id, stripeInfo.product_id)) - .limit(1) - } - catch (error) { - planError = error - } + if (!stripeInfo?.product_id) { + cloudlog({ + requestId: c.get('requestId'), + message: 'credit_plan_missing', + customerId, + }) + return await getFallbackCreditProductId(c, customerId, fetchFallbackPlan) + } - if (planError || !plan?.credit_id) { - cloudlog({ - requestId: c.get('requestId'), - message: 'credit_top_up_product_missing', + const { data: plan, error: planError } = await supabaseAdmin(c) + .from('plans') + .select('*') + .or(planProductIdOrFilter(stripeInfo.product_id)) + .maybeSingle() + + if (planError) { + if (isRetryablePostgrestError(planError)) { + const retryStatus = getRetryablePostgrestStatus(planError) ?? 503 + throw quickError(retryStatus, 'credit_plan_lookup_failed', 'Temporary credit plan lookup failure', { customerId, planStripeId: stripeInfo.product_id, - error: planError, - }) - return await getFallbackCreditProductId(c, customerId, async () => { - const [fallbackPlan] = await drizzleClient - .select({ credit_id: schema.plans.credit_id }) - .from(schema.plans) - .where(eq(schema.plans.name, 'Solo')) - .limit(1) - return fallbackPlan ?? null + planError, }) } - - return plan.credit_id + throw simpleError('credit_plan_lookup_failed', 'credit plan lookup failed', { + customerId, + planStripeId: stripeInfo.product_id, + planError, + }) } - finally { - closeClient(c, pgClient) + + if (!plan) { + cloudlog({ + requestId: c.get('requestId'), + message: 'credit_top_up_product_missing', + customerId, + planStripeId: stripeInfo.product_id, + }) + return await getFallbackCreditProductId(c, customerId, fetchFallbackPlan) } + + const creditProductId = resolvePlanCreditProductId(plan, billingAccount) + if (!creditProductId) + return await getFallbackCreditProductId(c, customerId, fetchFallbackPlan) + + return creditProductId } async function handleCheckoutSessionCompleted( @@ -926,8 +972,9 @@ async function handleCheckoutSessionCompleted( : null const creditProductId = metadataProductId ?? await getCreditTopUpProductIdFromCustomer(c, customerId) + const billingAccount = getWebhookBillingAccount(c) - const { creditQuantity, itemsSummary } = await getCreditCheckoutDetails(c, session, creditProductId) + const { creditQuantity, itemsSummary } = await getCreditCheckoutDetails(c, session, creditProductId, billingAccount) if (creditQuantity <= 0) { throw simpleError('credit_product_not_found', 'Checkout session does not include the credit product', { @@ -1040,7 +1087,8 @@ async function invoiceCreatedOrUpdated(c: Context, stripeEvent: Stripe.InvoiceCr } try { - const liveInvoice = await getStripe(c).invoices.retrieve(eventInvoice.id) + const billingAccount = getWebhookBillingAccount(c) + const liveInvoice = await getStripe(c, billingAccount).invoices.retrieve(eventInvoice.id) const footer = getTransferInvoiceFooterUpdate(liveInvoice) if (!footer) { cloudlog({ @@ -1057,7 +1105,7 @@ async function invoiceCreatedOrUpdated(c: Context, stripeEvent: Stripe.InvoiceCr return c.json(BRES) } - await getStripe(c).invoices.update(eventInvoice.id, { + await getStripe(c, billingAccount).invoices.update(eventInvoice.id, { footer, }) cloudlog({ @@ -1089,14 +1137,14 @@ async function invoiceUpcoming(c: Context, org: Org, stripeEvent: Stripe.Invoice if (stripeData.data.product_id) { const { data: plan } = await supabaseAdmin(c) .from('plans') - .select('name, price_y_id') - .eq('stripe_id', stripeData.data.product_id) + .select('name, price_y_id, price_y_id_us') + .or(planProductIdOrFilter(stripeData.data.product_id)) .single() if (!plan) { throw simpleError('failed_to_get_plan', 'failed to get plan', { stripeData }) } planName = plan.name - if (plan.price_y_id === stripeData.data.price_id) { + if (plan.price_y_id === stripeData.data.price_id || plan.price_y_id_us === stripeData.data.price_id) { planType = 'yearly' } } @@ -1143,7 +1191,7 @@ async function createdOrUpdated( const { data: plan } = await supabaseAdmin(c) .from('plans') .select() - .eq('stripe_id', stripeData.data.product_id) + .or(planProductIdOrFilter(stripeData.data.product_id!)) .single() if (plan) { const trackingState = getSubscriptionTrackingState(stripeData, status) @@ -1204,7 +1252,7 @@ async function createdOrUpdated( const previousProduct = await supabaseAdmin(c) .from('plans') .select() - .eq('stripe_id', stripeData.previousProductId) + .or(planProductIdOrFilter(stripeData.previousProductId)) .single() previousPlan = previousProduct.data const planChangeMetadata = buildSubscriptionEventMetadata(stripeData, plan, previousPlan) @@ -1232,7 +1280,7 @@ async function createdOrUpdated( } const segment = await customerToSegmentOrg(c, org.id, stripeData.data.price_id, plan, billingPlans.map(candidate => candidate.name)) - const isMonthly = plan.price_m_id === stripeData.data.price_id + const isMonthly = getPlanType(plan, stripeData.data.price_id) === 'monthly' const eventName = `user:subscribe_${statusName}:${isMonthly ? 'monthly' : 'yearly'}` const subscriptionMetadata = buildSubscriptionEventMetadata(stripeData, plan, previousPlan) await syncBillingBentoTags(c, org, stripeData.data.customer_id, segment) @@ -1460,160 +1508,269 @@ async function cancelingOrFinished( return c.json(BRES) } -app.post('/', middlewareStripeWebhook(), async (c) => { - const stripeData = c.get('stripeData')! - const stripeEvent = c.get('stripeEvent')! - const isCheckoutSession = isCheckoutSessionEvent(stripeEvent) - - if (isCustomerProfileEvent(stripeEvent)) { - await syncStripeCustomerCountry(c, stripeData.data.customer_id) - const org = await getOrgForCustomerId(c, stripeData.data.customer_id) - if (org) { - await ensureCustomerMetadata(c, stripeData.data.customer_id, org.id, org.created_by) - const billingOrg = await syncOrgManagementEmailFromStripeCustomer( - c, - org, - stripeData.data.customer_id, - stripeEvent, - ) - await syncBillingBentoTagsFromStoredStripeInfo(c, billingOrg, stripeData.data.customer_id) - } - return c.json(BRES) - } - - // find email from user with customer_id - const org = await getOrg(c, stripeData) - - await ensureCustomerMetadata(c, stripeData.data.customer_id, org.id, org.created_by) - stripeData.data.customer_country = await syncStripeCustomerCountry(c, stripeData.data.customer_id) - - if (isCheckoutSession) { - return handleCheckoutSessionCompleted(c, stripeEvent, org, stripeData.data.customer_id) - } - - if (stripeEvent.type === 'payment_intent.succeeded') { - await handleAutoTopUpPaymentIntent(c, stripeEvent, org.id) - return c.json(BRES) +async function handleCustomerProfileStripeEvent( + c: Context, + stripeEvent: Stripe.Event, + stripeData: StripeData, +) { + await syncStripeCustomerCountry(c, stripeData.data.customer_id) + const org = await getOrgForCustomerId(c, stripeData.data.customer_id) + if (org) { + await ensureCustomerMetadata(c, stripeData.data.customer_id, org.id, org.created_by) + const billingOrg = await syncOrgManagementEmailFromStripeCustomer( + c, + org, + stripeData.data.customer_id, + stripeEvent, + ) + await syncBillingBentoTagsFromStoredStripeInfo(c, billingOrg, stripeData.data.customer_id) } + return c.json(BRES) +} - const { data: customer } = await supabaseAdmin(c) +async function loadStripeCustomer(c: Context, stripeData: StripeData) { + const { data: customer, error: customerError } = await supabaseAdmin(c) .from('stripe_info') .select() .eq('customer_id', stripeData.data.customer_id) .single() + if (customerError) { + if (isRetryablePostgrestError(customerError)) { + const retryStatus = getRetryablePostgrestStatus(customerError) ?? 503 + throw quickError(retryStatus, 'stripe_info_lookup_failed', 'Temporary stripe_info lookup failure', { + stripeData, + customerError, + }) + } + throw simpleError('stripe_info_lookup_failed', 'stripe_info lookup failed', { stripeData, customerError }) + } + if (!customer) { throw simpleError('no_customer_found', 'no customer found', { stripeData }) } - if (stripeEvent.type === 'customer.source.expiring') { + return customer +} + +async function handlePaymentIntentStripeEvent( + c: Context, + stripeEvent: Stripe.Event, + org: Org, +) { + await handleAutoTopUpPaymentIntent(c, stripeEvent, org.id) + return c.json(BRES) +} + +async function handleCustomerSourceStripeEvent( + c: Context, + org: Org, + stripeEvent: Stripe.Event, +) { + if (stripeEvent.type === 'customer.source.expiring') return customerSourceExpiring(c, org) - } - else if (stripeEvent.type === 'customer.source.created') { + if (stripeEvent.type === 'customer.source.created') return customerSourceCreated(c, org, stripeEvent) - } - else if (stripeEvent.type === 'invoice.upcoming') { + return null +} + +async function handleInvoiceStripeEvent( + c: Context, + org: Org, + stripeEvent: Stripe.Event, + stripeData: StripeData, +) { + if (stripeEvent.type === 'invoice.upcoming') return invoiceUpcoming(c, org, stripeEvent, stripeData) - } - else if (stripeEvent.type === 'invoice.created' || stripeEvent.type === 'invoice.updated') { + if (stripeEvent.type === 'invoice.created' || stripeEvent.type === 'invoice.updated') return invoiceCreatedOrUpdated(c, stripeEvent) + return null +} + +async function handleChargeSucceededStripeEvent( + c: Context, + org: Org, + stripeData: StripeData, +) { + // Canonical dunning exit. Do not also emit this from subscription.updated: + // Stripe sends both, and a plan change is not proof of payment recovery. + await trackBillingBentoEvent(c, org, stripeData.data.customer_id, BENTO_CHARGE_SUCCEEDED_EVENT) + return c.json(BRES) +} + +async function handleActiveSubscriptionUpdate( + c: Context, + stripeEvent: Stripe.Event, + stripeData: StripeData, + org: Org, + customer: StripeInfoRow, +) { + const originalStatus = stripeData.data.status + const eventOccurredAtIso = new Date(stripeEvent.created * 1000).toISOString() + stripeData.data.status = 'succeeded' + const createdOrUpdatedResponse = await createdOrUpdated(c, stripeData, org, customer, eventOccurredAtIso, originalStatus) + if (createdOrUpdatedResponse) + return createdOrUpdatedResponse + return null +} + +async function handleFailedSubscriptionPayment( + c: Context, + stripeData: StripeData, + org: Org, +) { + if (await orgHasActiveUsageCredits(c, org.id)) { + cloudlog({ requestId: c.get('requestId'), message: 'Skipping failed payment email because org has active usage credits', orgId: org.id }) } - else if (stripeEvent.type === 'charge.succeeded') { - // Canonical dunning exit. Do not also emit this from subscription.updated: - // Stripe sends both, and a plan change is not proof of payment recovery. - await trackBillingBentoEvent(c, org, stripeData.data.customer_id, BENTO_CHARGE_SUCCEEDED_EVENT) + else { + await trackBillingBentoEvent(c, org, stripeData.data.customer_id, BENTO_FAILED_PAYMENT_EVENT) + } + await updateStripeInfo(c, stripeData) + return null +} + +function handleSubscriptionMissingPriceData( + c: Context, + stripeData: StripeData, +) { + cloudlog({ requestId: c.get('requestId'), message: 'Subscription webhook missing price_id or product_id', stripeData, subscriptionId: stripeData.data.subscription_id }) + return null +} + +async function handleCanceledSubscriptionStripeEvent( + c: Context, + stripeEvent: Stripe.Event, + stripeData: StripeData, + org: Org, + customer: StripeInfoRow, +) { + const eventOccurredAtIso = new Date(stripeEvent.created * 1000).toISOString() + if (isStaleStripeEvent(customer, eventOccurredAtIso)) { + cloudlog({ + requestId: c.get('requestId'), + message: 'Skipping stale Stripe cancellation event', + customerId: stripeData.data.customer_id, + eventOccurredAtIso, + currentStripeInfoLastStripeEventAt: customer?.last_stripe_event_at, + subscriptionId: stripeData.data.subscription_id, + }) return c.json(BRES) } - if (isSubscriptionUpdateStatus(stripeData.data.status) && stripeData.data.price_id && stripeData.data.product_id) { - const originalStatus = stripeData.data.status - const eventOccurredAtIso = new Date(stripeEvent.created * 1000).toISOString() + if (customer.subscription_id !== stripeData.data.subscription_id) { + cloudlog({ requestId: c.get('requestId'), message: 'Ignoring canceled/deleted webhook for subscription not in database', subscriptionInDb: customer?.subscription_id, webhookSubscription: stripeData.data.subscription_id }) + return null + } + + if (stripeData.data.subscription_anchor_end && new Date(stripeData.data.subscription_anchor_end) > new Date()) { stripeData.data.status = 'succeeded' - const createdOrUpdatedResponse = await createdOrUpdated(c, stripeData, org, customer, eventOccurredAtIso, originalStatus) - if (createdOrUpdatedResponse) - return createdOrUpdatedResponse + } + const updateData = toStripeInfoUpdate(stripeData.data) + const revenuePlans = await getRevenuePlans(c) + const revenueMovement = classifyRevenueMovement(customer, updateData, revenuePlans) + const didPersist = await persistStripeInfoAndRevenueMovement( + c, + stripeData.data.customer_id, + stripeEvent.id, + updateData, + eventOccurredAtIso, + revenueMovement, + ) + if (didPersist === 'duplicate') { + cloudlog({ + requestId: c.get('requestId'), + message: 'Skipping duplicate Stripe cancellation event', + customerId: stripeData.data.customer_id, + eventOccurredAtIso, + stripeEventId: stripeEvent.id, + subscriptionId: stripeData.data.subscription_id, + }) + return c.json(BRES) + } + if (didPersist === 'stale') { + cloudlog({ + requestId: c.get('requestId'), + message: 'Skipping stale Stripe cancellation event after row lock', + customerId: stripeData.data.customer_id, + eventOccurredAtIso, + subscriptionId: stripeData.data.subscription_id, + }) + return c.json(BRES) + } + if (didPersist === 'missing') + return quickError(404, 'canceled_customer_id_not_found', `canceled: customer_id not found`, { stripeData }) + + await didCancel(c, org, stripeData.data.customer_id) + return null +} + +async function handleSubscriptionStripeEvent( + c: Context, + stripeEvent: Stripe.Event, + stripeData: StripeData, + org: Org, + customer: StripeInfoRow, +) { + if (isSubscriptionUpdateStatus(stripeData.data.status) && stripeData.data.price_id && stripeData.data.product_id) { + const response = await handleActiveSubscriptionUpdate(c, stripeEvent, stripeData, org, customer) + if (response) + return response } else if (stripeData.data.status === 'failed') { - if (await orgHasActiveUsageCredits(c, org.id)) { - cloudlog({ requestId: c.get('requestId'), message: 'Skipping failed payment email because org has active usage credits', orgId: org.id }) - } - else { - await trackBillingBentoEvent(c, org, stripeData.data.customer_id, BENTO_FAILED_PAYMENT_EVENT) - } - // Update the database with failed status - await updateStripeInfo(c, stripeData) + await handleFailedSubscriptionPayment(c, stripeData, org) } else if (isSubscriptionUpdateStatus(stripeData.data.status) && (!stripeData.data.price_id || !stripeData.data.product_id)) { - // Subscription event without price/product data - log warning but don't process - cloudlog({ requestId: c.get('requestId'), message: 'Subscription webhook missing price_id or product_id', stripeData, subscriptionId: stripeData.data.subscription_id }) + handleSubscriptionMissingPriceData(c, stripeData) } else if (['canceled', 'deleted'].includes(stripeData.data.status ?? '')) { - const eventOccurredAtIso = new Date(stripeEvent.created * 1000).toISOString() - if (isStaleStripeEvent(customer, eventOccurredAtIso)) { - cloudlog({ - requestId: c.get('requestId'), - message: 'Skipping stale Stripe cancellation event', - customerId: stripeData.data.customer_id, - eventOccurredAtIso, - currentStripeInfoLastStripeEventAt: customer?.last_stripe_event_at, - subscriptionId: stripeData.data.subscription_id, - }) - return c.json(BRES) - } - // Check if this is the subscription currently in the database - if (customer && customer.subscription_id === stripeData.data.subscription_id) { - // Only mark as 'succeeded' if subscription is still active until period end - // Check if subscription_anchor_end is in the future - if (stripeData.data.subscription_anchor_end && new Date(stripeData.data.subscription_anchor_end) > new Date()) { - stripeData.data.status = 'succeeded' - } - const updateData = toStripeInfoUpdate(stripeData.data) - const revenuePlans = await getRevenuePlans(c) - const revenueMovement = classifyRevenueMovement(customer, updateData, revenuePlans) - // Otherwise keep it as 'canceled' since the period has ended - const didPersist = await persistStripeInfoAndRevenueMovement( - c, - stripeData.data.customer_id, - stripeEvent.id, - updateData, - eventOccurredAtIso, - revenueMovement, - ) - if (didPersist === 'duplicate') { - cloudlog({ - requestId: c.get('requestId'), - message: 'Skipping duplicate Stripe cancellation event', - customerId: stripeData.data.customer_id, - eventOccurredAtIso, - stripeEventId: stripeEvent.id, - subscriptionId: stripeData.data.subscription_id, - }) - return c.json(BRES) - } - if (didPersist === 'stale') { - cloudlog({ - requestId: c.get('requestId'), - message: 'Skipping stale Stripe cancellation event after row lock', - customerId: stripeData.data.customer_id, - eventOccurredAtIso, - subscriptionId: stripeData.data.subscription_id, - }) - return c.json(BRES) - } - if (didPersist === 'missing') - return quickError(404, 'canceled_customer_id_not_found', `canceled: customer_id not found`, { stripeData }) + const response = await handleCanceledSubscriptionStripeEvent(c, stripeEvent, stripeData, org, customer) + if (response) + return response + } + return null +} - // This is the known subscription being cancelled. - await didCancel(c, org, stripeData.data.customer_id) - } - // If it's a different subscription (not the one in DB), ignore it - // This prevents old subscription webhooks from overwriting newer active subscriptions - else { - cloudlog({ requestId: c.get('requestId'), message: 'Ignoring canceled/deleted webhook for subscription not in database', subscriptionInDb: customer?.subscription_id, webhookSubscription: stripeData.data.subscription_id }) - } +async function stripeEventHandler(c: Context) { + const stripeData = c.get('stripeData')! + const stripeEvent = c.get('stripeEvent')! + + if (isCustomerProfileEvent(stripeEvent)) { + return handleCustomerProfileStripeEvent(c, stripeEvent, stripeData) + } + + const org = await getOrg(c, stripeData) + const customer = await loadStripeCustomer(c, stripeData) + + await assertStripeBillingAccount(c, customer) + await ensureCustomerMetadata(c, stripeData.data.customer_id, org.id, org.created_by) + stripeData.data.customer_country = await syncStripeCustomerCountry(c, stripeData.data.customer_id) + + if (isCheckoutSessionEvent(stripeEvent)) { + return handleCheckoutSessionCompleted(c, stripeEvent, org, stripeData.data.customer_id) + } + + if (stripeEvent.type === 'payment_intent.succeeded') { + return handlePaymentIntentStripeEvent(c, stripeEvent, org) } + + const customerSourceResponse = await handleCustomerSourceStripeEvent(c, org, stripeEvent) + if (customerSourceResponse) + return customerSourceResponse + + const invoiceResponse = await handleInvoiceStripeEvent(c, org, stripeEvent, stripeData) + if (invoiceResponse) + return invoiceResponse + + if (stripeEvent.type === 'charge.succeeded') { + return handleChargeSucceededStripeEvent(c, org, stripeData) + } + + const subscriptionResponse = await handleSubscriptionStripeEvent(c, stripeEvent, stripeData, org, customer) + if (subscriptionResponse) + return subscriptionResponse + return cancelingOrFinished(c, stripeEvent, stripeData.data, customer) -}) +} export const stripeEventTestUtils = { BENTO_CHARGE_SUCCEEDED_EVENT, diff --git a/supabase/functions/_backend/triggers/stripe_event_us.ts b/supabase/functions/_backend/triggers/stripe_event_us.ts new file mode 100644 index 0000000000..3ce52df318 --- /dev/null +++ b/supabase/functions/_backend/triggers/stripe_event_us.ts @@ -0,0 +1,3 @@ +import { createStripeEventApp } from './stripe_event.ts' + +export const app = createStripeEventApp('us') diff --git a/supabase/functions/_backend/utils/credit_auto_top_up.ts b/supabase/functions/_backend/utils/credit_auto_top_up.ts index 3344e187c7..4cc6273d56 100644 --- a/supabase/functions/_backend/utils/credit_auto_top_up.ts +++ b/supabase/functions/_backend/utils/credit_auto_top_up.ts @@ -2,9 +2,8 @@ import type { Context } from 'hono' import Stripe from 'stripe' import { getFallbackCreditProductId } from './credits.ts' import { cloudlog, cloudlogErr } from './logging.ts' -import { getOneTimePriceId, getStripe, isStripeEmulatorEnabled } from './stripe.ts' +import { getBillingAccountForCustomer, getOneTimePriceId, getStripe, isStripeEmulatorEnabled, isStripeConfiguredForAccount, planProductIdOrFilter, resolvePlanCreditProductId, type BillingAccount } from './stripe.ts' import { supabaseAdmin } from './supabase.ts' -import { isStripeConfigured } from './utils.ts' export const MIN_AUTO_TOP_UP_THRESHOLD = 10 export const AUTO_TOP_UP_KIND = 'credit_auto_top_up' @@ -66,7 +65,8 @@ async function getAvailableCredits(c: Context, orgId: string): Promise { } export async function customerHasSavedPaymentMethod(c: Context, customerId: string): Promise { - if (!isStripeConfigured(c)) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return false try { return Boolean(await getDefaultPaymentMethodId(c, customerId)) @@ -78,7 +78,8 @@ export async function customerHasSavedPaymentMethod(c: Context, customerId: stri } async function getDefaultPaymentMethodId(c: Context, customerId: string): Promise { - const stripe = getStripe(c) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + const stripe = getStripe(c, billingAccount) const customer = await stripe.customers.retrieve(customerId) if (customer.deleted) return null @@ -104,16 +105,21 @@ async function getDefaultPaymentMethodId(c: Context, customerId: string): Promis return cards.data[0]?.id ?? null } -async function getCreditProductIdForCustomer(c: Context, customerId: string): Promise { +async function getCreditProductIdForCustomer( + c: Context, + customerId: string, + billingAccountOverride?: BillingAccount, +): Promise { + const billingAccount = billingAccountOverride ?? await getBillingAccountForCustomer(c, customerId) const loadSoloPlan = async () => { const { data, error } = await supabaseAdmin(c) .from('plans') - .select('credit_id') + .select('*') .eq('name', 'Solo') .maybeSingle() if (error) throw error - return data ?? null + return data ? { credit_id: resolvePlanCreditProductId(data, billingAccount) } : null } const { data: stripeInfo, error: stripeInfoError } = await supabaseAdmin(c) @@ -127,14 +133,18 @@ async function getCreditProductIdForCustomer(c: Context, customerId: string): Pr const { data: plan, error: planError } = await supabaseAdmin(c) .from('plans') - .select('credit_id, name') - .eq('stripe_id', stripeInfo.product_id) + .select('*') + .or(planProductIdOrFilter(stripeInfo.product_id)) .maybeSingle() - if (planError || !plan?.credit_id) + if (planError || !plan) + return await getFallbackCreditProductId(c, customerId, loadSoloPlan) + + const creditProductId = resolvePlanCreditProductId(plan, billingAccount) + if (!creditProductId) return await getFallbackCreditProductId(c, customerId, loadSoloPlan) - return plan.credit_id + return creditProductId } export async function grantCreditsFromAutoTopUpPayment( @@ -175,6 +185,9 @@ async function chargeOffSessionCredits( orgId: string, customerId: string, quantity: number, + billingAccount: BillingAccount, + productId: string, + priceId: string, ): Promise { const paymentMethodId = await getDefaultPaymentMethodId(c, customerId) if (!paymentMethodId) { @@ -182,14 +195,7 @@ async function chargeOffSessionCredits( return null } - const productId = await getCreditProductIdForCustomer(c, customerId) - const priceId = await getOneTimePriceId(c, productId) - if (!priceId) { - cloudlogErr({ requestId: c.get('requestId'), message: 'credit_auto_top_up_missing_price', orgId, productId }) - return null - } - - const stripe = getStripe(c) + const stripe = getStripe(c, billingAccount) const price = await stripe.prices.retrieve(priceId) const unitAmount = price.unit_amount if (!unitAmount || unitAmount <= 0) { @@ -296,8 +302,34 @@ export async function saveAutoTopUpSettings( } export async function maybeAutoTopUpCredits(c: Context, orgId: string): Promise { - if (!isStripeConfigured(c)) + const { data: org, error: orgError } = await supabaseAdmin(c) + .from('orgs') + .select('customer_id') + .eq('id', orgId) + .maybeSingle() + + if (orgError || !org?.customer_id) + return + + const billingAccount = await getBillingAccountForCustomer(c, org.customer_id) + if (!isStripeConfiguredForAccount(c, billingAccount)) + return + + let productId: string + let priceId: string | null + try { + productId = await getCreditProductIdForCustomer(c, org.customer_id, billingAccount) + priceId = await getOneTimePriceId(c, productId, billingAccount) + } + catch (error) { + cloudlogErr({ requestId: c.get('requestId'), message: 'credit_auto_top_up_product_lookup_failed', orgId, error }) return + } + + if (!priceId) { + cloudlogErr({ requestId: c.get('requestId'), message: 'credit_auto_top_up_missing_price', orgId, productId }) + return + } const { data: claim, error: claimError } = await supabaseAdmin(c) .rpc('try_claim_credit_auto_top_up', { p_org_id: orgId }) @@ -315,7 +347,7 @@ export async function maybeAutoTopUpCredits(c: Context, orgId: string): Promise< if (quantity < MIN_AUTO_TOP_UP_THRESHOLD) return - const paymentIntent = await chargeOffSessionCredits(c, orgId, claim.customer_id, quantity) + const paymentIntent = await chargeOffSessionCredits(c, orgId, claim.customer_id, quantity, billingAccount, productId, priceId) if (!paymentIntent || paymentIntent.status !== 'succeeded') return diff --git a/supabase/functions/_backend/utils/hono_middleware_stripe.ts b/supabase/functions/_backend/utils/hono_middleware_stripe.ts index e8d1bd527a..e2952c9c96 100644 --- a/supabase/functions/_backend/utils/hono_middleware_stripe.ts +++ b/supabase/functions/_backend/utils/hono_middleware_stripe.ts @@ -1,43 +1,47 @@ import type Stripe from 'stripe' import type { Bindings } from './cloudflare.ts' import type { StripeData } from './stripe.ts' +import type { BillingAccount } from './stripe_billing.ts' import { createFactory } from 'hono/factory' import { simpleError } from './hono.ts' import { cloudlog } from './logging.ts' import { extractDataEvent, parseStripeEvent } from './stripe_event.ts' -import { getEnv } from './utils.ts' +import { getStripeWebhookSecret } from './stripe_billing.ts' export interface MiddlewareKeyVariablesStripe { Bindings: Bindings Variables: { stripeEvent?: Stripe.Event stripeData?: StripeData + stripeBillingAccount?: BillingAccount } } export const honoFactory = createFactory() -export function middlewareStripeWebhook() { +export function middlewareStripeWebhook(billingAccount: BillingAccount = 'ee') { return honoFactory.createMiddleware(async (c, next) => { - if (!getEnv(c, 'STRIPE_WEBHOOK_SECRET')) { - cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no secret found' }) + const webhookSecret = getStripeWebhookSecret(c, billingAccount) + if (!webhookSecret) { + cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no secret found', billingAccount }) throw simpleError('webhook_error_no_secret', 'Webhook Error: no secret found') } const signature = c.req.raw.headers.get('stripe-signature') - if (!signature || !getEnv(c, 'STRIPE_WEBHOOK_SECRET')) { - cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no signature' }) + if (!signature) { + cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no signature', billingAccount }) throw simpleError('webhook_error_no_signature', 'Webhook Error: no signature') } const body = await c.req.text() - const stripeEvent = await parseStripeEvent(c, body, signature) + const stripeEvent = await parseStripeEvent(c, body, signature, billingAccount) const stripeDataEvent = extractDataEvent(c, stripeEvent) const stripeData = stripeDataEvent.data if (stripeData.customer_id === '') { - cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no customer found' }) + cloudlog({ requestId: c.get('requestId'), message: 'Webhook Error: no customer found', billingAccount }) throw simpleError('webhook_error_no_customer', 'Webhook Error: no customer found') } c.set('stripeEvent', stripeEvent) c.set('stripeData', stripeDataEvent) + c.set('stripeBillingAccount', billingAccount) await next() }) } diff --git a/supabase/functions/_backend/utils/pg.ts b/supabase/functions/_backend/utils/pg.ts index ba9ee56da8..aeb51c9689 100644 --- a/supabase/functions/_backend/utils/pg.ts +++ b/supabase/functions/_backend/utils/pg.ts @@ -2016,6 +2016,8 @@ export interface AdminPayingOrgBreakdown { paying_orgs_total: number } +const adminPlanJoinOnStripeProductId = sql`(si.billing_account = 'us' AND p.stripe_id_us = si.product_id) OR (COALESCE(si.billing_account, 'ee') <> 'us' AND p.stripe_id = si.product_id)` + export async function getAdminPayingOrgBreakdown(c: Context): Promise { const emptyResult: AdminPayingOrgBreakdown = { paying_orgs_subscription: 0, @@ -2032,7 +2034,7 @@ export async function getAdminPayingOrgBreakdown(c: Context): Promise= INTERVAL '330 days' @@ -2556,7 +2558,7 @@ export async function getAdminOrganizationInsights( ) AS has_sso FROM orgs o LEFT JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE true ${planFilter} ${billingFilter} @@ -2758,7 +2760,7 @@ export async function getAdminOrganizationInsights( SELECT COUNT(*)::int AS total FROM orgs o LEFT JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE true ${planFilter} ${billingFilter} @@ -2864,7 +2866,7 @@ export async function getAdminEnterpriseAdoption( COALESCE(si.paid_at, o.created_at)::date AS started_on FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - INNER JOIN plans p ON p.stripe_id = si.product_id + INNER JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE p.name = 'Enterprise' AND si.status = 'succeeded' ), @@ -3116,6 +3118,7 @@ export interface AdminCancelledOrganizationRow { plan_name: string | null billing_type: 'monthly' | 'yearly' | null subscription_or_signup_date: string + billing_account: 'ee' | 'us' } export interface AdminCancelledOrganizationsResult { @@ -3153,10 +3156,11 @@ export async function getAdminCancelledOrganizations( si.customer_id, si.churn_reason, si.subscription_id, + COALESCE(si.billing_account, 'ee') AS billing_account, p.name AS plan_name, CASE - WHEN si.price_id = p.price_y_id THEN 'yearly' - WHEN si.price_id = p.price_m_id THEN 'monthly' + WHEN si.price_id IN (p.price_y_id, p.price_y_id_us) THEN 'yearly' + WHEN si.price_id IN (p.price_m_id, p.price_m_id_us) THEN 'monthly' WHEN si.subscription_anchor_start IS NOT NULL AND si.subscription_anchor_end IS NOT NULL AND si.subscription_anchor_end::timestamp - si.subscription_anchor_start::timestamp >= INTERVAL '330 days' @@ -3169,7 +3173,7 @@ export async function getAdminCancelledOrganizations( COALESCE(si.paid_at, u.created_at, o.created_at) AS subscription_or_signup_date FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} LEFT JOIN users u ON u.id = o.created_by WHERE si.canceled_at IS NOT NULL ${dateFilter} @@ -3202,6 +3206,7 @@ export async function getAdminCancelledOrganizations( plan_name: row.plan_name ?? null, billing_type: row.billing_type ?? null, subscription_or_signup_date: normalizeTimestamp(row.subscription_or_signup_date) ?? '', + billing_account: row.billing_account === 'us' ? 'us' : 'ee', })) const total = Number((countResult.rows[0] as any)?.total) || 0 @@ -3284,7 +3289,7 @@ export async function getAdminTrialOrganizations( lbu.last_bundle_upload_at FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} LEFT JOIN latest_bundle_uploads lbu ON lbu.owner_org = o.id WHERE si.trial_at::date >= CURRENT_DATE AND (si.status IS NULL OR si.status != 'succeeded') @@ -3379,7 +3384,7 @@ export async function getAdminTrialPlanBreakdown( COUNT(DISTINCT o.id)::int AS trials FROM orgs o INNER JOIN stripe_info si ON si.customer_id = o.customer_id - LEFT JOIN plans p ON p.stripe_id = si.product_id + LEFT JOIN plans p ON ${adminPlanJoinOnStripeProductId} WHERE o.created_at >= ${startDay.toISOString()}::timestamptz AND o.created_at < ${endExclusive.toISOString()}::timestamptz AND si.trial_at IS NOT NULL diff --git a/supabase/functions/_backend/utils/plan-gating.ts b/supabase/functions/_backend/utils/plan-gating.ts index 24e6ea9d89..2a938dbeec 100644 --- a/supabase/functions/_backend/utils/plan-gating.ts +++ b/supabase/functions/_backend/utils/plan-gating.ts @@ -1,6 +1,7 @@ import type { Context } from 'hono' import { quickError } from './hono.ts' import { cloudlog, cloudlogErr } from './logging.ts' +import { planProductIdOrFilter } from './stripe.ts' import { getCurrentPlanNameOrg, supabaseAdmin } from './supabase.ts' function isActivePlanStatus(status: string | null | undefined): boolean { @@ -45,7 +46,7 @@ async function getActivePlanNameOrg(c: Context, orgId: string): Promise[1]>['apiVersion'] - return new Stripe(getEnv(c, 'STRIPE_SECRET_KEY'), { + return new Stripe(secretKey, { // Keep the pinned runtime API version even when the installed SDK types lag behind it. apiVersion: '2026-03-25.dahlia' as StripeApiVersion, httpClient: Stripe.createFetchHttpClient(), @@ -114,6 +181,28 @@ function getLicensedSubscriptionItem(items: Stripe.SubscriptionItem[] | undefine return items?.find(item => item.plan.usage_type === 'licensed') ?? items?.[0] ?? null } +function buildStripeContext(c: Context, billingAccount: BillingAccount) { + const configured = isStripeConfiguredForAccount(c, billingAccount) + if (!configured) { + return { + billingAccount, + configured: false as const, + stripe: null, + } + } + + return { + billingAccount, + configured: true as const, + stripe: getStripe(c, billingAccount), + } +} + +async function getStripeContextForCustomer(c: Context, customerId: string) { + const billingAccount = await getBillingAccountForCustomer(c, customerId) + return buildStripeContext(c, billingAccount) +} + function getSubscriptionProductId(c: Context, item: Stripe.SubscriptionItem | null) { if (!item) return null @@ -139,14 +228,20 @@ function getSubscriptionEndDate(subscription: Stripe.Subscription, item: Stripe. return stripeTimestampToIso(endSeconds) } -export async function getSubscriptionData(c: Context, customerId: string, subscriptionId: string | null) { +export async function getSubscriptionData(c: Context, customerId: string, subscriptionId: string | null, billingAccount?: BillingAccount) { if (!subscriptionId) return null try { cloudlog({ requestId: c.get('requestId'), message: 'Fetching subscription data', customerId, subscriptionId }) + const stripeContext = billingAccount + ? buildStripeContext(c, billingAccount) + : await getStripeContextForCustomer(c, customerId) + if (!stripeContext.configured) + return null + // Retrieve the specific subscription from Stripe - const subscription = await getStripe(c).subscriptions.retrieve(subscriptionId, { + const subscription = await stripeContext.stripe.subscriptions.retrieve(subscriptionId, { expand: ['items.data.price'], // Correct expand path for retrieve }) @@ -189,14 +284,14 @@ export async function getSubscriptionData(c: Context, customerId: string, subscr /** * Fetches cancellation details for a Stripe subscription, if available. */ -export async function getCancellationDetails(c: Context, subscriptionId: string | null): Promise { +export async function getCancellationDetails(c: Context, subscriptionId: string | null, billingAccount: BillingAccount = 'ee'): Promise { if (!subscriptionId) return null - if (!isStripeConfigured(c)) + if (!isStripeConfiguredForAccount(c, billingAccount)) return null try { - const subscription = await getStripe(c).subscriptions.retrieve(subscriptionId) + const subscription = await getStripe(c, billingAccount).subscriptions.retrieve(subscriptionId) return subscription.cancellation_details ?? null } catch (error) { @@ -205,11 +300,17 @@ export async function getCancellationDetails(c: Context, subscriptionId: string } } -async function getActiveSubscription(c: Context, customerId: string, subscriptionId: string | null) { +async function getActiveSubscription(c: Context, customerId: string, subscriptionId: string | null, billingAccount?: BillingAccount) { cloudlog({ requestId: c.get('requestId'), message: 'Stored subscription not tracked or not found, checking for others.', customerId, storedSubscriptionId: subscriptionId }) + const stripeContext = billingAccount + ? buildStripeContext(c, billingAccount) + : await getStripeContextForCustomer(c, customerId) + if (!stripeContext.configured) + return null + for (const status of TRACKED_STRIPE_SUBSCRIPTION_STATUSES) { - const subscriptions = await getStripe(c).subscriptions.list({ + const subscriptions = await stripeContext.stripe.subscriptions.list({ customer: customerId, status, limit: 1, @@ -218,7 +319,7 @@ async function getActiveSubscription(c: Context, customerId: string, subscriptio if (subscriptions.data.length > 0) { const activeSub = subscriptions.data[0] cloudlog({ requestId: c.get('requestId'), message: 'Found a tracked subscription, fetching its data.', activeSubscriptionId: activeSub.id, status: activeSub.status }) - return getSubscriptionData(c, customerId, activeSub.id) + return getSubscriptionData(c, customerId, activeSub.id, stripeContext.billingAccount) } } @@ -227,17 +328,18 @@ async function getActiveSubscription(c: Context, customerId: string, subscriptio } export async function syncSubscriptionData(c: Context, customerId: string, subscriptionId: string | null): Promise { - if (!isStripeConfigured(c)) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return try { // Get subscription data from Stripe using the ID stored in our DB - let subscriptionData = await getSubscriptionData(c, customerId, subscriptionId) + let subscriptionData = await getSubscriptionData(c, customerId, subscriptionId, billingAccount) if (!subscriptionData) { - subscriptionData = await getActiveSubscription(c, customerId, subscriptionId) + subscriptionData = await getActiveSubscription(c, customerId, subscriptionId, billingAccount) } else if (!TRACKED_STRIPE_SUBSCRIPTION_STATUSES.includes(subscriptionData.status as typeof TRACKED_STRIPE_SUBSCRIPTION_STATUSES[number])) { - const replacementSubscriptionData = await getActiveSubscription(c, customerId, subscriptionId) + const replacementSubscriptionData = await getActiveSubscription(c, customerId, subscriptionId, billingAccount) if (replacementSubscriptionData || subscriptionData.status !== 'canceled') subscriptionData = replacementSubscriptionData } @@ -312,35 +414,41 @@ export async function syncSubscriptionData(c: Context, customerId: string, subsc } export async function createPortal(c: Context, customerId: string, callbackUrl: string) { - if (!isStripeConfigured(c)) + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) return { url: '' } const allowedReturnUrl = getAllowedRedirectUrl(c, callbackUrl, 'return_url') - const session = await getStripe(c).billingPortal.sessions.create({ + const session = await stripe.billingPortal.sessions.create({ customer: customerId, return_url: allowedReturnUrl, }) return { url: session.url } } -export function updateCustomerEmail(c: Context, customerId: string, newEmail: string) { - if (!isStripeConfigured(c)) - return Promise.resolve() - return getStripe(c).customers.update(customerId, { email: newEmail, metadata: { email: newEmail } }, +export async function updateCustomerEmail(c: Context, customerId: string, newEmail: string) { + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) + return + return stripe.customers.update(customerId, { email: newEmail, metadata: { email: newEmail } }, ) } -export function updateCustomerOrganizationName(c: Context, customerId: string, newName: string) { - if (!isStripeConfigured(c)) - return Promise.resolve() - return getStripe(c).customers.update(customerId, { name: newName }) +export async function updateCustomerOrganizationName(c: Context, customerId: string, newName: string) { + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) + return + return stripe.customers.update(customerId, { name: newName }) } export async function getStripeCustomerName(c: Context, customerId: string | null | undefined): Promise { - if (!customerId || !isStripeConfigured(c)) + if (!customerId) + return undefined + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) return undefined try { - const customer = await getStripe(c).customers.retrieve(customerId) + const customer = await stripe.customers.retrieve(customerId) if (customer.deleted) return null return customer.name ?? null @@ -370,11 +478,14 @@ export function normalizeStripeCountryCode(country: string | null | undefined): } export async function getStripeCustomerCountry(c: Context, customerId: string | null | undefined): Promise { - if (!customerId || !isStripeConfigured(c)) + if (!customerId) + return undefined + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) return undefined try { - const customer = await getStripe(c).customers.retrieve(customerId) + const customer = await stripe.customers.retrieve(customerId) if (customer.deleted) return null return normalizeStripeCountryCode(customer.address?.country ?? null) @@ -386,7 +497,10 @@ export async function getStripeCustomerCountry(c: Context, customerId: string | } export async function syncStripeCustomerCountry(c: Context, customerId: string | null | undefined): Promise { - if (!customerId || !isStripeConfigured(c)) + if (!customerId) + return undefined + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return undefined const customerCountry = await getStripeCustomerCountry(c, customerId) @@ -410,16 +524,17 @@ export async function syncStripeCustomerCountry(c: Context, customerId: string | } export async function cancelSubscription(c: Context, customerId: string) { - if (!isStripeConfigured(c)) + const { stripe, configured } = await getStripeContextForCustomer(c, customerId) + if (!configured) return let succeeded = true - for await (const subscription of getStripe(c).subscriptions.list({ customer: customerId, status: 'all' })) { + for await (const subscription of stripe.subscriptions.list({ customer: customerId, status: 'all' })) { if (subscription.status === 'canceled' || subscription.status === 'incomplete_expired') continue try { - await getStripe(c).subscriptions.cancel(subscription.id) + await stripe.subscriptions.cancel(subscription.id) } catch (error) { succeeded = false @@ -429,20 +544,27 @@ export async function cancelSubscription(c: Context, customerId: string) { return succeeded } -async function getStoredPlanPriceId(c: Context, planId: string, recurrence: string): Promise { +async function getStoredPlanPriceId(c: Context, planId: string, recurrence: string, billingAccount: BillingAccount): Promise { try { - const { data, error } = await supabaseAdmin(c) + const admin = supabaseAdmin(c) + if (!admin?.from) + return null + + const baseQuery = admin .from('plans') - .select('price_m_id, price_y_id') - .eq('stripe_id', planId) - .single() + .select('price_m_id, price_y_id, price_m_id_us, price_y_id_us, stripe_id, stripe_id_us') + + const { data, error } = await (typeof baseQuery.or === 'function' + ? baseQuery.or(planProductIdOrFilter(planId)) + : baseQuery.eq('stripe_id', planId) + ).single() if (error) { cloudlogErr({ requestId: c.get('requestId'), message: 'getStoredPlanPriceId', planId, recurrence, error }) return null } - return recurrence === 'year' ? data.price_y_id : data.price_m_id + return getPlanPriceId(data, billingAccount, recurrence) } catch (error) { cloudlogErr({ requestId: c.get('requestId'), message: 'getStoredPlanPriceId', planId, recurrence, error }) @@ -450,12 +572,12 @@ async function getStoredPlanPriceId(c: Context, planId: string, recurrence: stri } } -async function getPriceIds(c: Context, planId: string, recurrence: string): Promise<{ priceId: string | null }> { +async function getPriceIds(c: Context, planId: string, recurrence: string, billingAccount: BillingAccount): Promise<{ priceId: string | null }> { let priceId = null - if (!isStripeConfigured(c)) + if (!isStripeConfiguredForAccount(c, billingAccount)) return { priceId } try { - const prices = await listPricesByProduct(c, planId) + const prices = await listPricesByProduct(c, planId, billingAccount) cloudlog({ requestId: c.get('requestId'), message: 'prices stripe', prices }) prices.data.forEach((price) => { if (price.recurring?.interval === recurrence && price.active && price.recurring?.usage_type === 'licensed') @@ -466,7 +588,7 @@ async function getPriceIds(c: Context, planId: string, recurrence: string): Prom cloudlog({ requestId: c.get('requestId'), message: 'search err', error: err }) } if (!priceId) { - priceId = await getStoredPlanPriceId(c, planId, recurrence) + priceId = await getStoredPlanPriceId(c, planId, recurrence, billingAccount) cloudlog({ requestId: c.get('requestId'), message: 'prices fallback', planId, recurrence, priceId }) } return { priceId } @@ -542,10 +664,12 @@ function getAffonsoReferralMetadata(affonsoReferral?: string | null): Record { - if (!isStripeConfigured(c)) +export async function getOneTimePriceId(c: Context, productId: string, billingAccount: BillingAccount = 'ee'): Promise { + if (!isStripeConfiguredForAccount(c, billingAccount)) return null try { - const prices = await listPricesByProduct(c, productId, true) + const prices = await listPricesByProduct(c, productId, billingAccount, true) for (const price of prices.data) { if (price.type === 'one_time' && price.active) @@ -616,10 +740,11 @@ export async function createOneTimeCheckout( datafastAttribution?: DatafastAttribution, affonsoReferral?: string | null, ) { - if (!isStripeConfigured(c)) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return { url: '' } - const priceId = await getOneTimePriceId(c, productId) + const priceId = await getOneTimePriceId(c, productId, billingAccount) if (!priceId) throw new Error(`Cannot find one-time price for product ${productId}`) @@ -627,7 +752,7 @@ export async function createOneTimeCheckout( const allowedCancelUrl = getAllowedRedirectUrl(c, cancelUrl, 'cancel_url') const successUrlWithFlag = allowedSuccessUrl.includes('?') ? `${allowedSuccessUrl}&success=true` : `${allowedSuccessUrl}?success=true` - const session = await getStripe(c).checkout.sessions.create({ + const session = await getStripe(c, billingAccount).checkout.sessions.create({ billing_address_collection: 'auto', mode: 'payment', customer: customerId, @@ -675,9 +800,9 @@ export async function createOneTimeCheckout( return { url: session.url } } -export async function getCreditCheckoutDetails(c: Context, session: Stripe.Checkout.Session, expectedProductId: string): Promise { +export async function getCreditCheckoutDetails(c: Context, session: Stripe.Checkout.Session, expectedProductId: string, billingAccount: BillingAccount = 'ee'): Promise { try { - const lineItems = await getStripe(c).checkout.sessions.listLineItems(session.id, { + const lineItems = await getStripe(c, billingAccount).checkout.sessions.listLineItems(session.id, { expand: ['data.price.product'], limit: 100, }) @@ -802,9 +927,9 @@ function customerMatchesOrg(customer: Stripe.Customer, orgId: string) { return customer.metadata?.org_id === orgId } -async function searchOrgStripeCustomer(c: Context, orgId: string) { +async function searchOrgStripeCustomer(c: Context, orgId: string, billingAccount: BillingAccount) { try { - const result = await getStripe(c).customers.search({ + const result = await getStripe(c, billingAccount).customers.search({ query: `metadata['org_id']:'${orgId.replaceAll('\'', '')}'`, limit: 10, }) @@ -822,23 +947,23 @@ async function searchOrgStripeCustomer(c: Context, orgId: string) { } } -async function findExistingOrgStripeCustomer(c: Context, orgId: string, email: string) { - const fromSearch = await searchOrgStripeCustomer(c, orgId) +async function findExistingOrgStripeCustomer(c: Context, orgId: string, email: string, billingAccount: BillingAccount) { + const fromSearch = await searchOrgStripeCustomer(c, orgId, billingAccount) if (fromSearch) return fromSearch // Search is eventually consistent; list by email is a bounded read-after-write fallback. - const listed = await getStripe(c).customers.list({ email, limit: 100 }) + const listed = await getStripe(c, billingAccount).customers.list({ email, limit: 100 }) const fromList = oldestCustomer(listed.data.filter(customer => customerMatchesOrg(customer, orgId))) if (fromList) return fromList // Email may have changed before Search indexed the original customer. - return await searchOrgStripeCustomer(c, orgId) + return await searchOrgStripeCustomer(c, orgId, billingAccount) } -export async function createCustomer(c: Context, email: string, userId: string, orgId: string, name: string) { - cloudlog({ requestId: c.get('requestId'), message: 'createCustomer', email, userId, orgId, name }) +export async function createCustomer(c: Context, email: string, userId: string, orgId: string, name: string, billingAccount: BillingAccount = 'ee') { + cloudlog({ requestId: c.get('requestId'), message: 'createCustomer', email, userId, orgId, name, billingAccount }) const baseConsoleUrl = trimTrailingSlashes(getEnv(c, 'WEBAPP_URL') || '') const metadata: Record = { user_id: userId, @@ -847,14 +972,14 @@ export async function createCustomer(c: Context, email: string, userId: string, if (baseConsoleUrl) { metadata.log_as = `${baseConsoleUrl}/log-as/${userId}` } - if (!isStripeConfigured(c)) { - cloudlog({ requestId: c.get('requestId'), message: 'createCustomer no stripe key', email, userId, name }) + if (!isStripeConfiguredForAccount(c, billingAccount)) { + cloudlog({ requestId: c.get('requestId'), message: 'createCustomer no stripe key', email, userId, name, billingAccount }) return { id: localOrgStripeCustomerId(orgId), email, name, metadata } } // Org-create queue retries must return the same customer instead of minting duplicates. let customer: Stripe.Customer try { - customer = await getStripe(c).customers.create({ + customer = await getStripe(c, billingAccount).customers.create({ email, name, metadata, @@ -868,7 +993,7 @@ export async function createCustomer(c: Context, email: string, userId: string, message: 'createCustomer idempotency mismatch, searching existing customer', orgId, }) - const existing = await findExistingOrgStripeCustomer(c, orgId, email) + const existing = await findExistingOrgStripeCustomer(c, orgId, email, billingAccount) if (!existing) { cloudlogErr({ requestId: c.get('requestId'), @@ -884,7 +1009,7 @@ export async function createCustomer(c: Context, email: string, userId: string, const supabaseLink = buildSupabaseDashboardLink(c, customer.id) if (supabaseLink) { metadata.supabase = supabaseLink - await getStripe(c).customers.update(customer.id, { metadata }) + await getStripe(c, billingAccount).customers.update(customer.id, { metadata }) } return customer } @@ -892,7 +1017,8 @@ export async function createCustomer(c: Context, email: string, userId: string, export async function ensureCustomerMetadata(c: Context, customerId: string, orgId: string, userId?: string | null) { if (!customerId) return - if (!isStripeConfigured(c)) + const billingAccount = await getBillingAccountForCustomer(c, customerId) + if (!isStripeConfiguredForAccount(c, billingAccount)) return const baseConsoleUrl = trimTrailingSlashes(getEnv(c, 'WEBAPP_URL') || '') @@ -911,17 +1037,17 @@ export async function ensureCustomerMetadata(c: Context, customerId: string, org metadata.supabase = supabaseLink try { - await getStripe(c).customers.update(customerId, { metadata }) + await getStripe(c, billingAccount).customers.update(customerId, { metadata }) } catch (error) { cloudlogErr({ requestId: c.get('requestId'), message: 'ensureCustomerMetadata', error }) } } -export async function removeOldSubscription(c: Context, subscriptionId: string) { - if (!isStripeConfigured(c)) +export async function removeOldSubscription(c: Context, subscriptionId: string, billingAccount: BillingAccount = 'ee') { + if (!isStripeConfiguredForAccount(c, billingAccount)) return Promise.resolve() - cloudlog({ requestId: c.get('requestId'), message: 'removeOldSubscription', id: subscriptionId }) - const deletedSubscription = await getStripe(c).subscriptions.cancel(subscriptionId) + cloudlog({ requestId: c.get('requestId'), message: 'removeOldSubscription', id: subscriptionId, billingAccount }) + const deletedSubscription = await getStripe(c, billingAccount).subscriptions.cancel(subscriptionId) return deletedSubscription } diff --git a/supabase/functions/_backend/utils/stripe_billing.ts b/supabase/functions/_backend/utils/stripe_billing.ts new file mode 100644 index 0000000000..371f0bc1b7 --- /dev/null +++ b/supabase/functions/_backend/utils/stripe_billing.ts @@ -0,0 +1,178 @@ +import type { Context } from 'hono' +import { cloudlogErr } from './logging.ts' +import { supabaseAdmin } from './supabase.ts' +import { getEnv } from './utils.ts' + +export type BillingAccount = 'ee' | 'us' + +export interface PlanStripeIds { + stripe_id: string + price_m_id: string + price_y_id: string + credit_id?: string + stripe_id_us?: string | null + price_m_id_us?: string | null + price_y_id_us?: string | null + credit_id_us?: string | null +} + +export function normalizeBillingAccount(value: string | null | undefined): BillingAccount { + return value === 'us' ? 'us' : 'ee' +} + +export function getNewCustomersBillingAccount(c: Context): BillingAccount { + const flag = getEnv(c, 'STRIPE_NEW_CUSTOMERS_ACCOUNT').trim().toLowerCase() + if (!flag || flag === 'ee') + return 'ee' + if (flag === 'us') + return 'us' + throw new Error(`Invalid STRIPE_NEW_CUSTOMERS_ACCOUNT value: ${JSON.stringify(flag)}`) +} + +export function getStripeSecretKeyEnvName(account: BillingAccount): string { + return account === 'us' ? 'STRIPE_SECRET_KEY_US' : 'STRIPE_SECRET_KEY' +} + +export function getStripeWebhookSecretEnvName(account: BillingAccount): string { + return account === 'us' ? 'STRIPE_WEBHOOK_SECRET_US' : 'STRIPE_WEBHOOK_SECRET' +} + +export function getStripeSecretKey(c: Context, account: BillingAccount = 'ee'): string { + return getEnv(c, getStripeSecretKeyEnvName(account)) +} + +export function getStripeWebhookSecret(c: Context, account: BillingAccount = 'ee'): string { + return getEnv(c, getStripeWebhookSecretEnvName(account)) +} + +export function isStripeConfiguredForAccount(c: Context, account: BillingAccount = 'ee'): boolean { + const secretKey = getStripeSecretKey(c, account).trim() + if (!secretKey) + return false + return secretKey.startsWith('sk_') || secretKey.startsWith('rk_') +} + +export class IncompleteUsPlanConfigError extends Error { + constructor(field: string) { + super(`Plan missing US Stripe identifier: ${field}`) + this.name = 'IncompleteUsPlanConfigError' + } +} + +function normalizeUsPlanField(value: string | null | undefined): string | null { + const trimmed = value?.trim() + return trimmed ? trimmed : null +} + +function requireUsPlanField(value: string | null | undefined, field: string): string { + const normalized = normalizeUsPlanField(value) + if (!normalized) + throw new IncompleteUsPlanConfigError(field) + return normalized +} + +export async function getBillingAccountForCustomer(c: Context, customerId: string): Promise { + const admin = supabaseAdmin(c) + if (!admin?.from) { + // Unit/emulator tests stub supabaseAdmin without a client. Production always + // has service-role access; default to ee here instead of failing checkout. + cloudlogErr({ + requestId: c.get('requestId'), + message: 'getBillingAccountForCustomer unavailable admin client, defaulting to ee', + customerId, + }) + return 'ee' + } + + const { data, error } = await admin + .from('stripe_info') + .select('billing_account') + .eq('customer_id', customerId) + .maybeSingle() + + if (error) { + cloudlogErr({ + requestId: c.get('requestId'), + message: 'getBillingAccountForCustomer', + customerId, + error, + }) + throw error + } + + return normalizeBillingAccount(data?.billing_account) +} + +export function getPlanProductId(plan: Pick, account: BillingAccount): string { + if (account === 'us') + return requireUsPlanField(plan.stripe_id_us, 'stripe_id_us') + return plan.stripe_id +} + +export function getPlanPriceId(plan: PlanStripeIds, account: BillingAccount, recurrence: string): string { + const yearly = recurrence === 'year' + if (account === 'us') { + return yearly + ? requireUsPlanField(plan.price_y_id_us, 'price_y_id_us') + : requireUsPlanField(plan.price_m_id_us, 'price_m_id_us') + } + return yearly ? plan.price_y_id : plan.price_m_id +} + +export function resolvePlanCreditProductId(plan: PlanStripeIds, account: BillingAccount): string { + if (account === 'us') + return normalizeUsPlanField(plan.credit_id_us) ?? '' + return plan.credit_id?.trim() ?? '' +} + +export function getPlanCreditProductId(plan: PlanStripeIds, account: BillingAccount): string { + if (account === 'us') + return requireUsPlanField(plan.credit_id_us, 'credit_id_us') + return plan.credit_id ?? '' +} + +// Stripe product ids: prod_ + alphanumeric/underscore/hyphen (see Stripe 2018-05-21 id rules). +const STRIPE_PRODUCT_ID_REGEX = /^prod_[A-Za-z0-9_-]+$/ + +export function planProductIdOrFilter(productId: string): string { + if (!STRIPE_PRODUCT_ID_REGEX.test(productId)) + throw new Error('invalid_stripe_product_id') + + return `stripe_id.eq.${productId},stripe_id_us.eq.${productId}` +} + +export async function findPlanByProductId(c: Context, productId: string) { + try { + const admin = supabaseAdmin(c) + if (!admin?.from) + return { data: null, error: null } + + return await admin + .from('plans') + .select('*') + .or(planProductIdOrFilter(productId)) + .maybeSingle() + } + catch (error) { + cloudlogErr({ + requestId: c.get('requestId'), + message: 'findPlanByProductId unavailable admin client', + productId, + error, + }) + return { data: null, error } + } +} + +export async function resolveCheckoutPlanProductId( + c: Context, + planProductId: string, + billingAccount: BillingAccount, +): Promise { + const { data: plan, error } = await findPlanByProductId(c, planProductId) + if (error) + throw error + if (!plan) + return planProductId + return getPlanProductId(plan, billingAccount) +} diff --git a/supabase/functions/_backend/utils/stripe_event.ts b/supabase/functions/_backend/utils/stripe_event.ts index c3d51c7aa8..e9c109efac 100644 --- a/supabase/functions/_backend/utils/stripe_event.ts +++ b/supabase/functions/_backend/utils/stripe_event.ts @@ -1,14 +1,15 @@ import type { Context } from 'hono' import type { StripeData } from './stripe.ts' +import type { BillingAccount } from './stripe_billing.ts' import Stripe from 'stripe' import { cloudlog, cloudlogErr } from './logging.ts' import { getStripe, parsePriceIds } from './stripe.ts' -import { getEnv } from './utils.ts' +import { getStripeWebhookSecret } from './stripe_billing.ts' -export function parseStripeEvent(c: Context, body: string, signature: string) { - const webhookKey = getEnv(c, 'STRIPE_WEBHOOK_SECRET') +export function parseStripeEvent(c: Context, body: string, signature: string, billingAccount: BillingAccount = 'ee') { + const webhookKey = getStripeWebhookSecret(c, billingAccount) - return getStripe(c).webhooks.constructEventAsync( + return getStripe(c, billingAccount).webhooks.constructEventAsync( body, signature, webhookKey, diff --git a/supabase/functions/_backend/utils/stripe_org.ts b/supabase/functions/_backend/utils/stripe_org.ts index 0beba260e7..34700ee99d 100644 --- a/supabase/functions/_backend/utils/stripe_org.ts +++ b/supabase/functions/_backend/utils/stripe_org.ts @@ -1,7 +1,7 @@ import type { Context } from 'hono' import type { Database } from './supabase.types.ts' import { cloudlog, cloudlogErr } from './logging.ts' -import { createCustomer } from './stripe.ts' +import { createCustomer, getNewCustomersBillingAccount, getPlanProductId, normalizeBillingAccount, planProductIdOrFilter, type BillingAccount } from './stripe.ts' import { getDefaultPlan, getStripeCustomer, supabaseAdmin } from './supabase.ts' type OrgRow = Database['public']['Tables']['orgs']['Row'] @@ -85,20 +85,37 @@ async function deleteUnusedStripeInfo(c: Context, customerId: string) { } } -async function resolveTrialPlan(c: Context, org: OrgRow) { - const pendingPlan = isPendingStripeCustomerId(org.customer_id) - ? await getStripeCustomer(c, org.customer_id!).then(async (pendingStripeInfo) => { - if (!pendingStripeInfo?.product_id) +async function resolveTrialPlan(c: Context, org: OrgRow, billingAccount: BillingAccount) { + const existingStripeInfoPlan = org.customer_id && !isProvisionedStripeCustomerId(org.customer_id) + ? await getStripeCustomer(c, org.customer_id).then(async (stripeInfo) => { + if (!stripeInfo?.product_id) return null - const { data } = await supabaseAdmin(c) + const { data, error } = await supabaseAdmin(c) .from('plans') .select() - .eq('stripe_id', pendingStripeInfo.product_id) - .single() + .or(planProductIdOrFilter(stripeInfo.product_id)) + .maybeSingle() + if (error) + throw error return data }) : null - return pendingPlan ?? await getDefaultPlan(c) + const plan = existingStripeInfoPlan ?? await getDefaultPlan(c) + if (!plan) + return null + return { + ...plan, + stripe_id: getPlanProductId(plan, billingAccount), + } +} + +async function resolveBillingAccountForCreate(c: Context, org: OrgRow): Promise { + if (org.customer_id && !isProvisionedStripeCustomerId(org.customer_id)) { + const stripeInfo = await getStripeCustomer(c, org.customer_id) + if (stripeInfo?.billing_account) + return normalizeBillingAccount(stripeInfo.billing_account) + } + return getNewCustomersBillingAccount(c) } async function trialPlanNameForCustomer(c: Context, customerId: string, fallbackPlanName?: string | null) { @@ -107,7 +124,7 @@ async function trialPlanNameForCustomer(c: Context, customerId: string, fallback const { data } = await supabaseAdmin(c) .from('plans') .select('name') - .eq('stripe_id', stripeInfo.product_id) + .or(planProductIdOrFilter(stripeInfo.product_id)) .maybeSingle() if (data?.name) return data.name @@ -131,16 +148,17 @@ export async function createStripeCustomer(c: Context, org: OrgRow) { return await trialPlanNameForCustomer(c, current.customer_id!) } - const selectedPlan = await resolveTrialPlan(c, current) + const billingAccount = await resolveBillingAccountForCreate(c, current) + const selectedPlan = await resolveTrialPlan(c, current, billingAccount) if (!selectedPlan) { cloudlog({ requestId: c.get('requestId'), message: 'no default plan' }) throw new Error('no default plan') } - const customer = await createCustomer(c, current.management_email, current.created_by, current.id, current.name) + const customer = await createCustomer(c, current.management_email, current.created_by, current.id, current.name, billingAccount) const trial_at = new Date() trial_at.setDate(trial_at.getDate() + 15) - cloudlog({ requestId: c.get('requestId'), message: 'createInfo', plan: selectedPlan, customer }) + cloudlog({ requestId: c.get('requestId'), message: 'createInfo', plan: selectedPlan, customer, billingAccount }) const { error: createInfoError } = await supabaseAdmin(c) .from('stripe_info') @@ -148,6 +166,7 @@ export async function createStripeCustomer(c: Context, org: OrgRow) { product_id: selectedPlan.stripe_id, customer_id: customer.id, trial_at: trial_at.toISOString(), + billing_account: billingAccount, }) if (createInfoError && !isUniqueViolation(createInfoError)) { cloudlog({ requestId: c.get('requestId'), message: 'createInfoError', createInfoError }) diff --git a/supabase/functions/_backend/utils/supabase.ts b/supabase/functions/_backend/utils/supabase.ts index d0360ec5ae..e6026b6ff3 100644 --- a/supabase/functions/_backend/utils/supabase.ts +++ b/supabase/functions/_backend/utils/supabase.ts @@ -1254,12 +1254,14 @@ function processSegments(segmentsObj: any): { segments: string[], deleteSegments } export async function getStripeCustomer(c: Context, customerId: string) { - const { data: stripeInfo } = await supabaseAdmin(c) + const { data, error } = await supabaseAdmin(c) .from('stripe_info') .select('*') .eq('customer_id', customerId) - .single() - return stripeInfo + .maybeSingle() + if (error) + throw error + return data } export async function getDefaultPlan(c: Context) { diff --git a/supabase/functions/_backend/utils/supabase.types.ts b/supabase/functions/_backend/utils/supabase.types.ts index b1fa498a92..92820d213e 100644 --- a/supabase/functions/_backend/utils/supabase.types.ts +++ b/supabase/functions/_backend/utils/supabase.types.ts @@ -2698,6 +2698,7 @@ export type Database = { build_time_unit: number created_at: string credit_id: string + credit_id_us: string | null description: string id: string market_desc: string | null @@ -2708,8 +2709,11 @@ export type Database = { price_m_id: string price_y: number price_y_id: string + price_y_id_us: string | null + price_m_id_us: string | null storage: number stripe_id: string + stripe_id_us: string | null updated_at: string } Insert: { @@ -2717,6 +2721,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2727,8 +2732,11 @@ export type Database = { price_m_id: string price_y?: number price_y_id: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Update: { @@ -2736,6 +2744,7 @@ export type Database = { build_time_unit?: number created_at?: string credit_id?: string + credit_id_us?: string | null description?: string id?: string market_desc?: string | null @@ -2746,8 +2755,11 @@ export type Database = { price_m_id?: string price_y?: number price_y_id?: string + price_y_id_us?: string | null + price_m_id_us?: string | null storage?: number stripe_id?: string + stripe_id_us?: string | null updated_at?: string } Relationships: [] @@ -3097,6 +3109,7 @@ export type Database = { stripe_info: { Row: { bandwidth_exceeded: boolean | null + billing_account: string build_time_exceeded: boolean | null canceled_at: string | null churn_reason: string | null @@ -3125,6 +3138,7 @@ export type Database = { } Insert: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -3153,6 +3167,7 @@ export type Database = { } Update: { bandwidth_exceeded?: boolean | null + billing_account?: string build_time_exceeded?: boolean | null canceled_at?: string | null churn_reason?: string | null @@ -3179,15 +3194,7 @@ export type Database = { updated_at?: string upgraded_at?: string | null } - Relationships: [ - { - foreignKeyName: "stripe_info_product_id_fkey" - columns: ["product_id"] - isOneToOne: false - referencedRelation: "plans" - referencedColumns: ["stripe_id"] - }, - ] + Relationships: [] } tmp_users: { Row: { diff --git a/supabase/functions/triggers/index.ts b/supabase/functions/triggers/index.ts index ae1eff6dbb..0e7478d28f 100644 --- a/supabase/functions/triggers/index.ts +++ b/supabase/functions/triggers/index.ts @@ -32,6 +32,7 @@ import { app as pluginNotifications } from '../_backend/triggers/plugin_notifica import { app as queue_consumer } from '../_backend/triggers/queue_consumer.ts' import { app as send_email } from '../_backend/triggers/send_email.ts' import { app as stripe_event } from '../_backend/triggers/stripe_event.ts' +import { app as stripe_event_us } from '../_backend/triggers/stripe_event_us.ts' import { app as webhook_delivery } from '../_backend/triggers/webhook_delivery.ts' import { app as webhook_dispatcher } from '../_backend/triggers/webhook_dispatcher.ts' import { createAllCatch, createHono } from '../_backend/utils/hono.ts' @@ -73,6 +74,7 @@ appGlobal.route('/on_version_update', on_version_update) appGlobal.route('/on_version_delete', on_version_delete) appGlobal.route('/on_manifest_create', on_manifest_create) appGlobal.route('/stripe_event', stripe_event) +appGlobal.route('/stripe_event_us', stripe_event_us) appGlobal.route('/on_organization_create', on_organization_create) appGlobal.route('/on_org_update', on_org_update) appGlobal.route('/cron_stat_app', cron_stat_app) diff --git a/supabase/migrations/20260923105200_dual_stripe_billing_account.sql b/supabase/migrations/20260923105200_dual_stripe_billing_account.sql new file mode 100644 index 0000000000..a9136f7af3 --- /dev/null +++ b/supabase/migrations/20260923105200_dual_stripe_billing_account.sql @@ -0,0 +1,452 @@ +-- Dual Stripe billing: EE (legacy) + US (new orgs). Existing rows stay on EE. + +ALTER TABLE public.stripe_info + ADD COLUMN IF NOT EXISTS billing_account text; + +ALTER TABLE public.stripe_info + ALTER COLUMN billing_account SET DEFAULT 'ee'; + +ALTER TABLE public.stripe_info + DROP CONSTRAINT IF EXISTS stripe_info_billing_account_check; + +ALTER TABLE public.stripe_info + ADD CONSTRAINT stripe_info_billing_account_check + CHECK (billing_account IN ('ee', 'us')) NOT VALID; + +UPDATE public.stripe_info +SET billing_account = 'ee' +WHERE billing_account IS NULL; + +ALTER TABLE public.stripe_info + VALIDATE CONSTRAINT stripe_info_billing_account_check; + +ALTER TABLE public.stripe_info + ALTER COLUMN billing_account SET NOT NULL; + +COMMENT ON COLUMN public.stripe_info.billing_account IS + 'Stripe account: ee (Capgo OÜ legacy) or us (CodepushGo LLC).'; + +ALTER TABLE public.plans + ADD COLUMN IF NOT EXISTS stripe_id_us character varying, + ADD COLUMN IF NOT EXISTS price_m_id_us character varying, + ADD COLUMN IF NOT EXISTS price_y_id_us character varying, + ADD COLUMN IF NOT EXISTS credit_id_us text; + +COMMENT ON COLUMN public.plans.stripe_id_us IS + 'Stripe product id on the US Stripe account.'; +COMMENT ON COLUMN public.plans.price_m_id_us IS + 'Monthly Stripe price id on the US Stripe account.'; +COMMENT ON COLUMN public.plans.price_y_id_us IS + 'Yearly Stripe price id on the US Stripe account.'; +COMMENT ON COLUMN public.plans.credit_id_us IS + 'Stripe product id for credit top-ups on the US Stripe account.'; + +UPDATE public.plans SET + stripe_id_us = 'prod_VDt1FTF7XJxyMR', + price_m_id_us = 'price_1UDRGPLr632EP5z4ufTRBBzf', + price_y_id_us = 'price_1UDRGULr632EP5z4OcZr5xpe', + credit_id_us = 'prod_VDt2YB5GrYFnII' +WHERE name = 'Solo'; + +UPDATE public.plans SET + stripe_id_us = 'prod_VDt2cnktX7IDVV', + price_m_id_us = 'price_1UDRGQLr632EP5z4YI3A5cPV', + price_y_id_us = 'price_1UDRGSLr632EP5z45R9qxkU8', + credit_id_us = 'prod_VDt2YB5GrYFnII' +WHERE name = 'Maker'; + +UPDATE public.plans SET + stripe_id_us = 'prod_VDt2xM7OyLzhqV', + price_m_id_us = 'price_1UDRGSLr632EP5z4n0Npf7P1', + price_y_id_us = 'price_1UDRGSLr632EP5z4jYu6vC42', + credit_id_us = 'prod_VDt2YB5GrYFnII' +WHERE name = 'Team'; + +UPDATE public.plans SET + stripe_id_us = 'prod_VDt2pia049SqpU', + price_m_id_us = 'price_1UDRGaLr632EP5z4D02F4rLu', + price_y_id_us = 'price_1UDRGcLr632EP5z4IxDyZuUK', + credit_id_us = 'prod_VDt2YB5GrYFnII' +WHERE name = 'Enterprise'; + +-- product_id references EE plans.stripe_id or US plans.stripe_id_us. +ALTER TABLE public.stripe_info + DROP CONSTRAINT IF EXISTS stripe_info_product_id_fkey; + +CREATE OR REPLACE FUNCTION public.validate_stripe_info_product_id() +RETURNS trigger +LANGUAGE plpgsql +SECURITY DEFINER +SET search_path = '' +AS $$ +BEGIN + IF NEW.product_id IS NULL THEN + RETURN NEW; + END IF; + + IF NEW.billing_account = 'us' THEN + PERFORM pg_advisory_xact_lock(hashtext('us:' || NEW.product_id)); + IF NOT EXISTS ( + SELECT 1 + FROM public.plans + WHERE public.plans.stripe_id_us = NEW.product_id + ) THEN + RAISE EXCEPTION + 'stripe_info.product_id % is not a known US plan product id', + NEW.product_id; + END IF; + ELSE + PERFORM pg_advisory_xact_lock(hashtext('ee:' || NEW.product_id)); + IF NOT EXISTS ( + SELECT 1 + FROM public.plans + WHERE public.plans.stripe_id = NEW.product_id + ) THEN + RAISE EXCEPTION + 'stripe_info.product_id % is not a known EE plan product id', + NEW.product_id; + END IF; + END IF; + + RETURN NEW; +END; +$$; + +ALTER FUNCTION public.validate_stripe_info_product_id() OWNER TO postgres; + +REVOKE ALL ON FUNCTION public.validate_stripe_info_product_id() FROM PUBLIC; +GRANT ALL ON FUNCTION public.validate_stripe_info_product_id() TO service_role; + +DROP TRIGGER IF EXISTS validate_stripe_info_product_id ON public.stripe_info; + +CREATE TRIGGER validate_stripe_info_product_id + BEFORE INSERT OR UPDATE OF product_id, billing_account ON public.stripe_info + FOR EACH ROW + EXECUTE FUNCTION public.validate_stripe_info_product_id(); + +CREATE OR REPLACE FUNCTION public.prevent_orphan_stripe_info_plan_ids() +RETURNS trigger +LANGUAGE plpgsql +SET search_path = '' +AS $$ +BEGIN + IF TG_OP = 'DELETE' THEN + PERFORM pg_advisory_xact_lock(hashtext('ee:' || OLD.stripe_id)); + IF OLD.stripe_id_us IS NOT NULL THEN + PERFORM pg_advisory_xact_lock(hashtext('us:' || OLD.stripe_id_us)); + END IF; + + IF EXISTS ( + SELECT 1 + FROM public.stripe_info + WHERE public.stripe_info.billing_account = 'ee' + AND public.stripe_info.product_id = OLD.stripe_id + ) THEN + RAISE EXCEPTION + 'Cannot delete plan: stripe_info rows reference plans.stripe_id %', + OLD.stripe_id; + END IF; + + IF OLD.stripe_id_us IS NOT NULL AND EXISTS ( + SELECT 1 + FROM public.stripe_info + WHERE public.stripe_info.billing_account = 'us' + AND public.stripe_info.product_id = OLD.stripe_id_us + ) THEN + RAISE EXCEPTION + 'Cannot delete plan: stripe_info rows reference plans.stripe_id_us %', + OLD.stripe_id_us; + END IF; + + RETURN OLD; + END IF; + + IF OLD.stripe_id IS DISTINCT FROM NEW.stripe_id THEN + PERFORM pg_advisory_xact_lock(hashtext('ee:' || OLD.stripe_id)); + END IF; + + IF OLD.stripe_id IS DISTINCT FROM NEW.stripe_id AND EXISTS ( + SELECT 1 + FROM public.stripe_info + WHERE public.stripe_info.billing_account = 'ee' + AND public.stripe_info.product_id = OLD.stripe_id + ) THEN + RAISE EXCEPTION + 'Cannot change plans.stripe_id %: referenced by stripe_info (ee)', + OLD.stripe_id; + END IF; + + IF OLD.stripe_id_us IS DISTINCT FROM NEW.stripe_id_us AND OLD.stripe_id_us IS NOT NULL THEN + PERFORM pg_advisory_xact_lock(hashtext('us:' || OLD.stripe_id_us)); + END IF; + + IF OLD.stripe_id_us IS DISTINCT FROM NEW.stripe_id_us + AND OLD.stripe_id_us IS NOT NULL + AND EXISTS ( + SELECT 1 + FROM public.stripe_info + WHERE public.stripe_info.billing_account = 'us' + AND public.stripe_info.product_id = OLD.stripe_id_us + ) THEN + RAISE EXCEPTION + 'Cannot change plans.stripe_id_us %: referenced by stripe_info (us)', + OLD.stripe_id_us; + END IF; + + RETURN NEW; +END; +$$; + +ALTER FUNCTION public.prevent_orphan_stripe_info_plan_ids() OWNER TO postgres; + +REVOKE ALL ON FUNCTION public.prevent_orphan_stripe_info_plan_ids() FROM PUBLIC; +GRANT ALL ON FUNCTION public.prevent_orphan_stripe_info_plan_ids() + TO service_role; + +DROP TRIGGER IF EXISTS prevent_orphan_stripe_info_plan_ids ON public.plans; + +CREATE TRIGGER prevent_orphan_stripe_info_plan_ids + BEFORE UPDATE OF stripe_id, stripe_id_us OR DELETE ON public.plans + FOR EACH ROW + EXECUTE FUNCTION public.prevent_orphan_stripe_info_plan_ids(); + +-- apps channel_device_count / manifest_bundle_count bumps are bookkeeping for +-- any actor. Stale capgkey headers in the same SQL transaction must not turn +-- recount work into audit noise +-- (see supabase/tests/40_test_audit_log_apikey.sql). +CREATE OR REPLACE FUNCTION public.audit_log_trigger() +RETURNS trigger +LANGUAGE plpgsql +SECURITY DEFINER +SET search_path = '' +AS $$ +DECLARE + v_old_record jsonb; + v_new_record jsonb; + v_changed_fields text[]; + v_org_id uuid; + v_record_id text; + v_user_id uuid; + v_key text; + v_api_key_text text; + v_api_key public.apikeys%ROWTYPE; + v_actor_type text := 'system'; + v_actor_user_id uuid; + v_actor_user_email text; + v_actor_apikey_id bigint; + v_actor_apikey_name text; + v_stats_refresh_fields constant text[] := ARRAY['stats_refresh_requested_at', 'stats_updated_at', 'updated_at']; + v_background_counter_fields constant text[] := ARRAY['channel_device_count', 'manifest_bundle_count', 'updated_at']; + v_onboarding_progress_fields constant text[] := ARRAY['onboarding', 'updated_at']; + v_fat_app_version_fields constant text[] := ARRAY['manifest', 'native_packages']; +BEGIN + SELECT auth.uid() INTO v_actor_user_id; + + IF v_actor_user_id IS NOT NULL THEN + v_actor_type := 'user'; + ELSE + SELECT public.get_apikey_header() INTO v_api_key_text; + + IF v_api_key_text IS NOT NULL THEN + SELECT * + INTO v_api_key + FROM public.find_apikey_by_value(v_api_key_text) + LIMIT 1; + + IF v_api_key.id IS NOT NULL + AND NOT public.is_apikey_expired(v_api_key.expires_at) + AND ( + public.is_allowed_capgkey(v_api_key_text, '{upload}'::text[]) + OR public.is_allowed_capgkey(v_api_key_text, '{write}'::text[]) + OR public.is_allowed_capgkey(v_api_key_text, '{all}'::text[]) + ) THEN + v_actor_type := 'apikey'; + v_actor_user_id := v_api_key.user_id; + v_actor_apikey_id := v_api_key.id; + v_actor_apikey_name := v_api_key.name; + END IF; + END IF; + END IF; + + IF v_actor_user_id IS NOT NULL THEN + SELECT users.email + INTO v_actor_user_email + FROM public.users AS users + WHERE users.id = v_actor_user_id; + END IF; + + v_user_id := v_actor_user_id; + + IF TG_OP = 'UPDATE' AND TG_TABLE_NAME = 'app_versions' THEN + v_old_record := pg_catalog.to_jsonb(OLD); + v_new_record := pg_catalog.to_jsonb(NEW); + IF ( + v_old_record + - 'manifest' + - 'updated_at' + - 'manifest_count' + - 'storage_provider' + - 'r2_path' + ) IS NOT DISTINCT FROM ( + v_new_record + - 'manifest' + - 'updated_at' + - 'manifest_count' + - 'storage_provider' + - 'r2_path' + ) THEN + RETURN NEW; + END IF; + END IF; + + IF TG_OP = 'DELETE' THEN + v_old_record := pg_catalog.to_jsonb(OLD); + v_new_record := NULL; + ELSIF TG_OP = 'INSERT' THEN + v_old_record := NULL; + v_new_record := pg_catalog.to_jsonb(NEW); + ELSE + v_old_record := pg_catalog.to_jsonb(OLD); + v_new_record := pg_catalog.to_jsonb(NEW); + + FOR v_key IN SELECT pg_catalog.jsonb_object_keys(v_new_record) + LOOP + IF v_old_record->v_key IS DISTINCT FROM v_new_record->v_key THEN + v_changed_fields := pg_catalog.array_append(v_changed_fields, v_key); + END IF; + END LOOP; + + IF v_changed_fields IS NOT NULL + AND NOT EXISTS ( + SELECT 1 + FROM pg_catalog.unnest(v_changed_fields) AS changed_field(field_name) + WHERE changed_field.field_name IS DISTINCT FROM 'updated_at' + ) THEN + RETURN NEW; + END IF; + + IF TG_TABLE_NAME = ANY(ARRAY['apps', 'orgs']) + AND v_changed_fields && ARRAY['stats_refresh_requested_at', 'stats_updated_at'] + AND NOT EXISTS ( + SELECT 1 + FROM pg_catalog.unnest(v_changed_fields) AS changed_field(field_name) + WHERE changed_field.field_name <> ALL(v_stats_refresh_fields) + ) THEN + RETURN NEW; + END IF; + + IF TG_TABLE_NAME = 'apps' + AND v_changed_fields && ARRAY['channel_device_count', 'manifest_bundle_count'] + AND NOT EXISTS ( + SELECT 1 + FROM pg_catalog.unnest(v_changed_fields) AS changed_field(field_name) + WHERE changed_field.field_name <> ALL(v_background_counter_fields) + ) THEN + RETURN NEW; + END IF; + + IF TG_TABLE_NAME = 'apps' + AND v_changed_fields && ARRAY['onboarding'] + AND NOT EXISTS ( + SELECT 1 + FROM pg_catalog.unnest(v_changed_fields) AS changed_field(field_name) + WHERE changed_field.field_name <> ALL(v_onboarding_progress_fields) + ) THEN + RETURN NEW; + END IF; + END IF; + + IF TG_TABLE_NAME = 'app_versions' THEN + IF v_old_record IS NOT NULL THEN + v_old_record := v_old_record - v_fat_app_version_fields; + END IF; + IF v_new_record IS NOT NULL THEN + v_new_record := v_new_record - v_fat_app_version_fields; + END IF; + END IF; + + IF TG_OP = 'DELETE' THEN + CASE TG_TABLE_NAME + WHEN 'orgs' THEN + v_org_id := OLD.id; + v_record_id := OLD.id::text; + WHEN 'apps' THEN + v_org_id := OLD.owner_org; + v_record_id := OLD.app_id::text; + WHEN 'channels' THEN + v_org_id := OLD.owner_org; + v_record_id := OLD.id::text; + WHEN 'app_versions' THEN + v_org_id := OLD.owner_org; + v_record_id := OLD.id::text; + WHEN 'org_users' THEN + v_org_id := OLD.org_id; + v_record_id := OLD.id::text; + ELSE + v_org_id := NULL; + v_record_id := NULL; + END CASE; + ELSE + CASE TG_TABLE_NAME + WHEN 'orgs' THEN + v_org_id := NEW.id; + v_record_id := NEW.id::text; + WHEN 'apps' THEN + v_org_id := NEW.owner_org; + v_record_id := NEW.app_id::text; + WHEN 'channels' THEN + v_org_id := NEW.owner_org; + v_record_id := NEW.id::text; + WHEN 'app_versions' THEN + v_org_id := NEW.owner_org; + v_record_id := NEW.id::text; + WHEN 'org_users' THEN + v_org_id := NEW.org_id; + v_record_id := NEW.id::text; + ELSE + v_org_id := NULL; + v_record_id := NULL; + END CASE; + END IF; + + IF v_org_id IS NOT NULL THEN + INSERT INTO public.audit_logs ( + table_name, + record_id, + operation, + user_id, + org_id, + old_record, + new_record, + changed_fields, + actor_type, + actor_user_id, + actor_user_email, + actor_apikey_id, + actor_apikey_name + ) VALUES ( + TG_TABLE_NAME, + v_record_id, + TG_OP, + v_user_id, + v_org_id, + v_old_record, + v_new_record, + v_changed_fields, + v_actor_type, + v_actor_user_id, + v_actor_user_email, + v_actor_apikey_id, + v_actor_apikey_name + ); + END IF; + + IF TG_OP = 'DELETE' THEN + RETURN OLD; + END IF; + + RETURN NEW; +END; +$$; + +ALTER FUNCTION public.audit_log_trigger() OWNER TO postgres; diff --git a/supabase/seed.sql b/supabase/seed.sql index 8fb515dc54..0f81a29e6b 100644 --- a/supabase/seed.sql +++ b/supabase/seed.sql @@ -81,11 +81,11 @@ BEGIN INSERT INTO "public"."deleted_account" ("created_at", "email", "id") VALUES (NOW(), encode(extensions.digest('deleted@capgo.app'::bytea, 'sha256'::text)::bytea, 'hex'::text), '00000000-0000-0000-0000-000000000001'); - INSERT INTO "public"."plans" ("created_at", "updated_at", "name", "description", "price_m", "price_y", "stripe_id", "credit_id", "id", "price_m_id", "price_y_id", "storage", "bandwidth", "mau", "market_desc", "build_time_unit", "native_build_concurrency") VALUES - (NOW(), NOW(), 'Maker', 'plan.maker.desc', 39, 396, 'prod_LQIs1Yucml9ChU', 'prod_TJRd2hFHZsBIPK', '440cfd69-0cfd-486e-b59b-cb99f7ae76a0', 'price_1KjSGyGH46eYKnWwL4h14DsK', 'price_1KjSKIGH46eYKnWwFG9u4tNi', 3221225472, 268435456000, 10000, 'Best for small business owners', 7200, 3), - (NOW(), NOW(), 'Enterprise', 'plan.payasyougo.desc', 239, 2490, 'prod_MH5Jh6ajC9e7ZH', 'prod_TJRd2hFHZsBIPK', '745d7ab3-6cd6-4d65-b257-de6782d5ba50', 'price_1LYX8yGH46eYKnWwzeBjISvW', 'price_1LYX8yGH46eYKnWwzeBjISvW', 12884901888, 3221225472000, 1000000, 'Best for scalling enterprises', 1200000, 6), - (NOW(), NOW(), 'Solo', 'plan.solo.desc', 14, 146, 'prod_LQIregjtNduh4q', 'prod_TJRd2hFHZsBIPK', '526e11d8-3c51-4581-ac92-4770c602f47c', 'price_1LVvuZGH46eYKnWwuGKOf4DK', 'price_1LVvuIGH46eYKnWwHMDCrxcH', 1073741824, 13958643712, 2000, 'Best for independent developers', 3600, 2), - (NOW(), NOW(), 'Team', 'plan.team.desc', 99, 998, 'prod_LQIugvJcPrxhda', 'prod_TJRd2hFHZsBIPK', 'abd76414-8f90-49a5-b3a4-8ff4d2e12c77', 'price_1KjSIUGH46eYKnWwWHvg8XYs', 'price_1KjSLlGH46eYKnWwAwMW2wiW', 6442450944, 536870912000, 100000, 'Best for medium enterprises', 36000, 4); + INSERT INTO "public"."plans" ("created_at", "updated_at", "name", "description", "price_m", "price_y", "stripe_id", "credit_id", "id", "price_m_id", "price_y_id", "stripe_id_us", "price_m_id_us", "price_y_id_us", "credit_id_us", "storage", "bandwidth", "mau", "market_desc", "build_time_unit", "native_build_concurrency") VALUES + (NOW(), NOW(), 'Maker', 'plan.maker.desc', 39, 396, 'prod_LQIs1Yucml9ChU', 'prod_TJRd2hFHZsBIPK', '440cfd69-0cfd-486e-b59b-cb99f7ae76a0', 'price_1KjSGyGH46eYKnWwL4h14DsK', 'price_1KjSKIGH46eYKnWwFG9u4tNi', 'prod_VDt2cnktX7IDVV', 'price_1UDRGQLr632EP5z4YI3A5cPV', 'price_1UDRGSLr632EP5z45R9qxkU8', 'prod_VDt2YB5GrYFnII', 3221225472, 268435456000, 10000, 'Best for small business owners', 7200, 3), + (NOW(), NOW(), 'Enterprise', 'plan.payasyougo.desc', 239, 2490, 'prod_MH5Jh6ajC9e7ZH', 'prod_TJRd2hFHZsBIPK', '745d7ab3-6cd6-4d65-b257-de6782d5ba50', 'price_1LYX8yGH46eYKnWwzeBjISvW', 'price_1LYX8yGH46eYKnWwzeBjISvW', 'prod_VDt2pia049SqpU', 'price_1UDRGaLr632EP5z4D02F4rLu', 'price_1UDRGcLr632EP5z4IxDyZuUK', 'prod_VDt2YB5GrYFnII', 12884901888, 3221225472000, 1000000, 'Best for scalling enterprises', 1200000, 6), + (NOW(), NOW(), 'Solo', 'plan.solo.desc', 14, 146, 'prod_LQIregjtNduh4q', 'prod_TJRd2hFHZsBIPK', '526e11d8-3c51-4581-ac92-4770c602f47c', 'price_1LVvuZGH46eYKnWwuGKOf4DK', 'price_1LVvuIGH46eYKnWwHMDCrxcH', 'prod_VDt1FTF7XJxyMR', 'price_1UDRGPLr632EP5z4ufTRBBzf', 'price_1UDRGULr632EP5z4OcZr5xpe', 'prod_VDt2YB5GrYFnII', 1073741824, 13958643712, 2000, 'Best for independent developers', 3600, 2), + (NOW(), NOW(), 'Team', 'plan.team.desc', 99, 998, 'prod_LQIugvJcPrxhda', 'prod_TJRd2hFHZsBIPK', 'abd76414-8f90-49a5-b3a4-8ff4d2e12c77', 'price_1KjSIUGH46eYKnWwWHvg8XYs', 'price_1KjSLlGH46eYKnWwAwMW2wiW', 'prod_VDt2xM7OyLzhqV', 'price_1UDRGSLr632EP5z4n0Npf7P1', 'price_1UDRGSLr632EP5z4jYu6vC42', 'prod_VDt2YB5GrYFnII', 6442450944, 536870912000, 100000, 'Best for medium enterprises', 36000, 4); INSERT INTO "public"."capgo_credits_steps" ( diff --git a/supabase/tests/40_test_audit_log_apikey.sql b/supabase/tests/40_test_audit_log_apikey.sql index f57a8cdc5d..c631dfa8b0 100644 --- a/supabase/tests/40_test_audit_log_apikey.sql +++ b/supabase/tests/40_test_audit_log_apikey.sql @@ -427,13 +427,17 @@ SELECT throws_ok( -- Test 18: background counter updates must not create audit or webhook work. +SELECT set_config('request.headers', NULL, true); +SELECT set_config('request.jwt.claim.sub', NULL, true); +SELECT set_config('request.jwt.claims', NULL, true); + DO $$ DECLARE v_before_audit_id bigint; BEGIN - PERFORM set_config('request.headers', '{}', true); - PERFORM set_config('request.jwt.claim.sub', '', true); - PERFORM set_config('request.jwt.claims', '{}', true); + PERFORM set_config('request.headers', NULL, true); + PERFORM set_config('request.jwt.claim.sub', NULL, true); + PERFORM set_config('request.jwt.claims', NULL, true); SELECT COALESCE(MAX(id), 0) INTO v_before_audit_id FROM public.audit_logs; diff --git a/tests/apikeys.test.ts b/tests/apikeys.test.ts index 751c0d5b5d..1187abdcf5 100644 --- a/tests/apikeys.test.ts +++ b/tests/apikeys.test.ts @@ -8,6 +8,7 @@ import { appApiKeyBindings, BASE_URL, executeSQL, + fetchTestRequest, getAuthHeaders, getAuthHeadersForCredentials, getSupabaseClient, @@ -42,11 +43,20 @@ async function appKeyBody(name: string, appId = APPNAME, extra: Record, headers: Record = authHeaders) { + return fetchTestRequest(`${BASE_URL}/apikey`, { + method: 'POST', + headers, + body: JSON.stringify(body), + }) +} + beforeAll(async () => { authHeaders = await getAuthHeaders() await resetAndSeedAppData(APPNAME) - // Load the apikey isolate before concurrent POSTs from this file. await warmEdgeEndpoint('/apikey', { method: 'GET', headers: authHeaders }) + // GET alone does not compile the POST handler; warm create path before concurrent POSTs. + await warmEdgeEndpoint('/apikey', { method: 'POST', headers: authHeaders, body: '{}' }) const warmupPostResponse = await fetch(`${BASE_URL}/apikey`, { method: 'POST', headers: { @@ -274,13 +284,9 @@ describe('[POST] /apikey operations', () => { it.concurrent('creates an app-only preview key bound to its owning organization', async () => { const appBindings = await appApiKeyBindings(APPNAME, 'app_preview') - const response = await fetch(`${BASE_URL}/apikey`, { - method: 'POST', - headers: authHeaders, - body: JSON.stringify({ - name: `app-preview-key-${id.slice(0, 8)}`, - bindings: appBindings, - }), + const response = await postApikey({ + name: `app-preview-key-${id.slice(0, 8)}`, + bindings: appBindings, }) expect(response.status).toBe(200) const data = await response.json<{ id: number, rbac_id: string }>() @@ -346,20 +352,12 @@ describe('[POST] /apikey operations', () => { const createdKeyIds: number[] = [] try { - const limitedResponse = await fetch(`${BASE_URL}/apikey`, { - method: 'POST', - headers: authHeaders, - body: JSON.stringify(await appKeyBody('app-management-blocked')), - }) + const limitedResponse = await postApikey(await appKeyBody('app-management-blocked')) expect(limitedResponse.status).toBe(200) const limitedData = await limitedResponse.json<{ id: number, key: string }>() createdKeyIds.push(limitedData.id) - const siblingResponse = await fetch(`${BASE_URL}/apikey`, { - method: 'POST', - headers: authHeaders, - body: JSON.stringify(orgKeyBody('sibling-management-target')), - }) + const siblingResponse = await postApikey(orgKeyBody('sibling-management-target')) expect(siblingResponse.status).toBe(200) const siblingData = await siblingResponse.json<{ id: number }>() createdKeyIds.push(siblingData.id) @@ -420,22 +418,14 @@ describe('[POST] /apikey operations', () => { const orgId = orgApiKeyBindings()[0].org_id try { - const managerResponse = await fetch(`${BASE_URL}/apikey`, { - method: 'POST', - headers: authHeaders, - body: JSON.stringify(orgKeyBody('org-management-blocked', { - bindings: orgApiKeyBindings(orgId, 'org_member'), - })), - }) + const managerResponse = await postApikey(orgKeyBody('org-management-blocked', { + bindings: orgApiKeyBindings(orgId, 'org_member'), + })) expect(managerResponse.status).toBe(200) const managerData = await managerResponse.json<{ id: number, key: string }>() createdKeyIds.push(managerData.id) - const siblingResponse = await fetch(`${BASE_URL}/apikey`, { - method: 'POST', - headers: authHeaders, - body: JSON.stringify(orgKeyBody('org-sibling-management-target')), - }) + const siblingResponse = await postApikey(orgKeyBody('org-sibling-management-target')) expect(siblingResponse.status).toBe(200) const siblingData = await siblingResponse.json<{ id: number }>() createdKeyIds.push(siblingData.id) diff --git a/tests/app-error-cases.test.ts b/tests/app-error-cases.test.ts index 4cf6811b56..d19da309ac 100644 --- a/tests/app-error-cases.test.ts +++ b/tests/app-error-cases.test.ts @@ -1,6 +1,6 @@ import { randomUUID } from 'node:crypto' import { afterAll, beforeAll, describe, expect, it } from 'vitest' -import { BASE_URL, createDirectApiKeyWithBindings, fetchTestRequest, getAuthHeaders, getSupabaseClient, NON_ACCESS_APP_NAME, resetAndSeedAppData, resetAppData, USER_EMAIL, USER_ID } from './test-utils.ts' +import { BASE_URL, createDirectApiKeyWithBindings, fetchTestRequest, getAuthHeaders, getSupabaseClient, NON_ACCESS_APP_NAME, resetAndSeedAppData, resetAppData, USER_EMAIL, USER_ID, warmEdgeEndpoint } from './test-utils.ts' const id = randomUUID().replace(/-/g, '').slice(0, 12) const APPNAME = `com.app.error.${id}` @@ -40,6 +40,12 @@ beforeAll(async () => { 'Content-Type': 'application/json', 'Authorization': createKey.key, } + + await warmEdgeEndpoint(`${BASE_URL}/app`, { + method: 'POST', + headers: testHeaders, + body: JSON.stringify({}), + }) }, 60000) afterAll(async () => { diff --git a/tests/stripe-billing-account.unit.test.ts b/tests/stripe-billing-account.unit.test.ts new file mode 100644 index 0000000000..4371d06220 --- /dev/null +++ b/tests/stripe-billing-account.unit.test.ts @@ -0,0 +1,180 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' + +const mockedEnv: Record = { + STRIPE_NEW_CUSTOMERS_ACCOUNT: 'ee', +} + +vi.mock('hono/adapter', async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + env: () => mockedEnv, + } +}) + +import { + getNewCustomersBillingAccount, + getPlanCreditProductId, + getPlanPriceId, + getPlanProductId, + IncompleteUsPlanConfigError, + getStripeSecretKeyEnvName, + getStripeWebhookSecretEnvName, + normalizeBillingAccount, + planProductIdOrFilter, + resolveCheckoutPlanProductId, + resolvePlanCreditProductId, +} from '../supabase/functions/_backend/utils/stripe_billing.ts' + +function createContext() { + return { + get: (key: string) => key === 'requestId' ? 'stripe-billing-test' : undefined, + } as any +} + +afterEach(() => { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = 'ee' +}) + +const SOLO_PLAN = { + stripe_id: 'prod_LQIregjtNduh4q', + price_m_id: 'price_1LVvuZGH46eYKnWwuGKOf4DK', + price_y_id: 'price_1LVvuIGH46eYKnWwHMDCrxcH', + credit_id: 'prod_TJRd2hFHZsBIPK', + stripe_id_us: 'prod_VDt1FTF7XJxyMR', + price_m_id_us: 'price_1UDRGPLr632EP5z4ufTRBBzf', + price_y_id_us: 'price_1UDRGULr632EP5z4OcZr5xpe', + credit_id_us: 'prod_VDt2YB5GrYFnII', +} + +describe('stripe billing account helpers', () => { + it('defaults new customers to ee', () => { + expect(getNewCustomersBillingAccount(createContext())).toBe('ee') + const previousFlag = mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT + try { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = '' + expect(getNewCustomersBillingAccount(createContext())).toBe('ee') + } + finally { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = previousFlag + } + }) + + it('routes new customers to us when flag is set', () => { + const previousFlag = mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT + try { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = 'us' + expect(getNewCustomersBillingAccount(createContext())).toBe('us') + } + finally { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = previousFlag + } + }) + + it('throws on invalid STRIPE_NEW_CUSTOMERS_ACCOUNT values', () => { + const previousFlag = mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT + try { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = 'typo' + expect(() => getNewCustomersBillingAccount(createContext())).toThrow(/Invalid STRIPE_NEW_CUSTOMERS_ACCOUNT/) + } + finally { + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = previousFlag + } + }) + + it('normalizes billing account values', () => { + expect(normalizeBillingAccount('us')).toBe('us') + expect(normalizeBillingAccount('ee')).toBe('ee') + expect(normalizeBillingAccount(null)).toBe('ee') + }) + + it('maps env var names per account', () => { + expect(getStripeSecretKeyEnvName('ee')).toBe('STRIPE_SECRET_KEY') + expect(getStripeSecretKeyEnvName('us')).toBe('STRIPE_SECRET_KEY_US') + expect(getStripeWebhookSecretEnvName('us')).toBe('STRIPE_WEBHOOK_SECRET_US') + }) + + it('resolves EE and US plan ids', () => { + expect(getPlanProductId(SOLO_PLAN, 'ee')).toBe('prod_LQIregjtNduh4q') + expect(getPlanProductId(SOLO_PLAN, 'us')).toBe('prod_VDt1FTF7XJxyMR') + expect(getPlanPriceId(SOLO_PLAN, 'us', 'month')).toBe('price_1UDRGPLr632EP5z4ufTRBBzf') + expect(getPlanPriceId(SOLO_PLAN, 'ee', 'year')).toBe('price_1LVvuIGH46eYKnWwHMDCrxcH') + expect(getPlanCreditProductId(SOLO_PLAN, 'us')).toBe('prod_VDt2YB5GrYFnII') + expect(resolvePlanCreditProductId({ ...SOLO_PLAN, credit_id_us: null }, 'us')).toBe('') + }) + + it('builds dual-product lookup filter', () => { + expect(planProductIdOrFilter('prod_VDt1FTF7XJxyMR')).toBe('stripe_id.eq.prod_VDt1FTF7XJxyMR,stripe_id_us.eq.prod_VDt1FTF7XJxyMR') + }) + + it('rejects invalid stripe product ids in lookup filter', () => { + expect(() => planProductIdOrFilter('')).toThrow('invalid_stripe_product_id') + expect(() => planProductIdOrFilter('price_123')).toThrow('invalid_stripe_product_id') + expect(() => planProductIdOrFilter('prod_bad),stripe_id_us.eq.x')).toThrow('invalid_stripe_product_id') + expect(planProductIdOrFilter('prod_solo_us')).toBe('stripe_id.eq.prod_solo_us,stripe_id_us.eq.prod_solo_us') + expect(planProductIdOrFilter('prod_cHwBt-3ULYAoLArt')).toBe('stripe_id.eq.prod_cHwBt-3ULYAoLArt,stripe_id_us.eq.prod_cHwBt-3ULYAoLArt') + }) + + it('rejects incomplete US plan config instead of falling back to EE ids', () => { + const incompleteUsPlan = { ...SOLO_PLAN, stripe_id_us: null, price_m_id_us: null } + expect(() => getPlanProductId(incompleteUsPlan, 'us')).toThrow(IncompleteUsPlanConfigError) + expect(() => getPlanPriceId(incompleteUsPlan, 'us', 'month')).toThrow(IncompleteUsPlanConfigError) + expect(() => getPlanCreditProductId({ ...SOLO_PLAN, credit_id_us: null }, 'us')).toThrow(IncompleteUsPlanConfigError) + expect(() => getPlanProductId({ ...SOLO_PLAN, stripe_id_us: ' ' }, 'us')).toThrow(IncompleteUsPlanConfigError) + }) + + it('defaults to ee when admin client is unavailable but throws on lookup errors', async () => { + const context = createContext() + const lookupError = { message: 'connection refused', code: 'PGRST000' } + + const adminModule = await import('../supabase/functions/_backend/utils/supabase.ts') + const billingModule = await import('../supabase/functions/_backend/utils/stripe_billing.ts') + + mockedEnv.STRIPE_NEW_CUSTOMERS_ACCOUNT = 'ee' + const missingAdminSpy = vi.spyOn(adminModule, 'supabaseAdmin').mockReturnValueOnce(undefined as any) + try { + await expect(billingModule.getBillingAccountForCustomer(context, 'cus_test')).resolves.toBe('ee') + } + finally { + missingAdminSpy.mockRestore() + } + + const lookupErrorSpy = vi.spyOn(adminModule, 'supabaseAdmin').mockReturnValueOnce({ + from: () => ({ + select: () => ({ + eq: () => ({ + maybeSingle: async () => ({ data: null, error: lookupError }), + }), + }), + }), + } as any) + try { + await expect(billingModule.getBillingAccountForCustomer(context, 'cus_test')).rejects.toEqual(lookupError) + } + finally { + lookupErrorSpy.mockRestore() + } + }) + + it('propagates plan lookup errors before selecting checkout product id', async () => { + const context = createContext() + const lookupError = { message: 'connection refused', code: 'PGRST000' } + const adminModule = await import('../supabase/functions/_backend/utils/supabase.ts') + + const lookupErrorSpy = vi.spyOn(adminModule, 'supabaseAdmin').mockReturnValue({ + from: () => ({ + select: () => ({ + or: () => ({ + maybeSingle: async () => ({ data: null, error: lookupError }), + }), + }), + }), + } as any) + try { + await expect(resolveCheckoutPlanProductId(context, 'prod_test', 'ee')).rejects.toEqual(lookupError) + } + finally { + lookupErrorSpy.mockRestore() + } + }) +}) diff --git a/tests/stripe-emulator.test.ts b/tests/stripe-emulator.test.ts index 4fcfee36f7..46e201f6d6 100644 --- a/tests/stripe-emulator.test.ts +++ b/tests/stripe-emulator.test.ts @@ -4,6 +4,7 @@ import { createServer } from 'node:net' import { createEmulator } from 'emulate' import { afterAll, afterEach, beforeAll, describe, expect, it, vi } from 'vitest' import { createCheckout, createOneTimeCheckout, getCreditCheckoutDetails, getStripe } from '../supabase/functions/_backend/utils/stripe.ts' +import type { BillingAccount } from '../supabase/functions/_backend/utils/stripe_billing.ts' const { mockedSupabaseAdmin } = vi.hoisted(() => ({ mockedSupabaseAdmin: vi.fn(), @@ -33,20 +34,60 @@ function expectCheckoutUrlOnEmulator(url: string, baseUrl: string) { expect(checkoutUrl.pathname).toMatch(/^\/checkout\/cs_/) } -function mockStoredPlanPrices(priceMonthId: string, priceYearId: string) { +function mockBillingAccountLookup(billingAccount: BillingAccount | null = 'ee') { mockedSupabaseAdmin.mockReturnValue({ - from: vi.fn().mockReturnValue({ - select: vi.fn().mockReturnValue({ - eq: vi.fn().mockReturnValue({ - single: vi.fn().mockResolvedValue({ - data: { - price_m_id: priceMonthId, - price_y_id: priceYearId, - }, - error: null, + from: vi.fn().mockImplementation((table: string) => { + if (table === 'stripe_info') { + return { + select: vi.fn().mockReturnValue({ + eq: vi.fn().mockReturnValue({ + maybeSingle: vi.fn().mockResolvedValue({ + data: billingAccount ? { billing_account: billingAccount } : null, + error: null, + }), + }), + }), + } + } + throw new Error(`unexpected table ${table}`) + }), + }) +} + +function mockCheckoutAdmin(planProductId: string, priceMonthId: string, priceYearId: string) { + const planRow = { + stripe_id: planProductId, + stripe_id_us: null, + price_m_id: priceMonthId, + price_y_id: priceYearId, + price_m_id_us: null, + price_y_id_us: null, + } + + mockedSupabaseAdmin.mockReturnValue({ + from: vi.fn().mockImplementation((table: string) => { + if (table === 'stripe_info') { + return { + select: vi.fn().mockReturnValue({ + eq: vi.fn().mockReturnValue({ + maybeSingle: vi.fn().mockResolvedValue({ data: { billing_account: 'ee' }, error: null }), + }), + }), + } + } + + return { + select: vi.fn().mockReturnValue({ + eq: vi.fn().mockReturnValue({ + single: vi.fn().mockResolvedValue({ data: planRow, error: null }), + maybeSingle: vi.fn().mockResolvedValue({ data: planRow, error: null }), + }), + or: vi.fn().mockReturnValue({ + single: vi.fn().mockResolvedValue({ data: planRow, error: null }), + maybeSingle: vi.fn().mockResolvedValue({ data: planRow, error: null }), }), }), - }), + } }), }) } @@ -160,7 +201,7 @@ describe('stripe emulator integration', () => { }, }) - mockStoredPlanPrices(monthlyPrice.id, yearlyPrice.id) + mockCheckoutAdmin(product.id, monthlyPrice.id, yearlyPrice.id) const checkout = await createCheckout( context, @@ -173,7 +214,7 @@ describe('stripe emulator integration', () => { expect(checkout.url).toBeTruthy() expectCheckoutUrlOnEmulator(checkout.url as string, stripeApiBaseUrl) - expect(mockedSupabaseAdmin).toHaveBeenCalledTimes(1) + expect(mockedSupabaseAdmin).toHaveBeenCalled() const sessions = await stripe.checkout.sessions.list({ limit: 10 }) const session = sessions.data.find(candidate => candidate.url === checkout.url) @@ -188,6 +229,7 @@ describe('stripe emulator integration', () => { it('falls back to checkout metadata when emulate does not implement line item reads', async () => { stubStripeEnv(stripeApiBaseUrl) + mockBillingAccountLookup() const context = createContext() const stripe = getStripe(context) diff --git a/tests/stripe-org-customer.unit.test.ts b/tests/stripe-org-customer.unit.test.ts index 39fc66c391..1d9f437020 100644 --- a/tests/stripe-org-customer.unit.test.ts +++ b/tests/stripe-org-customer.unit.test.ts @@ -12,9 +12,13 @@ const { supabaseAdminMock: vi.fn(), })) -vi.mock('../supabase/functions/_backend/utils/stripe.ts', () => ({ - createCustomer: createCustomerMock, -})) +vi.mock('../supabase/functions/_backend/utils/stripe.ts', async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + createCustomer: createCustomerMock, + } +}) vi.mock('../supabase/functions/_backend/utils/supabase.ts', () => ({ getDefaultPlan: getDefaultPlanMock, @@ -35,7 +39,14 @@ const USER_ID = 'a1bb59b7-34b3-4e06-a0f1-2cc696f043dc' const PENDING_ID = `pending_${ORG_ID}` const LOCAL_ID = `cus_local_${ORG_ID.replaceAll('-', '')}` const CUSTOMER_ID = 'cus_VAgMn1agG4iQSC' -const SOLO_PLAN = { name: 'Solo', stripe_id: 'prod_solo' } +const SOLO_PLAN = { + name: 'Solo', + stripe_id: 'prod_solo', + stripe_id_us: 'prod_solo_us', + price_m_id_us: 'price_solo_m_us', + price_y_id_us: 'price_solo_y_us', + credit_id_us: 'prod_credits_us', +} function createContext() { return { @@ -117,12 +128,14 @@ function mockSupabase(options: { } } if (table === 'plans') { + const planQueryResult = { + single: async () => ({ data: SOLO_PLAN, error: null }), + maybeSingle: async () => ({ data: SOLO_PLAN, error: null }), + } return { select: () => ({ - eq: () => ({ - single: async () => ({ data: SOLO_PLAN, error: null }), - maybeSingle: async () => ({ data: { name: SOLO_PLAN.name }, error: null }), - }), + eq: () => planQueryResult, + or: () => planQueryResult, }), } } @@ -162,7 +175,10 @@ describe('createStripeCustomer', () => { beforeEach(() => { vi.clearAllMocks() getDefaultPlanMock.mockResolvedValue(SOLO_PLAN) - getStripeCustomerMock.mockResolvedValue({ product_id: SOLO_PLAN.stripe_id }) + getStripeCustomerMock.mockResolvedValue({ + product_id: SOLO_PLAN.stripe_id, + billing_account: 'ee', + }) createCustomerMock.mockResolvedValue({ id: CUSTOMER_ID }) }) @@ -181,13 +197,63 @@ describe('createStripeCustomer', () => { const planName = await createStripeCustomer(createContext(), createOrg(LOCAL_ID)) expect(planName).toBe('Solo') - expect(createCustomerMock).toHaveBeenCalledTimes(1) + expect(createCustomerMock).toHaveBeenCalled() + expect(createCustomerMock).toHaveBeenCalledWith( + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + 'ee', + ) expect(orgUpdate).toHaveBeenCalledWith({ customer_id: CUSTOMER_ID }, expect.objectContaining({ id: ORG_ID, customer_id: LOCAL_ID, })) }) + it('uses pending stripe_info billing_account when finalizing a pending org', async () => { + getStripeCustomerMock.mockResolvedValue({ + product_id: 'prod_solo_us', + billing_account: 'us', + }) + const { stripeInfoInsert } = mockSupabase({ orgCustomerId: PENDING_ID }) + + await createStripeCustomer(createContext(), createOrg(PENDING_ID)) + + expect(createCustomerMock).toHaveBeenCalledWith( + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + 'us', + ) + expect(stripeInfoInsert).toHaveBeenCalledWith(expect.objectContaining({ + billing_account: 'us', + product_id: SOLO_PLAN.stripe_id_us, + })) + }) + + it('uses local stripe_info billing_account when replacing a fake customer id', async () => { + getStripeCustomerMock.mockResolvedValue({ + product_id: 'prod_solo_us', + billing_account: 'us', + }) + mockSupabase({ orgCustomerId: LOCAL_ID }) + + await createStripeCustomer(createContext(), createOrg(LOCAL_ID)) + + expect(createCustomerMock).toHaveBeenCalledWith( + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + expect.anything(), + 'us', + ) + }) + it('creates a real customer when the org has a pre-PR 24-hex fake id', async () => { const legacy24HexLocalId = `cus_${crypto.randomUUID().replaceAll('-', '').slice(0, 24)}` const { orgUpdate } = mockSupabase({ orgCustomerId: legacy24HexLocalId }) @@ -195,13 +261,61 @@ describe('createStripeCustomer', () => { const planName = await createStripeCustomer(createContext(), createOrg(legacy24HexLocalId)) expect(planName).toBe('Solo') - expect(createCustomerMock).toHaveBeenCalledTimes(1) + expect(createCustomerMock).toHaveBeenCalled() expect(orgUpdate).toHaveBeenCalledWith({ customer_id: CUSTOMER_ID }, expect.objectContaining({ id: ORG_ID, customer_id: legacy24HexLocalId, })) }) + it('throws when stored plan lookup fails so the queue can retry', async () => { + getStripeCustomerMock.mockResolvedValue({ + product_id: SOLO_PLAN.stripe_id, + billing_account: 'ee', + }) + const planLookupError = { message: 'timeout' } + supabaseAdminMock.mockImplementation(() => ({ + from: (table: string) => { + if (table === 'orgs') { + return { + select: () => ({ + eq: () => ({ + single: async () => ({ + data: createOrg(PENDING_ID), + error: null, + }), + }), + }), + update: (payload: { customer_id: string }) => orgUpdateQuery(payload, vi.fn()), + } + } + if (table === 'stripe_info') { + return { + insert: vi.fn(async () => ({ error: null })), + delete: () => ({ + eq: vi.fn(async () => ({ error: null })), + }), + } + } + if (table === 'plans') { + return { + select: () => ({ + or: () => ({ + maybeSingle: async () => ({ data: null, error: planLookupError }), + }), + }), + } + } + throw new Error(`unexpected table ${table}`) + }, + })) + + await expect(createStripeCustomer(createContext(), createOrg(PENDING_ID))) + .rejects + .toMatchObject(planLookupError) + expect(createCustomerMock).not.toHaveBeenCalled() + }) + it('throws when org reload fails so the queue can retry', async () => { supabaseAdminMock.mockImplementation(() => ({ from: (table: string) => { @@ -233,8 +347,8 @@ describe('createStripeCustomer', () => { const planName = await createStripeCustomer(createContext(), createOrg(PENDING_ID)) expect(planName).toBe('Solo') - expect(createCustomerMock).toHaveBeenCalledTimes(1) - expect(stripeInfoInsert).toHaveBeenCalledTimes(1) + expect(createCustomerMock).toHaveBeenCalled() + expect(stripeInfoInsert).toHaveBeenCalled() expect(orgUpdate).toHaveBeenCalledWith({ customer_id: CUSTOMER_ID }, expect.objectContaining({ id: ORG_ID, customer_id: PENDING_ID, @@ -282,12 +396,14 @@ describe('createStripeCustomer', () => { } } if (table === 'plans') { + const planQueryResult = { + single: async () => ({ data: SOLO_PLAN, error: null }), + maybeSingle: async () => ({ data: { name: SOLO_PLAN.name }, error: null }), + } return { select: () => ({ - eq: () => ({ - single: async () => ({ data: SOLO_PLAN, error: null }), - maybeSingle: async () => ({ data: { name: SOLO_PLAN.name }, error: null }), - }), + eq: () => planQueryResult, + or: () => planQueryResult, }), } } @@ -298,7 +414,7 @@ describe('createStripeCustomer', () => { const planName = await createStripeCustomer(createContext(), createOrg(PENDING_ID)) expect(planName).toBe('Solo') - expect(createCustomerMock).toHaveBeenCalledTimes(1) + expect(createCustomerMock).toHaveBeenCalled() expect(orgState.customer_id).toBe(existingId) expect(stripeInfoDeleteEq).toHaveBeenCalledWith('customer_id', CUSTOMER_ID) }) @@ -308,7 +424,10 @@ describe('finalizePendingStripeCustomer', () => { beforeEach(() => { vi.clearAllMocks() getDefaultPlanMock.mockResolvedValue(SOLO_PLAN) - getStripeCustomerMock.mockResolvedValue({ product_id: SOLO_PLAN.stripe_id }) + getStripeCustomerMock.mockResolvedValue({ + product_id: SOLO_PLAN.stripe_id, + billing_account: 'ee', + }) createCustomerMock.mockResolvedValue({ id: CUSTOMER_ID }) }) @@ -329,7 +448,7 @@ describe('finalizePendingStripeCustomer', () => { const planName = await finalizePendingStripeCustomer(createContext(), createOrg(PENDING_ID)) expect(planName).toBe('Solo') - expect(createCustomerMock).toHaveBeenCalledTimes(1) + expect(createCustomerMock).toHaveBeenCalled() expect(stripeInfoDeleteEq).toHaveBeenCalledWith('customer_id', PENDING_ID) }) }) diff --git a/tests/stripe-redirects.unit.test.ts b/tests/stripe-redirects.unit.test.ts index fb223b8d4a..d251f0fd7b 100644 --- a/tests/stripe-redirects.unit.test.ts +++ b/tests/stripe-redirects.unit.test.ts @@ -1,7 +1,8 @@ import Stripe from 'stripe' -import { afterEach, describe, expect, it, vi } from 'vitest' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' const mockedEnv: Record = { + ENV: 'local', WEBAPP_URL: 'https://capgo.test', STRIPE_SECRET_KEY: 'sk_test_123', } @@ -53,13 +54,40 @@ function createPriceList(recurringInterval = 'month', type = 'recurring') { ] } +function mockBillingAccountLookup(billingAccount = 'ee') { + mockedSupabaseAdmin.mockReturnValue({ + from: vi.fn().mockReturnValue({ + select: vi.fn().mockReturnValue({ + eq: vi.fn().mockReturnValue({ + maybeSingle: vi.fn().mockResolvedValue({ + data: { billing_account: billingAccount }, + error: null, + }), + }), + or: vi.fn().mockReturnValue({ + maybeSingle: vi.fn().mockResolvedValue({ + data: null, + error: null, + }), + }), + }), + }), + }) +} + afterEach(() => { delete mockedEnv.STRIPE_API_BASE_URL + mockedEnv.ENV = 'local' + mockedEnv.STRIPE_SECRET_KEY = 'sk_test_123' mockedSupabaseAdmin.mockReset() vi.restoreAllMocks() }) describe('stripe redirect URL allowlist', () => { + beforeEach(() => { + mockBillingAccountLookup() + }) + it('allows same-origin return URLs for billing portal', async () => { const createSession = vi.fn().mockResolvedValue({ url: 'https://pay.capgo.test/p/session' }) const stripeClient = { @@ -84,6 +112,25 @@ describe('stripe redirect URL allowlist', () => { }) }) + it('returns empty portal url when stripe is not configured for the billing account', async () => { + mockedEnv.STRIPE_SECRET_KEY = '' + + const createSession = vi.fn() + vi.mocked(Stripe).mockImplementation(function () { + return { + billingPortal: { + sessions: { create: createSession }, + }, + } as any + } as any) + + const { createPortal } = await import('../supabase/functions/_backend/utils/stripe.ts') + const result = await createPortal(createContext(), 'cus_123', '/app/usage') + + expect(result.url).toBe('') + expect(createSession).not.toHaveBeenCalled() + }) + it('rejects external return URLs for billing portal', async () => { const createSession = vi.fn() const stripeClient = { @@ -129,7 +176,7 @@ describe('stripe redirect URL allowlist', () => { createContext(), 'cus_123', 'month', - 'plan_test', + 'prod_test', '/app/success', '/app/cancel', 'org_123', @@ -178,7 +225,7 @@ describe('stripe redirect URL allowlist', () => { createContext(), 'cus_123', 'month', - 'plan_test', + 'prod_test', 'https://example.com/phishing', '/app/cancel', ).catch(error => error) @@ -244,6 +291,54 @@ describe('stripe redirect URL allowlist', () => { })) }) + it('allows host.docker.internal for Playwright Stripe emulator base URL', async () => { + mockedEnv.STRIPE_API_BASE_URL = 'http://host.docker.internal:4520' + mockedEnv.STRIPE_SECRET_KEY = 'sk_test_emulator' + + const stripeClient = { + checkout: { + sessions: {}, + }, + } as any + + vi.mocked(Stripe).mockImplementation(function () { + return stripeClient + } as any) + + const { getStripe } = await import('../supabase/functions/_backend/utils/stripe.ts') + getStripe(createContext()) + + expect(Stripe).toHaveBeenCalledWith('sk_test_emulator', expect.objectContaining({ + host: 'host.docker.internal', + port: 4520, + protocol: 'http', + })) + }) + + it('rejects http Stripe API base URL when using live credentials', async () => { + mockedEnv.STRIPE_API_BASE_URL = 'http://host.docker.internal:4520' + mockedEnv.STRIPE_SECRET_KEY = 'sk_live_123' + + const { getStripe } = await import('../supabase/functions/_backend/utils/stripe.ts') + expect(() => getStripe(createContext())).toThrow('STRIPE_API_BASE_URL must use https when using live Stripe credentials') + }) + + it('rejects host.docker.internal http base URL outside local emulator config', async () => { + mockedEnv.STRIPE_API_BASE_URL = 'http://host.docker.internal:4520' + mockedEnv.STRIPE_SECRET_KEY = 'sk_test_123' + mockedEnv.ENV = 'production' + + const { getStripe } = await import('../supabase/functions/_backend/utils/stripe.ts') + expect(() => getStripe(createContext())).toThrow('STRIPE_API_BASE_URL host.docker.internal is only allowed for local emulator config') + }) + + it('rejects non-local http Stripe API base URLs', async () => { + mockedEnv.STRIPE_API_BASE_URL = 'http://stripe.example.com' + + const { getStripe } = await import('../supabase/functions/_backend/utils/stripe.ts') + expect(() => getStripe(createContext())).toThrow('STRIPE_API_BASE_URL must use https for non-loopback hosts') + }) + it('falls back to checkout metadata for credit top-ups when line items are unavailable in emulator mode', async () => { mockedEnv.STRIPE_API_BASE_URL = 'http://127.0.0.1:4510' @@ -355,6 +450,28 @@ describe('stripe redirect URL allowlist', () => { }, error: null, }), + maybeSingle: vi.fn().mockResolvedValue({ + data: { billing_account: 'ee' }, + error: null, + }), + }), + or: vi.fn().mockReturnValue({ + single: vi.fn().mockResolvedValue({ + data: { + stripe_id: 'prod_test', + price_m_id: 'price_monthly_from_plan', + price_y_id: 'price_yearly_from_plan', + }, + error: null, + }), + maybeSingle: vi.fn().mockResolvedValue({ + data: { + stripe_id: 'prod_test', + price_m_id: 'price_monthly_from_plan', + price_y_id: 'price_yearly_from_plan', + }, + error: null, + }), }), }), }), @@ -389,7 +506,7 @@ describe('stripe redirect URL allowlist', () => { createContext(), 'cus_123', 'month', - 'plan_test', + 'prod_test', '/app/success', '/app/cancel', ) diff --git a/tests/test-utils.ts b/tests/test-utils.ts index 909f97a763..1973286278 100644 --- a/tests/test-utils.ts +++ b/tests/test-utils.ts @@ -907,7 +907,27 @@ export function getUpdateBaseData(appId: string): ReturnType '') + console.error(`[postUpdate] non-200 status=${response.status} body=${body.slice(0, 800)}`) + } + return response + } + if (attempt < 3) + await new Promise(resolve => setTimeout(resolve, attempt * 500)) + } + + return await fetchTestRequest( getEndpointUrl('/updates'), { method: 'POST', @@ -915,11 +935,6 @@ export async function postUpdate(data: object) { body: JSON.stringify(data), }, ) - if (response.status !== 200) { - const body = await response.clone().text().catch(() => '') - console.error(`[postUpdate] non-200 status=${response.status} body=${body.slice(0, 800)}`) - } - return response } export interface DeviceLink {