summaryrefslogtreecommitdiff
path: root/internal/billing/operations_test.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/operations_test.go58
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")
+ }
+}