core/cookies_test.go

294 lines
8.7 KiB
Go

// Copyright (c) 2026 Zeni Kim <zenik@smarteching.com>
// Use of this source code is governed by MIT-style
// license that can be found in the LICENSE file.
package core
import (
"encoding/hex"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
)
// testCookieSecret is a 32 byte (AES-256) key expressed as a hex string, suitable
// for the COOKIE_SECRET env var used by SetCookie/ClearCookie.
const testCookieSecret = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"
// enableTemplateCookieEnv enables the template engine (required for cookie handling)
// and sets a deterministic cookie secret, restoring both env vars on cleanup.
func enableTemplateCookieEnv(t *testing.T) {
t.Helper()
prevTemplate := os.Getenv("TEMPLATE_ENABLE")
prevSecret := os.Getenv("COOKIE_SECRET")
prevSecure := os.Getenv("COOKIE_SECURE")
os.Setenv("TEMPLATE_ENABLE", "true")
os.Setenv("COOKIE_SECRET", testCookieSecret)
os.Setenv("COOKIE_SECURE", "false")
t.Cleanup(func() {
os.Setenv("TEMPLATE_ENABLE", prevTemplate)
os.Setenv("COOKIE_SECRET", prevSecret)
os.Setenv("COOKIE_SECURE", prevSecure)
})
}
func decodeSecret(t *testing.T) []byte {
t.Helper()
key, err := hex.DecodeString(testCookieSecret)
if err != nil {
t.Fatalf("failed decoding test cookie secret: %v", err)
}
return key
}
func TestCookieWrite(t *testing.T) {
w := httptest.NewRecorder()
cookie := http.Cookie{Name: "goffee", Value: "hello", Path: "/"}
if err := CookieWrite(w, cookie); err != nil {
t.Fatalf("failed testing cookie write: %v", err)
}
rsp := w.Result()
cookies := rsp.Cookies()
if len(cookies) != 1 {
t.Fatalf("expected 1 cookie, got %d", len(cookies))
}
if cookies[0].Name != "goffee" {
t.Errorf("expected cookie name 'goffee', got %q", cookies[0].Name)
}
}
func TestCookieWriteTooLong(t *testing.T) {
w := httptest.NewRecorder()
// A value large enough that the base64-encoded cookie string exceeds 4096 bytes.
cookie := http.Cookie{Name: "goffee", Value: strings.Repeat("a", 5000), Path: "/"}
err := CookieWrite(w, cookie)
if err == nil {
t.Errorf("expected ErrValueTooLong, got nil")
}
}
func TestCookieRead(t *testing.T) {
// Build a request carrying a base64-encoded cookie value.
w := httptest.NewRecorder()
cookie := http.Cookie{Name: "goffee", Value: "hello-world", Path: "/"}
if err := CookieWrite(w, cookie); err != nil {
t.Fatalf("failed writing cookie: %v", err)
}
r := httptest.NewRequest(GET, LOCALHOST, nil)
for _, c := range w.Result().Cookies() {
r.AddCookie(c)
}
val, err := CookieRead(r, "goffee")
if err != nil {
t.Fatalf("failed testing cookie read: %v", err)
}
if val != "hello-world" {
t.Errorf("expected 'hello-world', got %q", val)
}
}
func TestCookieReadMissing(t *testing.T) {
r := httptest.NewRequest(GET, LOCALHOST, nil)
_, err := CookieRead(r, "goffee")
if err == nil {
t.Errorf("expected error reading missing cookie")
}
}
func TestCookieReadInvalidBase64(t *testing.T) {
r := httptest.NewRequest(GET, LOCALHOST, nil)
r.AddCookie(&http.Cookie{Name: "goffee", Value: "!!!not-base64!!!"})
_, err := CookieRead(r, "goffee")
if err == nil {
t.Errorf("expected error reading invalid base64 cookie")
}
}
func TestCookieWriteReadEncrypted(t *testing.T) {
key := decodeSecret(t)
w := httptest.NewRecorder()
cookie := http.Cookie{Name: "goffee", Value: "secret-value", Path: "/"}
if err := CookieWriteEncrypted(w, cookie, key); err != nil {
t.Fatalf("failed writing encrypted cookie: %v", err)
}
r := httptest.NewRequest(GET, LOCALHOST, nil)
for _, c := range w.Result().Cookies() {
r.AddCookie(c)
}
val, err := CookieReadEncrypted(r, "goffee", key)
if err != nil {
t.Fatalf("failed reading encrypted cookie: %v", err)
}
if val != "secret-value" {
t.Errorf("expected 'secret-value', got %q", val)
}
}
func TestCookieReadEncryptedWrongKey(t *testing.T) {
key := decodeSecret(t)
w := httptest.NewRecorder()
cookie := http.Cookie{Name: "goffee", Value: "secret-value", Path: "/"}
if err := CookieWriteEncrypted(w, cookie, key); err != nil {
t.Fatalf("failed writing encrypted cookie: %v", err)
}
// A different but valid AES key must fail decryption.
wrongKey, _ := hex.DecodeString("ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff")
r := httptest.NewRequest(GET, LOCALHOST, nil)
for _, c := range w.Result().Cookies() {
r.AddCookie(c)
}
if _, err := CookieReadEncrypted(r, "goffee", wrongKey); err == nil {
t.Errorf("expected error decrypting with wrong key")
}
}
func TestCookieReadEncryptedWrongName(t *testing.T) {
key := decodeSecret(t)
w := httptest.NewRecorder()
cookie := http.Cookie{Name: "goffee", Value: "secret-value", Path: "/"}
if err := CookieWriteEncrypted(w, cookie, key); err != nil {
t.Fatalf("failed writing encrypted cookie: %v", err)
}
r := httptest.NewRequest(GET, LOCALHOST, nil)
for _, c := range w.Result().Cookies() {
r.AddCookie(c)
}
// Reading with a different name must fail the authenticated name check.
if _, err := CookieReadEncrypted(r, "other", key); err == nil {
t.Errorf("expected error decrypting with wrong cookie name")
}
}
func TestSetCookie(t *testing.T) {
enableTemplateCookieEnv(t)
w := httptest.NewRecorder()
if err := SetCookie(w, "user@example.com", "token-123"); err != nil {
t.Fatalf("failed testing set cookie: %v", err)
}
cookies := w.Result().Cookies()
if len(cookies) != 1 {
t.Fatalf("expected 1 cookie, got %d", len(cookies))
}
if cookies[0].Name != "goffee" {
t.Errorf("expected cookie name 'goffee', got %q", cookies[0].Name)
}
if !cookies[0].HttpOnly {
t.Errorf("expected cookie to be HttpOnly")
}
if cookies[0].MaxAge <= 0 {
t.Errorf("expected a positive MaxAge, got %d", cookies[0].MaxAge)
}
}
func TestSetCookieUsesJWT_Lifetime(t *testing.T) {
enableTemplateCookieEnv(t)
prev := os.Getenv("JWT_LIFESPAN_MINUTES")
os.Setenv("JWT_LIFESPAN_MINUTES", "5")
t.Cleanup(func() { os.Setenv("JWT_LIFESPAN_MINUTES", prev) })
w := httptest.NewRecorder()
if err := SetCookie(w, "user@example.com", "token-123"); err != nil {
t.Fatalf("failed testing set cookie: %v", err)
}
cookies := w.Result().Cookies()
if cookies[0].MaxAge != 5*60 {
t.Errorf("expected MaxAge of %d, got %d", 5*60, cookies[0].MaxAge)
}
}
func TestSetCookieWithMaxAge(t *testing.T) {
enableTemplateCookieEnv(t)
w := httptest.NewRecorder()
if err := SetCookieWithMaxAge(w, "user@example.com", "token-abc", 120); err != nil {
t.Fatalf("failed testing set cookie with max age: %v", err)
}
cookies := w.Result().Cookies()
if cookies[0].MaxAge != 120 {
t.Errorf("expected MaxAge of 120, got %d", cookies[0].MaxAge)
}
}
func TestSetCookieTemplatesDisabledPanics(t *testing.T) {
prev := os.Getenv("TEMPLATE_ENABLE")
os.Setenv("TEMPLATE_ENABLE", "false")
t.Cleanup(func() { os.Setenv("TEMPLATE_ENABLE", prev) })
defer func() {
if r := recover(); r == nil {
t.Errorf("expected panic when templates are disabled")
}
}()
w := httptest.NewRecorder()
_ = SetCookieWithMaxAge(w, "user@example.com", "token", 60)
}
func TestClearCookie(t *testing.T) {
prevSecure := os.Getenv("COOKIE_SECURE")
os.Setenv("COOKIE_SECURE", "false")
t.Cleanup(func() { os.Setenv("COOKIE_SECURE", prevSecure) })
w := httptest.NewRecorder()
if err := ClearCookie(w); err != nil {
t.Fatalf("failed testing clear cookie: %v", err)
}
cookies := w.Result().Cookies()
if len(cookies) != 1 {
t.Fatalf("expected 1 cookie, got %d", len(cookies))
}
c := cookies[0]
if c.Name != "goffee" {
t.Errorf("expected cookie name 'goffee', got %q", c.Name)
}
if c.Value != "" {
t.Errorf("expected empty cookie value, got %q", c.Value)
}
if c.MaxAge != -1 {
t.Errorf("expected MaxAge of -1 to delete the cookie, got %d", c.MaxAge)
}
if !c.HttpOnly {
t.Errorf("expected cleared cookie to be HttpOnly")
}
}
func TestGetCookieRoundTrip(t *testing.T) {
enableTemplateCookieEnv(t)
// Set a cookie and feed it back into a request, then decrypt it.
w := httptest.NewRecorder()
if err := SetCookie(w, "user@example.com", "token-xyz"); err != nil {
t.Fatalf("failed setting cookie: %v", err)
}
r := httptest.NewRequest(GET, LOCALHOST, nil)
for _, c := range w.Result().Cookies() {
r.AddCookie(c)
}
user, err := GetCookie(r)
if err != nil {
t.Fatalf("failed testing get cookie: %v", err)
}
if user.Email != "user@example.com" {
t.Errorf("expected email 'user@example.com', got %q", user.Email)
}
if user.Token != "token-xyz" {
t.Errorf("expected token 'token-xyz', got %q", user.Token)
}
}
func TestGetCookieTemplatesDisabledPanics(t *testing.T) {
prev := os.Getenv("TEMPLATE_ENABLE")
os.Setenv("TEMPLATE_ENABLE", "false")
t.Cleanup(func() { os.Setenv("TEMPLATE_ENABLE", prev) })
defer func() {
if r := recover(); r == nil {
t.Errorf("expected panic when templates are disabled")
}
}()
r := httptest.NewRequest(GET, LOCALHOST, nil)
_, _ = GetCookie(r)
}