From 3e8dba6072877c31e4d6014afba0ad32b4bb5c45 Mon Sep 17 00:00:00 2001 From: Daniel Redetzke Date: Mon, 5 Oct 2026 23:07:30 +0300 Subject: [PATCH] Fixes from the audit: input checks, apply order, sign-in limits MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - The server endpoint must be a plain host name or IP address. It is written into client configs as is, so a newline could add lines such as PreUp, which wg-quick runs as root on the client. - Listen addresses and the session length (1–720 hours) are checked. Before web settings or a restore are saved, the server tries the new listen addresses and certificate files, so a value it cannot start with is refused instead of stopping the service at the next restart. - Kernel applies run one at a time and read the config once it is their turn, so an older config can no longer be applied last. - Pending passkey sign-ins are capped: 10 per address, 1000 in total. - Behind a local proxy, the last X-Forwarded-For entry is the client; earlier ones come from the client and are ignored. - With LAN access off, peers are also kept from the IPv6 networks on the uplink, not only from its private IPv4 networks. - A change that leaves no user with a password is refused, and so is a backup without one or from a newer version. --- api.go | 78 +++++++++++++- app.js | 2 +- auth.go | 9 +- config.go | 41 +++++++- kernel.go | 29 ++++++ kernel_linux.go | 39 +++---- main.go | 4 + main_test.go | 268 ++++++++++++++++++++++++++++++++++++++++++++++++ mfa.go | 43 +++++++- 9 files changed, 485 insertions(+), 28 deletions(-) diff --git a/api.go b/api.go index 89fb701..46f73c4 100644 --- a/api.go +++ b/api.go @@ -3,15 +3,18 @@ package main import ( "cmp" "context" + "crypto/tls" "encoding/json" "errors" "fmt" "io" "log/slog" + "net" "net/http" "net/netip" "slices" "strings" + "syscall" "time" ) @@ -29,7 +32,8 @@ type App struct { geo *Geo // nil in tests updates *Updater // nil in tests started time.Time - shutdown func() // graceful stop; systemd restarts the service + shutdown func() // graceful stop; systemd restarts the service + webAddrs []string // the addresses the web server listens on now } // --- helpers --- @@ -1090,11 +1094,15 @@ func (a *App) patchSettings(w http.ResponseWriter, r *http.Request) { b, _ := json.Marshal(w) return string(b) } - before := listen() + before, oldWeb := listen(), c.Web if err := field(m, "web", &c.Web); err != nil { return err } - restart = listen() != before + if restart = listen() != before; restart { + if err := a.checkWebStart(oldWeb, c.Web); err != nil { + return err + } + } if err := field(m, "stats", &c.Stats); err != nil { return err } @@ -1236,7 +1244,20 @@ func (a *App) restore(w http.ResponseWriter, r *http.Request) { writeErr(w, badRequest("this file has no server key; is it a backup of this app?")) return } - if err := a.store.Update(func(c *Config) error { *c = in; return nil }); err != nil { + if in.Version > configVersion { + writeErr(w, badRequest("this backup is from a newer version of %s; update this server first", appName)) + return + } + in.applyDefaults() + if !in.passwordSet() { + writeErr(w, badRequest("this backup has no user with a password; restoring it would lock everyone out")) + return + } + if err := a.store.Update(func(c *Config) error { + old := c.Web + *c = in + return a.checkWebStart(old, c.Web) + }); err != nil { writeErr(w, err) return } @@ -1244,6 +1265,55 @@ func (a *App) restore(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, map[string]any{"ok": true, "applyError": a.apply(), "restartRequired": true}) } +// checkWebStart refuses web settings the service could not start with: an +// address it cannot listen on, or certificate files it cannot read. The +// service would stop at the next restart, and the web interface and the API +// with it. +func (a *App) checkWebStart(old, next WebConfig) error { + if err := validateListen(next.Listen, "listen address", false); err != nil { + return &userError{err.Error()} + } + if err := validateListen(next.HTTPListen, "HTTP listen address", true); err != nil { + return &userError{err.Error()} + } + if next.TLS.Mode == "files" && next.TLS != old.TLS { + if _, err := tls.LoadX509KeyPair(next.TLS.CertFile, next.TLS.KeyFile); err != nil { + return badRequest("the certificate files cannot be used: %v", err) + } + } + addrs := []string{next.Listen} + if next.HTTPListen != "" && next.TLS.Mode != "off" { + addrs = append(addrs, next.HTTPListen) + } + for _, addr := range addrs { + if err := a.canListen(addr); err != nil { + return badRequest("cannot listen on %s: %v", addr, err) + } + } + return nil +} + +// canListen tries to listen on addr. An address the service listens on now, +// or one whose port it holds, is fine: it is free again after the restart. +func (a *App) canListen(addr string) error { + if slices.Contains(a.webAddrs, addr) { + return nil + } + ln, err := net.Listen("tcp", addr) + if err == nil { + return ln.Close() + } + if errors.Is(err, syscall.EADDRINUSE) { + _, port, _ := net.SplitHostPort(addr) + for _, own := range a.webAddrs { + if _, p, _ := net.SplitHostPort(own); p == port { + return nil + } + } + } + return err +} + // applyRuntime applies the settings that take effect without a restart: log // level and log rotation. Traffic retention is read by the stats sampler. func (a *App) applyRuntime(c *Config) { diff --git a/app.js b/app.js index f51777c..fbb0f19 100644 --- a/app.js +++ b/app.js @@ -1988,7 +1988,7 @@ fieldEl('up6', 'IPv6 uplink interface', h('input', { id: 'up6', class: 'mono', value: draft.uplinkV6, placeholder: 'auto: ' + (srv.detectedUplinkV6 || 'none found'), onInput: str('uplinkV6') })), cb('nat', 'Masquerade (NAT) peer traffic to the internet'), cb('peerToPeer', 'Allow peers to reach each other'), - cb('lanAccess', 'Allow peers to reach the server\'s LAN', 'Private networks on the uplink interface'), + cb('lanAccess', 'Allow peers to reach the server\'s LAN', 'Private IPv4 and the IPv6 networks on the uplink interface'), cb('openPort', 'Accept UDP ' + draft.listenPort + ' in the input chain'))), h('section', { class: 'card', 'aria-labelledby': 'ky' }, diff --git a/auth.go b/auth.go index 60e4c19..f38814c 100644 --- a/auth.go +++ b/auth.go @@ -287,9 +287,14 @@ func remoteIP(r *http.Request) string { host = r.RemoteAddr } // Behind a local reverse proxy the real client is in X-Forwarded-For. + // The proxy appends the address it saw, so only the last entry counts: + // earlier ones come from the client and can be anything. if ip := net.ParseIP(host); ip != nil && ip.IsLoopback() { - if xff := r.Header.Get("X-Forwarded-For"); xff != "" { - return strings.TrimSpace(strings.Split(xff, ",")[0]) + if xff := r.Header.Values("X-Forwarded-For"); len(xff) > 0 { + list := strings.Split(xff[len(xff)-1], ",") + if last := strings.TrimSpace(list[len(list)-1]); net.ParseIP(last) != nil { + return last + } } } return host diff --git a/config.go b/config.go index 0e80f98..38e50b4 100644 --- a/config.go +++ b/config.go @@ -11,6 +11,7 @@ import ( "path/filepath" "regexp" "slices" + "strconv" "strings" "sync" "syscall" @@ -76,8 +77,29 @@ const ( minLogFiles, maxLogFiles = 1, 100 minHourlyHours, maxHourlyHrs = 24, 24 * 31 minDailyDays, maxDailyDays = 7, 3660 + minSessionHours = 1 + maxSessionHours = 30 * 24 ) +// validateListen checks a listen address like ":443" or "192.0.2.1:443". +// Empty is allowed when optional (the HTTP listener is then off). +func validateListen(addr, field string, optional bool) error { + if addr == "" && optional { + return nil + } + host, port, err := net.SplitHostPort(addr) + if err != nil { + return fmt.Errorf("%s %q must look like :443 or 192.0.2.1:443", field, addr) + } + if n, err := strconv.Atoi(port); err != nil || n < 1 || n > 65535 { + return fmt.Errorf("%s %q: the port must be 1–65535", field, addr) + } + if host != "" && host != "localhost" && checkEndpoint(host) != nil { + return fmt.Errorf("%s %q: %q is not an IP address or host name", field, addr, host) + } + return nil +} + type WebConfig struct { Listen string `json:"listen"` // HTTPS (or HTTP when tls.mode is "off") listen address HTTPListen string `json:"httpListen"` // plain HTTP for ACME http-01 and redirects; "" disables @@ -364,7 +386,9 @@ func (c *Config) validate() error { if v6.Masked() != v6 { return fmt.Errorf("IPv6 network must be the network address, e.g. %s", v6.Masked()) } - if s.Endpoint != "" && strings.ContainsAny(s.Endpoint, " /:") && net.ParseIP(s.Endpoint) == nil { + // The endpoint is written into client configs as is, so it must be a + // plain host name or IP: anything else could add lines to them. + if s.Endpoint != "" && checkEndpoint(s.Endpoint) != nil { return errors.New("endpoint must be a host name or IP address without port") } if err := validateHostList(s.ClientDefaults.DNS, "DNS", false); err != nil { @@ -397,6 +421,15 @@ func (c *Config) validate() error { if _, ok := updateSources[c.Updates.Source]; !ok { return fmt.Errorf("update source must be gitea or github") } + if err := validateListen(c.Web.Listen, "listen address", false); err != nil { + return err + } + if err := validateListen(c.Web.HTTPListen, "HTTP listen address", true); err != nil { + return err + } + if h := c.Web.SessionHours; h < minSessionHours || h > maxSessionHours { + return fmt.Errorf("session length must be %d–%d hours", minSessionHours, maxSessionHours) + } switch c.Web.TLS.Mode { case "acme": if c.Web.TLS.Domain == "" { @@ -586,6 +619,12 @@ func (s *Store) Update(fn func(c *Config) error) error { s.mu.Unlock() return &userError{err.Error()} } + // With no user left (applyDefaults then adds an "admin" without a + // password), nobody could sign in until someone ran "passwd" on the server. + if old.passwordSet() && !next.passwordSet() { + s.mu.Unlock() + return &userError{"this would leave no user with a password, and nobody could sign in"} + } if err := writeFileAtomic(s.path, next, 0o600); err != nil { s.mu.Unlock() return err diff --git a/kernel.go b/kernel.go index 09ce6f8..664231c 100644 --- a/kernel.go +++ b/kernel.go @@ -4,6 +4,7 @@ import ( "log/slog" "net/netip" "os" + "slices" "strings" "sync" "time" @@ -40,6 +41,28 @@ type Kernel interface { Close() error } +// lanBlock picks, from the networks on the uplinks, the ones peers must not +// reach while LAN access is off: private IPv4 networks, and IPv6 networks +// except link-local, since a home LAN uses global IPv6 addresses. IPv6 +// prefixes shorter than /48 are left out: they are no LAN. +func lanBlock(nets []netip.Prefix) []netip.Prefix { + var out []netip.Prefix + for _, p := range nets { + a := p.Addr().Unmap() + p = netip.PrefixFrom(a, min(p.Bits(), a.BitLen())).Masked() + switch { + case a.Is4() && !a.IsPrivate(): + continue + case a.Is6() && (a.IsLinkLocalUnicast() || a.IsLoopback() || p.Bits() < 48): + continue + } + if !slices.Contains(out, p) { + out = append(out, p) + } + } + return out +} + // readSysctl returns the trimmed content of a /proc/sys file, or "". func readSysctl(path string) string { b, err := os.ReadFile(path) @@ -56,6 +79,10 @@ type Reconciler struct { store *Store trigger chan struct{} + // applyMu runs one apply at a time. Each reads the config once it holds + // the lock, so the last apply always uses the newest config. + applyMu sync.Mutex + mu sync.Mutex lastErr error lastApply time.Time @@ -76,6 +103,8 @@ func (r *Reconciler) Kick() { // ApplyNow applies synchronously and returns the result, so an API call can // report kernel errors to the user. func (r *Reconciler) ApplyNow() error { + r.applyMu.Lock() + defer r.applyMu.Unlock() err := r.kernel.Apply(r.store.Get()) r.mu.Lock() r.lastErr, r.lastApply = err, time.Now() diff --git a/kernel_linux.go b/kernel_linux.go index 008c01c..8d96d55 100644 --- a/kernel_linux.go +++ b/kernel_linux.go @@ -217,7 +217,8 @@ func (k *linuxKernel) Apply(c *Config) error { if c.Server.IPv6Enabled { _ = os.WriteFile("/proc/sys/net/ipv6/conf/all/forwarding", []byte("1"), 0o644) } - return applyFirewall(c, k.Uplink(c, false), k.Uplink(c, true), lanNetworks(k.Uplink(c, false))) + up4, up6 := k.Uplink(c, false), k.Uplink(c, true) + return applyFirewall(c, up4, up6, lanNetworks(up4, up6)) } func (k *linuxKernel) Sample(iface string) ([]PeerSample, error) { @@ -264,27 +265,27 @@ func (k *linuxKernel) Uplink(c *Config, v6 bool) string { return l.Attrs().Name } -// lanNetworks returns the private IPv4 networks on the uplink, used to block -// peers from the server's LAN when LAN access is off. -func lanNetworks(uplink string) []netip.Prefix { - if uplink == "" { - return nil - } - l, err := netlink.LinkByName(uplink) - if err != nil { - return nil - } - addrs, _ := netlink.AddrList(l, netlink.FAMILY_V4) - var out []netip.Prefix - for _, a := range addrs { - if !a.IP.IsPrivate() { +// lanNetworks returns the LAN networks on the IPv4 and IPv6 uplinks (see +// lanBlock), used to block peers from the server's LAN when LAN access is off. +func lanNetworks(uplinks ...string) []netip.Prefix { + var nets []netip.Prefix + for i, uplink := range uplinks { + if uplink == "" || slices.Contains(uplinks[:i], uplink) { continue } - ones, _ := a.Mask.Size() - ip, _ := netip.AddrFromSlice(a.IP.To4()) - out = append(out, netip.PrefixFrom(ip, ones).Masked()) + l, err := netlink.LinkByName(uplink) + if err != nil { + continue + } + addrs, _ := netlink.AddrList(l, netlink.FAMILY_ALL) + for _, a := range addrs { + ones, _ := a.Mask.Size() + if ip, ok := netip.AddrFromSlice(a.IP); ok { + nets = append(nets, netip.PrefixFrom(ip.Unmap(), ones)) + } + } } - return out + return lanBlock(nets) } // publicAddr reports the uplink's address for the health check: the first diff --git a/main.go b/main.go index ddb5b31..db6c578 100644 --- a/main.go +++ b/main.go @@ -224,6 +224,10 @@ func run(configPath string) error { app := &App{ store: store, kernel: kernel, recon: recon, stats: stats, speeds: speeds, auth: auth, tls: webTLS, logPath: logPath, logw: logw, geo: geo, updates: newUpdater(cfg.Updates), started: time.Now(), shutdown: shutdown, + webAddrs: []string{cfg.Web.Listen}, + } + if cfg.Web.HTTPListen != "" && cfg.Web.TLS.Mode != "off" { + app.webAddrs = append(app.webAddrs, cfg.Web.HTTPListen) } var wg sync.WaitGroup diff --git a/main_test.go b/main_test.go index 910094f..e58e790 100644 --- a/main_test.go +++ b/main_test.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "io" + "net" "net/http" "net/http/cookiejar" "net/http/httptest" @@ -79,6 +80,15 @@ func TestValidate(t *testing.T) { "bad port": func(c *Config) { c.Server.ListenPort = 70000 }, "unmasked net": func(c *Config) { c.Server.IPv4 = "10.84.12.5/24" }, "update source": func(c *Config) { c.Updates.Source = "sourceforge" }, + // The endpoint goes into client configs: no extra lines. + "endpoint newline": func(c *Config) { c.Server.Endpoint = "vpn.example.net\n[Interface]\nPreUp=id;#" }, + "endpoint tab": func(c *Config) { c.Server.Endpoint = "vpn.example.net\tx" }, + "endpoint port": func(c *Config) { c.Server.Endpoint = "vpn.example.net:51820" }, + "listen": func(c *Config) { c.Web.Listen = "not-an-address" }, + "listen port": func(c *Config) { c.Web.Listen = ":70000" }, + "http listen": func(c *Config) { c.Web.HTTPListen = "80" }, + "session hours": func(c *Config) { c.Web.SessionHours = -1 }, + "session too long": func(c *Config) { c.Web.SessionHours = 100000 }, } { cc := c.clone() mutate(cc) @@ -86,6 +96,20 @@ func TestValidate(t *testing.T) { t.Errorf("%s: expected an error", name) } } + for _, ep := range []string{"vpn.example.net", "203.0.113.7", "2001:db8::1"} { + cc := c.clone() + cc.Server.Endpoint = ep + if err := cc.validate(); err != nil { + t.Errorf("endpoint %q rejected: %v", ep, err) + } + } + for _, l := range []string{":443", "0.0.0.0:8443", "[::]:443", "localhost:8080"} { + cc := c.clone() + cc.Web.Listen = l + if err := cc.validate(); err != nil { + t.Errorf("listen %q rejected: %v", l, err) + } + } } func TestClientConfig(t *testing.T) { @@ -1339,3 +1363,247 @@ func TestSpeeds(t *testing.T) { t.Fatalf("kept %d points, want %d", n, speedPoints) } } + +// signedInApp starts the API with a signed-in admin and returns the app +// and a call function. +func signedInApp(t *testing.T) (*App, func(method, path string, body any, want int) map[string]any) { + t.Helper() + dir := t.TempDir() + store, err := openStore(filepath.Join(dir, "config.json")) + if err != nil { + t.Fatal(err) + } + hash, _ := hashPassword("a long test password") + _ = store.Update(func(c *Config) error { c.Users[0].PasswordHash = hash; return nil }) + k := &fakeKernel{} + st, _ := openStats(filepath.Join(dir, "stats.json"), store, k) + app := &App{store: store, kernel: k, recon: newReconciler(k, store), stats: st, auth: newAuth(store), + tls: &webTLS{}, logPath: filepath.Join(dir, "log.jsonl"), started: time.Now(), shutdown: func() {}} + srv := httptest.NewServer(app.routes()) + t.Cleanup(srv.Close) + jar, _ := cookiejar.New(nil) + cl := &http.Client{Jar: jar} + call := func(method, path string, body any, want int) map[string]any { + t.Helper() + var rd io.Reader + if body != nil { + b, _ := json.Marshal(body) + rd = bytes.NewReader(b) + } + req, _ := http.NewRequest(method, srv.URL+"/api/v1"+path, rd) + req.Header.Set("Content-Type", "application/json") + resp, err := cl.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + var out map[string]any + _ = json.NewDecoder(resp.Body).Decode(&out) + if resp.StatusCode != want { + t.Fatalf("%s %s: status %d, want %d: %v", method, path, resp.StatusCode, want, out) + } + return out + } + call("POST", "/auth/login", map[string]string{"username": "admin", "password": "a long test password"}, 200) + return app, call +} + +// Web settings the service could not start with are refused before they +// are saved: a restart would otherwise take the web interface and the API +// down for good. +func TestWebSettingsCheck(t *testing.T) { + app, call := signedInApp(t) + web := func(change func(w *WebConfig)) map[string]any { + w := app.store.Get().Web + w.HTTPListen = "" + change(&w) + return map[string]any{"web": w} + } + call("PATCH", "/settings", web(func(w *WebConfig) { w.Listen = "not-an-address" }), 400) + call("PATCH", "/settings", web(func(w *WebConfig) { w.SessionHours = -1 }), 400) + call("PATCH", "/settings", web(func(w *WebConfig) { + w.TLS = TLSConfig{Mode: "files", CertFile: "/nonexistent/cert.pem", KeyFile: "/nonexistent/key.pem"} + }), 400) + + // A port another program holds is refused; a free one is saved. + busy, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer busy.Close() + call("PATCH", "/settings", web(func(w *WebConfig) { w.Listen = busy.Addr().String() }), 400) + free, _ := net.Listen("tcp", "127.0.0.1:0") + addr := free.Addr().String() + free.Close() + call("PATCH", "/settings", web(func(w *WebConfig) { w.Listen = addr; w.TLS = TLSConfig{Mode: "off"} }), 200) + if app.store.Get().Web.Listen != addr { + t.Fatal("valid listen address not saved") + } + + // The address the service listens on now is in use by itself: fine. + app.webAddrs = []string{busy.Addr().String()} + call("PATCH", "/settings", web(func(w *WebConfig) { w.Listen = busy.Addr().String() }), 200) + + // Restore runs the same check. + backup := app.store.Get() + backup.Web.Listen = "not-an-address" + call("POST", "/restore", backup, 400) +} + +func TestRemoteIP(t *testing.T) { + for _, c := range []struct { + remote string + xff []string + want string + }{ + {"203.0.113.5:1234", nil, "203.0.113.5"}, + {"203.0.113.5:1234", []string{"198.51.100.1"}, "203.0.113.5"}, // not from a local proxy + {"127.0.0.1:1234", []string{"198.51.100.1"}, "198.51.100.1"}, + // The client sent its own header; the proxy appended the real address. + {"127.0.0.1:1234", []string{"1.2.3.4, 198.51.100.1"}, "198.51.100.1"}, + {"127.0.0.1:1234", []string{"1.2.3.4", "198.51.100.1"}, "198.51.100.1"}, + {"127.0.0.1:1234", []string{"garbage"}, "127.0.0.1"}, + } { + r := httptest.NewRequest("GET", "/", nil) + r.RemoteAddr = c.remote + for _, v := range c.xff { + r.Header.Add("X-Forwarded-For", v) + } + if got := remoteIP(r); got != c.want { + t.Errorf("%s %v: got %s, want %s", c.remote, c.xff, got, c.want) + } + } +} + +// Anyone can start a passkey sign-in, so pending ones are capped per +// address and in total. +func TestPasskeyLoginCap(t *testing.T) { + a := newAuth(nil) + start := func(id, ip string, expires time.Time) bool { + a.mu.Lock() + defer a.mu.Unlock() + return a.addPasskeyLoginLocked(id, &ceremony{ip: lockKey(ip), expires: expires}) + } + later := time.Now().Add(ticketTTL) + for i := range maxPasskeyLoginsPerIP { + if !start(fmt.Sprint("a", i), "198.51.100.1", later) { + t.Fatalf("sign-in %d refused", i) + } + } + if start("a-more", "198.51.100.1", later) { + t.Fatal("too many sign-ins from one address accepted") + } + if !start("b0", "198.51.100.2", later) { + t.Fatal("another address refused") + } + // Expired ones make room again. + a.mfa.logins = map[string]*ceremony{} + start("old", "198.51.100.3", time.Now().Add(-time.Second)) + if !start("new", "198.51.100.3", later) || len(a.mfa.logins) != 1 { + t.Fatalf("expired sign-in not dropped: %d pending", len(a.mfa.logins)) + } + // In total, the oldest makes room. + a.mfa.logins = map[string]*ceremony{} + for i := range maxPasskeyLogins { + start(fmt.Sprint("c", i), fmt.Sprintf("10.0.%d.%d", i/250, i%250), later.Add(time.Duration(i)*time.Millisecond)) + } + start("last", "192.0.2.1", later.Add(time.Hour)) + if _, ok := a.mfa.logins["c0"]; ok || len(a.mfa.logins) != maxPasskeyLogins { + t.Fatalf("cap not kept: %d pending, oldest kept %v", len(a.mfa.logins), ok) + } +} + +func TestLanBlock(t *testing.T) { + got := lanBlock([]netip.Prefix{ + netip.MustParsePrefix("192.168.1.20/24"), + netip.MustParsePrefix("203.0.113.9/24"), // public IPv4: not a LAN + netip.MustParsePrefix("2001:db8:1:2::20/64"), + netip.MustParsePrefix("fd00:1:2:3::20/64"), + netip.MustParsePrefix("fe80::1/64"), + netip.MustParsePrefix("2001:db8::1/32"), // no LAN + netip.MustParsePrefix("192.168.1.30/24"), // same network twice + }) + want := []netip.Prefix{ + netip.MustParsePrefix("192.168.1.0/24"), + netip.MustParsePrefix("2001:db8:1:2::/64"), + netip.MustParsePrefix("fd00:1:2:3::/64"), + } + if !slices.Equal(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +// A change that would leave no user with a password is refused: restoring +// a backup without users, or the last users deleting each other. +func TestNoUserLeftWithPassword(t *testing.T) { + app, call := signedInApp(t) + if err := app.store.Update(func(c *Config) error { c.Users = nil; return nil }); err == nil { + t.Fatal("removing every user was accepted") + } + if !app.store.Get().passwordSet() { + t.Fatal("password lost") + } + + backup := app.store.Get() + backup.Users, backup.APITokens = nil, nil + call("POST", "/restore", backup, 400) + backup = app.store.Get() + backup.Version = configVersion + 1 + call("POST", "/restore", backup, 400) + call("POST", "/restore", app.store.Get(), 200) + if !app.store.Get().passwordSet() { + t.Fatal("password lost") + } +} + +// slowKernel records the configs it applied; the first apply takes a while. +type slowKernel struct { + fakeKernel + mu sync.Mutex + calls int + applied []string // peer names, per apply +} + +func (k *slowKernel) Apply(c *Config) error { + k.mu.Lock() + k.calls++ + first := k.calls == 1 + k.mu.Unlock() + if first { + time.Sleep(200 * time.Millisecond) + } + var names []string + for _, p := range c.Peers { + names = append(names, p.Name) + } + k.mu.Lock() + k.applied = append(k.applied, strings.Join(names, ",")) + k.mu.Unlock() + return nil +} + +// Applies run one at a time, so a slow apply of an older config cannot +// finish after the newest one and undo it in the kernel. +func TestApplyOrder(t *testing.T) { + store, err := openStore(filepath.Join(t.TempDir(), "config.json")) + if err != nil { + t.Fatal(err) + } + k := &slowKernel{} + r := newReconciler(k, store) + var wg sync.WaitGroup + wg.Add(1) + go func() { defer wg.Done(); _ = r.ApplyNow() }() // the old config, slowly + time.Sleep(50 * time.Millisecond) + if err := store.Update(func(c *Config) error { + c.Peers = append(c.Peers, Peer{ID: newID(), Name: "phone", IPv4: serverIPv4(netip.MustParsePrefix(c.Server.IPv4)).Next().String()}) + return nil + }); err != nil { + t.Fatal(err) + } + _ = r.ApplyNow() + wg.Wait() + if last := k.applied[len(k.applied)-1]; last != "phone" { + t.Fatalf("the kernel ended with %q, not the newest config; applies: %q", last, k.applied) + } +} diff --git a/mfa.go b/mfa.go index 3dad0a8..39db703 100644 --- a/mfa.go +++ b/mfa.go @@ -217,6 +217,7 @@ type ticket struct { type ceremony struct { userID string // "" for a passkey sign-in + ip string // lockKey of who started a passkey sign-in data *webauthn.SessionData expires time.Time } @@ -551,12 +552,52 @@ func (a *App) loginPasskeyBegin(w http.ResponseWriter, r *http.Request) { return } id := randomString(24) + ip := remoteIP(r) a.auth.mu.Lock() - a.auth.mfa.logins[id] = &ceremony{data: data, expires: time.Now().Add(ticketTTL)} + ok := a.auth.addPasskeyLoginLocked(id, &ceremony{data: data, ip: lockKey(ip), expires: time.Now().Add(ticketTTL)}) a.auth.mu.Unlock() + if !ok { + writeJSON(w, http.StatusTooManyRequests, map[string]string{"error": errBusy.Error()}) + return + } writeJSON(w, http.StatusOK, map[string]any{"id": id, "options": opts}) } +// Anyone can start a passkey sign-in, so the pending ones are capped: per +// address, and in total, where the oldest makes room. +const ( + maxPasskeyLogins = 1000 + maxPasskeyLoginsPerIP = 10 +) + +// addPasskeyLoginLocked stores a started passkey sign-in, or reports false +// when its address has too many pending. a.mu must be held. +func (a *Auth) addPasskeyLoginLocked(id string, c *ceremony) bool { + now := time.Now() + var fromIP int + var oldestID string + for k, x := range a.mfa.logins { + if now.After(x.expires) { + delete(a.mfa.logins, k) + continue + } + if x.ip == c.ip { + fromIP++ + } + if oldestID == "" || x.expires.Before(a.mfa.logins[oldestID].expires) { + oldestID = k + } + } + if fromIP >= maxPasskeyLoginsPerIP { + return false + } + if len(a.mfa.logins) >= maxPasskeyLogins { + delete(a.mfa.logins, oldestID) + } + a.mfa.logins[id] = c + return true +} + func (a *App) loginPasskeyFinish(w http.ResponseWriter, r *http.Request) { id := r.URL.Query().Get("id") ip := remoteIP(r)