improve auth by using sessions (refresh tokens)
This commit is contained in:
@@ -29,7 +29,7 @@ func InitDatabase() {
|
|||||||
isNewDatabase := len(tables) == 0
|
isNewDatabase := len(tables) == 0
|
||||||
|
|
||||||
// Auto migrate the schema
|
// 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 {
|
if err != nil {
|
||||||
log.Fatal("Failed to migrate database:", err)
|
log.Fatal("Failed to migrate database:", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+143
-22
@@ -1,6 +1,10 @@
|
|||||||
package handlers
|
package handlers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/hex"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sheetless-server/config"
|
"sheetless-server/config"
|
||||||
"sheetless-server/database"
|
"sheetless-server/database"
|
||||||
@@ -9,13 +13,13 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
"github.com/google/uuid"
|
||||||
"golang.org/x/crypto/bcrypt"
|
"golang.org/x/crypto/bcrypt"
|
||||||
)
|
)
|
||||||
|
|
||||||
type RegisterRequest struct {
|
type RegisterRequest struct {
|
||||||
Username string `json:"username" binding:"required"`
|
Username string `json:"username" binding:"required,min=4"`
|
||||||
Email string `json:"email" binding:"required,email"`
|
Password string `json:"password" binding:"required,min=8"`
|
||||||
Password string `json:"password" binding:"required,min=6"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type LoginRequest struct {
|
type LoginRequest struct {
|
||||||
@@ -23,6 +27,16 @@ type LoginRequest struct {
|
|||||||
Password string `json:"password" binding:"required"`
|
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) {
|
func Register(c *gin.Context) {
|
||||||
var req RegisterRequest
|
var req RegisterRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
@@ -32,8 +46,8 @@ func Register(c *gin.Context) {
|
|||||||
|
|
||||||
// Check if user already exists
|
// Check if user already exists
|
||||||
var existingUser models.User
|
var existingUser models.User
|
||||||
if err := database.DB.Where("username = ? OR email = ?", req.Username, req.Email).First(&existingUser).Error; err == nil {
|
if err := database.DB.Where("username = ?", req.Username).First(&existingUser).Error; err == nil {
|
||||||
c.JSON(http.StatusConflict, gin.H{"error": "User already exists"})
|
c.JSON(http.StatusConflict, gin.H{"error": "Invalid user name"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -45,12 +59,14 @@ func Register(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Create user
|
// Create user
|
||||||
|
now := time.Now()
|
||||||
user := models.User{
|
user := models.User{
|
||||||
Username: req.Username,
|
Uuid: uuid.New(),
|
||||||
Email: req.Email,
|
Username: req.Username,
|
||||||
Password: string(hashedPassword),
|
HashedPassword: string(hashedPassword),
|
||||||
|
CreatedAt: now,
|
||||||
|
UpdatedAt: now,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := database.DB.Create(&user).Error; err != nil {
|
if err := database.DB.Create(&user).Error; err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to create user"})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to create user"})
|
||||||
return
|
return
|
||||||
@@ -66,30 +82,135 @@ func Login(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find user
|
// Check usename and password
|
||||||
var user models.User
|
var user models.User
|
||||||
if err := database.DB.Where("username = ?", req.Username).First(&user).Error; err != nil {
|
if err := database.DB.Where("username = ?", req.Username).First(&user).Error; err != nil {
|
||||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"})
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if err := bcrypt.CompareHashAndPassword([]byte(user.HashedPassword), []byte(req.Password)); err != nil {
|
||||||
// Check password
|
|
||||||
if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil {
|
|
||||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"})
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "Invalid credentials"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Generate JWT token
|
// Generate refresh token
|
||||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
refreshToken, err := generateRefreshToken()
|
||||||
"user_id": user.ID,
|
|
||||||
"exp": time.Now().Add(time.Hour * 24).Unix(),
|
|
||||||
})
|
|
||||||
|
|
||||||
tokenString, err := token.SignedString([]byte(config.AppConfig.JWT.Secret))
|
|
||||||
if err != nil {
|
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
|
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
|
package middleware
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sheetless-server/config"
|
"sheetless-server/config"
|
||||||
|
"sheetless-server/handlers"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -13,7 +15,7 @@ func AuthMiddleware() gin.HandlerFunc {
|
|||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
authHeader := c.GetHeader("Authorization")
|
authHeader := c.GetHeader("Authorization")
|
||||||
if authHeader == "" {
|
if authHeader == "" {
|
||||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "Authorization header required"})
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "no_authorization_header"})
|
||||||
c.Abort()
|
c.Abort()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -21,28 +23,23 @@ func AuthMiddleware() gin.HandlerFunc {
|
|||||||
// Extract token from "Bearer <token>" format
|
// Extract token from "Bearer <token>" format
|
||||||
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
||||||
if tokenString == authHeader {
|
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()
|
c.Abort()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse and validate token
|
// 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
|
return []byte(config.AppConfig.JWT.Secret), nil
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil || !token.Valid {
|
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()
|
c.Abort()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
c.Set("user_uuid", claims.UserUuid)
|
||||||
// Extract claims
|
c.Set("session_uuid", claims.SessionUuid)
|
||||||
if claims, ok := token.Claims.(jwt.MapClaims); ok {
|
|
||||||
if userID, ok := claims["user_id"].(float64); ok {
|
|
||||||
c.Set("user_id", uint(userID))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
c.Next()
|
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("/register", handlers.Register)
|
||||||
auth.POST("/login", handlers.Login)
|
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
|
// Protected routes
|
||||||
|
|||||||
Reference in New Issue
Block a user