session_test.go
⎇
Raw
1package web
2
3import (
4 "context"
5 "net/http"
6 "net/http/httptest"
7 "net/url"
8 "strconv"
9 "strings"
10 "sync"
11 "testing"
12 "time"
13
14 "github.com/go-webauthn/webauthn/webauthn"
15
16 "hearthforge/internal/db"
17)
18
19func getWithCookie(h http.Handler, path string, c *http.Cookie) *httptest.ResponseRecorder {
20 req := httptest.NewRequest(http.MethodGet, path, nil)
21 req.AddCookie(c)
22 rec := httptest.NewRecorder()
23 h.ServeHTTP(rec, req)
24 return rec
25}
26
27func TestPasswordChangeDropsOtherSessions(t *testing.T) {
28 _, h := newTestServer(t)
29 current, other := loginAsAdmin(t, h), loginAsAdmin(t, h)
30 rec := postForm(t, h, "/settings/password", url.Values{
31 "current_password": {testAdminPassword},
32 "new_password": {"new-password-1"}, "confirm_password": {"new-password-1"},
33 }, current)
34 if loc := rec.Header().Get("Location"); loc != "/settings?success=password" {
35 t.Fatalf("Location = %q", loc)
36 }
37 if rec := getWithCookie(h, "/settings", other); rec.Code != http.StatusFound {
38 t.Errorf("other session still works: status %d", rec.Code)
39 }
40 if rec := getWithCookie(h, "/settings", current); rec.Code != http.StatusOK {
41 t.Errorf("current session was dropped: status %d", rec.Code)
42 }
43}
44
45func TestStaleSessionMustSignInAgain(t *testing.T) {
46 s, h := newTestServer(t)
47 cookie := loginAsAdmin(t, h)
48 const note = "needs a recent sign-in"
49 if strings.Contains(getWithCookie(h, "/settings", cookie).Body.String(), note) {
50 t.Error("fresh session shows the sign-in note")
51 }
52 old := time.Now().UTC().Add(-time.Hour).Format(db.ISOLayout)
53 if _, err := s.DB.ExecContext(context.Background(), `UPDATE sessions SET created_at = ?`, old); err != nil {
54 t.Fatal(err)
55 }
56 if page := getWithCookie(h, "/settings", cookie).Body.String(); !strings.Contains(page, note) ||
57 !strings.Contains(page, `href="/login?next=%2Fsettings"`) {
58 t.Error("stale session lacks the sign-in note")
59 }
60
61 rec := postForm(t, h, "/settings/password/remove", nil, cookie)
62 if loc := rec.Header().Get("Location"); loc != "/login?next=%2Fsettings" {
63 t.Fatalf("remove password Location = %q", loc)
64 }
65 rec = postForm(t, h, "/settings/passkey/revoke", url.Values{"id": {"1"}}, cookie)
66 if loc := rec.Header().Get("Location"); loc != "/login?next=%2Fsettings" {
67 t.Fatalf("revoke passkey Location = %q", loc)
68 }
69 rec = postForm(t, h, "/auth/passkey/register/options", nil, cookie)
70 if rec.Code != http.StatusForbidden || !strings.Contains(rec.Body.String(), `"reauth"`) {
71 t.Fatalf("passkey options: %d %s", rec.Code, rec.Body.String())
72 }
73
74 page := getWithCookie(h, "/login?next=%2Fsettings", cookie).Body.String()
75 if !strings.Contains(page, `name="next" value="/settings"`) || !strings.Contains(page, "sign in again") {
76 t.Error("login page lacks the return path or the notice")
77 }
78 for next, want := range map[string]string{"/settings": "/settings", "//evil.example": "/", "/\\evil.example": "/"} {
79 rec = postForm(t, h, "/login", url.Values{
80 "username": {db.AdminUsername}, "password": {testAdminPassword}, "next": {next},
81 })
82 if loc := rec.Header().Get("Location"); loc != want {
83 t.Errorf("next %q: Location = %q, want %q", next, loc, want)
84 }
85 }
86}
87
88func TestPasswordChangeIsRateLimited(t *testing.T) {
89 s, h := newTestServer(t)
90 s.Cfg.RateLimitDisabled = false
91 cookie := loginAsAdmin(t, h)
92 var loc string
93 for range 11 {
94 rec := postForm(t, h, "/settings/password", url.Values{
95 "new_password": {"new-password-1"}, "confirm_password": {"mismatch-1"},
96 }, cookie)
97 loc, _ = url.QueryUnescape(rec.Header().Get("Location"))
98 }
99 if !strings.Contains(loc, "Too many attempts") {
100 t.Errorf("Location = %q", loc)
101 }
102}
103
104func TestReservedUsernameIgnoresCase(t *testing.T) {
105 _, h := newTestServer(t)
106 rec := postForm(t, h, "/register", url.Values{
107 "username": {"Admin"}, "password": {"password123"}, "password2": {"password123"},
108 })
109 if !strings.Contains(rec.Body.String(), "reserved") {
110 t.Errorf("body = %q", rec.Body.String())
111 }
112}
113
114func registryRequest(user, pass, remote string) *http.Request {
115 req := httptest.NewRequest(http.MethodGet, "/v2/", nil)
116 req.SetBasicAuth(user, pass)
117 req.RemoteAddr = remote
118 return req
119}
120
121func TestRegistryUserCachesVerifiedCredentials(t *testing.T) {
122 s, _ := newTestServer(t)
123 s.Cfg.RateLimitDisabled = false
124 const ip = "198.51.100.7:1"
125 // More requests than the limiter allows: only the first runs argon2.
126 for i := range 15 {
127 if _, _, ok := s.registryUser(registryRequest(db.AdminUsername, testAdminPassword, ip)); !ok {
128 t.Fatalf("request %d rejected", i)
129 }
130 }
131 hash, err := db.HashPassword("changed-password")
132 if err != nil {
133 t.Fatal(err)
134 }
135 if err := s.DB.SetPasswordHash(context.Background(), 1, &hash); err != nil {
136 t.Fatal(err)
137 }
138 if _, _, ok := s.registryUser(registryRequest(db.AdminUsername, testAdminPassword, ip)); ok {
139 t.Error("cached credentials survived a password change")
140 }
141}
142
143func TestRegistryUserCountsParallelGuesses(t *testing.T) {
144 s, _ := newTestServer(t)
145 s.Cfg.RateLimitDisabled = false
146 const ip = "198.51.100.8:1"
147 var wg sync.WaitGroup
148 for range 10 {
149 wg.Go(func() { s.registryUser(registryRequest(db.AdminUsername, "wrong", ip)) })
150 }
151 wg.Wait()
152 if _, _, ok := s.registryUser(registryRequest(db.AdminUsername, testAdminPassword, ip)); ok {
153 t.Error("correct password accepted after the limit was used up")
154 }
155}
156
157func TestChallengeStoreEvictsOldestWhenFull(t *testing.T) {
158 c := challengeStore{entries: map[string]challengeEntry{}}
159 c.put("first", webauthn.SessionData{})
160 later := time.Now().Add(time.Hour)
161 for i := range maxChallenges - 1 {
162 c.entries[strconv.Itoa(i)] = challengeEntry{expires: later}
163 }
164 c.put("new", webauthn.SessionData{})
165 if _, ok := c.peek("new"); !ok {
166 t.Error("new entry was refused")
167 }
168 if _, ok := c.peek("first"); ok {
169 t.Error("oldest entry was not evicted")
170 }
171 if len(c.entries) != maxChallenges {
172 t.Errorf("entries = %d, want %d", len(c.entries), maxChallenges)
173 }
174}
175