diff options
Diffstat (limited to 'internal/billing/stripe_preflight_test.go')
| -rw-r--r-- | internal/billing/stripe_preflight_test.go | 90 |
1 files changed, 90 insertions, 0 deletions
diff --git a/internal/billing/stripe_preflight_test.go b/internal/billing/stripe_preflight_test.go new file mode 100644 index 0000000..e8f1841 --- /dev/null +++ b/internal/billing/stripe_preflight_test.go @@ -0,0 +1,90 @@ +package billing + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/stripe/stripe-go/v86" +) + +func TestStripePermissionPreflightChecksRequiredResourcesWithoutLeakingKey(t *testing.T) { + const key = "rk_test_do_not_log_this_value" + seen := make(map[string]bool) + client := &http.Client{Transport: stripeRoundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.Method != http.MethodGet || r.URL.Query().Get("limit") != "1" { + t.Errorf("unexpected request %s %s", r.Method, r.URL.String()) + } + if r.Header.Get("Authorization") != "Bearer "+key { + t.Errorf("missing Stripe bearer authentication") + } + if r.Header.Get("Stripe-Version") != stripe.APIVersion { + t.Errorf("Stripe-Version = %q", r.Header.Get("Stripe-Version")) + } + seen[r.URL.Path] = true + return stripeTestResponse(http.StatusOK, `{"object":"list","data":[]}`), nil + })} + + result, err := checkStripePermissions(context.Background(), key, "https://stripe.test", client) + if err != nil { + t.Fatal(err) + } + if !result.Ready || !result.TestMode || result.APIVersion != stripe.APIVersion { + t.Fatalf("unexpected result %+v", result) + } + if len(result.Checks) != len(stripeRequiredReadChecks) { + t.Fatalf("checks = %d, want %d", len(result.Checks), len(stripeRequiredReadChecks)) + } + for _, check := range stripeRequiredReadChecks { + if !seen[check.path] { + t.Errorf("endpoint %s was not checked", check.path) + } + } +} + +func TestStripePermissionPreflightReportsSanitizedStripeError(t *testing.T) { + client := &http.Client{Transport: stripeRoundTripFunc(func(*http.Request) (*http.Response, error) { + return stripeTestResponse(http.StatusForbidden, `{"error":{"type":"invalid_request_error","code":"permission_denied","message":"secret details"}}`), nil + })} + + result, err := checkStripePermissions(context.Background(), "rk_test_placeholder", "https://stripe.test", client) + if err != nil { + t.Fatal(err) + } + if result.Ready || len(result.Checks) == 0 { + t.Fatalf("unexpected result %+v", result) + } + for _, check := range result.Checks { + if check.OK || check.StatusCode != http.StatusForbidden || check.ErrorCode != "permission_denied" || check.ErrorType != "invalid_request_error" { + t.Fatalf("unexpected check %+v", check) + } + } +} + +type stripeRoundTripFunc func(*http.Request) (*http.Response, error) + +func (fn stripeRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return fn(request) +} + +func stripeTestResponse(status int, body string) *http.Response { + return &http.Response{ + StatusCode: status, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func TestStripePermissionPreflightRejectsLiveAndMalformedKeys(t *testing.T) { + for _, key := range []string{"", "rk_live_forbidden", "not-a-stripe-key"} { + result, err := checkStripePermissions(context.Background(), key, "http://unused", nil) + if err != ErrLiveStripeKey { + t.Fatalf("key %q: error = %v, want ErrLiveStripeKey", key, err) + } + if result.Ready || result.Checks == nil || len(result.Checks) != 0 { + t.Fatalf("key %q: unexpected rejected result %+v", key, result) + } + } +} |
