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