diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/billing/operations_test.go | 58 |
1 files changed, 58 insertions, 0 deletions
diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go index 46bedeb..4d90383 100644 --- a/internal/billing/operations_test.go +++ b/internal/billing/operations_test.go @@ -1,8 +1,15 @@ package billing import ( + "context" + "fmt" + "os" "testing" "time" + + "aigw/internal/controlplane" + + "github.com/stripe/stripe-go/v86" ) func TestOperationalStatusReadiness(t *testing.T) { @@ -34,3 +41,54 @@ func TestOperationalStatusReadiness(t *testing.T) { t.Fatal("unmetered success must fail readiness") } } + +func TestCustomerPortalSessionContractPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + service, err := New(ctx, Options{ + DatabaseURL: databaseURL, Currency: "usd", StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", + StripePortalReturnURL: "https://console.example.test/billing", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + + suffix := time.Now().UnixNano() + var tenantID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Portal contract') RETURNING id::text`, fmt.Sprintf("portal-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM tenants WHERE id=$1`, tenantID); cleanupErr != nil { + t.Errorf("cleanup portal contract tenant: %v", cleanupErr) + } + }) + customerID := fmt.Sprintf("cus_portal_%d", suffix) + if _, err := service.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,'billing@example.test')`, tenantID, customerID); err != nil { + t.Fatal(err) + } + service.createStripePortalSession = func(_ context.Context, params *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) { + if params.Customer == nil || *params.Customer != customerID || params.ReturnURL == nil || *params.ReturnURL != service.stripePortalReturnURL { + t.Fatalf("unexpected Portal params %+v", params) + } + return &stripe.BillingPortalSession{URL: "https://billing.stripe.test/session"}, nil + } + result, err := service.CreatePortalSession(ctx, tenantID) + if err != nil || result.URL != "https://billing.stripe.test/session" { + t.Fatalf("CreatePortalSession result=%+v err=%v", result, err) + } + service.createStripePortalSession = func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) { + return nil, nil + } + if _, err := service.CreatePortalSession(ctx, tenantID); err == nil { + t.Fatal("incomplete Stripe Portal response was accepted") + } +} |
