mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 13:57:03 +02:00
Compare commits
101 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| a5cb635db4 | |||
| bf2f26b66f | |||
| 7c00c74550 | |||
| 738b4fa044 | |||
| 0e1f5342ac | |||
| 00b498f9f9 | |||
| 69e5e8eeb4 | |||
| fc950cf765 | |||
| 0d91e3e4d8 | |||
| e8322505ab | |||
| 009d42b4f5 | |||
| f8775bb6ca | |||
| 8d8cc4f73d | |||
| e16a31ff6a | |||
| 15f0f0d5d0 | |||
| 7902ce674e | |||
| c2fbe215e6 | |||
| 94ac6db4ca | |||
| a60350e34e | |||
| 58f3b6fda7 | |||
| a99699e370 | |||
| 3615a79b8b | |||
| 6be3d2264e | |||
| 048fa3fa70 | |||
| 147c202c27 | |||
| e290a6892f | |||
| 6c3972f464 | |||
| 861d3aabd7 | |||
| 86733864fe | |||
| ab736e4c6b | |||
| 5ecdd4eaa9 | |||
| 1b84bd0050 | |||
| 384610f81a | |||
| 3264ebf953 | |||
| 3c1dd52953 | |||
| 4e8733ca22 | |||
| 6675f5a361 | |||
| 14c9288127 | |||
| 90c7630270 | |||
| 81afed439e | |||
| d42407cf9a | |||
| ac5a18e56a | |||
| 2efdcb0366 | |||
| 92bd785f2b | |||
| 7d087a6d9f | |||
| 67da623e3b | |||
| 060b46f4b3 | |||
| 69c6132a76 | |||
| a383937bde | |||
| 3b30526379 | |||
| cb9547b6a1 | |||
| 69a6db705a | |||
| 9709893009 | |||
| 32db819315 | |||
| f437c7d372 | |||
| d725e53072 | |||
| 06fb503fc3 | |||
| 7cd3875934 | |||
| 90ea6346e9 | |||
| 37b752bb23 | |||
| c7e035479c | |||
| adb01f66a3 | |||
| cf5d73d625 | |||
| 56a09b7cbb | |||
| 1d17c785a8 | |||
| 7c8a70b0cc | |||
| c73245ca5c | |||
| f1a8054a9d | |||
| dd081ffeb6 | |||
| cfcdcb0546 | |||
| c3e2a2b3a7 | |||
| 0897f49576 | |||
| 8671a4ebbd | |||
| b5ad62aea1 | |||
| a1b8954a23 | |||
| 441a67fbf9 | |||
| 7f6976ba5d | |||
| 5dd566e220 | |||
| b79fdc272a | |||
| d4d747f219 | |||
| 4ff9e97dce | |||
| bdbd483063 | |||
| 1403473e4b | |||
| ae58085032 | |||
| 65450f9d21 | |||
| f9532f4244 | |||
| fcfbec4710 | |||
| e6eeb3709b | |||
| ef8b700474 | |||
| 243cf4a7c5 | |||
| 08ee2ab35f | |||
| 2d5d86f24d | |||
| 6fa1ecdc29 | |||
| b418a081a8 | |||
| fd3fa57e79 | |||
| 0eb92dcab4 | |||
| f6b206d96c | |||
| 31cc3a4169 | |||
| 6da314aa6b | |||
| 3da3d41fb5 | |||
| 4485c47641 |
@@ -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
|
||||
|
||||
+12
-32
@@ -20,44 +20,24 @@ 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
|
||||
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
||||
- name: Smoke Docker
|
||||
run: make smoke-docker
|
||||
|
||||
- name: setup docker image
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./build.sh
|
||||
- name: Smoke Docker IPv6 overlay
|
||||
run: make smoke-docker-ipv6
|
||||
|
||||
- name: run smoke
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./smoke.sh
|
||||
- name: Smoke Relay Docker
|
||||
run: make smoke-relay-docker
|
||||
|
||||
- name: setup docker image ipv6
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: SMOKE_OVERLAY_IPV6=1 ./build.sh
|
||||
- name: Smoke Docker boringcrypto
|
||||
run: make boringcrypto smoke-docker
|
||||
|
||||
- name: run smoke ipv6
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: SMOKE_OVERLAY_IPV6=1 ./smoke.sh
|
||||
|
||||
- name: setup relay docker image
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./build-relay.sh
|
||||
|
||||
- name: run smoke relay
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./smoke-relay.sh
|
||||
|
||||
- name: setup docker image for P256
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: NAME="smoke-p256" CURVE=P256 ./build.sh
|
||||
|
||||
- name: run smoke-p256
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: NAME="smoke-p256" ./smoke.sh
|
||||
- name: Smoke Docker fips140
|
||||
run: make fips140-all GOALS=smoke-docker
|
||||
|
||||
timeout-minutes: 10
|
||||
|
||||
+15
-10
@@ -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 }}
|
||||
@@ -58,9 +58,14 @@ jobs:
|
||||
e2e-cmd: make e2evv
|
||||
- name: linux-boringcrypto
|
||||
os: ubuntu-latest
|
||||
build-cmd: make bin-boringcrypto
|
||||
test-cmd: make test-boringcrypto
|
||||
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||
build-cmd: make boringcrypto
|
||||
test-cmd: make boringcrypto test
|
||||
e2e-cmd: make boringcrypto e2evv
|
||||
- name: linux-fips140
|
||||
os: ubuntu-latest
|
||||
build-cmd: make fips140-all
|
||||
test-cmd: make fips140-all GOALS=test
|
||||
e2e-cmd: make fips140-all GOALS=e2evv
|
||||
- name: linux-pkcs11
|
||||
os: ubuntu-latest
|
||||
build-cmd: make bin-pkcs11
|
||||
@@ -80,9 +85,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 +130,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 }}
|
||||
|
||||
@@ -72,6 +72,17 @@ ALL_CROSS_LINUX_ARM = linux-arm-5 linux-arm-6 linux-arm-7 linux-arm64
|
||||
ALL_CROSS_LINUX_MIPS = linux-mips linux-mipsle linux-mips64 linux-mips64le linux-mips-softfloat
|
||||
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
|
||||
|
||||
# Based on section 2.2 of the Go Cryptographic Module CVMP Security Policy #5247
|
||||
ALL_FIPS140 = linux-amd64-fips140 \
|
||||
linux-arm64-fips140 \
|
||||
windows-amd64-fips140 \
|
||||
windows-arm64-fips140 \
|
||||
darwin-arm64-fips140 \
|
||||
freebsd-amd64-fips140 \
|
||||
linux-arm-7-fips140 \
|
||||
linux-mips64-fips140 \
|
||||
linux-ppc64le-fips140
|
||||
|
||||
e2e:
|
||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||
|
||||
@@ -137,6 +148,8 @@ release-netbsd: $(ALL_NETBSD:%=build/nebula-%.tar.gz)
|
||||
|
||||
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
||||
|
||||
release-fips140: $(ALL_FIPS140:%=build/nebula-%.tar.gz)
|
||||
|
||||
BUILD_ARGS += -trimpath
|
||||
|
||||
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
||||
@@ -157,17 +170,20 @@ bin-freebsd-arm64: build/freebsd-arm64/nebula build/freebsd-arm64/nebula-cert
|
||||
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
||||
mv $? .
|
||||
|
||||
bin-fips140: build/linux-$(shell go env GOARCH)-fips140/nebula build/linux-$(shell go env GOARCH)-fips140/nebula-cert
|
||||
mv $? .
|
||||
|
||||
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
||||
bin-pkcs11: CGO_ENABLED = 1
|
||||
bin-pkcs11: bin
|
||||
|
||||
bin:
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||
$(GOENV) go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||
|
||||
install:
|
||||
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||
|
||||
build/linux-arm-%: GOENV += GOARM=$(word 3, $(subst -, ,$*))
|
||||
build/linux-mips-%: GOENV += GOMIPS=$(word 3, $(subst -, ,$*))
|
||||
@@ -178,8 +194,11 @@ build/linux-mips-softfloat/%: LDFLAGS += -s -w
|
||||
# boringcrypto
|
||||
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||
build/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
|
||||
# fips140
|
||||
FIPSVERSION = v1.0.0
|
||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): GOENV += GOFIPS140=$(FIPSVERSION)
|
||||
$(foreach _rule, $(ALL_FIPS140), build/$(_rule)/%): BUILD_ARGS += -tags fips140enforce
|
||||
|
||||
build/%/nebula: .FORCE
|
||||
GOOS=$(firstword $(subst -, , $*)) \
|
||||
@@ -210,10 +229,7 @@ vet:
|
||||
go vet $(VET_FLAGS) -v ./...
|
||||
|
||||
test:
|
||||
go test -v ./...
|
||||
|
||||
test-boringcrypto:
|
||||
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
||||
$(TEST_ENV) go test $(TEST_FLAGS) -v ./...
|
||||
|
||||
test-pkcs11:
|
||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
||||
@@ -256,29 +272,72 @@ ifeq ($(words $(MAKECMDGOALS)),1)
|
||||
@$(MAKE) service ${.DEFAULT_GOAL} --no-print-directory
|
||||
endif
|
||||
|
||||
# Useful to chain together, like:
|
||||
# - make fips140 e2evv
|
||||
# - make fips140 smoke-docker
|
||||
# Use `release-fips140` to build release binaries
|
||||
fips140:
|
||||
@echo > $(NULL_FILE)
|
||||
ifeq ($(strip $(GOFIPS140)),)
|
||||
$(eval GOFIPS140 = $(FIPSVERSION))
|
||||
endif
|
||||
$(eval GOENV += GOFIPS140=$(GOFIPS140))
|
||||
$(eval BUILD_ARGS += -tags fips140enforce)
|
||||
$(eval TEST_ENV += $(GOENV))
|
||||
$(eval CURVE = P256)
|
||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
||||
@$(MAKE) fips140 GOFIPS140=$(GOFIPS140) ${.DEFAULT_GOAL} --no-print-directory
|
||||
endif
|
||||
|
||||
# To test the future pending module, use like `make fips140-latest test`
|
||||
ALL_GOFIPS140 = v1.0.0 v1.26.0 latest
|
||||
define FIPS140_rule
|
||||
fips140-$(1): GOFIPS140 = $(1)
|
||||
fips140-$(1): fips140
|
||||
endef
|
||||
$(foreach _rule, $(ALL_GOFIPS140), $(eval $(call FIPS140_rule,$(_rule))))
|
||||
|
||||
# Iterate and run the goals for all fips versions, like `make fips140-all GOALS=test`
|
||||
fips140-all:
|
||||
@$(foreach _v,$(ALL_GOFIPS140),$(MAKE) fips140-$(_v) $(GOALS) &&) true
|
||||
|
||||
# Useful to chain together, like:
|
||||
# - make boringcrypto e2evv
|
||||
# - make boringcrypto smoke-docker
|
||||
# Use `release-boringcrypto` or `bin-boringcrypto` to build release binaries
|
||||
boringcrypto:
|
||||
@echo > $(NULL_FILE)
|
||||
$(eval GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1)
|
||||
$(eval TEST_ENV += $(GOENV))
|
||||
$(eval CURVE = P256)
|
||||
ifeq ($(words $(MAKECMDGOALS)),1)
|
||||
@$(MAKE) boringcrypto ${.DEFAULT_GOAL} --no-print-directory
|
||||
endif
|
||||
|
||||
bin-docker: bin build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
||||
|
||||
smoke-docker: BUILD_ARGS += -race
|
||||
smoke-docker: GOENV += CGO_ENABLED=1
|
||||
smoke-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && ./build.sh
|
||||
cd .github/workflows/smoke/ && ./smoke.sh
|
||||
cd .github/workflows/smoke/ && NAME="smoke-p256" CURVE="P256" ./build.sh
|
||||
cd .github/workflows/smoke/ && NAME="smoke-p256" ./smoke.sh
|
||||
# This is so we can limit `fips140` smoke test to just P256 curve.
|
||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./build.sh; fi
|
||||
if [ "$(CURVE)" != "P256" ]; then cd .github/workflows/smoke/ && $(GOENV) ./smoke.sh; fi
|
||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" CURVE="P256" ./build.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) NAME="smoke-p256" ./smoke.sh
|
||||
|
||||
smoke-relay-docker: BUILD_ARGS += -race
|
||||
smoke-relay-docker: GOENV += CGO_ENABLED=1
|
||||
smoke-relay-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./smoke-relay.sh
|
||||
|
||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||
smoke-docker-ipv6: smoke-docker
|
||||
|
||||
smoke-docker-race: BUILD_ARGS = -race
|
||||
smoke-docker-race: CGO_ENABLED = 1
|
||||
smoke-docker-race: smoke-docker
|
||||
|
||||
smoke-vagrant/%: bin-docker build/%/nebula
|
||||
cd .github/workflows/smoke/ && ./build.sh $*
|
||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||
|
||||
.FORCE:
|
||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 test test-pkcs11 test-cov-html vet smoke-vagrant/%
|
||||
.DEFAULT_GOAL := bin
|
||||
|
||||
@@ -145,17 +145,27 @@ To build nebula for a specific platform (ex, Windows):
|
||||
|
||||
See the [Makefile](Makefile) for more details on build targets
|
||||
|
||||
## Curve P256 and BoringCrypto
|
||||
## Curve P256 and FIPS 140-3 mode
|
||||
|
||||
The default curve used for cryptographic handshakes and signatures is Curve25519. This is the recommended setting for most users. If your deployment has certain compliance requirements, you have the option of creating your CA using `nebula-cert ca -curve P256` to use NIST Curve P256. The CA will then sign certificates using ECDSA P256, and any hosts using these certificates will use P256 for ECDH handshakes.
|
||||
|
||||
In addition, Nebula can be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets:
|
||||
Nebula can be built to support the [FIPS 140-3](https://go.dev/doc/security/fips140) mode of Go by running either of the following make targets. (This sets GOFIPS140=v1.0.0, which must be done at compile time so that the correct AES-GCM can be used for FIPS 140-3 enforcement mode).
|
||||
|
||||
```sh
|
||||
make fips140
|
||||
make fips140 test
|
||||
make release-fips140
|
||||
```
|
||||
|
||||
Nebula can also be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets.
|
||||
|
||||
```sh
|
||||
make bin-boringcrypto
|
||||
make release-boringcrypto
|
||||
```
|
||||
|
||||
NOTE: boringcrypto support is deprecated and will be removed in the next release. Users should migrate to the native FIPS 140-3 mode described above.
|
||||
|
||||
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
||||
|
||||
## Credits
|
||||
|
||||
@@ -3,6 +3,7 @@ package main
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/fips140"
|
||||
"crypto/rand"
|
||||
"flag"
|
||||
"fmt"
|
||||
@@ -43,6 +44,13 @@ type caFlags struct {
|
||||
subnets *string
|
||||
}
|
||||
|
||||
func defaultCurve() string {
|
||||
if fips140.Enforced() {
|
||||
return "P256"
|
||||
}
|
||||
return "25519"
|
||||
}
|
||||
|
||||
func newCaFlags() *caFlags {
|
||||
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
||||
cf.set.Usage = func() {}
|
||||
@@ -59,7 +67,7 @@ func newCaFlags() *caFlags {
|
||||
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
||||
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
||||
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
||||
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", defaultCurve(), "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
|
||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -24,7 +24,7 @@ func newKeygenFlags() *keygenFlags {
|
||||
cf.set.Usage = func() {}
|
||||
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
||||
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
||||
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", defaultCurve(), "ECDH Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
return &cf
|
||||
}
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
+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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+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")
|
||||
}
|
||||
|
||||
+8
-4
@@ -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.
|
||||
@@ -242,10 +250,6 @@ tun:
|
||||
# When tun is disabled, a lighthouse can be started without a local tun interface (and therefore without root)
|
||||
disabled: false
|
||||
# Name of the device. If not set, a default will be chosen by the OS.
|
||||
# For Linux: a single `%d` anywhere in the name is treated as a template and replaced with the
|
||||
# lowest number that yields an unused device name (e.g. `nebula%d` becomes `nebula0`, then `nebula1`, and so on, `neb%dprod` becomes `neb0prod`).
|
||||
# Only on Linux: `nebula%d` is the default if tun.dev is unset.
|
||||
# The name, both before and after %d substitution, must be shorter than the kernel limit of 16 characters.
|
||||
# For macOS: if set, must be in the form `utun[0-9]+`.
|
||||
# For NetBSD: Required to be set, must be in the form `tun[0-9]+`
|
||||
dev: nebula1
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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=
|
||||
|
||||
+13
-8
@@ -295,7 +295,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
||||
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,
|
||||
@@ -430,14 +436,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 +535,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 {
|
||||
@@ -1080,7 +1085,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])...)
|
||||
}
|
||||
|
||||
+152
-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
|
||||
}
|
||||
@@ -266,10 +275,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 +287,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 +338,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 +387,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 +447,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 +530,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 +539,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 +584,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 +615,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 +672,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) {
|
||||
|
||||
@@ -87,7 +87,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 +103,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
|
||||
}
|
||||
|
||||
@@ -408,7 +408,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
"udpAddr", hr,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
|
||||
+33
-14
@@ -2,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -215,6 +216,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
|
||||
}
|
||||
|
||||
@@ -235,6 +239,9 @@ func (f *Interface) activate() error {
|
||||
"build", f.version,
|
||||
"udpAddr", addr,
|
||||
"boringcrypto", boringEnabled(),
|
||||
"fips140Version", fips140.Version(),
|
||||
"fips140Enabled", fips140.Enabled(),
|
||||
"fips140Enforced", fips140.Enforced(),
|
||||
)
|
||||
|
||||
if f.routines > 1 {
|
||||
@@ -258,17 +265,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 +289,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 +329,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 +351,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)
|
||||
}
|
||||
@@ -542,9 +553,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 +577,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
|
||||
@@ -257,6 +268,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 +280,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
|
||||
}
|
||||
|
||||
+4
-65
@@ -4,77 +4,16 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
|
||||
// unsafe needed for go:linkname
|
||||
_ "unsafe"
|
||||
"crypto/boring"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
var CipherAESGCM noise.CipherFunc = CipherAESGCMFIPS140
|
||||
|
||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||
// This is true for boringcrypto because the Seal function verifies that the
|
||||
// nonce is strictly increasing.
|
||||
const EncryptLockNeeded = true
|
||||
|
||||
// NewGCMTLS is no longer exposed in go1.19+, so we need to link it in
|
||||
// See: https://github.com/golang/go/issues/56326
|
||||
//
|
||||
// NewGCMTLS is the internal method used with boringcrypto that provides a
|
||||
// validated mode of AES-GCM which enforces the nonce is strictly
|
||||
// monotonically increasing. This is the TLS 1.2 specification for nonce
|
||||
// generation (which also matches the method used by the Noise Protocol)
|
||||
//
|
||||
// - https://github.com/golang/go/blob/go1.19/src/crypto/tls/cipher_suites.go#L520-L522
|
||||
// - https://github.com/golang/go/blob/go1.19/src/crypto/internal/boring/aes.go#L235-L237
|
||||
// - https://github.com/golang/go/blob/go1.19/src/crypto/internal/boring/aes.go#L250
|
||||
// - https://github.com/google/boringssl/blob/ae223d6138807a13006342edfeef32e813246b39/include/openssl/aead.h#L379-L381
|
||||
// - https://github.com/google/boringssl/blob/ae223d6138807a13006342edfeef32e813246b39/crypto/fipsmodule/cipher/e_aes.c#L1082-L1093
|
||||
//
|
||||
//go:linkname newGCMTLS crypto/internal/boring.NewGCMTLS
|
||||
func newGCMTLS(c cipher.Block) (cipher.AEAD, error)
|
||||
|
||||
type cipherFn struct {
|
||||
fn func([32]byte) noise.Cipher
|
||||
name string
|
||||
}
|
||||
|
||||
func (c cipherFn) Cipher(k [32]byte) noise.Cipher { return c.fn(k) }
|
||||
func (c cipherFn) CipherName() string { return c.name }
|
||||
|
||||
// CipherAESGCM is the AES256-GCM AEAD cipher (using NewGCMTLS when GoBoring is present)
|
||||
var CipherAESGCM noise.CipherFunc = cipherFn{cipherAESGCMBoring, "AESGCM"}
|
||||
|
||||
func cipherAESGCMBoring(k [32]byte) noise.Cipher {
|
||||
c, err := aes.NewCipher(k[:])
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
gcm, err := newGCMTLS(c)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return aeadCipher{
|
||||
gcm,
|
||||
func(n uint64) []byte {
|
||||
var nonce [12]byte
|
||||
binary.BigEndian.PutUint64(nonce[4:], n)
|
||||
return nonce[:]
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type aeadCipher struct {
|
||||
cipher.AEAD
|
||||
nonce func(uint64) []byte
|
||||
}
|
||||
|
||||
func (c aeadCipher) Encrypt(out []byte, n uint64, ad, plaintext []byte) []byte {
|
||||
return c.Seal(out, c.nonce(n), plaintext, ad)
|
||||
}
|
||||
|
||||
func (c aeadCipher) Decrypt(out []byte, n uint64, ad, ciphertext []byte) ([]byte, error) {
|
||||
return c.Open(out, c.nonce(n), ciphertext, ad)
|
||||
}
|
||||
var boringEnabled = boring.Enabled()
|
||||
|
||||
@@ -4,8 +4,6 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/boring"
|
||||
"encoding/hex"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -14,33 +12,3 @@ import (
|
||||
func TestEncryptLockNeeded(t *testing.T) {
|
||||
assert.True(t, EncryptLockNeeded)
|
||||
}
|
||||
|
||||
// Ensure NewGCMTLS validates the nonce is non-repeating
|
||||
func TestNewGCMTLS(t *testing.T) {
|
||||
assert.True(t, boring.Enabled())
|
||||
|
||||
// Test Case 16 from GCM Spec:
|
||||
// - (now dead link): http://csrc.nist.gov/groups/ST/toolkit/BCM/documents/proposedmodes/gcm/gcm-spec.pdf
|
||||
// - as listed in boringssl tests: https://github.com/google/boringssl/blob/fips-20220613/crypto/cipher_extra/test/cipher_tests.txt#L412-L418
|
||||
key, _ := hex.DecodeString("feffe9928665731c6d6a8f9467308308feffe9928665731c6d6a8f9467308308")
|
||||
iv, _ := hex.DecodeString("cafebabefacedbaddecaf888")
|
||||
plaintext, _ := hex.DecodeString("d9313225f88406e5a55909c5aff5269a86a7a9531534f7da2e4c303d8a318a721c3c0c95956809532fcf0e2449a6b525b16aedf5aa0de657ba637b39")
|
||||
aad, _ := hex.DecodeString("feedfacedeadbeeffeedfacedeadbeefabaddad2")
|
||||
expected, _ := hex.DecodeString("522dc1f099567d07f47f37a32a84427d643a8cdcbfe5c0c97598a2bd2555d1aa8cb08e48590dbb3da7b08b1056828838c5f61e6393ba7a0abcc9f662")
|
||||
expectedTag, _ := hex.DecodeString("76fc6ece0f4e1768cddf8853bb2d551b")
|
||||
|
||||
expected = append(expected, expectedTag...)
|
||||
|
||||
var keyArray [32]byte
|
||||
copy(keyArray[:], key)
|
||||
c := CipherAESGCM.Cipher(keyArray)
|
||||
aead := c.(aeadCipher).AEAD
|
||||
|
||||
dst := aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
assert.Equal(t, expected, dst)
|
||||
|
||||
// We expect this to fail since we are re-encrypting with a repeat IV
|
||||
assert.PanicsWithError(t, "boringcrypto: EVP_AEAD_CTX_seal failed", func() {
|
||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -29,8 +29,11 @@ type CipherState interface {
|
||||
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
||||
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
||||
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
||||
if cs, ok := s.Cipher().(CipherState); ok {
|
||||
return cs
|
||||
}
|
||||
switch cipherFunc.CipherName() {
|
||||
case CipherAESGCM.CipherName():
|
||||
case noise.CipherAESGCM.CipherName():
|
||||
return NewCipherStateAESGCM(s)
|
||||
case noise.CipherChaChaPoly.CipherName():
|
||||
return NewCipherStateChaChaPoly(s)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"testing"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
@@ -10,24 +11,30 @@ import (
|
||||
|
||||
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
roundtrip(t, NewCipherState(enc, CipherAESGCM), NewCipherState(dec, CipherAESGCM))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
roundtrip(t, NewCipherState(enc, noise.CipherChaChaPoly), NewCipherState(dec, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func TestNewCipherStateDispatch(t *testing.T) {
|
||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
if !boringEnabled && !fips140.Enabled() {
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
} else {
|
||||
// fips140
|
||||
assert.IsType(t, encA.Cipher(), NewCipherState(encA, CipherAESGCM))
|
||||
}
|
||||
|
||||
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||
enc, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
assert.Panics(t, func() {
|
||||
NewCipherState(enc, fakeCipher{})
|
||||
})
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/cipher"
|
||||
"crypto/fips140"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"unsafe"
|
||||
|
||||
// unsafe needed for go:linkname
|
||||
_ "crypto/tls"
|
||||
_ "unsafe"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// TODO: Use NewGCMWithCounterNonce or NewGCMForQUIC once available:
|
||||
// - https://github.com/golang/go/issues/73110
|
||||
// - https://github.com/golang/go/issues/79219
|
||||
// Using tls.aeadAESGCMTLS13 gives us the TLS 1.3 GCM, which also verifies
|
||||
// that the nonce is strictly increasing. This works for both boringcrypto
|
||||
// and fips140.
|
||||
//
|
||||
//go:linkname aeadAESGCMTLS13 crypto/tls.aeadAESGCMTLS13
|
||||
func aeadAESGCMTLS13(key, noncePrefix []byte) cipher.AEAD
|
||||
|
||||
type cipherFn struct {
|
||||
fn func([32]byte) noise.Cipher
|
||||
name string
|
||||
}
|
||||
|
||||
func (c cipherFn) Cipher(k [32]byte) noise.Cipher { return c.fn(k) }
|
||||
func (c cipherFn) CipherName() string { return c.name }
|
||||
|
||||
// CipherAESGCMFIPS140 is the AES256-GCM AEAD cipher (using tls.aeadAESGCMTLS13, for both boringcrypto and fips140)
|
||||
var CipherAESGCMFIPS140 noise.CipherFunc = cipherFn{cipherAESGCMFIPS140, "AESGCM"}
|
||||
|
||||
// tls.aeadAESGCMTLS13 uses a 4 byte static prefix and an 8 byte XOR mask
|
||||
var emptyNonce = []byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}
|
||||
|
||||
func cipherAESGCMFIPS140(k [32]byte) noise.Cipher {
|
||||
gcm := aeadAESGCMTLS13(k[:], emptyNonce)
|
||||
gcm = extractFIPSAEAD(gcm)
|
||||
return &aeadGCMFIPS140Cipher{
|
||||
AEAD: gcm,
|
||||
}
|
||||
}
|
||||
|
||||
type aeadGCMFIPS140Cipher struct {
|
||||
cipher.AEAD
|
||||
ready bool
|
||||
}
|
||||
|
||||
// Extract the internal FIPS GCM implementation from the tls wrapper. The TLS
|
||||
// wrapper is not thread safe around Open, so instead of locking around it we
|
||||
// can grab the internal implementation that is thread safe. This is the FIPS
|
||||
// module implementation: `crypto/internal/fips140/aes/gcm.GCMWithXORCounterNonce`
|
||||
//
|
||||
// - https://github.com/golang/go/blob/go1.26.4/src/crypto/internal/fips140/aes/gcm/gcm_nonces.go#L212-L287
|
||||
//
|
||||
// The wrapper is struct `crypto/tls.xorNonceAEAD` , with field `aead`:
|
||||
//
|
||||
// - https://github.com/golang/go/blob/go1.26.4/src/crypto/tls/cipher_suites.go#L482-L487
|
||||
//
|
||||
// This can be cleaned up once these FIPS implementations are exposed directly:
|
||||
//
|
||||
// - https://github.com/golang/go/issues/73110
|
||||
func extractFIPSAEAD(xorNonceAEAD cipher.AEAD) cipher.AEAD {
|
||||
r := reflect.ValueOf(xorNonceAEAD)
|
||||
v := r.Elem().FieldByName("aead")
|
||||
if !v.IsValid() {
|
||||
// The internal crypto/tls.xorNonceAEAD struct no longer has an `aead`
|
||||
// field. This can only happen on a Go version this code was not built
|
||||
// against; the package init() self-test guards against ever reaching
|
||||
// this at runtime, so this is a defensive fail-fast.
|
||||
panic(fmt.Sprintf("noiseutil: could not extract FIPS AEAD from %T on %s: no `aead` field (incompatible Go version)", xorNonceAEAD, runtime.Version()))
|
||||
}
|
||||
v2 := reflect.NewAt(v.Type(), unsafe.Pointer(v.UnsafeAddr())).Elem()
|
||||
aead, ok := v2.Interface().(cipher.AEAD)
|
||||
if !ok {
|
||||
panic(fmt.Sprintf("noiseutil: extracted FIPS `aead` field is %s, not a cipher.AEAD, on %s (incompatible Go version)", v2.Type(), runtime.Version()))
|
||||
}
|
||||
return aead
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) init(nonce []byte) {
|
||||
// GCMWithXORCounterNonce expects that the first call to Seal
|
||||
// is with a counter of `0`, this is how it extracts the nonce mask.
|
||||
// We can clean this up in the future when NewGCMWithCounterNonce or
|
||||
// NewGCMForQUIC are available:
|
||||
if !bytes.Equal(emptyNonce, nonce) {
|
||||
c.AEAD.Seal([]byte{}, emptyNonce, []byte{}, []byte{})
|
||||
}
|
||||
c.ready = true
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) Seal(dst, nonce, plaintext, additionalData []byte) []byte {
|
||||
if !c.ready {
|
||||
c.init(nonce)
|
||||
}
|
||||
return c.AEAD.Seal(dst, nonce, plaintext, additionalData)
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) Encrypt(out []byte, n uint64, ad, plaintext []byte) []byte {
|
||||
return c.Seal(out, aeadGCMFIPS140CipherNonce(n), plaintext, ad)
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) Decrypt(out []byte, n uint64, ad, ciphertext []byte) ([]byte, error) {
|
||||
return c.Open(out, aeadGCMFIPS140CipherNonce(n), ciphertext, ad)
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
binary.BigEndian.PutUint64(nb[4:], n)
|
||||
out = c.Seal(out, nb, plaintext, ad)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
binary.BigEndian.PutUint64(nb[4:], n)
|
||||
return c.Open(out, nb, ciphertext, ad)
|
||||
}
|
||||
|
||||
func aeadGCMFIPS140CipherNonce(n uint64) []byte {
|
||||
// GCMWithXORCounterNonce uses a 4 byte static prefix and an 8 byte nonce
|
||||
var nonce [12]byte
|
||||
binary.BigEndian.PutUint64(nonce[4:], n)
|
||||
return nonce[:]
|
||||
}
|
||||
|
||||
// init validates the go:linkname + reflection extraction and the nonce-reuse
|
||||
// protection at startup, in every build. cipherAESGCMFIPS140 relies on unexported
|
||||
// crypto/tls and crypto/internal/fips140 internals; if a future Go version changes
|
||||
// those, this fails fast with a clear message instead of panicking per-handshake
|
||||
// (or, worse, silently losing the strictly-increasing nonce check that is the whole
|
||||
// point of using this cipher). Because this file has no build tag, this self-test
|
||||
// runs even in non-FIPS builds, so the default CI lane catches an incompatible Go.
|
||||
func init() {
|
||||
var key [32]byte
|
||||
c := cipherAESGCMFIPS140(key)
|
||||
|
||||
// Verify the extracted AEAD produces a working encrypt/decrypt roundtrip.
|
||||
plaintext := []byte("nebula fips140 self-test")
|
||||
ad := []byte("ad")
|
||||
ct := c.Encrypt(nil, 1, ad, plaintext)
|
||||
pt, err := c.Decrypt(nil, 1, ad, ct)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test roundtrip failed on %s: %v", runtime.Version(), err))
|
||||
}
|
||||
if !bytes.Equal(pt, plaintext) {
|
||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test roundtrip returned wrong plaintext on %s", runtime.Version()))
|
||||
}
|
||||
|
||||
// Verify the nonce-reuse protection still fires: re-encrypting with the same
|
||||
// counter must panic. This is the guarantee we depend on for nonce safety, so
|
||||
// if the extraction ever silently yields an AEAD without it, refuse to start.
|
||||
// The strictly-increasing nonce check only exists under boringcrypto/fips140;
|
||||
// in a plain build aeadAESGCMTLS13 wraps a standard GCM that does not enforce
|
||||
// it (and CipherAESGCMFIPS140 is unused there anyway), so only assert it when
|
||||
// one of those modes is active.
|
||||
if (boringEnabled || fips140.Enabled()) && !reusePanics(c) {
|
||||
panic(fmt.Sprintf("noiseutil: FIPS AES-GCM self-test did not reject a reused nonce on %s; nonce-reuse protection is missing (incompatible Go version)", runtime.Version()))
|
||||
}
|
||||
}
|
||||
|
||||
// reusePanics reports whether re-encrypting with an already-used counter panics,
|
||||
// as GCMWithXORCounterNonce is expected to.
|
||||
func reusePanics(c noise.Cipher) (panicked bool) {
|
||||
c.Encrypt(nil, 2, nil, nil)
|
||||
defer func() {
|
||||
if recover() != nil {
|
||||
panicked = true
|
||||
}
|
||||
}()
|
||||
c.Encrypt(nil, 2, nil, nil)
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/cipher"
|
||||
"crypto/fips140"
|
||||
"encoding/hex"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// Ensure NewAESGCM validates the nonce is non-repeating
|
||||
func TestNewAESGCM(t *testing.T) {
|
||||
if !boringEnabled && !fips140.Enabled() {
|
||||
t.Skip("TestNewAESGCM is only for fips140/boringcrypto")
|
||||
}
|
||||
|
||||
key, _ := hex.DecodeString("feffe9928665731c6d6a8f9467308308feffe9928665731c6d6a8f9467308308")
|
||||
iv, _ := hex.DecodeString("00000000facedbaddecaf888")
|
||||
plaintext, _ := hex.DecodeString("d9313225f88406e5a55909c5aff5269a86a7a9531534f7da2e4c303d8a318a721c3c0c95956809532fcf0e2449a6b525b16aedf5aa0de657ba637b39")
|
||||
aad, _ := hex.DecodeString("feedfacedeadbeeffeedfacedeadbeefabaddad2")
|
||||
expected, _ := hex.DecodeString("6a65c2edd45bd63c7e29f40e3d2ed8ba2b99f4c83135383d5676652f255059ceb24863ff10afb1089db701245da87fb88d3acd5f9dd0770cac220c3c04145caf25e190aeb775e7080401c628")
|
||||
|
||||
var keyArray [32]byte
|
||||
copy(keyArray[:], key)
|
||||
c := CipherAESGCM.Cipher(keyArray)
|
||||
aead := c.(cipher.AEAD)
|
||||
|
||||
dst := aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
t.Logf("%x", dst)
|
||||
assert.Equal(t, expected, dst)
|
||||
|
||||
// We expect this to fail since we are re-encrypting with a repeat IV
|
||||
switch {
|
||||
case boringEnabled:
|
||||
assert.PanicsWithError(t, "boringcrypto: EVP_AEAD_CTX_seal failed", func() {
|
||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
})
|
||||
case fips140.Version() == "v1.0.0":
|
||||
assert.PanicsWithValue(t, "crypto/cipher: counter decreased", func() {
|
||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
})
|
||||
default:
|
||||
assert.PanicsWithValue(t, "crypto/cipher: counter decreased or remained the same", func() {
|
||||
dst = aead.Seal([]byte{}, iv, plaintext, aad)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
)
|
||||
|
||||
func init() {
|
||||
if !fips140.Enforced() {
|
||||
panic("Nebula compiled with fips140 expects FIPS140 to be enforced. Do not set GODEBUG=fips140, or if you do it must be set as GODEBUG=fips140=only")
|
||||
}
|
||||
}
|
||||
+15
-4
@@ -1,14 +1,25 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||
const EncryptLockNeeded = false
|
||||
var EncryptLockNeeded = fips140.Enabled()
|
||||
|
||||
// CipherAESGCM is the standard noise.CipherAESGCM when boringcrypto is not enabled
|
||||
var CipherAESGCM noise.CipherFunc = noise.CipherAESGCM
|
||||
var CipherAESGCM noise.CipherFunc = initAESGCM()
|
||||
|
||||
func initAESGCM() noise.CipherFunc {
|
||||
if fips140.Enabled() {
|
||||
return CipherAESGCMFIPS140
|
||||
} else {
|
||||
return noise.CipherAESGCM
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
var boringEnabled = false
|
||||
|
||||
@@ -1,14 +0,0 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestEncryptLockNeeded(t *testing.T) {
|
||||
assert.False(t, EncryptLockNeeded)
|
||||
}
|
||||
+24
-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
|
||||
@@ -504,20 +489,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
|
||||
}
|
||||
|
||||
|
||||
+4
-35
@@ -5,7 +5,6 @@ package overlay
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
@@ -251,17 +250,6 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
||||
}
|
||||
|
||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||
// Validate the device name up front so a bad tun.dev fails fast, before we
|
||||
// open /dev/net/tun or leak a file descriptor. A single %d in the name is
|
||||
// substituted by the kernel during TUNSETIFF (dev_alloc_name) with the
|
||||
// lowest number that yields an unused device name. Resolving the template
|
||||
// in the kernel keeps the pick-a-name/create-the-device pair atomic, so
|
||||
// concurrent callers can never race each other to the same name.
|
||||
tunName := c.GetString("tun.dev", "nebula%d")
|
||||
if err := validateTunName(tunName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||
if err != nil {
|
||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
||||
@@ -289,11 +277,12 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
if multiqueue {
|
||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
||||
}
|
||||
copy(req.Name[:], tunName)
|
||||
nameStr := c.GetString("tun.dev", "")
|
||||
copy(req.Name[:], nameStr)
|
||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, &NameError{
|
||||
Name: tunName,
|
||||
Name: nameStr,
|
||||
Underlying: err,
|
||||
}
|
||||
}
|
||||
@@ -309,27 +298,6 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
return t, nil
|
||||
}
|
||||
|
||||
func validateTunName(tunName string) error {
|
||||
if !strings.Contains(tunName, "%d") {
|
||||
if len(tunName) >= unix.IFNAMSIZ {
|
||||
return fmt.Errorf("tun.dev %q is not shorter than the maximum device name length of %d", tunName, unix.IFNAMSIZ)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if strings.Count(tunName, "%d") > 1 {
|
||||
return fmt.Errorf("tun.dev template %q may only contain a single %%d", tunName)
|
||||
}
|
||||
if tunName == "%d" {
|
||||
return errors.New("please don't name your tun device '%d'")
|
||||
}
|
||||
// The kernel substitutes the %d itself and requires the template, like a
|
||||
// literal name, to be NUL-terminated within IFNAMSIZ bytes.
|
||||
if len(tunName) >= unix.IFNAMSIZ {
|
||||
return fmt.Errorf("tun.dev template %q is not shorter than the maximum device name length of %d", tunName, unix.IFNAMSIZ)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||
tfd, err := newTunFd(fd)
|
||||
@@ -800,6 +768,7 @@ func (t *tun) isGatewayInVpnNetworks(gwAddr netip.Addr) bool {
|
||||
|
||||
func (t *tun) getGatewaysFromRoute(r *netlink.Route) routing.Gateways {
|
||||
var gateways routing.Gateways
|
||||
|
||||
link, err := netlink.LinkByName(t.Device)
|
||||
if err != nil {
|
||||
t.l.Error("Ignoring route update: failed to get link by name", "deviceName", t.Device)
|
||||
|
||||
@@ -3,12 +3,7 @@
|
||||
|
||||
package overlay
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
import "testing"
|
||||
|
||||
var runAdvMSSTests = []struct {
|
||||
name string
|
||||
@@ -37,39 +32,3 @@ func TestTunAdvMSS(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTunName(t *testing.T) {
|
||||
// A device name must be shorter than IFNAMSIZ (i.e. IFNAMSIZ-1 chars max).
|
||||
maxLenName := strings.Repeat("a", unix.IFNAMSIZ-1)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
tmpl string
|
||||
wantErr bool
|
||||
}{
|
||||
{"short literal name is fine", "nebula1", false},
|
||||
{"literal name at the max length is fine", maxLenName, false},
|
||||
{"literal name at IFNAMSIZ is rejected", strings.Repeat("a", unix.IFNAMSIZ), true},
|
||||
{"trailing template is fine", "nebula%d", false},
|
||||
{"mid-string template is fine", "neb%dprod", false},
|
||||
{"leading template is fine", "%dnebula", false},
|
||||
{"template at the max length is fine", strings.Repeat("a", unix.IFNAMSIZ-3) + "%d", false},
|
||||
{"template at IFNAMSIZ is rejected", strings.Repeat("a", unix.IFNAMSIZ-2) + "%d", true},
|
||||
{"bare %d is rejected", "%d", true},
|
||||
{"multiple %d is rejected", "neb%d%dprod", true},
|
||||
{"over-long template is rejected", strings.Repeat("a", unix.IFNAMSIZ-1) + "%d", true},
|
||||
{"over-long mid-string template is rejected", "neb%d" + strings.Repeat("a", unix.IFNAMSIZ-3), true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := validateTunName(tt.tmpl)
|
||||
if tt.wantErr && err == nil {
|
||||
t.Fatalf("expected an error for %q, got none", tt.tmpl)
|
||||
}
|
||||
if !tt.wantErr && err != nil {
|
||||
t.Fatalf("unexpected error for %q: %v", tt.tmpl, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -237,10 +238,17 @@ func (cs *CertState) getCertificate(v cert.Version) cert.Certificate {
|
||||
return nil
|
||||
}
|
||||
|
||||
func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string) (noise.CipherSuite, error) {
|
||||
// newCipherSuite builds the noise.CipherSuite for the given curve and cipher.
|
||||
// When fips140Enforced is true (FIPS 140-only mode), non-approved algorithms
|
||||
// (Curve25519 and ChaChaPoly) are rejected with an error. Callers pass
|
||||
// fips140.Enforced() for fips140Enforced.
|
||||
func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string, fips140Enforced bool) (noise.CipherSuite, error) {
|
||||
var dhFunc noise.DHFunc
|
||||
switch curve {
|
||||
case cert.Curve_CURVE25519:
|
||||
if fips140Enforced {
|
||||
return nil, errors.New("pki: use of Curve25519 is not allowed in FIPS 140-only mode")
|
||||
}
|
||||
dhFunc = noise.DH25519
|
||||
case cert.Curve_P256:
|
||||
if pkcs11backed {
|
||||
@@ -253,6 +261,9 @@ func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string) (noise.C
|
||||
}
|
||||
|
||||
if cipher == "chachapoly" {
|
||||
if fips140Enforced {
|
||||
return nil, errors.New("pki: use of ChaChaPoly is not allowed in FIPS 140-only mode")
|
||||
}
|
||||
return noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256), nil
|
||||
}
|
||||
return noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256), nil
|
||||
@@ -326,6 +337,10 @@ func newCertStateFromConfig(c *config.C, cipher string) (*CertState, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if fips140.Enforced() && crt.Curve() != cert.Curve_P256 {
|
||||
return nil, fmt.Errorf("pki: use of %s is not allowed in FIPS 140-only mode", crt.Curve())
|
||||
}
|
||||
|
||||
switch crt.Version() {
|
||||
case cert.Version1:
|
||||
if v1 != nil {
|
||||
@@ -405,7 +420,7 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
||||
//NOTE: We do not currently have a method to verify a public private key pair when the private key is in an hsm
|
||||
} else {
|
||||
if err := v1.VerifyPrivateKey(privateKeyCurve, privateKey); err != nil {
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert")
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -413,7 +428,7 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error marshalling v1 certificate for handshake: %w", err)
|
||||
}
|
||||
ncs, err := newCipherSuite(v1.Curve(), pkcs11backed, cipher)
|
||||
ncs, err := newCipherSuite(v1.Curve(), pkcs11backed, cipher, fips140.Enforced())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -430,7 +445,7 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
||||
//NOTE: We do not currently have a method to verify a public private key pair when the private key is in an hsm
|
||||
} else {
|
||||
if err := v2.VerifyPrivateKey(privateKeyCurve, privateKey); err != nil {
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert")
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -438,7 +453,7 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error marshalling v2 certificate for handshake: %w", err)
|
||||
}
|
||||
ncs, err := newCipherSuite(v2.Curve(), pkcs11backed, cipher)
|
||||
ncs, err := newCipherSuite(v2.Curve(), pkcs11backed, cipher, fips140.Enforced())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
+94
@@ -0,0 +1,94 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNewCipherSuite(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
curve cert.Curve
|
||||
cipher string
|
||||
fips140Enforced bool
|
||||
wantErr string
|
||||
// wantName is the full expected CipherSuite name (<DH>_<Cipher>_<Hash>),
|
||||
// only checked when wantErr is empty. Asserting the whole name makes both
|
||||
// the curve and cipher selection load-bearing.
|
||||
wantName string
|
||||
}{
|
||||
{
|
||||
name: "curve25519 aesgcm, not enforced",
|
||||
curve: cert.Curve_CURVE25519,
|
||||
cipher: "aesgcm",
|
||||
wantName: "25519_AESGCM_SHA256",
|
||||
},
|
||||
{
|
||||
name: "curve25519 chachapoly, not enforced",
|
||||
curve: cert.Curve_CURVE25519,
|
||||
cipher: "chachapoly",
|
||||
wantName: "25519_ChaChaPoly_SHA256",
|
||||
},
|
||||
{
|
||||
name: "p256 aesgcm, not enforced",
|
||||
curve: cert.Curve_P256,
|
||||
cipher: "aesgcm",
|
||||
wantName: "P256_AESGCM_SHA256",
|
||||
},
|
||||
{
|
||||
name: "p256 aesgcm, enforced is allowed",
|
||||
curve: cert.Curve_P256,
|
||||
cipher: "aesgcm",
|
||||
fips140Enforced: true,
|
||||
wantName: "P256_AESGCM_SHA256",
|
||||
},
|
||||
{
|
||||
name: "curve25519 rejected when enforced",
|
||||
curve: cert.Curve_CURVE25519,
|
||||
cipher: "aesgcm",
|
||||
fips140Enforced: true,
|
||||
wantErr: "pki: use of Curve25519 is not allowed in FIPS 140-only mode",
|
||||
},
|
||||
{
|
||||
name: "chachapoly rejected when enforced",
|
||||
curve: cert.Curve_P256,
|
||||
cipher: "chachapoly",
|
||||
fips140Enforced: true,
|
||||
wantErr: "pki: use of ChaChaPoly is not allowed in FIPS 140-only mode",
|
||||
},
|
||||
{
|
||||
// Curve is checked before cipher, so a Curve25519+ChaChaPoly
|
||||
// request reports the Curve25519 rejection.
|
||||
name: "curve25519 chachapoly rejected on curve when enforced",
|
||||
curve: cert.Curve_CURVE25519,
|
||||
cipher: "chachapoly",
|
||||
fips140Enforced: true,
|
||||
wantErr: "pki: use of Curve25519 is not allowed in FIPS 140-only mode",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cs, err := newCipherSuite(tt.curve, false, tt.cipher, tt.fips140Enforced)
|
||||
if tt.wantErr != "" {
|
||||
require.EqualError(t, err, tt.wantErr)
|
||||
assert.Nil(t, cs)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cs)
|
||||
assert.Equal(t, tt.wantName, string(cs.Name()))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewCipherSuiteUnsupportedCurve(t *testing.T) {
|
||||
cs, err := newCipherSuite(cert.Curve(99), false, "aesgcm", false)
|
||||
require.Error(t, err)
|
||||
assert.True(t, strings.HasPrefix(err.Error(), "unsupported curve:"), "got: %v", err)
|
||||
assert.Nil(t, cs)
|
||||
}
|
||||
+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
|
||||
|
||||
@@ -2,6 +2,7 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -288,9 +289,12 @@ func (s *statsServer) buildRuntime(cfg statsConfig) ([]func(), *http.Server) {
|
||||
Name: "info",
|
||||
Help: "Version information for the Nebula binary",
|
||||
ConstLabels: prometheus.Labels{
|
||||
"version": s.buildVersion,
|
||||
"goversion": runtime.Version(),
|
||||
"boringcrypto": strconv.FormatBool(boringEnabled()),
|
||||
"version": s.buildVersion,
|
||||
"goversion": runtime.Version(),
|
||||
"boringcrypto": strconv.FormatBool(boringEnabled()),
|
||||
"fips140Version": fips140.Version(),
|
||||
"fips140Enabled": strconv.FormatBool(fips140.Enabled()),
|
||||
"fips140Enforced": strconv.FormatBool(fips140.Enforced()),
|
||||
},
|
||||
})
|
||||
pr.MustRegister(g)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
+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