231 lines
No EOL
5.9 KiB
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
|
|
}
|
|
} |