Files
OwnCord/Server/api/auth_handler_test.go
T
J3vbandClaude Fable 5 75d64dd412 refactor(b3-2): auth vertical slice — service.AuthService behind a consumer-owned interface (S-10) + HP-3 draft (#1450)
* docs(b3-1): record PR #1449 = 71d867cb in the status line, step table and evidence block

Pre-squash SHAs completed with the coverage commit a0356ee1 and the three
Codex rounds (head 8614603b).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* refactor(b3-2): auth_deps.go — the consumer-owned AuthService interface

Eight methods beside the handlers that need them: Register, Login,
VerifyTOTP, Logout, DeleteAccount, EnableTOTP, ConfirmTOTP, DisableTOTP —
fewer than the ten *db.DB methods the two handlers call today. The input
and result types they name (Principal, RegisterInput, LoginInput,
AuthResult, TOTPChangeResult) and the AuthBroadcaster the delete path needs
live in service/auth.go. Nothing implements or calls the interface yet.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* refactor(b3-2): service.AuthService — the auth orchestration, moved verbatim

Register, Login, VerifyTOTP, Logout, DeleteAccount, EnableTOTP, ConfirmTOTP,
DisableTOTP and the RegistrationPolicy gate two characterization rows pin
ahead of the body read. The enumeration guard, the F3 reserve-before-compare,
the audit writes, the best-effort custom-status clear and the 200+warning
partial-success contract move line for line; persistence stays in db behind
Store. Each refusal is a named service.Err* whose Error() is the exact
public message the handler wrote and whose category (ErrUnauthorized and
ErrInvalidInput join the message.go set) the transport maps to a status.
The auth rate multiplier moves to auth/ratescale.go so the route mounts
and the login failure accounting read one value; api keeps its wrappers.
Nothing calls the service yet — the handlers still own their copies.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* refactor(b3-2): thin auth handlers — decode, call AuthService, encode

*db.DB leaves every handler signature in auth_handler.go and
totp_handler.go; MountAuthRoutes takes the interface and the
AuthMiddleware the caller builds, and router.go constructs the service
after the hub. Each refusal is encoded by one writeAuthError switch on the
service's error categories. The principal helper in middleware.go hands
the handlers the caller as service.Principal, and userResponse moves next
to the profile handler, so neither auth file names db any more: their two
DBImportAllow rows go in this commit (TestDBImportAllowIsLive proves the
rows could not outlive the import) and the boundary fixture points at
middleware.go instead. The auth-slice limits leave api/constants.go with
the code that reads them; profile_handler.go reads the shared pw_confirm
budget from the service. Test files change only where they mount the
routes (four helper lines + two direct mounts); no assertion or row moves.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* docs(b3-2): after-state boundary inventory — api db importers 12 → 10

Regenerated table (49 files; move 28 → 26), the auth slice's after-state
dependency rows, and the honest reading of the plan's "neither db nor
service" target: met for db, not for service — the handlers import service
for the interface's types and Err* categories.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* docs(b3-2): evidence block — pre-squash SHAs, graph deltas, gates, coverage

Characterization green at each SHA in a detached worktree with the frozen
files byte-identical to 71d867cb; nine-method interface vs ten db methods;
api db importers 12 → 10; slice coverage 392/433 = 90.5% → 392/427 = 91.8%;
the five behaviour notes (decode-before-gate corner cases, shared
AuthMiddleware, folded confirmation block, moved limits, moved converter).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* docs(hp-3): scorecard draft and the D4 vertical-slice pattern in server.md

Five questions answered with commands and outputs at fe1d11b8/3f0d24ec;
owner sign-off line left blank. server.md gains D4 — the eight-step
interface/service/handler rule for B3-8 with the awkward step
(gate-before-decode) named — and its D3 deviation note drops the auth
routes. Plans README indexes the scorecard.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

* docs(b3-2): record PR #1450 in the evidence block and the HP-3 fetch line

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01A17Uq3d2C36rN82Jitf3wo

---------

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-08-30 09:49:05 +02:00

1834 lines
66 KiB
Go

package api_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync"
"testing"
"testing/fstest"
"time"
"github.com/J3vb/OwnCord/Server/api"
"github.com/J3vb/OwnCord/Server/auth"
"github.com/J3vb/OwnCord/Server/db"
"github.com/J3vb/OwnCord/Server/service"
"github.com/go-chi/chi/v5"
)
// newAuthTestDB builds an in-memory DB with the full schema needed for auth tests.
func newAuthTestDB(t *testing.T) *db.DB {
t.Helper()
database, err := db.Open(":memory:")
if err != nil {
t.Fatalf("db.Open: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
migrFS := fstest.MapFS{
"001_schema.sql": {Data: apiTestSchema},
}
if err := db.MigrateFS(database, migrFS); err != nil {
t.Fatalf("MigrateFS: %v", err)
}
return database
}
// buildAuthRouter returns a chi router with auth routes mounted on /api/v1/auth.
func buildAuthRouter(database *db.DB, limiter *auth.RateLimiter) http.Handler {
return buildAuthRouterWithProxies(database, limiter, nil)
}
func buildAuthRouterWithProxies(database *db.DB, limiter *auth.RateLimiter, trustedProxies []string) http.Handler {
r := chi.NewRouter()
api.MountAuthRoutes(r, service.NewAuthService(database, limiter, testTOTPKey, nil), api.AuthMiddleware(database), limiter, trustedProxies)
return r
}
// testTOTPKey is a fixed 32-byte AES-256 key used in tests.
var testTOTPKey = make([]byte, 32)
// postJSON is a test helper that POSTs JSON to the given router.
func postJSON(t *testing.T, router http.Handler, path string, body any) *httptest.ResponseRecorder {
t.Helper()
raw, _ := json.Marshal(body)
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
return rr
}
// postJSONWithToken posts with an Authorization header.
func postJSONWithToken(t *testing.T, router http.Handler, path, token string, body any) *httptest.ResponseRecorder {
t.Helper()
raw, _ := json.Marshal(body)
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+token)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
return rr
}
func postJSONFromIP(t *testing.T, router http.Handler, path string, body any, ip string) *httptest.ResponseRecorder {
t.Helper()
raw, _ := json.Marshal(body)
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = ip + ":9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
return rr
}
// getWithToken performs a GET with an Authorization header.
func getWithToken(t *testing.T, router http.Handler, path, token string) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodGet, path, nil)
req.Header.Set("Authorization", "Bearer "+token)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
return rr
}
// ─── Register tests ───────────────────────────────────────────────────────────
func TestRegister_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
// Create an invite first.
ownerID, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
"password": "securePass1",
"invite_code": code,
})
if rr.Code != http.StatusCreated {
t.Errorf("Register status = %d, want 201; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(rr.Body).Decode(&resp)
if resp["token"] == nil {
t.Error("Register response missing token")
}
if resp["user"] == nil {
t.Error("Register response missing user")
}
}
func TestRegister_RegistrationClosed(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = '0' WHERE key = 'registration_open'`); err != nil {
t.Fatalf("close registration: %v", err)
}
ownerID, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "closeduser",
"password": "securePass1",
"invite_code": code,
})
if rr.Code != http.StatusForbidden {
t.Fatalf("Register status = %d, want 403; body = %s", rr.Code, rr.Body.String())
}
}
func TestRegister_InvalidInvite(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
"password": "securePass1",
"invite_code": "bogus",
})
if rr.Code != http.StatusBadRequest {
t.Errorf("Register invalid invite status = %d, want 400", rr.Code)
}
}
func TestRegister_WeakPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "owner2", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
"password": "short",
"invite_code": code,
})
if rr.Code != http.StatusBadRequest {
t.Errorf("Register weak password status = %d, want 400", rr.Code)
}
}
func TestRegister_InviteUsedUp(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "owner3", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil) // max 1 use
// First registration should succeed.
postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "user1",
"password": "securePass1",
"invite_code": code,
})
// Second should fail — invite exhausted.
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "user2",
"password": "securePass2",
"invite_code": code,
})
if rr.Code != http.StatusBadRequest {
t.Errorf("Register exhausted invite status = %d, want 400", rr.Code)
}
}
func TestRegister_DuplicateUsername_DoesNotConsumeInvite(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "owner4", "hash", 1)
_, _ = database.CreateUser(context.Background(), "takenuser", "hash", 4)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
duplicate := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "takenuser",
"password": "securePass1",
"invite_code": code,
})
if duplicate.Code != http.StatusBadRequest {
t.Fatalf("duplicate username status = %d, want 400; body = %s", duplicate.Code, duplicate.Body.String())
}
success := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "freshuser",
"password": "securePass2",
"invite_code": code,
})
if success.Code != http.StatusCreated {
t.Fatalf("invite should remain usable after failed registration, status = %d, want 201; body = %s", success.Code, success.Body.String())
}
}
func TestRegister_MissingFields(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{})
if rr.Code != http.StatusBadRequest {
t.Errorf("Register missing fields status = %d, want 400", rr.Code)
}
}
func TestRegister_ErrorNeverRevealUsername(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "someone",
"password": "securePass1",
"invite_code": "bogus",
})
body := rr.Body.String()
// Must not hint that the username doesn't exist or the invite is invalid specifically
if contains(body, "username") && contains(body, "taken") {
t.Error("Register error message reveals username status")
}
}
// ─── Login tests ──────────────────────────────────────────────────────────────
func TestLogin_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "loginuser", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "loginuser",
"password": "correctPass1",
})
if rr.Code != http.StatusOK {
t.Errorf("Login status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(rr.Body).Decode(&resp)
if resp["token"] == nil {
t.Error("Login response missing token")
}
}
func TestLogin_WrongPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "loginuser2", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "loginuser2",
"password": "wrongpassword",
})
if rr.Code != http.StatusUnauthorized {
t.Errorf("Login wrong password status = %d, want 401", rr.Code)
}
}
func TestLogin_UnknownUser(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "nobody",
"password": "anypass123",
})
if rr.Code != http.StatusUnauthorized {
t.Errorf("Login unknown user status = %d, want 401", rr.Code)
}
}
func TestLogin_LockoutUsesTrustedForwardedIP(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouterWithProxies(database, limiter, []string{"127.0.0.0/8"})
for range 10 {
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", bytes.NewReader([]byte(`{"username":"nobody","password":"wrongpass123"}`)))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Forwarded-For", "198.51.100.10")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
}
// Use a different username AND IP to verify per-IP lockout isolation.
// (Same username would be locked by per-username lockout — BUG-110 fix.)
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", bytes.NewReader([]byte(`{"username":"other","password":"wrongpass123"}`)))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Forwarded-For", "198.51.100.11")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Fatalf("different forwarded client+user should not inherit another client's lockout, got %d", rr.Code)
}
}
func TestLogin_UsernameLockoutAcrossDifferentIPs(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "lockoutuser", hash, 4)
for i := range 10 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "lockoutuser",
"password": "wrongpassword",
}, fmt.Sprintf("198.51.100.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "lockoutuser",
"password": "wrongpassword",
}, "198.51.100.250")
if rr.Code != http.StatusTooManyRequests {
t.Fatalf("username lockout status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
}
func TestLogin_UsernameLockoutBlocksCorrectPasswordFromFreshIP(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "lockoutcorrect", hash, 4)
for i := range 10 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "lockoutcorrect",
"password": "wrongpassword",
}, fmt.Sprintf("203.0.113.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "lockoutcorrect",
"password": "correctPass1",
}, "203.0.113.250")
if rr.Code != http.StatusTooManyRequests {
t.Fatalf("locked correct-password status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
}
// TestLogin_UsernameLockoutIgnoresCasing locks F1: the per-username lockout key
// must be case-folded so it matches the DB's COLLATE NOCASE username lookup.
// Otherwise an attacker splits the 9-attempt lockout budget across case variants
// of one account (admin, Admin, ADMIN, …), all of which authenticate the same row.
func TestLogin_UsernameLockoutIgnoresCasing(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "casehunt", hash, 4)
// Trip the per-username lockout using the lowercase spelling, from many IPs
// so the per-IP limiter is never the binding cap.
for i := range 10 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "casehunt",
"password": "wrongpassword",
}, fmt.Sprintf("198.51.100.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
// A different casing of the SAME account must land in the same lockout bucket.
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "CASEHUNT",
"password": "wrongpassword",
}, "198.51.100.250")
if rr.Code != http.StatusTooManyRequests {
t.Fatalf("case-variant username bypassed the per-username lockout: status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
}
// TestLogin_ConcurrentBurstCannotExceedUsernameBudget locks F3: the
// per-username failure cap (the only cross-IP brute-force defence) must bind
// concurrent attackers, not just sequential ones. Before the fix the lockout
// was a read-only check followed by a post-bcrypt record, so N concurrent
// requests all passed the stale check and landed N guesses; a distributed
// burst must now land at most the same 10-attempt budget a sequential
// attacker gets.
func TestLogin_ConcurrentBurstCannotExceedUsernameBudget(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "bursttarget", hash, 4)
const burst = 40
start := make(chan struct{})
codes := make([]int, burst)
var wg sync.WaitGroup
for i := range burst {
wg.Go(func() {
raw, _ := json.Marshal(map[string]string{
"username": "bursttarget",
"password": "wrongpassword",
})
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
// Unique IP per request so only the per-username limiter binds.
req.RemoteAddr = fmt.Sprintf("203.0.113.%d:9999", i+1)
rr := httptest.NewRecorder()
<-start
router.ServeHTTP(rr, req)
codes[i] = rr.Code
})
}
close(start)
wg.Wait()
landed, limited := 0, 0
for i, code := range codes {
switch code {
case http.StatusUnauthorized:
landed++
case http.StatusTooManyRequests:
limited++
default:
t.Fatalf("request %d: unexpected status %d", i, code)
}
}
// Sequential budget is 10 landed guesses (9 recorded + the one that trips
// the lockout); a concurrent burst must not exceed it.
if landed > 10 {
t.Fatalf("concurrent burst landed %d password guesses (%d rate-limited), want at most 10", landed, limited)
}
}
// TestLogin_NineFailuresThenCorrectPasswordSucceeds pins the sequential
// accepted-input boundary that the F3 fix must not move: after 9 wrong
// guesses the account owner's correct password still logs in (attempt 10 is
// inside the budget). A fix that reserves attempts pre-compare at the
// original threshold would return 429 here and hand attackers a 9-request
// victim lockout.
func TestLogin_NineFailuresThenCorrectPasswordSucceeds(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "boundaryuser", hash, 4)
for i := range 9 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "boundaryuser",
"password": "wrongpassword",
}, fmt.Sprintf("198.51.100.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "boundaryuser",
"password": "correctPass1",
}, "198.51.100.250")
if rr.Code != http.StatusOK {
t.Fatalf("correct password on attempt 10 status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
}
func TestLogin_SuccessResetsUsernameFailureCounter(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "resetuser", hash, 4)
for i := range 8 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "resetuser",
"password": "wrongpassword",
}, fmt.Sprintf("192.0.2.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "resetuser",
"password": "correctPass1",
}, "192.0.2.200")
if rr.Code != http.StatusOK {
t.Fatalf("success status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
for i := range 3 {
rr = postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "resetuser",
"password": "wrongpassword",
}, fmt.Sprintf("192.0.2.%d", 201+i))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("post-reset attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
}
func TestLogin_GenericErrorOnBadCredentials(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "nobody",
"password": "anypass123",
})
body := rr.Body.String()
// The response must never reveal whether the user exists
if contains(body, "user not found") || contains(body, "does not exist") {
t.Errorf("Login error reveals user existence: %s", body)
}
}
func TestLogin_RequiresTOTPChallenge(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "totpuser", hash, 4)
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "totpuser",
"password": "correctPass1",
})
if rr.Code != http.StatusOK {
t.Fatalf("Login status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil {
t.Fatalf("decode response: %v", err)
}
if resp["requires_2fa"] != true {
t.Fatalf("requires_2fa = %v, want true", resp["requires_2fa"])
}
if resp["partial_token"] == nil || resp["partial_token"] == "" {
t.Fatal("expected partial_token in TOTP challenge response")
}
if token := resp["token"]; token != nil && token != "" {
t.Fatalf("expected no full session token before TOTP verification, got %v", token)
}
}
func TestLogin_UsernameLockoutBlocksTOTPChallenge(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "totplocked", hash, 4)
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
for i := range 10 {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "totplocked",
"password": "wrongpassword",
}, fmt.Sprintf("198.18.0.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "totplocked",
"password": "correctPass1",
}, "198.18.0.250")
if rr.Code != http.StatusTooManyRequests {
t.Fatalf("locked TOTP login status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil {
t.Fatalf("decode response: %v", err)
}
if resp["partial_token"] != nil {
t.Fatalf("partial_token = %v, want nil when username is locked out", resp["partial_token"])
}
}
func TestVerifyTotp_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "totpverify", hash, 4)
secret := "JBSWY3DPEHPK3PXP"
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
login := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "totpverify",
"password": "correctPass1",
})
if login.Code != http.StatusOK {
t.Fatalf("Login status = %d, want 200; body = %s", login.Code, login.Body.String())
}
var loginResp map[string]any
if err := json.NewDecoder(login.Body).Decode(&loginResp); err != nil {
t.Fatalf("decode login response: %v", err)
}
partialToken, _ := loginResp["partial_token"].(string)
if partialToken == "" {
t.Fatal("expected partial_token from login")
}
code, err := auth.GenerateTOTPCode(secret, time.Now().UTC())
if err != nil {
t.Fatalf("GenerateTOTPCode: %v", err)
}
verify := postJSONWithToken(t, router, "/api/v1/auth/verify-totp", partialToken, map[string]string{"code": code})
if verify.Code != http.StatusOK {
t.Fatalf("verify status = %d, want 200; body = %s", verify.Code, verify.Body.String())
}
var verifyResp map[string]any
if err := json.NewDecoder(verify.Body).Decode(&verifyResp); err != nil {
t.Fatalf("decode verify response: %v", err)
}
if verifyResp["token"] == nil || verifyResp["token"] == "" {
t.Fatal("expected full session token after successful TOTP verification")
}
if verifyResp["requires_2fa"] != false {
t.Fatalf("requires_2fa after verify = %v, want false", verifyResp["requires_2fa"])
}
}
func TestEnableConfirmDisableTotp(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "enrolltotp", hash, 4)
token, _ := auth.GenerateToken()
if _, err := database.CreateSession(context.Background(), userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
enable := postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
if enable.Code != http.StatusOK {
t.Fatalf("enable status = %d, want 200; body = %s", enable.Code, enable.Body.String())
}
var enableResp map[string]any
if err := json.NewDecoder(enable.Body).Decode(&enableResp); err != nil {
t.Fatalf("decode enable response: %v", err)
}
qrURI, _ := enableResp["qr_uri"].(string)
if qrURI == "" {
t.Fatal("expected qr_uri from enable response")
}
userBeforeConfirm, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID before confirm: %v", err)
}
if userBeforeConfirm.TOTPSecret != nil {
t.Fatal("TOTP secret should not be persisted before confirmation")
}
parsed, err := url.Parse(qrURI)
if err != nil {
t.Fatalf("parse qr uri: %v", err)
}
secret := parsed.Query().Get("secret")
if secret == "" {
t.Fatal("expected secret query param in qr_uri")
}
code, err := auth.GenerateTOTPCode(secret, time.Now().UTC())
if err != nil {
t.Fatalf("GenerateTOTPCode: %v", err)
}
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{"password": "correctPass1", "code": code})
if confirm.Code != http.StatusNoContent {
t.Fatalf("confirm status = %d, want 204; body = %s", confirm.Code, confirm.Body.String())
}
userAfterConfirm, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after confirm: %v", err)
}
if userAfterConfirm.TOTPSecret == nil || *userAfterConfirm.TOTPSecret == "" {
t.Fatal("TOTP secret should be persisted after confirmation")
}
deleteBody, err := json.Marshal(map[string]string{"password": "correctPass1"})
if err != nil {
t.Fatalf("marshal delete body: %v", err)
}
deleteReq := httptest.NewRequest(http.MethodDelete, "/api/v1/users/me/totp", bytes.NewReader(deleteBody))
deleteReq.Header.Set("Authorization", "Bearer "+token)
deleteReq.Header.Set("Content-Type", "application/json")
deleteReq.RemoteAddr = "127.0.0.1:9999"
deleteRec := httptest.NewRecorder()
router.ServeHTTP(deleteRec, deleteReq)
if deleteRec.Code != http.StatusNoContent {
t.Fatalf("disable status = %d, want 204; body = %s", deleteRec.Code, deleteRec.Body.String())
}
userAfterDelete, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after delete: %v", err)
}
if userAfterDelete.TOTPSecret != nil {
t.Fatal("TOTP secret should be cleared after disable")
}
}
func TestTOTPManagement_RequiresPasswordConfirmation(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "totppassword", hash, 4)
token, _ := auth.GenerateToken()
if _, err := database.CreateSession(context.Background(), userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
enable := postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "wrongPass"})
if enable.Code != http.StatusBadRequest {
t.Fatalf("enable status = %d, want 400; body = %s", enable.Code, enable.Body.String())
}
userAfterEnable, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after failed enable: %v", err)
}
if userAfterEnable.TOTPSecret != nil {
t.Fatal("TOTP secret should remain unset after failed password confirmation")
}
deleteBody, err := json.Marshal(map[string]string{"password": "wrongPass"})
if err != nil {
t.Fatalf("marshal delete body: %v", err)
}
deleteReq := httptest.NewRequest(http.MethodDelete, "/api/v1/users/me/totp", bytes.NewReader(deleteBody))
deleteReq.Header.Set("Authorization", "Bearer "+token)
deleteReq.Header.Set("Content-Type", "application/json")
deleteReq.RemoteAddr = "127.0.0.1:9999"
deleteRec := httptest.NewRecorder()
router.ServeHTTP(deleteRec, deleteReq)
if deleteRec.Code != http.StatusBadRequest {
t.Fatalf("disable status = %d, want 400; body = %s", deleteRec.Code, deleteRec.Body.String())
}
}
func TestVerifyTotp_ConsumesChallengeAfterRepeatedFailures(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser(context.Background(), "totplockout", hash, 4)
secret := "JBSWY3DPEHPK3PXP"
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
login := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "totplockout",
"password": "correctPass1",
})
if login.Code != http.StatusOK {
t.Fatalf("Login status = %d, want 200; body = %s", login.Code, login.Body.String())
}
var loginResp map[string]any
if err := json.NewDecoder(login.Body).Decode(&loginResp); err != nil {
t.Fatalf("decode login response: %v", err)
}
partialToken, _ := loginResp["partial_token"].(string)
if partialToken == "" {
t.Fatal("expected partial_token from login")
}
for i := range 5 {
verify := postJSONWithToken(t, router, "/api/v1/auth/verify-totp", partialToken, map[string]string{"code": "000000"})
if verify.Code != http.StatusUnauthorized {
t.Fatalf("attempt %d status = %d, want 401; body = %s", i+1, verify.Code, verify.Body.String())
}
}
code, err := auth.GenerateTOTPCode(secret, time.Now().UTC())
if err != nil {
t.Fatalf("GenerateTOTPCode: %v", err)
}
verify := postJSONWithToken(t, router, "/api/v1/auth/verify-totp", partialToken, map[string]string{"code": code})
if verify.Code != http.StatusUnauthorized {
t.Fatalf("verify after lockout status = %d, want 401; body = %s", verify.Code, verify.Body.String())
}
}
func TestLogin_Require2FASettingRejectsUsersWithoutEnrollment(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
t.Fatalf("enable require_2fa: %v", err)
}
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'false' WHERE key = 'registration_open'`); err != nil {
t.Fatalf("disable registration_open: %v", err)
}
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "needsenrollment", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "needsenrollment",
"password": "correctPass1",
})
if rr.Code != http.StatusForbidden {
t.Fatalf("Login status = %d, want 403; body = %s", rr.Code, rr.Body.String())
}
}
func TestLogin_BannedUser(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
id, _ := database.CreateUser(context.Background(), "banned", hash, 4)
_ = database.BanUser(context.Background(), id, "violated rules", nil)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "banned",
"password": "correctPass1",
})
if rr.Code != http.StatusForbidden {
t.Errorf("Login banned user status = %d, want 403", rr.Code)
}
}
func TestLogin_MissingFields(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{})
if rr.Code != http.StatusBadRequest {
t.Errorf("Login missing fields status = %d, want 400", rr.Code)
}
}
// ─── Logout tests ─────────────────────────────────────────────────────────────
func TestLogout_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "logoutuser", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
rr := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil)
if rr.Code != http.StatusNoContent {
t.Errorf("Logout status = %d, want 204", rr.Code)
}
// Session should be gone.
sess, _ := database.GetSessionByTokenHash(context.Background(), tokenHash)
if sess != nil {
t.Error("Session still exists after logout")
}
}
func TestLogout_NoAuth(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/logout", nil)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("Logout no auth status = %d, want 401", rr.Code)
}
}
// ─── Me tests ─────────────────────────────────────────────────────────────────
func TestMe_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "meuser", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := getWithToken(t, router, "/api/v1/auth/me", token)
if rr.Code != http.StatusOK {
t.Errorf("Me status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(rr.Body).Decode(&resp)
if resp["id"] == nil {
t.Error("Me response missing id")
}
if resp["username"] != "meuser" {
t.Errorf("Me username = %v, want meuser", resp["username"])
}
}
func TestMe_NoAuth(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/me", nil)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("Me no auth status = %d, want 401", rr.Code)
}
}
// ─── Fix 2.5: Password trim fix ───────────────────────────────────────────────
// TestLogin_PasswordWithLeadingSpaceIsPreserved verifies that a password with
// leading whitespace is NOT trimmed, so a user who set " securePass1" can log
// in with " securePass1" and NOT with "securePass1".
func TestLogin_PasswordWithLeadingSpaceIsPreserved(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
// Hash the password WITH the leading space — this is what was registered.
hash, _ := auth.HashPassword(" securePass1")
_, _ = database.CreateUser(context.Background(), "spacepassuser", hash, 4)
// Login with the exact same password (including space) must succeed.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "spacepassuser",
"password": " securePass1",
})
if rr.Code != http.StatusOK {
t.Errorf("Login space-prefixed password status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
}
// TestLogin_PasswordWithLeadingSpaceTrimmedFails verifies that logging in with
// the trimmed version of a space-prefixed password correctly fails.
func TestLogin_PasswordWithLeadingSpaceTrimmedFails(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
// Register with password that has a leading space.
hash, _ := auth.HashPassword(" securePass1")
_, _ = database.CreateUser(context.Background(), "spacepassuser2", hash, 4)
// Login without the leading space must fail.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "spacepassuser2",
"password": "securePass1",
})
if rr.Code != http.StatusUnauthorized {
t.Errorf("Login trimmed space password status = %d, want 401; body = %s", rr.Code, rr.Body.String())
}
}
// TestLogin_PasswordWithTrailingSpaceIsPreserved verifies that a password with
// trailing whitespace is NOT trimmed.
func TestLogin_PasswordWithTrailingSpaceIsPreserved(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("securePass1 ")
_, _ = database.CreateUser(context.Background(), "trailingspaceuser", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "trailingspaceuser",
"password": "securePass1 ",
})
if rr.Code != http.StatusOK {
t.Errorf("Login trailing-space password status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
}
// TestLogin_UsernameIsStillTrimmed verifies that the username IS still trimmed
// (only the password trim was removed).
func TestLogin_UsernameIsStillTrimmed(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "trimuser", hash, 4)
// Username with surrounding spaces should resolve to "trimuser".
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": " trimuser ",
"password": "correctPass1",
})
if rr.Code != http.StatusOK {
t.Errorf("Login space-padded username status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
}
// TestRegister_UsernameNotHTMLEscaped pins OC-0099: handleRegister must not
// persist an HTML-escaped username. A bare bluemonday sanitizer.Sanitize call
// HTML-escapes survivors (' -> &#39;, & -> &amp;, " -> &#34;), so a name like
// "O'Brien" would be stored as "O&#39;Brien" — different from what the user
// typed and from what handleLogin looks up (which only trims).
func TestRegister_UsernameNotHTMLEscaped(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "owner2", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "O'Brien",
"password": "securePass1",
"invite_code": code,
})
if rr.Code != http.StatusCreated {
t.Fatalf("Register status = %d, want 201; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(rr.Body).Decode(&resp)
user, _ := resp["user"].(map[string]any)
if got, want := user["username"], "O'Brien"; got != want {
t.Errorf("registered username = %q, want %q (must not be HTML-escaped)", got, want)
}
stored, err := database.GetUserByUsername(context.Background(), "O'Brien")
if err != nil || stored == nil {
t.Fatalf("GetUserByUsername(%q) = (%v, %v), want a match", "O'Brien", stored, err)
}
}
// TestLogin_UsernameWithApostropheSucceeds pins OC-0099 end-to-end: a user
// who registers with an apostrophe/quote/ampersand in their name must be able
// to log back in with the exact same name. Before the fix, handleRegister
// stored the HTML-escaped form while handleLogin looked up the raw form, so
// this login permanently 401s for any such account.
func TestLogin_UsernameWithApostropheSucceeds(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "owner3", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
regRR := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "O'Brien",
"password": "securePass1",
"invite_code": code,
})
if regRR.Code != http.StatusCreated {
t.Fatalf("Register status = %d, want 201; body = %s", regRR.Code, regRR.Body.String())
}
loginRR := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "O'Brien",
"password": "securePass1",
})
if loginRR.Code != http.StatusOK {
t.Fatalf("Login with registered username status = %d, want 200; body = %s", loginRR.Code, loginRR.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(loginRR.Body).Decode(&resp)
if resp["token"] == nil {
t.Error("Login response missing token")
}
}
// TestLogin_OversizedUsernameRejectedBeforeRateLimiterKey pins OC-0021:
// handleLogin never length-checks req.Username before using it to build
// RateLimiter map keys ("login_user_fail:"+username, "login_user_lock:"+...).
// An unauthenticated caller could otherwise pin an arbitrarily large key in
// the limiter's maps (retained for hours by Cleanup's window) on every
// attempt. The oversized username must be rejected with 400 before any such
// key is ever recorded, mirroring the same 32-rune cap register enforces via
// auth.ValidateUsername before touching anything.
func TestLogin_OversizedUsernameRejectedBeforeRateLimiterKey(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hugeUsername := strings.Repeat("a", 1<<20) // 1 MiB, as in the repro
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": hugeUsername,
"password": "anypass123",
})
if rr.Code != http.StatusBadRequest {
t.Errorf("Login oversized username status = %d, want 400; body = %s", rr.Code, rr.Body.String())
}
// The route's own per-IP RateLimitMiddleware ("login:"+ip) legitimately
// records one small, IP-bounded window entry regardless of this fix.
// What must NOT happen is handleLogin additionally recording a second
// entry keyed on the 1 MiB username itself (failKey/userFailKey) — that
// unbounded entry is the actual leak, so anything beyond the single
// expected IP entry means the oversized username reached key-building
// code before being rejected.
if windows, lockouts := limiter.Len(); windows > 1 || lockouts != 0 {
t.Errorf("RateLimiter retained state after oversized-username login: windows=%d lockouts=%d, want at most 1/0", windows, lockouts)
}
}
// OC-0151: registerReadRequest ran the fixpoint sanitizer
// (service.SanitizeText) over the raw username *before* auth.ValidateUsername
// applies its 32-rune cap. The sanitizer loops sanitizePass to a fixpoint,
// and nested HTML entities force roughly one extra pass per two nesting
// levels, so the cost is quadratic in the attacker-controlled field length.
// A 16 KB adversarial username measurably takes ~200ms to sanitize on this
// tree (measured up to ~3.4s at 64 KB) — all of it spent before any bound on
// the field is applied, and unauthenticated. The fix must reject an
// oversized username on a cheap byte-length check *before* sanitizing, so
// the rejection is near-instant regardless of payload size.
func TestRegister_OversizedUsernameRejectedBeforeSanitizing(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
// Adversarial nested-entity payload (16 KB) — see service.sanitizeToFixpoint's
// doc comment for why this shape is quadratic to sanitize.
hugeUsername := "&" + strings.Repeat("amp;", 4000) + "lt;"
start := time.Now()
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": hugeUsername,
"password": "securePass1",
"invite_code": "whatever",
})
elapsed := time.Since(start)
if rr.Code != http.StatusBadRequest {
t.Errorf("Register oversized username status = %d, want 400; body = %s", rr.Code, rr.Body.String())
}
// A guard that runs before sanitizing rejects in well under a
// millisecond; the pre-fix code spends ~200ms in sanitizeToFixpoint on
// this payload before it ever reaches auth.ValidateUsername's length
// check. 150ms gives generous margin over noise while still being far
// below the unguarded cost.
if elapsed > 150*time.Millisecond {
t.Errorf("Register oversized username took %v, want well under 150ms (raw field must be bounded before sanitizing, not after)", elapsed)
}
}
// ─── Rate limiting integration test ──────────────────────────────────────────
func TestRegister_RateLimit(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser(context.Background(), "rl_owner", "hash", 1)
// Attempt register 4 times (limit=3) — 4th should be rate-limited.
var lastCode int
for i := range 4 {
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "rl_user" + string(rune('0'+i)),
"password": "securePass1",
"invite_code": code,
})
lastCode = rr.Code
}
if lastCode != http.StatusTooManyRequests {
t.Errorf("Register rate limit: last attempt status = %d, want 429", lastCode)
}
}
// ─── DeleteAccount tests ─────────────────────────────────────────────────────
// deleteJSONWithToken sends a DELETE request with an Authorization header and JSON body.
func deleteJSONWithToken(t *testing.T, router http.Handler, path, token string, body any) *httptest.ResponseRecorder {
t.Helper()
raw, _ := json.Marshal(body)
req := httptest.NewRequest(http.MethodDelete, path, bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+token)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
return rr
}
func TestDeleteAccount_Success(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
// Create as Member (role_id=4) so the last-admin check does not block deletion.
uid, _ := database.CreateUser(context.Background(), "deleteuser", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "correctPass1",
})
if rr.Code != http.StatusNoContent {
t.Fatalf("DeleteAccount status = %d, want 204; body = %s", rr.Code, rr.Body.String())
}
// User should be anonymised (banned, username changed).
user, err := database.GetUserByID(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserByID after delete: %v", err)
}
if user == nil {
t.Fatal("user row should still exist (soft-delete), got nil")
}
if !user.Banned {
t.Error("expected user to be banned after deletion")
}
if user.Username != "[deleted-1]" && user.Username != "[deleted-"+fmt.Sprintf("%d", uid)+"]" {
t.Errorf("expected anonymised username, got %q", user.Username)
}
// Session should be gone.
sess, _ := database.GetSessionByTokenHash(context.Background(), tokenHash)
if sess != nil {
t.Error("session should be deleted after account deletion")
}
}
func TestDeleteAccount_MissingPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "delnopass", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{})
if rr.Code != http.StatusBadRequest {
t.Errorf("DeleteAccount missing password status = %d, want 400; body = %s", rr.Code, rr.Body.String())
}
}
func TestDeleteAccount_WrongPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "delwrong", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "wrongPassword1",
})
if rr.Code != http.StatusBadRequest {
t.Errorf("DeleteAccount wrong password status = %d, want 400; body = %s", rr.Code, rr.Body.String())
}
// Verify user is NOT deleted.
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.Banned {
t.Error("user should not be deleted after wrong password")
}
}
func TestDeleteAccount_LastAdmin(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
// Create as Owner (role_id=1) — the only admin-class user.
uid, _ := database.CreateUser(context.Background(), "lastadmin", hash, 1)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "correctPass1",
})
if rr.Code != http.StatusForbidden {
t.Errorf("DeleteAccount last admin status = %d, want 403; body = %s", rr.Code, rr.Body.String())
}
// User should still be intact.
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.Banned {
t.Error("last admin should not be deleted")
}
}
func TestDeleteAccount_NoAuth(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
req := httptest.NewRequest(http.MethodDelete, "/api/v1/auth/account", bytes.NewReader([]byte(`{"password":"x"}`)))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("DeleteAccount no auth status = %d, want 401", rr.Code)
}
}
func TestDeleteAccount_LockoutAfterRepeatedFailures(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "dellockout", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// 3 failures should trigger lockout on the 4th attempt.
for range 4 {
deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "wrongPassword1",
})
}
// Even with correct password, should now be locked out.
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "correctPass1",
})
if rr.Code != http.StatusTooManyRequests {
t.Errorf("DeleteAccount lockout status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
}
// ─── ConfirmTOTP additional tests ────────────────────────────────────────────
func TestConfirmTOTP_InvalidCode(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "totpbadcode", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first to get a pending secret.
enable := postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
if enable.Code != http.StatusOK {
t.Fatalf("enable status = %d, want 200; body = %s", enable.Code, enable.Body.String())
}
// Confirm with an invalid code.
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{
"password": "correctPass1",
"code": "000000",
})
if confirm.Code != http.StatusUnauthorized {
t.Errorf("ConfirmTOTP invalid code status = %d, want 401; body = %s", confirm.Code, confirm.Body.String())
}
// Secret should NOT be persisted.
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret != nil {
t.Error("TOTP secret should not be persisted after invalid code")
}
}
func TestConfirmTOTP_NoPendingSecret(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "totpnopending", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Confirm without enabling first — no pending secret.
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{
"password": "correctPass1",
"code": "123456",
})
if confirm.Code != http.StatusBadRequest {
t.Errorf("ConfirmTOTP no pending status = %d, want 400; body = %s", confirm.Code, confirm.Body.String())
}
}
func TestConfirmTOTP_MissingPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "totpnoconfirmpass", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first.
postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
// Confirm without password.
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{
"code": "123456",
})
if confirm.Code != http.StatusBadRequest {
t.Errorf("ConfirmTOTP missing password status = %d, want 400; body = %s", confirm.Code, confirm.Body.String())
}
}
func TestConfirmTOTP_WrongPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "totpwrongconfirm", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first.
postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
// Confirm with wrong password.
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{
"password": "wrongPass",
"code": "123456",
})
if confirm.Code != http.StatusBadRequest {
t.Errorf("ConfirmTOTP wrong password status = %d, want 400; body = %s", confirm.Code, confirm.Body.String())
}
}
func TestConfirmTOTP_NoAuth(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
req := httptest.NewRequest(http.MethodPost, "/api/v1/users/me/totp/confirm", bytes.NewReader([]byte(`{"password":"x","code":"123456"}`)))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("ConfirmTOTP no auth status = %d, want 401", rr.Code)
}
}
// ─── DisableTOTP additional tests ────────────────────────────────────────────
func TestDisableTOTP_WrongPassword(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "disabletotpwrong", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Set TOTP secret directly.
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
t.Fatalf("set totp secret: %v", err)
}
deleteBody, _ := json.Marshal(map[string]string{"password": "wrongPass"})
req := httptest.NewRequest(http.MethodDelete, "/api/v1/users/me/totp", bytes.NewReader(deleteBody))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusBadRequest {
t.Errorf("DisableTOTP wrong password status = %d, want 400; body = %s", rr.Code, rr.Body.String())
}
// TOTP should still be enabled.
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret == nil {
t.Error("TOTP secret should still be set after wrong password")
}
}
func TestDisableTOTP_Require2FABlocksDisable(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
// Enable require_2fa setting.
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
t.Fatalf("enable require_2fa: %v", err)
}
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "disabletotpreq", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Set TOTP secret directly.
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
t.Fatalf("set totp secret: %v", err)
}
deleteBody, _ := json.Marshal(map[string]string{"password": "correctPass1"})
req := httptest.NewRequest(http.MethodDelete, "/api/v1/users/me/totp", bytes.NewReader(deleteBody))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusForbidden {
t.Errorf("DisableTOTP require_2fa status = %d, want 403; body = %s", rr.Code, rr.Body.String())
}
// TOTP should still be enabled.
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret == nil {
t.Error("TOTP secret should still be set when require_2fa is enabled")
}
}
func TestDisableTOTP_NoAuth(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
req := httptest.NewRequest(http.MethodDelete, "/api/v1/users/me/totp", bytes.NewReader([]byte(`{"password":"x"}`)))
req.Header.Set("Content-Type", "application/json")
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("DisableTOTP no auth status = %d, want 401", rr.Code)
}
}
// ─── Logout additional tests ─────────────────────────────────────────────────
func TestLogout_InvalidToken(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSONWithToken(t, router, "/api/v1/auth/logout", "invalid-token-value", nil)
if rr.Code != http.StatusUnauthorized {
t.Errorf("Logout invalid token status = %d, want 401", rr.Code)
}
}
func TestLogout_SessionGoneAfterLogout(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "logoutsess", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
// First logout should succeed.
rr := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil)
if rr.Code != http.StatusNoContent {
t.Fatalf("first logout status = %d, want 204", rr.Code)
}
// Second logout with the same token should fail (session already deleted).
rr2 := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil)
if rr2.Code != http.StatusUnauthorized {
t.Errorf("second logout status = %d, want 401", rr2.Code)
}
}
// ─── Me additional tests ─────────────────────────────────────────────────────
func TestMe_ReturnsCorrectUserFields(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser(context.Background(), "medetailed", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := getWithToken(t, router, "/api/v1/auth/me", token)
if rr.Code != http.StatusOK {
t.Fatalf("Me status = %d, want 200; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil {
t.Fatalf("decode response: %v", err)
}
// Verify all expected fields are present.
for _, field := range []string{"id", "username", "status", "role_id", "totp_enabled", "created_at"} {
if _, ok := resp[field]; !ok {
t.Errorf("Me response missing field %q", field)
}
}
if resp["username"] != "medetailed" {
t.Errorf("username = %v, want medetailed", resp["username"])
}
if resp["totp_enabled"] != false {
t.Errorf("totp_enabled = %v, want false", resp["totp_enabled"])
}
}
func TestMe_InvalidToken(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := getWithToken(t, router, "/api/v1/auth/me", "not-a-real-token")
if rr.Code != http.StatusUnauthorized {
t.Errorf("Me invalid token status = %d, want 401", rr.Code)
}
}
// ─── Helpers ─────────────────────────────────────────────────────────────────
func contains(s, sub string) bool {
return len(s) >= len(sub) && (s == sub || len(s) > 0 && containsStr(s, sub))
}
func containsStr(s, sub string) bool {
for i := 0; i <= len(s)-len(sub); i++ {
if s[i:i+len(sub)] == sub {
return true
}
}
return false
}
// expiredInviteDB creates a DB with an already-expired invite.
func expiredInviteDB(t *testing.T) (*db.DB, string) {
t.Helper()
database := newAuthTestDB(t)
ownerID, _ := database.CreateUser(context.Background(), "expowner", "hash", 1)
past := time.Now().Add(-time.Hour)
code, _ := database.CreateInvite(context.Background(), ownerID, 0, &past)
return database, code
}
func TestRegister_ExpiredInvite(t *testing.T) {
database, code := expiredInviteDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
"password": "securePass1",
"invite_code": code,
})
if rr.Code != http.StatusBadRequest {
t.Errorf("Register expired invite status = %d, want 400", rr.Code)
}
}
// TestRegister_UsesTrustedForwardedIP pins OC-0093: handleRegister must
// resolve the client IP through the same trusted-proxy list handleLogin
// uses, not unconditionally use RemoteAddr. Behind a trusted reverse proxy,
// the sessions.ip row registration creates must record the real client, not
// the proxy's own address — otherwise the same client shows two different
// IPs on the "active sessions" screen depending on whether they registered
// or logged in.
func TestRegister_UsesTrustedForwardedIP(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouterWithProxies(database, limiter, []string{"127.0.0.0/8"})
ownerID, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/register", bytes.NewReader([]byte(
`{"username":"newuser","password":"securePass1","invite_code":"`+code+`"}`)))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Forwarded-For", "203.0.113.9")
req.RemoteAddr = "127.0.0.1:9999" // the trusted reverse proxy's own hop
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusCreated {
t.Fatalf("Register status = %d, want 201; body = %s", rr.Code, rr.Body.String())
}
var resp map[string]any
_ = json.NewDecoder(rr.Body).Decode(&resp)
token, _ := resp["token"].(string)
if token == "" {
t.Fatal("Register response missing token")
}
sess, err := database.GetSessionByTokenHash(context.Background(), auth.HashToken(token))
if err != nil || sess == nil {
t.Fatalf("GetSessionByTokenHash: %v", err)
}
if sess.IP != "203.0.113.9" {
t.Errorf("session IP = %q, want the trusted-forwarded client IP %q — registration behind a reverse proxy must not record the proxy's own address", sess.IP, "203.0.113.9")
}
}