package auth import ( "context" "database/sql" "errors" "net/http" "net/http/httptest" "testing" "time" "github.com/golang-jwt/jwt/v5" _ "modernc.org/sqlite" ) func TestJWT_Generate(t *testing.T) { mgr := NewJWTManager("test-secret", 24) token, expiresAt, err := mgr.Generate(1, "testuser", "admin") if err != nil { t.Fatalf("Generate() error = %v", err) } if token == "" { t.Fatal("Generate() returned empty token") } if expiresAt.Before(time.Now()) { t.Fatal("Generate() returned past expiration time") } if expiresAt.Before(time.Now().Add(23 * time.Hour)) { t.Fatal("Generate() expiration time is too early") } claims, err := mgr.Validate(token) if err != nil { t.Fatalf("Validate() error = %v", err) } if claims.UserID != 1 { t.Errorf("claims.UserID = %d, want 1", claims.UserID) } if claims.Username != "testuser" { t.Errorf("claims.Username = %s, want testuser", claims.Username) } if claims.Role != "admin" { t.Errorf("claims.Role = %s, want admin", claims.Role) } } func TestJWT_Validate_ValidToken(t *testing.T) { mgr := NewJWTManager("test-secret", 24) token, _, err := mgr.Generate(42, "alice", "user") if err != nil { t.Fatalf("Generate() error = %v", err) } claims, err := mgr.Validate(token) if err != nil { t.Fatalf("Validate() error = %v", err) } if claims.UserID != 42 { t.Errorf("claims.UserID = %d, want 42", claims.UserID) } if claims.Username != "alice" { t.Errorf("claims.Username = %s, want alice", claims.Username) } if claims.Role != "user" { t.Errorf("claims.Role = %s, want user", claims.Role) } } func TestJWT_Validate_ExpiredToken(t *testing.T) { mgr := NewJWTManager("test-secret", 0) token := jwtWithExpiry(time.Now().Add(-1 * time.Hour)) _, err := mgr.Validate(token) if !errors.Is(err, ErrExpiredToken) { t.Errorf("Validate() error = %v, want ErrExpiredToken", err) } } func TestJWT_Validate_InvalidToken(t *testing.T) { mgr := NewJWTManager("test-secret", 24) tests := []struct { name string token string }{ {"malformed token", "not.a.token"}, {"empty token", ""}, {"random string", "abcdef123456"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, err := mgr.Validate(tt.token) if err == nil { t.Error("Validate() expected error for invalid token") } }) } } func TestJWT_Validate_WrongSecret(t *testing.T) { mgr1 := NewJWTManager("secret-one", 24) mgr2 := NewJWTManager("secret-two", 24) token, _, err := mgr1.Generate(1, "user", "admin") if err != nil { t.Fatalf("Generate() error = %v", err) } _, err = mgr2.Validate(token) if err == nil { t.Error("Validate() expected error for token signed with different secret") } } func TestPassword_HashPassword_RandomSalts(t *testing.T) { password := "samepassword123" hash1, err := HashPassword(password) if err != nil { t.Fatalf("HashPassword() error = %v", err) } hash2, err := HashPassword(password) if err != nil { t.Fatalf("HashPassword() error = %v", err) } if string(hash1) == string(hash2) { t.Error("HashPassword() produced identical hashes for same password") } } func TestPassword_CheckPassword_Correct(t *testing.T) { password := "mysecretpassword" hash, err := HashPassword(password) if err != nil { t.Fatalf("HashPassword() error = %v", err) } if !VerifyPassword(hash, password) { t.Error("VerifyPassword() returned false for correct password") } } func TestPassword_CheckPassword_Wrong(t *testing.T) { password := "mysecretpassword" wrongPassword := "wrongpassword" hash, err := HashPassword(password) if err != nil { t.Fatalf("HashPassword() error = %v", err) } if VerifyPassword(hash, wrongPassword) { t.Error("VerifyPassword() returned true for wrong password") } } func TestMiddleware_RequireAuth_NoToken(t *testing.T) { InitJWTManager("test-secret", 24) handler := RequireAuth(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Error("next handler should not be called") })) req := httptest.NewRequest(http.MethodGet, "/", nil) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusUnauthorized { t.Errorf("RequireAuth() status = %d, want %d", rr.Code, http.StatusUnauthorized) } } func TestMiddleware_RequireAuth_InvalidToken(t *testing.T) { InitJWTManager("test-secret", 24) handler := RequireAuth(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Error("next handler should not be called") })) req := httptest.NewRequest(http.MethodGet, "/", nil) req.AddCookie(&http.Cookie{Name: CookieName, Value: "invalid-token"}) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusUnauthorized { t.Errorf("RequireAuth() status = %d, want %d", rr.Code, http.StatusUnauthorized) } } func TestMiddleware_RequireAuth_ValidToken(t *testing.T) { secret := "test-secret" InitJWTManager(secret, 24) mgr := GetJWTManager() token, _, err := mgr.Generate(99, "testuser", "admin") if err != nil { t.Fatalf("Generate() error = %v", err) } var capturedClaims *Claims handler := RequireAuth(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { capturedClaims = GetClaims(r.Context()) w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodGet, "/", nil) req.AddCookie(&http.Cookie{Name: CookieName, Value: token}) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusOK { t.Errorf("RequireAuth() status = %d, want %d", rr.Code, http.StatusOK) } if capturedClaims == nil { t.Fatal("RequireAuth() did not set claims in context") } if capturedClaims.UserID != 99 { t.Errorf("capturedClaims.UserID = %d, want 99", capturedClaims.UserID) } if capturedClaims.Username != "testuser" { t.Errorf("capturedClaims.Username = %s, want testuser", capturedClaims.Username) } if capturedClaims.Role != "admin" { t.Errorf("capturedClaims.Role = %s, want admin", capturedClaims.Role) } } func TestMiddleware_RequireAdmin_NonAdmin(t *testing.T) { InitJWTManager("test-secret", 24) handler := RequireAdmin(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Error("next handler should not be called for non-admin") })) req := httptest.NewRequest(http.MethodGet, "/", nil) ctx := context.WithValue(req.Context(), ClaimsCtxKey, &Claims{Role: "user"}) req = req.WithContext(ctx) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusForbidden { t.Errorf("RequireAdmin() status = %d, want %d", rr.Code, http.StatusForbidden) } } func TestMiddleware_RequireAdmin_NoClaims(t *testing.T) { InitJWTManager("test-secret", 24) handler := RequireAdmin(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Error("next handler should not be called when no claims in context") })) req := httptest.NewRequest(http.MethodGet, "/", nil) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusForbidden { t.Errorf("RequireAdmin() status = %d, want %d", rr.Code, http.StatusForbidden) } } func TestMiddleware_RequireAdmin_Admin(t *testing.T) { InitJWTManager("test-secret", 24) called := false handler := RequireAdmin(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { called = true w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodGet, "/", nil) ctx := context.WithValue(req.Context(), ClaimsCtxKey, &Claims{Role: "admin"}) req = req.WithContext(ctx) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) if rr.Code != http.StatusOK { t.Errorf("RequireAdmin() status = %d, want %d", rr.Code, http.StatusOK) } if !called { t.Error("RequireAdmin() did not call next handler for admin role") } } func TestSeed_SeedsAdminUser(t *testing.T) { db, err := sql.Open("sqlite", ":memory:") if err != nil { t.Fatalf("sql.Open() error = %v", err) } defer db.Close() _, err = db.Exec(`CREATE TABLE users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT UNIQUE NOT NULL, password_hash TEXT NOT NULL, role TEXT NOT NULL DEFAULT 'user', created_at DATETIME DEFAULT CURRENT_TIMESTAMP )`) if err != nil { t.Fatalf("CREATE TABLE error = %v", err) } err = SeedAdmin(db, "admin", "secretpassword123") if err != nil { t.Fatalf("SeedAdmin() error = %v", err) } var username, role string err = db.QueryRow("SELECT username, role FROM users WHERE username = 'admin'").Scan(&username, &role) if err != nil { t.Fatalf("QueryRow() error = %v", err) } if username != "admin" { t.Errorf("username = %s, want admin", username) } if role != "admin" { t.Errorf("role = %s, want admin", role) } } func TestSeed_Idempotent(t *testing.T) { db, err := sql.Open("sqlite", ":memory:") if err != nil { t.Fatalf("sql.Open() error = %v", err) } defer db.Close() _, err = db.Exec(`CREATE TABLE users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT UNIQUE NOT NULL, password_hash TEXT NOT NULL, role TEXT NOT NULL DEFAULT 'user', created_at DATETIME DEFAULT CURRENT_TIMESTAMP )`) if err != nil { t.Fatalf("CREATE TABLE error = %v", err) } _, err = db.Exec("INSERT INTO users (username, password_hash, role) VALUES ('admin', 'existing-hash', 'admin')") if err != nil { t.Fatalf("INSERT error = %v", err) } err = SeedAdmin(db, "admin", "newpassword") if err != nil { t.Fatalf("SeedAdmin() error = %v", err) } var count int err = db.QueryRow("SELECT COUNT(*) FROM users WHERE username = 'admin'").Scan(&count) if err != nil { t.Fatalf("QueryRow() error = %v", err) } if count != 1 { t.Errorf("user count = %d, want 1 (idempotent behavior)", count) } } func TestSeed_EmptyPasswordError(t *testing.T) { db, err := sql.Open("sqlite", ":memory:") if err != nil { t.Fatalf("sql.Open() error = %v", err) } defer db.Close() _, err = db.Exec(`CREATE TABLE users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT UNIQUE NOT NULL, password_hash TEXT NOT NULL, role TEXT NOT NULL DEFAULT 'user', created_at DATETIME DEFAULT CURRENT_TIMESTAMP )`) if err != nil { t.Fatalf("CREATE TABLE error = %v", err) } err = SeedAdmin(db, "admin", "") if err == nil { t.Error("SeedAdmin() expected error for empty password on first run") } } func jwtWithExpiry(expiry time.Time) string { claims := &Claims{ UserID: 1, Username: "testuser", Role: "admin", RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(expiry), IssuedAt: jwt.NewNumericDate(time.Now().Add(-2 * time.Hour)), NotBefore: jwt.NewNumericDate(time.Now().Add(-2 * time.Hour)), }, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) signed, _ := token.SignedString([]byte("test-secret")) return signed }