package main import ( "crypto/tls" "net/http" "net/http/httptest" "strings" "testing" "time" ) func TestSessionRoundTrip(t *testing.T) { key := sessionKey("token-abc") now := time.Now().UnixMilli() value := signSession(key, now+60_000) if !verifySession(key, value, now) { t.Fatal("verifySession = false for a freshly signed cookie, want true") } } func TestSessionRejects(t *testing.T) { key := sessionKey("token-abc") now := time.Now().UnixMilli() valid := signSession(key, now+60_000) payload, sig, _ := strings.Cut(valid, ".") cases := []struct { name string value string }{ {"empty", ""}, {"no separator", payload + sig}, {"unparseable expiry", "notanumber." + sig}, {"expired", signSession(key, now-1)}, {"tampered signature", payload + "." + flipLastChar(sig)}, {"tampered expiry", "99999999999999." + sig}, {"signed with another key", signSession(sessionKey("other-token"), now+60_000)}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { if verifySession(key, tc.value, now) { t.Fatalf("verifySession(%q) = true, want false", tc.value) } }) } } func flipLastChar(s string) string { if s == "" { return "x" } last := s[len(s)-1] if last == 'A' { return s[:len(s)-1] + "B" } return s[:len(s)-1] + "A" } func TestSessionKeyDependsOnToken(t *testing.T) { a := sessionKey("token-a") b := sessionKey("token-b") if string(a) == string(b) { t.Fatal("sessionKey collided for different API tokens") } } func TestSetSessionCookieAttributes(t *testing.T) { cases := []struct { name string tls bool forwarded string wantSecure bool }{ {"plain http dev", false, "", false}, {"direct tls", true, "", true}, {"behind https proxy", false, "https", true}, {"behind http proxy", false, "http", false}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { r := httptest.NewRequest(http.MethodPost, "/login", nil) if tc.tls { r.TLS = &tls.ConnectionState{} } if tc.forwarded != "" { r.Header.Set("X-Forwarded-Proto", tc.forwarded) } rr := httptest.NewRecorder() setSessionCookie(rr, r, sessionKey("token-abc")) cookies := rr.Result().Cookies() if len(cookies) != 1 { t.Fatalf("got %d cookies, want 1", len(cookies)) } c := cookies[0] if c.Name != sessionCookieName { t.Fatalf("cookie name = %q, want %q", c.Name, sessionCookieName) } if !c.HttpOnly { t.Fatal("cookie HttpOnly = false, want true") } if c.SameSite != http.SameSiteLaxMode { t.Fatalf("cookie SameSite = %v, want Lax", c.SameSite) } if c.Path != "/" { t.Fatalf("cookie Path = %q, want /", c.Path) } if c.Secure != tc.wantSecure { t.Fatalf("cookie Secure = %v, want %v", c.Secure, tc.wantSecure) } if c.MaxAge != int(sessionTTL/time.Second) { t.Fatalf("cookie MaxAge = %d, want %d", c.MaxAge, int(sessionTTL/time.Second)) } }) } } func TestClearSessionCookie(t *testing.T) { r := httptest.NewRequest(http.MethodPost, "/logout", nil) rr := httptest.NewRecorder() clearSessionCookie(rr, r) cookies := rr.Result().Cookies() if len(cookies) != 1 { t.Fatalf("got %d cookies, want 1", len(cookies)) } if cookies[0].MaxAge >= 0 { t.Fatalf("cleared cookie MaxAge = %d, want negative", cookies[0].MaxAge) } }