49 lines
1.2 KiB
Go
49 lines
1.2 KiB
Go
package auth
|
|
|
|
import (
|
|
"database/sql"
|
|
"net/http"
|
|
)
|
|
|
|
// RoleChecker validates that the session user has the required role.
|
|
type RoleChecker struct {
|
|
db *sql.DB
|
|
}
|
|
|
|
// NewRoleChecker creates a role checker backed by the database.
|
|
func NewRoleChecker(db *sql.DB) *RoleChecker {
|
|
return &RoleChecker{db: db}
|
|
}
|
|
|
|
// RequireAdmin is middleware that allows only users with the "admin" role.
|
|
// Must run after SessionMiddleware has populated the context.
|
|
func (rc *RoleChecker) RequireAdmin(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
userID, ok := GetUserID(r)
|
|
if !ok {
|
|
http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized)
|
|
return
|
|
}
|
|
|
|
isAdmin, err := rc.IsAdmin(userID)
|
|
if err != nil || !isAdmin {
|
|
http.Error(w, `{"error":"forbidden"}`, http.StatusForbidden)
|
|
return
|
|
}
|
|
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
// IsAdmin checks if a user has the admin role.
|
|
func (rc *RoleChecker) IsAdmin(username string) (bool, error) {
|
|
var role string
|
|
err := rc.db.QueryRow("SELECT role FROM users WHERE username = ?", username).Scan(&role)
|
|
if err == sql.ErrNoRows {
|
|
return false, nil
|
|
}
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
return role == "admin", nil
|
|
}
|