Files
wx-win-agent/control-plane/user_store.go
T
rogee 8760fa49f0
Build web service image / build (push) Successful in 2m7s
feat: add control-plane identity and management console
2026-09-27 20:27:31 +08:00

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
}