Interactive install, with flags kept for unattended installs

In a terminal, install now asks for the domain, Let's Encrypt email,
endpoint, WireGuard port and admin password, shows a summary and changes
nothing until confirmed. Re-runs offer the current settings and say how
many devices need a new config when the endpoint or port changes.

Flags still work and skip their question; -y or no terminal skips all
questions. New -port flag sets the WireGuard UDP port. Every value is
checked before the system is touched.
This commit is contained in:
Daniel Redetzke
2026-10-03 23:46:04 +03:00
parent 5feb73100f
commit 7d4cd48ca4
5 changed files with 577 additions and 38 deletions
+398
View File
@@ -0,0 +1,398 @@
package main
import (
"bufio"
"context"
"errors"
"fmt"
"io"
"net"
"regexp"
"strconv"
"strings"
)
// installPlan is what install changes in config.json. It comes from the
// flags, and in a terminal from the answers to its questions.
type installPlan struct {
domain, email, endpoint string
port int
noDomain bool // turn an existing domain off (self-signed certificate)
noEmail bool // remove an existing Let's Encrypt email
passwordHash string // asked in the terminal; empty: asked later
ipv4 string // tunnel network of a new install
}
var domainRe = regexp.MustCompile(`^([a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)+[a-zA-Z]{2,63}$`)
func checkDomain(d string) error {
if !domainRe.MatchString(d) {
return fmt.Errorf("%q is not a domain name like vpn.example.net", d)
}
return nil
}
func checkEmail(e string) error {
if at := strings.Index(e, "@"); at < 1 || at == len(e)-1 || strings.ContainsAny(e, " ,;") {
return fmt.Errorf("%q is not an email address", e)
}
return nil
}
func checkEndpoint(h string) error {
if net.ParseIP(h) != nil || domainRe.MatchString(h) {
return nil
}
return fmt.Errorf("%q is not a host name or IP address (without port)", h)
}
func checkPort(p int) error {
if p < 1 || p > 65535 {
return errors.New("the port must be 1–65535")
}
return nil
}
// check validates the flags before anything is changed.
func (p installPlan) check() error {
if p.domain != "" {
if err := checkDomain(p.domain); err != nil {
return err
}
}
if p.email != "" {
if err := checkEmail(p.email); err != nil {
return err
}
}
if p.endpoint != "" {
if err := checkEndpoint(p.endpoint); err != nil {
return err
}
}
if p.port != 0 {
return checkPort(p.port)
}
return nil
}
func (p installPlan) changes() bool {
return p.domain != "" || p.email != "" || p.endpoint != "" || p.port != 0 || p.noDomain || p.noEmail || p.passwordHash != ""
}
func (p installPlan) apply(c *Config) {
if p.noDomain {
if c.Web.TLS.Mode == "acme" {
c.Web.TLS.Mode = "selfsigned"
}
c.Web.TLS.Domain, c.Web.TLS.Email = "", ""
}
if p.domain != "" {
c.Web.TLS.Mode, c.Web.TLS.Domain = "acme", p.domain
if c.Server.Endpoint == "" {
c.Server.Endpoint = p.domain
}
}
if p.noEmail {
c.Web.TLS.Email = ""
}
if p.email != "" {
c.Web.TLS.Email = p.email
}
if p.endpoint != "" {
c.Server.Endpoint = p.endpoint
}
if p.port != 0 {
c.Server.ListenPort = p.port
}
if p.passwordHash != "" {
c.Admin.PasswordHash = p.passwordHash
}
}
// reissueCount is the number of devices whose config stops working because
// the endpoint host or port changes.
func (p installPlan) reissueCount(cur *Config, existing bool) int {
if !existing {
return 0
}
next := cur.clone()
p.apply(next)
if endpointString(next) == endpointString(cur) {
return 0
}
n := 0
for i := range cur.Peers {
if cur.Peers[i].hasKey() {
n++
}
}
return n
}
// --- questions ---
var errCancelled = errors.New("install cancelled; nothing was changed")
type prompter struct{ r *bufio.Reader }
// ask prints a question and returns the answer, def on Enter. check may
// reject an answer; the question is then asked again.
func (pr prompter) ask(question, def string, check func(string) error) (string, error) {
for {
if def != "" {
fmt.Printf(" %s [%s]\n > ", question, def)
} else {
fmt.Printf(" %s\n > ", question)
}
line, err := pr.r.ReadString('\n')
if err != nil && (err != io.EOF || line == "") {
fmt.Println()
return "", errCancelled
}
v := strings.TrimSpace(line)
if v == "" {
v = def
}
if check != nil {
if err := check(v); err != nil {
fmt.Printf(" ✗ %v\n", err)
continue
}
}
return v, nil
}
}
func (pr prompter) confirm(question string) (bool, error) {
fmt.Printf("%s [Y/n] ", question)
line, err := pr.r.ReadString('\n')
if err != nil && line == "" {
fmt.Println()
return false, errCancelled
}
switch strings.ToLower(strings.TrimSpace(line)) {
case "", "y", "yes":
return true, nil
}
return false, nil
}
// askInstall asks for every setting no flag gave, shows a summary and asks
// for confirmation. Nothing on the system has changed when it returns.
func askInstall(in io.Reader, cur *Config, existing bool, given map[string]bool, p installPlan) (installPlan, error) {
pr := prompter{bufio.NewReader(in)}
fmt.Printf("\n%s %s — WireGuard server manager\n", appName, version)
if existing {
fmt.Println("Already installed: the current settings are in [brackets]. Press Enter to keep them.")
} else {
fmt.Println("Press Enter to accept the value in [brackets].")
}
// Web interface
fmt.Println("\nWeb interface")
curDomain := ""
if cur.Web.TLS.Mode == "acme" {
curDomain = cur.Web.TLS.Domain
}
domain := p.domain
if !given["domain"] {
q := "Domain name for the web interface (empty: no domain, self-signed certificate)"
if curDomain != "" {
q = "Domain name for the web interface (none: no domain, self-signed certificate)"
}
v, err := pr.ask(q, curDomain, func(v string) error {
if v == "" || v == "none" {
return nil
}
return checkDomain(v)
})
if err != nil {
return p, err
}
switch {
case v == "none" || v == "":
p.noDomain, domain = curDomain != "", ""
case v != curDomain:
p.domain, domain = v, v
default:
domain = v
}
}
if domain == "" && cur.Web.TLS.Mode == "acme" && !p.noDomain {
domain = curDomain
}
if domain != "" && !given["email"] {
def := cur.Web.TLS.Email
q := "Email for Let's Encrypt expiry warnings (optional)"
if def != "" {
q = "Email for Let's Encrypt expiry warnings (none: no email)"
}
v, err := pr.ask(q, def, func(v string) error {
if v == "" || v == "none" {
return nil
}
return checkEmail(v)
})
if err != nil {
return p, err
}
switch {
case v == "none":
p.noEmail = true
case v != cur.Web.TLS.Email:
p.email = v
}
}
// WireGuard
fmt.Println("\nWireGuard")
if !given["endpoint"] {
def, hint := cur.Server.Endpoint, ""
// An endpoint that followed the old domain follows the new one.
if def == "" || (domain != "" && def == curDomain) {
def = domain
}
if def == "" {
if ip, err := detectPublicIP(context.Background()); err == nil {
def, hint = ip.String(), " (detected public IP)"
}
}
v, err := pr.ask("Address devices connect to"+hint, def, func(v string) error {
if v == "" {
return errors.New("devices need an address to connect to")
}
return checkEndpoint(v)
})
if err != nil {
return p, err
}
if v != cur.Server.Endpoint {
p.endpoint = v
}
}
if !given["port"] {
v, err := pr.ask("UDP port", strconv.Itoa(cur.Server.ListenPort), func(v string) error {
n, err := strconv.Atoi(v)
if err != nil {
return errors.New("the port must be a number")
}
return checkPort(n)
})
if err != nil {
return p, err
}
if n, _ := strconv.Atoi(v); n != cur.Server.ListenPort {
p.port = n
}
}
// Admin account, only when no password is set yet.
if cur.Admin.PasswordHash == "" {
fmt.Println("\nAdmin account")
for {
pw, err := readSecret(fmt.Sprintf(" Password for %q (at least 12 characters): ", cur.Admin.Username))
if err != nil {
return p, errCancelled
}
if err := validatePassword(pw); err != nil {
fmt.Printf(" ✗ %v\n", err)
continue
}
again, err := readSecret(" Repeat password: ")
if err != nil {
return p, errCancelled
}
if again != pw {
fmt.Println(" ✗ the passwords do not match")
continue
}
if p.passwordHash, err = hashPassword(pw); err != nil {
return p, err
}
break
}
}
printInstallSummary(cur, existing, p)
ok, err := pr.confirm("\nInstall with these settings?")
if err != nil {
return p, err
}
if !ok {
return p, errCancelled
}
fmt.Println()
return p, nil
}
func printInstallSummary(cur *Config, existing bool, p installPlan) {
next := cur.clone()
if !existing {
next.Server.IPv4 = p.ipv4
next.Server.IPv6Enabled = hasGlobalIPv6()
}
p.apply(next)
if next.Server.Endpoint == "" {
next.Server.Endpoint = next.Web.TLS.Domain
}
host := next.Web.TLS.Domain
if host == "" {
host = next.Server.Endpoint
}
_, webPort, _ := strings.Cut(next.Web.Listen, ":")
if webPort != "" && webPort != "443" {
host += ":" + webPort
}
web := "https://" + host + "/"
switch next.Web.TLS.Mode {
case "acme":
web += " (Let's Encrypt"
if next.Web.TLS.Email != "" {
web += ", " + next.Web.TLS.Email
}
web += ")"
case "selfsigned":
web += " (self-signed certificate: the browser warns once)"
case "files":
web += " (your certificate files)"
case "off":
web = "http://" + host + "/ (plain HTTP behind a reverse proxy)"
}
ep := endpointString(next) + "/udp"
if n := p.reissueCount(cur, existing); n > 0 {
ep += fmt.Sprintf(" (was %s — %d existing device(s) need a new config)", endpointString(cur), n)
}
tunnel := next.Server.IPv4
if !existing {
tunnel += " (random free range)"
}
if next.Server.IPv6Enabled {
tunnel += " · IPv6 on"
} else {
tunnel += " · IPv6 off (no public IPv6 address)"
}
var ports []string
if next.Web.TLS.Mode != "off" {
if webPort == "" {
webPort = "443"
}
ports = append(ports, webPort+"/tcp")
if next.Web.TLS.Mode == "acme" {
ports = append(ports, "80/tcp")
}
}
ports = append(ports, strconv.Itoa(next.Server.ListenPort)+"/udp")
fmt.Println("\nSummary")
fmt.Printf(" Web interface %s\n", web)
fmt.Printf(" Endpoint %s\n", ep)
fmt.Printf(" Tunnel network %s\n", tunnel)
fmt.Printf(" Firewall %s must be reachable\n", strings.Join(ports, ", "))
if p.passwordHash != "" {
fmt.Printf(" Admin %s (password set)\n", next.Admin.Username)
}
}