mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 14:06:58 +02:00
Compare commits
50 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8e320d5384 | |||
| 1617897043 | |||
| f8775bb6ca | |||
| e5d17af0de | |||
| 15f0f0d5d0 | |||
| 7902ce674e | |||
| c2fbe215e6 | |||
| 94ac6db4ca | |||
| a60350e34e | |||
| 58f3b6fda7 | |||
| a99699e370 | |||
| 3615a79b8b | |||
| 147c202c27 | |||
| e290a6892f | |||
| 6c3972f464 | |||
| 861d3aabd7 | |||
| 86733864fe | |||
| ab736e4c6b | |||
| 5ecdd4eaa9 | |||
| 1b84bd0050 | |||
| 384610f81a | |||
| f671406b82 | |||
| 448f06a378 | |||
| 8e607a91f4 | |||
| c72a37c16f | |||
| 610fcdb9bf | |||
| bb3c70da2e | |||
| 2f50b3c54f | |||
| 0824035906 | |||
| 510a8912a9 | |||
| 0496ef101e | |||
| ae9de47dd9 | |||
| 4eb86afa54 | |||
| f36db374ac | |||
| 7ac51c1af2 | |||
| dabce8a1b4 | |||
| 6b78e9cdb3 | |||
| b445d14ddb | |||
| 6606124bf9 | |||
| 05405bc261 | |||
| b033267d6e | |||
| 659d7fece6 | |||
| f2aef0d6eb | |||
| a2b9747b0f | |||
| 0e593ad582 | |||
| 28ecfcbc03 | |||
| e71059a410 | |||
| aec7f5f865 | |||
| 6d8e939648 | |||
| 326fc8758d |
@@ -25,9 +25,9 @@ inputs:
|
||||
required: false
|
||||
default: "code-signer"
|
||||
key-prefix:
|
||||
description: "S3 key prefix the caller is authorized to write under"
|
||||
description: "S3 key prefix to write under; defaults to code-signing/<owner>/<repo> of the calling repo"
|
||||
required: false
|
||||
default: "code-signing/slackhq/nebula"
|
||||
default: ""
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
@@ -57,6 +57,9 @@ runs:
|
||||
KEY_PREFIX: ${{ inputs.key-prefix }}
|
||||
run: |
|
||||
set -eu
|
||||
# Default the prefix to this repo so the S3 key attributes the sign correctly.
|
||||
# nebula-nightly runs this same action but writes under its own repo's prefix.
|
||||
KEY_PREFIX="${KEY_PREFIX:-code-signing/$GITHUB_REPOSITORY}"
|
||||
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
|
||||
|
||||
find "$SIGN_PATH" -name '*.exe' -print | while read -r path
|
||||
|
||||
@@ -12,9 +12,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -38,9 +38,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -78,9 +78,9 @@ jobs:
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Import certificates
|
||||
|
||||
@@ -32,9 +32,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -64,9 +64,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
@@ -90,9 +90,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||
|
||||
@@ -20,9 +20,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: build
|
||||
@@ -60,4 +60,12 @@ jobs:
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: NAME="smoke-p256" ./smoke.sh
|
||||
|
||||
- name: setup docker image for multiport
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: NAME="smoke-multiport" MULTIPORT_TX=true MULTIPORT_RX=true MULTIPORT_HANDSHAKE=true ./build.sh
|
||||
|
||||
- name: run smoke
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: NAME="smoke-multiport" ./smoke.sh
|
||||
|
||||
timeout-minutes: 10
|
||||
|
||||
@@ -48,6 +48,10 @@ listen:
|
||||
|
||||
tun:
|
||||
dev: ${TUN_DEV:-tun0}
|
||||
multiport:
|
||||
tx_enabled: ${MULTIPORT_TX:-false}
|
||||
rx_enabled: ${MULTIPORT_RX:-false}
|
||||
tx_handshake: ${MULTIPORT_HANDSHAKE:-false}
|
||||
|
||||
firewall:
|
||||
inbound_action: reject
|
||||
|
||||
@@ -20,9 +20,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Install goimports
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
- name: golangci-lint
|
||||
uses: golangci/golangci-lint-action@v9
|
||||
with:
|
||||
version: v2.5
|
||||
version: v2.12
|
||||
|
||||
test:
|
||||
name: Test ${{ matrix.name }}
|
||||
@@ -80,9 +80,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build
|
||||
@@ -125,9 +125,9 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.25'
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Build ${{ matrix.name }}
|
||||
|
||||
@@ -7,6 +7,88 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [1.11.0] - 2026-07-23
|
||||
|
||||
See the [v1.11.0](https://github.com/slackhq/nebula/milestone/25?closed=1) milestone for a complete list of changes.
|
||||
|
||||
### Breaking
|
||||
|
||||
- Logging has switched from logrus to Go's structured `slog`. Log output changes: levels are upper case
|
||||
(`level=INFO`), trace prints as `level=DEBUG-4`, timestamps are always RFC3339Nano and `logging.timestamp_format`
|
||||
is ignored, and some messages were reworded. Review any log parsing before upgrading. This is also an API break
|
||||
for embedders, as constructors now take a `*slog.Logger`. (#1672, #1734, #1621)
|
||||
- `firewall.inbound_action` and `firewall.outbound_action` (used to set reject vs. drop policy) were each being
|
||||
applied to the opposite direction, that is now corrected. This only affects how blocked packets are answered, not
|
||||
which packets the firewall allows or denies. If you set either of these you are getting the behavior of the other
|
||||
one today and likely want to swap them before upgrading. (#1798)
|
||||
- On Windows, Nebula now installs WFP PERMIT filters for the nebula adapter and the listener port by default. WFP
|
||||
sits below Windows Defender Firewall, so any WDF inbound rules you rely on for either will no longer apply. Set
|
||||
`tun.windows_bypass_wdf` and `listen.windows_bypass_wdf` to false to leave WDF in charge. (#1710)
|
||||
- On Windows, the nebula device is now set to the `private` network category instead of whatever Windows decided,
|
||||
which is usually `Public`. This makes the host firewall less restrictive on the overlay. Set
|
||||
`tun.network_category` to `unset` to keep the old behavior. (#1710)
|
||||
- Reject packets for non-TCP now use ICMP code 13, communication administratively prohibited, instead of code 3,
|
||||
port unreachable. Anything keying off the old code needs updating. (#1766, #1768)
|
||||
- The SSH debug server's profiling commands are now confined to `sshd.sandbox_dir`, which defaults to
|
||||
`$TMP/nebula-debug`. Relative paths resolve inside it and absolute paths outside it are rejected, so anything
|
||||
scripting `start-cpu-profile`, `save-heap-profile`, or `save-mutex-profile` with a path elsewhere needs the
|
||||
directory set. The directory is not created for you. (#1622)
|
||||
|
||||
### Added
|
||||
|
||||
- Sign the Windows release binaries. (#1718)
|
||||
- Generate IPv6 reject packets, matching the existing IPv4 behavior. (#1766, #1767, #1768)
|
||||
- Accept `-` in `nebula-cert` to read from stdin or write to stdout. (#1714)
|
||||
- Search for both `config.yml` and `config.yaml` in service and command line modes. (#1717)
|
||||
- Add version labels to the Docker/OCI images. (#1772)
|
||||
- Rebind the listener and re-query lighthouses on macOS when the underlay network changes, so devices moving
|
||||
between wifi and wired or between networks recover without waiting for dead tunnel detection. Controlled by
|
||||
`listen.rebind_on_network_change` (default `true`, not reloadable). (#1816)
|
||||
|
||||
### Changed
|
||||
|
||||
- Reload the firewall when the unsafe networks in the certificate change. (#1719)
|
||||
- Reconfigure, start, and stop the stats listener on a config reload instead of requiring a restart. (#1670)
|
||||
- Update a static host's addresses when they change on reload. (#1713)
|
||||
- Don't require a port on ICMP firewall rules. (#1609)
|
||||
- Connection track ICMP traffic. (#1602)
|
||||
- Return `NODATA` instead of `NXDOMAIN` from the DNS server for a name that exists but has no record of the
|
||||
requested type, so clients that query `AAAA` first (busybox/Alpine) fall through to `A`. (#1668)
|
||||
- Record the local host's details in the DNS server. (#1716)
|
||||
- Install Windows unsafe routes as link routes. (#1709)
|
||||
- Reduce relay handshake log spam, and only log a handshake send error at error level when the remote list
|
||||
changes. (#1733, #1765, #1810)
|
||||
- Start, stop, and reload subsystems (DNS, stats, conntrack, ssh, punchy) cleanly without leaking goroutines. (#1640, #1654, #1661, #1667, #1669, #1708, #1806, #1815)
|
||||
- `Control` is now safe to stop and wait on from any lifecycle state, and a new `Control.Wait` blocks until nebula
|
||||
has fully stopped and returns the first fatal reader error. Failed starts release the udp sockets and tun fd
|
||||
instead of leaking them. (#1794)
|
||||
- Trigger an immediate lighthouse update when reconnecting to or adding a lighthouse instead of waiting for the next update tick. (#1645)
|
||||
- Bring the Darwin and OpenBSD tun implementations in line with the other BSDs. (#1703)
|
||||
- Update to build against go v1.26. (#1818)
|
||||
- Various dependency updates. (#1586, #1587, #1604, #1617, #1618, #1627, #1628, #1629, #1652, #1664, #1665, #1697, #1721, #1732, #1742, #1743, #1750, #1763, #1771, #1782, #1800, #1807)
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix a data race on a host's remote address that could send packets to the wrong address during a roam. (#1773)
|
||||
- Fix tunnels that could permanently escape connection manager monitoring. (#1752)
|
||||
- Fix a crash when reloading the SSH server's trusted keys. (#1787)
|
||||
- Fix hostmap corruption when a host has multiple overlay addresses. Each address now gets its own list instead of
|
||||
a single shared chain, which also fixes two latent bugs on the add and makePrimary paths. (#1788, #1790)
|
||||
- Apply `remote_allow_list` IPv4 rules to 4-in-6 mapped addresses. (#1786)
|
||||
- Don't panic in the DNS server on a short or empty query name. (#1635)
|
||||
- Advance the replay window on relayed packets so a relay drops replayed frames instead of re-forwarding them. (#1751)
|
||||
- Fix a race in relay state handling. (#1753)
|
||||
- Lock replay window updates so concurrent readers can't corrupt it. (#1802)
|
||||
- Reject malformed handshakes more reliably, including invalid ed25519 key lengths. (#1601, #1756)
|
||||
- Properly handle `closetunnel` packets. (#1638)
|
||||
- Fix an IPv6 extension-header length overflow that could make the firewall parse the wrong protocol and ports. (#1789)
|
||||
- Fix relay re-establishment when a handshake arrives over a relay entry that a one-sided teardown left
|
||||
`Disestablished`, which silently dropped every send until dead tunnel detection forced a re-handshake. (#1805)
|
||||
- Don't build new relay state on a tunnel that was just discarded. (#1796)
|
||||
- Don't delete the wrong pending hostinfo in the handshake manager. (#1811)
|
||||
- Don't call the packet reader after a UDP error on Darwin. (#1755)
|
||||
- Open the FreeBSD tun device non blocking. (#1666)
|
||||
|
||||
## [1.10.3] - 2026-02-06
|
||||
|
||||
### Security
|
||||
|
||||
@@ -268,6 +268,10 @@ smoke-relay-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||
|
||||
smoke-multiport-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && NAME="smoke-multiport" MULTIPORT_TX=true MULTIPORT_RX=true MULTIPORT_HANDSHAKE=true ./build.sh
|
||||
cd .github/workflows/smoke/ && NAME="smoke-multiport" ./smoke.sh
|
||||
|
||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||
smoke-docker-ipv6: smoke-docker
|
||||
|
||||
|
||||
@@ -53,7 +53,12 @@ func main() {
|
||||
l := logging.NewLogger(os.Stdout)
|
||||
|
||||
if *serviceFlag != "" {
|
||||
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
||||
if *configTest {
|
||||
fmt.Println("-test is not supported with -service, run the config test without -service")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if err := doService(configPath, Build, serviceFlag); err != nil {
|
||||
l.Error("Service command failed", "error", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
@@ -93,15 +98,14 @@ func main() {
|
||||
}
|
||||
|
||||
if !*configTest {
|
||||
wait, err := ctrl.Start()
|
||||
if err != nil {
|
||||
if err := ctrl.Start(); err != nil {
|
||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
go ctrl.ShutdownBlock()
|
||||
|
||||
if err := wait(); err != nil {
|
||||
if err := ctrl.Wait(); err != nil {
|
||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package main
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
"github.com/slackhq/nebula"
|
||||
@@ -14,7 +15,6 @@ var logger service.Logger
|
||||
|
||||
type program struct {
|
||||
configPath *string
|
||||
configTest *bool
|
||||
build string
|
||||
control *nebula.Control
|
||||
}
|
||||
@@ -40,22 +40,41 @@ func (p *program) Start(s service.Service) error {
|
||||
}
|
||||
})
|
||||
|
||||
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
||||
p.control, err = nebula.Main(c, false, Build, l, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
p.control.Start()
|
||||
if err := p.control.Start(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens.
|
||||
go func() {
|
||||
if err := p.control.Wait(); err != nil {
|
||||
logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err))
|
||||
os.Exit(2)
|
||||
}
|
||||
}()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *program) Stop(s service.Service) error {
|
||||
logger.Info("Nebula service stopping.")
|
||||
if p.control == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
p.control.Stop()
|
||||
|
||||
// block until nebula has fully drained before reporting stopped.
|
||||
// error logging is handled by Start.
|
||||
_ = p.control.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
||||
func doService(configPath *string, build string, serviceFlag *string) error {
|
||||
if *configPath == "" {
|
||||
p, err := config.DefaultPath()
|
||||
if err != nil {
|
||||
@@ -73,7 +92,6 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
||||
|
||||
prg := &program{
|
||||
configPath: configPath,
|
||||
configTest: configTest,
|
||||
build: build,
|
||||
}
|
||||
|
||||
@@ -105,8 +123,9 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
||||
switch *serviceFlag {
|
||||
case "run":
|
||||
if err := s.Run(); err != nil {
|
||||
// Route any errors to the system logger
|
||||
// Route any errors to the system logger and report the failure
|
||||
logger.Error(err)
|
||||
return err
|
||||
}
|
||||
default:
|
||||
if err := service.Control(s, *serviceFlag); err != nil {
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
cert_test "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as
|
||||
// a library, and on a config update dnclient calls Stop() in-process to tear the
|
||||
// old instance down before starting a new one. This boots a real nebula (real
|
||||
// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer
|
||||
// and asserts it actually closes. If the reader goroutines parked in recvmmsg
|
||||
// don't wake on Close(), Wait() blocks forever and this fails with a goroutine
|
||||
// dump instead of relying on a process signal to unstick them.
|
||||
func TestControlStopClosesOnTimer(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
dir := t.TempDir()
|
||||
|
||||
before := time.Now().Add(-time.Hour)
|
||||
after := time.Now().Add(time.Hour)
|
||||
ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil)
|
||||
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
||||
_, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil)
|
||||
|
||||
caPath := filepath.Join(dir, "ca.pem")
|
||||
certPath := filepath.Join(dir, "cert.pem")
|
||||
keyPath := filepath.Join(dir, "key.pem")
|
||||
require.NoError(t, os.WriteFile(caPath, caPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(certPath, certPEM, 0o600))
|
||||
require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600))
|
||||
|
||||
// tun disabled so no device/root is needed; routines: 2 so we exercise the
|
||||
// multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs.
|
||||
configBody := fmt.Sprintf(`
|
||||
pki:
|
||||
ca: %s
|
||||
cert: %s
|
||||
key: %s
|
||||
listen:
|
||||
host: 127.0.0.1
|
||||
port: 0
|
||||
tun:
|
||||
disabled: true
|
||||
firewall:
|
||||
outbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
inbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
routines: 2
|
||||
`, caPath, certPath, keyPath)
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600))
|
||||
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.Load(dir))
|
||||
|
||||
ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, ctrl.Start())
|
||||
|
||||
// Run like a live nebula, then close on a timer, exactly as dnclient does.
|
||||
<-time.NewTimer(5 * time.Second).C
|
||||
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun
|
||||
ctrl.Wait() // blocks until every reader goroutine has returned
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-stopped:
|
||||
t.Log("nebula closed cleanly on timer")
|
||||
case <-time.After(10 * time.Second):
|
||||
buf := make([]byte, 1<<20)
|
||||
n := runtime.Stack(buf, true)
|
||||
t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n])
|
||||
}
|
||||
}
|
||||
+2
-3
@@ -84,8 +84,7 @@ func main() {
|
||||
}
|
||||
|
||||
if !*configTest {
|
||||
wait, err := ctrl.Start()
|
||||
if err != nil {
|
||||
if err := ctrl.Start(); err != nil {
|
||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||
os.Exit(1)
|
||||
}
|
||||
@@ -93,7 +92,7 @@ func main() {
|
||||
go ctrl.ShutdownBlock()
|
||||
notifyReady(l)
|
||||
|
||||
if err := wait(); err != nil {
|
||||
if err := ctrl.Wait(); err != nil {
|
||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
package config
|
||||
|
||||
type MultiPortConfig struct {
|
||||
Tx bool
|
||||
Rx bool
|
||||
TxBasePort uint16
|
||||
TxPorts int
|
||||
TxHandshake bool
|
||||
TxHandshakeDelay int64
|
||||
}
|
||||
@@ -25,6 +25,7 @@ func newTestLighthouse() *LightHouse {
|
||||
lighthouses := []netip.Addr{}
|
||||
staticList := map[netip.Addr]struct{}{}
|
||||
|
||||
lh.localAddrsFn = func(*LocalAllowList) []netip.Addr { return nil }
|
||||
lh.lighthouses.Store(&lighthouses)
|
||||
lh.staticList.Store(&staticList)
|
||||
|
||||
|
||||
@@ -2,11 +2,13 @@ package nebula
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
)
|
||||
|
||||
@@ -20,6 +22,7 @@ type ConnectionState struct {
|
||||
initiator bool
|
||||
messageCounter atomic.Uint64
|
||||
window *Bits
|
||||
decryptLock sync.Mutex
|
||||
writeLock sync.Mutex
|
||||
}
|
||||
|
||||
@@ -54,3 +57,52 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||
func (cs *ConnectionState) Curve() cert.Curve {
|
||||
return cs.myCert.Curve()
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) Decrypt(l *slog.Logger, messageCounter uint64, out []byte, packet []byte, nb []byte) ([]byte, error) {
|
||||
var err error
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
|
||||
out, err = cs.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return nil, ErrAlreadySeen
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// VerifyRelay verifies AEAD protected (but not encrypted) relay frames. packet must be length-checked by the caller.
|
||||
func (cs *ConnectionState) VerifyRelay(l *slog.Logger, messageCounter uint64, packet []byte, nb []byte) error {
|
||||
cs.decryptLock.Lock()
|
||||
result := cs.window.Check(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
signedPayload := packet[:len(packet)-cs.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-cs.dKey.Overhead():]
|
||||
_, err := cs.dKey.DecryptDanger(nil, signedPayload, signatureValue, messageCounter, nb)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cs.decryptLock.Lock()
|
||||
result = cs.window.Update(l, messageCounter)
|
||||
cs.decryptLock.Unlock()
|
||||
if !result {
|
||||
return ErrAlreadySeen
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -55,6 +56,7 @@ func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
||||
cert.Version2, initCreds, verifier,
|
||||
func() (uint32, error) { return 1000, nil },
|
||||
true, header.HandshakeIXPSK0,
|
||||
config.MultiPortConfig{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -62,6 +64,7 @@ func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
||||
cert.Version2, respCreds, verifier,
|
||||
func() (uint32, error) { return 2000, nil },
|
||||
false, header.HandshakeIXPSK0,
|
||||
config.MultiPortConfig{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
+58
-24
@@ -53,6 +53,7 @@ type Control struct {
|
||||
statsStart func()
|
||||
dnsStart func()
|
||||
lighthouseStart func()
|
||||
networkChangeStart func(rebind func())
|
||||
connectionManagerStart func(context.Context)
|
||||
}
|
||||
|
||||
@@ -69,29 +70,29 @@ type ControlHostInfo struct {
|
||||
}
|
||||
|
||||
// Start actually runs nebula, this is a nonblocking call.
|
||||
// The returned function blocks until nebula has fully stopped and returns the
|
||||
// first fatal reader error (if any). A nil error means nebula shut down
|
||||
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
||||
// triggered the shutdown.
|
||||
func (c *Control) Start() (func() error, error) {
|
||||
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown.
|
||||
func (c *Control) Start() error {
|
||||
c.stateLock.Lock()
|
||||
defer c.stateLock.Unlock()
|
||||
switch c.state {
|
||||
case StateReady:
|
||||
//yay!
|
||||
case StateStopped, StateStopping:
|
||||
return nil, ErrAlreadyStopped
|
||||
return ErrAlreadyStopped
|
||||
case StateStarted:
|
||||
return nil, ErrAlreadyStarted
|
||||
return ErrAlreadyStarted
|
||||
default:
|
||||
return nil, ErrUnknownState
|
||||
return ErrUnknownState
|
||||
}
|
||||
|
||||
// Activate the interface
|
||||
err := c.f.activate()
|
||||
if err != nil {
|
||||
// Cancel before Close so a caller returning from Wait always observes a dead Context
|
||||
c.cancel()
|
||||
_ = c.f.Close()
|
||||
c.state = StateStopped
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
// Call all the delayed funcs that waited patiently for the interface to be created.
|
||||
@@ -104,6 +105,9 @@ func (c *Control) Start() (func() error, error) {
|
||||
if c.dnsStart != nil {
|
||||
go c.dnsStart()
|
||||
}
|
||||
if c.networkChangeStart != nil {
|
||||
go c.networkChangeStart(c.RebindUDPServer)
|
||||
}
|
||||
if c.connectionManagerStart != nil {
|
||||
go c.connectionManagerStart(c.ctx)
|
||||
}
|
||||
@@ -114,13 +118,9 @@ func (c *Control) Start() (func() error, error) {
|
||||
c.f.triggerShutdown = c.Stop
|
||||
|
||||
// Start reading packets.
|
||||
out, err := c.f.run()
|
||||
if err != nil {
|
||||
c.state = StateStopped
|
||||
return nil, err
|
||||
}
|
||||
c.f.run()
|
||||
c.state = StateStarted
|
||||
return out, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Control) State() RunState {
|
||||
@@ -133,10 +133,26 @@ func (c *Control) Context() context.Context {
|
||||
return c.ctx
|
||||
}
|
||||
|
||||
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
||||
// Stop tears nebula down, closing all tunnels and releasing everything it holds.
|
||||
// Use Wait to block until the shutdown has completed.
|
||||
// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped.
|
||||
func (c *Control) Stop() {
|
||||
c.stateLock.Lock()
|
||||
if c.state != StateStarted {
|
||||
switch c.state {
|
||||
case StateStarted:
|
||||
// Fall through to the full teardown below
|
||||
|
||||
case StateReady:
|
||||
// Never started
|
||||
c.cancel()
|
||||
c.state = StateStopped
|
||||
if err := c.f.Close(); err != nil {
|
||||
c.l.Error("Close interface failed", "error", err)
|
||||
}
|
||||
c.stateLock.Unlock()
|
||||
return
|
||||
|
||||
default:
|
||||
c.stateLock.Unlock()
|
||||
// We are stopping or stopped already
|
||||
return
|
||||
@@ -145,19 +161,26 @@ func (c *Control) Stop() {
|
||||
c.state = StateStopping
|
||||
c.stateLock.Unlock()
|
||||
|
||||
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
||||
// being created while we're shutting them all down.
|
||||
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it
|
||||
c.cancel()
|
||||
|
||||
c.CloseAllTunnels(false)
|
||||
|
||||
c.stateLock.Lock()
|
||||
c.state = StateStopped
|
||||
if err := c.f.Close(); err != nil {
|
||||
c.l.Error("Close interface failed", "error", err)
|
||||
}
|
||||
c.stateLock.Lock()
|
||||
c.state = StateStopped
|
||||
c.stateLock.Unlock()
|
||||
}
|
||||
|
||||
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
|
||||
// and returns the first fatal packet reader error if there was one.
|
||||
// It is safe to call from multiple goroutines and at any point in the lifecycle,
|
||||
// but a Wait on a Control that is never started and never stopped will block forever.
|
||||
func (c *Control) Wait() error {
|
||||
return c.f.wait()
|
||||
}
|
||||
|
||||
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
||||
func (c *Control) ShutdownBlock() {
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
@@ -170,9 +193,20 @@ func (c *Control) ShutdownBlock() {
|
||||
c.Stop()
|
||||
}
|
||||
|
||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
|
||||
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change.
|
||||
func (c *Control) RebindUDPServer() {
|
||||
_ = c.f.outside.Rebind()
|
||||
c.stateLock.Lock()
|
||||
defer c.stateLock.Unlock()
|
||||
|
||||
if c.state != StateStarted {
|
||||
return
|
||||
}
|
||||
|
||||
// A failure here means we are likely still pinned to the interface we came up on, so the rest of this is
|
||||
// unlikely to help. Say so instead of silently carrying on as if we rebound.
|
||||
if err := c.f.outside.Rebind(); err != nil {
|
||||
c.l.Error("Failed to rebind udp socket", "error", err)
|
||||
}
|
||||
|
||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||
c.f.lightHouse.SendUpdate()
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type fakeDevice struct {
|
||||
closeOnce sync.Once
|
||||
closedCh chan struct{}
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newFakeDevice() *fakeDevice {
|
||||
return &fakeDevice{closedCh: make(chan struct{})}
|
||||
}
|
||||
|
||||
// Read blocks until Close like a real tun with no traffic, then reports EOF
|
||||
// the same way a closed device does
|
||||
func (d *fakeDevice) Read(p []byte) (int, error) {
|
||||
<-d.closedCh
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
|
||||
|
||||
func (d *fakeDevice) Close() error {
|
||||
d.closeOnce.Do(func() {
|
||||
d.closed = true
|
||||
close(d.closedCh)
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *fakeDevice) Activate() error { return nil }
|
||||
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
|
||||
func (d *fakeDevice) Name() string { return "fake" }
|
||||
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
|
||||
func (d *fakeDevice) SupportsMultiqueue() bool { return false }
|
||||
func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||
return nil, errors.New("unsupported")
|
||||
}
|
||||
|
||||
// newReadyControl hand-builds the minimum Control that Main would have
|
||||
// produced right before Start, including the construction token NewInterface
|
||||
// takes so waiters block until Close releases the resources
|
||||
func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
|
||||
l := test.NewLogger()
|
||||
dev := newFakeDevice()
|
||||
conn := &fakeConn{}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
f := &Interface{
|
||||
ctx: ctx,
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
readers: make([]io.ReadWriteCloser, 1),
|
||||
routines: 1,
|
||||
hostMap: newHostMap(l),
|
||||
lightHouse: lh,
|
||||
l: l,
|
||||
}
|
||||
f.wg.Add(1)
|
||||
|
||||
return &Control{
|
||||
state: StateReady,
|
||||
f: f,
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}, dev, conn
|
||||
}
|
||||
|
||||
func TestControl_StopBeforeStart(t *testing.T) {
|
||||
c, dev, conn := newReadyControl(t)
|
||||
|
||||
// A Stop on a never started control must release everything Main acquired
|
||||
c.Stop()
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
|
||||
|
||||
// Wait must return promptly now that the resources are released
|
||||
require.NoError(t, c.Wait())
|
||||
|
||||
// A stopped control can never be started
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
|
||||
// A second Stop is a harmless no-op
|
||||
c.Stop()
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
require.NoError(t, c.Wait())
|
||||
}
|
||||
|
||||
func TestControl_WaitBlocksUntilStop(t *testing.T) {
|
||||
c, _, _ := newReadyControl(t)
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() { done <- c.Wait() }()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("Wait returned before Stop")
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
|
||||
c.Stop()
|
||||
select {
|
||||
case err := <-done:
|
||||
require.NoError(t, err)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Wait did not return after Stop")
|
||||
}
|
||||
}
|
||||
|
||||
type fakeConn struct {
|
||||
closed bool
|
||||
rebinds int
|
||||
}
|
||||
|
||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
||||
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
||||
|
||||
type multiqueueDevice struct {
|
||||
*fakeDevice
|
||||
}
|
||||
|
||||
func (d *multiqueueDevice) SupportsMultiqueue() bool { return true }
|
||||
|
||||
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
|
||||
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
|
||||
conn := &fakeConn{}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
f := &Interface{
|
||||
ctx: ctx,
|
||||
inside: dev,
|
||||
outside: conn,
|
||||
writers: []udp.Conn{conn},
|
||||
readers: make([]io.ReadWriteCloser, 2),
|
||||
routines: 2,
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
f.wg.Add(1)
|
||||
|
||||
c := &Control{
|
||||
state: StateReady,
|
||||
f: f,
|
||||
l: test.NewLogger(),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}
|
||||
|
||||
// The second reader fails to open, everything must be released
|
||||
require.Error(t, c.Start())
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||
|
||||
// And Wait must not hang on the construction token
|
||||
require.NoError(t, c.Wait())
|
||||
}
|
||||
|
||||
func TestInterface_CloseIsIdempotent(t *testing.T) {
|
||||
dev := newFakeDevice()
|
||||
f := &Interface{
|
||||
inside: dev,
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
f.wg.Add(1)
|
||||
|
||||
require.NoError(t, f.Close())
|
||||
assert.True(t, dev.closed)
|
||||
|
||||
// A second Close must not double release the wg token or the device
|
||||
require.NoError(t, f.Close())
|
||||
require.NoError(t, f.wait())
|
||||
}
|
||||
|
||||
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
|
||||
c, dev, conn := newReadyControl(t)
|
||||
|
||||
// Mirror what Start wires up, without needing real packet readers
|
||||
c.f.triggerShutdown = c.Stop
|
||||
c.state = StateStarted
|
||||
|
||||
boom := errors.New("boom")
|
||||
c.f.onFatal(boom)
|
||||
|
||||
require.ErrorIs(t, c.Wait(), boom)
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
assert.True(t, dev.closed)
|
||||
assert.True(t, conn.closed)
|
||||
|
||||
// A second fatal error must not fire the shutdown again or replace the first
|
||||
c.f.onFatal(errors.New("later"))
|
||||
require.ErrorIs(t, c.Wait(), boom)
|
||||
|
||||
// Wait stays factual, a Stop after the death does not mask the error
|
||||
c.Stop()
|
||||
require.ErrorIs(t, c.Wait(), boom)
|
||||
}
|
||||
|
||||
func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
||||
c, _, _ := newReadyControl(t)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 2; i++ {
|
||||
wg.Go(func() { c.Stop() })
|
||||
}
|
||||
wg.Go(func() { _ = c.Start() })
|
||||
wg.Go(func() {
|
||||
_ = c.Wait()
|
||||
// A returned Wait must always observe the final state, no matter how
|
||||
// the race resolved
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
})
|
||||
wg.Wait()
|
||||
|
||||
// However the race resolves, the control must end fully stopped with no
|
||||
// panic and Wait must observe the final state
|
||||
require.NoError(t, c.Wait())
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
}
|
||||
|
||||
func TestControl_StartStopLifecycle(t *testing.T) {
|
||||
c, dev, conn := newReadyControl(t)
|
||||
|
||||
require.NoError(t, c.Start())
|
||||
assert.Equal(t, StateStarted, c.State())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
||||
|
||||
// Stop must unpark the reader blocked in the device and release everything
|
||||
c.Stop()
|
||||
assert.Equal(t, StateStopped, c.State())
|
||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
||||
|
||||
// The reader drained off a closed device, that is not a fatal error
|
||||
require.NoError(t, c.Wait())
|
||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
||||
}
|
||||
|
||||
func TestControl_RebindIsGatedByState(t *testing.T) {
|
||||
c, _, conn := newReadyControl(t)
|
||||
|
||||
// A rebind before Start reaches nothing, the interface is not up
|
||||
c.RebindUDPServer()
|
||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
||||
|
||||
require.NoError(t, c.Start())
|
||||
c.RebindUDPServer()
|
||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
||||
|
||||
// A rebind racing a completed stop must not touch the closed conn
|
||||
c.Stop()
|
||||
require.NoError(t, c.Wait())
|
||||
c.RebindUDPServer()
|
||||
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
|
||||
}
|
||||
+21
-1
@@ -108,7 +108,19 @@ func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||
}
|
||||
|
||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
||||
return c.f.outside.(*udp.TesterConn).Addr
|
||||
return c.f.outside.(*udp.TesterConn).GetAddr()
|
||||
}
|
||||
|
||||
// SetUDPAddr moves this node to a new underlay address, standing in for a laptop waking up on a different
|
||||
// network. Register the new address with the router as well or nothing will route back.
|
||||
func (c *Control) SetUDPAddr(addr netip.AddrPort) {
|
||||
c.f.outside.(*udp.TesterConn).SetAddr(addr)
|
||||
}
|
||||
|
||||
// SetLocalAddrsFn replaces underlay address discovery so a test can advertise its simulated address instead of
|
||||
// whatever this machine's NICs happen to be. Call it before Start, SendUpdate reads it from the update worker.
|
||||
func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
||||
c.f.lightHouse.localAddrsFn = fn
|
||||
}
|
||||
|
||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||
@@ -125,6 +137,14 @@ func (c *Control) GetHostmap() *HostMap {
|
||||
return c.f.hostMap
|
||||
}
|
||||
|
||||
// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding
|
||||
// the hostmap read lock so tests can poll it while connection manager churns tunnels.
|
||||
func (c *Control) GetHostmapIndexCount() int {
|
||||
c.f.hostMap.RLock()
|
||||
defer c.f.hostMap.RUnlock()
|
||||
return len(c.f.hostMap.Indexes)
|
||||
}
|
||||
|
||||
func (c *Control) GetF() *Interface {
|
||||
return c.f
|
||||
}
|
||||
|
||||
+16
-7
@@ -97,8 +97,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
newAddr := getDnsServerAddr(c)
|
||||
|
||||
d.serverMu.Lock()
|
||||
running := d.server
|
||||
runningStarted := d.started
|
||||
running := d.server != nil
|
||||
sameAddr := d.addr == newAddr
|
||||
d.addr = newAddr
|
||||
d.enabled.Store(enabled)
|
||||
@@ -112,7 +111,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
}
|
||||
|
||||
if !enabled {
|
||||
if running != nil {
|
||||
if running {
|
||||
d.Stop()
|
||||
}
|
||||
// Drop any records that accumulated while enabled; a later re-enable
|
||||
@@ -121,12 +120,12 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if running == nil {
|
||||
if !running {
|
||||
// Was disabled (or never started); bring it up now.
|
||||
go d.Start()
|
||||
} else if !sameAddr {
|
||||
d.shutdownServer(running, runningStarted, "reload")
|
||||
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
||||
// Stop clears the slot before shutting down, otherwise the Start below can find the dying server and refuse
|
||||
d.Stop()
|
||||
go d.Start()
|
||||
}
|
||||
|
||||
@@ -162,7 +161,9 @@ func (d *dnsServer) Start() {
|
||||
|
||||
started := make(chan struct{})
|
||||
d.serverMu.Lock()
|
||||
if d.ctx.Err() != nil {
|
||||
// Re-check enabled under the lock, a disable that raced our check above snapshots the slot under it too.
|
||||
// Two reloads in quick succession can both spawn a Start, the loser would orphan the live listener past Stop
|
||||
if d.ctx.Err() != nil || d.server != nil || !d.enabled.Load() {
|
||||
d.serverMu.Unlock()
|
||||
return
|
||||
}
|
||||
@@ -200,6 +201,14 @@ func (d *dnsServer) Start() {
|
||||
close(started)
|
||||
}
|
||||
|
||||
// Release our slot, unless a reload already replaced us, so a dead listener can't block a future Start
|
||||
d.serverMu.Lock()
|
||||
if d.server == server {
|
||||
d.server = nil
|
||||
d.started = nil
|
||||
}
|
||||
d.serverMu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||
}
|
||||
|
||||
+206
-4
@@ -194,14 +194,51 @@ func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
// No server running yet, no addr change. Reload should not spawn anything.
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
before := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, before)
|
||||
|
||||
// Same address, so the running listener must be left alone rather than rebuilt under live queries
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
assert.True(t, ds.enabled.Load())
|
||||
assert.Nil(t, ds.server)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
after := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, before, after, "a same-address reload must not restart the listener")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// The branch the old sameAddr test was accidentally hitting: enabled with nothing running means reload starts it.
|
||||
func TestDnsServer_reload_whenNotRunning_starts(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
|
||||
// initial only records config, it never starts anything
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "the initial reload must not start a listener")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a reload with nothing running should bring DNS up")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||
@@ -427,3 +464,168 @@ func waitFor(t *testing.T, cond func() bool) {
|
||||
}
|
||||
t.Fatal("timed out waiting for condition")
|
||||
}
|
||||
|
||||
// Two reloads in quick succession, or a HUP before Control.Start, can race two Starts at the same listener.
|
||||
func TestDnsServer_Start_isIdempotent(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
first := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, first)
|
||||
|
||||
// If the second Start replaces the tracked server, Stop kills the wrong one and the port leaks
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("second Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
second := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
assert.Same(t, first, second, "a second Start must not replace the running server")
|
||||
|
||||
// The real proof, after Stop the port must actually be free
|
||||
ds.Stop()
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
// An address change must actually end up listening on the new port. Start's guard refuses when a server is already
|
||||
// installed, so reload has to clear the slot before shutting the old one down.
|
||||
func TestDnsServer_reload_addrChange_restarts(t *testing.T) {
|
||||
first := freeUDPPort(t)
|
||||
second := freeUDPPort(t)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", first, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
go ds.Start()
|
||||
waitForBind(t, ds)
|
||||
|
||||
// Cycle a few times, the failure this guards against depends on which goroutine wins serverMu
|
||||
for i := range 8 {
|
||||
want := second
|
||||
if i%2 == 1 {
|
||||
want = first
|
||||
}
|
||||
setDnsConfig(c, "127.0.0.1", want, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
srv := ds.server
|
||||
ds.serverMu.Unlock()
|
||||
require.NotNil(t, srv, "reload left DNS down instead of restarting it")
|
||||
require.Equal(t, "127.0.0.1:"+want, srv.Addr, "reload should be serving the new address")
|
||||
}
|
||||
|
||||
// Land back on second so the port assertions below are meaningful
|
||||
setDnsConfig(c, "127.0.0.1", second, true, true)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
// The old port must be released and the new one actually held
|
||||
waitFor(t, func() bool {
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+first)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = pc.Close()
|
||||
return true
|
||||
})
|
||||
_, err := net.ListenPacket("udp", "127.0.0.1:"+second)
|
||||
require.Error(t, err, "the new address should be bound by the DNS responder")
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A listener that dies on its own must release the slot, or a later same-addr reload sees it as running and no-ops.
|
||||
func TestDnsServer_Start_bindFailure_releasesSlot(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
blocker, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
ds.Start() // returns once the bind fails
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "a listener that failed to bind must not stay parked in the slot")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
// With the slot released, a reload can retry once the port frees up
|
||||
require.NoError(t, blocker.Close())
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
waitForBind(t, ds)
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.NotNil(t, ds.server, "a same-addr reload should retry after a failed bind")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
ds.Stop()
|
||||
}
|
||||
|
||||
// A disable that lands while Start is between its unlocked check and the guard must not leave a listener behind.
|
||||
func TestDnsServer_Start_refusesWhenDisabledUnderLock(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
require.True(t, ds.enabled.Load())
|
||||
|
||||
// Holding serverMu parks Start on the lock, the only way to land the disable in that window on purpose
|
||||
ds.serverMu.Lock()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
ds.serverMu.Unlock()
|
||||
t.Fatal("Start returned early, the test never exercised the window")
|
||||
case <-time.After(time.Millisecond * 100):
|
||||
}
|
||||
|
||||
// The disable reload's critical section. It sees nothing running, so it never calls Stop.
|
||||
ds.enabled.Store(false)
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("Start never returned")
|
||||
}
|
||||
|
||||
ds.serverMu.Lock()
|
||||
assert.Nil(t, ds.server, "Start must not install a listener a disable already cancelled")
|
||||
ds.serverMu.Unlock()
|
||||
|
||||
pc, err := net.ListenPacket("udp", "127.0.0.1:"+port)
|
||||
require.NoError(t, err, "an orphaned listener is still holding the port")
|
||||
_ = pc.Close()
|
||||
}
|
||||
|
||||
+98
-30
@@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) {
|
||||
|
||||
r.Log("Spin until connection manager tears down a tunnel")
|
||||
|
||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
||||
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
t.Log("Connection manager hasn't ticked yet")
|
||||
time.Sleep(time.Second)
|
||||
@@ -453,9 +453,11 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
|
||||
r.Log("Nuke my hostmap")
|
||||
myHostmap := myControl.GetHostmap()
|
||||
myHostmap.Lock()
|
||||
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||
myHostmap.Unlock()
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
||||
p = r.RouteForAllUntilTxTun(theirControl)
|
||||
@@ -465,10 +467,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
r.Log("Wait for the dead index to go away")
|
||||
start := len(theirControl.GetHostmap().Indexes)
|
||||
start := theirControl.GetHostmapIndexCount()
|
||||
for {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
if len(theirControl.GetHostmap().Indexes) < start {
|
||||
if theirControl.GetHostmapIndexCount() < start {
|
||||
break
|
||||
}
|
||||
time.Sleep(time.Second)
|
||||
@@ -504,9 +506,11 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
|
||||
r.Log("Nuke my hostmap")
|
||||
theirHostmap := theirControl.GetHostmap()
|
||||
theirHostmap.Lock()
|
||||
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
|
||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||
theirHostmap.Unlock()
|
||||
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
@@ -517,10 +521,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
r.Log("Wait for the dead index to go away")
|
||||
start := len(myControl.GetHostmap().Indexes)
|
||||
start := myControl.GetHostmapIndexCount()
|
||||
for {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
if len(myControl.GetHostmap().Indexes) < start {
|
||||
if myControl.GetHostmapIndexCount() < start {
|
||||
break
|
||||
}
|
||||
time.Sleep(time.Second)
|
||||
@@ -628,10 +632,10 @@ func TestReestablishRelays(t *testing.T) {
|
||||
r.Log("Close the tunnel")
|
||||
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||
|
||||
start := len(myControl.GetHostmap().Indexes)
|
||||
curIndexes := len(myControl.GetHostmap().Indexes)
|
||||
start := myControl.GetHostmapIndexCount()
|
||||
curIndexes := myControl.GetHostmapIndexCount()
|
||||
for curIndexes >= start {
|
||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
||||
curIndexes = myControl.GetHostmapIndexCount()
|
||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
||||
|
||||
@@ -721,6 +725,70 @@ func TestReestablishRelays(t *testing.T) {
|
||||
|
||||
}
|
||||
|
||||
func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) {
|
||||
t.Parallel()
|
||||
// If them tears down the tunnel while me keeps Established relay state, me's next
|
||||
// handshake flows through the relay with no fresh CreateRelayRequest and lands on
|
||||
// them's Disestablished terminal relay entry. them must re-establish that entry, or
|
||||
// its first transmit deletes its only relay and the tunnel is born transmit-dead:
|
||||
// them can receive but every send is silently dropped.
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||
|
||||
// Teach my how to get to the relay and that their can be reached via the relay
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
|
||||
// Build a router so we don't have to reason who gets which packet
|
||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
// Start the servers
|
||||
myControl.Start()
|
||||
relayControl.Start()
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex
|
||||
|
||||
t.Log("Close the tunnel on them only, marking their relay entry Disestablished")
|
||||
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||
|
||||
t.Log("Re-handshake from me, riding the still-Established relay state")
|
||||
myControl.ReHandshake(theirVpnIpNet[0].Addr())
|
||||
for {
|
||||
h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||
if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 {
|
||||
break
|
||||
}
|
||||
r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType {
|
||||
return router.RouteAndExit
|
||||
})
|
||||
}
|
||||
|
||||
hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
require.NotNil(t, hAtThem, "them should have completed the relayed handshake")
|
||||
require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel")
|
||||
|
||||
t.Log("Send from them to me; their only relay entry must survive the transmit")
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
require.Never(t, func() bool {
|
||||
h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
||||
return h == nil || len(h.CurrentRelaysToMe) == 0
|
||||
}, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead")
|
||||
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
||||
}
|
||||
|
||||
func TestStage1RaceRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||
@@ -819,18 +887,18 @@ func TestStage1RaceRelays2(t *testing.T) {
|
||||
|
||||
t.Log("Wait until we remove extra tunnels")
|
||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||
len(myControl.GetHostmap().Indexes),
|
||||
len(theirControl.GetHostmap().Indexes),
|
||||
len(relayControl.GetHostmap().Indexes),
|
||||
myControl.GetHostmapIndexCount(),
|
||||
theirControl.GetHostmapIndexCount(),
|
||||
relayControl.GetHostmapIndexCount(),
|
||||
)
|
||||
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||
retries := 60
|
||||
for hostInfos > 6 && retries > 0 {
|
||||
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount()
|
||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||
len(myControl.GetHostmap().Indexes),
|
||||
len(theirControl.GetHostmap().Indexes),
|
||||
len(relayControl.GetHostmap().Indexes),
|
||||
myControl.GetHostmapIndexCount(),
|
||||
theirControl.GetHostmapIndexCount(),
|
||||
relayControl.GetHostmapIndexCount(),
|
||||
)
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
t.Log("Connection manager hasn't ticked yet")
|
||||
@@ -924,24 +992,24 @@ func TestRehandshakingRelays(t *testing.T) {
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||
// We should have two hostinfos on all sides
|
||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
||||
for myControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Logf("myControl hostinfos got cleaned up!")
|
||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
||||
for theirControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Logf("theirControl hostinfos got cleaned up!")
|
||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
||||
for relayControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
@@ -1029,24 +1097,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
|
||||
// We should have two hostinfos on all sides
|
||||
for len(myControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
|
||||
for myControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Logf("myControl hostinfos got cleaned up!")
|
||||
for len(theirControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
|
||||
for theirControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Logf("theirControl hostinfos got cleaned up!")
|
||||
for len(relayControl.GetHostmap().Indexes) != 2 {
|
||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
|
||||
for relayControl.GetHostmapIndexCount() != 2 {
|
||||
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount())
|
||||
r.Log("Assert the relay tunnel still works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
r.Log("yupitdoes")
|
||||
@@ -1123,7 +1191,7 @@ func TestRehandshaking(t *testing.T) {
|
||||
theirConfig.ReloadConfigString(string(rc))
|
||||
|
||||
r.Log("Spin until there is only 1 tunnel")
|
||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
||||
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
t.Log("Connection manager hasn't ticked yet")
|
||||
time.Sleep(time.Second)
|
||||
@@ -1223,7 +1291,7 @@ func TestRehandshakingLoser(t *testing.T) {
|
||||
myConfig.ReloadConfigString(string(rc))
|
||||
|
||||
r.Log("Spin until there is only 1 tunnel")
|
||||
for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
|
||||
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 {
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
t.Log("Connection manager hasn't ticked yet")
|
||||
time.Sleep(time.Second)
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// reportedAddrs is what the lighthouse would hand a peer asking where vpnAddr is.
|
||||
func reportedAddrs(t *testing.T, lh *nebula.Control, vpnAddr netip.Addr) []netip.AddrPort {
|
||||
t.Helper()
|
||||
cm := lh.QueryLighthouse(vpnAddr)
|
||||
if cm == nil {
|
||||
return nil
|
||||
}
|
||||
var out []netip.AddrPort
|
||||
for _, c := range *cm {
|
||||
out = append(out, c.Reported...)
|
||||
out = append(out, c.Learned...)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// waitForLighthouseMsg routes until a lighthouse message lands on lh, or gives up. Reports whether one arrived.
|
||||
func waitForLighthouseMsg(t *testing.T, r *router.R, lh *nebula.Control, wait time.Duration) bool {
|
||||
t.Helper()
|
||||
h := &header.H{}
|
||||
return r.RouteForAllExitFuncOrTimeout(wait, func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
if c != lh {
|
||||
return router.KeepRouting
|
||||
}
|
||||
// Punches are a single byte and never parse, they are just not what we are after
|
||||
if err := h.Parse(p.Data); err != nil {
|
||||
return router.KeepRouting
|
||||
}
|
||||
if h.Type == header.LightHouse {
|
||||
return router.RouteAndExit
|
||||
}
|
||||
return router.KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// A laptop that changes networks has to tell the lighthouse promptly, otherwise the lighthouse keeps handing peers
|
||||
// the old address and their punches land nowhere. On a long lighthouse interval the only thing that closes that
|
||||
// window is the rebind, which on darwin the network change monitor drives. The e2e build compiles the monitor out,
|
||||
// so we call RebindUDPServer directly, which is the same thing the monitor does.
|
||||
func TestRebindSendsLighthouseUpdate(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
// 600s interval, so nothing scheduled can send an update during this test. A rebind is the only thing that can.
|
||||
myControl, _, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
// Let the startup registration finish, then clear everything it left behind
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
// Nothing should be talking to the lighthouse on its own now
|
||||
require.False(t, waitForLighthouseMsg(t, r, lhControl, time.Millisecond*200),
|
||||
"nothing should reach the lighthouse before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
|
||||
assert.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5),
|
||||
"a rebind should push an update to the lighthouse rather than waiting out the interval")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
// The other half of a rebind: every live tunnel requeries the lighthouse on its next send. That query is what makes
|
||||
// the lighthouse tell the peer to punch toward our new address, which is the part that actually revives a tunnel
|
||||
// whose remote NAT state died while we were on a different network.
|
||||
func TestRebindRequeriesPeersOnNextSend(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
lhCfg := m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
// Without this the peers advertise this machine's real addresses and then try to punch at them,
|
||||
// which the router has no route for.
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
}
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", lhCfg)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.3/24", lhCfg)
|
||||
|
||||
r := router.NewR(t, lhControl, myControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
// Point the peers at each other directly, this test is about the rebind and not about lighthouse discovery
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("initial")))
|
||||
r.RouteFor(time.Second)
|
||||
require.NotNil(t, myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false), "expected a tunnel to them")
|
||||
r.RouteFor(time.Millisecond * 300)
|
||||
|
||||
// Assert on what the peer sees rather than on lighthouse traffic. A query for them makes the lighthouse send
|
||||
// them a punch notification, which is the whole point. Our own update to the lighthouse sends them nothing,
|
||||
// so this cannot be satisfied by the update the rebind itself pushes.
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("quiet")))
|
||||
require.False(t, waitForLighthouseMsg(t, r, theirControl, time.Millisecond*300),
|
||||
"an ordinary send should not requery the lighthouse")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
r.RouteFor(time.Millisecond * 300) // let the update the rebind itself sends pass by
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("after rebind")))
|
||||
assert.True(t, waitForLighthouseMsg(t, r, theirControl, time.Second*5),
|
||||
"the first send after a rebind should requery the lighthouse, which then tells the peer to punch at us")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
}
|
||||
|
||||
// The scenario this whole thing exists for: a laptop sleeps at the office and wakes up at home on a new address.
|
||||
// Until it tells the lighthouse, the lighthouse keeps handing peers the office address, so their punches land
|
||||
// nowhere and the tunnel stays dead. On a long interval the rebind is the only thing that closes that window.
|
||||
func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{
|
||||
"lighthouse": m{"am_lighthouse": true},
|
||||
})
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
},
|
||||
"static_host_map": m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
},
|
||||
})
|
||||
|
||||
// Advertise wherever we currently are rather than this machine's real NICs, read fresh each time so a move
|
||||
// is picked up.
|
||||
myControl.SetLocalAddrsFn(func(*nebula.LocalAllowList) []netip.Addr {
|
||||
return []netip.Addr{myControl.GetUDPAddr().Addr()}
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an initial registration")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
require.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), myUdpAddr,
|
||||
"the lighthouse should know the address we started on")
|
||||
|
||||
// Wake up somewhere else
|
||||
newAddr := netip.MustParseAddrPort("10.0.0.99:4242")
|
||||
myControl.SetUDPAddr(newAddr)
|
||||
r.AddRoute(newAddr.Addr(), newAddr.Port(), myControl)
|
||||
|
||||
// Nothing has told the lighthouse, and with interval 600 nothing scheduled will
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
require.NotContains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"the lighthouse should still be handing out the old address before the rebind")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
require.True(t, waitForLighthouseMsg(t, r, lhControl, time.Second*5), "expected an update after the rebind")
|
||||
r.RouteFor(time.Millisecond * 400)
|
||||
|
||||
assert.Contains(t, reportedAddrs(t, lhControl, myVpnIpNet[0].Addr()), newAddr,
|
||||
"after the rebind the lighthouse should hand peers our new address")
|
||||
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
+129
-27
@@ -114,6 +114,28 @@ type packet struct {
|
||||
packet *udp.Packet
|
||||
tun bool // a packet pulled off a tun device
|
||||
rx bool // the packet was received by a udp device
|
||||
|
||||
// h is the nebula header, parsed once when the packet is recorded. parseErr says why there isn't one, which
|
||||
// the flow log reports rather than hiding. Punchy sends a single byte, so an unparseable packet is normal.
|
||||
h header.H
|
||||
parseErr error
|
||||
}
|
||||
|
||||
// fromAddr and toAddr are the addresses this packet actually travelled between. Reading them off the control
|
||||
// instead would misreport the whole history once a test moves a node. Tun packets are synthesized without
|
||||
// addresses, so they fall back to the control.
|
||||
func (p *packet) fromAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.From.IsValid() {
|
||||
return p.from.GetUDPAddr()
|
||||
}
|
||||
return p.packet.From
|
||||
}
|
||||
|
||||
func (p *packet) toAddr() netip.AddrPort {
|
||||
if p.tun || !p.packet.To.IsValid() {
|
||||
return p.to.GetUDPAddr()
|
||||
}
|
||||
return p.packet.To
|
||||
}
|
||||
|
||||
func (p *packet) WasReceived() {
|
||||
@@ -249,7 +271,7 @@ func (r *R) renderFlow() {
|
||||
continue
|
||||
}
|
||||
|
||||
addr := e.packet.from.GetUDPAddr()
|
||||
addr := e.packet.fromAddr()
|
||||
if _, ok := participants[addr]; ok {
|
||||
continue
|
||||
}
|
||||
@@ -268,7 +290,6 @@ func (r *R) renderFlow() {
|
||||
}
|
||||
|
||||
// Print packets
|
||||
h := &header.H{}
|
||||
for _, e := range r.flow {
|
||||
if e.packet == nil {
|
||||
//fmt.Fprintf(f, " note over %s: %s\n", strings.Join(participantsVals, ", "), e.note)
|
||||
@@ -280,21 +301,22 @@ func (r *R) renderFlow() {
|
||||
fmt.Fprintln(f, r.formatUdpPacket(p))
|
||||
|
||||
} else {
|
||||
if err := h.Parse(p.packet.Data); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
line := "--x"
|
||||
if p.rx {
|
||||
line = "->>"
|
||||
}
|
||||
|
||||
fmt.Fprintf(f,
|
||||
" %s%s%s: %s(%s), index %v, counter: %v\n",
|
||||
normalizeName(p.from.GetUDPAddr().String()),
|
||||
detail := fmt.Sprintf("%s(%s), index %v, counter: %v",
|
||||
p.h.TypeName(), p.h.SubTypeName(), p.h.RemoteIndex, p.h.MessageCounter)
|
||||
if p.parseErr != nil {
|
||||
detail = fmt.Sprintf("unparsed, %v (%d bytes)", p.parseErr, len(p.packet.Data))
|
||||
}
|
||||
|
||||
fmt.Fprintf(f, " %s%s%s: %s\n",
|
||||
normalizeName(p.fromAddr().String()),
|
||||
line,
|
||||
normalizeName(p.to.GetUDPAddr().String()),
|
||||
h.TypeName(), h.SubTypeName(), h.RemoteIndex, h.MessageCounter,
|
||||
normalizeName(p.toAddr().String()),
|
||||
detail,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -408,29 +430,34 @@ func (r *R) unlockedInjectFlow(from, to *nebula.Control, p *udp.Packet, tun bool
|
||||
|
||||
r.renderHostmaps(fmt.Sprintf("Packet %v", len(r.flow)))
|
||||
|
||||
if len(r.ignoreFlows) > 0 {
|
||||
var h header.H
|
||||
err := h.Parse(p.Data)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
var h header.H
|
||||
var parseErr error
|
||||
if !tun {
|
||||
parseErr = h.Parse(p.Data)
|
||||
}
|
||||
|
||||
for _, i := range r.ignoreFlows {
|
||||
if !tun {
|
||||
if i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
}
|
||||
} else if i.tun.HasValue && i.tun.IsTrue {
|
||||
// Decide before copying, the copy comes from a freelist and an ignored packet would never be released
|
||||
for _, i := range r.ignoreFlows {
|
||||
if tun {
|
||||
if i.tun.HasValue && i.tun.IsTrue {
|
||||
return nil
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// A packet we could not parse has no type to match against, so no rule can ignore it
|
||||
if parseErr == nil && i.messageType == h.Type && i.subType == h.Subtype {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
fp := &packet{
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
from: from,
|
||||
to: to,
|
||||
packet: p.Copy(),
|
||||
tun: tun,
|
||||
h: h,
|
||||
parseErr: parseErr,
|
||||
}
|
||||
|
||||
r.flow = append(r.flow, flowEntry{packet: fp})
|
||||
@@ -690,6 +717,81 @@ func (r *R) RouteUntilAfterMsgType(sender *nebula.Control, msgType header.Messag
|
||||
})
|
||||
}
|
||||
|
||||
// RouteFor routes everything that shows up for the given duration and then returns. Use it to let a test settle
|
||||
// deterministically rather than sleeping and hoping: a single FlushAll races a completing handshake, which queues
|
||||
// more packets right behind it.
|
||||
func (r *R) RouteFor(d time.Duration) {
|
||||
r.RouteForAllExitFuncOrTimeout(d, func(*udp.Packet, *nebula.Control) ExitType {
|
||||
return KeepRouting
|
||||
})
|
||||
}
|
||||
|
||||
// RouteForAllExitFuncOrTimeout is RouteForAllExitFunc with a deadline, reporting whether whatDo asked to exit
|
||||
// before time ran out. The unbounded version blocks forever on a quiet network, so this is what a test needs to
|
||||
// assert that something does NOT happen, or to route for a fixed settling period.
|
||||
func (r *R) RouteForAllExitFuncOrTimeout(timeout time.Duration, whatDo ExitFunc) bool {
|
||||
sc := make([]reflect.SelectCase, 0, len(r.controls)+1)
|
||||
cm := make([]*nebula.Control, 0, len(r.controls))
|
||||
|
||||
for _, c := range r.controls {
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
cm = append(cm, c)
|
||||
}
|
||||
|
||||
timer := time.NewTimer(timeout)
|
||||
defer timer.Stop()
|
||||
sc = append(sc, reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(timer.C),
|
||||
Send: reflect.Value{},
|
||||
})
|
||||
|
||||
for {
|
||||
x, rx, _ := reflect.Select(sc)
|
||||
if x == len(cm) {
|
||||
return false
|
||||
}
|
||||
|
||||
r.Lock()
|
||||
p := rx.Interface().(*udp.Packet)
|
||||
receiver := r.getControl(cm[x].GetUDPAddr(), p.To, p)
|
||||
if receiver == nil {
|
||||
r.Unlock()
|
||||
panic("Can't RouteForAllExitFuncOrTimeout for host: " + p.To.String())
|
||||
}
|
||||
|
||||
e := whatDo(p, receiver)
|
||||
switch e {
|
||||
case ExitNow:
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case RouteAndExit:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return true
|
||||
|
||||
case KeepRouting:
|
||||
fp := r.unlockedInjectFlow(cm[x], receiver, p, false)
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
|
||||
default:
|
||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||
}
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *R) RouteForAllUntilAfterMsgTypeTo(receiver *nebula.Control, msgType header.MessageType, subType header.MessageSubType) {
|
||||
h := &header.H{}
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, r *nebula.Control) ExitType {
|
||||
|
||||
+6
-6
@@ -43,8 +43,8 @@ func TestDropInactiveTunnels(t *testing.T) {
|
||||
r.Log("Go inactive and wait for the tunnels to get dropped")
|
||||
waitStart := time.Now()
|
||||
for {
|
||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||
myIndexes := myControl.GetHostmapIndexCount()
|
||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||
if myIndexes == 0 && theirIndexes == 0 {
|
||||
break
|
||||
}
|
||||
@@ -493,8 +493,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||
|
||||
waitStart := time.Now()
|
||||
for {
|
||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||
myIndexes := myControl.GetHostmapIndexCount()
|
||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||
if myIndexes == 0 && theirIndexes == 0 {
|
||||
break
|
||||
}
|
||||
@@ -548,8 +548,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||
r.Log("Injected bogus close tunnel. Let's see!")
|
||||
waitStart = time.Now()
|
||||
for {
|
||||
myIndexes := len(myControl.GetHostmap().Indexes)
|
||||
theirIndexes := len(theirControl.GetHostmap().Indexes)
|
||||
myIndexes := myControl.GetHostmapIndexCount()
|
||||
theirIndexes := theirControl.GetHostmapIndexCount()
|
||||
if myIndexes == 0 {
|
||||
t.Fatal("myIndexes should not be 0")
|
||||
}
|
||||
|
||||
@@ -146,6 +146,14 @@ listen:
|
||||
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
||||
#windows_bypass_wdf: true
|
||||
|
||||
# On macOS only
|
||||
# macOS scopes the udp socket to the interface it was created on, so moving between networks (wifi to wired,
|
||||
# office to home) leaves Nebula sending out an interface that no longer has a route. When true, Nebula watches
|
||||
# the routing socket and rebinds the listener once the change settles.
|
||||
# iOS does not use this, the host app drives the same rebind itself.
|
||||
# Default true. Not reloadable.
|
||||
#rebind_on_network_change: true
|
||||
|
||||
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
|
||||
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
|
||||
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
|
||||
@@ -320,6 +328,47 @@ tun:
|
||||
# SO_RCVBUFFORCE is used to avoid having to raise the system wide max
|
||||
#use_system_route_table_buffer_size: 0
|
||||
|
||||
# EXPERIMENTAL: This option may change or disappear in the future.
|
||||
# Multiport spreads outgoing UDP packets across multiple UDP send ports,
|
||||
# which allows nebula to work around any issues on the underlay network.
|
||||
# Some example issues this could work around:
|
||||
# - UDP rate limits on a per flow basis.
|
||||
# - Partial underlay network failure in which some flows work and some don't
|
||||
# Agreement is done during the handshake to decide if multiport mode will
|
||||
# be used for a given tunnel (one side must have tx_enabled set, the other
|
||||
# side must have rx_enabled set)
|
||||
#
|
||||
# NOTE: you cannot use multiport on a host if you are relying on UDP hole
|
||||
# punching to get through a NAT or firewall.
|
||||
#
|
||||
# NOTE: Linux only (uses raw sockets to send). Also currently only works
|
||||
# with IPv4 underlay network remotes.
|
||||
#
|
||||
# The default values are listed below:
|
||||
#multiport:
|
||||
# This host support sending via multiple UDP ports.
|
||||
#tx_enabled: false
|
||||
#
|
||||
# This host supports receiving packets sent from multiple UDP ports.
|
||||
#rx_enabled: false
|
||||
#
|
||||
# How many UDP ports to use when sending. The lowest source port will be
|
||||
# listen.port and go up to (but not including) listen.port + tx_ports.
|
||||
#tx_ports: 100
|
||||
#
|
||||
# NOTE: All of your hosts must be running a version of Nebula that supports
|
||||
# multiport if you want to enable this feature. Older versions of Nebula
|
||||
# will be confused by these multiport handshakes.
|
||||
#
|
||||
# If handshakes are not getting a response, attempt to transmit handshakes
|
||||
# using random UDP source ports (to get around partial underlay network
|
||||
# failures).
|
||||
#tx_handshake: false
|
||||
#
|
||||
# How many unresponded handshakes we should send before we attempt to
|
||||
# send multiport handshakes.
|
||||
#tx_handshake_delay: 2
|
||||
|
||||
# Configure logging level
|
||||
logging:
|
||||
# trace, debug, info, warn, or error. Default is info and is reloadable.
|
||||
|
||||
@@ -8,6 +8,15 @@ Before=sshd.service
|
||||
Type=notify
|
||||
NotifyAccess=main
|
||||
SyslogIdentifier=nebula
|
||||
|
||||
# Uncomment to run as an unprivileged user with only CAP_NET_ADMIN. Requires a
|
||||
# nebula user that owns the config directory. Add CAP_NET_BIND_SERVICE to both
|
||||
# lines if any listener (lighthouse DNS, listen.port, stats, sshd) binds <1024.
|
||||
#User=nebula
|
||||
#Group=nebula
|
||||
#CapabilityBoundingSet=CAP_NET_ADMIN
|
||||
#AmbientCapabilities=CAP_NET_ADMIN
|
||||
|
||||
ExecReload=/bin/kill -HUP $MAINPID
|
||||
ExecStart=/usr/local/bin/nebula -config /etc/nebula/config.yml
|
||||
Restart=always
|
||||
|
||||
+8
-8
@@ -44,8 +44,8 @@ type Firewall struct {
|
||||
InRules *FirewallTable
|
||||
OutRules *FirewallTable
|
||||
|
||||
InSendReject bool
|
||||
OutSendReject bool
|
||||
InboundSendReject bool
|
||||
OutboundSendReject bool
|
||||
|
||||
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
|
||||
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
|
||||
@@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal
|
||||
inboundAction := c.GetString("firewall.inbound_action", "drop")
|
||||
switch inboundAction {
|
||||
case "reject":
|
||||
fw.InSendReject = true
|
||||
fw.InboundSendReject = true
|
||||
case "drop":
|
||||
fw.InSendReject = false
|
||||
fw.InboundSendReject = false
|
||||
default:
|
||||
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
||||
fw.InSendReject = false
|
||||
fw.InboundSendReject = false
|
||||
}
|
||||
|
||||
outboundAction := c.GetString("firewall.outbound_action", "drop")
|
||||
switch outboundAction {
|
||||
case "reject":
|
||||
fw.OutSendReject = true
|
||||
fw.OutboundSendReject = true
|
||||
case "drop":
|
||||
fw.OutSendReject = false
|
||||
fw.OutboundSendReject = false
|
||||
default:
|
||||
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
||||
fw.OutSendReject = false
|
||||
fw.OutboundSendReject = false
|
||||
}
|
||||
|
||||
err := AddFirewallRulesFromConfig(l, false, c, fw)
|
||||
|
||||
@@ -3,6 +3,7 @@ package firewall
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
mathrand "math/rand"
|
||||
"net/netip"
|
||||
)
|
||||
|
||||
@@ -65,3 +66,30 @@ func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
"Fragment": fp.Fragment,
|
||||
})
|
||||
}
|
||||
|
||||
// UDPSendPort calculates the UDP port to send from when using multiport mode.
|
||||
// The result will be from [0, numBuckets)
|
||||
func (fp Packet) UDPSendPort(numBuckets int) uint16 {
|
||||
if numBuckets <= 1 {
|
||||
return 0
|
||||
}
|
||||
|
||||
// If there is no port (like an ICMP packet), pick a random UDP send port
|
||||
if fp.LocalPort == 0 {
|
||||
return uint16(mathrand.Intn(numBuckets))
|
||||
}
|
||||
|
||||
// A decent enough 32bit hash function
|
||||
// Prospecting for Hash Functions
|
||||
// - https://nullprogram.com/blog/2018/07/31/
|
||||
// - https://github.com/skeeto/hash-prospector
|
||||
// [16 21f0aaad 15 d35a2d97 15] = 0.10760229515479501
|
||||
x := (uint32(fp.LocalPort) << 16) | uint32(fp.RemotePort)
|
||||
x ^= x >> 16
|
||||
x *= 0x21f0aaad
|
||||
x ^= x >> 15
|
||||
x *= 0xd35a2d97
|
||||
x ^= x >> 15
|
||||
|
||||
return uint16(x) % uint16(numBuckets)
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
module github.com/slackhq/nebula
|
||||
|
||||
go 1.25.0
|
||||
go 1.26.0
|
||||
|
||||
require (
|
||||
dario.cat/mergo v1.0.2
|
||||
@@ -24,12 +24,12 @@ require (
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
go.uber.org/goleak v1.3.0
|
||||
go.yaml.in/yaml/v3 v3.0.4
|
||||
golang.org/x/crypto v0.53.0
|
||||
golang.org/x/crypto v0.54.0
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/sync v0.21.0
|
||||
golang.org/x/sys v0.46.0
|
||||
golang.org/x/term v0.44.0
|
||||
golang.org/x/net v0.57.0
|
||||
golang.org/x/sync v0.22.0
|
||||
golang.org/x/sys v0.47.0
|
||||
golang.org/x/term v0.45.0
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||
|
||||
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
||||
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
||||
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
||||
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
|
||||
@@ -24,6 +24,14 @@ message NebulaHandshakeDetails {
|
||||
uint64 Cookie = 4 [deprecated = true];
|
||||
uint64 Time = 5;
|
||||
uint32 CertVersion = 8;
|
||||
// reserved for WIP multiport
|
||||
reserved 6, 7;
|
||||
|
||||
MultiPortDetails InitiatorMultiPort = 6;
|
||||
MultiPortDetails ResponderMultiPort = 7;
|
||||
}
|
||||
|
||||
message MultiPortDetails {
|
||||
bool RxSupported = 1;
|
||||
bool TxSupported = 2;
|
||||
uint32 BasePort = 3;
|
||||
uint32 TotalPorts = 4;
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -71,6 +72,7 @@ func newTestMachine(
|
||||
cs.version, cs.getCredential,
|
||||
verifier, func() (uint32, error) { return localIndex, nil },
|
||||
initiator, header.HandshakeIXPSK0,
|
||||
config.MultiPortConfig{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
return m
|
||||
|
||||
+43
-1
@@ -3,11 +3,13 @@ package handshake
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"math"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
)
|
||||
|
||||
@@ -39,6 +41,10 @@ type Result struct {
|
||||
HandshakeTime uint64
|
||||
MessageIndex uint64 // number of messages exchanged during the handshake
|
||||
Initiator bool
|
||||
|
||||
MultiportRx bool
|
||||
MultiportTx bool
|
||||
MultiportBasePort uint16
|
||||
}
|
||||
|
||||
// Machine drives a Noise handshake through N messages. It handles Noise
|
||||
@@ -67,6 +73,8 @@ type Machine struct {
|
||||
remoteCertSet bool
|
||||
payloadSet bool
|
||||
failed bool
|
||||
|
||||
multiport config.MultiPortConfig
|
||||
}
|
||||
|
||||
// NewMachine creates a handshake state machine. The subtype determines both
|
||||
@@ -80,6 +88,7 @@ func NewMachine(
|
||||
allocIndex IndexAllocator,
|
||||
initiator bool,
|
||||
subtype header.MessageSubType,
|
||||
multiport config.MultiPortConfig,
|
||||
) (*Machine, error) {
|
||||
info, err := subtypeInfoFor(subtype)
|
||||
if err != nil {
|
||||
@@ -108,6 +117,8 @@ func NewMachine(
|
||||
Initiator: initiator,
|
||||
Cipher: cred.cipherSuite,
|
||||
},
|
||||
|
||||
multiport: multiport,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -298,7 +309,7 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
||||
}
|
||||
|
||||
// Assert the payload contains exactly what we expect
|
||||
hasPayloadData := payload.InitiatorIndex != 0 || payload.ResponderIndex != 0 || payload.Time != 0
|
||||
hasPayloadData := payload.InitiatorIndex != 0 || payload.ResponderIndex != 0 || payload.Time != 0 || payload.InitiatorMultiPort != nil || payload.ResponderMultiPort != nil
|
||||
if hasPayloadData != flags.expectsPayload {
|
||||
m.failed = true
|
||||
return ErrUnexpectedContent
|
||||
@@ -315,8 +326,22 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
||||
var remoteIndex uint32
|
||||
if m.result.Initiator {
|
||||
remoteIndex = payload.ResponderIndex
|
||||
if payload.ResponderMultiPort != nil {
|
||||
m.result.MultiportRx = payload.ResponderMultiPort.RxSupported
|
||||
m.result.MultiportTx = payload.ResponderMultiPort.TxSupported
|
||||
if payload.ResponderMultiPort.BasePort <= math.MaxUint16 {
|
||||
m.result.MultiportBasePort = uint16(payload.ResponderMultiPort.BasePort)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
remoteIndex = payload.InitiatorIndex
|
||||
if payload.InitiatorMultiPort != nil {
|
||||
m.result.MultiportRx = payload.InitiatorMultiPort.RxSupported
|
||||
m.result.MultiportTx = payload.InitiatorMultiPort.TxSupported
|
||||
if payload.InitiatorMultiPort.BasePort <= math.MaxUint16 {
|
||||
m.result.MultiportBasePort = uint16(payload.InitiatorMultiPort.BasePort)
|
||||
}
|
||||
}
|
||||
}
|
||||
// The payload presence check above can be satisfied by Time alone, so a payload
|
||||
// could still carry a zero index here. We need to reject it.
|
||||
@@ -397,11 +422,28 @@ func (m *Machine) marshalOutgoing(flags msgFlags) ([]byte, error) {
|
||||
|
||||
if m.result.Initiator {
|
||||
p.InitiatorIndex = m.result.LocalIndex
|
||||
if m.multiport.Rx || m.multiport.Tx {
|
||||
p.InitiatorMultiPort = &PayloadMultiPortDetails{
|
||||
RxSupported: m.multiport.Rx,
|
||||
TxSupported: m.multiport.Tx,
|
||||
BasePort: uint32(m.multiport.TxBasePort),
|
||||
TotalPorts: uint32(m.multiport.TxPorts),
|
||||
}
|
||||
}
|
||||
} else {
|
||||
p.ResponderIndex = m.result.LocalIndex
|
||||
p.InitiatorIndex = m.result.RemoteIndex
|
||||
if m.multiport.Rx || m.multiport.Tx {
|
||||
p.ResponderMultiPort = &PayloadMultiPortDetails{
|
||||
RxSupported: m.multiport.Rx,
|
||||
TxSupported: m.multiport.Tx,
|
||||
BasePort: uint32(m.multiport.TxBasePort),
|
||||
TotalPorts: uint32(m.multiport.TxPorts),
|
||||
}
|
||||
}
|
||||
}
|
||||
p.Time = uint64(time.Now().UnixNano())
|
||||
|
||||
}
|
||||
if flags.expectsCert {
|
||||
cred := m.getCred(m.myVersion)
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -444,6 +445,7 @@ func TestMachineThreeMessagePattern(t *testing.T) {
|
||||
initCS.getCredential, v,
|
||||
func() (uint32, error) { return 1000, nil },
|
||||
true, header.HandshakeXXPSK0,
|
||||
config.MultiPortConfig{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -452,6 +454,7 @@ func TestMachineThreeMessagePattern(t *testing.T) {
|
||||
respCS.getCredential, v,
|
||||
func() (uint32, error) { return 2000, nil },
|
||||
false, header.HandshakeXXPSK0,
|
||||
config.MultiPortConfig{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -20,6 +20,16 @@ type Payload struct {
|
||||
ResponderIndex uint32
|
||||
Time uint64
|
||||
CertVersion uint32
|
||||
|
||||
InitiatorMultiPort *PayloadMultiPortDetails
|
||||
ResponderMultiPort *PayloadMultiPortDetails
|
||||
}
|
||||
|
||||
type PayloadMultiPortDetails struct {
|
||||
RxSupported bool
|
||||
TxSupported bool
|
||||
BasePort uint32
|
||||
TotalPorts uint32
|
||||
}
|
||||
|
||||
// Proto field numbers for NebulaHandshakeDetails
|
||||
@@ -29,6 +39,17 @@ const (
|
||||
fieldResponderIndex = 3 // uint32
|
||||
fieldTime = 5 // uint64
|
||||
fieldCertVersion = 8 // uint32
|
||||
|
||||
fieldInitiatorMultiPort = 6 // MultiPortDetails
|
||||
fieldResponderMultiPort = 7 // MultiPortDetails
|
||||
)
|
||||
|
||||
// Proto field numbers for MultiPortDetails
|
||||
const (
|
||||
fieldMultiportRxSupported = 1 // bool
|
||||
fieldMultiportTxSupported = 2 // bool
|
||||
fieldMultiportBasePort = 3 // uint32
|
||||
fieldMultiportTotalPorts = 4 // uint32
|
||||
)
|
||||
|
||||
// MarshalPayload encodes a handshake payload in protobuf wire format compatible
|
||||
@@ -57,6 +78,16 @@ func MarshalPayload(out []byte, p Payload) []byte {
|
||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.CertVersion))
|
||||
}
|
||||
if p.InitiatorMultiPort != nil {
|
||||
details = protowire.AppendTag(details, fieldInitiatorMultiPort, protowire.BytesType)
|
||||
details = protowire.AppendVarint(details, uint64(p.InitiatorMultiPort.size()))
|
||||
details = p.InitiatorMultiPort.marshal(details)
|
||||
}
|
||||
if p.ResponderMultiPort != nil {
|
||||
details = protowire.AppendTag(details, fieldResponderMultiPort, protowire.BytesType)
|
||||
details = protowire.AppendVarint(details, uint64(p.ResponderMultiPort.size()))
|
||||
details = p.ResponderMultiPort.marshal(details)
|
||||
}
|
||||
|
||||
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
||||
out = protowire.AppendBytes(out, details)
|
||||
@@ -64,6 +95,23 @@ func MarshalPayload(out []byte, p Payload) []byte {
|
||||
return out
|
||||
}
|
||||
|
||||
func (p PayloadMultiPortDetails) marshal(details []byte) []byte {
|
||||
details = protowire.AppendTag(details, fieldMultiportRxSupported, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, protowire.EncodeBool(p.RxSupported))
|
||||
details = protowire.AppendTag(details, fieldMultiportTxSupported, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, protowire.EncodeBool(p.TxSupported))
|
||||
details = protowire.AppendTag(details, fieldMultiportBasePort, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.BasePort))
|
||||
details = protowire.AppendTag(details, fieldMultiportTotalPorts, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.TotalPorts))
|
||||
|
||||
return details
|
||||
}
|
||||
|
||||
func (p PayloadMultiPortDetails) size() int {
|
||||
return 4 + 2 + protowire.SizeVarint(uint64(p.BasePort)) + protowire.SizeVarint(uint64(p.TotalPorts))
|
||||
}
|
||||
|
||||
// UnmarshalPayload decodes a protobuf-encoded NebulaHandshake message.
|
||||
func UnmarshalPayload(b []byte) (Payload, error) {
|
||||
var p Payload
|
||||
@@ -161,6 +209,97 @@ func unmarshalPayloadDetails(p *Payload, b []byte) error {
|
||||
}
|
||||
p.CertVersion = uint32(v)
|
||||
b = b[n:]
|
||||
case fieldInitiatorMultiPort:
|
||||
if typ != protowire.BytesType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
d, n := protowire.ConsumeBytes(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeMessage
|
||||
}
|
||||
b = b[n:]
|
||||
p.InitiatorMultiPort = new(PayloadMultiPortDetails)
|
||||
if err := unmarshalPayloadMultiPortDetails(p.InitiatorMultiPort, d); err != nil {
|
||||
return err
|
||||
}
|
||||
case fieldResponderMultiPort:
|
||||
if typ != protowire.BytesType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
d, n := protowire.ConsumeBytes(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeMessage
|
||||
}
|
||||
b = b[n:]
|
||||
p.ResponderMultiPort = new(PayloadMultiPortDetails)
|
||||
if err := unmarshalPayloadMultiPortDetails(p.ResponderMultiPort, d); err != nil {
|
||||
return err
|
||||
}
|
||||
default:
|
||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
b = b[n:]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func unmarshalPayloadMultiPortDetails(p *PayloadMultiPortDetails, b []byte) error {
|
||||
for len(b) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
b = b[n:]
|
||||
|
||||
// For known field numbers, reject any non-matching wire type as a
|
||||
// hard error rather than silently skipping. The caller will catch
|
||||
// missing-field cases downstream, but a wire-type mismatch on a tag
|
||||
// we know is a peer protocol violation worth flagging here.
|
||||
// Repeated occurrences of a singular field follow proto3 last-wins.
|
||||
switch num {
|
||||
case fieldMultiportRxSupported:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.RxSupported = protowire.DecodeBool(v)
|
||||
b = b[n:]
|
||||
case fieldMultiportTxSupported:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.TxSupported = protowire.DecodeBool(v)
|
||||
b = b[n:]
|
||||
case fieldMultiportBasePort:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.BasePort = uint32(v)
|
||||
b = b[n:]
|
||||
case fieldMultiportTotalPorts:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.TotalPorts = uint32(v)
|
||||
b = b[n:]
|
||||
default:
|
||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||
if n < 0 {
|
||||
|
||||
+16
-16
@@ -117,24 +117,24 @@ func TestPayloadUnknownFields(t *testing.T) {
|
||||
assert.Equal(t, uint32(88), got.ResponderIndex)
|
||||
})
|
||||
|
||||
t.Run("reserved fields 6 and 7 are skipped", func(t *testing.T) {
|
||||
// Fields 6 and 7 are reserved in the proto definition
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 100)
|
||||
details = protowire.AppendTag(details, 6, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 1)
|
||||
details = protowire.AppendTag(details, 7, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 2)
|
||||
// t.Run("reserved fields 6 and 7 are skipped", func(t *testing.T) {
|
||||
// // Fields 6 and 7 are reserved in the proto definition
|
||||
// var details []byte
|
||||
// details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
// details = protowire.AppendVarint(details, 100)
|
||||
// details = protowire.AppendTag(details, 6, protowire.VarintType)
|
||||
// details = protowire.AppendVarint(details, 1)
|
||||
// details = protowire.AppendTag(details, 7, protowire.VarintType)
|
||||
// details = protowire.AppendVarint(details, 2)
|
||||
|
||||
var data []byte
|
||||
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||
data = protowire.AppendBytes(data, details)
|
||||
// var data []byte
|
||||
// data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||
// data = protowire.AppendBytes(data, details)
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(100), got.InitiatorIndex)
|
||||
})
|
||||
// got, err := UnmarshalPayload(data)
|
||||
// require.NoError(t, err)
|
||||
// assert.Equal(t, uint32(100), got.InitiatorIndex)
|
||||
// })
|
||||
}
|
||||
|
||||
func TestPayloadBytesConsumed(t *testing.T) {
|
||||
|
||||
+68
-8
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -71,6 +72,9 @@ type HandshakeManager struct {
|
||||
f *Interface
|
||||
l *slog.Logger
|
||||
|
||||
multiPort config.MultiPortConfig
|
||||
udpRaw *udp.RawConn
|
||||
|
||||
// can be used to trigger outbound handshake for the given vpnIp
|
||||
trigger chan netip.Addr
|
||||
}
|
||||
@@ -291,11 +295,18 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
|
||||
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
||||
var sentTo []netip.AddrPort
|
||||
var sentMultiport bool
|
||||
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||
err := hm.outside.WriteTo(stage0, addr)
|
||||
if err != nil {
|
||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||
// These repeat every attempt, so match the success log below and only shout when the remotes changed
|
||||
level := slog.LevelDebug
|
||||
if remotesHaveChanged {
|
||||
level = slog.LevelError
|
||||
}
|
||||
|
||||
hostinfo.logger(hm.l).Log(context.Background(), level, "Failed to send handshake message",
|
||||
"udpAddr", addr,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
@@ -305,6 +316,29 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
} else {
|
||||
sentTo = append(sentTo, addr)
|
||||
}
|
||||
|
||||
// Attempt a multiport handshake if we are past the TxHandshakeDelay attempts
|
||||
if hm.multiPort.TxHandshake && hm.udpRaw != nil && hh.counter >= hm.multiPort.TxHandshakeDelay {
|
||||
sentMultiport = true
|
||||
// We need to re-allocate with 8 bytes at the start of SOCK_RAW
|
||||
raw := hostinfo.HandshakePacket[0x80]
|
||||
if raw == nil {
|
||||
raw = make([]byte, len(hostinfo.HandshakePacket[0])+udp.RawOverhead)
|
||||
copy(raw[udp.RawOverhead:], hostinfo.HandshakePacket[0])
|
||||
hostinfo.HandshakePacket[0x80] = raw
|
||||
}
|
||||
|
||||
hm.messageMetrics.Tx(header.Handshake, header.MessageSubType(hostinfo.HandshakePacket[0][1]), 1)
|
||||
err = hm.udpRaw.WriteTo(raw, udp.RandomSendPort.UDPSendPort(hm.multiPort.TxPorts), addr)
|
||||
if err != nil {
|
||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||
"error", err,
|
||||
"udpAddr", addr,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// Don't be too noisy or confusing if we fail to send a handshake - if we don't get through we'll eventually log a timeout,
|
||||
@@ -314,6 +348,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
"udpAddrs", sentTo,
|
||||
"initiatorIndex", hostinfo.localIndexId,
|
||||
"handshake", hsFields,
|
||||
"multiportHandshake", sentMultiport,
|
||||
)
|
||||
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
||||
@@ -430,14 +465,11 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
||||
// Check if we already have a tunnel with this vpn ip
|
||||
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
|
||||
if found && existingHostInfo != nil {
|
||||
testHostInfo := existingHostInfo
|
||||
for testHostInfo != nil {
|
||||
// Is it just a delayed handshake packet?
|
||||
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address.
|
||||
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) {
|
||||
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
|
||||
return testHostInfo, ErrAlreadySeen
|
||||
}
|
||||
|
||||
testHostInfo = testHostInfo.next
|
||||
}
|
||||
|
||||
// Is this a newer handshake?
|
||||
@@ -532,7 +564,9 @@ func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||
|
||||
func (hm *HandshakeManager) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
delete(hm.vpnIps, addr)
|
||||
if cur, ok := hm.vpnIps[addr]; ok && cur.hostinfo == hostinfo {
|
||||
delete(hm.vpnIps, addr)
|
||||
}
|
||||
}
|
||||
|
||||
if len(hm.vpnIps) == 0 {
|
||||
@@ -667,6 +701,7 @@ func (hm *HandshakeManager) buildStage0Packet(hh *HandshakeHostInfo) bool {
|
||||
v, cs.GetCredential,
|
||||
hm.certVerifier(), func() (uint32, error) { return hm.allocateIndex(hh) },
|
||||
true, header.HandshakeIXPSK0,
|
||||
hm.multiPort,
|
||||
)
|
||||
if err != nil {
|
||||
hm.f.l.Error("Failed to create handshake machine",
|
||||
@@ -708,6 +743,7 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
v, cs.GetCredential,
|
||||
hm.certVerifier(), func() (uint32, error) { return generateIndex(f.l) },
|
||||
false, header.HandshakeIXPSK0,
|
||||
hm.multiPort,
|
||||
)
|
||||
if err != nil {
|
||||
f.l.Error("Failed to create handshake machine", "from", via, "error", err)
|
||||
@@ -732,6 +768,12 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
return
|
||||
}
|
||||
|
||||
if !via.IsRelayed && result.MultiportTx && result.MultiportBasePort != via.UdpAddr.Port() {
|
||||
// The other side sent us a handshake from a different port, make sure
|
||||
// we send responses back to the BasePort
|
||||
via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), result.MultiportBasePort)
|
||||
}
|
||||
|
||||
remoteCert := result.RemoteCert
|
||||
if remoteCert == nil {
|
||||
f.l.Error("Handshake did not produce a peer certificate", "from", via)
|
||||
@@ -756,6 +798,8 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
relayForByAddr: map[netip.Addr]*Relay{},
|
||||
relayForByIdx: map[uint32]*Relay{},
|
||||
},
|
||||
multiportTx: hm.multiPort.Tx && result.MultiportRx,
|
||||
multiportRx: hm.multiPort.Rx && result.MultiportTx,
|
||||
}
|
||||
|
||||
msg := "Handshake message received"
|
||||
@@ -772,6 +816,8 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
||||
"initiatorIndex", result.RemoteIndex,
|
||||
"responderIndex", result.LocalIndex,
|
||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||
"multiportTx", hostinfo.multiportTx,
|
||||
"multiportRx", hostinfo.multiportRx,
|
||||
)
|
||||
|
||||
// packet aliases the listener's incoming buffer, so this copy must stay.
|
||||
@@ -862,6 +908,14 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
return
|
||||
}
|
||||
|
||||
if !via.IsRelayed && result.MultiportTx && result.MultiportBasePort != via.UdpAddr.Port() {
|
||||
// The other side sent us a handshake from a different port, make sure
|
||||
// we send responses back to the BasePort
|
||||
via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), result.MultiportBasePort)
|
||||
}
|
||||
hostinfo.multiportTx = hm.multiPort.Tx && result.MultiportRx
|
||||
hostinfo.multiportRx = hm.multiPort.Rx && result.MultiportTx
|
||||
|
||||
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
||||
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
||||
|
||||
@@ -956,6 +1010,8 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||
"durationNs", duration,
|
||||
"sentCachedPackets", len(hh.packetStore),
|
||||
"multiportTx", hostinfo.multiportTx,
|
||||
"multiportRx", hostinfo.multiportRx,
|
||||
)
|
||||
|
||||
hostinfo.vpnAddrs = vpnAddrs
|
||||
@@ -1080,7 +1136,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
// We received a valid handshake on this relay, so make sure the relay
|
||||
// state reflects that, in case it had been marked Disestablished.
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||
}
|
||||
@@ -1097,6 +1153,10 @@ func (hm *HandshakeManager) handleCheckAndCompleteError(err error, existing, hos
|
||||
|
||||
switch err {
|
||||
case ErrAlreadySeen:
|
||||
if hostinfo.multiportRx {
|
||||
// The other host is sending to us with multiport, so only grab the IP
|
||||
via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), hostinfo.GetRemote().Port())
|
||||
}
|
||||
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||
}
|
||||
|
||||
+158
-98
@@ -56,11 +56,20 @@ type Relay struct {
|
||||
}
|
||||
|
||||
type HostMap struct {
|
||||
sync.RWMutex //Because we concurrently read and write to our maps
|
||||
Indexes map[uint32]*HostInfo
|
||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||
RemoteIndexes map[uint32]*HostInfo
|
||||
sync.RWMutex //Because we concurrently read and write to our maps
|
||||
Indexes map[uint32]*HostInfo
|
||||
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
|
||||
RemoteIndexes map[uint32]*HostInfo
|
||||
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
|
||||
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
|
||||
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
|
||||
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
|
||||
// never corrupt another address's ordering the way the old shared next/prev chain could.
|
||||
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
|
||||
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
|
||||
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
|
||||
Hosts map[netip.Addr]*HostInfo
|
||||
moreHosts map[netip.Addr][]*HostInfo
|
||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||
l *slog.Logger
|
||||
}
|
||||
@@ -245,6 +254,12 @@ type HostInfo struct {
|
||||
networks *bart.Table[NetworkType]
|
||||
relayState RelayState
|
||||
|
||||
// If true, we should send to this remote using multiport
|
||||
multiportTx bool
|
||||
|
||||
// If true, we will receive from this remote using multiport
|
||||
multiportRx bool
|
||||
|
||||
// HandshakePacket records the packets used to create this hostinfo
|
||||
// We need these to avoid replayed handshake packets creating new hostinfos which causes churn
|
||||
HandshakePacket map[uint8][]byte
|
||||
@@ -266,10 +281,6 @@ type HostInfo struct {
|
||||
lastRoam time.Time
|
||||
lastRoamRemote netip.AddrPort
|
||||
|
||||
// Used to track other hostinfos for this vpn ip since only 1 can be primary
|
||||
// Synchronised via hostmap lock and not the hostinfo lock.
|
||||
next, prev *HostInfo
|
||||
|
||||
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
|
||||
in, out, pendingDeletion atomic.Bool
|
||||
|
||||
@@ -282,7 +293,6 @@ type HostInfo struct {
|
||||
type ViaSender struct {
|
||||
UdpAddr netip.AddrPort
|
||||
relayHI *HostInfo // relayHI is the host info object of the relay
|
||||
remoteIdx uint32 // remoteIdx is the index included in the header of the received packet
|
||||
relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us.
|
||||
IsRelayed bool // IsRelayed is true if the packet was sent through a relay
|
||||
}
|
||||
@@ -334,6 +344,7 @@ func newHostMap(l *slog.Logger) *HostMap {
|
||||
Relays: map[uint32]*HostInfo{},
|
||||
RemoteIndexes: map[uint32]*HostInfo{},
|
||||
Hosts: map[netip.Addr]*HostInfo{},
|
||||
moreHosts: map[netip.Addr][]*HostInfo{},
|
||||
l: l,
|
||||
}
|
||||
}
|
||||
@@ -382,13 +393,55 @@ func (hm *HostMap) EmitStats() {
|
||||
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
|
||||
}
|
||||
|
||||
// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
|
||||
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty
|
||||
// list removes the address. This is the one place Hosts and moreHosts are written together, keep
|
||||
// it that way. Callers must hold the write lock.
|
||||
func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) {
|
||||
if len(list) == 0 {
|
||||
delete(hm.Hosts, addr)
|
||||
delete(hm.moreHosts, addr)
|
||||
return
|
||||
}
|
||||
hm.Hosts[addr] = list[0]
|
||||
if len(list) > 1 {
|
||||
hm.moreHosts[addr] = list
|
||||
} else {
|
||||
delete(hm.moreHosts, addr)
|
||||
}
|
||||
}
|
||||
|
||||
// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no
|
||||
// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this
|
||||
// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read
|
||||
// or write).
|
||||
func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo {
|
||||
if list, ok := hm.moreHosts[addr]; ok {
|
||||
return list
|
||||
}
|
||||
if h, ok := hm.Hosts[addr]; ok {
|
||||
return []*HostInfo{h}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is
|
||||
// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever
|
||||
// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to
|
||||
// invalidate.
|
||||
func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo {
|
||||
idx := slices.Index(list, hi)
|
||||
if idx < 0 {
|
||||
return list
|
||||
}
|
||||
return slices.Delete(list, idx, idx+1)
|
||||
}
|
||||
|
||||
// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds
|
||||
// any of its vpn addrs, meaning we no longer have a tunnel to the peer
|
||||
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
|
||||
// Delete the host itself, ensuring it's not modified anymore
|
||||
hm.Lock()
|
||||
// If we have a previous or next hostinfo then we are not the last one for this vpn ip
|
||||
final := (hostinfo.next == nil && hostinfo.prev == nil)
|
||||
hm.unlockedDeleteHostInfo(hostinfo)
|
||||
final := hm.unlockedDeleteHostInfo(hostinfo)
|
||||
hm.Unlock()
|
||||
|
||||
return final
|
||||
@@ -400,71 +453,66 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
|
||||
hm.unlockedMakePrimary(hostinfo)
|
||||
}
|
||||
|
||||
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
|
||||
// Get the current primary, if it exists
|
||||
oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
|
||||
|
||||
// Every address in the hostinfo gets elevated to primary
|
||||
for _, vpnAddr := range hostinfo.vpnAddrs {
|
||||
//NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
|
||||
// indexes so it should be fine.
|
||||
hm.Hosts[vpnAddr] = hostinfo
|
||||
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses,
|
||||
// false only when it is no longer in the hostmap at all.
|
||||
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool {
|
||||
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
|
||||
// tunnel teardown, deciding to promote under the read lock and only taking the write lock
|
||||
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every
|
||||
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test.
|
||||
if hm.Indexes[hostinfo.localIndexId] != hostinfo {
|
||||
return false
|
||||
}
|
||||
|
||||
// If we are already primary then we won't bother re-linking
|
||||
if oldHostinfo == hostinfo {
|
||||
return
|
||||
}
|
||||
|
||||
// Unlink this hostinfo
|
||||
if hostinfo.prev != nil {
|
||||
hostinfo.prev.next = hostinfo.next
|
||||
}
|
||||
if hostinfo.next != nil {
|
||||
hostinfo.next.prev = hostinfo.prev
|
||||
}
|
||||
|
||||
// If there wasn't a previous primary then clear out any links
|
||||
if oldHostinfo == nil {
|
||||
hostinfo.next = nil
|
||||
hostinfo.prev = nil
|
||||
return
|
||||
}
|
||||
|
||||
// Relink the hostinfo as primary
|
||||
hostinfo.next = oldHostinfo
|
||||
oldHostinfo.prev = hostinfo
|
||||
hostinfo.prev = nil
|
||||
}
|
||||
|
||||
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil
|
||||
|
||||
// Move hostinfo to the front (primary) of each of its address lists. The lists are
|
||||
// independent per address, so this can never leave a dangling entry the way promoting
|
||||
// against a single shared chain could.
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
if hm.Hosts[addr] != hostinfo {
|
||||
if hm.Hosts[addr] == hostinfo {
|
||||
// Already primary for this address, the list is already in the right order
|
||||
continue
|
||||
}
|
||||
if hostinfo.next != nil {
|
||||
// Promote the next hostinfo in the shared chain to primary for this address
|
||||
hm.Hosts[addr] = hostinfo.next
|
||||
} else {
|
||||
delete(hm.Hosts, addr)
|
||||
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
|
||||
list = append([]*HostInfo{hostinfo}, list...)
|
||||
hm.unlockedSetHostsForAddr(addr, list)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
|
||||
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
|
||||
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
|
||||
// state and disestablish relays.
|
||||
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
||||
// Remove this hostinfo from each of its address lists. The lists are independent, so a
|
||||
// sibling is never promoted to an address it does not own and no other list is touched.
|
||||
final := true
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
if list, ok := hm.moreHosts[addr]; ok {
|
||||
list = removeHostInfo(list, hostinfo)
|
||||
hm.unlockedSetHostsForAddr(addr, list)
|
||||
if len(list) > 0 {
|
||||
final = false
|
||||
}
|
||||
} else if existing, ok := hm.Hosts[addr]; ok {
|
||||
if existing == hostinfo {
|
||||
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
|
||||
delete(hm.Hosts, addr)
|
||||
} else {
|
||||
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
||||
final = false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
||||
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
||||
if len(hm.Hosts) == 0 {
|
||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||
}
|
||||
|
||||
// Splice this hostinfo out of the shared chain exactly once
|
||||
if hostinfo.prev != nil {
|
||||
hostinfo.prev.next = hostinfo.next
|
||||
if len(hm.moreHosts) == 0 {
|
||||
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
||||
}
|
||||
if hostinfo.next != nil {
|
||||
hostinfo.next.prev = hostinfo.prev
|
||||
}
|
||||
|
||||
hostinfo.next = nil
|
||||
hostinfo.prev = nil
|
||||
|
||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
||||
// the remote index pointer here if it points to the hostinfo we are deleting
|
||||
@@ -488,7 +536,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
)
|
||||
}
|
||||
|
||||
if isLastHostinfo {
|
||||
if final {
|
||||
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
|
||||
// hops as 'Requested' so that new relay tunnels are created in the future.
|
||||
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
|
||||
@@ -497,6 +545,8 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
|
||||
delete(hm.Relays, localRelayIdx)
|
||||
}
|
||||
|
||||
return final
|
||||
}
|
||||
|
||||
func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
|
||||
@@ -540,19 +590,30 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
||||
hm.RLock()
|
||||
defer hm.RUnlock()
|
||||
|
||||
// This runs per relayed packet, so check the primary with a single map probe and only consult
|
||||
// moreHosts when the primary can't relay for us.
|
||||
h, ok := hm.Hosts[relayHostIp]
|
||||
if !ok {
|
||||
return nil, nil, errors.New("unable to find host")
|
||||
}
|
||||
|
||||
for h != nil {
|
||||
for _, targetIp := range targetIps {
|
||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||
if ok && r.State == Established {
|
||||
return h, r, nil
|
||||
for _, targetIp := range targetIps {
|
||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||
if ok && r.State == Established {
|
||||
return h, r, nil
|
||||
}
|
||||
}
|
||||
|
||||
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
||||
// list[0] is the primary we already checked
|
||||
for _, h := range list[1:] {
|
||||
for _, targetIp := range targetIps {
|
||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||
if ok && r.State == Established {
|
||||
return h, r, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
h = h.next
|
||||
}
|
||||
|
||||
return nil, nil, errors.New("unable to find host with relay")
|
||||
@@ -560,20 +621,14 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
|
||||
|
||||
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
|
||||
for _, relayHostIp := range hi.relayState.CopyRelayIps() {
|
||||
if h, ok := hm.Hosts[relayHostIp]; ok {
|
||||
for h != nil {
|
||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||
h = h.next
|
||||
}
|
||||
for _, h := range hm.unlockedGetHostList(relayHostIp) {
|
||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||
}
|
||||
}
|
||||
for _, rs := range hi.relayState.CopyAllRelayFor() {
|
||||
if rs.Type == ForwardingType {
|
||||
if h, ok := hm.Hosts[rs.PeerAddr]; ok {
|
||||
for h != nil {
|
||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||
h = h.next
|
||||
}
|
||||
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) {
|
||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -623,22 +678,27 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
||||
}
|
||||
|
||||
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
|
||||
existing := hm.Hosts[vpnAddr]
|
||||
hm.Hosts[vpnAddr] = hostinfo
|
||||
|
||||
if existing != nil && existing != hostinfo {
|
||||
hostinfo.next = existing
|
||||
existing.prev = hostinfo
|
||||
existing, ok := hm.Hosts[vpnAddr]
|
||||
if !ok {
|
||||
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
||||
hm.Hosts[vpnAddr] = hostinfo
|
||||
return
|
||||
}
|
||||
|
||||
i := 1
|
||||
check := hostinfo
|
||||
for check != nil {
|
||||
if i > MaxHostInfosPerVpnIp {
|
||||
hm.unlockedDeleteHostInfo(check)
|
||||
}
|
||||
check = check.next
|
||||
i++
|
||||
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so
|
||||
// we never hold a duplicate, then prepend.
|
||||
list, ok := hm.moreHosts[vpnAddr]
|
||||
if !ok {
|
||||
list = []*HostInfo{existing}
|
||||
}
|
||||
list = removeHostInfo(list, hostinfo)
|
||||
list = append([]*HostInfo{hostinfo}, list...)
|
||||
hm.unlockedSetHostsForAddr(vpnAddr, list)
|
||||
|
||||
// Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it.
|
||||
// Deleting it removes it from all of its addresses and the index maps, matching prior behavior.
|
||||
if len(list) > MaxHostInfosPerVpnIp {
|
||||
hm.unlockedDeleteHostInfo(list[len(list)-1])
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+237
-181
@@ -2,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
@@ -10,78 +11,84 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It
|
||||
// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it
|
||||
// fails fast.
|
||||
func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 {
|
||||
t.Helper()
|
||||
assertHostMapInvariants(t, hm)
|
||||
list := hm.unlockedGetHostList(addr)
|
||||
ids := make([]uint32, len(list))
|
||||
for i, h := range list {
|
||||
ids[i] = h.localIndexId
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses
|
||||
// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold
|
||||
// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every
|
||||
// indexed hostinfo is reachable through each of its addresses.
|
||||
func assertHostMapInvariants(t *testing.T, hm *HostMap) {
|
||||
t.Helper()
|
||||
for addr, list := range hm.moreHosts {
|
||||
require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr)
|
||||
require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr)
|
||||
seen := map[*HostInfo]bool{}
|
||||
for _, h := range list {
|
||||
require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr)
|
||||
require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId)
|
||||
seen[h] = true
|
||||
require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId)
|
||||
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId)
|
||||
}
|
||||
}
|
||||
for addr, h := range hm.Hosts {
|
||||
require.NotNilf(t, h, "Hosts[%s] must never be nil", addr)
|
||||
require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId)
|
||||
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId)
|
||||
}
|
||||
for idx, h := range hm.Indexes {
|
||||
require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId)
|
||||
for _, va := range h.vpnAddrs {
|
||||
require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHostMap_MakePrimary(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
hm := newHostMap(l)
|
||||
|
||||
f := &Interface{}
|
||||
a := netip.MustParseAddr("0.0.0.1")
|
||||
|
||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||
|
||||
hm.unlockedAddHostInfo(h4, f)
|
||||
hm.unlockedAddHostInfo(h3, f)
|
||||
hm.unlockedAddHostInfo(h2, f)
|
||||
hm.unlockedAddHostInfo(h1, f)
|
||||
|
||||
// Make sure we go h1 -> h2 -> h3 -> h4
|
||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
||||
assert.Nil(t, h4.next)
|
||||
// Most-recently-added is primary: h1, h2, h3, h4
|
||||
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a))
|
||||
assert.Equal(t, h1, hm.QueryVpnAddr(a))
|
||||
|
||||
// Swap h3/middle to primary
|
||||
// Swap the middle to primary: h3, h1, h2, h4
|
||||
hm.MakePrimary(h3)
|
||||
assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a))
|
||||
assert.Equal(t, h3, hm.QueryVpnAddr(a))
|
||||
|
||||
// Make sure we go h3 -> h1 -> h2 -> h4
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h3.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h1.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
||||
assert.Nil(t, h4.next)
|
||||
|
||||
// Swap h4/tail to primary
|
||||
// Swap the tail to primary: h4, h3, h1, h2
|
||||
hm.MakePrimary(h4)
|
||||
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||
|
||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
||||
assert.Nil(t, h2.next)
|
||||
|
||||
// Swap h4 again should be no-op
|
||||
// Swapping the current primary again is a no-op
|
||||
hm.MakePrimary(h4)
|
||||
|
||||
// Make sure we go h4 -> h3 -> h1 -> h2
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
|
||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
||||
assert.Nil(t, h2.next)
|
||||
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
|
||||
}
|
||||
|
||||
func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||
@@ -89,13 +96,14 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||
hm := newHostMap(l)
|
||||
|
||||
f := &Interface{}
|
||||
a := netip.MustParseAddr("0.0.0.1")
|
||||
|
||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
|
||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
|
||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
|
||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
|
||||
h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
|
||||
h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
|
||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3}
|
||||
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4}
|
||||
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5}
|
||||
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6}
|
||||
|
||||
hm.unlockedAddHostInfo(h6, f)
|
||||
hm.unlockedAddHostInfo(h5, f)
|
||||
@@ -104,94 +112,110 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||
hm.unlockedAddHostInfo(h2, f)
|
||||
hm.unlockedAddHostInfo(h1, f)
|
||||
|
||||
// h6 should be deleted
|
||||
assert.Nil(t, h6.next)
|
||||
assert.Nil(t, h6.prev)
|
||||
h := hm.QueryIndex(h6.localIndexId)
|
||||
assert.Nil(t, h)
|
||||
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first.
|
||||
assert.Nil(t, hm.QueryIndex(h6.localIndexId))
|
||||
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a))
|
||||
|
||||
// Make sure we go h1 -> h2 -> h3 -> h4 -> h5
|
||||
prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h1.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
||||
assert.Nil(t, h5.next)
|
||||
// Delete primary; not final since siblings remain.
|
||||
assert.False(t, hm.DeleteHostInfo(h1))
|
||||
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||
|
||||
// Delete primary
|
||||
hm.DeleteHostInfo(h1)
|
||||
assert.Nil(t, h1.prev)
|
||||
assert.Nil(t, h1.next)
|
||||
// Deleting the same hostinfo again must not report final while siblings remain and must not
|
||||
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a
|
||||
// second delete looked final and wiped lighthouse state out from under the live sibling.
|
||||
assert.False(t, hm.DeleteHostInfo(h1))
|
||||
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
|
||||
|
||||
// Make sure we go h2 -> h3 -> h4 -> h5
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
|
||||
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
|
||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
||||
assert.Nil(t, h5.next)
|
||||
// Delete a middle node.
|
||||
assert.False(t, hm.DeleteHostInfo(h3))
|
||||
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a))
|
||||
|
||||
// Delete in the middle
|
||||
hm.DeleteHostInfo(h3)
|
||||
assert.Nil(t, h3.prev)
|
||||
assert.Nil(t, h3.next)
|
||||
// Delete the tail.
|
||||
assert.False(t, hm.DeleteHostInfo(h5))
|
||||
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a))
|
||||
|
||||
// Make sure we go h2 -> h4 -> h5
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
||||
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
|
||||
assert.Nil(t, h5.next)
|
||||
// Delete the head; h4 remains and becomes primary.
|
||||
assert.False(t, hm.DeleteHostInfo(h2))
|
||||
assert.Equal(t, []uint32{4}, chainIds(t, hm, a))
|
||||
assert.Equal(t, h4, hm.QueryVpnAddr(a))
|
||||
|
||||
// Delete the tail
|
||||
hm.DeleteHostInfo(h5)
|
||||
assert.Nil(t, h5.prev)
|
||||
assert.Nil(t, h5.next)
|
||||
// Delete the only remaining item; final is true and the address is gone.
|
||||
assert.True(t, hm.DeleteHostInfo(h4))
|
||||
assert.Empty(t, chainIds(t, hm, a))
|
||||
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||
|
||||
// Make sure we go h2 -> h4
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h2.localIndexId, prim.localIndexId)
|
||||
assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
|
||||
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
|
||||
assert.Nil(t, h4.next)
|
||||
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore.
|
||||
assert.True(t, hm.DeleteHostInfo(h4))
|
||||
assert.Empty(t, chainIds(t, hm, a))
|
||||
}
|
||||
|
||||
// Delete the head
|
||||
hm.DeleteHostInfo(h2)
|
||||
assert.Nil(t, h2.prev)
|
||||
assert.Nil(t, h2.next)
|
||||
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with
|
||||
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and
|
||||
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a
|
||||
// no-op, not a resurrection that installs an unmanaged primary.
|
||||
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
hm := newHostMap(l)
|
||||
f := &Interface{}
|
||||
a := netip.MustParseAddr("0.0.0.1")
|
||||
|
||||
// Make sure we only have h4
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Equal(t, h4.localIndexId, prim.localIndexId)
|
||||
assert.Nil(t, prim.prev)
|
||||
assert.Nil(t, prim.next)
|
||||
assert.Nil(t, h4.next)
|
||||
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2}
|
||||
hm.unlockedAddHostInfo(h1, f)
|
||||
hm.unlockedAddHostInfo(h2, f)
|
||||
|
||||
// Delete the only item
|
||||
hm.DeleteHostInfo(h4)
|
||||
assert.Nil(t, h4.prev)
|
||||
assert.Nil(t, h4.next)
|
||||
// h1 is fully deleted while another goroutine still holds a pointer to it.
|
||||
assert.False(t, hm.DeleteHostInfo(h1))
|
||||
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||
|
||||
// Make sure we have nil
|
||||
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
|
||||
assert.Nil(t, prim)
|
||||
// The stale promote must not bring it back.
|
||||
hm.MakePrimary(h1)
|
||||
assert.Equal(t, []uint32{2}, chainIds(t, hm, a))
|
||||
assert.Equal(t, h2, hm.QueryVpnAddr(a))
|
||||
assert.Nil(t, hm.QueryIndex(h1.localIndexId))
|
||||
}
|
||||
|
||||
// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older
|
||||
// hostinfo is still found after a newer tunnel without relay state takes primary for the same
|
||||
// address. The lookup checks the primary first and falls back to the rest of the list.
|
||||
func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
hm := newHostMap(l)
|
||||
f := &Interface{}
|
||||
relayAddr := netip.MustParseAddr("0.0.0.9")
|
||||
target := netip.MustParseAddr("0.0.0.1")
|
||||
|
||||
older := &HostInfo{
|
||||
vpnAddrs: []netip.Addr{relayAddr},
|
||||
localIndexId: 1,
|
||||
relayState: RelayState{
|
||||
relayForByAddr: map[netip.Addr]*Relay{},
|
||||
relayForByIdx: map[uint32]*Relay{},
|
||||
},
|
||||
}
|
||||
older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target})
|
||||
hm.unlockedAddHostInfo(older, f)
|
||||
|
||||
// The relay is found on the primary.
|
||||
h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, older, h)
|
||||
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||
|
||||
// A re-handshake with no relay state takes primary; the established relay on the older
|
||||
// hostinfo must still be found through the fallback.
|
||||
newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2}
|
||||
hm.unlockedAddHostInfo(newer, f)
|
||||
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr))
|
||||
|
||||
h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, older, h)
|
||||
assert.Equal(t, uint32(100), r.LocalIndex)
|
||||
|
||||
// No hostinfo at all is a plain miss.
|
||||
_, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42"))
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
|
||||
@@ -216,32 +240,82 @@ func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) {
|
||||
hm.unlockedAddHostInfo(other, f)
|
||||
hm.unlockedAddHostInfo(head, f)
|
||||
|
||||
// head is primary for both addresses, other is next in the shared chain
|
||||
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId)
|
||||
assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId)
|
||||
assert.Equal(t, other.localIndexId, head.next.localIndexId)
|
||||
assert.Equal(t, head.localIndexId, other.prev.localIndexId)
|
||||
// head is primary for both addresses, other is next in each address's list.
|
||||
assert.Equal(t, head, hm.QueryVpnAddr(a))
|
||||
assert.Equal(t, head, hm.QueryVpnAddr(b))
|
||||
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b))
|
||||
|
||||
// Delete the head. other is still live, so it must become primary for BOTH addresses.
|
||||
hm.DeleteHostInfo(head)
|
||||
assert.False(t, hm.DeleteHostInfo(head))
|
||||
assert.Equal(t, other, hm.QueryVpnAddr(a))
|
||||
assert.Equal(t, other, hm.QueryVpnAddr(b))
|
||||
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||
assert.Equal(t, []uint32{1}, chainIds(t, hm, b))
|
||||
|
||||
// Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than
|
||||
// promoted, leaving other unreachable at b.
|
||||
require.NotNil(t, hm.QueryVpnAddr(a))
|
||||
require.NotNil(t, hm.QueryVpnAddr(b))
|
||||
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId)
|
||||
assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId)
|
||||
|
||||
// other is now the only hostinfo in the chain
|
||||
assert.Nil(t, other.prev)
|
||||
assert.Nil(t, other.next)
|
||||
|
||||
// head is fully detached
|
||||
assert.Nil(t, head.prev)
|
||||
assert.Nil(t, head.next)
|
||||
// head is fully removed from the index map.
|
||||
assert.Nil(t, hm.QueryIndex(head.localIndexId))
|
||||
}
|
||||
|
||||
// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose
|
||||
// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node
|
||||
// must not promote a sibling to an address it does not own.
|
||||
func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
hm := newHostMap(l)
|
||||
f := &Interface{}
|
||||
a := netip.MustParseAddr("0.0.0.1")
|
||||
b := netip.MustParseAddr("0.0.0.2")
|
||||
|
||||
// sub owns only a; super (a newer handshake) owns a and b.
|
||||
sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||
super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
|
||||
hm.unlockedAddHostInfo(sub, f)
|
||||
hm.unlockedAddHostInfo(super, f)
|
||||
|
||||
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
|
||||
assert.Equal(t, []uint32{2}, chainIds(t, hm, b))
|
||||
|
||||
// Delete super: a promotes to sub (which owns it); b has no remaining owner and must be
|
||||
// removed, not dangled at sub (which does not own b).
|
||||
assert.False(t, hm.DeleteHostInfo(super))
|
||||
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
|
||||
assert.Empty(t, chainIds(t, hm, b))
|
||||
assert.Equal(t, sub, hm.QueryVpnAddr(a))
|
||||
assert.Nil(t, hm.QueryVpnAddr(b))
|
||||
assert.Nil(t, hm.QueryIndex(super.localIndexId))
|
||||
|
||||
// Deleting sub cleans up fully.
|
||||
assert.True(t, hm.DeleteHostInfo(sub))
|
||||
assert.Nil(t, hm.QueryVpnAddr(a))
|
||||
assertHostMapInvariants(t, hm)
|
||||
}
|
||||
|
||||
// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two
|
||||
// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one
|
||||
// of them (in Indexes but unreachable via its address); independent per-address lists cannot.
|
||||
func TestHostMap_AddDivergentOverlap(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
hm := newHostMap(l)
|
||||
f := &Interface{}
|
||||
a := netip.MustParseAddr("0.0.0.1")
|
||||
b := netip.MustParseAddr("0.0.0.2")
|
||||
|
||||
hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
|
||||
hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2}
|
||||
hm.unlockedAddHostInfo(hiA, f)
|
||||
hm.unlockedAddHostInfo(hiP, f)
|
||||
|
||||
hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3}
|
||||
hm.unlockedAddHostInfo(hiB, f)
|
||||
|
||||
assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a))
|
||||
assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b))
|
||||
// hiA is still reachable via its address (not orphaned) and still indexed.
|
||||
assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId)
|
||||
assert.NotNil(t, hm.QueryIndex(hiA.localIndexId))
|
||||
}
|
||||
|
||||
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
|
||||
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
|
||||
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
|
||||
@@ -267,32 +341,14 @@ func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) {
|
||||
|
||||
oldest := hostinfos[len(hostinfos)-1]
|
||||
|
||||
// The oldest hostinfo should have been pruned and fully detached
|
||||
assert.Nil(t, oldest.next)
|
||||
assert.Nil(t, oldest.prev)
|
||||
// The oldest hostinfo was pruned from both lists and the index map.
|
||||
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
|
||||
|
||||
// Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one)
|
||||
primA := hm.QueryVpnAddr(a)
|
||||
primB := hm.QueryVpnAddr(b)
|
||||
require.NotNil(t, primA)
|
||||
require.NotNil(t, primB)
|
||||
assert.Equal(t, primA.localIndexId, primB.localIndexId)
|
||||
assert.NotEqual(t, oldest.localIndexId, primA.localIndexId)
|
||||
|
||||
// Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent
|
||||
seen := map[uint32]struct{}{}
|
||||
for h := primA; h != nil; h = h.next {
|
||||
_, dup := seen[h.localIndexId]
|
||||
require.False(t, dup, "cycle detected in hostinfo chain")
|
||||
seen[h.localIndexId] = struct{}{}
|
||||
if h.next != nil {
|
||||
assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next")
|
||||
}
|
||||
}
|
||||
assert.Len(t, seen, MaxHostInfosPerVpnIp)
|
||||
_, prunedStillPresent := seen[oldest.localIndexId]
|
||||
assert.False(t, prunedStillPresent)
|
||||
// Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent.
|
||||
require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp)
|
||||
assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order")
|
||||
assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId)
|
||||
assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b))
|
||||
}
|
||||
|
||||
func TestHostMap_reload(t *testing.T) {
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
||||
@@ -73,7 +74,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
|
||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||
if dropReason == nil {
|
||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q, fwPacket)
|
||||
|
||||
} else {
|
||||
f.rejectInside(packet, out, q)
|
||||
@@ -87,7 +88,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
}
|
||||
|
||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
if !f.firewall.InSendReject {
|
||||
if !f.firewall.OutboundSendReject {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -103,7 +104,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
}
|
||||
|
||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
||||
if !f.firewall.OutSendReject {
|
||||
if !f.firewall.InboundSendReject {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -122,7 +123,7 @@ func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *
|
||||
return
|
||||
}
|
||||
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q, nil)
|
||||
}
|
||||
|
||||
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
||||
@@ -235,7 +236,7 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
||||
return
|
||||
}
|
||||
|
||||
f.sendNoMetrics(header.Message, st, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, p, nb, out, 0)
|
||||
f.sendNoMetrics(header.Message, st, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, p, nb, out, 0, nil)
|
||||
}
|
||||
|
||||
// SendMessageToVpnAddr handles real addr:port lookup and sends to the current best known address for vpnAddr.
|
||||
@@ -267,12 +268,12 @@ func (f *Interface) SendMessageToHostInfo(t header.MessageType, st header.Messag
|
||||
|
||||
func (f *Interface) send(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, p, nb, out []byte) {
|
||||
f.messageMetrics.Tx(t, st, 1)
|
||||
f.sendNoMetrics(t, st, ci, hostinfo, netip.AddrPort{}, p, nb, out, 0)
|
||||
f.sendNoMetrics(t, st, ci, hostinfo, netip.AddrPort{}, p, nb, out, 0, nil)
|
||||
}
|
||||
|
||||
func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte) {
|
||||
f.messageMetrics.Tx(t, st, 1)
|
||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0, nil)
|
||||
}
|
||||
|
||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||
@@ -340,10 +341,27 @@ func (f *Interface) SendVia(via *HostInfo,
|
||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||
}
|
||||
|
||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int, udpPortGetter udp.SendPortGetter) {
|
||||
if ci.eKey == nil {
|
||||
return
|
||||
}
|
||||
|
||||
multiport := f.multiPort.Tx && hostinfo.multiportTx
|
||||
rawOut := out
|
||||
if multiport {
|
||||
if len(out) < udp.RawOverhead {
|
||||
// NOTE: This is because some spots in the code send us `out[:0]`, so
|
||||
// we need to expand the slice back out to get our 8 bytes back.
|
||||
out = out[:udp.RawOverhead]
|
||||
}
|
||||
// Preserve bytes needed for the raw socket
|
||||
out = out[udp.RawOverhead:]
|
||||
|
||||
if udpPortGetter == nil {
|
||||
udpPortGetter = udp.RandomSendPort
|
||||
}
|
||||
}
|
||||
|
||||
useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid()
|
||||
fullOut := out
|
||||
|
||||
@@ -396,7 +414,13 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
}
|
||||
|
||||
if remote.IsValid() {
|
||||
err = f.writers[q].WriteTo(out, remote)
|
||||
if multiport {
|
||||
rawOut = rawOut[:len(out)+udp.RawOverhead]
|
||||
port := udpPortGetter.UDPSendPort(f.multiPort.TxPorts)
|
||||
err = f.udpRaw.WriteTo(rawOut, port, remote)
|
||||
} else {
|
||||
err = f.writers[q].WriteTo(out, remote)
|
||||
}
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
@@ -404,11 +428,17 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
)
|
||||
}
|
||||
} else if hr := hostinfo.GetRemote(); hr.IsValid() {
|
||||
err = f.writers[q].WriteTo(out, hr)
|
||||
if multiport {
|
||||
rawOut = rawOut[:len(out)+udp.RawOverhead]
|
||||
port := udpPortGetter.UDPSendPort(f.multiPort.TxPorts)
|
||||
err = f.udpRaw.WriteTo(rawOut, port, hr)
|
||||
} else {
|
||||
err = f.writers[q].WriteTo(out, hr)
|
||||
}
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
"udpAddr", hr,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
|
||||
+52
-14
@@ -99,6 +99,9 @@ type Interface struct {
|
||||
// triggerShutdown is a function that will be run exactly once, when onFatal swaps something non-nil into fatalErr
|
||||
triggerShutdown func()
|
||||
|
||||
udpRaw *udp.RawConn
|
||||
multiPort config.MultiPortConfig
|
||||
|
||||
metricHandshakes metrics.Histogram
|
||||
messageMetrics *MessageMetrics
|
||||
cachedPacketMetrics *cachedPacketMetrics
|
||||
@@ -106,6 +109,15 @@ type Interface struct {
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type MultiPortConfig struct {
|
||||
Tx bool
|
||||
Rx bool
|
||||
TxBasePort uint16
|
||||
TxPorts int
|
||||
TxHandshake bool
|
||||
TxHandshakeDelay int64
|
||||
}
|
||||
|
||||
type EncWriter interface {
|
||||
SendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
@@ -215,6 +227,9 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
|
||||
ifce.connectionManager.intf = ifce
|
||||
|
||||
// Held until Close so waiting on the interface blocks until the resources are actually released
|
||||
ifce.wg.Add(1)
|
||||
|
||||
return ifce, nil
|
||||
}
|
||||
|
||||
@@ -246,6 +261,8 @@ func (f *Interface) activate() error {
|
||||
|
||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||
|
||||
metrics.GetOrRegisterGauge("multiport.tx_ports", nil).Update(int64(f.multiPort.TxPorts))
|
||||
|
||||
// Prepare n tun queues
|
||||
var reader io.ReadWriteCloser = f.inside
|
||||
for i := 0; i < f.routines; i++ {
|
||||
@@ -258,17 +275,16 @@ func (f *Interface) activate() error {
|
||||
f.readers[i] = reader
|
||||
}
|
||||
|
||||
f.wg.Add(1) // for us to wait on Close() to return
|
||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
||||
// before releasing our resources so a waiter never observes a live context
|
||||
if err = f.inside.Activate(); err != nil {
|
||||
f.wg.Done()
|
||||
f.inside.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) run() (func() error, error) {
|
||||
func (f *Interface) run() {
|
||||
// Launch n queues to read packets from udp
|
||||
for i := 0; i < f.routines; i++ {
|
||||
f.wg.Go(func() {
|
||||
@@ -283,13 +299,14 @@ func (f *Interface) run() (func() error, error) {
|
||||
})
|
||||
}
|
||||
|
||||
return func() error {
|
||||
f.wg.Wait()
|
||||
if e := f.fatalErr.Load(); e != nil {
|
||||
return *e
|
||||
}
|
||||
return nil
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *Interface) wait() error {
|
||||
f.wg.Wait()
|
||||
if e := f.fatalErr.Load(); e != nil {
|
||||
return *e
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
||||
@@ -322,7 +339,10 @@ func (f *Interface) listenOut(i int) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||
})
|
||||
|
||||
if err != nil && !f.closed.Load() {
|
||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
||||
// reacting to it, like the user device pipes
|
||||
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
|
||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||
f.onFatal(err)
|
||||
}
|
||||
@@ -341,7 +361,8 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||
for {
|
||||
n, err := reader.Read(packet)
|
||||
if err != nil {
|
||||
if !f.closed.Load() {
|
||||
// Same shutdown noise handling as listenOut
|
||||
if !f.closed.Load() && f.ctx.Err() == nil {
|
||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||
f.onFatal(err)
|
||||
}
|
||||
@@ -498,6 +519,8 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||
|
||||
udpStats := udp.NewUDPStatsEmitter(f.writers)
|
||||
|
||||
var rawStats func()
|
||||
|
||||
certExpirationGauge := metrics.GetOrRegisterGauge("certificate.ttl_seconds", nil)
|
||||
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
||||
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
||||
@@ -512,6 +535,13 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
||||
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
||||
|
||||
if f.udpRaw != nil {
|
||||
if rawStats == nil {
|
||||
rawStats = udp.NewRawStatsEmitter(f.udpRaw)
|
||||
}
|
||||
rawStats()
|
||||
}
|
||||
|
||||
// Report the max certificate version we are capable of using
|
||||
if certState.v2Cert != nil {
|
||||
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
||||
@@ -542,9 +572,15 @@ func (f *Interface) GetCertState() *CertState {
|
||||
return f.pki.getCertState()
|
||||
}
|
||||
|
||||
// Close releases the interface's resources: the udp sockets and the tun device.
|
||||
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
|
||||
// calls after the first return nil without doing anything.
|
||||
func (f *Interface) Close() error {
|
||||
if !f.closed.CompareAndSwap(false, true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
var errs []error
|
||||
f.closed.Store(true)
|
||||
|
||||
// Release the udp readers
|
||||
for i, u := range f.writers {
|
||||
@@ -560,6 +596,8 @@ func (f *Interface) Close() error {
|
||||
if closeErr != nil {
|
||||
errs = append(errs, closeErr)
|
||||
}
|
||||
|
||||
// Release the construction token so waiters know the resources are gone
|
||||
f.wg.Done()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
+9
-1
@@ -36,6 +36,10 @@ type LightHouse struct {
|
||||
myVpnNetworksTable *bart.Lite
|
||||
punchy *Punchy
|
||||
|
||||
// localAddrsFn enumerates the underlay addresses we advertise. It is a field so tests can supply simulated
|
||||
// addresses rather than whatever this machine's NICs happen to be. Set it before Start.
|
||||
localAddrsFn func(*LocalAllowList) []netip.Addr
|
||||
|
||||
// Local cache of answers from light houses
|
||||
// map of vpn addr to answers
|
||||
addrMap map[netip.Addr]*RemoteList
|
||||
@@ -107,6 +111,10 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||
l: l,
|
||||
}
|
||||
h.localAddrsFn = func(al *LocalAllowList) []netip.Addr {
|
||||
return localAddrs(h.l, al)
|
||||
}
|
||||
|
||||
lighthouses := make([]netip.Addr, 0)
|
||||
h.lighthouses.Store(&lighthouses)
|
||||
staticList := make(map[netip.Addr]struct{})
|
||||
@@ -918,7 +926,7 @@ func (lh *LightHouse) SendUpdate() {
|
||||
}
|
||||
|
||||
lal := lh.GetLocalAllowList()
|
||||
for _, e := range localAddrs(lh.l, lal) {
|
||||
for _, e := range lh.localAddrsFn(lal) {
|
||||
if lh.myVpnNetworksTable.Contains(e) {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -130,6 +130,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
udpConns := make([]udp.Conn, routines)
|
||||
port := c.GetInt("listen.port", 0)
|
||||
|
||||
// Callers get no handle to these until the Control is returned, release them on any error.
|
||||
defer func() {
|
||||
if reterr != nil {
|
||||
for _, u := range udpConns {
|
||||
if u != nil {
|
||||
_ = u.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
if !configTest {
|
||||
rawListenHost := c.GetString("listen.host", "0.0.0.0")
|
||||
var listenHost netip.Addr
|
||||
@@ -233,6 +244,39 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
ifce.writers = udpConns
|
||||
lightHouse.ifce = ifce
|
||||
|
||||
loadMultiPortConfig := func(c *config.C) {
|
||||
ifce.multiPort.Rx = c.GetBool("tun.multiport.rx_enabled", false)
|
||||
|
||||
tx := c.GetBool("tun.multiport.tx_enabled", false)
|
||||
|
||||
if tx && ifce.udpRaw == nil {
|
||||
ifce.udpRaw, err = udp.NewRawConn(l, c.GetString("listen.host", "0.0.0.0"), port, uint16(port))
|
||||
if err != nil {
|
||||
l.Error("Failed to get raw socket for tun.multiport.tx_enabled", "error", err)
|
||||
ifce.udpRaw = nil
|
||||
tx = false
|
||||
}
|
||||
}
|
||||
|
||||
if tx {
|
||||
ifce.multiPort.TxBasePort = uint16(port)
|
||||
ifce.multiPort.TxPorts = c.GetInt("tun.multiport.tx_ports", 100)
|
||||
ifce.multiPort.TxHandshake = c.GetBool("tun.multiport.tx_handshake", false)
|
||||
ifce.multiPort.TxHandshakeDelay = int64(c.GetInt("tun.multiport.tx_handshake_delay", 2))
|
||||
ifce.udpRaw.ReloadConfig(c)
|
||||
}
|
||||
ifce.multiPort.Tx = tx
|
||||
|
||||
// TODO: if we upstream this, make this cleaner
|
||||
handshakeManager.udpRaw = ifce.udpRaw
|
||||
handshakeManager.multiPort = ifce.multiPort
|
||||
|
||||
l.Info("Multiport configured", "multiPort", ifce.multiPort)
|
||||
}
|
||||
|
||||
loadMultiPortConfig(c)
|
||||
c.RegisterReloadCallback(loadMultiPortConfig)
|
||||
|
||||
ifce.RegisterConfigChangeCallbacks(c)
|
||||
ifce.reloadDisconnectInvalid(c)
|
||||
ifce.reloadSendRecvError(c)
|
||||
@@ -257,6 +301,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
|
||||
attachCommands(l, c, ssh, ifce)
|
||||
|
||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
return &Control{
|
||||
state: StateReady,
|
||||
f: ifce,
|
||||
@@ -267,6 +313,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
statsStart: stats.Start,
|
||||
dnsStart: ds.Start,
|
||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||
networkChangeStart: networkChanges.Start,
|
||||
connectionManagerStart: connManager.Start,
|
||||
}, nil
|
||||
}
|
||||
|
||||
+33
-53
@@ -102,27 +102,31 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
return
|
||||
}
|
||||
|
||||
if len(packet) < header.Len+hostinfo.ConnectionState.dKey.Overhead() {
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("packet too small", "from", via, "length", len(packet))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// All remaining packets are encrypted
|
||||
ci := hostinfo.ConnectionState
|
||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||
return
|
||||
}
|
||||
|
||||
// Relay packets are special
|
||||
if isMessageRelay {
|
||||
// Relay packets are special, this branch should always early-return
|
||||
if err = hostinfo.ConnectionState.VerifyRelay(f.l, h.MessageCounter, packet, nb); err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to verify relay packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
out, err = hostinfo.ConnectionState.Decrypt(f.l, h.MessageCounter, out, packet, nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
|
||||
"error", err,
|
||||
"from", via,
|
||||
"header", h,
|
||||
)
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet", "error", err, "from", via, "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -151,7 +155,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||
case header.TestRequest:
|
||||
//recycle the input packet ciphertext as our output buffer
|
||||
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
|
||||
f.send(header.Test, header.TestReply, hostinfo.ConnectionState, hostinfo, out, nb, packet)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||
return
|
||||
@@ -170,27 +174,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// The entire body is sent as AD, not encrypted.
|
||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Advance the replay window now that the frame is authenticated
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Successfully validated the thing. Get rid of the Relay header.
|
||||
signedPayload = signedPayload[header.Len:]
|
||||
// Successfully validated the thing. Get rid of the Relay header and the AEAD tag
|
||||
signedPayload := packet[header.Len : len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||
@@ -214,7 +199,6 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
via = ViaSender{
|
||||
UdpAddr: via.UdpAddr,
|
||||
relayHI: hostinfo,
|
||||
remoteIdx: relay.RemoteIndex,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
@@ -235,9 +219,10 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
||||
if targetRelay.State == Established {
|
||||
switch targetRelay.Type {
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel
|
||||
// Find the target HostInfo
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||
// Forward this packet through the relay tunnel, rebuilding it in place.
|
||||
// Encode overwrites the old outer header, and the new AEAD tag lands where the old one was
|
||||
fwdBuf := packet[:0:len(packet)] // Cap to len(packet) to protect memory from a larger parent buffer
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, fwdBuf, true)
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
return
|
||||
@@ -279,6 +264,15 @@ func (f *Interface) sendCloseTunnel(h *HostInfo) {
|
||||
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
||||
curRemote := hostinfo.GetRemote()
|
||||
if !via.IsRelayed && curRemote != via.UdpAddr {
|
||||
if hostinfo.multiportRx {
|
||||
// If the remote is sending with multiport, we aren't roaming unless
|
||||
// the IP has changed
|
||||
if curRemote.Addr().Compare(via.UdpAddr.Addr()) == 0 {
|
||||
return
|
||||
}
|
||||
// Keep the port from the original hostinfo, because the remote is transmitting from multiport ports
|
||||
via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), curRemote.Port())
|
||||
}
|
||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr)
|
||||
@@ -504,20 +498,6 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
||||
return nil, ErrOutOfWindow
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||
err := newPacket(out, true, fwPacket)
|
||||
if err != nil {
|
||||
|
||||
@@ -40,6 +40,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
||||
|
||||
err := t.reload(c, true)
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
|
||||
@@ -659,7 +659,6 @@ func addRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
||||
return fmt.Errorf("failed to create route.RouteMessage for change: %w", err)
|
||||
}
|
||||
_, err = unix.Write(sock, data[:])
|
||||
fmt.Println("DOING CHANGE")
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err)
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
"github.com/slackhq/nebula/util"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
type tun struct {
|
||||
@@ -33,6 +34,12 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error)
|
||||
}
|
||||
|
||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
if err := unix.SetNonblock(deviceFd, true); err != nil {
|
||||
// We own the fd from the moment it is handed to us, same as the reload error path below
|
||||
_ = unix.Close(deviceFd)
|
||||
return nil, fmt.Errorf("failed to set the tun fd to non-blocking mode: %w", err)
|
||||
}
|
||||
|
||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||
t := &tun{
|
||||
vpnNetworks: vpnNetworks,
|
||||
@@ -42,6 +49,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
||||
|
||||
err := t.reload(c, true)
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
|
||||
+9
-1
@@ -107,7 +107,10 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
||||
if relayHostInfo.GetRemote().IsValid() {
|
||||
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
||||
if err != nil {
|
||||
// No local relay state was installed, so a CreateRelayRequest would hand the
|
||||
// peer an index we could never resolve. Skip it.
|
||||
hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
m := NebulaControl{
|
||||
@@ -237,7 +240,12 @@ func AddRelay(l *slog.Logger, relayHostInfo *HostInfo, hm *HostMap, vpnIp netip.
|
||||
// Avoid standing up a relay that can't be used since only the primary hostinfo
|
||||
// will be pointed to by the relay logic
|
||||
//TODO: if there was an existing primary and it had relay state, should we merge?
|
||||
hm.unlockedMakePrimary(relayHostInfo)
|
||||
if !hm.unlockedMakePrimary(relayHostInfo) {
|
||||
// The tunnel was torn down after the caller grabbed relayHostInfo. A relay standing
|
||||
// on an unlinked hostinfo would never carry traffic, and its Relays entry could
|
||||
// never be reclaimed since the delete-time cleanup has already run.
|
||||
return 0, errors.New("relay hostinfo is no longer in the hostmap")
|
||||
}
|
||||
|
||||
hm.Relays[index] = relayHostInfo
|
||||
newRelay := Relay{
|
||||
|
||||
+16
-8
@@ -43,12 +43,25 @@ type Service struct {
|
||||
}
|
||||
}
|
||||
|
||||
func New(control *nebula.Control) (*Service, error) {
|
||||
wait, err := control.Start()
|
||||
func New(control *nebula.Control) (_ *Service, reterr error) {
|
||||
// Check this before Start so a failure doesn't leave a running nebula
|
||||
device, ok := control.Device().(*overlay.UserDevice)
|
||||
if !ok {
|
||||
return nil, errors.New("must be using user device")
|
||||
}
|
||||
|
||||
err := control.Start()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Anything that fails after a successful Start must tear nebula back down
|
||||
defer func() {
|
||||
if reterr != nil {
|
||||
control.Stop()
|
||||
}
|
||||
}()
|
||||
|
||||
ctx := control.Context()
|
||||
eg, ctx := errgroup.WithContext(ctx)
|
||||
s := Service{
|
||||
@@ -57,11 +70,6 @@ func New(control *nebula.Control) (*Service, error) {
|
||||
}
|
||||
s.mu.listeners = map[uint16]*tcpListener{}
|
||||
|
||||
device, ok := control.Device().(*overlay.UserDevice)
|
||||
if !ok {
|
||||
return nil, errors.New("must be using user device")
|
||||
}
|
||||
|
||||
s.ipstack = stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||
TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol4, icmp.NewProtocol6},
|
||||
@@ -147,7 +155,7 @@ func New(control *nebula.Control) (*Service, error) {
|
||||
// Add the nebula wait function to the group so a fatal reader error
|
||||
// propagates out through errgroup.Wait().
|
||||
eg.Go(func() error {
|
||||
return wait()
|
||||
return control.Wait()
|
||||
})
|
||||
|
||||
return &s, nil
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
)
|
||||
|
||||
// NetworkChangeMonitor rebinds the udp listener when the local network moves out from under it.
|
||||
//
|
||||
// Detection lives here in the udp package, next to the socket it concerns and the platform matrix that already knows
|
||||
// which sockets go stale. What to do about a change — updating the lighthouse, requerying tunnels — is not the udp
|
||||
// package's business, so Start takes the reaction as a plain function. Passing it at Start rather than holding it
|
||||
// keeps this package from referencing whatever owns the rebind.
|
||||
//
|
||||
// On platforms whose sockets do not go stale, watchNetworkChanges hands back a nil channel and Start returns.
|
||||
type NetworkChangeMonitor struct {
|
||||
l *slog.Logger
|
||||
ctx context.Context
|
||||
enabled bool
|
||||
}
|
||||
|
||||
// NewNetworkChangeMonitor builds a monitor for local network changes. The returned monitor is always usable: Start
|
||||
// is safe to call unconditionally, it no-ops when disabled or on a platform that does not need it.
|
||||
func NewNetworkChangeMonitor(ctx context.Context, l *slog.Logger, c *config.C) *NetworkChangeMonitor {
|
||||
return &NetworkChangeMonitor{
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
enabled: c.GetBool("listen.rebind_on_network_change", true),
|
||||
}
|
||||
}
|
||||
|
||||
// Start watches for network changes until the context is cancelled, calling rebind once per settled change. It
|
||||
// blocks, so callers run it in a goroutine, and it no-ops when disabled, unsupported, or with nothing to rebind.
|
||||
func (m *NetworkChangeMonitor) Start(rebind func()) {
|
||||
if !m.enabled || rebind == nil || m.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
|
||||
changes, err := watchNetworkChanges(m.ctx, m.l)
|
||||
if err != nil {
|
||||
// Not fatal. Everything else still works, we just won't notice a network change on our own.
|
||||
m.l.Error("Failed to watch for network changes, will not rebind the udp listener when the network moves",
|
||||
"error", err,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if changes == nil {
|
||||
// This platform's sockets don't go stale, so there is nothing to watch for.
|
||||
return
|
||||
}
|
||||
|
||||
m.l.Info("Watching for network changes to rebind the udp listener")
|
||||
|
||||
for range changes {
|
||||
m.l.Info("Local network changed, rebinding the udp listener")
|
||||
rebind()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
//go:build darwin && !ios && !e2e_testing
|
||||
// +build darwin,!ios,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
const (
|
||||
// netChangeSettleWindow is how long we keep swallowing routing messages after the first interesting one. A
|
||||
// single network change is never a single message, it is a burst: the link drops, addresses go away, new ones
|
||||
// arrive, routes get rewritten. Reporting part way through that just means reporting again.
|
||||
netChangeSettleWindow = time.Second
|
||||
|
||||
// netChangeReadBuffer is sized well past any rt_msghdr plus its addresses. A short read would be discarded by
|
||||
// the kernel, so being generous here is how we avoid missing a message.
|
||||
netChangeReadBuffer = 4096
|
||||
)
|
||||
|
||||
// watchNetworkChanges reports when the local network moves out from under us, so the listener can be rebound.
|
||||
//
|
||||
// Darwin scopes a udp socket to whatever interface it came up on. Move between networks and we keep sending out an
|
||||
// interface that no longer has a route, which surfaces as an instant "no route to host" with no packet ever leaving
|
||||
// the box. Rebind clears that, but only if something notices the change and calls it. iOS has always been told by
|
||||
// the host app off NWPathMonitor. This is the equivalent for everything else that runs on darwin.
|
||||
//
|
||||
// The returned channel is buffered and coalescing: a send is dropped if one is already pending, since both mean the
|
||||
// same thing to a reader. It is closed when ctx is cancelled or the routing socket fails, so a caller can simply
|
||||
// range over it. Platforms whose sockets do not need rebinding return a nil channel and no error.
|
||||
func watchNetworkChanges(ctx context.Context, l *slog.Logger) (<-chan struct{}, error) {
|
||||
sock, err := openRouteSocket()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
go func() {
|
||||
defer close(changes)
|
||||
defer func() { _ = sock.Close() }()
|
||||
|
||||
// Closing the socket is what unblocks the read in watchRouteSocket, so this turns cancellation into a
|
||||
// close. It is scoped to this call so it cannot outlive the watch it belongs to.
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = sock.Close()
|
||||
case <-done:
|
||||
}
|
||||
}()
|
||||
|
||||
watchRouteSocket(l, sock, changes)
|
||||
}()
|
||||
|
||||
return changes, nil
|
||||
}
|
||||
|
||||
// watchRouteSocket blocks reading the routing socket, reporting once per settled burst of changes. It returns when
|
||||
// the socket is closed, which is how cancellation gets us out of here.
|
||||
func watchRouteSocket(l *slog.Logger, sock *os.File, changes chan<- struct{}) {
|
||||
buf := make([]byte, netChangeReadBuffer)
|
||||
|
||||
for {
|
||||
n, err := sock.Read(buf)
|
||||
if err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
if !isNetworkChange(buf[:n]) {
|
||||
continue
|
||||
}
|
||||
|
||||
// Swallow the rest of the burst. The deadline is absolute and not extended by what arrives, so this always
|
||||
// ends after the settle window no matter how chatty the socket is. Changes that land after the window
|
||||
// simply produce another report, which is the correct outcome anyway.
|
||||
deadline := time.Now().Add(netChangeSettleWindow)
|
||||
for {
|
||||
if err = sock.SetReadDeadline(deadline); err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
if _, err = sock.Read(buf); err != nil {
|
||||
if os.IsTimeout(err) {
|
||||
break
|
||||
}
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err = sock.SetReadDeadline(time.Time{}); err != nil {
|
||||
logRouteSocketError(l, err)
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case changes <- struct{}{}:
|
||||
default:
|
||||
// One already pending, and a second "the network moved" tells the reader nothing new.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// logRouteSocketError reports a routing socket failure unless it is just us shutting the socket down.
|
||||
func logRouteSocketError(l *slog.Logger, err error) {
|
||||
if errors.Is(err, os.ErrClosed) {
|
||||
return
|
||||
}
|
||||
|
||||
l.Error("Error reading the routing socket, will no longer notice local network changes", "error", err)
|
||||
}
|
||||
|
||||
// openRouteSocket returns the routing socket as a non blocking os.File. Going through os.File puts reads on the go
|
||||
// poller, which buys us both a working read deadline and a Close that unblocks a read in progress.
|
||||
func openRouteSocket() (*os.File, error) {
|
||||
fd, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err = unix.SetNonblock(fd, true); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return os.NewFile(uintptr(fd), "route"), nil
|
||||
}
|
||||
|
||||
// isNetworkChange reports whether a routing message means our local addressing may have moved out from under us.
|
||||
//
|
||||
// We read the header instead of parsing the message because the type is the only part we need, and a full parse can
|
||||
// fail on shapes we don't care about, which would turn "a message I can't parse" into "a change I missed".
|
||||
// rt_msghdr, if_msghdr and ifa_msghdr all begin with the same three fields, so this is the same for every type.
|
||||
func isNetworkChange(msg []byte) bool {
|
||||
if len(msg) < 4 {
|
||||
return false
|
||||
}
|
||||
|
||||
// u_short msglen, u_char version, u_char type
|
||||
if int(binary.NativeEndian.Uint16(msg[0:2])) > len(msg) || msg[2] != unix.RTM_VERSION {
|
||||
return false
|
||||
}
|
||||
|
||||
switch msg[3] {
|
||||
case unix.RTM_NEWADDR, unix.RTM_DELADDR, unix.RTM_IFINFO:
|
||||
// An address arrived or left, or a link changed state. Anything else on this socket is either a route
|
||||
// churning underneath us, which a rebind doesn't help with, or unrelated traffic.
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
//go:build darwin && !ios && !e2e_testing
|
||||
// +build darwin,!ios,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// routeMsg builds the first four bytes of a routing message, which is all isNetworkChange reads.
|
||||
func routeMsg(msgType uint8, extra int) []byte {
|
||||
msg := make([]byte, 4+extra)
|
||||
binary.NativeEndian.PutUint16(msg[0:2], uint16(len(msg)))
|
||||
msg[2] = unix.RTM_VERSION
|
||||
msg[3] = msgType
|
||||
return msg
|
||||
}
|
||||
|
||||
func TestIsNetworkChange(t *testing.T) {
|
||||
// The three that mean our addressing may have moved
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_NEWADDR, 0)))
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_DELADDR, 0)))
|
||||
assert.True(t, isNetworkChange(routeMsg(unix.RTM_IFINFO, 0)))
|
||||
|
||||
// Route churn is not something a rebind helps with
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_ADD, 0)))
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_DELETE, 0)))
|
||||
assert.False(t, isNetworkChange(routeMsg(unix.RTM_GET, 0)))
|
||||
|
||||
// Garbage must not be mistaken for a change
|
||||
assert.False(t, isNetworkChange(nil), "empty")
|
||||
assert.False(t, isNetworkChange([]byte{0, 0, 0}), "short header")
|
||||
|
||||
wrongVersion := routeMsg(unix.RTM_NEWADDR, 0)
|
||||
wrongVersion[2] = unix.RTM_VERSION + 1
|
||||
assert.False(t, isNetworkChange(wrongVersion), "wrong rtm_version")
|
||||
|
||||
lying := routeMsg(unix.RTM_NEWADDR, 0)
|
||||
binary.NativeEndian.PutUint16(lying[0:2], 512)
|
||||
assert.False(t, isNetworkChange(lying), "msglen longer than what we read")
|
||||
}
|
||||
|
||||
// socketPair returns a connected pair of datagram sockets, the first wrapped the same way the routing socket is. It
|
||||
// stands in for the kernel so the watch loop can be driven with synthetic messages.
|
||||
func socketPair(t *testing.T) (*os.File, int) {
|
||||
t.Helper()
|
||||
|
||||
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, unix.SetNonblock(fds[0], true))
|
||||
|
||||
f := os.NewFile(uintptr(fds[0]), "route")
|
||||
t.Cleanup(func() {
|
||||
_ = f.Close()
|
||||
_ = unix.Close(fds[1])
|
||||
})
|
||||
|
||||
return f, fds[1]
|
||||
}
|
||||
|
||||
func TestWatchRouteSocketCoalescesABurst(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// One network change is a burst of messages. All of these land inside the settle window, so they must produce
|
||||
// exactly one report rather than one apiece.
|
||||
for range 5 {
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
// Uninteresting messages in the middle of a burst must not add a report of their own either.
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_ADD, 8))
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-changes:
|
||||
case <-time.After(netChangeSettleWindow * 4):
|
||||
t.Fatal("a burst should have reported a change")
|
||||
}
|
||||
|
||||
// Nothing more from that burst
|
||||
select {
|
||||
case <-changes:
|
||||
t.Fatal("a burst should report exactly once")
|
||||
case <-time.After(netChangeSettleWindow):
|
||||
}
|
||||
|
||||
// A change after the window has closed is a separate event and gets its own report.
|
||||
_, err = unix.Write(kernel, routeMsg(unix.RTM_IFINFO, 8))
|
||||
require.NoError(t, err)
|
||||
select {
|
||||
case <-changes:
|
||||
case <-time.After(netChangeSettleWindow * 4):
|
||||
t.Fatal("a later change should report again")
|
||||
}
|
||||
|
||||
// Closing the socket is how the real thing shuts down
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWatchRouteSocketIgnoresUninterestingMessages(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
for _, msgType := range []uint8{unix.RTM_ADD, unix.RTM_DELETE, unix.RTM_GET, unix.RTM_MISS} {
|
||||
_, err := unix.Write(kernel, routeMsg(msgType, 8))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-changes:
|
||||
t.Fatal("route churn alone must not report a change")
|
||||
case <-time.After(netChangeSettleWindow * 2):
|
||||
}
|
||||
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket did not return after the socket was closed")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWatchRouteSocketDropsRatherThanBlocks covers the coalescing send. A reader that is busy rebinding must not
|
||||
// wedge the watcher, and a second pending "the network moved" tells it nothing new anyway.
|
||||
func TestWatchRouteSocketDropsRatherThanBlocks(t *testing.T) {
|
||||
sock, kernel := socketPair(t)
|
||||
changes := make(chan struct{}, 1)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
watchRouteSocket(test.NewLogger(), sock, changes)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Nobody is reading changes, so after the first report the buffer is full for the rest of this test
|
||||
for range 3 {
|
||||
_, err := unix.Write(kernel, routeMsg(unix.RTM_NEWADDR, 8))
|
||||
require.NoError(t, err)
|
||||
time.Sleep(netChangeSettleWindow + time.Millisecond*250)
|
||||
}
|
||||
|
||||
// The watcher must still be alive and responsive to a close
|
||||
require.NoError(t, sock.Close())
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("watchRouteSocket wedged on a full channel")
|
||||
}
|
||||
|
||||
assert.Len(t, changes, 1, "the pending report should have coalesced, not queued")
|
||||
}
|
||||
|
||||
// TestWatchNetworkChangesStopsWithContext covers the detection path against a real routing socket, including that
|
||||
// cancelling the context closes the channel so a ranging caller falls out of its loop.
|
||||
func TestWatchNetworkChangesStopsWithContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
changes, err := watchNetworkChanges(ctx, test.NewLogger())
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, changes, "darwin should support watching")
|
||||
|
||||
drained := make(chan struct{})
|
||||
go func() {
|
||||
for range changes {
|
||||
}
|
||||
close(drained)
|
||||
}()
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-drained:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("cancelling the context should close the changes channel")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetworkChangeMonitorStopsWithContext drives the whole monitor against a real routing socket: Start must block
|
||||
// watching, and cancelling the context (which is all Control does on shutdown, it never stops the monitor directly)
|
||||
// must return it and clean up the watch goroutines.
|
||||
func TestNetworkChangeMonitorStopsWithContext(t *testing.T) {
|
||||
// IgnoreCurrent because other tests in this package leave readers running; we only care about what this test
|
||||
// leaks itself.
|
||||
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.LoadString("listen:\n rebind_on_network_change: true\n"))
|
||||
m := NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
m.Start(func() {})
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Start should be sitting on the routing socket, not have fallen out. If it returned early it either failed to
|
||||
// watch or no-op'd, both of which we want to catch.
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("Start returned instead of watching")
|
||||
case <-time.After(time.Millisecond * 250):
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second * 5):
|
||||
t.Fatal("Start did not return after the context was cancelled")
|
||||
}
|
||||
|
||||
// Starting again after the context is dead must not open anything.
|
||||
m.Start(func() {})
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
//go:build !darwin || ios || e2e_testing
|
||||
// +build !darwin ios e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
// watchNetworkChanges is a no-op outside of darwin.
|
||||
//
|
||||
// Darwin is the platform that scopes a udp socket to the interface it came up on, so it is the platform whose socket
|
||||
// goes stale when the local network changes. Everywhere else Rebind has nothing to do, so there is nothing to watch
|
||||
// for. iOS is excluded on purpose even though it is darwin: the host app already drives the rebind off NWPathMonitor,
|
||||
// and two things racing to rebind the same socket is worse than one.
|
||||
//
|
||||
// A nil channel means "not supported here", which callers must treat as "do not start a watcher" rather than
|
||||
// selecting on it, since a receive from a nil channel blocks forever.
|
||||
func watchNetworkChanges(_ context.Context, _ *slog.Logger) (<-chan struct{}, error) {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newMonitor(t *testing.T, ctx context.Context, cfg string) *NetworkChangeMonitor {
|
||||
t.Helper()
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
require.NoError(t, c.LoadString(cfg))
|
||||
return NewNetworkChangeMonitor(ctx, l, c)
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorDefaultsOn(t *testing.T) {
|
||||
// Says nothing about rebinding, so this covers the default.
|
||||
m := newMonitor(t, context.Background(), "listen:\n host: 0.0.0.0\n")
|
||||
assert.True(t, m.enabled, "should default to on")
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorDisabledIsANoOp(t *testing.T) {
|
||||
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: false\n")
|
||||
require.False(t, m.enabled)
|
||||
|
||||
// Must return without opening a socket. If it watched anything this would block.
|
||||
m.Start(func() {})
|
||||
}
|
||||
|
||||
func TestNetworkChangeMonitorNilRebindIsANoOp(t *testing.T) {
|
||||
// Nothing to rebind, so there is no point watching, on any platform.
|
||||
m := newMonitor(t, context.Background(), "listen:\n rebind_on_network_change: true\n")
|
||||
m.Start(nil)
|
||||
}
|
||||
+4
-5
@@ -187,6 +187,9 @@ func (u *StdConn) SupportsMultipleReaders() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// Rebind clears the interface the kernel scoped this socket to, so that sends are routed against the current
|
||||
// routing table instead of the interface we happened to be on when the socket was created. Darwin pins sockets
|
||||
// this way on its own, which is what strands us after the underlying network changes.
|
||||
func (u *StdConn) Rebind() error {
|
||||
var err error
|
||||
if u.isV4 {
|
||||
@@ -195,9 +198,5 @@ func (u *StdConn) Rebind() error {
|
||||
err = syscall.SetsockoptInt(int(u.sysFd), syscall.IPPROTO_IPV6, syscall.IPV6_BOUND_IF, 0)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
u.l.Error("Failed to rebind udp socket", "error", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
return err
|
||||
}
|
||||
|
||||
+168
-167
@@ -4,12 +4,13 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
@@ -19,58 +20,51 @@ import (
|
||||
)
|
||||
|
||||
type StdConn struct {
|
||||
udpConn *net.UDPConn
|
||||
rawConn syscall.RawConn
|
||||
isV4 bool
|
||||
l *slog.Logger
|
||||
batch int
|
||||
}
|
||||
|
||||
func setReusePort(network, address string, c syscall.RawConn) error {
|
||||
var opErr error
|
||||
err := c.Control(func(fd uintptr) {
|
||||
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1)
|
||||
//CloseOnExec already set by the runtime
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return opErr
|
||||
sysFd int
|
||||
closed atomic.Bool
|
||||
isV4 bool
|
||||
l *slog.Logger
|
||||
batch int
|
||||
}
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||
listen := netip.AddrPortFrom(ip, uint16(port))
|
||||
lc := net.ListenConfig{}
|
||||
af := unix.AF_INET6
|
||||
if ip.Is4() {
|
||||
af = unix.AF_INET
|
||||
}
|
||||
syscall.ForkLock.RLock()
|
||||
fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP)
|
||||
if err == nil {
|
||||
unix.CloseOnExec(fd)
|
||||
}
|
||||
syscall.ForkLock.RUnlock()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to open socket: %w", err)
|
||||
}
|
||||
|
||||
if multi {
|
||||
lc.Control = setReusePort
|
||||
}
|
||||
//this context is only used during the bind operation, you can't cancel it to kill the socket
|
||||
pc, err := lc.ListenPacket(context.Background(), "udp", listen.String())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unable to open socket: %s", err)
|
||||
}
|
||||
udpConn := pc.(*net.UDPConn)
|
||||
rawConn, err := udpConn.SyscallConn()
|
||||
if err != nil {
|
||||
_ = udpConn.Close()
|
||||
return nil, err
|
||||
}
|
||||
//gotta find out if we got an AF_INET6 socket or not:
|
||||
out := &StdConn{
|
||||
udpConn: udpConn,
|
||||
rawConn: rawConn,
|
||||
l: l,
|
||||
batch: batch,
|
||||
if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, fmt.Errorf("unable to set SO_REUSEPORT: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
af, err := out.getSockOptInt(unix.SO_DOMAIN)
|
||||
if err != nil {
|
||||
_ = out.Close()
|
||||
return nil, err
|
||||
var sa unix.Sockaddr
|
||||
if ip.Is4() {
|
||||
sa4 := &unix.SockaddrInet4{Port: port}
|
||||
sa4.Addr = ip.As4()
|
||||
sa = sa4
|
||||
} else {
|
||||
sa6 := &unix.SockaddrInet6{Port: port}
|
||||
sa6.Addr = ip.As16()
|
||||
sa = sa6
|
||||
}
|
||||
if err = unix.Bind(fd, sa); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, fmt.Errorf("unable to bind to socket: %w", err)
|
||||
}
|
||||
out.isV4 = af == unix.AF_INET
|
||||
|
||||
return out, nil
|
||||
return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||
@@ -81,134 +75,111 @@ func (u *StdConn) Rebind() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *StdConn) getSockOptInt(opt int) (int, error) {
|
||||
if u.rawConn == nil {
|
||||
return 0, fmt.Errorf("no UDP connection")
|
||||
}
|
||||
var out int
|
||||
var opErr error
|
||||
err := u.rawConn.Control(func(fd uintptr) {
|
||||
out, opErr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt)
|
||||
})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return out, opErr
|
||||
}
|
||||
|
||||
func (u *StdConn) setSockOptInt(opt int, n int) error {
|
||||
if u.rawConn == nil {
|
||||
return fmt.Errorf("no UDP connection")
|
||||
}
|
||||
var opErr error
|
||||
err := u.rawConn.Control(func(fd uintptr) {
|
||||
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, opt, n)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return opErr
|
||||
}
|
||||
|
||||
func (u *StdConn) SetRecvBuffer(n int) error {
|
||||
return u.setSockOptInt(unix.SO_RCVBUFFORCE, n)
|
||||
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n)
|
||||
}
|
||||
|
||||
func (u *StdConn) SetSendBuffer(n int) error {
|
||||
return u.setSockOptInt(unix.SO_SNDBUFFORCE, n)
|
||||
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
|
||||
}
|
||||
|
||||
func (u *StdConn) SetSoMark(mark int) error {
|
||||
return u.setSockOptInt(unix.SO_MARK, mark)
|
||||
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK, mark)
|
||||
}
|
||||
|
||||
func (u *StdConn) GetRecvBuffer() (int, error) {
|
||||
return u.getSockOptInt(unix.SO_RCVBUF)
|
||||
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUF)
|
||||
}
|
||||
|
||||
func (u *StdConn) GetSendBuffer() (int, error) {
|
||||
return u.getSockOptInt(unix.SO_SNDBUF)
|
||||
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUF)
|
||||
}
|
||||
|
||||
func (u *StdConn) GetSoMark() (int, error) {
|
||||
return u.getSockOptInt(unix.SO_MARK)
|
||||
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK)
|
||||
}
|
||||
|
||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||
a := u.udpConn.LocalAddr()
|
||||
|
||||
switch v := a.(type) {
|
||||
case *net.UDPAddr:
|
||||
addr, ok := netip.AddrFromSlice(v.IP)
|
||||
if !ok {
|
||||
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned invalid IP address: %s", v.IP)
|
||||
}
|
||||
return netip.AddrPortFrom(addr, uint16(v.Port)), nil
|
||||
|
||||
sa, err := unix.Getsockname(u.sysFd)
|
||||
if err != nil {
|
||||
return netip.AddrPort{}, err
|
||||
}
|
||||
switch sa := sa.(type) {
|
||||
case *unix.SockaddrInet4:
|
||||
return netip.AddrPortFrom(netip.AddrFrom4(sa.Addr), uint16(sa.Port)), nil
|
||||
case *unix.SockaddrInet6:
|
||||
return netip.AddrPortFrom(netip.AddrFrom16(sa.Addr), uint16(sa.Port)), nil
|
||||
default:
|
||||
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned: %#v", a)
|
||||
return netip.AddrPort{}, fmt.Errorf("unsupported sock type: %T", sa)
|
||||
}
|
||||
}
|
||||
|
||||
func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
|
||||
var errno syscall.Errno
|
||||
n, _, errno := unix.Syscall6(
|
||||
// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams
|
||||
func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) {
|
||||
r, _, errno := unix.Syscall6(
|
||||
unix.SYS_RECVMMSG,
|
||||
fd,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&msgs[0])),
|
||||
uintptr(len(msgs)),
|
||||
unix.MSG_WAITFORONE,
|
||||
0,
|
||||
0,
|
||||
)
|
||||
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
||||
// No data available, block for I/O and try again.
|
||||
return int(n), false, nil
|
||||
}
|
||||
if errno != 0 {
|
||||
return int(n), true, &net.OpError{Op: "recvmmsg", Err: errno}
|
||||
}
|
||||
return int(n), true, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) listenOutSingle(r EncReader) error {
|
||||
var err error
|
||||
var n int
|
||||
var from netip.AddrPort
|
||||
buffer := make([]byte, MTU)
|
||||
|
||||
for {
|
||||
n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer)
|
||||
if err != nil {
|
||||
return err
|
||||
if u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
||||
r(from, buffer[:n])
|
||||
return 0, &net.OpError{Op: "recvmmsg", Err: errno}
|
||||
}
|
||||
n := int(r)
|
||||
if (n == 0 || msgs[0].Len == 0) && u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) listenOutBatch(r EncReader) error {
|
||||
// recvmsg does one blocking recvmsg into msgs[0]
|
||||
func (u *StdConn) recvmsg(msgs []rawMessage) (int, error) {
|
||||
r, _, errno := unix.Syscall6(
|
||||
unix.SYS_RECVMSG,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&msgs[0].Hdr)),
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
)
|
||||
if errno != 0 {
|
||||
if u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
return 0, &net.OpError{Op: "recvmsg", Err: errno}
|
||||
}
|
||||
if r == 0 && u.closed.Load() {
|
||||
return 0, net.ErrClosed
|
||||
}
|
||||
msgs[0].Len = uint32(r)
|
||||
return 1, nil
|
||||
}
|
||||
|
||||
func (u *StdConn) ListenOut(r EncReader) error {
|
||||
var ip netip.Addr
|
||||
var n int
|
||||
var operr error
|
||||
|
||||
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
||||
|
||||
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
||||
//defining it outside the loop so it gets re-used
|
||||
reader := func(fd uintptr) (done bool) {
|
||||
n, done, operr = recvmmsg(fd, msgs)
|
||||
return done
|
||||
read := u.recvmmsg
|
||||
if u.batch == 1 {
|
||||
read = u.recvmsg
|
||||
}
|
||||
|
||||
for {
|
||||
err := u.rawConn.Read(reader)
|
||||
n, err := read(msgs)
|
||||
if err != nil {
|
||||
if errors.Is(err, unix.EINTR) {
|
||||
continue // interrupted by a signal, retry the read
|
||||
}
|
||||
// net.ErrClosed after Close() is teardown, absorbed by the caller's
|
||||
// closed flag like the other platforms; anything else is a real error.
|
||||
return err
|
||||
}
|
||||
if operr != nil {
|
||||
return operr
|
||||
}
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||
@@ -222,26 +193,68 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) ListenOut(r EncReader) error {
|
||||
if u.batch == 1 {
|
||||
return u.listenOutSingle(r)
|
||||
} else {
|
||||
return u.listenOutBatch(r)
|
||||
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
||||
if u.isV4 {
|
||||
return u.writeTo4(b, ip)
|
||||
}
|
||||
return u.writeTo6(b, ip)
|
||||
}
|
||||
|
||||
func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error {
|
||||
var rsa unix.RawSockaddrInet6
|
||||
rsa.Family = unix.AF_INET6
|
||||
rsa.Addr = ip.Addr().As16()
|
||||
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||
|
||||
for {
|
||||
_, _, err := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&b[0])),
|
||||
uintptr(len(b)),
|
||||
uintptr(0),
|
||||
uintptr(unsafe.Pointer(&rsa)),
|
||||
uintptr(unix.SizeofSockaddrInet6),
|
||||
)
|
||||
if err != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: err}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
||||
_, err := u.udpConn.WriteToUDPAddrPort(b, ip)
|
||||
return err
|
||||
func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error {
|
||||
if !ip.Addr().Is4() {
|
||||
return ErrInvalidIPv6RemoteForSocket
|
||||
}
|
||||
|
||||
var rsa unix.RawSockaddrInet4
|
||||
rsa.Family = unix.AF_INET
|
||||
rsa.Addr = ip.Addr().As4()
|
||||
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||
|
||||
for {
|
||||
_, _, err := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&b[0])),
|
||||
uintptr(len(b)),
|
||||
uintptr(0),
|
||||
uintptr(unsafe.Pointer(&rsa)),
|
||||
uintptr(unix.SizeofSockaddrInet4),
|
||||
)
|
||||
if err != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: err}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (u *StdConn) ReloadConfig(c *config.C) {
|
||||
b := c.GetInt("listen.read_buffer", 0)
|
||||
if b > 0 {
|
||||
err := u.SetRecvBuffer(b)
|
||||
if err == nil {
|
||||
s, err := u.GetRecvBuffer()
|
||||
if err == nil {
|
||||
if err := u.SetRecvBuffer(b); err == nil {
|
||||
if s, err := u.GetRecvBuffer(); err == nil {
|
||||
u.l.Info("listen.read_buffer was set", "size", s)
|
||||
} else {
|
||||
u.l.Warn("Failed to get listen.read_buffer", "error", err)
|
||||
@@ -253,10 +266,8 @@ func (u *StdConn) ReloadConfig(c *config.C) {
|
||||
|
||||
b = c.GetInt("listen.write_buffer", 0)
|
||||
if b > 0 {
|
||||
err := u.SetSendBuffer(b)
|
||||
if err == nil {
|
||||
s, err := u.GetSendBuffer()
|
||||
if err == nil {
|
||||
if err := u.SetSendBuffer(b); err == nil {
|
||||
if s, err := u.GetSendBuffer(); err == nil {
|
||||
u.l.Info("listen.write_buffer was set", "size", s)
|
||||
} else {
|
||||
u.l.Warn("Failed to get listen.write_buffer", "error", err)
|
||||
@@ -269,10 +280,8 @@ func (u *StdConn) ReloadConfig(c *config.C) {
|
||||
b = c.GetInt("listen.so_mark", 0)
|
||||
s, err := u.GetSoMark()
|
||||
if b > 0 || (err == nil && s != 0) {
|
||||
err := u.SetSoMark(b)
|
||||
if err == nil {
|
||||
s, err := u.GetSoMark()
|
||||
if err == nil {
|
||||
if err := u.SetSoMark(b); err == nil {
|
||||
if s, err := u.GetSoMark(); err == nil {
|
||||
u.l.Info("listen.so_mark was set", "mark", s)
|
||||
} else {
|
||||
u.l.Warn("Failed to get listen.so_mark", "error", err)
|
||||
@@ -285,28 +294,20 @@ func (u *StdConn) ReloadConfig(c *config.C) {
|
||||
|
||||
func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
|
||||
var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
|
||||
|
||||
if u.rawConn == nil {
|
||||
return fmt.Errorf("no UDP connection")
|
||||
}
|
||||
var opErr error
|
||||
err := u.rawConn.Control(func(fd uintptr) {
|
||||
_, _, syserr := unix.Syscall6(unix.SYS_GETSOCKOPT, fd, uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
|
||||
if syserr != 0 {
|
||||
opErr = syserr
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
_, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
|
||||
if err != 0 {
|
||||
return err
|
||||
}
|
||||
return opErr
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *StdConn) Close() error {
|
||||
if u.udpConn != nil {
|
||||
return u.udpConn.Close()
|
||||
}
|
||||
return nil
|
||||
u.closed.Store(true)
|
||||
// Wake the reader parked in recvmmsg/recvmsg. shutdown(2) on an unconnected socket
|
||||
// returns ENOTCONN but still wakes it, so ignore the error.
|
||||
// The reader then sees closed and stops touching the fd, making the Close below safe.
|
||||
_ = unix.Shutdown(u.sysFd, unix.SHUT_RDWR)
|
||||
return unix.Close(u.sysFd)
|
||||
}
|
||||
|
||||
func NewUDPStatsEmitter(udpConns []Conn) func() {
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
//go:build linux && !android && !e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func testLogger() *slog.Logger {
|
||||
return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
|
||||
}
|
||||
|
||||
// TestShutdownWakesAfterRx_Mechanism exercises the kernel quirk our teardown
|
||||
// relies on: once a socket has received a packet, shutdown(2) wakes a blocked
|
||||
// recvmmsg with n>=1/Len==0 (not n==0). recvmmsg must turn that into net.ErrClosed
|
||||
// once Close set closed, so a parked reader exits instead of spinning.
|
||||
func TestShutdownWakesAfterRx_Mechanism(t *testing.T) {
|
||||
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, 64)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := c.(*StdConn)
|
||||
addr, err := sc.LocalAddr()
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
msgs, _, _ := sc.PrepareRawMessages(sc.batch)
|
||||
|
||||
// Receive a real packet so the socket has carried data.
|
||||
send, err := net.Dial("udp", addr.String())
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
if _, err := send.Write([]byte("hello")); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
n, err := sc.recvmmsg(msgs)
|
||||
t.Logf("drain of real packet: n=%d err=%v msgs[0].Len=%d", n, err, msgs[0].Len)
|
||||
_ = send.Close()
|
||||
|
||||
// Block a reader on the now-empty queue, then tear down as Close() does.
|
||||
// recvmmsg must return net.ErrClosed (not hang, not spin) even post-rx.
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := sc.recvmmsg(msgs)
|
||||
done <- err
|
||||
}()
|
||||
time.Sleep(150 * time.Millisecond) // let it park in recvmmsg
|
||||
|
||||
sc.closed.Store(true)
|
||||
if serr := unix.Shutdown(sc.sysFd, unix.SHUT_RDWR); serr != nil {
|
||||
t.Logf("shutdown returned %v (expected ENOTCONN on unconnected UDP)", serr)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, net.ErrClosed) {
|
||||
t.Errorf("recvmmsg after post-rx shutdown returned %v, want net.ErrClosed", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("HANG: recvmmsg did not return after shutdown following a received packet")
|
||||
}
|
||||
_ = unix.Close(sc.sysFd)
|
||||
}
|
||||
|
||||
// TestListenOutTeardown_TrafficPatterns reproduces the field report: a blocking
|
||||
// reader must tear down cleanly on Close() regardless of what the socket has
|
||||
// carried. The three cases the report called out:
|
||||
//
|
||||
// no traffic ever -> works (shutdown wakes recvmmsg with n==0)
|
||||
// ping once, then idle -> historically HUNG: once the socket has received a
|
||||
// packet, shutdown(2) wakes recvmmsg with n>=1/Len==0,
|
||||
// which an n==0-only teardown check misses
|
||||
// continuous traffic -> works (a real packet is always arriving)
|
||||
//
|
||||
// All three must return within the deadline; a hang dumps goroutines so the
|
||||
// stuck reader is visible.
|
||||
func TestListenOutTeardown_TrafficPatterns(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
traffic func(send net.Conn, stop <-chan struct{})
|
||||
}{
|
||||
{"no_traffic_ever", func(net.Conn, <-chan struct{}) {}},
|
||||
{"ping_once_then_idle", func(send net.Conn, _ <-chan struct{}) {
|
||||
_, _ = send.Write([]byte("hello"))
|
||||
}},
|
||||
{"continuous", func(send net.Conn, stop <-chan struct{}) {
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
_, _ = send.Write([]byte("hello"))
|
||||
time.Sleep(2 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
}},
|
||||
}
|
||||
|
||||
// batch 1 exercises the recvmsg path, batch 64 the recvmmsg path; both must
|
||||
// tear down cleanly.
|
||||
for _, batch := range []int{1, 64} {
|
||||
for _, tc := range cases {
|
||||
t.Run(fmt.Sprintf("batch%d/%s", batch, tc.name), func(t *testing.T) {
|
||||
runTeardownCase(t, batch, tc.name, tc.traffic)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net.Conn, stop <-chan struct{})) {
|
||||
c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, batch)
|
||||
if err != nil {
|
||||
t.Fatalf("NewListener: %v", err)
|
||||
}
|
||||
sc := c.(*StdConn)
|
||||
addr, err := sc.LocalAddr()
|
||||
if err != nil {
|
||||
t.Fatalf("LocalAddr: %v", err)
|
||||
}
|
||||
|
||||
var received atomic.Int64
|
||||
loopDone := make(chan error, 1)
|
||||
go func() {
|
||||
loopDone <- sc.ListenOut(func(netip.AddrPort, []byte) {
|
||||
received.Add(1)
|
||||
})
|
||||
}()
|
||||
|
||||
send, err := net.Dial("udp", addr.String())
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
defer send.Close()
|
||||
|
||||
stop := make(chan struct{})
|
||||
trafficDone := make(chan struct{})
|
||||
go func() {
|
||||
traffic(send, stop)
|
||||
close(trafficDone)
|
||||
}()
|
||||
|
||||
// Let the pattern run and, for the idle case, the reader park again on an
|
||||
// empty queue with the socket already having received a packet.
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
start := time.Now()
|
||||
if err := sc.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
close(stop)
|
||||
|
||||
select {
|
||||
case err := <-loopDone:
|
||||
// Clean teardown surfaces as net.ErrClosed (propagated like the other
|
||||
// platforms); the caller absorbs it via its closed flag.
|
||||
if err != nil && !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("%s: ListenOut returned unexpected error on teardown: %v", name, err)
|
||||
}
|
||||
t.Logf("%s: closed in %v (received %d packets)", name, time.Since(start), received.Load())
|
||||
case <-time.After(3 * time.Second):
|
||||
buf := make([]byte, 1<<20)
|
||||
n := runtime.Stack(buf, true)
|
||||
t.Fatalf("%s: HANG, ListenOut did not return within 3s of Close\n%s", name, buf[:n])
|
||||
}
|
||||
<-trafficDone
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
package udp
|
||||
|
||||
import mathrand "math/rand"
|
||||
|
||||
type SendPortGetter interface {
|
||||
// UDPSendPort returns the port to use
|
||||
UDPSendPort(maxPort int) uint16
|
||||
}
|
||||
|
||||
type randomSendPort struct{}
|
||||
|
||||
func (randomSendPort) UDPSendPort(maxPort int) uint16 {
|
||||
return uint16(mathrand.Intn(maxPort))
|
||||
}
|
||||
|
||||
var RandomSendPort = randomSendPort{}
|
||||
@@ -0,0 +1,191 @@
|
||||
//go:build !android && !e2e_testing
|
||||
// +build !android,!e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// RawOverhead is the number of bytes that need to be reserved at the start of
|
||||
// the raw bytes passed to (*RawConn).WriteTo. This is used by WriteTo to prefix
|
||||
// the IP and UDP headers.
|
||||
const RawOverhead = 28
|
||||
|
||||
type RawConn struct {
|
||||
sysFd int
|
||||
basePort uint16
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func NewRawConn(l *slog.Logger, ip string, port int, basePort uint16) (*RawConn, error) {
|
||||
syscall.ForkLock.RLock()
|
||||
// With IPPROTO_UDP, the linux kernel tries to deliver every UDP packet
|
||||
// received in the system to our socket. This constantly overflows our
|
||||
// buffer and marks our socket as having dropped packets. This makes the
|
||||
// stats on the socket useless.
|
||||
//
|
||||
// In contrast, IPPROTO_RAW is not delivered any packets and thus our read
|
||||
// buffer will not fill up and mark as having dropped packets. The only
|
||||
// difference is that we have to assemble the IP header as well, but this
|
||||
// is fairly easy since Linux does the checksum for us.
|
||||
//
|
||||
// TODO: How to get this working with Inet6 correctly? I was having issues
|
||||
// with the source address when testing before, probably need to `bind(2)`?
|
||||
fd, err := unix.Socket(unix.AF_INET, unix.SOCK_RAW, unix.IPPROTO_RAW)
|
||||
if err == nil {
|
||||
unix.CloseOnExec(fd)
|
||||
}
|
||||
syscall.ForkLock.RUnlock()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// We only want to send, not recv. This will hopefully help the kernel avoid
|
||||
// wasting time on us
|
||||
if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_RCVBUF, 0); err != nil {
|
||||
return nil, fmt.Errorf("unable to set SO_RCVBUF: %s", err)
|
||||
}
|
||||
|
||||
var lip [16]byte
|
||||
copy(lip[:], net.ParseIP(ip))
|
||||
|
||||
// TODO do we need to `bind(2)` so that we send from the correct address/interface?
|
||||
if err = unix.Bind(fd, &unix.SockaddrInet6{Addr: lip, Port: port}); err != nil {
|
||||
return nil, fmt.Errorf("unable to bind to socket: %s", err)
|
||||
}
|
||||
|
||||
return &RawConn{
|
||||
sysFd: fd,
|
||||
basePort: basePort,
|
||||
l: l,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// WriteTo must be called with raw leaving the first `udp.RawOverhead` bytes empty,
|
||||
// for the IP/UDP headers.
|
||||
func (u *RawConn) WriteTo(raw []byte, fromPort uint16, ip netip.AddrPort) error {
|
||||
var rsa unix.RawSockaddrInet4
|
||||
rsa.Family = unix.AF_INET
|
||||
rsa.Addr = ip.Addr().As4()
|
||||
|
||||
totalLen := len(raw)
|
||||
udpLen := totalLen - ipv4.HeaderLen
|
||||
|
||||
// IP header
|
||||
raw[0] = byte(ipv4.Version<<4 | (ipv4.HeaderLen >> 2 & 0x0f))
|
||||
raw[1] = 0 // tos
|
||||
binary.BigEndian.PutUint16(raw[2:4], uint16(totalLen))
|
||||
binary.BigEndian.PutUint16(raw[4:6], 0) // id (linux does it for us)
|
||||
binary.BigEndian.PutUint16(raw[6:8], 0) // frag options
|
||||
raw[8] = byte(64) // ttl
|
||||
raw[9] = byte(17) // protocol
|
||||
binary.BigEndian.PutUint16(raw[10:12], 0) // checksum (linux does it for us)
|
||||
binary.BigEndian.PutUint32(raw[12:16], 0) // src (linux does it for us)
|
||||
copy(raw[16:20], rsa.Addr[:]) // dst
|
||||
|
||||
// UDP header
|
||||
fromPort = u.basePort + fromPort
|
||||
binary.BigEndian.PutUint16(raw[20:22], uint16(fromPort)) // src port
|
||||
binary.BigEndian.PutUint16(raw[22:24], uint16(ip.Port())) // dst port
|
||||
binary.BigEndian.PutUint16(raw[24:26], uint16(udpLen)) // UDP length
|
||||
binary.BigEndian.PutUint16(raw[26:28], 0) // checksum (optional)
|
||||
|
||||
for {
|
||||
_, _, err := unix.Syscall6(
|
||||
unix.SYS_SENDTO,
|
||||
uintptr(u.sysFd),
|
||||
uintptr(unsafe.Pointer(&raw[0])),
|
||||
uintptr(len(raw)),
|
||||
uintptr(0),
|
||||
uintptr(unsafe.Pointer(&rsa)),
|
||||
uintptr(unix.SizeofSockaddrInet4),
|
||||
)
|
||||
|
||||
if err != 0 {
|
||||
return &net.OpError{Op: "sendto", Err: err}
|
||||
}
|
||||
|
||||
//TODO: handle incomplete writes
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (u *RawConn) ReloadConfig(c *config.C) {
|
||||
b := c.GetInt("listen.write_buffer", 0)
|
||||
if b <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if err := u.SetSendBuffer(b); err != nil {
|
||||
u.l.Error("Failed to set listen.write_buffer", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
s, err := u.GetSendBuffer()
|
||||
if err != nil {
|
||||
u.l.Warn("Failed to get listen.write_buffer", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
u.l.Info("listen.write_buffer was set", "size", s)
|
||||
}
|
||||
|
||||
func (u *RawConn) SetSendBuffer(n int) error {
|
||||
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
|
||||
}
|
||||
|
||||
func (u *RawConn) GetSendBuffer() (int, error) {
|
||||
return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUF)
|
||||
}
|
||||
|
||||
func (u *RawConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
|
||||
var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
|
||||
_, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
|
||||
if err != 0 {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func NewRawStatsEmitter(rawConn *RawConn) func() {
|
||||
// Check if our kernel supports SO_MEMINFO before registering the gauges
|
||||
var gauges [unix.SK_MEMINFO_VARS]metrics.Gauge
|
||||
var meminfo [unix.SK_MEMINFO_VARS]uint32
|
||||
if err := rawConn.getMemInfo(&meminfo); err == nil {
|
||||
gauges = [unix.SK_MEMINFO_VARS]metrics.Gauge{
|
||||
metrics.GetOrRegisterGauge("raw.rmem_alloc", nil),
|
||||
metrics.GetOrRegisterGauge("raw.rcvbuf", nil),
|
||||
metrics.GetOrRegisterGauge("raw.wmem_alloc", nil),
|
||||
metrics.GetOrRegisterGauge("raw.sndbuf", nil),
|
||||
metrics.GetOrRegisterGauge("raw.fwd_alloc", nil),
|
||||
metrics.GetOrRegisterGauge("raw.wmem_queued", nil),
|
||||
metrics.GetOrRegisterGauge("raw.optmem", nil),
|
||||
metrics.GetOrRegisterGauge("raw.backlog", nil),
|
||||
metrics.GetOrRegisterGauge("raw.drops", nil),
|
||||
}
|
||||
} else {
|
||||
// return no-op because we don't support SO_MEMINFO
|
||||
return func() {}
|
||||
}
|
||||
|
||||
return func() {
|
||||
if err := rawConn.getMemInfo(&meminfo); err == nil {
|
||||
for j := 0; j < unix.SK_MEMINFO_VARS; j++ {
|
||||
gauges[j].Update(int64(meminfo[j]))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
//go:build !linux || android || e2e_testing
|
||||
// +build !linux android e2e_testing
|
||||
|
||||
package udp
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
)
|
||||
|
||||
const RawOverhead = 0
|
||||
|
||||
type RawConn struct{}
|
||||
|
||||
func NewRawConn(l *slog.Logger, ip string, port int, basePort uint16) (*RawConn, error) {
|
||||
return nil, fmt.Errorf("multiport tx is not supported on %s", runtime.GOOS)
|
||||
}
|
||||
|
||||
func (u *RawConn) WriteTo(raw []byte, fromPort uint16, addr netip.AddrPort) error {
|
||||
return fmt.Errorf("multiport tx is not supported on %s", runtime.GOOS)
|
||||
}
|
||||
|
||||
func (u *RawConn) ReloadConfig(c *config.C) {}
|
||||
|
||||
func NewRawStatsEmitter(rawConn *RawConn) func() { return func() {} }
|
||||
+20
-6
@@ -10,6 +10,7 @@ import (
|
||||
"net/netip"
|
||||
"os"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
@@ -64,7 +65,9 @@ func acquirePacket() *Packet {
|
||||
}
|
||||
|
||||
type TesterConn struct {
|
||||
Addr netip.AddrPort
|
||||
// addr is read by nebula's own goroutines on every send and by the router's flow renderer, and a test can
|
||||
// move it mid-run to simulate roaming, so it is atomic rather than a plain field.
|
||||
addr atomic.Pointer[netip.AddrPort]
|
||||
|
||||
RxPackets chan *Packet // Packets to receive into nebula
|
||||
TxPackets chan *Packet // Packets transmitted outside by nebula
|
||||
@@ -82,13 +85,24 @@ type TesterConn struct {
|
||||
}
|
||||
|
||||
func NewListener(l *slog.Logger, ip netip.Addr, port int, _ bool, _ int) (Conn, error) {
|
||||
return &TesterConn{
|
||||
Addr: netip.AddrPortFrom(ip, uint16(port)),
|
||||
c := &TesterConn{
|
||||
RxPackets: make(chan *Packet, 10),
|
||||
TxPackets: make(chan *Packet, 10),
|
||||
done: make(chan struct{}),
|
||||
l: l,
|
||||
}, nil
|
||||
}
|
||||
c.SetAddr(netip.AddrPortFrom(ip, uint16(port)))
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// GetAddr returns the underlay address this conn currently sends from.
|
||||
func (u *TesterConn) GetAddr() netip.AddrPort {
|
||||
return *u.addr.Load()
|
||||
}
|
||||
|
||||
// SetAddr moves this conn to a new underlay address, standing in for a host waking up on a different network.
|
||||
func (u *TesterConn) SetAddr(addr netip.AddrPort) {
|
||||
u.addr.Store(&addr)
|
||||
}
|
||||
|
||||
// Send will place a UdpPacket onto the receive queue for nebula to consume
|
||||
@@ -147,7 +161,7 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
||||
p.Data = p.Data[:len(b)]
|
||||
}
|
||||
copy(p.Data, b)
|
||||
p.From = u.Addr
|
||||
p.From = u.GetAddr()
|
||||
p.To = addr
|
||||
select {
|
||||
case <-u.done:
|
||||
@@ -178,7 +192,7 @@ func NewUDPStatsEmitter(_ []Conn) func() {
|
||||
}
|
||||
|
||||
func (u *TesterConn) LocalAddr() (netip.AddrPort, error) {
|
||||
return u.Addr, nil
|
||||
return u.GetAddr(), nil
|
||||
}
|
||||
|
||||
func (u *TesterConn) SupportsMultipleReaders() bool {
|
||||
|
||||
Reference in New Issue
Block a user