182 lines
5.2 KiB
Go
182 lines
5.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: PolyForm-Noncommercial-1.0.0
|
||
|
|
|
||
|
|
package session
|
||
|
|
|
||
|
|
import (
|
||
|
|
"net/http"
|
||
|
|
"net/http/httptest"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
"testing/synctest"
|
||
|
|
"time"
|
||
|
|
)
|
||
|
|
|
||
|
|
func TestSetSaveLoadRoundTrip(t *testing.T) {
|
||
|
|
store := New("secret-key", time.Hour, false)
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
sess := &Session{data: map[string]string{}, store: store}
|
||
|
|
sess.Set("user", "petr")
|
||
|
|
store.Save(rec, sess)
|
||
|
|
|
||
|
|
cookies := rec.Result().Cookies()
|
||
|
|
if len(cookies) != 1 {
|
||
|
|
t.Fatalf("cookies = %v", cookies)
|
||
|
|
}
|
||
|
|
if cookies[0].Name != CookieName || cookies[0].HttpOnly != true ||
|
||
|
|
cookies[0].SameSite != http.SameSiteStrictMode {
|
||
|
|
t.Fatalf("cookie attrs = %+v", cookies[0])
|
||
|
|
}
|
||
|
|
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/admin/", nil)
|
||
|
|
req.AddCookie(cookies[0])
|
||
|
|
loaded := store.Load(req)
|
||
|
|
if loaded.Get("user") != "petr" {
|
||
|
|
t.Fatalf("loaded = %v", loaded.data)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestTamperedCookieRejected(t *testing.T) {
|
||
|
|
store := New("secret-key", time.Hour, false)
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
sess := &Session{data: map[string]string{}, store: store}
|
||
|
|
sess.Set("user", "petr")
|
||
|
|
store.Save(rec, sess)
|
||
|
|
value := rec.Result().Cookies()[0].Value
|
||
|
|
|
||
|
|
tampered := value[:len(value)-2] + "xx"
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
req.AddCookie(&http.Cookie{Name: CookieName, Value: tampered})
|
||
|
|
if got := store.Load(req).Get("user"); got != "" {
|
||
|
|
t.Fatalf("tampered cookie accepted: %q", got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestForeignSignatureRejected(t *testing.T) {
|
||
|
|
other := New("other-key", time.Hour, false)
|
||
|
|
store := New("secret-key", time.Hour, false)
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
sess := &Session{data: map[string]string{}, store: other}
|
||
|
|
sess.Set("user", "petr")
|
||
|
|
other.Save(rec, sess)
|
||
|
|
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
req.AddCookie(rec.Result().Cookies()[0])
|
||
|
|
if got := store.Load(req).Get("user"); got != "" {
|
||
|
|
t.Fatalf("foreign cookie accepted: %q", got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// synctest advances the clock past the one-nanosecond lifetime, so the
|
||
|
|
// test needs no sleep and cannot be flaky on a loaded machine.
|
||
|
|
func TestExpiredSessionRejected(t *testing.T) {
|
||
|
|
synctest.Test(t, func(t *testing.T) {
|
||
|
|
store := New("secret-key", time.Nanosecond, false)
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
sess := &Session{data: map[string]string{}, store: store}
|
||
|
|
sess.Set("user", "petr")
|
||
|
|
store.Save(rec, sess)
|
||
|
|
time.Sleep(time.Second)
|
||
|
|
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
req.AddCookie(rec.Result().Cookies()[0])
|
||
|
|
if got := store.Load(req).Get("user"); got != "" {
|
||
|
|
t.Fatalf("expired cookie accepted: %q", got)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestMissingCookieYieldsEmptySession(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
sess := store.Load(req)
|
||
|
|
if sess.Get("user") != "" || sess.dirty {
|
||
|
|
t.Fatalf("session = %v dirty=%v", sess.data, sess.dirty)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestCleanSessionNotPersisted(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
store.Save(rec, &Session{data: map[string]string{}, store: store})
|
||
|
|
if len(rec.Result().Cookies()) != 0 {
|
||
|
|
t.Fatal("cookie set for clean session")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestSetSameValueNotDirty(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
sess := &Session{data: map[string]string{"a": "1"}, store: store}
|
||
|
|
sess.Set("a", "1")
|
||
|
|
if sess.dirty {
|
||
|
|
t.Fatal("dirty for unchanged value")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestDeleteAndClear(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
sess := &Session{data: map[string]string{"a": "1", "b": "2"}, store: store}
|
||
|
|
sess.Delete("a")
|
||
|
|
if sess.Get("a") != "" || !sess.dirty {
|
||
|
|
t.Fatal("delete failed")
|
||
|
|
}
|
||
|
|
sess.dirty = false
|
||
|
|
sess.Delete("missing")
|
||
|
|
if sess.dirty {
|
||
|
|
t.Fatal("deleting missing key dirtied the session")
|
||
|
|
}
|
||
|
|
sess.Clear()
|
||
|
|
if sess.Get("b") != "" || !sess.dirty {
|
||
|
|
t.Fatal("clear failed")
|
||
|
|
}
|
||
|
|
sess.dirty = false
|
||
|
|
sess.Clear()
|
||
|
|
if sess.dirty {
|
||
|
|
t.Fatal("clearing empty session dirtied it")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestMiddlewareRoundTrip(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
handler := store.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||
|
|
sess := FromContext(r.Context())
|
||
|
|
sess.Set("user", "petr")
|
||
|
|
w.WriteHeader(http.StatusNoContent)
|
||
|
|
}))
|
||
|
|
|
||
|
|
rec := httptest.NewRecorder()
|
||
|
|
handler.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
|
||
|
|
cookies := rec.Result().Cookies()
|
||
|
|
if len(cookies) != 1 {
|
||
|
|
t.Fatalf("cookies = %v", cookies)
|
||
|
|
}
|
||
|
|
|
||
|
|
handler2 := store.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||
|
|
if got := FromContext(r.Context()).Get("user"); got != "petr" {
|
||
|
|
t.Errorf("user = %q", got)
|
||
|
|
}
|
||
|
|
}))
|
||
|
|
rec2 := httptest.NewRecorder()
|
||
|
|
req2 := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
req2.AddCookie(cookies[0])
|
||
|
|
handler2.ServeHTTP(rec2, req2)
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestFromContextWithoutMiddleware(t *testing.T) {
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
sess := FromContext(req.Context())
|
||
|
|
sess.Set("x", "y")
|
||
|
|
if sess.Get("x") != "y" {
|
||
|
|
t.Fatal("detached session broken")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestGarbageCookieIgnored(t *testing.T) {
|
||
|
|
store := New("k", time.Hour, false)
|
||
|
|
for _, value := range []string{"nodot", ".", "a.b", strings.Repeat("!", 50)} {
|
||
|
|
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||
|
|
req.AddCookie(&http.Cookie{Name: CookieName, Value: value})
|
||
|
|
store.Load(req) // must not panic
|
||
|
|
}
|
||
|
|
}
|