server_test.go (15358 bytes)
1 package main 2 3 import ( 4 "crypto/hmac" 5 "crypto/sha256" 6 "encoding/hex" 7 "fmt" 8 "io" 9 "net/http" 10 "net/http/cookiejar" 11 "net/http/httptest" 12 "net/url" 13 "os" 14 "path/filepath" 15 "strings" 16 "sync" 17 "testing" 18 "time" 19 ) 20 21 // fakeStripe records the calls to the Stripe API and answers them. 22 type fakeStripe struct { 23 mu sync.Mutex 24 calls []string // "METHOD /path form" 25 } 26 27 func (f *fakeStripe) ServeHTTP(w http.ResponseWriter, r *http.Request) { 28 r.ParseForm() 29 f.mu.Lock() 30 f.calls = append(f.calls, r.Method+" "+r.URL.Path+" "+r.PostForm.Encode()) 31 f.mu.Unlock() 32 switch { 33 case r.URL.Path == "/v1/checkout/sessions": 34 io.WriteString(w, `{"url": "https://checkout.stripe.test/s1"}`) 35 case r.URL.Path == "/v1/billing_portal/sessions": 36 io.WriteString(w, `{"url": "https://billing.stripe.test/p1"}`) 37 case strings.HasPrefix(r.URL.Path, "/v1/customers/"): 38 io.WriteString(w, `{"deleted": true}`) 39 default: 40 w.WriteHeader(http.StatusNotFound) 41 io.WriteString(w, `{"error": {"message": "no such thing"}}`) 42 } 43 } 44 45 func (f *fakeStripe) called(prefix string) string { 46 f.mu.Lock() 47 defer f.mu.Unlock() 48 for _, c := range f.calls { 49 if strings.HasPrefix(c, prefix) { 50 return c 51 } 52 } 53 return "" 54 } 55 56 type harness struct { 57 t *testing.T 58 s *server 59 web *httptest.Server 60 stripe *fakeStripe 61 } 62 63 func newEnv(t *testing.T, billing bool) *harness { 64 store, err := openStore(t.TempDir()) 65 if err != nil { 66 t.Fatal(err) 67 } 68 s := &server{store: store, trial: 30 * 24 * time.Hour, 69 prices: map[string]string{"month": "price_m", "year": "price_y"}, 70 labels: map[string]string{"month": "$3 a month", "year": "$30 a year"}} 71 e := &harness{t: t, s: s} 72 if billing { 73 e.stripe = &fakeStripe{} 74 api := httptest.NewServer(e.stripe) 75 t.Cleanup(api.Close) 76 s.stripe = &Stripe{Key: "sk_test_x", Webhook: "whsec_test", API: api.URL, Client: api.Client()} 77 } 78 e.web = httptest.NewServer(s.routes()) 79 t.Cleanup(e.web.Close) 80 s.site, s.origin = e.web.URL, e.web.URL 81 return e 82 } 83 84 // browser gives a client that keeps cookies and does not follow redirects. 85 func (e *harness) browser() *http.Client { 86 jar, _ := cookiejar.New(nil) 87 return &http.Client{Jar: jar, CheckRedirect: func(*http.Request, []*http.Request) error { 88 return http.ErrUseLastResponse 89 }} 90 } 91 92 func (e *harness) post(c *http.Client, path string, form url.Values) *http.Response { 93 req, _ := http.NewRequest("POST", e.web.URL+path, strings.NewReader(form.Encode())) 94 req.Header.Set("Content-Type", "application/x-www-form-urlencoded") 95 req.Header.Set("Origin", e.web.URL) 96 resp, err := c.Do(req) 97 if err != nil { 98 e.t.Fatal(err) 99 } 100 return resp 101 } 102 103 func (e *harness) signup(email string) *http.Client { 104 c := e.browser() 105 resp := e.post(c, "/signup", url.Values{"email": {email}, "password": {"password1"}}) 106 if resp.StatusCode != http.StatusSeeOther { 107 e.t.Fatalf("signup: %s", resp.Status) 108 } 109 return c 110 } 111 112 func (e *harness) token(email string) string { 113 resp := e.post(http.DefaultClient, "/api/login", url.Values{"email": {email}, "password": {"password1"}}) 114 b, _ := io.ReadAll(resp.Body) 115 if resp.StatusCode != http.StatusOK { 116 e.t.Fatalf("api login: %s %s", resp.Status, b) 117 } 118 return strings.TrimSpace(string(b)) 119 } 120 121 // sync sends a file as a device does and gives the answer. 122 func (e *harness) sync(token, base, text string) (int, string, string) { 123 req, _ := http.NewRequest("POST", e.web.URL+"/api/sync?base="+base, strings.NewReader(text)) 124 req.Header.Set("Authorization", "Bearer "+token) 125 resp, err := http.DefaultClient.Do(req) 126 if err != nil { 127 e.t.Fatal(err) 128 } 129 defer resp.Body.Close() 130 b, _ := io.ReadAll(resp.Body) 131 return resp.StatusCode, string(b), resp.Header.Get("Sbm-Version") 132 } 133 134 func (e *harness) body(c *http.Client, path string) string { 135 resp, err := c.Get(e.web.URL + path) 136 if err != nil { 137 e.t.Fatal(err) 138 } 139 defer resp.Body.Close() 140 b, _ := io.ReadAll(resp.Body) 141 return string(b) 142 } 143 144 func TestTwoDevicesSync(t *testing.T) { 145 e := newEnv(t, false) 146 e.signup("Me@Example.org") 147 a, b := e.token("me@example.org"), e.token(" me@example.org") 148 149 code, text, va := e.sync(a, "", "https://a.org\tA\t\n") 150 if code != 200 || text != "https://a.org\tA\t\n" || len(va) != 64 { 151 t.Fatalf("first sync: %d %q %q", code, text, va) 152 } 153 _, text, vb := e.sync(b, "", "https://b.org\tB\t\n") 154 if text != "https://a.org\tA\t\nhttps://b.org\tB\t\n" { 155 t.Fatalf("second device: %q", text) 156 } 157 // Device a removes its bookmark and adds another. 158 _, text, _ = e.sync(a, va, "https://c.org\tC\t\n") 159 if text != "https://b.org\tB\t\nhttps://c.org\tC\t\n" { 160 t.Fatalf("remove and add: %q", text) 161 } 162 // Device b has no changes and gets the file of device a. 163 _, text, _ = e.sync(b, vb, "https://a.org\tA\t\nhttps://b.org\tB\t\n") 164 if text != "https://b.org\tB\t\nhttps://c.org\tC\t\n" { 165 t.Fatalf("pull: %q", text) 166 } 167 } 168 169 func TestSyncRefusals(t *testing.T) { 170 e := newEnv(t, false) 171 e.signup("me@example.org") 172 tok := e.token("me@example.org") 173 if code, _, _ := e.sync("wrong", "", "x\n"); code != http.StatusUnauthorized { 174 t.Errorf("bad token: %d", code) 175 } 176 if code, _, _ := e.sync(tok, "../../accounts.json", "x\n"); code != http.StatusBadRequest { 177 t.Errorf("bad base: %d", code) 178 } 179 if code, _, _ := e.sync(tok, "", strings.Repeat("x", maxFile+1)); code != http.StatusRequestEntityTooLarge { 180 t.Errorf("large file: %d", code) 181 } 182 // An unknown base is no error: nothing is removed. 183 if code, text, _ := e.sync(tok, strings.Repeat("0", 64), "a\n"); code != 200 || text != "a\n" { 184 t.Errorf("unknown base: %d %q", code, text) 185 } 186 } 187 188 func TestSignupAndLogin(t *testing.T) { 189 e := newEnv(t, false) 190 c := e.signup("me@example.org") 191 if !strings.Contains(e.body(c, "/account"), "Signed in as me@example.org") { 192 t.Error("the account page does not show the account") 193 } 194 if resp := e.post(e.browser(), "/signup", url.Values{"email": {"ME@example.org"}, "password": {"password2"}}); resp.StatusCode != http.StatusBadRequest { 195 t.Errorf("second signup with the same email: %s", resp.Status) 196 } 197 for _, bad := range []url.Values{ 198 {"email": {"not an address"}, "password": {"password1"}}, 199 {"email": {"Me <x@example.org>"}, "password": {"password1"}}, 200 {"email": {"y@example.org"}, "password": {"short"}}, 201 {"email": {"y@example.org"}, "password": {strings.Repeat("p", 73)}}, 202 } { 203 if resp := e.post(e.browser(), "/signup", bad); resp.StatusCode != http.StatusBadRequest { 204 t.Errorf("signup %v: %s", bad, resp.Status) 205 } 206 } 207 if resp := e.post(e.browser(), "/login", url.Values{"email": {"me@example.org"}, "password": {"wrong-password"}}); resp.StatusCode != http.StatusUnauthorized { 208 t.Errorf("wrong password: %s", resp.Status) 209 } 210 e.post(c, "/logout", nil) 211 if strings.Contains(e.body(c, "/"), "Your account") { 212 t.Error("still signed in after sign out") 213 } 214 } 215 216 func TestOtherSiteCannotPostForms(t *testing.T) { 217 e := newEnv(t, false) 218 c := e.signup("me@example.org") 219 req, _ := http.NewRequest("POST", e.web.URL+"/account/delete", strings.NewReader("password=password1")) 220 req.Header.Set("Content-Type", "application/x-www-form-urlencoded") 221 req.Header.Set("Origin", "https://evil.example") 222 resp, err := c.Do(req) 223 if err != nil { 224 t.Fatal(err) 225 } 226 if resp.StatusCode != http.StatusForbidden { 227 t.Errorf("post from another site: %s", resp.Status) 228 } 229 if e.s.store.byTok(e.token("me@example.org")) == nil { 230 t.Error("the account is gone") 231 } 232 } 233 234 func TestLoginLimit(t *testing.T) { 235 e := newEnv(t, false) 236 e.signup("me@example.org") 237 var last int 238 for range 21 { 239 resp := e.post(http.DefaultClient, "/api/login", url.Values{"email": {"me@example.org"}, "password": {"guess"}}) 240 last = resp.StatusCode 241 } 242 if last != http.StatusTooManyRequests { 243 t.Errorf("attempt 21: %d", last) 244 } 245 } 246 247 func TestDeleteAccount(t *testing.T) { 248 e := newEnv(t, true) 249 c := e.signup("me@example.org") 250 tok := e.token("me@example.org") 251 e.sync(tok, "", "a\n") 252 a := e.s.store.byTok(tok) 253 e.s.store.setCustomer(a, "cus_1") 254 255 if resp := e.post(c, "/account/delete", url.Values{"password": {"wrong-password"}}); resp.StatusCode != http.StatusUnauthorized { 256 t.Fatalf("delete with a wrong password: %s", resp.Status) 257 } 258 if resp := e.post(c, "/account/delete", url.Values{"password": {"password1"}}); resp.StatusCode != http.StatusSeeOther { 259 t.Fatalf("delete: %s", resp.Status) 260 } 261 if e.stripe.called("DELETE /v1/customers/cus_1") == "" { 262 t.Error("the Stripe customer is not deleted") 263 } 264 if code, _, _ := e.sync(tok, "", "a\n"); code != http.StatusUnauthorized { 265 t.Errorf("sync after delete: %d", code) 266 } 267 if _, err := os.Stat(filepath.Join(e.s.store.dir, "files", a.ID)); !os.IsNotExist(err) { 268 t.Errorf("the files of the account are still there: %v", err) 269 } 270 } 271 272 func TestTrialAndSubscription(t *testing.T) { 273 e := newEnv(t, true) 274 c := e.signup("me@example.org") 275 tok := e.token("me@example.org") 276 if code, _, _ := e.sync(tok, "", "a\n"); code != 200 { 277 t.Fatalf("sync in the trial: %d", code) 278 } 279 if !strings.Contains(e.body(c, "/account"), "Free trial: 30 days left.") { 280 t.Error("the account page does not show the trial") 281 } 282 a := e.s.store.byTok(tok) 283 a.Created = time.Now().Add(-31 * 24 * time.Hour) 284 if code, text, _ := e.sync(tok, "", "a\n"); code != http.StatusPaymentRequired { 285 t.Fatalf("sync after the trial: %d %q", code, text) 286 } 287 288 resp := e.post(c, "/account/checkout", url.Values{"plan": {"year"}}) 289 if resp.StatusCode != http.StatusSeeOther || resp.Header.Get("Location") != "https://checkout.stripe.test/s1" { 290 t.Fatalf("checkout: %s to %q", resp.Status, resp.Header.Get("Location")) 291 } 292 call := e.stripe.called("POST /v1/checkout/sessions") 293 for _, want := range []string{"client_reference_id=" + a.ID, "price%5D=price_y", "customer_email=me%40example.org", "mode=subscription"} { 294 if !strings.Contains(call, want) { 295 t.Errorf("checkout call %q lacks %q", call, want) 296 } 297 } 298 299 now := time.Now().Unix() 300 e.event(now, "checkout.session.completed", fmt.Sprintf(`{"customer":"cus_9","client_reference_id":%q}`, a.ID)) 301 e.event(now+2, "customer.subscription.updated", `{"customer":"cus_9","status":"active"}`) 302 e.event(now+1, "customer.subscription.created", `{"customer":"cus_9","status":"incomplete"}`) // late 303 if code, _, _ := e.sync(tok, "", "a\n"); code != 200 { 304 t.Fatalf("sync with a subscription: %d", code) 305 } 306 if page := e.body(c, "/account"); !strings.Contains(page, "Your subscription is active.") || strings.Contains(page, "Subscribe:") { 307 t.Error("the account page does not show the subscription") 308 } 309 if resp := e.post(c, "/account/portal", nil); resp.Header.Get("Location") != "https://billing.stripe.test/p1" { 310 t.Errorf("portal: %s", resp.Status) 311 } 312 e.event(now+3, "customer.subscription.deleted", `{"customer":"cus_9","status":"canceled"}`) 313 if code, _, _ := e.sync(tok, "", "a\n"); code != http.StatusPaymentRequired { 314 t.Errorf("sync after the end of the subscription: %d", code) 315 } 316 } 317 318 func TestTrialStartsWithBilling(t *testing.T) { 319 e := newEnv(t, true) 320 e.signup("me@example.org") 321 tok := e.token("me@example.org") 322 a := e.s.store.byTok(tok) 323 a.Created = time.Now().Add(-100 * 24 * time.Hour) // from before billing 324 e.s.since = time.Now().Add(-10 * 24 * time.Hour) 325 if code, _, _ := e.sync(tok, "", "a\n"); code != 200 { 326 t.Errorf("sync 10 days after billing started: %d", code) 327 } 328 e.s.since = time.Now().Add(-31 * 24 * time.Hour) 329 if code, _, _ := e.sync(tok, "", "a\n"); code != http.StatusPaymentRequired { 330 t.Errorf("sync 31 days after billing started: %d", code) 331 } 332 } 333 334 func TestSubscriptionFoundByMetadata(t *testing.T) { 335 e := newEnv(t, true) 336 e.signup("me@example.org") 337 a := e.s.store.byTok(e.token("me@example.org")) 338 e.event(time.Now().Unix(), "customer.subscription.created", 339 fmt.Sprintf(`{"customer":"cus_5","status":"active","metadata":{"account":%q}}`, a.ID)) 340 if v := e.s.store.view(a); v.Customer != "cus_5" || v.Status != "active" { 341 t.Errorf("account: customer %q, status %q", v.Customer, v.Status) 342 } 343 } 344 345 // event sends a signed webhook event as Stripe does. 346 func (e *harness) event(created int64, typ, object string) { 347 body := fmt.Sprintf(`{"type":%q,"created":%d,"data":{"object":%s}}`, typ, created, object) 348 req, _ := http.NewRequest("POST", e.web.URL+"/stripe", strings.NewReader(body)) 349 req.Header.Set("Stripe-Signature", sign([]byte(body), "whsec_test", time.Now().Unix())) 350 resp, err := http.DefaultClient.Do(req) 351 if err != nil { 352 e.t.Fatal(err) 353 } 354 if resp.StatusCode != 200 { 355 e.t.Fatalf("event %s: %s", typ, resp.Status) 356 } 357 } 358 359 func sign(body []byte, secret string, t int64) string { 360 mac := hmac.New(sha256.New, []byte(secret)) 361 fmt.Fprintf(mac, "%d.%s", t, body) 362 return fmt.Sprintf("t=%d,v1=%s", t, hex.EncodeToString(mac.Sum(nil))) 363 } 364 365 func TestVerify(t *testing.T) { 366 body := []byte(`{"type":"x"}`) 367 now := time.Unix(1_800_000_000, 0) 368 good := sign(body, "whsec_a", now.Unix()) 369 if err := verify(body, good, "whsec_a", now); err != nil { 370 t.Errorf("good signature: %v", err) 371 } 372 if err := verify(body, "t=1,v1=00,"+strings.TrimPrefix(good, "t=1800000000,"), "whsec_a", now); err == nil { 373 t.Error("a signature for another time passes") 374 } 375 if err := verify(body, "v1=00,"+good, "whsec_a", now); err != nil { 376 t.Errorf("one good signature among more: %v", err) 377 } 378 for name, c := range map[string]struct { 379 body string 380 header string 381 secret string 382 now time.Time 383 }{ 384 "changed body": {`{"type":"y"}`, good, "whsec_a", now}, 385 "other secret": {string(body), good, "whsec_b", now}, 386 "old": {string(body), good, "whsec_a", now.Add(6 * time.Minute)}, 387 "no header": {string(body), "", "whsec_a", now}, 388 "no secret set": {string(body), good, "", now}, 389 } { 390 if err := verify([]byte(c.body), c.header, c.secret, c.now); err == nil { 391 t.Errorf("%s: passes", name) 392 } 393 } 394 e := newEnv(t, true) 395 req, _ := http.NewRequest("POST", e.web.URL+"/stripe", strings.NewReader(`{}`)) 396 req.Header.Set("Stripe-Signature", "t=1,v1=00") 397 if resp, _ := http.DefaultClient.Do(req); resp.StatusCode != http.StatusBadRequest { 398 t.Errorf("webhook with a bad signature: %s", resp.Status) 399 } 400 } 401 402 func TestMoney(t *testing.T) { 403 for cents, want := range map[int64]string{300: "$3", 3000: "$30", 299: "$2.99", 5: "$0.05"} { 404 if got := money(cents, "usd"); got != want { 405 t.Errorf("money(%d) = %q, want %q", cents, got, want) 406 } 407 } 408 if got := money(500, "eur"); got != "5 EUR" { 409 t.Errorf("euro: %q", got) 410 } 411 } 412 413 func TestStoreSurvivesRestart(t *testing.T) { 414 dir := t.TempDir() 415 s, _ := openStore(dir) 416 a, err := s.signup("me@example.org", "password1") 417 if err != nil { 418 t.Fatal(err) 419 } 420 tok, _ := s.signIn(a) 421 s.sync(a, "", "a\n") 422 s2, err := openStore(dir) 423 if err != nil { 424 t.Fatal(err) 425 } 426 b := s2.byTok(tok) 427 if b == nil || b.Email != "me@example.org" { 428 t.Fatal("the account or its token is lost") 429 } 430 if text, _, _ := s2.sync(b, "", ""); text != "a\n" { 431 t.Errorf("the file is lost: %q", text) 432 } 433 } 434 435 func TestOldVersionsGo(t *testing.T) { 436 s, _ := openStore(t.TempDir()) 437 a, _ := s.signup("me@example.org", "password1") 438 base := "" 439 for i := range maxVersions + 10 { 440 _, v, err := s.sync(a, base, fmt.Sprintf("line %d\n", i)) 441 if err != nil { 442 t.Fatal(err) 443 } 444 base = v 445 time.Sleep(2 * time.Millisecond) // distinct times of change 446 } 447 entries, _ := os.ReadDir(filepath.Join(s.dir, "files", a.ID)) 448 n := 0 449 for _, e := range entries { 450 if hashName.MatchString(e.Name()) { 451 n++ 452 } 453 } 454 if n != maxVersions { 455 t.Errorf("%d versions kept, want %d", n, maxVersions) 456 } 457 }