Traefik_Control/backend/internal/api/middleware/auth.go

231 lines
No EOL
5.9 KiB
Go

package middleware
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/traefik/traefik-gui/backend/internal/auth"
"github.com/traefik/traefik-gui/backend/internal/database/repositories"
"github.com/traefik/traefik-gui/backend/internal/models"
)
const (
SessionCookieName = "traefik_gui_session"
CSRFHeaderName = "X-CSRF-Token"
UserContextKey = "user"
SessionContextKey = "session"
)
type AuthMiddleware struct {
sessionRepo *repositories.SessionRepository
userRepo *repositories.UserRepository
}
func NewAuthMiddleware(sessionRepo *repositories.SessionRepository, userRepo *repositories.UserRepository) *AuthMiddleware {
return &AuthMiddleware{
sessionRepo: sessionRepo,
userRepo: userRepo,
}
}
func (m *AuthMiddleware) RequireAuth() gin.HandlerFunc {
return func(c *gin.Context) {
sessionID, err := c.Cookie(SessionCookieName)
if err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
return
}
session, err := m.sessionRepo.GetByID(sessionID)
if err != nil || session == nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid session"})
return
}
if session.ExpiresAt.Before(time.Now()) {
m.sessionRepo.Delete(sessionID)
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "session expired"})
return
}
user, err := m.userRepo.GetByID(session.UserID)
if err != nil || user == nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not found"})
return
}
c.Set(SessionContextKey, session)
c.Set(UserContextKey, user)
c.Next()
}
}
func (m *AuthMiddleware) RequireCSRF() gin.HandlerFunc {
return func(c *gin.Context) {
if c.Request.Method == "GET" || c.Request.Method == "HEAD" || c.Request.Method == "OPTIONS" {
c.Next()
return
}
sessionVal, exists := c.Get(SessionContextKey)
if !exists {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "no session"})
return
}
session := sessionVal.(*models.Session)
csrfToken := c.GetHeader(CSRFHeaderName)
if csrfToken == "" {
csrfToken = c.PostForm("_csrf")
}
if csrfToken != session.CSRFToken {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "invalid CSRF token"})
return
}
// Rotate CSRF token on successful validation
newToken, err := auth.GenerateCSRFToken()
if err == nil {
m.sessionRepo.RotateCSRFToken(session.ID, newToken)
c.Header(CSRFHeaderName, newToken)
}
c.Next()
}
}
func (m *AuthMiddleware) OptionalAuth() gin.HandlerFunc {
return func(c *gin.Context) {
sessionID, err := c.Cookie(SessionCookieName)
if err != nil {
c.Next()
return
}
session, err := m.sessionRepo.GetByID(sessionID)
if err != nil || session == nil {
c.Next()
return
}
if session.ExpiresAt.Before(time.Now()) {
m.sessionRepo.Delete(sessionID)
c.Next()
return
}
user, err := m.userRepo.GetByID(session.UserID)
if err != nil || user == nil {
c.Next()
return
}
c.Set(SessionContextKey, session)
c.Set(UserContextKey, user)
c.Next()
}
}
func GetUser(c *gin.Context) *models.User {
val, exists := c.Get(UserContextKey)
if !exists {
return nil
}
return val.(*models.User)
}
func GetSession(c *gin.Context) *models.Session {
val, exists := c.Get(SessionContextKey)
if !exists {
return nil
}
return val.(*models.Session)
}
func RequireRole(allowedRoles ...string) gin.HandlerFunc {
return func(c *gin.Context) {
user := GetUser(c)
if user == nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
return
}
for _, role := range allowedRoles {
if user.Role == role {
c.Next()
return
}
}
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "insufficient permissions"})
}
}
func CORSMiddleware(allowedOrigin string) gin.HandlerFunc {
// Reject wildcard when credentials are enabled — browsers will block it anyway.
// Only the explicitly configured origin is allowed.
isWildcard := allowedOrigin == "*"
return func(c *gin.Context) {
origin := c.Request.Header.Get("Origin")
if !isWildcard && origin != "" && origin == allowedOrigin {
c.Header("Access-Control-Allow-Origin", origin)
c.Header("Vary", "Origin")
}
// Explicitly do not set Access-Control-Allow-Origin to "*" when credentials are true
c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS")
c.Header("Access-Control-Allow-Headers", "Content-Type, Authorization, X-CSRF-Token")
c.Header("Access-Control-Allow-Credentials", "true")
c.Header("Access-Control-Max-Age", "86400")
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(http.StatusNoContent)
return
}
c.Next()
}
}
func SecurityHeadersMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("X-Frame-Options", "DENY")
c.Header("X-XSS-Protection", "1; mode=block")
c.Header("Referrer-Policy", "strict-origin-when-cross-origin")
c.Header("Strict-Transport-Security", "max-age=31536000; includeSubDomains; preload")
c.Header("Content-Security-Policy", "default-src 'self'; script-src 'self'; style-src 'self' 'unsafe-inline'; img-src 'self' data:; font-src 'self'; connect-src 'self'; frame-ancestors 'none';")
c.Header("Permissions-Policy", "camera=(), microphone=(), geolocation=(), payment=()")
c.Header("Cache-Control", "no-store, no-cache, must-revalidate, proxy-revalidate")
c.Header("Pragma", "no-cache")
c.Next()
}
}
func LoggingMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
start := time.Now()
path := c.Request.URL.Path
raw := c.Request.URL.RawQuery
c.Next()
latency := time.Since(start)
clientIP := c.ClientIP()
method := c.Request.Method
statusCode := c.Writer.Status()
if raw != "" {
path = path + "?" + raw
}
// Log via zerolog in production
_ = statusCode // avoid unused in dev
_ = clientIP
_ = method
_ = path
_ = latency
}
}