Files

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"))
}