356 lines
9.2 KiB
Go
356 lines
9.2 KiB
Go
package app
|
|
|
|
import (
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"golang.org/x/crypto/bcrypt"
|
|
)
|
|
|
|
type serverUserRecord struct {
|
|
ID int64
|
|
Email string
|
|
Username string
|
|
IsAdmin bool
|
|
CreatedAt time.Time
|
|
}
|
|
|
|
type serverUserEdit struct {
|
|
Email *string
|
|
Username *string
|
|
Password *string
|
|
Admin *bool
|
|
}
|
|
|
|
func cliServerUsers(args []string) error {
|
|
dsn, args, err := extractRequiredDatabaseDSN(args)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if len(args) < 1 {
|
|
return errors.New("usage: gitocean server users --dsn DSN <list|view|add|edit|remove>")
|
|
}
|
|
db, err := openMySQLAndCreateDatabaseIfMissing(dsn)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer db.Close()
|
|
if err := migrate(db); err != nil {
|
|
return err
|
|
}
|
|
switch args[0] {
|
|
case "list":
|
|
return serverUsersList(db)
|
|
case "view":
|
|
if len(args) != 2 {
|
|
return errors.New("usage: gitocean server users --dsn DSN view USER")
|
|
}
|
|
return serverUsersView(db, args[1])
|
|
case "add":
|
|
email, username, password, admin, err := parseServerUsersAddArgs(args[1:])
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return serverUsersAdd(db, email, username, password, admin)
|
|
case "edit":
|
|
ref, edit, err := parseServerUsersEditArgs(args[1:])
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return serverUsersEditUser(db, ref, edit)
|
|
case "remove":
|
|
ref, err := parseServerUsersRemoveArgs(args[1:])
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return serverUsersRemove(db, ref)
|
|
default:
|
|
return fmt.Errorf("unknown server users command %q", args[0])
|
|
}
|
|
}
|
|
|
|
func extractRequiredDatabaseDSN(args []string) (string, []string, error) {
|
|
out := make([]string, 0, len(args))
|
|
var dsn string
|
|
for i := 0; i < len(args); i++ {
|
|
arg := args[i]
|
|
switch {
|
|
case arg == "--dsn" || arg == "-dsn" || arg == "--database-url":
|
|
if i+1 >= len(args) {
|
|
return "", nil, fmt.Errorf("%s requires a value", arg)
|
|
}
|
|
i++
|
|
dsn = args[i]
|
|
case strings.HasPrefix(arg, "--dsn="):
|
|
dsn = strings.TrimPrefix(arg, "--dsn=")
|
|
case strings.HasPrefix(arg, "-dsn="):
|
|
dsn = strings.TrimPrefix(arg, "-dsn=")
|
|
case strings.HasPrefix(arg, "--database-url="):
|
|
dsn = strings.TrimPrefix(arg, "--database-url=")
|
|
default:
|
|
out = append(out, arg)
|
|
}
|
|
}
|
|
if strings.TrimSpace(dsn) == "" {
|
|
return "", nil, errors.New("database URL is required; pass --dsn DSN")
|
|
}
|
|
return dsn, out, nil
|
|
}
|
|
|
|
func parseServerUsersAddArgs(args []string) (email, username, password string, admin bool, err error) {
|
|
for i := 0; i < len(args); i++ {
|
|
switch args[i] {
|
|
case "--email":
|
|
if i+1 >= len(args) {
|
|
return "", "", "", false, errors.New("--email requires a value")
|
|
}
|
|
i++
|
|
email = args[i]
|
|
case "--username":
|
|
if i+1 >= len(args) {
|
|
return "", "", "", false, errors.New("--username requires a value")
|
|
}
|
|
i++
|
|
username = args[i]
|
|
case "--password":
|
|
if i+1 >= len(args) {
|
|
return "", "", "", false, errors.New("--password requires a value")
|
|
}
|
|
i++
|
|
password = args[i]
|
|
case "--admin":
|
|
admin = true
|
|
default:
|
|
return "", "", "", false, fmt.Errorf("unknown flag %s", args[i])
|
|
}
|
|
}
|
|
if email == "" {
|
|
email = prompt("Email: ")
|
|
}
|
|
if username == "" {
|
|
username = prompt("Username: ")
|
|
}
|
|
if password == "" {
|
|
password = prompt("Password: ")
|
|
}
|
|
return email, username, password, admin, nil
|
|
}
|
|
|
|
func parseServerUsersEditArgs(args []string) (string, serverUserEdit, error) {
|
|
var ref string
|
|
var edit serverUserEdit
|
|
for i := 0; i < len(args); i++ {
|
|
switch args[i] {
|
|
case "--email":
|
|
if i+1 >= len(args) {
|
|
return "", edit, errors.New("--email requires a value")
|
|
}
|
|
i++
|
|
v := args[i]
|
|
edit.Email = &v
|
|
case "--username":
|
|
if i+1 >= len(args) {
|
|
return "", edit, errors.New("--username requires a value")
|
|
}
|
|
i++
|
|
v := args[i]
|
|
edit.Username = &v
|
|
case "--password":
|
|
if i+1 >= len(args) {
|
|
return "", edit, errors.New("--password requires a value")
|
|
}
|
|
i++
|
|
v := args[i]
|
|
edit.Password = &v
|
|
case "--admin":
|
|
if i+1 >= len(args) {
|
|
return "", edit, errors.New("--admin requires true or false")
|
|
}
|
|
i++
|
|
v, err := parseFlexibleBool(args[i])
|
|
if err != nil {
|
|
return "", edit, err
|
|
}
|
|
edit.Admin = &v
|
|
default:
|
|
if strings.HasPrefix(args[i], "-") {
|
|
return "", edit, fmt.Errorf("unknown flag %s", args[i])
|
|
}
|
|
if ref != "" {
|
|
return "", edit, errors.New("usage: gitocean server users --dsn DSN edit USER [--email EMAIL] [--username USERNAME] [--password PASSWORD] [--admin true|false]")
|
|
}
|
|
ref = args[i]
|
|
}
|
|
}
|
|
if ref == "" {
|
|
return "", edit, errors.New("usage: gitocean server users --dsn DSN edit USER [--email EMAIL] [--username USERNAME] [--password PASSWORD] [--admin true|false]")
|
|
}
|
|
if edit.Email == nil && edit.Username == nil && edit.Password == nil && edit.Admin == nil {
|
|
return "", edit, errors.New("no user changes provided")
|
|
}
|
|
return ref, edit, nil
|
|
}
|
|
|
|
func parseServerUsersRemoveArgs(args []string) (string, error) {
|
|
force := false
|
|
var ref string
|
|
for _, arg := range args {
|
|
switch arg {
|
|
case "--force", "-force":
|
|
force = true
|
|
default:
|
|
if strings.HasPrefix(arg, "-") {
|
|
return "", fmt.Errorf("unknown flag %s", arg)
|
|
}
|
|
if ref != "" {
|
|
return "", errors.New("usage: gitocean server users --dsn DSN remove USER --force")
|
|
}
|
|
ref = arg
|
|
}
|
|
}
|
|
if ref == "" || !force {
|
|
return "", errors.New("usage: gitocean server users --dsn DSN remove USER --force")
|
|
}
|
|
return ref, nil
|
|
}
|
|
|
|
func parseFlexibleBool(s string) (bool, error) {
|
|
switch strings.ToLower(strings.TrimSpace(s)) {
|
|
case "1", "true", "t", "yes", "y":
|
|
return true, nil
|
|
case "0", "false", "f", "no", "n":
|
|
return false, nil
|
|
default:
|
|
return false, fmt.Errorf("invalid boolean %q", s)
|
|
}
|
|
}
|
|
|
|
func serverUsersList(db *sql.DB) error {
|
|
rows, err := db.Query(`SELECT id, email, username, is_admin, created_at FROM users ORDER BY username`)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer rows.Close()
|
|
fmt.Printf("%-6s %-24s %-34s %-7s %s\n", "ID", "USERNAME", "EMAIL", "ADMIN", "CREATED")
|
|
for rows.Next() {
|
|
var u serverUserRecord
|
|
if err := rows.Scan(&u.ID, &u.Email, &u.Username, &u.IsAdmin, &u.CreatedAt); err != nil {
|
|
return err
|
|
}
|
|
printServerUserRow(u)
|
|
}
|
|
return rows.Err()
|
|
}
|
|
|
|
func serverUsersView(db *sql.DB, ref string) error {
|
|
u, err := loadServerUser(db, ref)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
fmt.Printf("ID: %d\n", u.ID)
|
|
fmt.Printf("Username: %s\n", u.Username)
|
|
fmt.Printf("Email: %s\n", u.Email)
|
|
fmt.Printf("Admin: %v\n", u.IsAdmin)
|
|
fmt.Printf("Created: %s\n", u.CreatedAt.Format(time.RFC3339))
|
|
return nil
|
|
}
|
|
|
|
func serverUsersAdd(db *sql.DB, email, username, password string, admin bool) error {
|
|
email = strings.ToLower(strings.TrimSpace(email))
|
|
username = strings.ToLower(strings.TrimSpace(username))
|
|
if err := createUserDirect(db, email, username, password, admin); err != nil {
|
|
return err
|
|
}
|
|
fmt.Printf("Created user %s <%s> admin=%v\n", username, email, admin)
|
|
return nil
|
|
}
|
|
|
|
func serverUsersEditUser(db *sql.DB, ref string, edit serverUserEdit) error {
|
|
set := []string{}
|
|
args := []any{}
|
|
if edit.Email != nil {
|
|
email := strings.ToLower(strings.TrimSpace(*edit.Email))
|
|
if !strings.Contains(email, "@") || len(email) > 255 {
|
|
return errors.New("invalid email")
|
|
}
|
|
set = append(set, "email = ?")
|
|
args = append(args, email)
|
|
}
|
|
if edit.Username != nil {
|
|
username := strings.ToLower(strings.TrimSpace(*edit.Username))
|
|
if !usernameRE.MatchString(username) || isReservedName(username) {
|
|
return errors.New("invalid or reserved username")
|
|
}
|
|
set = append(set, "username = ?")
|
|
args = append(args, username)
|
|
}
|
|
if edit.Password != nil {
|
|
if len(*edit.Password) < 8 {
|
|
return errors.New("password must be at least 8 characters")
|
|
}
|
|
hash, err := bcrypt.GenerateFromPassword([]byte(*edit.Password), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
set = append(set, "password_hash = ?")
|
|
args = append(args, string(hash))
|
|
}
|
|
if edit.Admin != nil {
|
|
set = append(set, "is_admin = ?")
|
|
args = append(args, *edit.Admin)
|
|
}
|
|
where, whereArgs := serverUserWhere(ref)
|
|
args = append(args, whereArgs...)
|
|
res, err := db.Exec("UPDATE users SET "+strings.Join(set, ", ")+" WHERE "+where, args...)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
affected, _ := res.RowsAffected()
|
|
if affected == 0 {
|
|
return sql.ErrNoRows
|
|
}
|
|
fmt.Printf("Updated user %s\n", ref)
|
|
return nil
|
|
}
|
|
|
|
func serverUsersRemove(db *sql.DB, ref string) error {
|
|
where, args := serverUserWhere(ref)
|
|
res, err := db.Exec("DELETE FROM users WHERE "+where, args...)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
affected, _ := res.RowsAffected()
|
|
if affected == 0 {
|
|
return sql.ErrNoRows
|
|
}
|
|
fmt.Printf("Removed user %s\n", ref)
|
|
return nil
|
|
}
|
|
|
|
func loadServerUser(db *sql.DB, ref string) (serverUserRecord, error) {
|
|
where, args := serverUserWhere(ref)
|
|
var u serverUserRecord
|
|
err := db.QueryRow(`SELECT id, email, username, is_admin, created_at FROM users WHERE `+where, args...).Scan(&u.ID, &u.Email, &u.Username, &u.IsAdmin, &u.CreatedAt)
|
|
return u, err
|
|
}
|
|
|
|
func serverUserWhere(ref string) (string, []any) {
|
|
ref = strings.TrimSpace(ref)
|
|
if id, err := strconv.ParseInt(ref, 10, 64); err == nil {
|
|
return "id = ?", []any{id}
|
|
}
|
|
ref = strings.ToLower(ref)
|
|
if strings.Contains(ref, "@") {
|
|
return "email = ?", []any{ref}
|
|
}
|
|
return "username = ?", []any{ref}
|
|
}
|
|
|
|
func printServerUserRow(u serverUserRecord) {
|
|
fmt.Printf("%-6d %-24s %-34s %-7v %s\n", u.ID, u.Username, u.Email, u.IsAdmin, u.CreatedAt.Format("2006-01-02 15:04"))
|
|
}
|