From 532f3f5faa1c0ef9edf19158cbad0633a1e0218f Mon Sep 17 00:00:00 2001 From: Julian Mutter Date: Thu, 2 Jul 2026 22:39:58 +0200 Subject: [PATCH] improve auth by using sessions (refresh tokens) --- src/database/connection.go | 2 +- src/handlers/auth.go | 165 ++++++++++++++++++++++++++++++++----- src/middleware/auth.go | 21 ++--- src/models/session.go | 14 ++++ src/routes/routes.go | 2 + 5 files changed, 169 insertions(+), 35 deletions(-) create mode 100644 src/models/session.go diff --git a/src/database/connection.go b/src/database/connection.go index 99dc508..6a15743 100644 --- a/src/database/connection.go +++ b/src/database/connection.go @@ -29,7 +29,7 @@ func InitDatabase() { isNewDatabase := len(tables) == 0 // Auto migrate the schema - err = DB.AutoMigrate(&models.User{}, &models.Sheet{}, &models.Composer{}, &models.Annotation{}) + err = DB.AutoMigrate(&models.User{}, &models.Sheet{}, &models.Composer{}, &models.Annotation{}, &models.Session{}) if err != nil { log.Fatal("Failed to migrate database:", err) } diff --git a/src/handlers/auth.go b/src/handlers/auth.go index 719a246..bb44264 100644 --- a/src/handlers/auth.go +++ b/src/handlers/auth.go @@ -1,6 +1,10 @@ package handlers import ( + "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/hex" "net/http" "sheetless-server/config" "sheetless-server/database" @@ -9,13 +13,13 @@ import ( "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" "golang.org/x/crypto/bcrypt" ) type RegisterRequest struct { - Username string `json:"username" binding:"required"` - Email string `json:"email" binding:"required,email"` - Password string `json:"password" binding:"required,min=6"` + Username string `json:"username" binding:"required,min=4"` + Password string `json:"password" binding:"required,min=8"` } type LoginRequest struct { @@ -23,6 +27,16 @@ type LoginRequest struct { Password string `json:"password" binding:"required"` } +type RefreshRequest struct { + RefreshToken string `json:"refresh_token" binding:"required"` +} + +type AccessTokenClaims struct { + UserUuid uuid.UUID `json:"user_uuid"` + SessionUuid uuid.UUID `json:"session_uuid"` + jwt.RegisteredClaims +} + func Register(c *gin.Context) { var req RegisterRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -32,8 +46,8 @@ func Register(c *gin.Context) { // Check if user already exists var existingUser models.User - if err := database.DB.Where("username = ? OR email = ?", req.Username, req.Email).First(&existingUser).Error; err == nil { - c.JSON(http.StatusConflict, gin.H{"error": "User already exists"}) + if err := database.DB.Where("username = ?", req.Username).First(&existingUser).Error; err == nil { + c.JSON(http.StatusConflict, gin.H{"error": "Invalid user name"}) return } @@ -45,12 +59,14 @@ func Register(c *gin.Context) { } // Create user + now := time.Now() user := models.User{ - Username: req.Username, - Email: req.Email, - Password: string(hashedPassword), + Uuid: uuid.New(), + Username: req.Username, + HashedPassword: string(hashedPassword), + CreatedAt: now, + UpdatedAt: now, } - if err := database.DB.Create(&user).Error; err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to create user"}) return @@ -66,30 +82,135 @@ func Login(c *gin.Context) { return } - // Find user + // Check usename and password var user models.User if err := database.DB.Where("username = ?", req.Username).First(&user).Error; err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"}) return } - - // Check password - if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil { + if err := bcrypt.CompareHashAndPassword([]byte(user.HashedPassword), []byte(req.Password)); err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"}) return } - // Generate JWT token - token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ - "user_id": user.ID, - "exp": time.Now().Add(time.Hour * 24).Unix(), - }) - - tokenString, err := token.SignedString([]byte(config.AppConfig.JWT.Secret)) + // Generate refresh token + refreshToken, err := generateRefreshToken() if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to generate token"}) + c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to generate refresh token"}) return } - c.JSON(http.StatusOK, gin.H{"token": tokenString}) + // Create session + session := models.Session{ + Uuid: uuid.New(), + UserUuid: user.Uuid, + HashedRefreshToken: hashRefreshToken(refreshToken), + CreatedAt: time.Now(), + Revoked: false, + } + if err := database.DB.Create(&session).Error; err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to create session"}) + return + } + + // Generate access token (JWT) + accessToken, err := generateAccessToken(session) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to generate access token"}) + return + } + + c.JSON(http.StatusOK, gin.H{"refresh_token": refreshToken, "access_token": accessToken}) +} + +func Logout(c *gin.Context) { + userUuidVal, _ := c.Get("user_uuid") + userUuid, ok := userUuidVal.(uuid.UUID) + if !ok { + c.JSON(http.StatusInternalServerError, gin.H{"error": "invalid or missing user_uuid"}) + return + } + sessionUuidVal, _ := c.Get("session_uuid") + sessionUuid, ok := sessionUuidVal.(uuid.UUID) + if !ok { + c.JSON(http.StatusInternalServerError, gin.H{"error": "invalid or missing session_uuid"}) + return + } + + // Delete session + result := database.DB.Where("uuid = ? AND user_uuid = ?", sessionUuid, userUuid).Delete(&models.Session{}) + if result.Error != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Could not delete session"}) + return + } + + if result.RowsAffected == 0 { + c.JSON(http.StatusUnauthorized, gin.H{"error": "Session not existing"}) + return + } + + c.Status(http.StatusOK) +} + +func RefreshAccessToken(c *gin.Context) { + var req RefreshRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // Check session + var session models.Session + if err := database.DB.Where("hashed_refresh_token = ?", hashRefreshToken(req.RefreshToken)).First(&session).Error; err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid refresh token"}) + return + } + if session.Revoked { + c.JSON(http.StatusUnauthorized, gin.H{"error": "Session revoked"}) + return + } + + // Generate access token (JWT) + accessToken, err := generateAccessToken(session) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to generate access token"}) + return + } + + c.JSON(http.StatusOK, gin.H{"access_token": accessToken}) +} + +func generateRefreshToken() (string, error) { + length := 32 + b := make([]byte, length) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} + +func hashRefreshToken(plainToken string) string { + hash := sha256.Sum256([]byte(plainToken)) + return hex.EncodeToString(hash[:]) +} + +func verifyRefreshToken(plainToken string, dbHash string) bool { + incomingHash := hashRefreshToken(plainToken) + match := subtle.ConstantTimeCompare([]byte(incomingHash), []byte(dbHash)) + return match == 1 +} + +func generateAccessToken(session models.Session) (string, error) { + now := time.Now() + accessTokenExpiry := now.Add(15 * time.Minute) + token := jwt.NewWithClaims(jwt.SigningMethodHS256, + AccessTokenClaims{ + SessionUuid: session.Uuid, + UserUuid: session.UserUuid, + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(accessTokenExpiry), + IssuedAt: jwt.NewNumericDate(now), + }, + }) + return token.SignedString([]byte(config.AppConfig.JWT.Secret)) } diff --git a/src/middleware/auth.go b/src/middleware/auth.go index 2054789..07fb0ee 100644 --- a/src/middleware/auth.go +++ b/src/middleware/auth.go @@ -1,8 +1,10 @@ package middleware import ( + "log" "net/http" "sheetless-server/config" + "sheetless-server/handlers" "strings" "github.com/gin-gonic/gin" @@ -13,7 +15,7 @@ func AuthMiddleware() gin.HandlerFunc { return func(c *gin.Context) { authHeader := c.GetHeader("Authorization") if authHeader == "" { - c.JSON(http.StatusUnauthorized, gin.H{"error": "Authorization header required"}) + c.JSON(http.StatusUnauthorized, gin.H{"error": "no_authorization_header"}) c.Abort() return } @@ -21,28 +23,23 @@ func AuthMiddleware() gin.HandlerFunc { // Extract token from "Bearer " format tokenString := strings.TrimPrefix(authHeader, "Bearer ") if tokenString == authHeader { - c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid token format"}) + c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid_token_format"}) c.Abort() return } // Parse and validate token - token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) { + var claims handlers.AccessTokenClaims + token, err := jwt.ParseWithClaims(tokenString, &claims, func(token *jwt.Token) (any, error) { return []byte(config.AppConfig.JWT.Secret), nil }) - if err != nil || !token.Valid { - c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid token"}) + c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid_token"}) c.Abort() return } - - // Extract claims - if claims, ok := token.Claims.(jwt.MapClaims); ok { - if userID, ok := claims["user_id"].(float64); ok { - c.Set("user_id", uint(userID)) - } - } + c.Set("user_uuid", claims.UserUuid) + c.Set("session_uuid", claims.SessionUuid) c.Next() } diff --git a/src/models/session.go b/src/models/session.go new file mode 100644 index 0000000..657a012 --- /dev/null +++ b/src/models/session.go @@ -0,0 +1,14 @@ +package models + +import ( + "github.com/google/uuid" + "time" +) + +type Session struct { + Uuid uuid.UUID `json:"uuid" gorm:"type:uuid;primaryKey"` + UserUuid uuid.UUID `json:"user_uuid" gorm:"not null"` + HashedRefreshToken string `json:"-" gorm:"not null;unique"` + CreatedAt time.Time `json:"created_at"` + Revoked bool `json:"revoked"` +} diff --git a/src/routes/routes.go b/src/routes/routes.go index 6b107ce..7fa9ca1 100644 --- a/src/routes/routes.go +++ b/src/routes/routes.go @@ -13,6 +13,8 @@ func SetupRoutes(r *gin.Engine) { { // auth.POST("/register", handlers.Register) auth.POST("/login", handlers.Login) + auth.POST("/refresh", handlers.RefreshAccessToken) + auth.POST("/logout", middleware.AuthMiddleware(), handlers.Logout) // User needs to be logged in } // Protected routes