3482a03707
On a server running pivpn's WireGuard, a new install offers to take it over: the server key, port, MTU, tunnel networks, endpoint, DNS, AllowedIPs and keepalive, and every client with its public key, preshared key and addresses. Devices keep their configs. Clients pivpn switched off are imported switched off, with the note "Imported from pivpn". Client private keys in /etc/wireguard/configs are not read. - Install notes which peers are connected, stops wg-quick@wg0, starts GHOSTWIRE on the same wg0 and waits up to 30 s for those peers. The wait only reports; idle devices reconnect when they next send. - If the service does not stay running, install removes what it set up, including config.json, and starts pivpn's WireGuard again. - Without a terminal the takeover needs -import-pivpn; install refuses to run next to pivpn otherwise, and the flag is refused on an existing install. - Names GHOSTWIRE does not accept are renamed and listed in the summary. An IPv6 address that differs from the mapped one is kept on the peer until its config is issued again. - uninstall without a config of its own (e.g. after a takeover was undone) leaves the WireGuard interface alone and removes only the firewall table. - README: "Coming from pivpn?" under the intro, a Features entry and a "Moving from pivpn" section. Tested end to end on Ubuntu 24.04 with pivpn aa96de7.
444 lines
11 KiB
Go
444 lines
11 KiB
Go
package main
|
||
|
||
import (
|
||
"bufio"
|
||
"cmp"
|
||
"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
|
||
pivpn *pivpnSetup // pivpn's WireGuard server to take over
|
||
}
|
||
|
||
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 != ""
|
||
}
|
||
|
||
// errPivpnDeclined ends an install whose admin keeps pivpn: both would use
|
||
// the same interface.
|
||
var errPivpnDeclined = errors.New("install cancelled; nothing was changed. GHOSTWIRE and pivpn cannot both run the WireGuard interface: remove pivpn first, or install again and take it over")
|
||
|
||
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.Users[0].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 && p.pivpn == nil {
|
||
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].")
|
||
}
|
||
|
||
// pivpn: taken over first, so its settings become the defaults below.
|
||
if p.pivpn != nil && !given["import-pivpn"] {
|
||
s := cur.Server
|
||
nets := s.IPv4
|
||
if s.IPv6Enabled {
|
||
nets += " + " + s.IPv6
|
||
}
|
||
fmt.Println("\npivpn found")
|
||
fmt.Printf(" %s · %s · UDP %d · %s\n", s.Interface, nets, s.ListenPort, cmp.Or(s.Endpoint, "no endpoint"))
|
||
fmt.Printf(" %s: %s\n", plural(len(p.pivpn.Peers), "client"), cmp.Or(p.pivpn.names(), "none"))
|
||
ok, err := pr.confirm(" Take over this WireGuard server? The devices keep working without new configs.")
|
||
if err != nil {
|
||
return p, err
|
||
}
|
||
if !ok {
|
||
return p, errPivpnDeclined
|
||
}
|
||
}
|
||
|
||
// 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
|
||
}
|
||
}
|
||
|
||
// First user, only when nobody has a password yet.
|
||
if !cur.passwordSet() {
|
||
fmt.Println("\nAdmin account")
|
||
for {
|
||
pw, err := readSecret(fmt.Sprintf(" Password for %q (at least 12 characters): ", cur.Users[0].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 && p.pivpn == nil {
|
||
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
|
||
switch {
|
||
case p.pivpn != nil:
|
||
tunnel += " (from pivpn)"
|
||
case !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.Users[0].Username)
|
||
}
|
||
if pv := p.pivpn; pv != nil {
|
||
s := next.Server
|
||
fmt.Printf(" From pivpn server key, %s with their keys and addresses,\n", plural(len(pv.Peers), "peer"))
|
||
fmt.Printf(" DNS %s · keepalive %d s · MTU %d\n", strings.Join(s.ClientDefaults.DNS, ", "), s.ClientDefaults.Keepalive, s.MTU)
|
||
for _, r := range pv.Renamed {
|
||
fmt.Printf(" Renamed %s → %s (names here: 1–32 letters, digits, . @ _ -)\n", r[0], r[1])
|
||
}
|
||
fmt.Printf(" Not taken client private keys in %s (never stored here)\n", pv.ClientKeys)
|
||
}
|
||
}
|
||
|
||
// plural writes "1 peer" or "5 peers".
|
||
func plural(n int, word string) string {
|
||
if n == 1 {
|
||
return "1 " + word
|
||
}
|
||
return strconv.Itoa(n) + " " + word + "s"
|
||
}
|