259 lines
6 KiB
Go
259 lines
6 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
|
|
"update_server/internal/config"
|
|
"update_server/internal/db"
|
|
)
|
|
|
|
const sessionTokenBytes = 32
|
|
const sessionCookiePath = "/admin"
|
|
|
|
var (
|
|
ErrInvalidCredentials = errors.New("invalid credentials")
|
|
ErrUnauthenticated = errors.New("unauthenticated")
|
|
ErrUnauthorized = errors.New("unauthorized")
|
|
)
|
|
|
|
type Service struct {
|
|
config config.Config
|
|
logger *slog.Logger
|
|
store *db.Store
|
|
}
|
|
|
|
type SessionState struct {
|
|
User db.User
|
|
Session db.Session
|
|
}
|
|
|
|
func NewService(cfg config.Config, logger *slog.Logger, store *db.Store) *Service {
|
|
return &Service{
|
|
config: cfg,
|
|
logger: logger,
|
|
store: store,
|
|
}
|
|
}
|
|
|
|
func (s *Service) EnsureBootstrapAdmin(ctx context.Context) error {
|
|
hasActiveAdmin, err := s.store.Users.HasActiveAdmin(ctx)
|
|
if err != nil {
|
|
return fmt.Errorf("check active admin users: %w", err)
|
|
}
|
|
|
|
if hasActiveAdmin {
|
|
return nil
|
|
}
|
|
|
|
email := strings.TrimSpace(s.config.AdminEmail)
|
|
password := s.config.AdminPassword
|
|
if email == "" || strings.TrimSpace(password) == "" {
|
|
if s.logger != nil {
|
|
s.logger.Warn("no active admin user found; set ADMIN_EMAIL and ADMIN_PASSWORD to bootstrap the first admin")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
if existingUser, err := s.store.Users.GetByEmail(ctx, email); err == nil {
|
|
if s.logger != nil {
|
|
s.logger.Warn("bootstrap admin skipped because the configured email already exists", "email", existingUser.Email)
|
|
}
|
|
|
|
return nil
|
|
} else if !errors.Is(err, db.ErrNotFound) {
|
|
return fmt.Errorf("check bootstrap admin email: %w", err)
|
|
}
|
|
|
|
passwordHash, err := HashPassword(password)
|
|
if err != nil {
|
|
return fmt.Errorf("hash bootstrap admin password: %w", err)
|
|
}
|
|
|
|
if _, err := s.store.Users.Create(ctx, db.CreateUserParams{
|
|
Email: email,
|
|
PasswordHash: passwordHash,
|
|
Role: db.UserRoleAdmin,
|
|
IsActive: true,
|
|
}); err != nil {
|
|
return fmt.Errorf("create bootstrap admin user: %w", err)
|
|
}
|
|
|
|
if s.logger != nil {
|
|
s.logger.Info("bootstrapped admin user", "email", email)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *Service) Authenticate(ctx context.Context, email, password, ipAddress, userAgent string) (string, *SessionState, error) {
|
|
email = strings.TrimSpace(email)
|
|
if email == "" || password == "" {
|
|
return "", nil, ErrInvalidCredentials
|
|
}
|
|
|
|
user, err := s.store.Users.GetByEmail(ctx, email)
|
|
if err != nil {
|
|
if errors.Is(err, db.ErrNotFound) {
|
|
return "", nil, ErrInvalidCredentials
|
|
}
|
|
|
|
return "", nil, fmt.Errorf("load user by email: %w", err)
|
|
}
|
|
|
|
if !user.IsActive {
|
|
return "", nil, ErrInvalidCredentials
|
|
}
|
|
|
|
if err := ComparePassword(user.PasswordHash, password); err != nil {
|
|
return "", nil, ErrInvalidCredentials
|
|
}
|
|
|
|
now := time.Now().UTC()
|
|
expiresAt := now.Add(s.config.SessionTTL)
|
|
token, tokenHash, err := generateSessionToken()
|
|
if err != nil {
|
|
return "", nil, err
|
|
}
|
|
|
|
var session *db.Session
|
|
if err := s.store.WithTx(ctx, func(tx *db.TxStore) error {
|
|
createdSession, err := tx.Sessions.Create(ctx, db.CreateSessionParams{
|
|
UserID: user.ID,
|
|
TokenHash: tokenHash,
|
|
ExpiresAt: expiresAt,
|
|
IPAddress: ipAddress,
|
|
UserAgent: userAgent,
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
if err := tx.Users.UpdateLastLoginAt(ctx, user.ID, now); err != nil {
|
|
return err
|
|
}
|
|
|
|
session = createdSession
|
|
return nil
|
|
}); err != nil {
|
|
return "", nil, fmt.Errorf("create authenticated session: %w", err)
|
|
}
|
|
|
|
user.LastLoginAt = &now
|
|
|
|
return token, &SessionState{
|
|
User: *user,
|
|
Session: *session,
|
|
}, nil
|
|
}
|
|
|
|
func (s *Service) LoadSession(ctx context.Context, token string) (*SessionState, error) {
|
|
token = strings.TrimSpace(token)
|
|
if token == "" {
|
|
return nil, ErrUnauthenticated
|
|
}
|
|
|
|
now := time.Now().UTC()
|
|
record, err := s.store.Sessions.GetActiveWithUserByTokenHash(ctx, hashSessionToken(token), now)
|
|
if err != nil {
|
|
if errors.Is(err, db.ErrNotFound) {
|
|
return nil, ErrUnauthenticated
|
|
}
|
|
|
|
return nil, fmt.Errorf("lookup active session: %w", err)
|
|
}
|
|
|
|
if err := s.store.Sessions.Touch(ctx, record.Session.ID, now); err != nil {
|
|
return nil, fmt.Errorf("touch active session: %w", err)
|
|
}
|
|
|
|
record.Session.LastSeenAt = &now
|
|
|
|
return &SessionState{
|
|
User: record.User,
|
|
Session: record.Session,
|
|
}, nil
|
|
}
|
|
|
|
func (s *Service) InvalidateSession(ctx context.Context, token string) error {
|
|
token = strings.TrimSpace(token)
|
|
if token == "" {
|
|
return nil
|
|
}
|
|
|
|
if err := s.store.Sessions.InvalidateByTokenHash(ctx, hashSessionToken(token), time.Now().UTC()); err != nil && !errors.Is(err, db.ErrNotFound) {
|
|
return fmt.Errorf("invalidate session: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *Service) SessionCookie(token string, expiresAt time.Time) *http.Cookie {
|
|
maxAge := int(time.Until(expiresAt).Seconds())
|
|
if maxAge < 0 {
|
|
maxAge = 0
|
|
}
|
|
|
|
return &http.Cookie{
|
|
Name: s.config.SessionCookieName,
|
|
Value: token,
|
|
Path: sessionCookiePath,
|
|
HttpOnly: true,
|
|
SameSite: http.SameSiteLaxMode,
|
|
Secure: s.config.SecureCookies,
|
|
Expires: expiresAt.UTC(),
|
|
MaxAge: maxAge,
|
|
}
|
|
}
|
|
|
|
func (s *Service) ClearSessionCookie() *http.Cookie {
|
|
return &http.Cookie{
|
|
Name: s.config.SessionCookieName,
|
|
Value: "",
|
|
Path: sessionCookiePath,
|
|
HttpOnly: true,
|
|
SameSite: http.SameSiteLaxMode,
|
|
Secure: s.config.SecureCookies,
|
|
Expires: time.Unix(0, 0).UTC(),
|
|
MaxAge: -1,
|
|
}
|
|
}
|
|
|
|
func (s *Service) SessionCookieName() string {
|
|
return s.config.SessionCookieName
|
|
}
|
|
|
|
func RoleAllowed(actualRole, requiredRole db.UserRole) bool {
|
|
ranks := map[db.UserRole]int{
|
|
db.UserRoleViewer: 1,
|
|
db.UserRoleEditor: 2,
|
|
db.UserRoleAdmin: 3,
|
|
}
|
|
|
|
return ranks[actualRole] >= ranks[requiredRole] && ranks[requiredRole] > 0
|
|
}
|
|
|
|
func generateSessionToken() (string, string, error) {
|
|
bytes := make([]byte, sessionTokenBytes)
|
|
if _, err := rand.Read(bytes); err != nil {
|
|
return "", "", fmt.Errorf("generate session token: %w", err)
|
|
}
|
|
|
|
token := base64.RawURLEncoding.EncodeToString(bytes)
|
|
return token, hashSessionToken(token), nil
|
|
}
|
|
|
|
func hashSessionToken(token string) string {
|
|
sum := sha256.Sum256([]byte(token))
|
|
return hex.EncodeToString(sum[:])
|
|
}
|