mirror of
https://github.com/R0m1k3/FlowReader.git
synced 2026-10-11 17:28:05 +02:00
feat(story-1.4): add session-based authentication with login/logout
This commit is contained in:
1 parent
e186a4dd87
commit
d67e169a43
8 files changed
+446
-13
No files matched your search
@@ -38,8 +38,8 @@ development_status:
|
||||
epic-1: in-progress
|
||||
1-1-project-initialization-walking-skeleton: done
|
||||
1-2-database-migration-system: done
|
||||
1-3-user-registration-api: review
|
||||
1-4-session-authentication: backlog
|
||||
1-3-user-registration-api: done
|
||||
1-4-session-authentication: review
|
||||
1-5-frontend-auth-foundation: backlog
|
||||
epic-1-retrospective: optional
|
||||
|
||||
|
||||
+10
-2
@@ -37,9 +37,10 @@ func main() {
|
||||
|
||||
// Initialize repositories
|
||||
userRepo := repository.NewUserRepository(pool)
|
||||
sessionRepo := repository.NewSessionRepository(pool)
|
||||
|
||||
// Initialize services
|
||||
authService := service.NewAuthService(userRepo)
|
||||
authService := service.NewAuthService(userRepo, sessionRepo)
|
||||
|
||||
// Initialize handlers
|
||||
authHandler := handler.NewAuthHandler(authService)
|
||||
@@ -72,9 +73,16 @@ func main() {
|
||||
w.Write([]byte(`{"message":"FlowReader API v1"}`))
|
||||
})
|
||||
|
||||
// Auth routes
|
||||
// Auth routes (public)
|
||||
r.Route("/auth", func(r chi.Router) {
|
||||
r.Post("/register", authHandler.Register)
|
||||
r.Post("/login", authHandler.Login)
|
||||
r.Post("/logout", authHandler.Logout)
|
||||
})
|
||||
|
||||
// User routes
|
||||
r.Route("/users", func(r chi.Router) {
|
||||
r.Get("/me", authHandler.Me)
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Session represents an active user session.
|
||||
type Session struct {
|
||||
ID uuid.UUID `json:"id"`
|
||||
UserID uuid.UUID `json:"user_id"`
|
||||
Token string `json:"-"` // Never expose in JSON
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UserAgent string `json:"user_agent,omitempty"`
|
||||
IPAddress string `json:"ip_address,omitempty"`
|
||||
}
|
||||
|
||||
// SessionRepository defines the interface for session data access.
|
||||
type SessionRepository interface {
|
||||
Create(session *Session) error
|
||||
GetByToken(token string) (*Session, error)
|
||||
Delete(token string) error
|
||||
DeleteByUserID(userID uuid.UUID) error
|
||||
DeleteExpired() (int64, error)
|
||||
}
|
||||
+106
-1
@@ -1,10 +1,10 @@
|
||||
// Package handler contains HTTP handlers for the REST API.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/michael/flowreader/internal/service"
|
||||
)
|
||||
@@ -45,6 +45,111 @@ func (h *AuthHandler) Register(w http.ResponseWriter, r *http.Request) {
|
||||
respondJSON(w, http.StatusCreated, resp)
|
||||
}
|
||||
|
||||
// Login handles POST /api/v1/auth/login
|
||||
func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) {
|
||||
var req service.LoginRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
respondError(w, http.StatusBadRequest, "Invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
// Add request metadata
|
||||
req.UserAgent = r.UserAgent()
|
||||
req.IPAddress = getClientIP(r)
|
||||
|
||||
resp, err := h.authService.Login(req)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrInvalidCredentials) {
|
||||
respondError(w, http.StatusUnauthorized, "Invalid email or password")
|
||||
} else {
|
||||
respondError(w, http.StatusInternalServerError, "Login failed")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Set session cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "session_id",
|
||||
Value: resp.Token,
|
||||
Path: "/",
|
||||
Expires: resp.ExpiresAt,
|
||||
HttpOnly: true,
|
||||
Secure: r.TLS != nil, // Secure only if HTTPS
|
||||
SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
|
||||
respondJSON(w, http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Logout handles POST /api/v1/auth/logout
|
||||
func (h *AuthHandler) Logout(w http.ResponseWriter, r *http.Request) {
|
||||
cookie, err := r.Cookie("session_id")
|
||||
if err != nil {
|
||||
respondJSON(w, http.StatusOK, map[string]string{"message": "Already logged out"})
|
||||
return
|
||||
}
|
||||
|
||||
_ = h.authService.Logout(cookie.Value)
|
||||
|
||||
// Clear the cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "session_id",
|
||||
Value: "",
|
||||
Path: "/",
|
||||
MaxAge: -1,
|
||||
HttpOnly: true,
|
||||
Secure: r.TLS != nil,
|
||||
SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
|
||||
respondJSON(w, http.StatusOK, map[string]string{"message": "Logged out successfully"})
|
||||
}
|
||||
|
||||
// Me handles GET /api/v1/users/me
|
||||
func (h *AuthHandler) Me(w http.ResponseWriter, r *http.Request) {
|
||||
cookie, err := r.Cookie("session_id")
|
||||
if err != nil {
|
||||
respondError(w, http.StatusUnauthorized, "Not authenticated")
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.authService.GetUserByToken(cookie.Value)
|
||||
if err != nil {
|
||||
respondError(w, http.StatusInternalServerError, "Failed to get user")
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
respondError(w, http.StatusUnauthorized, "Session expired")
|
||||
return
|
||||
}
|
||||
|
||||
respondJSON(w, http.StatusOK, service.UserInfo{
|
||||
ID: user.ID,
|
||||
Email: user.Email,
|
||||
IsAdmin: user.IsAdmin,
|
||||
})
|
||||
}
|
||||
|
||||
// getClientIP extracts the client IP from the request.
|
||||
func getClientIP(r *http.Request) string {
|
||||
// Check X-Forwarded-For header
|
||||
forwarded := r.Header.Get("X-Forwarded-For")
|
||||
if forwarded != "" {
|
||||
// Take the first IP in the chain
|
||||
parts := strings.Split(forwarded, ",")
|
||||
return strings.TrimSpace(parts[0])
|
||||
}
|
||||
|
||||
// Check X-Real-IP header
|
||||
realIP := r.Header.Get("X-Real-IP")
|
||||
if realIP != "" {
|
||||
return realIP
|
||||
}
|
||||
|
||||
// Fall back to RemoteAddr
|
||||
return r.RemoteAddr
|
||||
}
|
||||
|
||||
// respondJSON writes a JSON response.
|
||||
func respondJSON(w http.ResponseWriter, status int, data interface{}) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"github.com/michael/flowreader/internal/domain"
|
||||
)
|
||||
|
||||
// SessionRepository implements domain.SessionRepository using PostgreSQL.
|
||||
type SessionRepository struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewSessionRepository creates a new session repository.
|
||||
func NewSessionRepository(pool *pgxpool.Pool) *SessionRepository {
|
||||
return &SessionRepository{pool: pool}
|
||||
}
|
||||
|
||||
// Create inserts a new session into the database.
|
||||
func (r *SessionRepository) Create(session *domain.Session) error {
|
||||
ctx := context.Background()
|
||||
|
||||
query := `
|
||||
INSERT INTO sessions (id, user_id, token, expires_at, created_at, user_agent, ip_address)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
`
|
||||
|
||||
_, err := r.pool.Exec(ctx, query,
|
||||
session.ID,
|
||||
session.UserID,
|
||||
session.Token,
|
||||
session.ExpiresAt,
|
||||
session.CreatedAt,
|
||||
session.UserAgent,
|
||||
session.IPAddress,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating session: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetByToken retrieves a session by its token.
|
||||
func (r *SessionRepository) GetByToken(token string) (*domain.Session, error) {
|
||||
ctx := context.Background()
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, token, expires_at, created_at, user_agent, ip_address
|
||||
FROM sessions
|
||||
WHERE token = $1 AND expires_at > NOW()
|
||||
`
|
||||
|
||||
var session domain.Session
|
||||
var userAgent, ipAddress *string
|
||||
err := r.pool.QueryRow(ctx, query, token).Scan(
|
||||
&session.ID,
|
||||
&session.UserID,
|
||||
&session.Token,
|
||||
&session.ExpiresAt,
|
||||
&session.CreatedAt,
|
||||
&userAgent,
|
||||
&ipAddress,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, nil // Session not found or expired
|
||||
}
|
||||
return nil, fmt.Errorf("getting session by token: %w", err)
|
||||
}
|
||||
|
||||
if userAgent != nil {
|
||||
session.UserAgent = *userAgent
|
||||
}
|
||||
if ipAddress != nil {
|
||||
session.IPAddress = *ipAddress
|
||||
}
|
||||
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
// Delete removes a session by its token.
|
||||
func (r *SessionRepository) Delete(token string) error {
|
||||
ctx := context.Background()
|
||||
|
||||
query := `DELETE FROM sessions WHERE token = $1`
|
||||
_, err := r.pool.Exec(ctx, query, token)
|
||||
if err != nil {
|
||||
return fmt.Errorf("deleting session: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteByUserID removes all sessions for a user.
|
||||
func (r *SessionRepository) DeleteByUserID(userID uuid.UUID) error {
|
||||
ctx := context.Background()
|
||||
|
||||
query := `DELETE FROM sessions WHERE user_id = $1`
|
||||
_, err := r.pool.Exec(ctx, query, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("deleting user sessions: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteExpired removes all expired sessions.
|
||||
func (r *SessionRepository) DeleteExpired() (int64, error) {
|
||||
ctx := context.Background()
|
||||
|
||||
query := `DELETE FROM sessions WHERE expires_at < NOW()`
|
||||
result, err := r.pool.Exec(ctx, query)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("deleting expired sessions: %w", err)
|
||||
}
|
||||
|
||||
return result.RowsAffected(), nil
|
||||
}
|
||||
+151
-8
@@ -3,10 +3,12 @@ package service
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
@@ -25,21 +27,27 @@ var (
|
||||
|
||||
// Argon2 parameters (OWASP recommended)
|
||||
const (
|
||||
argon2Time = 1
|
||||
argon2Memory = 64 * 1024 // 64MB
|
||||
argon2Threads = 4
|
||||
argon2KeyLen = 32
|
||||
saltLength = 16
|
||||
argon2Time = 1
|
||||
argon2Memory = 64 * 1024 // 64MB
|
||||
argon2Threads = 4
|
||||
argon2KeyLen = 32
|
||||
saltLength = 16
|
||||
tokenLength = 32
|
||||
sessionDuration = 7 * 24 * time.Hour // 7 days
|
||||
)
|
||||
|
||||
// AuthService handles user authentication business logic.
|
||||
type AuthService struct {
|
||||
userRepo domain.UserRepository
|
||||
userRepo domain.UserRepository
|
||||
sessionRepo domain.SessionRepository
|
||||
}
|
||||
|
||||
// NewAuthService creates a new authentication service.
|
||||
func NewAuthService(userRepo domain.UserRepository) *AuthService {
|
||||
return &AuthService{userRepo: userRepo}
|
||||
func NewAuthService(userRepo domain.UserRepository, sessionRepo domain.SessionRepository) *AuthService {
|
||||
return &AuthService{
|
||||
userRepo: userRepo,
|
||||
sessionRepo: sessionRepo,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRequest contains the data needed to register a new user.
|
||||
@@ -55,6 +63,28 @@ type RegisterResponse struct {
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// LoginRequest contains the data needed to log in.
|
||||
type LoginRequest struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
UserAgent string `json:"-"`
|
||||
IPAddress string `json:"-"`
|
||||
}
|
||||
|
||||
// LoginResponse contains the session token.
|
||||
type LoginResponse struct {
|
||||
Token string `json:"token"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
User UserInfo `json:"user"`
|
||||
}
|
||||
|
||||
// UserInfo contains basic user information.
|
||||
type UserInfo struct {
|
||||
ID uuid.UUID `json:"id"`
|
||||
Email string `json:"email"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
}
|
||||
|
||||
// Register creates a new user account.
|
||||
func (s *AuthService) Register(req RegisterRequest) (*RegisterResponse, error) {
|
||||
// Validate email format
|
||||
@@ -104,6 +134,78 @@ func (s *AuthService) Register(req RegisterRequest) (*RegisterResponse, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Login authenticates a user and creates a session.
|
||||
func (s *AuthService) Login(req LoginRequest) (*LoginResponse, error) {
|
||||
// Find user by email
|
||||
user, err := s.userRepo.GetByEmail(req.Email)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("finding user: %w", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
// Verify password
|
||||
if !verifyPassword(req.Password, user.PasswordHash) {
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
// Generate session token
|
||||
token, err := generateToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("generating token: %w", err)
|
||||
}
|
||||
|
||||
// Create session
|
||||
now := time.Now()
|
||||
session := &domain.Session{
|
||||
ID: uuid.New(),
|
||||
UserID: user.ID,
|
||||
Token: token,
|
||||
ExpiresAt: now.Add(sessionDuration),
|
||||
CreatedAt: now,
|
||||
UserAgent: req.UserAgent,
|
||||
IPAddress: req.IPAddress,
|
||||
}
|
||||
|
||||
if err := s.sessionRepo.Create(session); err != nil {
|
||||
return nil, fmt.Errorf("creating session: %w", err)
|
||||
}
|
||||
|
||||
return &LoginResponse{
|
||||
Token: token,
|
||||
ExpiresAt: session.ExpiresAt,
|
||||
User: UserInfo{
|
||||
ID: user.ID,
|
||||
Email: user.Email,
|
||||
IsAdmin: user.IsAdmin,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Logout invalidates a session.
|
||||
func (s *AuthService) Logout(token string) error {
|
||||
return s.sessionRepo.Delete(token)
|
||||
}
|
||||
|
||||
// GetUserByToken retrieves the user associated with a session token.
|
||||
func (s *AuthService) GetUserByToken(token string) (*domain.User, error) {
|
||||
session, err := s.sessionRepo.GetByToken(token)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("getting session: %w", err)
|
||||
}
|
||||
if session == nil {
|
||||
return nil, nil // Invalid or expired session
|
||||
}
|
||||
|
||||
user, err := s.userRepo.GetByID(session.UserID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("getting user: %w", err)
|
||||
}
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// hashPassword creates an Argon2id hash of the password.
|
||||
func hashPassword(password string) (string, error) {
|
||||
salt := make([]byte, saltLength)
|
||||
@@ -126,6 +228,47 @@ func hashPassword(password string) (string, error) {
|
||||
return encoded, nil
|
||||
}
|
||||
|
||||
// verifyPassword checks if the password matches the hash.
|
||||
func verifyPassword(password, encodedHash string) bool {
|
||||
// Parse the encoded hash
|
||||
parts := strings.Split(encodedHash, "$")
|
||||
if len(parts) != 6 {
|
||||
return false
|
||||
}
|
||||
|
||||
var memory, time uint32
|
||||
var threads uint8
|
||||
_, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
expectedHash, err := base64.RawStdEncoding.DecodeString(parts[5])
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// Compute hash with same parameters
|
||||
computedHash := argon2.IDKey([]byte(password), salt, time, memory, threads, uint32(len(expectedHash)))
|
||||
|
||||
// Constant-time comparison
|
||||
return subtle.ConstantTimeCompare(expectedHash, computedHash) == 1
|
||||
}
|
||||
|
||||
// generateToken creates a cryptographically secure random token.
|
||||
func generateToken() (string, error) {
|
||||
bytes := make([]byte, tokenLength)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.URLEncoding.EncodeToString(bytes), nil
|
||||
}
|
||||
|
||||
// isValidEmail checks if the email has a valid format.
|
||||
func isValidEmail(email string) bool {
|
||||
emailRegex := regexp.MustCompile(`^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$`)
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
-- Rollback: 002_create_sessions
|
||||
|
||||
DROP TABLE IF EXISTS sessions;
|
||||
@@ -0,0 +1,21 @@
|
||||
-- Migration: 002_create_sessions
|
||||
-- Description: Create the sessions table for stateful authentication
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
token VARCHAR(64) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMPTZ NOT NULL,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
user_agent TEXT,
|
||||
ip_address VARCHAR(45)
|
||||
);
|
||||
|
||||
-- Index for token lookup (authentication)
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_token ON sessions(token);
|
||||
|
||||
-- Index for user_id (session cleanup)
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions(user_id);
|
||||
|
||||
-- Index for expiration cleanup
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_expires_at ON sessions(expires_at);
|
||||
Reference in new issue
Block a user