487 lines
15 KiB
Go
487 lines
15 KiB
Go
package controlplane
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/subtle"
|
|
"database/sql"
|
|
"encoding/base64"
|
|
"errors"
|
|
"fmt"
|
|
"sort"
|
|
"strings"
|
|
"time"
|
|
"unicode/utf8"
|
|
|
|
"golang.org/x/crypto/argon2"
|
|
)
|
|
|
|
const (
|
|
RoleAdmin = "admin"
|
|
RoleMember = "member"
|
|
passwordMemory = 64 * 1024
|
|
passwordTime = 3
|
|
passwordThreads = 2
|
|
passwordKeySize = 32
|
|
passwordSaltLen = 16
|
|
)
|
|
|
|
var (
|
|
ErrUserNotFound = errors.New("user not found")
|
|
ErrUserAlreadyExists = errors.New("user already exists")
|
|
ErrLastAdministrator = errors.New("cannot disable or demote the last active administrator")
|
|
ErrInvalidUser = errors.New("invalid user")
|
|
ErrNoBootstrapAdmin = errors.New("no web users are configured to bootstrap the identity database")
|
|
)
|
|
|
|
type WebUser struct {
|
|
ID string `json:"id"`
|
|
Username string `json:"username"`
|
|
DisplayName string `json:"display_name"`
|
|
Role string `json:"role"`
|
|
Active bool `json:"active"`
|
|
NodeIDs []string `json:"node_ids"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
UpdatedAt time.Time `json:"updated_at"`
|
|
}
|
|
|
|
type WebPrincipal struct {
|
|
WebUser
|
|
}
|
|
|
|
func (p WebPrincipal) CanAccessNode(nodeID string) bool {
|
|
if p.Role == RoleAdmin {
|
|
return true
|
|
}
|
|
for _, assignedID := range p.NodeIDs {
|
|
if assignedID == nodeID {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
type webCredential struct {
|
|
user WebUser
|
|
passwordHash string
|
|
}
|
|
|
|
type UserUpdate struct {
|
|
DisplayName string
|
|
Role string
|
|
Active bool
|
|
NodeIDs []string
|
|
}
|
|
|
|
type UserStore struct {
|
|
db *sql.DB
|
|
}
|
|
|
|
func NewUserStore(db *sql.DB) (*UserStore, error) {
|
|
if db == nil {
|
|
return nil, errors.New("user store database is required")
|
|
}
|
|
store := &UserStore{db: db}
|
|
if err := store.init(context.Background()); err != nil {
|
|
return nil, err
|
|
}
|
|
return store, nil
|
|
}
|
|
|
|
func (s *UserStore) init(ctx context.Context) error {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
CREATE TABLE IF NOT EXISTS users (
|
|
id TEXT PRIMARY KEY,
|
|
username TEXT NOT NULL COLLATE NOCASE UNIQUE,
|
|
display_name TEXT NOT NULL,
|
|
password_hash TEXT NOT NULL,
|
|
role TEXT NOT NULL CHECK (role IN ('admin', 'member')),
|
|
active INTEGER NOT NULL CHECK (active IN (0, 1)),
|
|
created_at TEXT NOT NULL,
|
|
updated_at TEXT NOT NULL
|
|
);
|
|
CREATE TABLE IF NOT EXISTS user_nodes (
|
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
node_id TEXT NOT NULL,
|
|
created_at TEXT NOT NULL,
|
|
PRIMARY KEY (user_id, node_id)
|
|
);
|
|
CREATE INDEX IF NOT EXISTS user_nodes_node_id_idx ON user_nodes(node_id);
|
|
CREATE TABLE IF NOT EXISTS web_auth_schema (
|
|
id INTEGER PRIMARY KEY CHECK (id = 1),
|
|
version INTEGER NOT NULL
|
|
);
|
|
INSERT OR IGNORE INTO web_auth_schema (id, version) VALUES (1, 1);
|
|
`)
|
|
if err != nil {
|
|
return fmt.Errorf("initialize user database: %w", err)
|
|
}
|
|
var version int
|
|
if err := s.db.QueryRowContext(ctx, `SELECT version FROM web_auth_schema WHERE id = 1`).Scan(&version); err != nil {
|
|
return fmt.Errorf("read user schema version: %w", err)
|
|
}
|
|
if version != 1 {
|
|
return fmt.Errorf("unsupported user schema version %d", version)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// BootstrapAdmins imports legacy environment credentials only while the new user table is empty.
|
|
func (s *UserStore) BootstrapAdmins(ctx context.Context, credentials map[string]string) error {
|
|
var count int
|
|
if err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&count); err != nil {
|
|
return fmt.Errorf("count users: %w", err)
|
|
}
|
|
if count > 0 {
|
|
return nil
|
|
}
|
|
if len(credentials) == 0 {
|
|
return ErrNoBootstrapAdmin
|
|
}
|
|
users := make([]string, 0, len(credentials))
|
|
for username := range credentials {
|
|
users = append(users, username)
|
|
}
|
|
sort.Strings(users)
|
|
tx, err := s.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return fmt.Errorf("begin admin bootstrap: %w", err)
|
|
}
|
|
defer tx.Rollback()
|
|
now := time.Now().UTC().Format(time.RFC3339Nano)
|
|
for _, username := range users {
|
|
password := credentials[username]
|
|
if !validUsername(username) || password == "" {
|
|
return ErrInvalidUser
|
|
}
|
|
hash, err := hashPassword(password)
|
|
if err != nil {
|
|
return fmt.Errorf("hash bootstrap password: %w", err)
|
|
}
|
|
if _, err := tx.ExecContext(ctx, `
|
|
INSERT INTO users (id, username, display_name, password_hash, role, active, created_at, updated_at)
|
|
VALUES (?, ?, ?, ?, ?, 1, ?, ?)`, randomID(), strings.TrimSpace(username), strings.TrimSpace(username), hash, RoleAdmin, now, now); err != nil {
|
|
return fmt.Errorf("bootstrap administrator: %w", err)
|
|
}
|
|
}
|
|
return tx.Commit()
|
|
}
|
|
|
|
func (s *UserStore) Count(ctx context.Context) (int, error) {
|
|
var count int
|
|
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&count)
|
|
return count, err
|
|
}
|
|
|
|
func (s *UserStore) Authenticate(ctx context.Context, username, password string) (WebUser, bool, error) {
|
|
var credential webCredential
|
|
var active int
|
|
var created, updated string
|
|
err := s.db.QueryRowContext(ctx, `
|
|
SELECT id, username, display_name, password_hash, role, active, created_at, updated_at
|
|
FROM users WHERE username = ? COLLATE NOCASE`, strings.TrimSpace(username)).
|
|
Scan(&credential.user.ID, &credential.user.Username, &credential.user.DisplayName,
|
|
&credential.passwordHash, &credential.user.Role, &active, &created, &updated)
|
|
credential.user.Active = active == 1
|
|
if err == nil {
|
|
credential.user.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
|
|
if err == nil {
|
|
credential.user.UpdatedAt, err = time.Parse(time.RFC3339Nano, updated)
|
|
}
|
|
}
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return WebUser{}, false, nil
|
|
}
|
|
if err != nil {
|
|
return WebUser{}, false, fmt.Errorf("read user credentials: %w", err)
|
|
}
|
|
if !verifyPassword(credential.passwordHash, password) || !credential.user.Active {
|
|
return WebUser{}, false, nil
|
|
}
|
|
user, err := s.Get(ctx, credential.user.ID)
|
|
return user, err == nil, err
|
|
}
|
|
|
|
func (s *UserStore) List(ctx context.Context) ([]WebUser, error) {
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT u.id, u.username, u.display_name, u.role, u.active, u.created_at, u.updated_at, un.node_id
|
|
FROM users u LEFT JOIN user_nodes un ON un.user_id = u.id
|
|
ORDER BY u.username COLLATE NOCASE, un.node_id`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("list users: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
users := make([]WebUser, 0)
|
|
byID := make(map[string]int)
|
|
for rows.Next() {
|
|
var user WebUser
|
|
var active int
|
|
var created, updated string
|
|
var nodeID sql.NullString
|
|
if err := rows.Scan(&user.ID, &user.Username, &user.DisplayName, &user.Role, &active, &created, &updated, &nodeID); err != nil {
|
|
return nil, fmt.Errorf("scan user: %w", err)
|
|
}
|
|
index, exists := byID[user.ID]
|
|
if !exists {
|
|
user.Active = active == 1
|
|
user.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse user creation time: %w", err)
|
|
}
|
|
user.UpdatedAt, err = time.Parse(time.RFC3339Nano, updated)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse user update time: %w", err)
|
|
}
|
|
user.NodeIDs = []string{}
|
|
users = append(users, user)
|
|
index = len(users) - 1
|
|
byID[user.ID] = index
|
|
}
|
|
if nodeID.Valid {
|
|
users[index].NodeIDs = append(users[index].NodeIDs, nodeID.String)
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("read users: %w", err)
|
|
}
|
|
return users, nil
|
|
}
|
|
|
|
func (s *UserStore) Get(ctx context.Context, id string) (WebUser, error) {
|
|
users, err := s.listBy(ctx, id)
|
|
if err != nil {
|
|
return WebUser{}, err
|
|
}
|
|
if len(users) == 0 {
|
|
return WebUser{}, ErrUserNotFound
|
|
}
|
|
return users[0], nil
|
|
}
|
|
|
|
func (s *UserStore) listBy(ctx context.Context, id string) ([]WebUser, error) {
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT u.id, u.username, u.display_name, u.role, u.active, u.created_at, u.updated_at, un.node_id
|
|
FROM users u LEFT JOIN user_nodes un ON un.user_id = u.id WHERE u.id = ? ORDER BY un.node_id`, id)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("read user: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
var user WebUser
|
|
found := false
|
|
for rows.Next() {
|
|
var active int
|
|
var created, updated string
|
|
var nodeID sql.NullString
|
|
if err := rows.Scan(&user.ID, &user.Username, &user.DisplayName, &user.Role, &active, &created, &updated, &nodeID); err != nil {
|
|
return nil, fmt.Errorf("scan user: %w", err)
|
|
}
|
|
if !found {
|
|
user.Active = active == 1
|
|
user.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse user creation time: %w", err)
|
|
}
|
|
user.UpdatedAt, err = time.Parse(time.RFC3339Nano, updated)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parse user update time: %w", err)
|
|
}
|
|
user.NodeIDs = []string{}
|
|
found = true
|
|
}
|
|
if nodeID.Valid {
|
|
user.NodeIDs = append(user.NodeIDs, nodeID.String)
|
|
}
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("read user: %w", err)
|
|
}
|
|
if !found {
|
|
return nil, nil
|
|
}
|
|
return []WebUser{user}, nil
|
|
}
|
|
|
|
func (s *UserStore) Create(ctx context.Context, user WebUser, password string) (WebUser, error) {
|
|
user.Username = strings.TrimSpace(user.Username)
|
|
user.DisplayName = strings.TrimSpace(user.DisplayName)
|
|
if !validUsername(user.Username) || !validDisplayName(user.DisplayName) || !validRole(user.Role) || !validNewPassword(password) {
|
|
return WebUser{}, ErrInvalidUser
|
|
}
|
|
if user.Role == RoleAdmin {
|
|
user.NodeIDs = nil
|
|
} else {
|
|
user.NodeIDs = uniqueNodeIDs(user.NodeIDs)
|
|
}
|
|
hash, err := hashPassword(password)
|
|
if err != nil {
|
|
return WebUser{}, err
|
|
}
|
|
now := time.Now().UTC()
|
|
user.ID, user.Active, user.CreatedAt, user.UpdatedAt = randomID(), true, now, now
|
|
tx, err := s.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return WebUser{}, fmt.Errorf("begin user creation: %w", err)
|
|
}
|
|
defer tx.Rollback()
|
|
var exists int
|
|
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE username = ? COLLATE NOCASE`, user.Username).Scan(&exists); err != nil {
|
|
return WebUser{}, fmt.Errorf("check username: %w", err)
|
|
}
|
|
if exists > 0 {
|
|
return WebUser{}, ErrUserAlreadyExists
|
|
}
|
|
stamp := now.Format(time.RFC3339Nano)
|
|
if _, err := tx.ExecContext(ctx, `INSERT INTO users (id, username, display_name, password_hash, role, active, created_at, updated_at) VALUES (?, ?, ?, ?, ?, 1, ?, ?)`,
|
|
user.ID, user.Username, user.DisplayName, hash, user.Role, stamp, stamp); err != nil {
|
|
return WebUser{}, fmt.Errorf("create user: %w", err)
|
|
}
|
|
if err := insertUserNodes(ctx, tx, user.ID, user.NodeIDs, stamp); err != nil {
|
|
return WebUser{}, err
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
return WebUser{}, fmt.Errorf("commit user creation: %w", err)
|
|
}
|
|
return user, nil
|
|
}
|
|
|
|
func (s *UserStore) Update(ctx context.Context, id string, update UserUpdate) (WebUser, error) {
|
|
update.DisplayName = strings.TrimSpace(update.DisplayName)
|
|
if !validDisplayName(update.DisplayName) || !validRole(update.Role) {
|
|
return WebUser{}, ErrInvalidUser
|
|
}
|
|
update.NodeIDs = uniqueNodeIDs(update.NodeIDs)
|
|
if update.Role == RoleAdmin {
|
|
update.NodeIDs = nil
|
|
}
|
|
tx, err := s.db.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return WebUser{}, fmt.Errorf("begin user update: %w", err)
|
|
}
|
|
defer tx.Rollback()
|
|
var oldRole string
|
|
var oldActive int
|
|
if err := tx.QueryRowContext(ctx, `SELECT role, active FROM users WHERE id = ?`, id).Scan(&oldRole, &oldActive); err != nil {
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return WebUser{}, ErrUserNotFound
|
|
}
|
|
return WebUser{}, fmt.Errorf("read user before update: %w", err)
|
|
}
|
|
if oldRole == RoleAdmin && oldActive == 1 && (update.Role != RoleAdmin || !update.Active) {
|
|
var otherAdmins int
|
|
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = ? AND active = 1 AND id <> ?`, RoleAdmin, id).Scan(&otherAdmins); err != nil {
|
|
return WebUser{}, fmt.Errorf("count active administrators: %w", err)
|
|
}
|
|
if otherAdmins == 0 {
|
|
return WebUser{}, ErrLastAdministrator
|
|
}
|
|
}
|
|
stamp := time.Now().UTC().Format(time.RFC3339Nano)
|
|
if _, err := tx.ExecContext(ctx, `UPDATE users SET display_name = ?, role = ?, active = ?, updated_at = ? WHERE id = ?`,
|
|
update.DisplayName, update.Role, boolInt(update.Active), stamp, id); err != nil {
|
|
return WebUser{}, fmt.Errorf("update user: %w", err)
|
|
}
|
|
if _, err := tx.ExecContext(ctx, `DELETE FROM user_nodes WHERE user_id = ?`, id); err != nil {
|
|
return WebUser{}, fmt.Errorf("replace user Node assignments: %w", err)
|
|
}
|
|
if err := insertUserNodes(ctx, tx, id, update.NodeIDs, stamp); err != nil {
|
|
return WebUser{}, err
|
|
}
|
|
if err := tx.Commit(); err != nil {
|
|
return WebUser{}, fmt.Errorf("commit user update: %w", err)
|
|
}
|
|
return s.Get(ctx, id)
|
|
}
|
|
|
|
func (s *UserStore) SetPassword(ctx context.Context, id, password string) error {
|
|
if !validNewPassword(password) {
|
|
return ErrInvalidUser
|
|
}
|
|
hash, err := hashPassword(password)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
result, err := s.db.ExecContext(ctx, `UPDATE users SET password_hash = ?, updated_at = ? WHERE id = ?`, hash, time.Now().UTC().Format(time.RFC3339Nano), id)
|
|
if err != nil {
|
|
return fmt.Errorf("update user password: %w", err)
|
|
}
|
|
count, err := result.RowsAffected()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if count == 0 {
|
|
return ErrUserNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func insertUserNodes(ctx context.Context, tx *sql.Tx, userID string, nodeIDs []string, stamp string) error {
|
|
for _, nodeID := range nodeIDs {
|
|
if !validIdentifier(nodeID, 200) {
|
|
return ErrInvalidUser
|
|
}
|
|
if _, err := tx.ExecContext(ctx, `INSERT INTO user_nodes (user_id, node_id, created_at) VALUES (?, ?, ?)`, userID, nodeID, stamp); err != nil {
|
|
return fmt.Errorf("assign user Node: %w", err)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func hashPassword(password string) (string, error) {
|
|
salt := make([]byte, passwordSaltLen)
|
|
if _, err := rand.Read(salt); err != nil {
|
|
return "", fmt.Errorf("generate password salt: %w", err)
|
|
}
|
|
hash := argon2.IDKey([]byte(password), salt, passwordTime, passwordMemory, passwordThreads, passwordKeySize)
|
|
return "argon2id$v=19$m=65536,t=3,p=2$" + base64.RawStdEncoding.EncodeToString(salt) + "$" + base64.RawStdEncoding.EncodeToString(hash), nil
|
|
}
|
|
|
|
func verifyPassword(encoded, password string) bool {
|
|
parts := strings.Split(encoded, "$")
|
|
if len(parts) != 5 || parts[0] != "argon2id" || parts[1] != "v=19" || parts[2] != "m=65536,t=3,p=2" {
|
|
return false
|
|
}
|
|
salt, err := base64.RawStdEncoding.DecodeString(parts[3])
|
|
if err != nil || len(salt) != passwordSaltLen {
|
|
return false
|
|
}
|
|
want, err := base64.RawStdEncoding.DecodeString(parts[4])
|
|
if err != nil || len(want) != passwordKeySize {
|
|
return false
|
|
}
|
|
got := argon2.IDKey([]byte(password), salt, passwordTime, passwordMemory, passwordThreads, passwordKeySize)
|
|
return subtle.ConstantTimeCompare(got, want) == 1
|
|
}
|
|
|
|
func validUsername(username string) bool {
|
|
if len(username) < 3 || len(username) > 64 || strings.TrimSpace(username) != username {
|
|
return false
|
|
}
|
|
for _, char := range username {
|
|
if !(char >= 'a' && char <= 'z' || char >= 'A' && char <= 'Z' || char >= '0' && char <= '9' || strings.ContainsRune("._@-", char)) {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func validDisplayName(name string) bool {
|
|
return name != "" && len(name) <= 120 && utf8.ValidString(name) && strings.TrimSpace(name) == name && !strings.ContainsAny(name, "\r\n\x00")
|
|
}
|
|
|
|
func validRole(role string) bool { return role == RoleAdmin || role == RoleMember }
|
|
|
|
func validNewPassword(password string) bool { return len(password) >= 12 && len(password) <= 256 }
|
|
|
|
func uniqueNodeIDs(values []string) []string {
|
|
seen := make(map[string]struct{}, len(values))
|
|
result := make([]string, 0, len(values))
|
|
for _, value := range values {
|
|
if _, exists := seen[value]; exists {
|
|
continue
|
|
}
|
|
seen[value] = struct{}{}
|
|
result = append(result, value)
|
|
}
|
|
return result
|
|
}
|