summaryrefslogtreecommitdiff
path: root/internal/billing/stripe_preflight_test.go
blob: e8f1841b13378f7bdff9590ec54ec676687c9939 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
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)
		}
	}
}