package auth
import (
"crypto/pbkdf2"
"crypto/rand"
"crypto/sha256"
"crypto/sha3"
"encoding/base64"
"errors"
"fmt"
"hash"
"log/slog"
"net/http"
"net/url"
"regexp"
"strconv"
"strings"
"gopkg.awl.red/lx/internal/config"
)
var (
ErrTooManyUsers = errors.New("too many users")
ErrSessionExpired = errors.New("session expired")
ErrBadHashFields = errors.New("expected 5 hash fields, got other")
ErrUserNotApproved = errors.New("user not approved")
ErrUserPasswordLocked = errors.New("password locked")
ErrHashMismatch = errors.New("password hashes did not match")
)
type Session struct {
PersonId int `json:"person_id"`
ExpiresAt int `json:"expires_at"`
Roles map[string]struct{} `json:"roles"`
}
const (
RoleBusinessOwner = "business_owner"
RoleNewPersonApprover = "new_person_approver"
RoleAdmin = "site_admin"
)
var Roles = map[string]struct{}{
RoleBusinessOwner: {},
RoleNewPersonApprover: {},
RoleAdmin: {},
}
const (
PersonRequestTypeUserApprove = "user_approve"
PersonRequestTypePasswordReset = "password_reset"
PersonRequestTypeChangeEmail = "change_email"
PersonRequestTypeConfirmEmail = "confirm_email"
)
var (
ErrNoUsername = errors.New("no username provided")
ErrNoFirstname = errors.New("no firstname provided")
ErrNoLastname = errors.New("no lastname provided")
ErrPasswordMismatch = errors.New("password and confirm password do not match")
ErrPasswordTooShort = errors.New("password too short")
)
type VettedSignupRequest struct {
Firstname string `json:"firstname"`
Lastname string `json:"lastname"`
Username string `json:"username"`
Email string `json:"email"`
RequestText string `json:"request_text"`
Hashed string `json:"hashed"`
}
func HandleSignupForm(f url.Values, cfg *config.Config) (*VettedSignupRequest, error) {
signupFails := []error{}
firstname := f.Get("firstname")
lastname := f.Get("lastname")
username := f.Get("username")
email := f.Get("email")
password := f.Get("password")
passwordConfirm := f.Get("password_confirm")
requestText := f.Get("request_text")
emailReStrs := strings.Split(cfg.EmailRegex, "@@@")
if email != "" || cfg.EmailRequired {
for _, reStr := range emailReStrs {
re, err := regexp.Compile(reStr)
if err != nil {
slog.Error("bad email regex", "err", err, "re", reStr)
return nil, err
}
if !re.MatchString(email) {
signupFails = append(signupFails, fmt.Errorf("failed email regexp test %q", reStr))
}
}
}
if username == "" {
signupFails = append(signupFails, ErrNoUsername)
}
if firstname == "" {
signupFails = append(signupFails, ErrNoFirstname)
}
if lastname == "" {
signupFails = append(signupFails, ErrNoLastname)
}
if password != passwordConfirm {
signupFails = append(signupFails, ErrPasswordMismatch)
}
if password != "" && len(password) < cfg.PasswordMin {
signupFails = append(signupFails, ErrPasswordTooShort)
}
hashed, err := toPwHash(cfg.PasswordHashStrategy, cfg.PasswordSaltLength, password, cfg.PasswordRounds, cfg.PasswordKeyLength)
if err != nil {
signupFails = append(signupFails, err)
}
if allErrs := errors.Join(signupFails...); allErrs != nil {
return nil, allErrs
}
return &VettedSignupRequest{
Firstname: firstname,
Lastname: lastname,
Username: username,
Email: email,
RequestText: requestText,
Hashed: hashed,
}, nil
}
func StrategyToHash(strategy string) func() hash.Hash {
var hsh func() hash.Hash
switch strategy {
case config.HashStrategySha2_256:
hsh = sha256.New
case config.HashStrategySha3_256:
hsh = func() hash.Hash {
return sha3.New256()
}
case config.HashStrategySha3_512:
hsh = func() hash.Hash {
return sha3.New512()
}
default:
slog.Warn("unknown password hash strategy", "unknown", strategy, "chosen_default", "SHA_256")
hsh = sha256.New
}
return hsh
}
func PwMatchesHash(pw, hashed string) error {
hashParts := strings.Split(hashed, ":")
if len(hashParts) != 5 {
return ErrBadHashFields
}
strategy, saltBytesB64, roundsStr, keyLengthStr, hashedStr := hashParts[0],
hashParts[1],
hashParts[2],
hashParts[3],
hashParts[4]
hsh := StrategyToHash(strategy)
saltBytes, err := base64.StdEncoding.DecodeString(saltBytesB64)
if err != nil {
return err
}
rounds, err := strconv.Atoi(roundsStr)
if err != nil {
return err
}
keyLength, err := strconv.Atoi(keyLengthStr)
if err != nil {
return err
}
keyBytes, err := pbkdf2.Key(hsh, pw, saltBytes, rounds, keyLength)
if err != nil {
return err
}
if hashedStr != AsB64(keyBytes) {
return ErrHashMismatch
}
return nil
}
func toPwHash(strategy string, saltLength int, password string, rounds int, keyLength int) (string, error) {
hsh := StrategyToHash(strategy)
saltBytes, err := newRandBytes(saltLength)
if err != nil {
return "", err
}
keyBytes, err := pbkdf2.Key(hsh, password, saltBytes, rounds, keyLength)
if err != nil {
return "", err
}
// strategy:salt:rounds:key_length:hash
hashed := fmt.Sprintf("%s:%s:%d:%d:%s",
strategy,
AsB64(saltBytes),
rounds,
keyLength,
AsB64(keyBytes),
)
return hashed, nil
}
func AsB64(bs []byte) string {
return base64.StdEncoding.EncodeToString(bs)
}
func newRandBytes(length int) ([]byte, error) {
randomBytes := make([]byte, length)
_, err := rand.Read(randomBytes)
return randomBytes, err
}
func NewToken() string {
tokenBytes, _ := newRandBytes(64)
return AsB64(tokenBytes)
}
const DelSessionCookie = "Session=; Path=/; Max-Age=0; HttpOnly; Secure; Partitioned; SameSite=Lax; Expires=Thu, 01 Jan 1970 00:00:00 GMT"
func SetSession(h http.Header, token string, cfg *config.Config) {
cookieVal := fmt.Sprintf(
"Session=%s; Path=/; Max-Age=%d; HttpOnly; Secure; Partitioned; SameSite=Lax",
token,
cfg.SessionLengthMax,
)
h.Set("Set-Cookie", cookieVal)
}