package main import ( "context" "fmt" "net/http" "net/http/httptest" "net/url" "os" "strings" "testing" "time" "github.com/jackc/pgx/v5/pgxpool" ) // Set TEST_DATABASE_URL to a throwaway database — these drop and recreate the public schema. func testApp(t *testing.T) *app { t.Helper() dbURL := os.Getenv("TEST_DATABASE_URL") if dbURL == "" { t.Skip("TEST_DATABASE_URL not set") } ctx := context.Background() pool, err := pgxpool.New(ctx, dbURL) if err != nil { t.Fatal(err) } t.Cleanup(pool.Close) if _, err := pool.Exec(ctx, `drop schema public cascade; create schema public`); err != nil { t.Fatal(err) } if err := migrate(ctx, pool); err != nil { t.Fatal(err) } return &app{cfg: config{adminUser: "admin", adminPass: "s3cret"}, pool: pool} } func post(t *testing.T, h http.Handler, path string, form url.Values) *httptest.ResponseRecorder { t.Helper() r := httptest.NewRequest("POST", path, strings.NewReader(form.Encode())) r.Header.Set("Content-Type", "application/x-www-form-urlencoded") w := httptest.NewRecorder() h.ServeHTTP(w, r) return w } func (a *app) inviteValid(t *testing.T, code string) bool { t.Helper() var valid bool if err := a.pool.QueryRow(context.Background(), `select is_valid from invites where code = $1`, code).Scan(&valid); err != nil { t.Fatal(err) } return valid } // A failed registration must leave the code usable; a successful one must not. func TestInviteIsSpentOnlyBySuccess(t *testing.T) { a := testApp(t) ctx := context.Background() mux := a.withMember(a.memberMux()) if _, err := a.pool.Exec(ctx, `insert into invites (code) values ('kutsu1')`); err != nil { t.Fatal(err) } if _, err := a.pool.Exec(ctx, `insert into users (name, email, password_hash) values ('Esa', 'esa@example.com', 'x')`); err != nil { t.Fatal(err) } // Taken email — the insert fails after the invite has already been marked spent in the tx. w := post(t, mux, "/register", url.Values{ "code": {"kutsu1"}, "name": {"Toinen"}, "email": {"esa@example.com"}, "password": {"salasana1"}, }) if w.Code != http.StatusUnprocessableEntity { t.Fatalf("duplicate email: status = %d, want 422", w.Code) } if !a.inviteValid(t, "kutsu1") { t.Fatal("failed registration spent the invite") } // Missing password — rejected before the invite is touched at all. w = post(t, mux, "/register", url.Values{ "code": {"kutsu1"}, "name": {"Toinen"}, "email": {"toinen@example.com"}, "password": {""}, }) if w.Code != http.StatusUnprocessableEntity { t.Fatalf("short password: status = %d, want 422", w.Code) } if !a.inviteValid(t, "kutsu1") { t.Fatal("rejected registration spent the invite") } w = post(t, mux, "/register", url.Values{ "code": {"kutsu1"}, "name": {"Toinen"}, "email": {"toinen@example.com"}, "password": {"salasana1"}, }) if w.Code != http.StatusSeeOther { t.Fatalf("valid registration: status = %d, want 303", w.Code) } if a.inviteValid(t, "kutsu1") { t.Fatal("successful registration left the invite usable") } // And it cannot be used twice. w = post(t, mux, "/register", url.Values{ "code": {"kutsu1"}, "name": {"Kolmas"}, "email": {"kolmas@example.com"}, "password": {"salasana1"}, }) if w.Code != http.StatusUnprocessableEntity { t.Fatalf("reused invite: status = %d, want 422", w.Code) } } // The limiter has its own unit test; this covers the wiring into the handler. func TestLoginHandlerRefusesAfterTooManyFailures(t *testing.T) { a := testApp(t) mux := a.withMember(a.memberMux()) a.seedMember(t, "esa@example.com") bad := url.Values{"email": {"esa@example.com"}, "password": {"väärin"}} for i := range loginMaxFailures { if w := post(t, mux, "/login", bad); w.Code != http.StatusUnauthorized { t.Fatalf("attempt %d: status = %d, want 401", i+1, w.Code) } } if w := post(t, mux, "/login", bad); w.Code != http.StatusTooManyRequests { t.Fatalf("attempt %d: status = %d, want 429", loginMaxFailures+1, w.Code) } } func (a *app) seedMember(t *testing.T, email string) int64 { t.Helper() var id int64 err := a.pool.QueryRow(context.Background(), `insert into users (name, email, password_hash) values ('Esa', $1, 'x') returning id`, email).Scan(&id) if err != nil { t.Fatal(err) } return id } func (a *app) sessionFor(t *testing.T, token string) *member { t.Helper() r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: sessionCookie, Value: token}) return a.session(httptest.NewRecorder(), r) } func TestSessionIdleTimeout(t *testing.T) { a := testApp(t) ctx := context.Background() id := a.seedMember(t, "esa@example.com") live, _, err := a.startSession(ctx, id, false) if err != nil { t.Fatal(err) } if m := a.sessionFor(t, live); m == nil || m.ID != id { t.Fatal("fresh session did not resolve to its member") } // Age it past the idle window: the timeout is what expiry means, so this is the whole rule. if _, err := a.pool.Exec(ctx, `update sessions set expires_at = now() - interval '1 second' where token = $1`, live); err != nil { t.Fatal(err) } if m := a.sessionFor(t, live); m != nil { t.Fatal("expired session still resolved") } // A session used inside the window slides forward. fresh, _, err := a.startSession(ctx, id, false) if err != nil { t.Fatal(err) } if _, err := a.pool.Exec(ctx, `update sessions set expires_at = now() + interval '1 hour' where token = $1`, fresh); err != nil { t.Fatal(err) } if m := a.sessionFor(t, fresh); m == nil { t.Fatal("session inside the window did not resolve") } var expires time.Time if err := a.pool.QueryRow(ctx, `select expires_at from sessions where token = $1`, fresh).Scan(&expires); err != nil { t.Fatal(err) } if time.Until(expires) < 23*time.Hour { t.Fatalf("session was not extended: expires in %s", time.Until(expires)) } } func TestBanDropsSessionsAndBlocksLogin(t *testing.T) { a := testApp(t) ctx := context.Background() mux := a.withMember(a.memberMux()) if _, err := a.pool.Exec(ctx, `insert into invites (code) values ('kutsu2')`); err != nil { t.Fatal(err) } w := post(t, mux, "/register", url.Values{ "code": {"kutsu2"}, "name": {"Esa"}, "email": {"esa@example.com"}, "password": {"salasana1"}, }) if w.Code != http.StatusSeeOther { t.Fatalf("registration: status = %d, want 303", w.Code) } var id int64 if err := a.pool.QueryRow(ctx, `select id from users where email = 'esa@example.com'`).Scan(&id); err != nil { t.Fatal(err) } adminMux := a.adminMux() if w := post(t, adminMux, fmt.Sprintf("/admin/users/%d/ban", id), nil); w.Code != http.StatusSeeOther { t.Fatalf("ban: status = %d, want 303", w.Code) } var sessions int if err := a.pool.QueryRow(ctx, `select count(*) from sessions where user_id = $1`, id).Scan(&sessions); err != nil { t.Fatal(err) } if sessions != 0 { t.Fatalf("banned member kept %d sessions", sessions) } w = post(t, mux, "/login", url.Values{"email": {"esa@example.com"}, "password": {"salasana1"}}) if w.Code != http.StatusForbidden { t.Fatalf("banned login: status = %d, want 403", w.Code) } // Reversible: unban, and the same credentials work again. if w := post(t, adminMux, fmt.Sprintf("/admin/users/%d/ban", id), nil); w.Code != http.StatusSeeOther { t.Fatalf("unban: status = %d, want 303", w.Code) } w = post(t, mux, "/login", url.Values{"email": {"esa@example.com"}, "password": {"salasana1"}}) if w.Code != http.StatusSeeOther { t.Fatalf("login after unban: status = %d, want 303", w.Code) } }