summaryrefslogtreecommitdiff
path: root/internal/billing/stripe_preflight.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/stripe_preflight.go124
1 files changed, 124 insertions, 0 deletions
diff --git a/internal/billing/stripe_preflight.go b/internal/billing/stripe_preflight.go
new file mode 100644
index 0000000..4f9250c
--- /dev/null
+++ b/internal/billing/stripe_preflight.go
@@ -0,0 +1,124 @@
+package billing
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "strings"
+
+ "github.com/stripe/stripe-go/v86"
+)
+
+const (
+ stripeAPIBaseURL = "https://api.stripe.com"
+ maxStripeErrorBodySize = 64 << 10
+)
+
+var ErrLiveStripeKey = errors.New("Stripe permission preflight only accepts test-mode keys")
+
+type StripePermissionCheck struct {
+ Name string `json:"name"`
+ OK bool `json:"ok"`
+ StatusCode int `json:"status_code"`
+ ErrorCode string `json:"error_code,omitempty"`
+ ErrorType string `json:"error_type,omitempty"`
+}
+
+type StripePreflightResult struct {
+ APIVersion string `json:"api_version"`
+ TestMode bool `json:"test_mode"`
+ Ready bool `json:"ready"`
+ Checks []StripePermissionCheck `json:"checks"`
+}
+
+type stripeReadCheck struct {
+ name string
+ path string
+}
+
+var stripeRequiredReadChecks = []stripeReadCheck{
+ {name: "customers_read", path: "/v1/customers"},
+ {name: "checkout_sessions_read", path: "/v1/checkout/sessions"},
+ {name: "setup_intents_read", path: "/v1/setup_intents"},
+ {name: "payment_intents_read", path: "/v1/payment_intents"},
+ {name: "refunds_read", path: "/v1/refunds"},
+ {name: "charges_read", path: "/v1/charges"},
+ {name: "disputes_read", path: "/v1/disputes"},
+ {name: "invoices_read", path: "/v1/invoices"},
+ {name: "billing_portal_configurations_read", path: "/v1/billing_portal/configurations"},
+}
+
+// CheckStripePermissions validates the read side of the restricted-key contract
+// without creating Stripe objects. Write permissions are exercised by the
+// sandbox Checkout, Portal, automatic top-up, refund, and reconciliation flows.
+func CheckStripePermissions(ctx context.Context, apiKey string) (StripePreflightResult, error) {
+ return checkStripePermissions(ctx, apiKey, stripeAPIBaseURL, http.DefaultClient)
+}
+
+func checkStripePermissions(ctx context.Context, apiKey, baseURL string, client *http.Client) (StripePreflightResult, error) {
+ apiKey = strings.TrimSpace(apiKey)
+ result := StripePreflightResult{
+ APIVersion: stripe.APIVersion,
+ TestMode: isStripeTestKey(apiKey),
+ Checks: make([]StripePermissionCheck, 0, len(stripeRequiredReadChecks)),
+ }
+ if !result.TestMode {
+ return result, ErrLiveStripeKey
+ }
+ result.Ready = true
+ if client == nil {
+ client = http.DefaultClient
+ }
+ for _, check := range stripeRequiredReadChecks {
+ item := runStripeReadCheck(ctx, client, apiKey, baseURL, check)
+ result.Checks = append(result.Checks, item)
+ result.Ready = result.Ready && item.OK
+ }
+ return result, nil
+}
+
+func runStripeReadCheck(ctx context.Context, client *http.Client, apiKey, baseURL string, check stripeReadCheck) StripePermissionCheck {
+ endpoint, err := url.JoinPath(baseURL, check.path)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"}
+ }
+ query := url.Values{"limit": []string{"1"}}
+ request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint+"?"+query.Encode(), nil)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"}
+ }
+ request.Header.Set("Authorization", "Bearer "+apiKey)
+ request.Header.Set("Stripe-Version", stripe.APIVersion)
+ response, err := client.Do(request)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "network_error"}
+ }
+ defer response.Body.Close()
+ item := StripePermissionCheck{Name: check.name, OK: response.StatusCode >= 200 && response.StatusCode < 300, StatusCode: response.StatusCode}
+ if item.OK {
+ _, _ = io.Copy(io.Discard, response.Body)
+ return item
+ }
+ var envelope struct {
+ Error struct {
+ Code string `json:"code"`
+ Type string `json:"type"`
+ } `json:"error"`
+ }
+ if err := json.NewDecoder(io.LimitReader(response.Body, maxStripeErrorBodySize)).Decode(&envelope); err == nil {
+ item.ErrorCode = envelope.Error.Code
+ item.ErrorType = envelope.Error.Type
+ }
+ if item.ErrorType == "" {
+ item.ErrorType = fmt.Sprintf("http_%d", response.StatusCode)
+ }
+ return item
+}
+
+func isStripeTestKey(value string) bool {
+ return strings.HasPrefix(value, "rk_test_") || strings.HasPrefix(value, "sk_test_")
+}