improve auth by using sessions (refresh tokens)
This commit is contained in:
@@ -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
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user