improve auth by using sessions (refresh tokens)

This commit is contained in:
2026-07-02 22:39:58 +02:00
parent 0b2f042faa
commit 532f3f5faa
5 changed files with 169 additions and 35 deletions
+1 -1
View File
@@ -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)
}
+143 -22
View File
@@ -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))
}
+9 -12
View File
@@ -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 <token>" 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()
}
+14
View File
@@ -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"`
}
+2
View File
@@ -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