mirror of
https://github.com/slackhq/nebula.git
synced 2026-09-30 14:16:38 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8713d4acb1 | ||
|
|
75cbf9358d |
+31
-32
@@ -20,45 +20,44 @@ jobs:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: Smoke Docker
|
||||
run: make smoke-docker
|
||||
|
||||
- name: Smoke Docker IPv6 overlay
|
||||
run: make smoke-docker-ipv6
|
||||
|
||||
- name: Smoke Relay Docker
|
||||
run: make smoke-relay-docker
|
||||
|
||||
- name: Smoke Docker boringcrypto
|
||||
run: make boringcrypto smoke-docker
|
||||
|
||||
- name: Smoke Docker fips140
|
||||
run: make fips140-all GOALS=smoke-docker
|
||||
|
||||
timeout-minutes: 10
|
||||
|
||||
smoke-self:
|
||||
name: Run self traffic smoke test on macOS
|
||||
runs-on: macos-latest
|
||||
steps:
|
||||
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version: '1.26'
|
||||
check-latest: true
|
||||
|
||||
- name: build
|
||||
run: make bin
|
||||
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
||||
|
||||
- name: run smoke-self
|
||||
- name: setup docker image
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./smoke-self.sh
|
||||
run: ./build.sh
|
||||
|
||||
- name: run smoke
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./smoke.sh
|
||||
|
||||
- name: setup docker image ipv6
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: SMOKE_OVERLAY_IPV6=1 ./build.sh
|
||||
|
||||
- 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
|
||||
|
||||
timeout-minutes: 10
|
||||
|
||||
@@ -1,130 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# A host must be able to reach its own overlay address. Where the kernel sends
|
||||
# that traffic through the tun rather than over loopback, nebula sees it and
|
||||
# hands it straight back (immediatelyForwardToSelf), and whether the kernel
|
||||
# accepts what comes back is only answerable against a real kernel. Runs one
|
||||
# nebula on this machine as root and aims every probe at its own address.
|
||||
|
||||
set -e -x
|
||||
|
||||
set -o pipefail
|
||||
|
||||
V4=192.0.2.1
|
||||
V6=2001:db8::1
|
||||
|
||||
case "$(uname -s)" in
|
||||
Darwin) TUN_DEV=utun ;;
|
||||
*) TUN_DEV=tun0 ;;
|
||||
esac
|
||||
|
||||
ROOT="$(cd ../../.. && pwd)"
|
||||
|
||||
rm -rf build/self
|
||||
mkdir -p build/self
|
||||
cd build/self
|
||||
|
||||
cleanup() {
|
||||
echo
|
||||
echo " *** cleanup"
|
||||
echo
|
||||
|
||||
set +e
|
||||
if [ -n "$NEBULA_PID" ]
|
||||
then
|
||||
sudo kill "$NEBULA_PID"
|
||||
fi
|
||||
{ kill $(jobs -p); wait; } 2>/dev/null
|
||||
sed 's/^/ [self] /' nebula.log
|
||||
}
|
||||
|
||||
trap cleanup EXIT
|
||||
|
||||
# perl is on every platform this runs on; timeout(1) is not.
|
||||
alarm() {
|
||||
perl -e 'alarm shift; exec @ARGV' "$@"
|
||||
}
|
||||
|
||||
RESULTS=""
|
||||
FAILED=""
|
||||
probe() {
|
||||
local name="$1"
|
||||
shift
|
||||
if "$@"
|
||||
then
|
||||
RESULTS="$RESULTS $name=ok"
|
||||
else
|
||||
RESULTS="$RESULTS $name=FAIL"
|
||||
FAILED="$FAILED $name"
|
||||
fi
|
||||
}
|
||||
|
||||
# Send one datagram, then wait for the listener to have written it out.
|
||||
udp_probe() {
|
||||
echo self | alarm 5 nc -u -w1 "$1" 3000 || true
|
||||
set +x
|
||||
for _ in $(seq 1 20)
|
||||
do
|
||||
if grep -q self "$2"
|
||||
then
|
||||
set -x
|
||||
return 0
|
||||
fi
|
||||
sleep 0.25
|
||||
done
|
||||
set -x
|
||||
return 1
|
||||
}
|
||||
|
||||
"$ROOT/nebula-cert" ca -name "Smoke Test"
|
||||
"$ROOT/nebula-cert" sign -name self -networks "$V4/24,$V6/64"
|
||||
|
||||
HOST=self AM_LIGHTHOUSE=true TUN_DEV="$TUN_DEV" ../../genconfig.sh >self.yml
|
||||
|
||||
"$ROOT/nebula" -config self.yml -test
|
||||
|
||||
sudo -v
|
||||
sudo "$ROOT/nebula" -config self.yml >nebula.log 2>&1 &
|
||||
NEBULA_PID=$!
|
||||
|
||||
for _ in $(seq 1 40)
|
||||
do
|
||||
ifconfig | grep "inet6 $V6 " >/dev/null && break
|
||||
sleep 0.25
|
||||
done
|
||||
ifconfig | grep "inet $V4 "
|
||||
ifconfig | grep "inet6 $V6 "
|
||||
|
||||
nc -l "$V4" 2000 >/dev/null &
|
||||
nc -l "$V6" 2000 >/dev/null &
|
||||
nc -u -l "$V4" 3000 >udp4.txt &
|
||||
nc -u -l "$V6" 3000 >udp6.txt &
|
||||
sleep 1
|
||||
|
||||
set +x
|
||||
echo
|
||||
echo " *** Testing self traffic from $V4"
|
||||
echo
|
||||
set -x
|
||||
probe icmp4 alarm 5 ping -c1 "$V4"
|
||||
probe tcp4 alarm 5 nc -z "$V4" 2000
|
||||
probe udp4 udp_probe "$V4" udp4.txt
|
||||
|
||||
set +x
|
||||
echo
|
||||
echo " *** Testing self traffic from $V6"
|
||||
echo
|
||||
set -x
|
||||
probe icmp6 alarm 5 ping6 -c1 "$V6"
|
||||
probe tcp6 alarm 5 nc -z "$V6" 2000
|
||||
probe udp6 udp_probe "$V6" udp6.txt
|
||||
|
||||
set +x
|
||||
echo
|
||||
echo " *** self traffic:$RESULTS"
|
||||
echo
|
||||
if [ -n "$FAILED" ]
|
||||
then
|
||||
echo "self traffic failed:$FAILED" >&2
|
||||
exit 1
|
||||
fi
|
||||
@@ -228,6 +228,14 @@ try {
|
||||
Write-Host "OK: $DevName $family NlMtu=$Mtu"
|
||||
}
|
||||
|
||||
# Both are set on the same handle as the v6 NlMtu, so by now they are either applied or never will be.
|
||||
Wait-Until -TimeoutSec 30 -What "$DevName IPv6 DadTransmits=0 RouterDiscovery=Disabled" -Predicate {
|
||||
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before the v6 interface was configured" }
|
||||
$rows = @(Get-NetIPInterface -InterfaceAlias $DevName -AddressFamily IPv6 -ErrorAction SilentlyContinue)
|
||||
$rows.Count -gt 0 -and -not ($rows | Where-Object { $_.DadTransmits -ne 0 -or "$($_.RouterDiscovery)" -ne 'Disabled' })
|
||||
}
|
||||
Write-Host "OK: $DevName IPv6 DadTransmits=0 RouterDiscovery=Disabled"
|
||||
|
||||
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
|
||||
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
|
||||
|
||||
@@ -58,14 +58,9 @@ jobs:
|
||||
e2e-cmd: make e2evv
|
||||
- name: linux-boringcrypto
|
||||
os: ubuntu-latest
|
||||
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
|
||||
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"
|
||||
- name: linux-pkcs11
|
||||
os: ubuntu-latest
|
||||
build-cmd: make bin-pkcs11
|
||||
|
||||
+1
-36
@@ -7,32 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### Added
|
||||
|
||||
- New `nebula ctl <command>` subcommand, which runs any of the debug and administrative commands the sshd
|
||||
block exposes without requiring an ssh server, a host key, or authorized keys. Nebula serves them over a
|
||||
local unix socket, configured by the new `ctl` block and enabled by default at `/run/nebula/ctl.sock` on
|
||||
Linux and `/var/run/nebula/ctl.sock` elsewhere. The socket lives in a `0700` directory so filesystem
|
||||
permissions are the access control; failing to create it is logged and never prevents nebula from
|
||||
starting. Packagers running nebula under systemd will want `RuntimeDirectory=nebula` in the unit so the
|
||||
directory exists with the right ownership. Not supported on Windows yet, and never enabled on iOS or
|
||||
Android. Reloadable.
|
||||
|
||||
### Changed
|
||||
|
||||
- The ssh console now reports a real exit status for `ssh <host> <command>` rather than always reporting
|
||||
success, so commands run that way are scriptable.
|
||||
- The debug and administrative commands moved out of `ssh.go` into `commands.go` and are no longer tied to
|
||||
ssh: both the ssh console and `nebula ctl` dispatch against one shared registry, so a command added in
|
||||
one place is available over both. Embedders of the `sshd` package are affected: `sshd.NewSSHServer` now
|
||||
takes a `*diag.Registry`, `sshd.SSHServer.RegisterCommand` is gone in favor of registering on that
|
||||
registry directly, and the command types now live in the `diag` package rather than being re-exported
|
||||
from `sshd`.
|
||||
|
||||
## [1.11.1] - 2026-08-21
|
||||
|
||||
See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) milestone for a complete list of changes.
|
||||
|
||||
### Changed
|
||||
|
||||
- IPv6 packets whose next header is a protocol Nebula does not parse (SCTP, GRE, IP-in-IP, etc.) are now
|
||||
@@ -41,18 +15,11 @@ See the [v1.11.1](https://github.com/slackhq/nebula/milestone/30?closed=1) miles
|
||||
their true protocol, so only a `proto: any` rule allows them. If you carry one of these protocols over the
|
||||
overlay, confirm a `proto: any` rule covers it before upgrading, it may have been passing only through this
|
||||
bypass. (#1840)
|
||||
- Drop the dependency on `github.com/cyberdelia/go-metrics-graphite`, which has been unmaintained for over ten
|
||||
years, by inlining the small amount of code Nebula used. (#1832)
|
||||
|
||||
### Fixed
|
||||
|
||||
- The ICMPv6 type was read from the wrong byte when classifying IPv6 packets, so the echo identifier used
|
||||
for conntrack was never picked up. (#1840)
|
||||
- Enforce outbound message counter limits so a tunnel is rehandshaked before the counter can wrap, preventing
|
||||
nonce reuse. This is unreachable in practice, but is enforced as a defense-in-depth measure. (#1841)
|
||||
- Prevent `nebula-cert ca` from running out of memory on 32bit systems when generating encrypted private keys. (#1834)
|
||||
- Tolerate `ErrDumpInterrupted` when listing tun addresses on Linux, so a transient interrupted netlink dump
|
||||
no longer aborts startup. (#1835)
|
||||
|
||||
## [1.11.0] - 2026-07-23
|
||||
|
||||
@@ -917,9 +884,7 @@ created.)
|
||||
|
||||
- Initial public release.
|
||||
|
||||
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.11.1...HEAD
|
||||
[1.11.1]: https://github.com/slackhq/nebula/releases/tag/v1.11.1
|
||||
[1.11.0]: https://github.com/slackhq/nebula/releases/tag/v1.11.0
|
||||
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.10.3...HEAD
|
||||
[1.10.3]: https://github.com/slackhq/nebula/releases/tag/v1.10.3
|
||||
[1.10.2]: https://github.com/slackhq/nebula/releases/tag/v1.10.2
|
||||
[1.10.1]: https://github.com/slackhq/nebula/releases/tag/v1.10.1
|
||||
|
||||
@@ -72,17 +72,6 @@ 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
|
||||
|
||||
@@ -148,8 +137,6 @@ 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
|
||||
@@ -170,9 +157,6 @@ 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
|
||||
@@ -182,12 +166,12 @@ debug: BUILD_ARGS += -tags debug
|
||||
debug: bin
|
||||
|
||||
bin:
|
||||
$(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
|
||||
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
|
||||
|
||||
install:
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
$(GOENV) go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ./cmd/nebula-cert
|
||||
go install $(BUILD_ARGS) -ldflags "$(LDFLAGS)" ${NEBULA_CMD_PATH}
|
||||
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 -, ,$*))
|
||||
@@ -198,11 +182,8 @@ 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
|
||||
|
||||
# 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/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
||||
|
||||
build/%/nebula: .FORCE
|
||||
GOOS=$(firstword $(subst -, , $*)) \
|
||||
@@ -233,7 +214,10 @@ vet:
|
||||
go vet $(VET_FLAGS) -v ./...
|
||||
|
||||
test:
|
||||
$(TEST_ENV) go test $(TEST_FLAGS) -v ./...
|
||||
go test -v ./...
|
||||
|
||||
test-boringcrypto:
|
||||
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
||||
|
||||
test-pkcs11:
|
||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
||||
@@ -276,75 +260,29 @@ 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
|
||||
# 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
|
||||
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
|
||||
|
||||
smoke-relay-docker: BUILD_ARGS += -race
|
||||
smoke-relay-docker: GOENV += CGO_ENABLED=1
|
||||
smoke-relay-docker: bin-docker
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && $(GOENV) ./smoke-relay.sh
|
||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||
|
||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
||||
smoke-docker-ipv6: smoke-docker
|
||||
|
||||
smoke-self: bin
|
||||
cd .github/workflows/smoke/ && ./smoke-self.sh
|
||||
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 bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug 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 smoke-self test test-pkcs11 test-cov-html vet 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 debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||
.DEFAULT_GOAL := bin
|
||||
|
||||
@@ -145,27 +145,17 @@ To build nebula for a specific platform (ex, Windows):
|
||||
|
||||
See the [Makefile](Makefile) for more details on build targets
|
||||
|
||||
## Curve P256 and FIPS 140-3 mode
|
||||
## Curve P256 and BoringCrypto
|
||||
|
||||
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.
|
||||
|
||||
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.
|
||||
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:
|
||||
|
||||
```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
|
||||
|
||||
+1
-13
@@ -3,9 +3,7 @@ package main
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/fips140"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -46,13 +44,6 @@ type caFlags struct {
|
||||
subnets *string
|
||||
}
|
||||
|
||||
func defaultCurve() string {
|
||||
if fips140.Enforced() {
|
||||
return "P256"
|
||||
}
|
||||
return "25519"
|
||||
}
|
||||
|
||||
func newCaFlags() *caFlags {
|
||||
// prevent running out of memory on 32-bit systems by defaulting to
|
||||
// RFC9106's recommendation for memory-constrained environments
|
||||
@@ -83,7 +74,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", defaultArgonIterations, "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", defaultCurve(), "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
|
||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
||||
@@ -268,9 +259,6 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
||||
} else {
|
||||
switch *cf.curve {
|
||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||
if fips140.Enforced() {
|
||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
||||
}
|
||||
curve = cert.Curve_CURVE25519
|
||||
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -1,8 +1,6 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -26,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", defaultCurve(), "ECDH Curve (25519, P256)")
|
||||
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||
cf.p11url = p11Flag(cf.set)
|
||||
return &cf
|
||||
}
|
||||
@@ -63,9 +61,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
||||
} else {
|
||||
switch *cf.curve {
|
||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||
if fips140.Enforced() {
|
||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
||||
}
|
||||
pub, rawPriv = x25519Keypair()
|
||||
curve = cert.Curve_CURVE25519
|
||||
case "P256":
|
||||
|
||||
@@ -2,7 +2,6 @@ package main
|
||||
|
||||
import (
|
||||
"crypto/ecdh"
|
||||
"crypto/fips140"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"flag"
|
||||
@@ -269,10 +268,6 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
||||
}(p11Client)
|
||||
}
|
||||
|
||||
if fips140.Enforced() && curve == cert.Curve_CURVE25519 {
|
||||
return errors.New("use of Curve25519 is not allowed in FIPS 140-only mode")
|
||||
}
|
||||
|
||||
if *sf.inPubPath != "" {
|
||||
var pubCurve cert.Curve
|
||||
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -1,131 +0,0 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"os"
|
||||
"syscall"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
)
|
||||
|
||||
// ctlMain implements `nebula ctl <command> [args...]`, which runs a debug command against the
|
||||
// nebula already running on this host. Everything after the command name is forwarded to that
|
||||
// nebula verbatim and parsed there by the same flag sets the ssh console uses, so this side
|
||||
// deliberately understands as little as possible about it.
|
||||
//
|
||||
// Returns the process exit status.
|
||||
func ctlMain(argv []string) int {
|
||||
fl := flag.NewFlagSet("nebula ctl", flag.ContinueOnError)
|
||||
fl.Usage = func() {
|
||||
out := fl.Output()
|
||||
fmt.Fprintf(out, "Usage: nebula ctl [-config path] [-socket path] <command> [arguments]\n\n")
|
||||
fmt.Fprintf(out, "Runs a debug command against the running nebula on this host, over its local\n")
|
||||
fmt.Fprintf(out, "control socket. Run `nebula ctl` with no command for the list of commands.\n\n")
|
||||
fl.PrintDefaults()
|
||||
}
|
||||
|
||||
socket := fl.String("socket", "", "Path to the control socket. Overrides ctl.socket from the config")
|
||||
configPath := fl.String("config", "", "Path to the nebula config, read only to find ctl.socket")
|
||||
|
||||
// The flag package stops at the first non-flag argument, which is exactly the behaviour
|
||||
// wanted here: `nebula ctl -socket /x list-hostmap -json` consumes -socket, stops at
|
||||
// list-hostmap, and leaves the rest untouched for the daemon to parse.
|
||||
if err := fl.Parse(argv); err != nil {
|
||||
// -h is a request, not a failure.
|
||||
if errors.Is(err, flag.ErrHelp) {
|
||||
return diag.StatusOK
|
||||
}
|
||||
return diag.StatusUsage
|
||||
}
|
||||
|
||||
path := *socket
|
||||
if path == "" {
|
||||
path = ctlSocketPath(*configPath)
|
||||
}
|
||||
|
||||
if path == "" {
|
||||
fmt.Fprintln(os.Stderr, "nebula ctl: no control socket path is known for this platform, set ctl.socket in the config")
|
||||
return diag.StatusError
|
||||
}
|
||||
|
||||
client, err := diag.Dial(path)
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, ctlDialError(path, err))
|
||||
return diag.StatusError
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
args := fl.Args()
|
||||
status, err := client.Run(args, os.Stdout)
|
||||
if err != nil {
|
||||
if errors.Is(err, diag.ErrTruncated) {
|
||||
fmt.Fprintf(os.Stderr, "nebula ctl: nebula closed the connection before %s finished\n", ctlCommandName(args))
|
||||
return diag.StatusError
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "nebula ctl: %s\n", err)
|
||||
if status == diag.StatusOK {
|
||||
return diag.StatusError
|
||||
}
|
||||
}
|
||||
|
||||
return status
|
||||
}
|
||||
|
||||
// ctlSocketPath finds the socket to talk to. The platform default is the primary mechanism;
|
||||
// reading the config is the refinement for someone who moved the socket. It is best effort by
|
||||
// design, because config.DefaultPath resolves next to the nebula binary and a packaged install
|
||||
// keeps its config somewhere else entirely, so a config we cannot find is the normal case
|
||||
// rather than a failure.
|
||||
func ctlSocketPath(configPath string) string {
|
||||
if configPath == "" {
|
||||
p, err := config.DefaultPath()
|
||||
if err != nil {
|
||||
return diag.DefaultSocketPath()
|
||||
}
|
||||
configPath = p
|
||||
}
|
||||
|
||||
c := config.NewC(slog.New(slog.DiscardHandler))
|
||||
if err := c.Load(configPath); err != nil {
|
||||
return diag.DefaultSocketPath()
|
||||
}
|
||||
|
||||
return c.GetString("ctl.socket", diag.DefaultSocketPath())
|
||||
}
|
||||
|
||||
// ctlDialError turns a connect failure into something an operator can act on. These messages
|
||||
// are the entire user experience when things are not working, so they name the path and say
|
||||
// what to check.
|
||||
func ctlDialError(path string, err error) string {
|
||||
switch {
|
||||
case errors.Is(err, diag.ErrNotSupported):
|
||||
return "nebula ctl is not supported on this platform yet"
|
||||
|
||||
case errors.Is(err, fs.ErrNotExist):
|
||||
return fmt.Sprintf("nebula ctl: no control socket at %s. Is nebula running? Is ctl.enabled set to false, or ctl.socket set to another path?", path)
|
||||
|
||||
case errors.Is(err, syscall.ECONNREFUSED):
|
||||
return fmt.Sprintf("nebula ctl: found a stale socket at %s, nebula is not listening on it", path)
|
||||
|
||||
case errors.Is(err, fs.ErrPermission):
|
||||
return fmt.Sprintf("nebula ctl: permission denied opening %s. nebula ctl must run as the user nebula runs as, usually root", path)
|
||||
|
||||
default:
|
||||
return fmt.Sprintf("nebula ctl: %s", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ctlCommandName names the command for an error message, for the case where there isn't one.
|
||||
func ctlCommandName(args []string) string {
|
||||
if len(args) == 0 {
|
||||
return "the command"
|
||||
}
|
||||
|
||||
return args[0]
|
||||
}
|
||||
@@ -1,70 +0,0 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// The daemon parses the command's own flags, so this side must consume its own and forward
|
||||
// everything from the command name onwards untouched.
|
||||
func TestCtlSocketPath(t *testing.T) {
|
||||
t.Run("a config naming a socket is used", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "config.yml")
|
||||
require.NoError(t, os.WriteFile(path, []byte("ctl:\n socket: /run/somewhere/ctl.sock\n"), 0600))
|
||||
|
||||
assert.Equal(t, "/run/somewhere/ctl.sock", ctlSocketPath(path))
|
||||
})
|
||||
|
||||
t.Run("a config without a ctl block falls back to the platform default", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "config.yml")
|
||||
require.NoError(t, os.WriteFile(path, []byte("pki:\n ca: /dev/null\n"), 0600))
|
||||
|
||||
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(path))
|
||||
})
|
||||
|
||||
// A packaged install keeps its config somewhere config.DefaultPath will never look, so a
|
||||
// config we cannot read is the ordinary case and must not be fatal.
|
||||
t.Run("an unreadable config falls back to the platform default", func(t *testing.T) {
|
||||
assert.Equal(t, diag.DefaultSocketPath(), ctlSocketPath(filepath.Join(t.TempDir(), "nope.yml")))
|
||||
})
|
||||
}
|
||||
|
||||
func TestCtlDialError(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
wants string
|
||||
}{
|
||||
{"missing socket names the path and what to check", fs.ErrNotExist, "no control socket at /x/ctl.sock. Is nebula running?"},
|
||||
{"a stale socket is called stale", syscall.ECONNREFUSED, "found a stale socket at /x/ctl.sock"},
|
||||
{"permission denied suggests the right user", fs.ErrPermission, "must run as the user nebula runs as"},
|
||||
{"an unsupported platform says so", diag.ErrNotSupported, "not supported on this platform"},
|
||||
{"anything else is reported verbatim", errors.New("something else"), "something else"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Contains(t, ctlDialError("/x/ctl.sock", tt.err), tt.wants)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("a wrapped syscall error is still recognised", func(t *testing.T) {
|
||||
err := &os.SyscallError{Syscall: "connect", Err: syscall.ECONNREFUSED}
|
||||
assert.Contains(t, ctlDialError("/x/ctl.sock", err), "stale socket")
|
||||
})
|
||||
}
|
||||
|
||||
func TestCtlCommandName(t *testing.T) {
|
||||
assert.Equal(t, "print-cert", ctlCommandName([]string{"print-cert", "-json"}))
|
||||
assert.Equal(t, "the command", ctlCommandName(nil))
|
||||
}
|
||||
@@ -1,5 +0,0 @@
|
||||
//go:build fips140enforce
|
||||
|
||||
//go:debug fips140=only
|
||||
|
||||
package main
|
||||
@@ -32,26 +32,11 @@ func init() {
|
||||
}
|
||||
|
||||
func main() {
|
||||
// Subcommands are dispatched before flag.Parse, because flag.Parse stops at the first
|
||||
// non-flag argument and everything after `ctl` has to reach the running nebula's own flag
|
||||
// parser untouched. Nothing here looks at -json or a vpn address.
|
||||
if len(os.Args) > 1 && os.Args[1] == "ctl" {
|
||||
os.Exit(ctlMain(os.Args[2:]))
|
||||
}
|
||||
|
||||
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
|
||||
configTest := flag.Bool("test", false, "Test the config and print the end result. Non zero exit indicates a faulty config")
|
||||
printVersion := flag.Bool("version", false, "Print version")
|
||||
printUsage := flag.Bool("help", false, "Print command line usage")
|
||||
|
||||
flag.Usage = func() {
|
||||
out := flag.CommandLine.Output()
|
||||
fmt.Fprintf(out, "Usage of %s:\n", os.Args[0])
|
||||
flag.PrintDefaults()
|
||||
fmt.Fprintf(out, "\nCommands:\n")
|
||||
fmt.Fprintf(out, " ctl [command]\n\tRun a debug command against the running nebula on this host.\n\tRun `nebula ctl` on its own for the list of commands.\n")
|
||||
}
|
||||
|
||||
flag.Parse()
|
||||
|
||||
if *printVersion {
|
||||
|
||||
-922
@@ -1,922 +0,0 @@
|
||||
package nebula
|
||||
|
||||
// The commands nebula exposes for debugging and administration. They are transport neutral:
|
||||
// the ssh console in ssh.go and the `nebula ctl` socket in ctl.go both dispatch against the
|
||||
// registry attachCommands fills in, and a command cannot tell which one invoked it. Adding a
|
||||
// command here makes it available over both.
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"runtime/pprof"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
type listHostMapFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
ByIndex bool
|
||||
}
|
||||
|
||||
type printCertFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
Raw bool
|
||||
}
|
||||
|
||||
type printTunnelFlags struct {
|
||||
Pretty bool
|
||||
}
|
||||
|
||||
type changeRemoteFlags struct {
|
||||
Address string
|
||||
}
|
||||
|
||||
type closeTunnelFlags struct {
|
||||
LocalOnly bool
|
||||
}
|
||||
|
||||
type createTunnelFlags struct {
|
||||
Address string
|
||||
}
|
||||
|
||||
type deviceInfoFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
}
|
||||
|
||||
func attachCommands(l *slog.Logger, c *config.C, reg *diag.Registry, f *Interface) {
|
||||
// sandboxDir defaults to a dir in temp. The intention is that end user will
|
||||
// create this dir as needed. Overriding this config value to "" allows
|
||||
// writing to anywhere in the system.
|
||||
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
|
||||
// The key is spelled for both transports now: the profile writers are reachable over
|
||||
// `nebula ctl` as well, but sshd.sandbox_dir keeps working for anyone already setting it.
|
||||
sandboxDir := c.GetString("ctl.sandbox_dir", c.GetString("sshd.sandbox_dir", defaultDir))
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "list-hostmap",
|
||||
ShortDescription: "List all known previously connected hosts",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := listHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdListHostMap(f.hostMap, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "list-pending-hostmap",
|
||||
ShortDescription: "List all handshaking hosts",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := listHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdListHostMap(f.handshakeManager, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "list-lighthouse-addrmap",
|
||||
ShortDescription: "List all lighthouse map entries",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := listHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdListLighthouseMap(f.lightHouse, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "reload",
|
||||
ShortDescription: "Reloads configuration from disk, same as sending HUP to the process",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdReload(c, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "start-cpu-profile",
|
||||
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdStartCpuProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "stop-cpu-profile",
|
||||
ShortDescription: "Stops a cpu profile and writes output to the previously provided file",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
pprof.StopCPUProfile()
|
||||
return w.WriteLine("If a CPU profile was running it is now stopped")
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "save-heap-profile",
|
||||
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdGetHeapProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "mutex-profile-fraction",
|
||||
ShortDescription: "Gets or sets runtime.SetMutexProfileFraction",
|
||||
Callback: cmdMutexProfileFraction,
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "save-mutex-profile",
|
||||
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdGetMutexProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "log-level",
|
||||
ShortDescription: "Gets or sets the current log level",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdLogLevel(l, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "log-format",
|
||||
ShortDescription: "Gets or sets the current log format",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdLogFormat(l, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "version",
|
||||
ShortDescription: "Prints the currently running version of nebula",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdVersion(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "device-info",
|
||||
ShortDescription: "Prints information about the network device.",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := deviceInfoFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdDeviceInfo(f, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "print-cert",
|
||||
ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := printCertFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdPrintCert(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "print-tunnel",
|
||||
ShortDescription: "Prints json details about a tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := printTunnelFlags{}
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdPrintTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "print-relays",
|
||||
ShortDescription: "Prints json details about all relay info",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := printTunnelFlags{}
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdPrintRelays(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "change-remote",
|
||||
ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := changeRemoteFlags{}
|
||||
fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdChangeRemote(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "close-tunnel",
|
||||
ShortDescription: "Closes a tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := closeTunnelFlags{}
|
||||
fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdCloseTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "create-tunnel",
|
||||
ShortDescription: "Creates a tunnel for the provided vpn address",
|
||||
Help: "The lighthouses will be queried for real addresses but you can provide one as well.",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := createTunnelFlags{}
|
||||
fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdCreateTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
reg.RegisterCommand(&diag.Command{
|
||||
Name: "query-lighthouse",
|
||||
ShortDescription: "Query the lighthouses for the provided vpn address",
|
||||
Help: "This command is asynchronous. Only currently known udp addresses will be printed.",
|
||||
Callback: func(fs any, a []string, w diag.StringWriter) error {
|
||||
return cmdQueryLighthouse(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func cmdListHostMap(hl controlHostLister, a any, w diag.StringWriter) error {
|
||||
fs, ok := a.(*listHostMapFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
|
||||
}
|
||||
|
||||
var hm []ControlHostInfo
|
||||
if fs.ByIndex {
|
||||
hm = listHostMapIndexes(hl)
|
||||
} else {
|
||||
hm = listHostMapHosts(hl)
|
||||
}
|
||||
|
||||
sort.Slice(hm, func(i, j int) bool {
|
||||
return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0
|
||||
})
|
||||
|
||||
if fs.Json || fs.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if fs.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
err := js.Encode(hm)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
} else {
|
||||
for _, v := range hm {
|
||||
err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func cmdListLighthouseMap(lightHouse *LightHouse, a any, w diag.StringWriter) error {
|
||||
fs, ok := a.(*listHostMapFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be listHostMapFlags but was %+v", a)
|
||||
}
|
||||
|
||||
type lighthouseInfo struct {
|
||||
VpnAddr string `json:"vpnAddr"`
|
||||
Addrs *CacheMap `json:"addrs"`
|
||||
}
|
||||
|
||||
lightHouse.RLock()
|
||||
addrMap := make([]lighthouseInfo, len(lightHouse.addrMap))
|
||||
x := 0
|
||||
for k, v := range lightHouse.addrMap {
|
||||
addrMap[x] = lighthouseInfo{
|
||||
VpnAddr: k.String(),
|
||||
Addrs: v.CopyCache(),
|
||||
}
|
||||
x++
|
||||
}
|
||||
lightHouse.RUnlock()
|
||||
|
||||
sort.Slice(addrMap, func(i, j int) bool {
|
||||
return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0
|
||||
})
|
||||
|
||||
if fs.Json || fs.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if fs.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
err := js.Encode(addrMap)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
} else {
|
||||
for _, v := range addrMap {
|
||||
b, err := json.Marshal(v.Addrs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b)))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// sanitizeFilePath validates that the given file path is within the sandbox directory.
|
||||
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
|
||||
func sanitizeFilePath(sandboxDir, filePath string) (string, error) {
|
||||
if sandboxDir == "" {
|
||||
return filePath, nil
|
||||
}
|
||||
|
||||
// Clean and resolve the path relative to the sandbox directory
|
||||
if !filepath.IsAbs(filePath) {
|
||||
filePath = filepath.Join(sandboxDir, filePath)
|
||||
}
|
||||
cleaned := filepath.Clean(filePath)
|
||||
|
||||
// Ensure the resolved path is within the sandbox directory
|
||||
cleanedSandbox := filepath.Clean(sandboxDir)
|
||||
if cleaned == cleanedSandbox {
|
||||
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
|
||||
}
|
||||
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
|
||||
}
|
||||
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func cmdStartCpuProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
err := w.WriteLine("No path to write profile provided")
|
||||
return err
|
||||
}
|
||||
|
||||
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = pprof.StartCPUProfile(file)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a))
|
||||
return err
|
||||
}
|
||||
|
||||
func cmdVersion(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
return w.WriteLine(fmt.Sprintf("%s", ifce.version))
|
||||
}
|
||||
|
||||
func cmdQueryLighthouse(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
var cm *CacheMap
|
||||
rl := ifce.lightHouse.Query(vpnAddr)
|
||||
if rl != nil {
|
||||
cm = rl.CopyCache()
|
||||
}
|
||||
return json.NewEncoder(w.GetWriter()).Encode(cm)
|
||||
}
|
||||
|
||||
func cmdCloseTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
flags, ok := fs.(*closeTunnelFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be closeTunnelFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||
}
|
||||
|
||||
if !flags.LocalOnly {
|
||||
ifce.send(
|
||||
header.CloseTunnel,
|
||||
0,
|
||||
hostInfo.ConnectionState,
|
||||
hostInfo,
|
||||
[]byte{},
|
||||
make([]byte, 12, 12),
|
||||
make([]byte, mtu),
|
||||
)
|
||||
}
|
||||
|
||||
ifce.closeTunnel(hostInfo)
|
||||
return w.WriteLine("Closed")
|
||||
}
|
||||
|
||||
func cmdCreateTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
flags, ok := fs.(*createTunnelFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be createTunnelFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Tunnel already exists"))
|
||||
}
|
||||
|
||||
hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Tunnel already handshaking"))
|
||||
}
|
||||
|
||||
var addr netip.AddrPort
|
||||
if flags.Address != "" {
|
||||
addr, err = netip.ParseAddrPort(flags.Address)
|
||||
if err != nil {
|
||||
return w.WriteLine("Address could not be parsed")
|
||||
}
|
||||
}
|
||||
|
||||
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
|
||||
if addr.IsValid() {
|
||||
hostInfo.SetRemote(addr)
|
||||
}
|
||||
|
||||
return w.WriteLine("Created")
|
||||
}
|
||||
|
||||
func cmdChangeRemote(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
flags, ok := fs.(*changeRemoteFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be changeRemoteFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
if flags.Address == "" {
|
||||
return w.WriteLine("No address was provided")
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddrPort(flags.Address)
|
||||
if err != nil {
|
||||
return w.WriteLine("Address could not be parsed")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||
}
|
||||
|
||||
hostInfo.SetRemote(addr)
|
||||
return w.WriteLine("Changed")
|
||||
}
|
||||
|
||||
func cmdGetHeapProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No path to write profile provided")
|
||||
}
|
||||
|
||||
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = pprof.WriteHeapProfile(file)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a))
|
||||
return err
|
||||
}
|
||||
|
||||
func cmdMutexProfileFraction(fs any, a []string, w diag.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
rate := runtime.SetMutexProfileFraction(-1)
|
||||
return w.WriteLine(fmt.Sprintf("Current value: %d", rate))
|
||||
}
|
||||
|
||||
newRate, err := strconv.Atoi(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0]))
|
||||
}
|
||||
|
||||
oldRate := runtime.SetMutexProfileFraction(newRate)
|
||||
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
|
||||
}
|
||||
|
||||
func cmdGetMutexProfile(sandboxDir string, fs any, a []string, w diag.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No path to write profile provided")
|
||||
}
|
||||
|
||||
filePath, err := sanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
mutexProfile := pprof.Lookup("mutex")
|
||||
if mutexProfile == nil {
|
||||
return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")")
|
||||
}
|
||||
|
||||
err = mutexProfile.WriteTo(file, 0)
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||
}
|
||||
|
||||
return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a))
|
||||
}
|
||||
|
||||
func cmdLogLevel(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
|
||||
ctrl, ok := l.Handler().(interface {
|
||||
GetLevel() slog.Level
|
||||
SetLevel(slog.Level)
|
||||
})
|
||||
if !ok {
|
||||
return w.WriteLine("Log level is not reconfigurable on this logger")
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||
}
|
||||
|
||||
level, err := logging.ParseLevel(strings.ToLower(a[0]))
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a))
|
||||
}
|
||||
|
||||
ctrl.SetLevel(level)
|
||||
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||
}
|
||||
|
||||
func cmdLogFormat(l *slog.Logger, fs any, a []string, w diag.StringWriter) error {
|
||||
ctrl, ok := l.Handler().(interface {
|
||||
GetFormat() string
|
||||
SetFormat(string) error
|
||||
})
|
||||
if !ok {
|
||||
return w.WriteLine("Log format is not reconfigurable on this logger")
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||
}
|
||||
|
||||
if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil {
|
||||
return err
|
||||
}
|
||||
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||
}
|
||||
|
||||
func cmdPrintCert(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
args, ok := fs.(*printCertFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be printCertFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
cert := ifce.pki.getCertState().GetDefaultCertificate()
|
||||
if len(a) > 0 {
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||
}
|
||||
|
||||
cert = hostInfo.GetCert().Certificate
|
||||
}
|
||||
|
||||
if args.Json || args.Pretty {
|
||||
b, err := cert.MarshalJSON()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if args.Pretty {
|
||||
buf := new(bytes.Buffer)
|
||||
err := json.Indent(buf, b, "", " ")
|
||||
b = buf.Bytes()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return w.WriteBytes(b)
|
||||
}
|
||||
|
||||
if args.Raw {
|
||||
b, err := cert.MarshalPEM()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return w.WriteBytes(b)
|
||||
}
|
||||
|
||||
return w.WriteLine(cert.String())
|
||||
}
|
||||
|
||||
func cmdPrintRelays(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
args, ok := fs.(*printTunnelFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
relays := map[uint32]*HostInfo{}
|
||||
ifce.hostMap.Lock()
|
||||
maps.Copy(relays, ifce.hostMap.Relays)
|
||||
ifce.hostMap.Unlock()
|
||||
|
||||
type RelayFor struct {
|
||||
Error error
|
||||
Type string
|
||||
State string
|
||||
PeerAddr netip.Addr
|
||||
LocalIndex uint32
|
||||
RemoteIndex uint32
|
||||
RelayedThrough []netip.Addr
|
||||
}
|
||||
|
||||
type RelayOutput struct {
|
||||
NebulaAddr netip.Addr
|
||||
RelayForAddrs []RelayFor
|
||||
}
|
||||
|
||||
type CmdOutput struct {
|
||||
Relays []*RelayOutput
|
||||
}
|
||||
|
||||
co := CmdOutput{}
|
||||
|
||||
enc := json.NewEncoder(w.GetWriter())
|
||||
|
||||
if args.Pretty {
|
||||
enc.SetIndent("", " ")
|
||||
}
|
||||
|
||||
for k, v := range relays {
|
||||
ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]}
|
||||
co.Relays = append(co.Relays, &ro)
|
||||
relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0])
|
||||
if relayHI == nil {
|
||||
ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")})
|
||||
continue
|
||||
}
|
||||
for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() {
|
||||
rf := RelayFor{Error: nil}
|
||||
r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr)
|
||||
if ok {
|
||||
t := ""
|
||||
switch r.Type {
|
||||
case ForwardingType:
|
||||
t = "forwarding"
|
||||
case TerminalType:
|
||||
t = "terminal"
|
||||
default:
|
||||
t = "unknown"
|
||||
}
|
||||
|
||||
s := ""
|
||||
switch r.State {
|
||||
case Requested:
|
||||
s = "requested"
|
||||
case Established:
|
||||
s = "established"
|
||||
default:
|
||||
s = "unknown"
|
||||
}
|
||||
|
||||
rf.LocalIndex = r.LocalIndex
|
||||
rf.RemoteIndex = r.RemoteIndex
|
||||
rf.PeerAddr = r.PeerAddr
|
||||
rf.Type = t
|
||||
rf.State = s
|
||||
if rf.LocalIndex != k {
|
||||
rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k)
|
||||
}
|
||||
}
|
||||
relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if relayedHI != nil {
|
||||
rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...)
|
||||
}
|
||||
|
||||
ro.RelayForAddrs = append(ro.RelayForAddrs, rf)
|
||||
}
|
||||
}
|
||||
err := enc.Encode(co)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func cmdPrintTunnel(ifce *Interface, fs any, a []string, w diag.StringWriter) error {
|
||||
args, ok := fs.(*printTunnelFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be printTunnelFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||
}
|
||||
|
||||
enc := json.NewEncoder(w.GetWriter())
|
||||
if args.Pretty {
|
||||
enc.SetIndent("", " ")
|
||||
}
|
||||
|
||||
return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges()))
|
||||
}
|
||||
|
||||
func cmdDeviceInfo(ifce *Interface, fs any, w diag.StringWriter) error {
|
||||
|
||||
data := struct {
|
||||
Name string `json:"name"`
|
||||
Cidr []netip.Prefix `json:"cidr"`
|
||||
}{
|
||||
Name: ifce.inside.Name(),
|
||||
Cidr: make([]netip.Prefix, len(ifce.inside.Networks())),
|
||||
}
|
||||
|
||||
copy(data.Cidr, ifce.inside.Networks())
|
||||
|
||||
flags, ok := fs.(*deviceInfoFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be deviceInfoFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if flags.Json || flags.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if flags.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
return js.Encode(data)
|
||||
} else {
|
||||
return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr))
|
||||
}
|
||||
}
|
||||
|
||||
func cmdReload(c *config.C, w diag.StringWriter) error {
|
||||
err := w.WriteLine("Reloading config")
|
||||
c.ReloadConfig()
|
||||
return err
|
||||
}
|
||||
@@ -1,69 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"log/slog"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// attachedCommands is every command nebula exposes. The ssh console and `nebula ctl` dispatch
|
||||
// against this one set, so this list is the contract for both transports.
|
||||
var attachedCommands = []string{
|
||||
"change-remote",
|
||||
"close-tunnel",
|
||||
"create-tunnel",
|
||||
"device-info",
|
||||
"list-hostmap",
|
||||
"list-lighthouse-addrmap",
|
||||
"list-pending-hostmap",
|
||||
"log-format",
|
||||
"log-level",
|
||||
"mutex-profile-fraction",
|
||||
"print-cert",
|
||||
"print-relays",
|
||||
"print-tunnel",
|
||||
"query-lighthouse",
|
||||
"reload",
|
||||
"save-heap-profile",
|
||||
"save-mutex-profile",
|
||||
"start-cpu-profile",
|
||||
"stop-cpu-profile",
|
||||
"version",
|
||||
}
|
||||
|
||||
func TestAttachCommands(t *testing.T) {
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
reg := diag.NewRegistry()
|
||||
|
||||
// The callbacks capture these but do not touch them until a command runs, and this test
|
||||
// only registers and asks for help.
|
||||
attachCommands(l, config.NewC(l), reg, &Interface{})
|
||||
|
||||
t.Run("every command is registered", func(t *testing.T) {
|
||||
for _, name := range attachedCommands {
|
||||
assert.Equal(t, []string{name}, reg.Match(name), "%s is not registered", name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("help is available for every command", func(t *testing.T) {
|
||||
for _, name := range attachedCommands {
|
||||
buf := &bytes.Buffer{}
|
||||
require.NoError(t, reg.DispatchArgs([]string{"help", name}, diag.NewWriter(buf)), name)
|
||||
assert.Contains(t, buf.String(), name+" - ", name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("the command list names them all", func(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
require.NoError(t, reg.DispatchArgs(nil, diag.NewWriter(buf)))
|
||||
|
||||
for _, name := range attachedCommands {
|
||||
assert.Contains(t, buf.String(), name+" - ", name)
|
||||
}
|
||||
})
|
||||
}
|
||||
+8
-14
@@ -105,18 +105,11 @@ func (cm *connectionManager) getInactivityTimeout() time.Duration {
|
||||
}
|
||||
|
||||
func (cm *connectionManager) In(h *HostInfo) {
|
||||
h.markIn()
|
||||
h.in.Store(true)
|
||||
}
|
||||
|
||||
// OutNoRebind records outbound traffic without consuming the rebind epoch, for relayed sends: the direct path
|
||||
// to the relay consumes the edge, the via send must not.
|
||||
func (cm *connectionManager) OutNoRebind(h *HostInfo) {
|
||||
h.markOutOnly()
|
||||
}
|
||||
|
||||
// Out records outbound traffic and reports whether we rebound since this tunnel last sent
|
||||
func (cm *connectionManager) Out(h *HostInfo) bool {
|
||||
return h.markOut(cm.intf.rebindEpoch.Load())
|
||||
func (cm *connectionManager) Out(h *HostInfo) {
|
||||
h.out.Store(true)
|
||||
}
|
||||
|
||||
func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
||||
@@ -135,7 +128,8 @@ func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
||||
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
||||
// resets the state for this local index
|
||||
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
|
||||
in, out := h.takeTraffic()
|
||||
in := h.in.Swap(false)
|
||||
out := h.out.Swap(false)
|
||||
if in || out {
|
||||
h.lastUsed = now
|
||||
}
|
||||
@@ -352,7 +346,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
||||
)
|
||||
}
|
||||
hostinfo.setPendingDeletion(false)
|
||||
hostinfo.pendingDeletion.Store(false)
|
||||
|
||||
if mainHostInfo {
|
||||
decision = tryRehandshake
|
||||
@@ -375,7 +369,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
return decision, hostinfo, primary
|
||||
}
|
||||
|
||||
if hostinfo.isPendingDeletion() {
|
||||
if hostinfo.pendingDeletion.Load() {
|
||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
||||
@@ -426,7 +420,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
}
|
||||
}
|
||||
|
||||
hostinfo.setPendingDeletion(true)
|
||||
hostinfo.pendingDeletion.Store(true)
|
||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
|
||||
return decision, hostinfo, nil
|
||||
}
|
||||
|
||||
+36
-36
@@ -86,25 +86,25 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
||||
// We saw traffic out to vpnIp
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
|
||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
|
||||
// Do another traffic check tick, this host should be pending deletion now
|
||||
nc.Out(hostinfo)
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.True(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
@@ -168,33 +168,33 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||
// We saw traffic out to vpnIp
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
|
||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
|
||||
// Do another traffic check tick, this host should be pending deletion now
|
||||
nc.Out(hostinfo)
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.True(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
// We saw traffic, should no longer be pending deletion
|
||||
nc.In(hostinfo)
|
||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
}
|
||||
@@ -326,31 +326,31 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
|
||||
nc.Out(hostinfo)
|
||||
nc.In(hostinfo)
|
||||
assert.True(t, hostinfo.sentSinceCheck())
|
||||
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.True(t, hostinfo.out.Load())
|
||||
assert.True(t, hostinfo.in.Load())
|
||||
|
||||
now := time.Now()
|
||||
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
|
||||
assert.Equal(t, tryRehandshake, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
|
||||
assert.Equal(t, doNothing, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
|
||||
// Do another traffic check tick, should still not be pending deletion
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
|
||||
assert.Equal(t, doNothing, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
|
||||
@@ -358,9 +358,9 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
|
||||
assert.Equal(t, closeTunnel, decision)
|
||||
assert.Equal(t, now, hostinfo.lastUsed)
|
||||
assert.False(t, hostinfo.isPendingDeletion())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
|
||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
assert.False(t, hostinfo.in.Load())
|
||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||
}
|
||||
|
||||
@@ -12,7 +12,6 @@ import (
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -118,33 +117,7 @@ func TestSendNoMetricsDropsExhausted(t *testing.T) {
|
||||
|
||||
// The crossing send is refused: it records an exhaustion drop and never reaches connectionManager.Out.
|
||||
assert.Equal(t, int64(1), f.messageMetrics.txExhausted.Count())
|
||||
assert.False(t, hostinfo.sentSinceCheck())
|
||||
}
|
||||
|
||||
// TestSendNoMetricsCloseTunnelKeepsRebindEpoch pins that a closing tunnel does not consume a rebind, a later
|
||||
// packet on a re-established tunnel still needs that edge to trigger the far-side punch.
|
||||
func TestSendNoMetricsCloseTunnelKeepsRebindEpoch(t *testing.T) {
|
||||
initR, _ := runTestHandshake(t)
|
||||
ci, err := newConnectionStateFromResult(initR)
|
||||
require.NoError(t, err)
|
||||
|
||||
f := &Interface{
|
||||
l: test.NewLogger(),
|
||||
messageMetrics: &MessageMetrics{txExhausted: metrics.NewCounter()},
|
||||
writers: []udp.Conn{udp.NoopConn{}},
|
||||
connectionManager: &connectionManager{},
|
||||
}
|
||||
hostinfo := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")}, ConnectionState: ci}
|
||||
|
||||
// Tunnel is on epoch 0, then we rebind.
|
||||
hostinfo.markOut(0)
|
||||
f.rebindEpoch.Add(1)
|
||||
|
||||
remote := netip.MustParseAddrPort("10.0.0.2:4242")
|
||||
f.sendNoMetrics(header.CloseTunnel, 0, ci, hostinfo, remote, []byte{}, make([]byte, 12), make([]byte, mtu), 0)
|
||||
|
||||
// markOut at the new epoch still reports the move, so the edge was preserved.
|
||||
assert.True(t, hostinfo.markOut(1), "a CloseTunnel send must not consume the rebind epoch")
|
||||
assert.False(t, hostinfo.out.Load())
|
||||
}
|
||||
|
||||
func TestNewConnectionStateFromResult(t *testing.T) {
|
||||
|
||||
+1
-5
@@ -50,7 +50,6 @@ type Control struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
sshStart func()
|
||||
ctlStart func()
|
||||
statsStart func()
|
||||
dnsStart func()
|
||||
lighthouseStart func()
|
||||
@@ -100,9 +99,6 @@ func (c *Control) Start() error {
|
||||
if c.sshStart != nil {
|
||||
go c.sshStart()
|
||||
}
|
||||
if c.ctlStart != nil {
|
||||
go c.ctlStart()
|
||||
}
|
||||
if c.statsStart != nil {
|
||||
go c.statsStart()
|
||||
}
|
||||
@@ -216,7 +212,7 @@ func (c *Control) RebindUDPServer() {
|
||||
c.f.lightHouse.SendUpdate()
|
||||
|
||||
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
||||
c.f.rebindEpoch.Add(1)
|
||||
c.f.rebindCount++
|
||||
}
|
||||
|
||||
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
||||
|
||||
@@ -123,16 +123,6 @@ func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
|
||||
c.f.lightHouse.localAddrsFn = fn
|
||||
}
|
||||
|
||||
// GetRebindEpochFor returns the rebind epoch a tunnel last sent under, so a test can tell whether a send
|
||||
// consumed the epoch edge without having to infer it from lighthouse traffic.
|
||||
func (c *Control) GetRebindEpochFor(vpnAddr netip.Addr) (uint32, bool) {
|
||||
h := c.f.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if h == nil {
|
||||
return 0, false
|
||||
}
|
||||
return h.state.Load() >> stateEpochShift, true
|
||||
}
|
||||
|
||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
||||
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
|
||||
if hostinfo == nil {
|
||||
|
||||
@@ -1,234 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
|
||||
// ctlConfig is the parsed form of the `ctl` config block. It is comparable so that a reload
|
||||
// can tell "nothing changed" from "the socket moved" with ==.
|
||||
type ctlConfig struct {
|
||||
enabled bool
|
||||
socket string
|
||||
|
||||
// explicit records that the operator named a socket path rather than taking the platform
|
||||
// default. It only affects how loudly a failure to listen is reported: an unprivileged
|
||||
// nebula that cannot create /run/nebula is a normal deployment, not a problem to shout
|
||||
// about on every upgrade, but a path someone chose deliberately failing to bind is.
|
||||
explicit bool
|
||||
}
|
||||
|
||||
// ctlServer owns the unix socket `nebula ctl` connects to. It exposes the same command
|
||||
// registry the ssh console does, minus the ceremony of running an ssh server: the socket is
|
||||
// local only and guarded by filesystem permissions, so it needs no keys.
|
||||
//
|
||||
// The lifecycle mirrors statsServer: the constructor wires the reload callback, reload
|
||||
// records config and reconciles a running listener, Start builds and serves the runtime, and
|
||||
// Stop tears it down.
|
||||
type ctlServer struct {
|
||||
l *slog.Logger
|
||||
ctx context.Context
|
||||
srv *diag.Server
|
||||
|
||||
runMu sync.Mutex
|
||||
runCfg *ctlConfig
|
||||
run *ctlRuntime
|
||||
}
|
||||
|
||||
// ctlRuntime is the live state owned by a single Start invocation.
|
||||
type ctlRuntime struct {
|
||||
cancel context.CancelFunc
|
||||
listener net.Listener
|
||||
}
|
||||
|
||||
// newCtlServerFromConfig builds a ctlServer, parses the config, and registers a reload
|
||||
// callback. It deliberately does not start listening: there is no interface yet, and
|
||||
// Control.Start is what launches the first runtime. The callback is registered before the
|
||||
// config is parsed so a SIGHUP can fix a bad block even if the first parse failed.
|
||||
//
|
||||
// reg is only held, never read, until Start runs. That is what lets this be constructed
|
||||
// before attachCommands has populated the registry.
|
||||
func newCtlServerFromConfig(ctx context.Context, l *slog.Logger, c *config.C, reg *diag.Registry) (*ctlServer, error) {
|
||||
s := &ctlServer{
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
srv: diag.NewServer(l, reg),
|
||||
}
|
||||
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if err := s.reload(c, false); err != nil {
|
||||
s.l.Error("Failed to reload ctl from config", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
if err := s.reload(c, true); err != nil {
|
||||
return s, err
|
||||
}
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// loadCtlConfig parses and validates the `ctl` block. An empty socket path while enabled is
|
||||
// not an error: it means the platform has no default and the operator did not name one, so
|
||||
// there is simply nothing to listen on.
|
||||
func loadCtlConfig(c *config.C) (ctlConfig, error) {
|
||||
cfg := ctlConfig{
|
||||
enabled: c.GetBool("ctl.enabled", true),
|
||||
socket: c.GetString("ctl.socket", diag.DefaultSocketPath()),
|
||||
explicit: c.IsSet("ctl.socket"),
|
||||
}
|
||||
|
||||
if cfg.enabled && cfg.socket != "" && !filepath.IsAbs(cfg.socket) {
|
||||
return cfg, util.NewContextualError("ctl.socket must be an absolute path", m{"path": cfg.socket}, nil)
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// reload parses the config and records it, then reconciles the running listener against it:
|
||||
//
|
||||
// - newly enabled -> spawn Start
|
||||
// - newly disabled -> Stop the runtime
|
||||
// - socket moved (still enabled) -> Stop the old, Start the new
|
||||
// - no change -> no-op
|
||||
//
|
||||
// On the initial call it only records configuration; Control.Start is what launches the first
|
||||
// runtime via ctlStart. There is no interface to serve yet at that point.
|
||||
func (s *ctlServer) reload(c *config.C, initial bool) error {
|
||||
newCfg, err := loadCtlConfig(c)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
s.runMu.Lock()
|
||||
sameCfg := s.runCfg != nil && *s.runCfg == newCfg
|
||||
s.runCfg = &newCfg
|
||||
running := s.run != nil
|
||||
s.runMu.Unlock()
|
||||
|
||||
if initial || sameCfg {
|
||||
return nil
|
||||
}
|
||||
|
||||
if running {
|
||||
s.Stop()
|
||||
}
|
||||
|
||||
if newCfg.enabled && newCfg.socket != "" {
|
||||
go s.Start()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Start binds the socket and serves until Stop is called or ctx fires. Safe to call when ctl
|
||||
// is disabled or already running: both no-op.
|
||||
func (s *ctlServer) Start() {
|
||||
s.runMu.Lock()
|
||||
if s.ctx.Err() != nil || s.run != nil || s.runCfg == nil {
|
||||
s.runMu.Unlock()
|
||||
return
|
||||
}
|
||||
cfg := *s.runCfg
|
||||
s.runMu.Unlock()
|
||||
|
||||
if !cfg.enabled || cfg.socket == "" {
|
||||
if cfg.enabled {
|
||||
s.l.Info("ctl has no socket path on this platform, `nebula ctl` will not be available",
|
||||
"hint", "set ctl.socket to enable it",
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
listener, err := diag.Listen(cfg.socket)
|
||||
if err != nil {
|
||||
// A default path nebula cannot create is an ordinary state for an unprivileged
|
||||
// install; a path the operator chose failing to bind is something they want to know
|
||||
// about. Either way ctl is optional and nebula carries on without it.
|
||||
if cfg.explicit {
|
||||
s.l.Error("Failed to listen on the ctl socket", "ctlSocket", cfg.socket, "error", err)
|
||||
} else {
|
||||
s.l.Info("Not serving the ctl socket, `nebula ctl` will not be available",
|
||||
"ctlSocket", cfg.socket,
|
||||
"error", err,
|
||||
"hint", "set ctl.socket to a path nebula can write, or ctl.enabled to false",
|
||||
)
|
||||
}
|
||||
|
||||
// Drop the cached config so a SIGHUP retries once the underlying problem is fixed,
|
||||
// even when the config itself is unchanged.
|
||||
s.runMu.Lock()
|
||||
if s.runCfg != nil && *s.runCfg == cfg {
|
||||
s.runCfg = nil
|
||||
}
|
||||
s.runMu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
runCtx, cancel := context.WithCancel(s.ctx)
|
||||
rt := &ctlRuntime{cancel: cancel, listener: listener}
|
||||
|
||||
s.runMu.Lock()
|
||||
// Losing the race against a Stop or a competing Start means this listener is already
|
||||
// obsolete. Close it rather than serving a socket nobody will tear down.
|
||||
if s.ctx.Err() != nil || s.run != nil {
|
||||
s.runMu.Unlock()
|
||||
cancel()
|
||||
_ = listener.Close()
|
||||
return
|
||||
}
|
||||
s.run = rt
|
||||
s.runMu.Unlock()
|
||||
|
||||
s.l.Info("ctl socket is listening", "ctlSocket", cfg.socket)
|
||||
|
||||
err = s.srv.Serve(runCtx, listener)
|
||||
if err != nil {
|
||||
s.l.Error("The ctl listener stopped", "ctlSocket", cfg.socket, "error", err)
|
||||
}
|
||||
|
||||
// Clear our runtime only if nothing has replaced it.
|
||||
s.runMu.Lock()
|
||||
if s.run == rt {
|
||||
rt.cancel()
|
||||
s.run = nil
|
||||
if err != nil {
|
||||
// An unclean exit leaves runCfg cached as if it were applied, so drop it and let a
|
||||
// SIGHUP retry.
|
||||
s.runCfg = nil
|
||||
}
|
||||
}
|
||||
s.runMu.Unlock()
|
||||
}
|
||||
|
||||
// Stop closes the listener and unlinks the socket. It deliberately does not touch connections
|
||||
// that are already being served: `nebula ctl reload` runs every reload callback inline on its
|
||||
// own connection, including this one, and hanging up on it would truncate the response to a
|
||||
// reload that actually succeeded.
|
||||
//
|
||||
// The socket file is removed by net.UnixListener's unlink-on-close, so there is no os.Remove
|
||||
// here; doing it by hand would delete a successor's socket after a fast reload.
|
||||
func (s *ctlServer) Stop() {
|
||||
s.runMu.Lock()
|
||||
rt := s.run
|
||||
s.run = nil
|
||||
s.runMu.Unlock()
|
||||
|
||||
if rt == nil {
|
||||
return
|
||||
}
|
||||
|
||||
rt.cancel()
|
||||
if err := rt.listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||
s.l.Warn("Failed to close the ctl listener", "error", err)
|
||||
}
|
||||
}
|
||||
-305
@@ -1,305 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newTestCtlServer(t *testing.T) (*ctlServer, *config.C) {
|
||||
t.Helper()
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
|
||||
return &ctlServer{
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
srv: diag.NewServer(l, diag.NewRegistry()),
|
||||
}, config.NewC(l)
|
||||
}
|
||||
|
||||
func setCtlConfig(c *config.C, m map[string]any) {
|
||||
c.Settings["ctl"] = m
|
||||
}
|
||||
|
||||
func currentCtlRuntime(s *ctlServer) *ctlRuntime {
|
||||
s.runMu.Lock()
|
||||
defer s.runMu.Unlock()
|
||||
return s.run
|
||||
}
|
||||
|
||||
// testCtlSocket returns a short socket path, see the note in diag/server_test.go about
|
||||
// sun_path on darwin.
|
||||
func testCtlSocket(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
dir, err := os.MkdirTemp("/tmp", "nebctl")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||
|
||||
return filepath.Join(dir, "ctl.sock")
|
||||
}
|
||||
|
||||
func startCtl(t *testing.T, s *ctlServer) chan struct{} {
|
||||
t.Helper()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
s.Start()
|
||||
close(done)
|
||||
}()
|
||||
return done
|
||||
}
|
||||
|
||||
func requireCtlStopped(t *testing.T, done chan struct{}) {
|
||||
t.Helper()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("ctl Start did not return after Stop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCtlServer_loadConfig(t *testing.T) {
|
||||
t.Run("defaults to enabled at the platform path", func(t *testing.T) {
|
||||
_, c := newTestCtlServer(t)
|
||||
|
||||
cfg, err := loadCtlConfig(c)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, cfg.enabled)
|
||||
assert.Equal(t, diag.DefaultSocketPath(), cfg.socket)
|
||||
assert.False(t, cfg.explicit)
|
||||
})
|
||||
|
||||
t.Run("an operator chosen path is recorded as explicit", func(t *testing.T) {
|
||||
_, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"socket": "/run/somewhere/ctl.sock"})
|
||||
|
||||
cfg, err := loadCtlConfig(c)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "/run/somewhere/ctl.sock", cfg.socket)
|
||||
assert.True(t, cfg.explicit)
|
||||
})
|
||||
|
||||
t.Run("a relative path is rejected", func(t *testing.T) {
|
||||
_, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"socket": "ctl.sock"})
|
||||
|
||||
_, err := loadCtlConfig(c)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "must be an absolute path")
|
||||
})
|
||||
|
||||
t.Run("a relative path is not rejected when ctl is off", func(t *testing.T) {
|
||||
_, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"enabled": false, "socket": "ctl.sock"})
|
||||
|
||||
_, err := loadCtlConfig(c)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCtlServer_reload(t *testing.T) {
|
||||
t.Run("the initial reload records config without listening", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||
|
||||
require.NoError(t, s.reload(c, true))
|
||||
assert.Nil(t, currentCtlRuntime(s), "Control.Start is what starts listening")
|
||||
})
|
||||
|
||||
t.Run("enabling on reload starts listening", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
path := testCtlSocket(t)
|
||||
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
|
||||
require.NoError(t, s.reload(c, false))
|
||||
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
assert.FileExists(t, path)
|
||||
|
||||
s.Stop()
|
||||
})
|
||||
|
||||
t.Run("disabling on reload stops listening and unlinks", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
path := testCtlSocket(t)
|
||||
setCtlConfig(c, map[string]any{"enabled": true, "socket": path})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
done := startCtl(t, s)
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
|
||||
setCtlConfig(c, map[string]any{"enabled": false, "socket": path})
|
||||
require.NoError(t, s.reload(c, false))
|
||||
|
||||
requireCtlStopped(t, done)
|
||||
assert.Nil(t, currentCtlRuntime(s))
|
||||
assert.NoFileExists(t, path)
|
||||
})
|
||||
|
||||
t.Run("moving the socket restarts at the new path", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
oldPath := testCtlSocket(t)
|
||||
newPath := testCtlSocket(t)
|
||||
setCtlConfig(c, map[string]any{"socket": oldPath})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
done := startCtl(t, s)
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
require.FileExists(t, oldPath)
|
||||
|
||||
setCtlConfig(c, map[string]any{"socket": newPath})
|
||||
require.NoError(t, s.reload(c, false))
|
||||
requireCtlStopped(t, done)
|
||||
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
assert.FileExists(t, newPath)
|
||||
assert.NoFileExists(t, oldPath, "the old socket should have been unlinked")
|
||||
|
||||
s.Stop()
|
||||
})
|
||||
|
||||
t.Run("an unchanged config leaves the listener alone", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
startCtl(t, s)
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
before := currentCtlRuntime(s)
|
||||
|
||||
require.NoError(t, s.reload(c, false))
|
||||
assert.Same(t, before, currentCtlRuntime(s), "the runtime should not have been replaced")
|
||||
|
||||
s.Stop()
|
||||
})
|
||||
}
|
||||
|
||||
func TestCtlServer_Start(t *testing.T) {
|
||||
t.Run("a command can be run over the socket", func(t *testing.T) {
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
|
||||
reg := diag.NewRegistry()
|
||||
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, reg)}
|
||||
c := config.NewC(l)
|
||||
|
||||
path := testCtlSocket(t)
|
||||
setCtlConfig(c, map[string]any{"socket": path})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
startCtl(t, s)
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
|
||||
client, err := diag.Dial(path)
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
out := &testWriter{}
|
||||
status, err := client.Run([]string{"help"}, out)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, diag.StatusOK, status)
|
||||
assert.Contains(t, out.String(), "Available commands:")
|
||||
|
||||
s.Stop()
|
||||
})
|
||||
|
||||
t.Run("Start is a no-op when ctl is disabled", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"enabled": false, "socket": testCtlSocket(t)})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
s.Start()
|
||||
assert.Nil(t, currentCtlRuntime(s))
|
||||
})
|
||||
|
||||
t.Run("Start is a no-op with no socket path for this platform", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"enabled": true, "socket": ""})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
s.Start()
|
||||
assert.Nil(t, currentCtlRuntime(s))
|
||||
})
|
||||
|
||||
t.Run("Start is a no-op after the context is cancelled", func(t *testing.T) {
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
s := &ctlServer{l: l, ctx: ctx, srv: diag.NewServer(l, diag.NewRegistry())}
|
||||
c := config.NewC(l)
|
||||
|
||||
path := testCtlSocket(t)
|
||||
setCtlConfig(c, map[string]any{"socket": path})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
cancel()
|
||||
|
||||
s.Start()
|
||||
assert.Nil(t, currentCtlRuntime(s))
|
||||
assert.NoFileExists(t, path)
|
||||
})
|
||||
|
||||
// A path nebula cannot bind must not stop it from running, and a SIGHUP with the same
|
||||
// config has to be able to retry once the problem is fixed.
|
||||
t.Run("a listen failure is survivable and retried on the next reload", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
path := testCtlSocket(t)
|
||||
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
|
||||
require.NoError(t, os.WriteFile(path, []byte("in the way"), 0600))
|
||||
|
||||
setCtlConfig(c, map[string]any{"socket": path})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
s.Start()
|
||||
assert.Nil(t, currentCtlRuntime(s))
|
||||
|
||||
s.runMu.Lock()
|
||||
cachedCfg := s.runCfg
|
||||
s.runMu.Unlock()
|
||||
assert.Nil(t, cachedCfg, "the cached config should be dropped so a reload retries")
|
||||
|
||||
require.NoError(t, os.Remove(path))
|
||||
require.NoError(t, s.reload(c, false))
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
|
||||
s.Stop()
|
||||
})
|
||||
|
||||
t.Run("Stop is idempotent", func(t *testing.T) {
|
||||
s, c := newTestCtlServer(t)
|
||||
setCtlConfig(c, map[string]any{"socket": testCtlSocket(t)})
|
||||
require.NoError(t, s.reload(c, true))
|
||||
|
||||
done := startCtl(t, s)
|
||||
waitFor(t, func() bool { return currentCtlRuntime(s) != nil })
|
||||
|
||||
s.Stop()
|
||||
requireCtlStopped(t, done)
|
||||
assert.NotPanics(t, s.Stop)
|
||||
})
|
||||
}
|
||||
|
||||
// testWriter collects command output.
|
||||
type testWriter struct{ b []byte }
|
||||
|
||||
func (w *testWriter) Write(p []byte) (int, error) {
|
||||
w.b = append(w.b, p...)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (w *testWriter) String() string { return string(w.b) }
|
||||
@@ -1,41 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"io"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// dialTimeout bounds the connect only. A command may take as long as it likes to answer.
|
||||
const dialTimeout = 2 * time.Second
|
||||
|
||||
// Client is a connection to a nebula serving the ctl socket. It carries exactly one command.
|
||||
type Client struct {
|
||||
conn net.Conn
|
||||
}
|
||||
|
||||
// Dial connects to the nebula serving at path. On a platform without socket support the
|
||||
// returned error wraps ErrNotSupported.
|
||||
func Dial(path string) (*Client, error) {
|
||||
conn, err := dialSocket(path, dialTimeout)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Client{conn: conn}, nil
|
||||
}
|
||||
|
||||
// Run sends args and streams the command's output to out, returning the command's exit
|
||||
// status. A non-nil error means the exchange itself failed and the status means nothing.
|
||||
func (c *Client) Run(args []string, out io.Writer) (int, error) {
|
||||
if err := writeRequest(c.conn, args); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
return readResponse(bufio.NewReader(c.conn), out)
|
||||
}
|
||||
|
||||
func (c *Client) Close() error {
|
||||
return c.conn.Close()
|
||||
}
|
||||
-226
@@ -1,226 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// The ctl protocol is one request, one response, one connection.
|
||||
//
|
||||
// The request is a single JSON line. argv travels as a list rather than a joined string so
|
||||
// that a path with a space in it survives the trip; the client already has a real argv from
|
||||
// the operating system and re-splitting it would only ever lose information.
|
||||
//
|
||||
// The response is a stream of frames rather than raw bytes followed by a status line,
|
||||
// because there is no sentinel that is safe to look for: `print-cert -raw` emits arbitrary
|
||||
// PEM and `list-hostmap -json` emits arbitrary JSON, either of which could contain whatever
|
||||
// terminator we picked.
|
||||
const (
|
||||
// ProtoVersion is the only request version this build understands. An unknown version
|
||||
// gets a legible error rather than a hang, which is the whole point of sending it.
|
||||
ProtoVersion = 1
|
||||
|
||||
// frameOutput carries raw command output, destined for the client's stdout.
|
||||
frameOutput = 0x01
|
||||
// frameEnd carries a JSON endPayload and is the last frame on a connection.
|
||||
frameEnd = 0x02
|
||||
// frameStderr is reserved. Commands write to a single writer today, so there is nothing
|
||||
// to put in it, but holding the number means adding one later needs no version bump.
|
||||
frameStderr = 0x03
|
||||
|
||||
// maxFrame bounds a single frame's payload. Larger writes are split across frames.
|
||||
maxFrame = 64 * 1024
|
||||
// maxRequest bounds the request line, so a client that never sends a newline cannot make
|
||||
// nebula buffer without limit.
|
||||
maxRequest = 64 * 1024
|
||||
// outputBuffer is what keeps json.NewEncoder(w.GetWriter()) from emitting a frame per
|
||||
// token; output accumulates here and flushes in useful sized chunks.
|
||||
outputBuffer = 32 * 1024
|
||||
)
|
||||
|
||||
// ErrTruncated means the connection ended before the end frame arrived, which is how a
|
||||
// client notices that nebula died or was torn down partway through a command.
|
||||
var ErrTruncated = errors.New("connection closed before the command finished")
|
||||
|
||||
// request is the JSON line a client sends.
|
||||
type request struct {
|
||||
Version int `json:"version"`
|
||||
Args []string `json:"args"`
|
||||
}
|
||||
|
||||
// endPayload is the JSON body of the end frame. Error is set only when Status is non-zero
|
||||
// and describes a failure to run the command, not a failure the command itself reported.
|
||||
type endPayload struct {
|
||||
Status int `json:"status"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// writeRequest sends the request line.
|
||||
func writeRequest(w io.Writer, args []string) error {
|
||||
b, err := json.Marshal(request{Version: ProtoVersion, Args: args})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if len(b)+1 > maxRequest {
|
||||
return fmt.Errorf("command line is too long: %d bytes", len(b))
|
||||
}
|
||||
|
||||
_, err = w.Write(append(b, '\n'))
|
||||
return err
|
||||
}
|
||||
|
||||
// readRequest reads and validates one request line.
|
||||
func readRequest(r *bufio.Reader) (request, error) {
|
||||
var req request
|
||||
|
||||
line, err := readLimitedLine(r, maxRequest)
|
||||
if err != nil {
|
||||
return req, err
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(line, &req); err != nil {
|
||||
return req, fmt.Errorf("malformed request: %w", err)
|
||||
}
|
||||
|
||||
if req.Version != ProtoVersion {
|
||||
return req, fmt.Errorf("unsupported protocol version %d, this nebula speaks version %d", req.Version, ProtoVersion)
|
||||
}
|
||||
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// readLimitedLine reads through the next newline, refusing a line longer than limit rather
|
||||
// than buffering whatever an unfriendly client decides to send.
|
||||
func readLimitedLine(r *bufio.Reader, limit int) ([]byte, error) {
|
||||
line := make([]byte, 0, 256)
|
||||
for {
|
||||
b, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if b == '\n' {
|
||||
return line, nil
|
||||
}
|
||||
|
||||
if len(line) >= limit {
|
||||
return nil, fmt.Errorf("request exceeded %d bytes without a newline", limit)
|
||||
}
|
||||
|
||||
line = append(line, b)
|
||||
}
|
||||
}
|
||||
|
||||
// frameWriter turns writes into output frames. It is handed to commands wrapped in a
|
||||
// bufio.Writer, so a command that makes many small writes does not make many small frames.
|
||||
type frameWriter struct {
|
||||
w io.Writer
|
||||
}
|
||||
|
||||
func (f *frameWriter) Write(b []byte) (int, error) {
|
||||
written := 0
|
||||
for {
|
||||
chunk := b[written:]
|
||||
if len(chunk) > maxFrame {
|
||||
chunk = chunk[:maxFrame]
|
||||
}
|
||||
|
||||
if err := writeFrame(f.w, frameOutput, chunk); err != nil {
|
||||
return written, err
|
||||
}
|
||||
|
||||
written += len(chunk)
|
||||
if written == len(b) {
|
||||
return written, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writeFrame emits one frame: a type byte, a big endian length, then the payload.
|
||||
func writeFrame(w io.Writer, kind byte, payload []byte) error {
|
||||
var hdr [5]byte
|
||||
hdr[0] = kind
|
||||
binary.BigEndian.PutUint32(hdr[1:], uint32(len(payload)))
|
||||
|
||||
if _, err := w.Write(hdr[:]); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if len(payload) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, err := w.Write(payload)
|
||||
return err
|
||||
}
|
||||
|
||||
// writeEnd emits the final frame. A transport error here is unreportable by definition, the
|
||||
// connection is the only channel we have.
|
||||
func writeEnd(w io.Writer, status int, msg string) error {
|
||||
b, err := json.Marshal(endPayload{Status: status, Error: msg})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return writeFrame(w, frameEnd, b)
|
||||
}
|
||||
|
||||
// readResponse consumes frames until the end frame, copying output to out. It returns the
|
||||
// command's exit status. A non-nil error means the exchange failed and the status is
|
||||
// meaningless.
|
||||
func readResponse(r io.Reader, out io.Writer) (int, error) {
|
||||
var hdr [5]byte
|
||||
|
||||
for {
|
||||
if _, err := io.ReadFull(r, hdr[:]); err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
return 0, ErrTruncated
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
|
||||
length := binary.BigEndian.Uint32(hdr[1:])
|
||||
if length > maxFrame {
|
||||
return 0, fmt.Errorf("frame of %d bytes exceeds the %d byte maximum", length, maxFrame)
|
||||
}
|
||||
|
||||
payload := make([]byte, length)
|
||||
if _, err := io.ReadFull(r, payload); err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
return 0, ErrTruncated
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
|
||||
switch hdr[0] {
|
||||
case frameOutput:
|
||||
if _, err := out.Write(payload); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
case frameEnd:
|
||||
var end endPayload
|
||||
if err := json.Unmarshal(payload, &end); err != nil {
|
||||
return 0, fmt.Errorf("malformed end frame: %w", err)
|
||||
}
|
||||
|
||||
if end.Error != "" {
|
||||
return end.Status, errors.New(end.Error)
|
||||
}
|
||||
|
||||
return end.Status, nil
|
||||
|
||||
case frameStderr:
|
||||
// Reserved and unused by this build. Skipping rather than failing means an older
|
||||
// client stays usable against a newer nebula that starts sending them.
|
||||
|
||||
default:
|
||||
return 0, fmt.Errorf("unknown frame type 0x%02x", hdr[0])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,144 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestRequestRoundTrip(t *testing.T) {
|
||||
t.Run("argv survives a round trip, spaces and all", func(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
args := []string{"start-cpu-profile", "/tmp/a path.pb.gz", "-json"}
|
||||
require.NoError(t, writeRequest(buf, args))
|
||||
|
||||
req, err := readRequest(bufio.NewReader(buf))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, ProtoVersion, req.Version)
|
||||
assert.Equal(t, args, req.Args)
|
||||
})
|
||||
|
||||
t.Run("an unknown version is refused by name", func(t *testing.T) {
|
||||
r := bufio.NewReader(strings.NewReader(`{"version":99,"args":["version"]}` + "\n"))
|
||||
|
||||
_, err := readRequest(r)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "unsupported protocol version 99")
|
||||
})
|
||||
|
||||
t.Run("malformed json is refused", func(t *testing.T) {
|
||||
r := bufio.NewReader(strings.NewReader("not json\n"))
|
||||
|
||||
_, err := readRequest(r)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "malformed request")
|
||||
})
|
||||
|
||||
t.Run("a line without a newline is bounded rather than buffered forever", func(t *testing.T) {
|
||||
r := bufio.NewReader(strings.NewReader(strings.Repeat("a", maxRequest+10)))
|
||||
|
||||
_, err := readRequest(r)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "without a newline")
|
||||
})
|
||||
}
|
||||
|
||||
func TestResponseRoundTrip(t *testing.T) {
|
||||
t.Run("output and status survive a round trip", func(t *testing.T) {
|
||||
wire := &bytes.Buffer{}
|
||||
w := bufio.NewWriterSize(&frameWriter{w: wire}, outputBuffer)
|
||||
require.NoError(t, NewWriter(w).WriteLine("hello"))
|
||||
require.NoError(t, w.Flush())
|
||||
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||
|
||||
out := &bytes.Buffer{}
|
||||
status, err := readResponse(wire, out)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, "hello\n", out.String())
|
||||
})
|
||||
|
||||
// print-cert -raw and list-hostmap -json both emit arbitrary bytes, so a payload larger
|
||||
// than one frame has to reassemble exactly.
|
||||
t.Run("a payload larger than one frame reassembles byte for byte", func(t *testing.T) {
|
||||
big := bytes.Repeat([]byte("nebula"), maxFrame)
|
||||
|
||||
wire := &bytes.Buffer{}
|
||||
fw := &frameWriter{w: wire}
|
||||
n, err := fw.Write(big)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, len(big), n)
|
||||
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||
|
||||
out := &bytes.Buffer{}
|
||||
status, err := readResponse(wire, out)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, big, out.Bytes())
|
||||
})
|
||||
|
||||
t.Run("a non-zero status carries its message", func(t *testing.T) {
|
||||
wire := &bytes.Buffer{}
|
||||
require.NoError(t, writeEnd(wire, StatusError, "it went wrong"))
|
||||
|
||||
status, err := readResponse(wire, &bytes.Buffer{})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, StatusError, status)
|
||||
assert.Contains(t, err.Error(), "it went wrong")
|
||||
})
|
||||
|
||||
// This is how the CLI notices a nebula that died mid-command rather than silently
|
||||
// reporting whatever partial output it managed to read.
|
||||
t.Run("a stream ending without an end frame is truncated, not successful", func(t *testing.T) {
|
||||
wire := &bytes.Buffer{}
|
||||
_, err := (&frameWriter{w: wire}).Write([]byte("partial"))
|
||||
require.NoError(t, err)
|
||||
|
||||
out := &bytes.Buffer{}
|
||||
_, err = readResponse(wire, out)
|
||||
assert.ErrorIs(t, err, ErrTruncated)
|
||||
})
|
||||
|
||||
t.Run("a truncated frame header is truncated, not successful", func(t *testing.T) {
|
||||
_, err := readResponse(bytes.NewReader([]byte{frameOutput, 0x00}), &bytes.Buffer{})
|
||||
assert.ErrorIs(t, err, ErrTruncated)
|
||||
})
|
||||
|
||||
t.Run("an oversized frame is refused rather than allocated", func(t *testing.T) {
|
||||
var hdr [5]byte
|
||||
hdr[0] = frameOutput
|
||||
binary.BigEndian.PutUint32(hdr[1:], maxFrame+1)
|
||||
|
||||
_, err := readResponse(bytes.NewReader(hdr[:]), &bytes.Buffer{})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeds")
|
||||
})
|
||||
|
||||
// A reserved frame an older client does not understand must not break it.
|
||||
t.Run("a reserved frame type is skipped", func(t *testing.T) {
|
||||
wire := &bytes.Buffer{}
|
||||
require.NoError(t, writeFrame(wire, frameStderr, []byte("future")))
|
||||
require.NoError(t, writeFrame(wire, frameOutput, []byte("now")))
|
||||
require.NoError(t, writeEnd(wire, StatusOK, ""))
|
||||
|
||||
out := &bytes.Buffer{}
|
||||
status, err := readResponse(wire, out)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, "now", out.String())
|
||||
})
|
||||
|
||||
t.Run("an unknown frame type is an error", func(t *testing.T) {
|
||||
wire := &bytes.Buffer{}
|
||||
require.NoError(t, writeFrame(wire, 0x7f, nil))
|
||||
|
||||
_, err := readResponse(wire, &bytes.Buffer{})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "unknown frame type")
|
||||
})
|
||||
}
|
||||
@@ -1,125 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/anmitsu/go-shlex"
|
||||
"github.com/armon/go-radix"
|
||||
)
|
||||
|
||||
// Registry is the set of commands nebula exposes for debugging and administration. It is
|
||||
// transport neutral: the ssh console and the `nebula ctl` unix socket dispatch against the
|
||||
// same registry, and neither knows the other exists.
|
||||
//
|
||||
// Registration is expected to happen once during startup, before any transport is serving,
|
||||
// but the lock makes a late RegisterCommand safe rather than a data race waiting to happen.
|
||||
type Registry struct {
|
||||
mu sync.RWMutex
|
||||
commands *radix.Tree
|
||||
}
|
||||
|
||||
// NewRegistry returns a registry containing only `help`. Everything else is attached by
|
||||
// the caller, see attachCommands in the nebula package.
|
||||
func NewRegistry() *Registry {
|
||||
r := &Registry{commands: radix.New()}
|
||||
|
||||
r.RegisterCommand(&Command{
|
||||
Name: "help",
|
||||
ShortDescription: "prints available commands or help <command> for specific usage info",
|
||||
Callback: func(a any, args []string, w StringWriter) error {
|
||||
return r.help(args, w)
|
||||
},
|
||||
})
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// RegisterCommand adds a command that a user can run.
|
||||
func (r *Registry) RegisterCommand(c *Command) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.commands.Insert(c.Name, c)
|
||||
}
|
||||
|
||||
// Clone returns an independent copy sharing no tree with the original. The ssh session uses
|
||||
// this so the `logout` command it adds for itself is invisible to every other session, and
|
||||
// to `nebula ctl`.
|
||||
func (r *Registry) Clone() *Registry {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return &Registry{commands: radix.NewFromMap(r.commands.ToMap())}
|
||||
}
|
||||
|
||||
// Match returns every registered command name carrying the given prefix, for tab completion.
|
||||
func (r *Registry) Match(prefix string) []string {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return matchCommand(r.commands, prefix)
|
||||
}
|
||||
|
||||
// Dispatch splits line the way a shell would and runs the result. The ssh console uses this
|
||||
// because a terminal only ever hands it a line; a transport that already has a real argv
|
||||
// should call DispatchArgs instead rather than round tripping through a quoting parser.
|
||||
func (r *Registry) Dispatch(line string, w StringWriter) error {
|
||||
args, err := shlex.Split(line, true)
|
||||
if err != nil {
|
||||
if wErr := w.WriteLine(fmt.Sprintf("Unable to parse command: %s", err)); wErr != nil {
|
||||
return wErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
return r.DispatchArgs(args, w)
|
||||
}
|
||||
|
||||
// DispatchArgs runs args[0] with args[1:] as its arguments, writing everything the command
|
||||
// produces to w. An empty args dumps the command list, matching what an empty line does on
|
||||
// the ssh console.
|
||||
//
|
||||
// Callbacks report user facing problems as prose on w and return nil by convention, so a
|
||||
// non-nil error here means the command could not be run at all: ErrUnknownCommand, an
|
||||
// ErrUsage wrapped flag failure, or an internal failure a callback chose to surface.
|
||||
func (r *Registry) DispatchArgs(args []string, w StringWriter) error {
|
||||
if len(args) == 0 {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
dumpCommands(r.commands, w)
|
||||
return nil
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
cmd, err := lookupCommand(r.commands, args[0])
|
||||
r.mu.RUnlock()
|
||||
if err != nil {
|
||||
if wErr := w.WriteLine(fmt.Sprintf("Command lookup failed: %s", err)); wErr != nil {
|
||||
return wErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
if cmd == nil {
|
||||
if wErr := w.WriteLine(fmt.Sprintf("Did not understand: %s", args[0])); wErr != nil {
|
||||
return wErr
|
||||
}
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
dumpCommands(r.commands, w)
|
||||
return fmt.Errorf("%w: %s", ErrUnknownCommand, args[0])
|
||||
}
|
||||
|
||||
// -h and -help anywhere in the arguments mean the user wants to know how the command
|
||||
// works, not to run it.
|
||||
if checkHelpArgs(args) {
|
||||
return r.help([]string{cmd.Name}, w)
|
||||
}
|
||||
|
||||
return execCommand(cmd, args[1:], w)
|
||||
}
|
||||
|
||||
// help renders the command list, or one command's usage, onto w.
|
||||
func (r *Registry) help(args []string, w StringWriter) error {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return helpCallback(r.commands, args, w)
|
||||
}
|
||||
@@ -1,167 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"flag"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type testFlags struct {
|
||||
Json bool
|
||||
}
|
||||
|
||||
// testCommand builds a command carrying a flag set, recording what the callback was actually
|
||||
// handed so a test can assert on it.
|
||||
func testCommand(name string, seen *any, args *[]string) *Command {
|
||||
return &Command{
|
||||
Name: name,
|
||||
ShortDescription: name + " short description",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
f := &testFlags{}
|
||||
fl.BoolVar(&f.Json, "json", false, "outputs json")
|
||||
return fl, f
|
||||
},
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
if seen != nil {
|
||||
*seen = fs
|
||||
}
|
||||
if args != nil {
|
||||
*args = a
|
||||
}
|
||||
return w.WriteLine("ran " + name)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newTestRegistry(t *testing.T) (*Registry, *bytes.Buffer, StringWriter) {
|
||||
t.Helper()
|
||||
buf := &bytes.Buffer{}
|
||||
return NewRegistry(), buf, NewWriter(buf)
|
||||
}
|
||||
|
||||
func TestRegistryDispatch(t *testing.T) {
|
||||
t.Run("a new registry knows help and nothing else", func(t *testing.T) {
|
||||
r, buf, w := newTestRegistry(t)
|
||||
|
||||
require.NoError(t, r.DispatchArgs([]string{"help"}, w))
|
||||
assert.Contains(t, buf.String(), "help -")
|
||||
})
|
||||
|
||||
t.Run("empty args dump the command list, matching an empty line on the console", func(t *testing.T) {
|
||||
r, buf, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", nil, nil))
|
||||
|
||||
require.NoError(t, r.DispatchArgs(nil, w))
|
||||
assert.Contains(t, buf.String(), "Available commands:")
|
||||
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
|
||||
})
|
||||
|
||||
t.Run("an unknown command reports ErrUnknownCommand and still tells the user", func(t *testing.T) {
|
||||
r, buf, w := newTestRegistry(t)
|
||||
|
||||
err := r.DispatchArgs([]string{"nope"}, w)
|
||||
require.ErrorIs(t, err, ErrUnknownCommand)
|
||||
assert.Contains(t, buf.String(), "Did not understand: nope")
|
||||
assert.Contains(t, buf.String(), "Available commands:")
|
||||
})
|
||||
|
||||
// This is the hazard the ctl transport has to preserve: every callback in ssh.go begins by
|
||||
// type asserting fs to its own concrete flags struct. Reach a callback without going
|
||||
// through Command.Flags and every one of them fails.
|
||||
t.Run("a callback is handed the concrete struct its Flags callback returned", func(t *testing.T) {
|
||||
var seen any
|
||||
r, _, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||
|
||||
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json"}, w))
|
||||
|
||||
flags, ok := seen.(*testFlags)
|
||||
require.True(t, ok, "callback was handed %T, not *testFlags", seen)
|
||||
assert.True(t, flags.Json)
|
||||
})
|
||||
|
||||
t.Run("positional arguments survive flag parsing", func(t *testing.T) {
|
||||
var args []string
|
||||
r, _, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", nil, &args))
|
||||
|
||||
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-json", "10.0.0.1"}, w))
|
||||
assert.Equal(t, []string{"10.0.0.1"}, args)
|
||||
})
|
||||
|
||||
// Documents stdlib flag behaviour rather than endorsing it: parsing stops at the first
|
||||
// positional, so a flag written after one is silently a positional too.
|
||||
t.Run("a flag after a positional is not parsed as a flag", func(t *testing.T) {
|
||||
var seen any
|
||||
var args []string
|
||||
r, _, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", &seen, &args))
|
||||
|
||||
require.NoError(t, r.DispatchArgs([]string{"do-thing", "10.0.0.1", "-json"}, w))
|
||||
assert.False(t, seen.(*testFlags).Json)
|
||||
assert.Equal(t, []string{"10.0.0.1", "-json"}, args)
|
||||
})
|
||||
|
||||
t.Run("a bad flag reports ErrUsage and writes the usage text", func(t *testing.T) {
|
||||
r, buf, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", nil, nil))
|
||||
|
||||
err := r.DispatchArgs([]string{"do-thing", "-nope"}, w)
|
||||
require.ErrorIs(t, err, ErrUsage)
|
||||
assert.Contains(t, buf.String(), "flag provided but not defined")
|
||||
})
|
||||
|
||||
t.Run("-h anywhere routes to help instead of running the command", func(t *testing.T) {
|
||||
var seen any
|
||||
r, buf, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||
|
||||
require.NoError(t, r.DispatchArgs([]string{"do-thing", "-h"}, w))
|
||||
assert.Nil(t, seen, "the callback should not have run")
|
||||
assert.Contains(t, buf.String(), "do-thing - do-thing short description")
|
||||
assert.Contains(t, buf.String(), "-json")
|
||||
})
|
||||
|
||||
t.Run("Dispatch splits a line the way a shell would", func(t *testing.T) {
|
||||
var args []string
|
||||
r, _, w := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("do-thing", nil, &args))
|
||||
|
||||
require.NoError(t, r.Dispatch(`do-thing "/tmp/a path.pb.gz"`, w))
|
||||
assert.Equal(t, []string{"/tmp/a path.pb.gz"}, args)
|
||||
})
|
||||
|
||||
t.Run("Match returns names by prefix for tab completion", func(t *testing.T) {
|
||||
r, _, _ := newTestRegistry(t)
|
||||
r.RegisterCommand(testCommand("print-cert", nil, nil))
|
||||
r.RegisterCommand(testCommand("print-tunnel", nil, nil))
|
||||
r.RegisterCommand(testCommand("version", nil, nil))
|
||||
|
||||
assert.Equal(t, []string{"print-cert", "print-tunnel"}, r.Match("print-"))
|
||||
})
|
||||
}
|
||||
|
||||
// A clone is what keeps the ssh session's `logout` command from being visible to every other
|
||||
// session, and to nebula ctl.
|
||||
func TestRegistryCloneIsolation(t *testing.T) {
|
||||
parent, _, w := newTestRegistry(t)
|
||||
parent.RegisterCommand(testCommand("shared", nil, nil))
|
||||
|
||||
child := parent.Clone()
|
||||
child.RegisterCommand(testCommand("logout", nil, nil))
|
||||
|
||||
require.NoError(t, child.DispatchArgs([]string{"logout"}, w))
|
||||
|
||||
buf := &bytes.Buffer{}
|
||||
err := parent.DispatchArgs([]string{"logout"}, NewWriter(buf))
|
||||
assert.ErrorIs(t, err, ErrUnknownCommand)
|
||||
|
||||
buf.Reset()
|
||||
require.NoError(t, child.DispatchArgs([]string{"shared"}, NewWriter(buf)))
|
||||
assert.True(t, strings.HasPrefix(buf.String(), "ran shared"))
|
||||
}
|
||||
-137
@@ -1,137 +0,0 @@
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Exit statuses the client reports. They follow shell convention closely enough that a
|
||||
// script can tell "you asked for something that does not exist" from "it ran and failed".
|
||||
const (
|
||||
// StatusOK means the command ran. Note that commands report their own user facing
|
||||
// problems as prose and still exit 0, matching the ssh console.
|
||||
StatusOK = 0
|
||||
// StatusError means the command could not be completed.
|
||||
StatusError = 1
|
||||
// StatusUsage means the arguments were not valid for that command.
|
||||
StatusUsage = 2
|
||||
// StatusUnknownCommand means there is no such command.
|
||||
StatusUnknownCommand = 127
|
||||
)
|
||||
|
||||
// requestTimeout bounds how long a connected client may take to send its request line. There
|
||||
// is deliberately no timeout on the response: `reload` runs every reload callback inline
|
||||
// before it returns, and a slow one is not a reason to hang up on the operator.
|
||||
const requestTimeout = 5 * time.Second
|
||||
|
||||
// Server serves a Registry over a stream listener. It knows nothing about unix sockets, so
|
||||
// tests can drive it over a net.Pipe.
|
||||
type Server struct {
|
||||
l *slog.Logger
|
||||
reg *Registry
|
||||
}
|
||||
|
||||
func NewServer(l *slog.Logger, reg *Registry) *Server {
|
||||
return &Server{l: l, reg: reg}
|
||||
}
|
||||
|
||||
// Serve accepts connections until ln is closed. Cancelling ctx closes ln, which is what ends
|
||||
// the accept loop; a listener closed underneath us is a normal shutdown, not an error.
|
||||
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
if err := ln.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||
s.l.Warn("Failed to close the ctl listener", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
go s.ServeConn(ctx, conn)
|
||||
}
|
||||
}
|
||||
|
||||
// ServeConn handles one request and closes c.
|
||||
func (s *Server) ServeConn(ctx context.Context, c net.Conn) {
|
||||
defer func() {
|
||||
if err := c.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||
s.l.Debug("Failed to close a ctl connection", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := c.SetReadDeadline(time.Now().Add(requestTimeout)); err != nil {
|
||||
s.l.Debug("Failed to set a ctl read deadline", "error", err)
|
||||
}
|
||||
|
||||
req, err := readRequest(bufio.NewReaderSize(c, maxRequest))
|
||||
if err != nil {
|
||||
s.l.Debug("Rejected a ctl request", "error", err)
|
||||
// Best effort: the client may already be gone, and there is nowhere else to report it.
|
||||
_ = writeEnd(c, StatusError, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// The request is in hand, so the command owns the rest of the connection's lifetime.
|
||||
if err := c.SetReadDeadline(time.Time{}); err != nil {
|
||||
s.l.Debug("Failed to clear the ctl read deadline", "error", err)
|
||||
}
|
||||
|
||||
s.l.Debug("Running a ctl command", "args", req.Args)
|
||||
|
||||
buf := bufio.NewWriterSize(&frameWriter{w: c}, outputBuffer)
|
||||
dispatchErr := s.reg.DispatchArgs(req.Args, NewWriter(buf))
|
||||
|
||||
if err := buf.Flush(); err != nil {
|
||||
s.l.Debug("Failed to flush ctl output", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
status, msg := statusFor(dispatchErr)
|
||||
if err := writeEnd(c, status, msg); err != nil {
|
||||
s.l.Debug("Failed to write the ctl end frame", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// StatusFor maps a dispatch error onto an exit status, for a transport that has somewhere to
|
||||
// put one.
|
||||
func StatusFor(err error) int {
|
||||
status, _ := statusFor(err)
|
||||
return status
|
||||
}
|
||||
|
||||
// statusFor maps a dispatch error onto an exit status and, when the failure is ours to
|
||||
// explain rather than one the command already wrote as prose, a message to go with it.
|
||||
func statusFor(err error) (int, string) {
|
||||
switch {
|
||||
case err == nil:
|
||||
return StatusOK, ""
|
||||
case errors.Is(err, ErrUnknownCommand):
|
||||
return StatusUnknownCommand, ""
|
||||
case errors.Is(err, ErrUsage):
|
||||
return StatusUsage, ""
|
||||
default:
|
||||
return StatusError, fmt.Sprintf("%s", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ErrNotSupported means this platform has no ctl transport. Windows is waiting on a named
|
||||
// pipe implementation; mobile has no daemon for a CLI to attach to in the first place.
|
||||
var ErrNotSupported = errors.New("nebula ctl is not supported on this platform")
|
||||
|
||||
// Listen creates the ctl listener at path. It is the platform boundary: everything above it
|
||||
// in this package is portable.
|
||||
func Listen(path string) (net.Listener, error) {
|
||||
return listenSocket(path)
|
||||
}
|
||||
@@ -1,273 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package diag
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testSocketPath returns a short socket path. t.TempDir on darwin lives under
|
||||
// /var/folders/... and readily exceeds the 104 byte sun_path limit, which fails as a bare
|
||||
// "invalid argument" a long way from the cause.
|
||||
func testSocketPath(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
dir, err := os.MkdirTemp("/tmp", "nebctl")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||
|
||||
path := filepath.Join(dir, "sub", "ctl.sock")
|
||||
require.LessOrEqual(t, len(path), maxSocketPath, "test socket path is too long for sun_path")
|
||||
return path
|
||||
}
|
||||
|
||||
func newTestServer(t *testing.T) (*Registry, string) {
|
||||
t.Helper()
|
||||
|
||||
reg := NewRegistry()
|
||||
path := testSocketPath(t)
|
||||
|
||||
ln, err := Listen(path)
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
srv := NewServer(slog.New(slog.DiscardHandler), reg)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
assert.NoError(t, srv.Serve(ctx, ln))
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
wg.Wait()
|
||||
})
|
||||
|
||||
return reg, path
|
||||
}
|
||||
|
||||
func run(t *testing.T, path string, args ...string) (string, int, error) {
|
||||
t.Helper()
|
||||
|
||||
c, err := Dial(path)
|
||||
require.NoError(t, err)
|
||||
defer c.Close()
|
||||
|
||||
out := &bytes.Buffer{}
|
||||
status, err := c.Run(args, out)
|
||||
return out.String(), status, err
|
||||
}
|
||||
|
||||
func TestServeConn(t *testing.T) {
|
||||
t.Run("a command runs and its output comes back", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
reg.RegisterCommand(&Command{
|
||||
Name: "version",
|
||||
ShortDescription: "prints a version",
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
return w.WriteLine("1.2.3")
|
||||
},
|
||||
})
|
||||
|
||||
out, status, err := run(t, path, "version")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, "1.2.3\n", out)
|
||||
})
|
||||
|
||||
t.Run("no args gets the command list", func(t *testing.T) {
|
||||
_, path := newTestServer(t)
|
||||
|
||||
out, status, err := run(t, path)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Contains(t, out, "Available commands:")
|
||||
})
|
||||
|
||||
t.Run("an unknown command exits 127", func(t *testing.T) {
|
||||
_, path := newTestServer(t)
|
||||
|
||||
out, status, err := run(t, path, "nope")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusUnknownCommand, status)
|
||||
assert.Contains(t, out, "Did not understand: nope")
|
||||
})
|
||||
|
||||
t.Run("a bad flag exits 2", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
var seen any
|
||||
reg.RegisterCommand(testCommand("do-thing", &seen, nil))
|
||||
|
||||
out, status, err := run(t, path, "do-thing", "-nope")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusUsage, status)
|
||||
assert.Contains(t, out, "flag provided but not defined")
|
||||
})
|
||||
|
||||
t.Run("a callback error exits 1 and reports why", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
reg.RegisterCommand(&Command{
|
||||
Name: "explode",
|
||||
ShortDescription: "fails",
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
return errors.New("boom")
|
||||
},
|
||||
})
|
||||
|
||||
_, status, err := run(t, path, "explode")
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, StatusError, status)
|
||||
assert.Contains(t, err.Error(), "boom")
|
||||
})
|
||||
|
||||
t.Run("output larger than the buffer arrives intact", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
want := bytes.Repeat([]byte("x"), outputBuffer*3+7)
|
||||
reg.RegisterCommand(&Command{
|
||||
Name: "big",
|
||||
ShortDescription: "writes a lot",
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
return w.WriteBytes(want)
|
||||
},
|
||||
})
|
||||
|
||||
out, status, err := run(t, path, "big")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, string(want), out)
|
||||
})
|
||||
|
||||
t.Run("concurrent clients are all served", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
reg.RegisterCommand(&Command{
|
||||
Name: "slow",
|
||||
ShortDescription: "takes a moment",
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
return w.WriteLine("done")
|
||||
},
|
||||
})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 8; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
out, status, err := run(t, path, "slow")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, "done\n", out)
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
|
||||
t.Run("a client that hangs up mid command does not take the server down", func(t *testing.T) {
|
||||
reg, path := newTestServer(t)
|
||||
reg.RegisterCommand(&Command{
|
||||
Name: "version",
|
||||
ShortDescription: "prints a version",
|
||||
Callback: func(fs any, a []string, w StringWriter) error {
|
||||
return w.WriteLine("1.2.3")
|
||||
},
|
||||
})
|
||||
|
||||
c, err := Dial(path)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writeRequest(c.conn, []string{"version"}))
|
||||
require.NoError(t, c.Close())
|
||||
|
||||
// The next client still gets served.
|
||||
out, status, err := run(t, path, "version")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, StatusOK, status)
|
||||
assert.Equal(t, "1.2.3\n", out)
|
||||
})
|
||||
}
|
||||
|
||||
func TestListenSocket(t *testing.T) {
|
||||
t.Run("the socket is 0600 inside a 0700 directory", func(t *testing.T) {
|
||||
path := testSocketPath(t)
|
||||
ln, err := Listen(path)
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
fi, err := os.Stat(path)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, os.FileMode(0600), fi.Mode().Perm(), "socket mode")
|
||||
|
||||
di, err := os.Stat(filepath.Dir(path))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, os.FileMode(0700), di.Mode().Perm(), "socket directory mode")
|
||||
})
|
||||
|
||||
t.Run("the socket is unlinked when the listener closes", func(t *testing.T) {
|
||||
path := testSocketPath(t)
|
||||
ln, err := Listen(path)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, ln.Close())
|
||||
|
||||
_, err = os.Stat(path)
|
||||
assert.ErrorIs(t, err, fs.ErrNotExist)
|
||||
})
|
||||
|
||||
// A crashed nebula leaves its socket behind, and the next one has to be able to start.
|
||||
t.Run("a socket left behind by a dead nebula is replaced", func(t *testing.T) {
|
||||
path := testSocketPath(t)
|
||||
ln, err := Listen(path)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Close the listener without unlinking, the way a killed process leaves things.
|
||||
unix, ok := ln.(*net.UnixListener)
|
||||
require.True(t, ok)
|
||||
unix.SetUnlinkOnClose(false)
|
||||
require.NoError(t, ln.Close())
|
||||
require.FileExists(t, path)
|
||||
|
||||
ln2, err := Listen(path)
|
||||
require.NoError(t, err)
|
||||
assert.NoError(t, ln2.Close())
|
||||
})
|
||||
|
||||
// Silently stealing it would break the nebula that got there first.
|
||||
t.Run("a socket another nebula is serving is refused", func(t *testing.T) {
|
||||
_, path := newTestServer(t)
|
||||
|
||||
_, err := Listen(path)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already being served")
|
||||
})
|
||||
|
||||
t.Run("a path that is not a socket is refused rather than removed", func(t *testing.T) {
|
||||
path := testSocketPath(t)
|
||||
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0700))
|
||||
require.NoError(t, os.WriteFile(path, []byte("precious"), 0600))
|
||||
|
||||
_, err := Listen(path)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "is not a socket")
|
||||
assert.FileExists(t, path, "the file must not have been removed")
|
||||
})
|
||||
|
||||
t.Run("a path too long for sun_path says so", func(t *testing.T) {
|
||||
_, err := Listen("/tmp/" + fmt.Sprintf("%0*d", maxSocketPath, 0) + "/ctl.sock")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "the maximum is")
|
||||
})
|
||||
}
|
||||
@@ -1,107 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package diag
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"time"
|
||||
)
|
||||
|
||||
// maxSocketPath is the smallest sun_path across the platforms nebula ships on: 104 bytes on
|
||||
// darwin and the BSDs, 108 on Linux. Checking it ourselves turns a bare "invalid argument"
|
||||
// into something an operator can act on.
|
||||
const maxSocketPath = 103
|
||||
|
||||
// DefaultSocketPath is where nebula listens when ctl.socket is unset. An empty string means
|
||||
// the platform has no sensible default and ctl stays off unless an operator names a path.
|
||||
func DefaultSocketPath() string {
|
||||
switch runtime.GOOS {
|
||||
case "ios", "android":
|
||||
// No daemon to attach to and no shell to attach from, and nowhere writable that
|
||||
// would survive being guessed. Mobile embedders drive nebula through Control.
|
||||
return ""
|
||||
case "linux":
|
||||
return "/run/nebula/ctl.sock"
|
||||
default:
|
||||
// /run does not exist on darwin, and /var/run is the portable spelling everywhere
|
||||
// else nebula builds.
|
||||
return "/var/run/nebula/ctl.sock"
|
||||
}
|
||||
}
|
||||
|
||||
// listenSocket creates the listening socket at path, taking over one a previous nebula left
|
||||
// behind but refusing one that is still being served.
|
||||
func listenSocket(path string) (net.Listener, error) {
|
||||
if len(path) > maxSocketPath {
|
||||
return nil, fmt.Errorf("socket path is %d bytes, the maximum is %d", len(path), maxSocketPath)
|
||||
}
|
||||
|
||||
// The directory, not the socket, is what enforces access control. net.Listen creates the
|
||||
// socket with 0777&^umask, so with a typical 0022 umask it is world connectable for the
|
||||
// window between bind and chmod. Nobody can traverse into a 0700 directory to reach it in
|
||||
// that window, and unlike the socket's own mode, directory traversal is enforced
|
||||
// consistently across every platform this file builds for.
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0700); err != nil {
|
||||
return nil, fmt.Errorf("failed to create %s: %w", dir, err)
|
||||
}
|
||||
if err := os.Chmod(dir, 0700); err != nil {
|
||||
return nil, fmt.Errorf("failed to set permissions on %s: %w", dir, err)
|
||||
}
|
||||
|
||||
if err := clearStaleSocket(path); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ln, err := net.Listen("unix", path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Defence in depth behind the directory, for anyone who relocates the socket somewhere
|
||||
// more permissive.
|
||||
if err := os.Chmod(path, 0600); err != nil {
|
||||
_ = ln.Close()
|
||||
return nil, fmt.Errorf("failed to set permissions on %s: %w", path, err)
|
||||
}
|
||||
|
||||
return ln, nil
|
||||
}
|
||||
|
||||
// dialSocket connects to a nebula serving at path.
|
||||
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
|
||||
return net.DialTimeout("unix", path, timeout)
|
||||
}
|
||||
|
||||
// clearStaleSocket removes a socket a crashed nebula left behind, but refuses to steal one
|
||||
// another nebula is still serving. Two instances on one host need two paths; they cannot
|
||||
// share one, and silently taking the socket would break the instance that got there first.
|
||||
func clearStaleSocket(path string) error {
|
||||
fi, err := os.Lstat(path)
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if fi.Mode()&fs.ModeSocket == 0 {
|
||||
return fmt.Errorf("%s exists and is not a socket, refusing to remove it", path)
|
||||
}
|
||||
|
||||
// A successful dial is the only reliable way to tell a live socket from an abandoned
|
||||
// one; the inode looks identical either way.
|
||||
c, err := net.DialTimeout("unix", path, 100*time.Millisecond)
|
||||
if err == nil {
|
||||
_ = c.Close()
|
||||
return fmt.Errorf("%s is already being served, is another nebula running?", path)
|
||||
}
|
||||
|
||||
return os.Remove(path)
|
||||
}
|
||||
@@ -1,27 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package diag
|
||||
|
||||
import (
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Windows has AF_UNIX since Windows 10 1803, but no way to secure the socket that resembles
|
||||
// what the unix build does: os.Chmod cannot express an ACL, and a socket's reachability comes
|
||||
// down to whatever its directory inherited. Doing this properly means a named pipe with an
|
||||
// explicit security descriptor, which is a dependency and a design this change does not carry.
|
||||
// Until then the stub keeps the package building and gives operators a real answer.
|
||||
|
||||
// DefaultSocketPath returns an empty string: there is no path worth defaulting to here.
|
||||
func DefaultSocketPath() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func listenSocket(path string) (net.Listener, error) {
|
||||
return nil, ErrNotSupported
|
||||
}
|
||||
|
||||
func dialSocket(path string, timeout time.Duration) (net.Conn, error) {
|
||||
return nil, ErrNotSupported
|
||||
}
|
||||
@@ -116,9 +116,6 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific
|
||||
"key": string(myPrivKey),
|
||||
},
|
||||
//"tun": m{"disabled": true},
|
||||
// Several tests bring up more than one nebula in this process, and they would all
|
||||
// contend for the same default ctl socket path. None of them exercise it.
|
||||
"ctl": m{"enabled": false},
|
||||
"firewall": m{
|
||||
"outbound": []m{{
|
||||
"proto": "any",
|
||||
@@ -216,9 +213,6 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o
|
||||
"key": string(key),
|
||||
},
|
||||
//"tun": m{"disabled": true},
|
||||
// Several tests bring up more than one nebula in this process, and they would all
|
||||
// contend for the same default ctl socket path. None of them exercise it.
|
||||
"ctl": m{"enabled": false},
|
||||
"firewall": m{
|
||||
"outbound": []m{{
|
||||
"proto": "any",
|
||||
|
||||
@@ -223,60 +223,3 @@ func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
// A relayed send records traffic but must not consume the rebind epoch. If it does, the next direct send to the
|
||||
// relay host sees the epoch already current and never requeries, so the far side is never told to punch at our
|
||||
// new address. This pins the SendVia call site, which the unit tests cannot reach.
|
||||
func TestRebindRequeriesAfterRelayedSend(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{})
|
||||
|
||||
// No lighthouse on purpose: it would hand out a direct address for them and nothing would relay.
|
||||
// Long connection manager timers so it never fires a direct test packet at the relay tunnel and bumps its
|
||||
// epoch mid-test, which is the only other thing that touches that tunnel and would flake the assertion below.
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24",
|
||||
m{"relay": m{"use_relays": true}, "timers": m{"connection_alive_interval": 3600, "pending_deletion_interval": 3600}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
|
||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
myControl.Start()
|
||||
relayControl.Start()
|
||||
theirControl.Start()
|
||||
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
r.RouteFor(time.Millisecond * 500)
|
||||
|
||||
hi := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
||||
require.NotNil(t, hi, "expected a tunnel to them")
|
||||
require.NotEmpty(t, hi.CurrentRelaysToMe, "them must be reachable only via the relay for this test to mean anything")
|
||||
// sendNoMetrics only reaches SendVia when there is no direct remote, so pin that too. Without this the test
|
||||
// keeps passing while quietly sending direct and never exercising the relay path.
|
||||
require.False(t, hi.CurrentRemote.IsValid(), "them must have no direct remote, otherwise SendVia is never called")
|
||||
|
||||
before, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
|
||||
require.True(t, ok, "expected a tunnel to the relay")
|
||||
|
||||
myControl.RebindUDPServer()
|
||||
|
||||
// Traffic to them goes through SendVia on the relay tunnel. That must record traffic without consuming the
|
||||
// relay tunnel's own epoch edge, which belongs to the direct path.
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("relayed")))
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
|
||||
after, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, before, after,
|
||||
"a relayed send consumed the relay tunnel's rebind epoch, so the next direct send will not requery")
|
||||
|
||||
myControl.Stop()
|
||||
relayControl.Stop()
|
||||
theirControl.Stop()
|
||||
}
|
||||
|
||||
+2
-26
@@ -173,7 +173,7 @@ listen:
|
||||
# This setting is reloadable.
|
||||
#so_mark: 0
|
||||
# the udp_offloads setting controls if Nebula will attempt to enable GSO and GRO for its UDP socket(s). Linux only, not reloadable.
|
||||
# udp_offloads: false
|
||||
# udp_offloads: true
|
||||
|
||||
# Routines is the number of thread pairs to run that consume from the tun and UDP queues.
|
||||
# Currently, this defaults to 1 which means we have 1 tun queue reader and 1
|
||||
@@ -236,30 +236,6 @@ punchy:
|
||||
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
|
||||
#sandbox_dir: /var/tmp/nebula-debug
|
||||
|
||||
# ctl exposes nebula's debug and administrative commands over a local unix socket, so that `nebula ctl <command>` can
|
||||
# reach the same commands the sshd block offers above without running an ssh server. Run `nebula ctl` on its own for the
|
||||
# list of commands. Anyone who can open the socket can do everything the ssh console can, including closing tunnels,
|
||||
# changing remotes, and writing profile data to disk, so the socket lives in a directory only the user nebula runs as
|
||||
# can enter. Enabled by default. Not supported on Windows yet, and never enabled on iOS or Android.
|
||||
#ctl:
|
||||
# Toggles the feature. This setting is reloadable.
|
||||
#enabled: true
|
||||
|
||||
# socket is the unix socket to listen on. The parent directory is created if it is missing and made readable only by
|
||||
# the user nebula runs as, and a socket left behind by a crashed nebula is replaced. Defaults to /run/nebula/ctl.sock
|
||||
# on Linux and /var/run/nebula/ctl.sock everywhere else; running nebula as a non-root user means picking a path it can
|
||||
# write. Two nebulas on one host need two paths, the second to start will log that the socket is already being served
|
||||
# and carry on without one. `nebula ctl` reads this value from the same config file when it is given -config, and
|
||||
# otherwise assumes the default above. This setting is reloadable.
|
||||
#socket: /run/nebula/ctl.sock
|
||||
|
||||
# sandbox_dir restricts the file paths the profiling commands (start-cpu-profile, save-heap-profile,
|
||||
# save-mutex-profile) may write, exactly like sshd.sandbox_dir above, which it defaults to. Note that these paths are
|
||||
# resolved by the nebula process and not by the shell running `nebula ctl`, so a relative path lands in this directory
|
||||
# rather than in your working directory, and under a systemd unit with PrivateTmp=yes it lands somewhere your shell
|
||||
# cannot see at all. The directory is NOT automatically created.
|
||||
#sandbox_dir: /var/tmp/nebula-debug
|
||||
|
||||
# EXPERIMENTAL: relay support for networks that can't establish direct connections.
|
||||
relay:
|
||||
# Relays are a list of Nebula IP's that peers can use to relay packets to me.
|
||||
@@ -292,7 +268,7 @@ tun:
|
||||
mtu: 1300
|
||||
|
||||
# the use_offloads setting controls if Nebula will attempt to enable GSO and GRO for the tun device. Linux only, not reloadable.
|
||||
#use_offloads: false
|
||||
#use_offloads: true
|
||||
|
||||
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
|
||||
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
|
||||
|
||||
+14
-15
@@ -21,7 +21,6 @@ import (
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
)
|
||||
|
||||
type FirewallInterface interface {
|
||||
@@ -263,11 +262,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
||||
}
|
||||
|
||||
switch proto {
|
||||
case iputil.IPProtocolTCP:
|
||||
case firewall.ProtoTCP:
|
||||
fp = ft.TCP
|
||||
case iputil.IPProtocolUDP:
|
||||
case firewall.ProtoUDP:
|
||||
fp = ft.UDP
|
||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
||||
if startPort != firewall.PortAny {
|
||||
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
||||
@@ -365,13 +364,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
|
||||
proto = firewall.ProtoAny
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "tcp":
|
||||
proto = iputil.IPProtocolTCP
|
||||
proto = firewall.ProtoTCP
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "udp":
|
||||
proto = iputil.IPProtocolUDP
|
||||
proto = firewall.ProtoUDP
|
||||
startPort, endPort, err = parsePort(sPort)
|
||||
case "icmp":
|
||||
proto = iputil.IPProtocolICMP
|
||||
proto = firewall.ProtoICMP
|
||||
startPort = firewall.PortAny
|
||||
endPort = firewall.PortAny
|
||||
if sPort != "" {
|
||||
@@ -561,9 +560,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
}
|
||||
|
||||
switch fp.Protocol {
|
||||
case iputil.IPProtocolTCP:
|
||||
case firewall.ProtoTCP:
|
||||
c.Expires = time.Now().Add(f.TCPTimeout)
|
||||
case iputil.IPProtocolUDP:
|
||||
case firewall.ProtoUDP:
|
||||
c.Expires = time.Now().Add(f.UDPTimeout)
|
||||
default:
|
||||
c.Expires = time.Now().Add(f.DefaultTimeout)
|
||||
@@ -583,9 +582,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
||||
c := &conn{}
|
||||
|
||||
switch fp.Protocol {
|
||||
case iputil.IPProtocolTCP:
|
||||
case firewall.ProtoTCP:
|
||||
timeout = f.TCPTimeout
|
||||
case iputil.IPProtocolUDP:
|
||||
case firewall.ProtoUDP:
|
||||
timeout = f.UDPTimeout
|
||||
default:
|
||||
timeout = f.DefaultTimeout
|
||||
@@ -636,15 +635,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
|
||||
}
|
||||
|
||||
switch p.Protocol {
|
||||
case iputil.IPProtocolTCP:
|
||||
case firewall.ProtoTCP:
|
||||
if ft.TCP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
case iputil.IPProtocolUDP:
|
||||
case firewall.ProtoUDP:
|
||||
if ft.UDP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
|
||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||
if ft.ICMP.match(p, incoming, c, caPool) {
|
||||
return true
|
||||
}
|
||||
@@ -681,7 +680,7 @@ func (fp firewallPort) match(p firewall.Packet, incoming bool, c *cert.CachedCer
|
||||
}
|
||||
|
||||
// this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match
|
||||
if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
|
||||
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 {
|
||||
// port numbers are re-used for connection tracking of ICMP,
|
||||
// but we don't want to actually filter on them.
|
||||
return fp[firewall.PortAny].match(p, c, caPool)
|
||||
|
||||
+10
-7
@@ -4,14 +4,17 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
)
|
||||
|
||||
type m = map[string]any
|
||||
|
||||
const (
|
||||
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
||||
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
|
||||
ProtoTCP = 6
|
||||
ProtoUDP = 17
|
||||
ProtoICMP = 1
|
||||
ProtoICMPv6 = 58
|
||||
|
||||
PortAny = 0 // Special value for matching `port: any`
|
||||
PortFragment = -1 // Special value for matching `port: fragment`
|
||||
)
|
||||
@@ -42,13 +45,13 @@ func (fp *Packet) Copy() *Packet {
|
||||
func (fp Packet) MarshalJSON() ([]byte, error) {
|
||||
var proto string
|
||||
switch fp.Protocol {
|
||||
case iputil.IPProtocolTCP:
|
||||
case ProtoTCP:
|
||||
proto = "tcp"
|
||||
case iputil.IPProtocolICMP:
|
||||
case ProtoICMP:
|
||||
proto = "icmp"
|
||||
case iputil.IPProtocolICMPv6:
|
||||
case ProtoICMPv6:
|
||||
proto = "icmpv6"
|
||||
case iputil.IPProtocolUDP:
|
||||
case ProtoUDP:
|
||||
proto = "udp"
|
||||
default:
|
||||
proto = fmt.Sprintf("unknown %v", fp.Protocol)
|
||||
|
||||
+33
-34
@@ -13,7 +13,6 @@ import (
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -73,20 +72,20 @@ func TestFirewall_AddRule(t *testing.T) {
|
||||
ti6, err := netip.ParsePrefix("fd12::34/128")
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", ""))
|
||||
// An empty rule is any
|
||||
assert.True(t, fw.InRules.TCP[1].Any.Any.Any)
|
||||
assert.Empty(t, fw.InRules.TCP[1].Any.Groups)
|
||||
assert.Empty(t, fw.InRules.TCP[1].Any.Hosts)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
|
||||
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
|
||||
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
|
||||
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", ""))
|
||||
//no matter what port is given for icmp, it should end up as "any"
|
||||
assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any)
|
||||
assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
|
||||
@@ -117,11 +116,11 @@ func TestFirewall_AddRule(t *testing.T) {
|
||||
assert.True(t, ok)
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
|
||||
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
|
||||
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
|
||||
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
@@ -186,7 +185,7 @@ func TestFirewall_Drop(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -264,7 +263,7 @@ func TestFirewall_DropV6(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -351,7 +350,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
Certificate: &dummyCert{},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -361,7 +360,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
Certificate: &dummyCert{},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -371,7 +370,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
}
|
||||
ip := netip.MustParsePrefix("9.254.254.254/32")
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
|
||||
@@ -380,7 +379,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
}
|
||||
ip := netip.MustParsePrefix("fd99::99/128")
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -393,7 +392,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
||||
@@ -405,7 +404,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -418,7 +417,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
|
||||
@@ -430,7 +429,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -442,7 +441,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -454,7 +453,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
|
||||
@@ -465,7 +464,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"good-group": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
|
||||
}
|
||||
})
|
||||
|
||||
@@ -477,7 +476,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
InvertedGroups: map[string]struct{}{"nope": {}},
|
||||
}
|
||||
for n := 0; n < b.N; n++ {
|
||||
ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
|
||||
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -493,7 +492,7 @@ func TestFirewall_Drop2(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -551,7 +550,7 @@ func TestFirewall_Drop3(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -639,7 +638,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("fd12::34"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
@@ -676,7 +675,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
network := netip.MustParsePrefix("1.2.3.4/24")
|
||||
@@ -759,13 +758,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
templ := firewall.Packet{
|
||||
LocalAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
|
||||
Protocol: iputil.IPProtocolICMP,
|
||||
Protocol: firewall.ProtoICMP,
|
||||
Fragment: false,
|
||||
}
|
||||
|
||||
t.Run("ICMP allowed", func(t *testing.T) {
|
||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
|
||||
t.Run("zero ports", func(t *testing.T) {
|
||||
p := templ.Copy()
|
||||
p.LocalPort = 0
|
||||
@@ -911,7 +910,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
|
||||
LocalPort: 1,
|
||||
RemotePort: 1,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||
@@ -962,7 +961,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||
LocalPort: 443,
|
||||
RemotePort: 55000,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
}
|
||||
|
||||
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
|
||||
@@ -1032,7 +1031,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
|
||||
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
|
||||
LocalPort: 443,
|
||||
RemotePort: 55000,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
@@ -1318,28 +1317,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
mf := &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding udp rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule no port
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding any rule
|
||||
conf = config.NewC(test.NewLogger())
|
||||
@@ -1583,7 +1582,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
|
||||
RemoteAddr: theirPrefixes[0].Addr(),
|
||||
LocalPort: 10,
|
||||
RemotePort: 90,
|
||||
Protocol: iputil.IPProtocolUDP,
|
||||
Protocol: firewall.ProtoUDP,
|
||||
Fragment: false,
|
||||
}
|
||||
return testcase{
|
||||
|
||||
+13
-67
@@ -239,15 +239,11 @@ const (
|
||||
|
||||
type HostInfo struct {
|
||||
remote atomic.Pointer[netip.AddrPort]
|
||||
remotes *RemoteList
|
||||
promoteCounter atomic.Uint32
|
||||
ConnectionState *ConnectionState
|
||||
|
||||
// Traffic bits, pendingDeletion, and the rebind epoch we last sent under
|
||||
state atomic.Uint32
|
||||
|
||||
promoteCounter atomic.Uint32
|
||||
remoteIndexId uint32
|
||||
localIndexId uint32
|
||||
remotes *RemoteList
|
||||
remoteIndexId uint32
|
||||
localIndexId uint32
|
||||
|
||||
// vpnAddrs is a list of vpn addresses assigned to this host that are within our own vpn networks
|
||||
// The host may have other vpn addresses that are outside our
|
||||
@@ -266,6 +262,11 @@ type HostInfo struct {
|
||||
// This is used to limit lighthouse re-queries in chatty clients
|
||||
nextLHQuery atomic.Int64
|
||||
|
||||
// lastRebindCount is the other side of Interface.rebindCount, if these values don't match then we need to ask LH
|
||||
// for a punch from the remote end of this tunnel. The goal being to prime their conntrack for our traffic just like
|
||||
// with a handshake
|
||||
lastRebindCount int8
|
||||
|
||||
// lastHandshakeTime records the time the remote side told us about at the stage when the handshake was completed locally
|
||||
// Stage 1 packet will contain it if I am a responder, stage 2 packet if I am an initiator
|
||||
// This is used to avoid an attack where a handshake packet is replayed after some time
|
||||
@@ -274,6 +275,9 @@ type HostInfo struct {
|
||||
lastRoam time.Time
|
||||
lastRoamRemote netip.AddrPort
|
||||
|
||||
//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
|
||||
|
||||
// lastUsed tracks the last time ConnectionManager checked the tunnel and it was in use.
|
||||
// This value will be behind against actual tunnel utilization in the hot path.
|
||||
// This should only be used by the ConnectionManagers ticker routine.
|
||||
@@ -665,7 +669,7 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||
|
||||
hostinfo.markOut(f.rebindEpoch.Load())
|
||||
hostinfo.out.Store(true)
|
||||
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
|
||||
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
|
||||
}
|
||||
@@ -766,64 +770,6 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
|
||||
}
|
||||
}
|
||||
|
||||
// Bits within HostInfo.state, everything above stateEpochShift is the epoch
|
||||
const (
|
||||
stateIn uint32 = 1 << iota
|
||||
stateOut
|
||||
statePendingDeletion
|
||||
|
||||
stateFlags = stateIn | stateOut | statePendingDeletion
|
||||
// The epoch is the top 29 bits, it would take 2^29 rebinds to wrap and we will never get there
|
||||
stateEpochShift = 3
|
||||
)
|
||||
|
||||
// markIn records inbound traffic
|
||||
func (i *HostInfo) markIn() {
|
||||
if i.state.Load()&stateIn == 0 {
|
||||
i.state.Or(stateIn)
|
||||
}
|
||||
}
|
||||
|
||||
// markOut records a send and reports whether the epoch moved, meaning we want a punch from the far side
|
||||
func (i *HostInfo) markOut(epoch uint32) bool {
|
||||
e := epoch << stateEpochShift
|
||||
for {
|
||||
old := i.state.Load()
|
||||
if old&stateOut != 0 && old&^stateFlags == e {
|
||||
return false
|
||||
}
|
||||
|
||||
if i.state.CompareAndSwap(old, old&stateFlags|stateOut|e) {
|
||||
return old&^stateFlags != e
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// markOutOnly records a send without consuming the rebind epoch, for paths that cannot act on a requery
|
||||
func (i *HostInfo) markOutOnly() {
|
||||
if i.state.Load()&stateOut == 0 {
|
||||
i.state.Or(stateOut)
|
||||
}
|
||||
}
|
||||
|
||||
// takeTraffic clears both traffic bits, leaving the epoch alone, and reports what they were
|
||||
func (i *HostInfo) takeTraffic() (in bool, out bool) {
|
||||
old := i.state.And(^(stateIn | stateOut))
|
||||
return old&stateIn != 0, old&stateOut != 0
|
||||
}
|
||||
|
||||
func (i *HostInfo) setPendingDeletion(v bool) {
|
||||
if v {
|
||||
i.state.Or(statePendingDeletion)
|
||||
} else {
|
||||
i.state.And(^statePendingDeletion)
|
||||
}
|
||||
}
|
||||
|
||||
func (i *HostInfo) isPendingDeletion() bool {
|
||||
return i.state.Load()&statePendingDeletion != 0
|
||||
}
|
||||
|
||||
func (i *HostInfo) GetCert() *cert.CachedCertificate {
|
||||
if i.ConnectionState != nil {
|
||||
return i.ConnectionState.peerCert
|
||||
|
||||
@@ -401,49 +401,3 @@ func TestHostMap_RelayState(t *testing.T) {
|
||||
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
|
||||
|
||||
}
|
||||
|
||||
// sentSinceCheck reports whether anything has been sent since the connection manager last looked. Test only:
|
||||
// production reads the out bit through takeTraffic on the connection manager tick.
|
||||
func (i *HostInfo) sentSinceCheck() bool {
|
||||
return i.state.Load()&stateOut != 0
|
||||
}
|
||||
|
||||
func TestHostInfo_markOut(t *testing.T) {
|
||||
h := &HostInfo{}
|
||||
h.markOut(5) // stamped when the tunnel was added
|
||||
|
||||
// A tunnel already on the current epoch has nothing to report, which is what keeps a fresh tunnel from
|
||||
// requerying on its first packet
|
||||
assert.False(t, h.markOut(5), "an unchanged epoch should not report a move")
|
||||
assert.True(t, h.sentSinceCheck(), "the send is still recorded as traffic")
|
||||
|
||||
// A rebind is observed exactly once, so we requery once per rebind
|
||||
assert.True(t, h.markOut(6), "a bumped epoch should report a move")
|
||||
assert.False(t, h.markOut(6), "the epoch move should only be reported once")
|
||||
|
||||
// Traffic and pendingDeletion live in the same word and must survive an epoch change
|
||||
h.setPendingDeletion(true)
|
||||
h.markIn()
|
||||
assert.True(t, h.markOut(7))
|
||||
assert.True(t, h.isPendingDeletion(), "pendingDeletion must survive an epoch change")
|
||||
in, out := h.takeTraffic()
|
||||
assert.True(t, in, "inbound traffic must survive an epoch change")
|
||||
assert.True(t, out)
|
||||
|
||||
// Clearing the traffic bits leaves the epoch alone, otherwise an idle tunnel would requery forever
|
||||
assert.False(t, h.markOut(7), "takeTraffic must not disturb the epoch")
|
||||
}
|
||||
|
||||
// A relayed send records traffic but must leave the rebind epoch for the direct path to consume, otherwise
|
||||
// relaying to a host swallows the requery that gets the far side punching at our new address.
|
||||
func TestHostInfo_markOutOnly(t *testing.T) {
|
||||
h := &HostInfo{}
|
||||
h.markOut(5)
|
||||
|
||||
h.markOutOnly()
|
||||
assert.True(t, h.sentSinceCheck(), "a relayed send is still outbound traffic")
|
||||
assert.False(t, h.markOut(5), "a relayed send must not disturb the epoch")
|
||||
|
||||
assert.True(t, h.markOut(6), "a relayed send must not consume the epoch edge")
|
||||
assert.False(t, h.markOut(6))
|
||||
}
|
||||
|
||||
@@ -58,9 +58,6 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
|
||||
// kernel as one giant blob; segment first so the loopback
|
||||
// path sees one IP datagram per Write.
|
||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||
// The kernel may have left the transport checksum for hardware
|
||||
// offload to finish; nothing between here and the tun will.
|
||||
iputil.SetTransportChecksum(seg)
|
||||
_, werr := f.queues[q].Write(seg)
|
||||
return werr
|
||||
})
|
||||
@@ -163,18 +160,21 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
||||
// One traffic-out mark covers every segment of the superpacket; doing it
|
||||
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
|
||||
// times per TSO packet, inside writeLock under boring crypto.
|
||||
//
|
||||
// We rebound since this tunnel last sent, ask the lighthouse to get the far side punching at us again
|
||||
if f.connectionManager.Out(hostinfo) {
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
remote := hostinfo.GetRemote()
|
||||
if hostinfo.lastRebindCount != f.rebindCount {
|
||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||
hostinfo.lastRebindCount = f.rebindCount
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind epoch",
|
||||
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
remote := hostinfo.GetRemote()
|
||||
if !remote.IsValid() { //the relay path
|
||||
//first, find our relay hostinfo:
|
||||
var relayHostInfo *HostInfo
|
||||
@@ -460,7 +460,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
}
|
||||
|
||||
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
|
||||
f.connectionManager.OutNoRebind(via)
|
||||
f.connectionManager.Out(via)
|
||||
|
||||
// Authenticate the header and payload, but do not encrypt for this message type.
|
||||
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
|
||||
@@ -553,13 +553,17 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
|
||||
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
|
||||
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
|
||||
// A closing tunnel is torn down right after this, so skip the connection manager entirely: no point recording
|
||||
// traffic or asking the lighthouse for a punch. Otherwise, if we rebound since this tunnel last sent, ask the
|
||||
// lighthouse to get the far side punching at us again.
|
||||
if t != header.CloseTunnel && f.connectionManager.Out(hostinfo) {
|
||||
f.connectionManager.Out(hostinfo)
|
||||
|
||||
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
|
||||
// all our addrs and enable a faster roaming.
|
||||
if t != header.CloseTunnel && hostinfo.lastRebindCount != f.rebindCount {
|
||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||
hostinfo.lastRebindCount = f.rebindCount
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Lighthouse update triggered for punch due to rebind epoch",
|
||||
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
|
||||
-265
@@ -1,265 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
ipv4HeaderLen = 20
|
||||
ipv6HeaderLen = 40
|
||||
)
|
||||
|
||||
// capturingTun is a tio.Queue that records what is written to it. A queue that
|
||||
// discards writes is indistinguishable from a packet that was never forwarded.
|
||||
type capturingTun struct {
|
||||
writes [][]byte
|
||||
}
|
||||
|
||||
func (c *capturingTun) Read() ([]tio.Packet, error) { return nil, io.EOF }
|
||||
func (c *capturingTun) Close() error { return nil }
|
||||
|
||||
func (c *capturingTun) Write(b []byte) (int, error) {
|
||||
c.writes = append(c.writes, append([]byte(nil), b...))
|
||||
return len(b), nil
|
||||
}
|
||||
|
||||
func newSelfForwardInterface(myAddrs ...netip.Addr) (*Interface, *capturingTun) {
|
||||
vpnAddrs := &bart.Lite{}
|
||||
for _, a := range myAddrs {
|
||||
vpnAddrs.Insert(netip.PrefixFrom(a, a.BitLen()))
|
||||
}
|
||||
|
||||
tun := &capturingTun{}
|
||||
return &Interface{
|
||||
l: test.NewLogger(),
|
||||
myVpnAddrsTable: vpnAddrs,
|
||||
myBroadcastAddrsTable: &bart.Lite{},
|
||||
queues: []tio.Queue{tun},
|
||||
}, tun
|
||||
}
|
||||
|
||||
func consumeInside(f *Interface, packet []byte) {
|
||||
f.consumeInsidePacket(tio.Packet{Bytes: packet}, &firewall.ParsedPacket{}, make([]byte, 12), nil, make([]byte, mtu), 0, nil)
|
||||
}
|
||||
|
||||
// l4Proto describes one upper-layer header for these tests: its IP next-header
|
||||
// value, where its checksum field sits within the header, and how to build a
|
||||
// minimal instance of it.
|
||||
type l4Proto struct {
|
||||
name string
|
||||
nextHdr uint8
|
||||
cksumAt int
|
||||
build func() []byte
|
||||
}
|
||||
|
||||
var (
|
||||
tcpSyn = l4Proto{"tcp", iputil.IPProtocolTCP, 16, func() []byte {
|
||||
h := make([]byte, 20)
|
||||
binary.BigEndian.PutUint16(h[0:2], 49152)
|
||||
binary.BigEndian.PutUint16(h[2:4], 443)
|
||||
binary.BigEndian.PutUint32(h[4:8], 0x11223344) // sequence
|
||||
h[12] = 5 << 4 // data offset, no options
|
||||
h[13] = 0x02 // SYN
|
||||
binary.BigEndian.PutUint16(h[14:16], 65535) // window
|
||||
return h
|
||||
}}
|
||||
|
||||
udpDatagram = l4Proto{"udp", iputil.IPProtocolUDP, 6, func() []byte {
|
||||
h := make([]byte, 8+4)
|
||||
binary.BigEndian.PutUint16(h[0:2], 49152)
|
||||
binary.BigEndian.PutUint16(h[2:4], 53)
|
||||
binary.BigEndian.PutUint16(h[4:6], uint16(len(h)))
|
||||
copy(h[8:], "ping")
|
||||
return h
|
||||
}}
|
||||
|
||||
icmpEcho = l4Proto{"icmp", iputil.IPProtocolICMP, 2, func() []byte { return echoRequest(8) }}
|
||||
icmpv6Echo = l4Proto{"icmpv6", iputil.IPProtocolICMPv6, 2, func() []byte { return echoRequest(128) }}
|
||||
)
|
||||
|
||||
// echoRequest builds an echo request body. The type differs between ICMP and
|
||||
// ICMPv6, the rest of the header does not.
|
||||
func echoRequest(typ uint8) []byte {
|
||||
h := make([]byte, 8)
|
||||
h[0] = typ
|
||||
binary.BigEndian.PutUint16(h[4:6], 0xbeef) // identifier
|
||||
binary.BigEndian.PutUint16(h[6:8], 1) // sequence
|
||||
return h
|
||||
}
|
||||
|
||||
func buildIPv6(src, dst netip.Addr, p l4Proto) []byte {
|
||||
l4 := p.build()
|
||||
pkt := make([]byte, ipv6HeaderLen+len(l4))
|
||||
pkt[0] = 0x60
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(l4)))
|
||||
pkt[6] = p.nextHdr
|
||||
pkt[7] = 64
|
||||
copy(pkt[8:24], src.AsSlice())
|
||||
copy(pkt[24:40], dst.AsSlice())
|
||||
copy(pkt[ipv6HeaderLen:], l4)
|
||||
if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
|
||||
sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4)))
|
||||
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
|
||||
}
|
||||
return pkt
|
||||
}
|
||||
|
||||
func buildIPv4(src, dst netip.Addr, p l4Proto) []byte {
|
||||
l4 := p.build()
|
||||
pkt := make([]byte, ipv4HeaderLen+len(l4))
|
||||
pkt[0] = 0x45
|
||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
|
||||
pkt[8] = 64
|
||||
pkt[9] = p.nextHdr
|
||||
copy(pkt[12:16], src.AsSlice())
|
||||
copy(pkt[16:20], dst.AsSlice())
|
||||
copy(pkt[ipv4HeaderLen:], l4)
|
||||
if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
|
||||
sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4)))
|
||||
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
|
||||
}
|
||||
return pkt
|
||||
}
|
||||
|
||||
// ipv6PseudoheaderSum is the RFC 2460 section 8.1 pseudo-header sum: source,
|
||||
// destination, a 32 bit upper-layer packet length and a 32 bit zero-padded next
|
||||
// header. Kept local to the test so these assertions do not check nebula's
|
||||
// checksum code against itself.
|
||||
func ipv6PseudoheaderSum(src, dst netip.Addr, nextHeader, length uint32) uint32 {
|
||||
var csum uint32
|
||||
s, d := src.AsSlice(), dst.AsSlice()
|
||||
for i := 0; i < 16; i += 2 {
|
||||
csum += uint32(s[i])<<8 | uint32(s[i+1])
|
||||
csum += uint32(d[i])<<8 | uint32(d[i+1])
|
||||
}
|
||||
return csum + length + nextHeader
|
||||
}
|
||||
|
||||
func sumBytes(b []byte, csum uint32) uint32 {
|
||||
for i := 0; i+1 < len(b); i += 2 {
|
||||
csum += uint32(b[i])<<8 | uint32(b[i+1])
|
||||
}
|
||||
if len(b)%2 == 1 {
|
||||
csum += uint32(b[len(b)-1]) << 8
|
||||
}
|
||||
return csum
|
||||
}
|
||||
|
||||
func fold(csum uint32) uint16 {
|
||||
for csum > 0xffff {
|
||||
csum = (csum >> 16) + (csum & 0xffff)
|
||||
}
|
||||
return uint16(csum)
|
||||
}
|
||||
|
||||
// l4ChecksumValid6 verifies an IPv6 upper-layer checksum the way a receiver
|
||||
// does: the pseudo-header plus the whole upper-layer segment, checksum field
|
||||
// included, folds to 0xffff. The next header field is the upper-layer protocol
|
||||
// only while there are no extension headers, which is all this file builds.
|
||||
func l4ChecksumValid6(pkt []byte) bool {
|
||||
src, _ := netip.AddrFromSlice(pkt[8:24])
|
||||
dst, _ := netip.AddrFromSlice(pkt[24:40])
|
||||
l4 := pkt[ipv6HeaderLen:]
|
||||
return fold(sumBytes(l4, ipv6PseudoheaderSum(src, dst, uint32(pkt[6]), uint32(len(l4))))) == 0xffff
|
||||
}
|
||||
|
||||
// l4ChecksumValid4 is the IPv4 counterpart: the RFC 793/768 pseudo-header is
|
||||
// source, destination, a zero byte, the protocol and the upper-layer length.
|
||||
func l4ChecksumValid4(pkt []byte) bool {
|
||||
ihl := int(pkt[0]&0x0f) << 2
|
||||
l4 := pkt[ihl:]
|
||||
return fold(sumBytes(l4, sumBytes(pkt[12:20], uint32(pkt[9])+uint32(len(l4))))) == 0xffff
|
||||
}
|
||||
|
||||
// TestConsumeInsidePacketSelfTraffic covers the self-addressed branch of
|
||||
// consumeInsidePacket, taken where immediatelyForwardToSelf is set (see
|
||||
// inside_bsd.go): the packet goes straight back to the tun, ahead of the
|
||||
// firewall and the handshake.
|
||||
func TestConsumeInsidePacketSelfTraffic(t *testing.T) {
|
||||
v4 := netip.MustParseAddr("100.100.1.42")
|
||||
v6 := netip.MustParseAddr("fd00::42")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
addr netip.Addr
|
||||
pkt []byte
|
||||
}{
|
||||
{"ipv4/tcp", v4, buildIPv4(v4, v4, tcpSyn)},
|
||||
{"ipv4/udp", v4, buildIPv4(v4, v4, udpDatagram)},
|
||||
{"ipv4/icmp", v4, buildIPv4(v4, v4, icmpEcho)},
|
||||
{"ipv6/tcp", v6, buildIPv6(v6, v6, tcpSyn)},
|
||||
{"ipv6/udp", v6, buildIPv6(v6, v6, udpDatagram)},
|
||||
{"ipv6/icmpv6", v6, buildIPv6(v6, v6, icmpv6Echo)},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
f, tun := newSelfForwardInterface(tt.addr)
|
||||
// consumeInsidePacket writes through the slice it is handed, so a
|
||||
// packet that arrived with a valid checksum must come back out of
|
||||
// bytes taken before the call, unchanged.
|
||||
want := append([]byte(nil), tt.pkt...)
|
||||
consumeInside(f, tt.pkt)
|
||||
|
||||
if immediatelyForwardToSelf {
|
||||
require.Len(t, tun.writes, 1)
|
||||
assert.Equal(t, want, tun.writes[0])
|
||||
} else {
|
||||
assert.Empty(t, tun.writes, "self traffic reaches the tun over loopback here and must be dropped")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestConsumeInsidePacketSelfTrafficChecksum shows that the self-forward
|
||||
// returns the bytes it was handed, so a packet that arrived with a wrong
|
||||
// upper-layer checksum is written back with that same wrong checksum and the
|
||||
// kernel drops it on re-entry.
|
||||
//
|
||||
// This is how a macOS host loses TCP and UDP to its own IPv6 overlay address:
|
||||
// the kernel writes only the pseudo-header sum into the checksum field and
|
||||
// defers completion to hardware offload, state that does not survive the
|
||||
// crossing into userspace. Which kernels do this, for which protocols and IP
|
||||
// versions, is a property of the kernel and belongs to a test against a live
|
||||
// one; here the checksum is simply wrong, and the forward must make it right.
|
||||
func TestConsumeInsidePacketSelfTrafficChecksum(t *testing.T) {
|
||||
if !immediatelyForwardToSelf {
|
||||
t.Skip("self traffic never reaches the tun on this platform")
|
||||
}
|
||||
versions := []struct {
|
||||
name string
|
||||
addr netip.Addr
|
||||
build func(src, dst netip.Addr, p l4Proto) []byte
|
||||
l4At int
|
||||
valid func(pkt []byte) bool
|
||||
}{
|
||||
{"v4", netip.MustParseAddr("100.100.1.42"), buildIPv4, ipv4HeaderLen, l4ChecksumValid4},
|
||||
{"v6", netip.MustParseAddr("fd00::42"), buildIPv6, ipv6HeaderLen, l4ChecksumValid6},
|
||||
}
|
||||
for _, v := range versions {
|
||||
for _, p := range []l4Proto{tcpSyn, udpDatagram} {
|
||||
t.Run(v.name+"/"+p.name, func(t *testing.T) {
|
||||
pkt := v.build(v.addr, v.addr, p)
|
||||
binary.BigEndian.PutUint16(pkt[v.l4At+p.cksumAt:], 0x1234)
|
||||
require.False(t, v.valid(pkt), "the packet under test must start with a wrong checksum")
|
||||
f, tun := newSelfForwardInterface(v.addr)
|
||||
consumeInside(f, pkt)
|
||||
require.Len(t, tun.writes, 1)
|
||||
assert.True(t, v.valid(tun.writes[0]),
|
||||
"a forwarded %s packet must carry a valid checksum, got 0x%04x",
|
||||
p.name, binary.BigEndian.Uint16(tun.writes[0][v.l4At+p.cksumAt:]))
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
+2
-6
@@ -2,7 +2,6 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -107,8 +106,8 @@ type Interface struct {
|
||||
sendRecvErrorConfig recvErrorConfig
|
||||
acceptRecvErrorConfig recvErrorConfig
|
||||
|
||||
// Bumped on every udp rebind, tunnels compare it to decide they need a punch from the far side
|
||||
rebindEpoch atomic.Uint32
|
||||
// rebindCount is used to decide if an active tunnel should trigger a punch notification through a lighthouse
|
||||
rebindCount int8
|
||||
version string
|
||||
|
||||
conntrackCacheTimeout time.Duration
|
||||
@@ -271,9 +270,6 @@ 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 && !f.outside.SupportsMultipleReaders() {
|
||||
|
||||
@@ -1,146 +0,0 @@
|
||||
package iputil
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
|
||||
"github.com/slackhq/nebula/overlay/checksum"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
const udpHeaderLen = 8
|
||||
|
||||
// SetTransportChecksum recomputes the TCP or UDP checksum of an IPv4 or IPv6
|
||||
// packet in place.
|
||||
//
|
||||
// A kernel that offloads checksums to the NIC hands a packet to a tun with the
|
||||
// transport checksum unfinished: only the pseudo-header sum is in the field and
|
||||
// the rest is left for hardware that a tun does not have. A packet written
|
||||
// straight back to that tun is dropped on re-entry unless the checksum is
|
||||
// completed first. ICMP is left alone; it arrived complete on the kernels this
|
||||
// was measured against.
|
||||
//
|
||||
// So is any packet whose transport header cannot be located: fragments, unknown
|
||||
// extension headers and truncated packets. An IPv6 fragment header is declined
|
||||
// even when it carries the whole datagram (RFC 6946 atomic fragment), because
|
||||
// the walk reports only that a fragment header was present.
|
||||
func SetTransportChecksum(packet []byte) {
|
||||
if len(packet) < 1 {
|
||||
return
|
||||
}
|
||||
switch int(packet[0] >> 4) {
|
||||
case ipv4.Version:
|
||||
setTransportChecksum4(packet)
|
||||
case ipv6.Version:
|
||||
setTransportChecksum6(packet)
|
||||
}
|
||||
}
|
||||
|
||||
func setTransportChecksum4(packet []byte) {
|
||||
if len(packet) < ipv4.HeaderLen {
|
||||
return
|
||||
}
|
||||
ihl := int(packet[0]&0x0f) << 2
|
||||
end := int(binary.BigEndian.Uint16(packet[2:4]))
|
||||
if ihl < ipv4.HeaderLen || end < ihl || end > len(packet) {
|
||||
return
|
||||
}
|
||||
// The checksum covers the whole datagram, which a fragment (MF set or a
|
||||
// non-zero offset) does not carry.
|
||||
if binary.BigEndian.Uint16(packet[6:8])&0x3fff != 0 {
|
||||
return
|
||||
}
|
||||
|
||||
transport, ok := transportExtent(packet[ihl:end], packet[9])
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport)))
|
||||
writeTransportChecksum(transport, packet[9], csum)
|
||||
}
|
||||
|
||||
func setTransportChecksum6(packet []byte) {
|
||||
if len(packet) < ipv6.HeaderLen {
|
||||
return
|
||||
}
|
||||
end := ipv6.HeaderLen + int(binary.BigEndian.Uint16(packet[4:6]))
|
||||
if end > len(packet) {
|
||||
return
|
||||
}
|
||||
|
||||
// The checksum covers the whole datagram, which a fragment does not carry.
|
||||
// An unknown extension header hides where the transport header starts. A
|
||||
// chain longer than the walk's budget ends it early, at an offset that was
|
||||
// never checked against the packet.
|
||||
proto, offset, _, anyFragment, err := IPv6FindUpperProtocol(packet[:end])
|
||||
if err != nil || anyFragment || offset >= end {
|
||||
return
|
||||
}
|
||||
|
||||
transport, ok := transportExtent(packet[offset:end], proto)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport)))
|
||||
writeTransportChecksum(transport, proto, csum)
|
||||
}
|
||||
|
||||
// transportExtent narrows a segment to the length its own header declares. UDP
|
||||
// carries a Length field, and RFC 768 and RFC 8200 section 8.1 both make that
|
||||
// field, not the IP payload extent, the length the pseudo-header counts and the
|
||||
// checksum covers; a datagram padded out to a link's minimum frame is the usual
|
||||
// way the two differ. TCP has no such field, so its segment runs to the end of
|
||||
// the IP payload. A Length that overruns the bytes IP delivered describes a
|
||||
// datagram that is not there.
|
||||
func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
|
||||
if proto != IPProtocolUDP {
|
||||
return transport, true
|
||||
}
|
||||
if len(transport) < udpHeaderLen {
|
||||
return nil, false
|
||||
}
|
||||
ulen := int(binary.BigEndian.Uint16(transport[4:6]))
|
||||
if ulen < udpHeaderLen || ulen > len(transport) {
|
||||
return nil, false
|
||||
}
|
||||
return transport[:ulen], true
|
||||
}
|
||||
|
||||
// writeTransportChecksum stores the checksum of transport, taken over the
|
||||
// pseudo-header sum csum, in the header's checksum field. A UDP checksum that
|
||||
// computes to zero goes on the wire as 0xffff: zero means no checksum was
|
||||
// computed (RFC 768), and over IPv6 the checksum is mandatory (RFC 8200
|
||||
// section 8.1).
|
||||
func writeTransportChecksum(transport []byte, proto uint8, csum uint32) {
|
||||
var at, minLen int
|
||||
switch proto {
|
||||
case IPProtocolTCP:
|
||||
at, minLen = 16, 20
|
||||
case IPProtocolUDP:
|
||||
at, minLen = 6, udpHeaderLen
|
||||
default:
|
||||
return
|
||||
}
|
||||
if len(transport) < minLen {
|
||||
return
|
||||
}
|
||||
|
||||
transport[at], transport[at+1] = 0, 0
|
||||
sum := ^checksum.Checksum(transport, fold(csum))
|
||||
if sum == 0 && proto == IPProtocolUDP {
|
||||
sum = 0xffff
|
||||
}
|
||||
binary.BigEndian.PutUint16(transport[at:], sum)
|
||||
}
|
||||
|
||||
// fold reduces a pseudo-header sum to the 16 bit seed Checksum takes. Carrying
|
||||
// the high half back into the low half is what keeps the reduction lossless, so
|
||||
// the seed sums exactly as the wider value would; 0xffff is its fixed point.
|
||||
// Every term of that sum comes from a 16 bit field, so it stays far below the
|
||||
// width at which the accumulator would wrap.
|
||||
func fold(csum uint32) uint16 {
|
||||
for csum > 0xffff {
|
||||
csum = (csum >> 16) + (csum & 0xffff)
|
||||
}
|
||||
return uint16(csum)
|
||||
}
|
||||
@@ -1,242 +0,0 @@
|
||||
package iputil
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
// serialize builds a packet with gopacket, whose checksums are computed
|
||||
// independently of this package.
|
||||
func serialize(t *testing.T, ls ...gopacket.SerializableLayer) []byte {
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
require.NoError(t, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, ls...))
|
||||
return append([]byte(nil), buf.Bytes()...)
|
||||
}
|
||||
|
||||
// withExtensionHeader inserts an 8 byte IPv6 extension header of the given
|
||||
// type between the IPv6 header and its payload. The transport checksum does not
|
||||
// change: the pseudo-header counts only upper-layer bytes.
|
||||
func withExtensionHeader(pkt []byte, typ layers.IPProtocol, hdr [8]byte) []byte {
|
||||
hdr[0] = pkt[6]
|
||||
out := make([]byte, 0, len(pkt)+8)
|
||||
out = append(out, pkt[:40]...)
|
||||
out = append(out, hdr[:]...)
|
||||
out = append(out, pkt[40:]...)
|
||||
out[6] = byte(typ)
|
||||
binary.BigEndian.PutUint16(out[4:6], binary.BigEndian.Uint16(pkt[4:6])+8)
|
||||
return out
|
||||
}
|
||||
|
||||
// truncate copies the first n bytes into a buffer of exactly that capacity, so
|
||||
// a read past the length panics instead of quietly succeeding.
|
||||
func truncate(pkt []byte, n int) []byte {
|
||||
out := make([]byte, n)
|
||||
copy(out, pkt)
|
||||
return out
|
||||
}
|
||||
|
||||
// extChain builds an IPv6 packet fronted by n Destination Options headers. Each
|
||||
// points at another one, so the walk spends its whole budget without reaching a
|
||||
// transport header. lastExtLen inflates the final header's declared length,
|
||||
// which is how the walk ends up past the end of the packet.
|
||||
func extChain(n int, lastExtLen byte) []byte {
|
||||
pkt := make([]byte, ipv6.HeaderLen)
|
||||
pkt[0], pkt[6], pkt[7] = 0x60, 60, 64
|
||||
for i := range n {
|
||||
h := make([]byte, 8)
|
||||
h[0] = 60
|
||||
if i == n-1 {
|
||||
h[1] = lastExtLen
|
||||
}
|
||||
pkt = append(pkt, h...)
|
||||
}
|
||||
pkt = append(pkt, make([]byte, 20)...)
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(pkt)-ipv6.HeaderLen))
|
||||
return pkt
|
||||
}
|
||||
|
||||
func TestSetTransportChecksum(t *testing.T) {
|
||||
// Source and destination differ so that a pseudo-header built from the wrong
|
||||
// one, or from the two swapped, does not land on the same checksum anyway.
|
||||
v4 := func(proto layers.IPProtocol) *layers.IPv4 {
|
||||
return &layers.IPv4{Version: 4, TTL: 64, Id: 0x1234, Protocol: proto, SrcIP: net.IPv4(192, 0, 2, 1).To4(), DstIP: net.IPv4(198, 51, 100, 2).To4()}
|
||||
}
|
||||
v6 := func(proto layers.IPProtocol) *layers.IPv6 {
|
||||
return &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: net.ParseIP("2001:db8::1"), DstIP: net.ParseIP("2001:db8:1::2")}
|
||||
}
|
||||
tcp := func(ip gopacket.NetworkLayer) *layers.TCP {
|
||||
l := &layers.TCP{SrcPort: 49152, DstPort: 443, SYN: true, Window: 65535}
|
||||
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
|
||||
return l
|
||||
}
|
||||
udp := func(ip gopacket.NetworkLayer) *layers.UDP {
|
||||
l := &layers.UDP{SrcPort: 49152, DstPort: 53}
|
||||
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
|
||||
return l
|
||||
}
|
||||
payload := gopacket.Payload("self")
|
||||
nop := layers.IPv4Option{OptionType: 1, OptionLength: 1}
|
||||
|
||||
ip4tcp := v4(layers.IPProtocolTCP)
|
||||
ip4opts := v4(layers.IPProtocolTCP)
|
||||
ip4opts.Options = []layers.IPv4Option{nop, nop, nop, nop}
|
||||
ip4udp := v4(layers.IPProtocolUDP)
|
||||
ip6tcp := v6(layers.IPProtocolTCP)
|
||||
ip6udp := v6(layers.IPProtocolUDP)
|
||||
hopByHop := [8]byte{0, 0, 1, 4} // next header, length 0, PadN of 4
|
||||
|
||||
// Bytes past the length the IP header declares are not part of the
|
||||
// datagram and must not be summed.
|
||||
trailing4 := append(serialize(t, ip4tcp, tcp(ip4tcp), payload), []byte("trailing")...)
|
||||
trailing6 := append(serialize(t, ip6tcp, tcp(ip6tcp), payload), []byte("trailing")...)
|
||||
|
||||
// A datagram padded out past the length UDP declares: the pseudo-header
|
||||
// counts the UDP Length field, so the checksum is the unpadded one.
|
||||
padded4 := append(serialize(t, ip4udp, udp(ip4udp), payload), []byte("pad!")...)
|
||||
binary.BigEndian.PutUint16(padded4[2:4], uint16(len(padded4)))
|
||||
padded6 := append(serialize(t, ip6udp, udp(ip6udp), payload), []byte("pad!")...)
|
||||
binary.BigEndian.PutUint16(padded6[4:6], uint16(len(padded6)-ipv6.HeaderLen))
|
||||
|
||||
// Corrupting the checksum and asking for it back must yield gopacket's
|
||||
// packet, byte for byte.
|
||||
recomputed := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
cksum int
|
||||
}{
|
||||
{"v4 tcp", serialize(t, ip4tcp, tcp(ip4tcp), payload), 20 + 16},
|
||||
{"v4 tcp with ip options", serialize(t, ip4opts, tcp(ip4opts), payload), 24 + 16},
|
||||
{"v4 udp", serialize(t, ip4udp, udp(ip4udp), payload), 20 + 6},
|
||||
{"v4 tcp header only", serialize(t, ip4tcp, tcp(ip4tcp)), 20 + 16},
|
||||
{"v4 udp header only", serialize(t, ip4udp, udp(ip4udp)), 20 + 6},
|
||||
{"v6 tcp", serialize(t, ip6tcp, tcp(ip6tcp), payload), 40 + 16},
|
||||
{"v6 udp", serialize(t, ip6udp, udp(ip6udp), payload), 40 + 6},
|
||||
{"v6 udp header only", serialize(t, ip6udp, udp(ip6udp)), 40 + 6},
|
||||
{"v6 tcp behind hop-by-hop", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6HopByHop, hopByHop), 48 + 16},
|
||||
{"v4 tcp with bytes past the total length", trailing4, 20 + 16},
|
||||
{"v6 tcp with bytes past the payload length", trailing6, 40 + 16},
|
||||
{"v4 udp padded past its declared length", padded4, 20 + 6},
|
||||
{"v6 udp padded past its declared length", padded6, 40 + 6},
|
||||
}
|
||||
for _, tt := range recomputed {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := append([]byte(nil), tt.pkt...)
|
||||
binary.BigEndian.PutUint16(got[tt.cksum:], 0x1234)
|
||||
require.NotEqual(t, tt.pkt, got)
|
||||
SetTransportChecksum(got)
|
||||
assert.Equal(t, tt.pkt, got)
|
||||
})
|
||||
}
|
||||
|
||||
ip4frag := v4(layers.IPProtocolTCP)
|
||||
ip4frag.Flags = layers.IPv4MoreFragments
|
||||
ip4later := v4(layers.IPProtocolTCP)
|
||||
ip4later.FragOffset = 1
|
||||
ip4icmp := v4(layers.IPProtocolICMPv4)
|
||||
|
||||
badIHL := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
||||
badIHL[0] = 0x44 // header length 16, shorter than an ipv4 header
|
||||
shortTotalLen := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
||||
binary.BigEndian.PutUint16(shortTotalLen[2:4], 10) // shorter than the header it introduces
|
||||
cutTCP := serialize(t, ip4tcp, tcp(ip4tcp), payload)
|
||||
binary.BigEndian.PutUint16(cutTCP[2:4], 20+19) // one byte short of a tcp header
|
||||
cutTCP = truncate(cutTCP, 20+19)
|
||||
cutUDP := serialize(t, ip4udp, udp(ip4udp), payload)
|
||||
binary.BigEndian.PutUint16(cutUDP[2:4], 20+7) // one byte short of a udp header
|
||||
cutUDP = truncate(cutUDP, 20+7)
|
||||
// Two bytes short, so a transport header survives whole and the minimum
|
||||
// length check cannot stand in for the bounds check.
|
||||
cutV6 := truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 62)
|
||||
fragment := [8]byte{0, 0, 0, 1, 0, 0, 0, 1} // next header, reserved, offset 0 with M set, id
|
||||
overrun4 := serialize(t, ip4udp, udp(ip4udp), payload)
|
||||
binary.BigEndian.PutUint16(overrun4[24:26], uint16(len(overrun4)-20+1)) // one byte past what ip delivered
|
||||
overrun6 := serialize(t, ip6udp, udp(ip6udp), payload)
|
||||
binary.BigEndian.PutUint16(overrun6[44:46], uint16(len(overrun6)-ipv6.HeaderLen+1))
|
||||
shortUDPLen := serialize(t, ip4udp, udp(ip4udp), payload)
|
||||
binary.BigEndian.PutUint16(shortUDPLen[24:26], 7) // shorter than the header it counts
|
||||
|
||||
// Where the checksum cannot be completed the packet is left as it came.
|
||||
untouched := []struct {
|
||||
name string
|
||||
pkt []byte
|
||||
cksum int
|
||||
}{
|
||||
{"v4 first fragment", serialize(t, ip4frag, tcp(ip4frag), payload), 20 + 16},
|
||||
{"v4 later fragment", serialize(t, ip4later, tcp(ip4later), payload), 20 + 16},
|
||||
{"v4 icmp", serialize(t, ip4icmp, &layers.ICMPv4{TypeCode: layers.CreateICMPv4TypeCode(8, 0), Id: 1, Seq: 1}, payload), 20 + 2},
|
||||
{"v4 header length below the minimum", badIHL, 20 + 16},
|
||||
{"v4 total length below the header length", shortTotalLen, 20 + 16},
|
||||
{"v4 truncated below its total length", truncate(serialize(t, ip4tcp, tcp(ip4tcp), payload), 30), -1},
|
||||
{"v4 tcp header cut short", cutTCP, 20 + 16},
|
||||
{"v4 udp header cut short", cutUDP, -1},
|
||||
{"v6 fragment", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6Fragment, fragment), 48 + 16},
|
||||
{"v6 truncated below its payload length", truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 50), -1},
|
||||
{"v6 truncated with a whole transport header still present", cutV6, 40 + 16},
|
||||
{"v6 extension header chain longer than the walk", extChain(9, 0), 112 + 16},
|
||||
{"v6 extension header chain running past the packet", extChain(8, 255), 104 + 16},
|
||||
{"v4 udp length past the end of the datagram", overrun4, 20 + 6},
|
||||
{"v6 udp length past the end of the datagram", overrun6, 40 + 6},
|
||||
{"v4 udp length below a udp header", shortUDPLen, 20 + 6},
|
||||
}
|
||||
for _, tt := range untouched {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if tt.cksum >= 0 {
|
||||
binary.BigEndian.PutUint16(tt.pkt[tt.cksum:], 0x1234)
|
||||
}
|
||||
want := append([]byte(nil), tt.pkt...)
|
||||
SetTransportChecksum(tt.pkt)
|
||||
assert.Equal(t, want, tt.pkt)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("too short to carry a header", func(t *testing.T) {
|
||||
for _, pkt := range [][]byte{nil, {}, {0x45}, {0x60}} {
|
||||
assert.NotPanics(t, func() { SetTransportChecksum(pkt) })
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp checksum of zero goes out as zero", func(t *testing.T) {
|
||||
pkt := serialize(t, ip4tcp, tcp(ip4tcp), gopacket.Payload{0, 0})
|
||||
c := binary.BigEndian.Uint16(pkt[36:38])
|
||||
require.NotZero(t, c)
|
||||
// Only udp reserves zero to mean "not computed", so tcp keeps it.
|
||||
binary.BigEndian.PutUint16(pkt[40:42], c)
|
||||
SetTransportChecksum(pkt)
|
||||
assert.Zero(t, binary.BigEndian.Uint16(pkt[36:38]))
|
||||
})
|
||||
|
||||
t.Run("udp checksum of zero goes out as 0xffff", func(t *testing.T) {
|
||||
pkt := serialize(t, ip4udp, udp(ip4udp), gopacket.Payload{0, 0})
|
||||
c := binary.BigEndian.Uint16(pkt[26:28])
|
||||
require.NotZero(t, c)
|
||||
// The one's complement sum is now 0xffff - c; adding c to the payload
|
||||
// makes it 0xffff, whose complement is zero.
|
||||
binary.BigEndian.PutUint16(pkt[28:30], c)
|
||||
SetTransportChecksum(pkt)
|
||||
assert.Equal(t, uint16(0xffff), binary.BigEndian.Uint16(pkt[26:28]))
|
||||
})
|
||||
}
|
||||
|
||||
func TestFold(t *testing.T) {
|
||||
// 0xffff is the fold's fixed point, so a loop bound one notch tight never
|
||||
// terminates on it.
|
||||
for _, tt := range []struct {
|
||||
in uint32
|
||||
want uint16
|
||||
}{
|
||||
{0, 0},
|
||||
{0xffff, 0xffff},
|
||||
{0x10000, 1},
|
||||
{0x1fffe, 0xffff},
|
||||
{0xffffffff, 0xffff},
|
||||
} {
|
||||
assert.Equal(t, tt.want, fold(tt.in))
|
||||
}
|
||||
}
|
||||
@@ -27,13 +27,6 @@ const (
|
||||
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
||||
|
||||
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
||||
|
||||
IPProtocolICMP = 1
|
||||
IPProtocolICMPv6 = 58
|
||||
IPProtocolTCP = 6
|
||||
IPProtocolUDP = 17
|
||||
ICMPv6TypeEchoRequest = 128
|
||||
ICMPv6TypeEchoReply = 129
|
||||
)
|
||||
|
||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||
|
||||
+1
-15
@@ -34,9 +34,7 @@ type LightHouse struct {
|
||||
|
||||
myVpnNetworks []netip.Prefix
|
||||
myVpnNetworksTable *bart.Lite
|
||||
// myVpnAddrsTable contains our overlay host addrs, as opposed to the overlay networks
|
||||
myVpnAddrsTable *bart.Lite
|
||||
punchy *Punchy
|
||||
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.
|
||||
@@ -106,7 +104,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
||||
amLighthouse: amLighthouse,
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
||||
addrMap: make(map[netip.Addr]*RemoteList),
|
||||
nebulaPort: nebulaPort,
|
||||
punchy: p,
|
||||
@@ -1161,17 +1158,6 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
|
||||
return
|
||||
}
|
||||
|
||||
// Don't respond to requests for us.
|
||||
if lhh.lh.myVpnAddrsTable.Contains(queryVpnAddr) {
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("Ignoring HostQuery for one of my own addresses",
|
||||
"fromVpnAddrs", fromVpnAddrs,
|
||||
"queryVpnAddr", queryVpnAddr,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
found, ln, err := lhh.lh.queryAndPrepMessage(queryVpnAddr, func(c *cache) (int, error) {
|
||||
n = lhh.resetMeta()
|
||||
n.Type = NebulaMeta_HostQueryReply
|
||||
|
||||
+53
-79
@@ -27,27 +27,15 @@ func TestOldIPv4Only(t *testing.T) {
|
||||
assert.Equal(t, binary.BigEndian.Uint32(bp[:]), m.GetAddr())
|
||||
}
|
||||
|
||||
func testCertState(networks ...netip.Prefix) *CertState {
|
||||
cs := &CertState{
|
||||
myVpnNetworks: networks,
|
||||
myVpnNetworksTable: new(bart.Lite),
|
||||
myVpnAddrs: make([]netip.Addr, 0, len(networks)),
|
||||
myVpnAddrsTable: new(bart.Lite),
|
||||
}
|
||||
|
||||
for _, n := range networks {
|
||||
cs.myVpnNetworksTable.Insert(n)
|
||||
cs.myVpnAddrs = append(cs.myVpnAddrs, n.Addr())
|
||||
cs.myVpnAddrsTable.Insert(netip.PrefixFrom(n.Addr(), n.Addr().BitLen()))
|
||||
}
|
||||
|
||||
return cs
|
||||
}
|
||||
|
||||
func Test_lhStaticMapping(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh1 := "10.128.0.2"
|
||||
|
||||
c := config.NewC(l)
|
||||
@@ -67,7 +55,12 @@ func Test_lhStaticMapping(t *testing.T) {
|
||||
func TestReloadLighthouseInterval(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh1 := "10.128.0.2"
|
||||
|
||||
c := config.NewC(l)
|
||||
@@ -97,7 +90,12 @@ func TestReloadLighthouseInterval(t *testing.T) {
|
||||
func BenchmarkLighthouseHandleRequest(b *testing.B) {
|
||||
l := test.NewLogger()
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/0")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
|
||||
c := config.NewC(l)
|
||||
lh, err := NewLightHouseFromConfig(b.Context(), l, c, cs, nil, nil)
|
||||
@@ -197,7 +195,12 @@ func TestLighthouse_Memory(t *testing.T) {
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
require.NoError(t, err)
|
||||
@@ -277,7 +280,12 @@ func TestLighthouse_reload(t *testing.T) {
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
@@ -307,7 +315,12 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
@@ -416,9 +429,7 @@ func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
||||
}
|
||||
|
||||
// sendLHHostRequest delivers a HostQuery to lhh and hands back the writer that
|
||||
// captured what it emitted. Pass a nil filter to see every message.
|
||||
func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler, filter *NebulaMeta_MessageType) *testEncWriter {
|
||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||
req := &NebulaMeta{
|
||||
Type: NebulaMeta_HostQuery,
|
||||
Details: &NebulaMetaDetails{},
|
||||
@@ -436,59 +447,12 @@ func sendLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr,
|
||||
panic(err)
|
||||
}
|
||||
|
||||
w := &testEncWriter{metaFilter: filter}
|
||||
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
|
||||
return w
|
||||
}
|
||||
|
||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||
filter := NebulaMeta_HostQueryReply
|
||||
return sendLHHostRequest(fromAddr, myVpnIp, queryVpnIp, lhh, &filter).lastReply
|
||||
}
|
||||
|
||||
func TestLighthouse_IgnoresHostQueryForItself(t *testing.T) {
|
||||
// Validate that we don't answer host queries for our own address.
|
||||
l := test.NewLogger()
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
myVpnIp := myVpnNet.Addr()
|
||||
|
||||
c := config.NewC(l)
|
||||
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
||||
// Add a static_host_map entry for ourselves, so our address
|
||||
// is in the addrMap.
|
||||
c.Settings["static_host_map"] = map[string]any{
|
||||
myVpnIp.String(): []any{"192.168.100.1:4242"},
|
||||
w := &testEncWriter{
|
||||
metaFilter: &filter,
|
||||
}
|
||||
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, testCertState(myVpnNet), nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
lhh := lh.NewRequestHandler()
|
||||
|
||||
peerVpnIp := netip.MustParseAddr("10.128.0.2")
|
||||
peerUdpAddr := netip.MustParseAddrPort("10.0.0.2:4242")
|
||||
otherVpnIp := netip.MustParseAddr("10.128.0.3")
|
||||
otherUdpAddr := netip.MustParseAddrPort("10.0.0.3:4242")
|
||||
|
||||
newLHHostUpdate(peerUdpAddr, peerVpnIp, []netip.AddrPort{peerUdpAddr}, lhh)
|
||||
newLHHostUpdate(otherUdpAddr, otherVpnIp, []netip.AddrPort{otherUdpAddr}, lhh)
|
||||
|
||||
// Control: a query about a real peer is still answered, and still ends with
|
||||
// the punch notification aimed at the host that was asked about.
|
||||
w := sendLHHostRequest(peerUdpAddr, peerVpnIp, otherVpnIp, lhh, nil)
|
||||
require.NotNil(t, w.lastReply.msg)
|
||||
assert.Equal(t, NebulaMeta_HostPunchNotification, w.lastReply.msg.Type)
|
||||
assert.Equal(t, otherVpnIp, w.lastReply.vpnIp)
|
||||
|
||||
// Now validate that we don't send to ourselves.
|
||||
found, _, err := lh.queryAndPrepMessage(myVpnIp, func(*cache) (int, error) { return 0, nil })
|
||||
require.NoError(t, err)
|
||||
require.True(t, found, "the lighthouse should hold a cache entry for its own address")
|
||||
|
||||
w = sendLHHostRequest(peerUdpAddr, peerVpnIp, myVpnIp, lhh, nil)
|
||||
assert.Nil(t, w.lastReply.msg, "a query about our own address must produce no reply and no punch notification")
|
||||
lhh.HandleRequest(fromAddr, []netip.Addr{myVpnIp}, b, w)
|
||||
return w.lastReply
|
||||
}
|
||||
|
||||
func newLHHostUpdate(fromAddr netip.AddrPort, vpnIp netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
|
||||
@@ -678,7 +642,12 @@ func TestLighthouse_Dont_Delete_Static_Hosts(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
@@ -739,7 +708,12 @@ func TestLighthouse_DeletesWork(t *testing.T) {
|
||||
}
|
||||
|
||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
||||
cs := testCertState(myVpnNet)
|
||||
nt := new(bart.Lite)
|
||||
nt.Insert(myVpnNet)
|
||||
cs := &CertState{
|
||||
myVpnNetworks: []netip.Prefix{myVpnNet},
|
||||
myVpnNetworksTable: nt,
|
||||
}
|
||||
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
|
||||
require.NoError(t, err)
|
||||
lh.ifce = &mockEncWriter{}
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/cpupick"
|
||||
"github.com/slackhq/nebula/diag"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
@@ -69,9 +68,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
}
|
||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||
|
||||
commands := diag.NewRegistry()
|
||||
|
||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"), commands)
|
||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||
}
|
||||
@@ -191,7 +188,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
Listen: listen,
|
||||
Multi: routines > 1,
|
||||
Batch: batchSize,
|
||||
Offloads: c.GetBool("listen.udp_offloads", false),
|
||||
Offloads: c.GetBool("listen.udp_offloads", true),
|
||||
}
|
||||
udpServer, err := udp.NewListener(l, udpSettings)
|
||||
if err != nil {
|
||||
@@ -247,12 +244,10 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
|
||||
// The operator didn't choose pin CPUs, so pick a default set that
|
||||
// prefers performance cores and doesn't stack co-located instances
|
||||
// onto allowed[0].
|
||||
|
||||
// key is used to seed the spreading of routines->cores.
|
||||
// use PID if you want to ensure many different Nebulas in VMs or containers land on different cores
|
||||
// use port if you want to always end up on the same cores, ideal for benchmarking.
|
||||
key := uint64(os.Getpid()) //default to PID
|
||||
// onto allowed[0]. The bound UDP port keys the per-instance spread:
|
||||
// distinct across instances sharing a box, stable across restarts.
|
||||
// A nil result keeps listenIn's stock allowed[i] fallback.
|
||||
key := uint64(os.Getpid())
|
||||
pinKeyStr := strings.ToLower(c.GetString("tun.pin_threads_key", ""))
|
||||
switch pinKeyStr {
|
||||
case "":
|
||||
@@ -260,16 +255,14 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
case "pid":
|
||||
l.Debug("tun.pin_threads_key is PID")
|
||||
case "port":
|
||||
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
|
||||
l.Info("tun.pin_threads_key is port number")
|
||||
key = uint64(ap.Port())
|
||||
} else {
|
||||
l.Warn("Failed to get a port number for tun.pin_threads_key, falling back to PID", "err", err)
|
||||
}
|
||||
l.Info("tun.pin_threads_key is port number")
|
||||
default:
|
||||
l.Warn("tun.pin_threads_key is invalid, using PID")
|
||||
}
|
||||
|
||||
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
|
||||
key = uint64(ap.Port())
|
||||
}
|
||||
cpuAffinity = cpupick.Default(routines, key, l)
|
||||
}
|
||||
|
||||
@@ -325,20 +318,13 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
return nil, util.ContextualizeIfNeeded("Failed to start stats emitter", err)
|
||||
}
|
||||
|
||||
// Built before the configTest return so that a bad ctl block fails `nebula -test`. It only
|
||||
// holds the registry, which attachCommands populates below, and reads nothing until Start.
|
||||
ctlServer, err := newCtlServerFromConfig(ctx, l.With("subsystem", "ctl"), c, commands)
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Failed to configure the ctl socket", err)
|
||||
}
|
||||
|
||||
if configTest {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
go ifce.emitStats(ctx, c.GetDuration("stats.interval", time.Second*10))
|
||||
|
||||
attachCommands(l, c, commands, ifce)
|
||||
attachCommands(l, c, ssh, ifce)
|
||||
|
||||
networkChanges := udp.NewNetworkChangeMonitor(ctx, l, c)
|
||||
|
||||
@@ -349,7 +335,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
sshStart: sshStart,
|
||||
ctlStart: ctlServer.Start,
|
||||
statsStart: stats.Start,
|
||||
dnsStart: ds.Start,
|
||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||
|
||||
+65
-4
@@ -4,16 +4,77 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/boring"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
|
||||
// unsafe needed for go:linkname
|
||||
_ "unsafe"
|
||||
|
||||
"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
|
||||
|
||||
var boringEnabled = boring.Enabled()
|
||||
// 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)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/boring"
|
||||
"encoding/hex"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -12,3 +14,33 @@ 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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -40,11 +40,8 @@ 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 noise.CipherAESGCM.CipherName():
|
||||
case CipherAESGCM.CipherName():
|
||||
return NewCipherStateAESGCM(s)
|
||||
case noise.CipherChaChaPoly.CipherName():
|
||||
return NewCipherStateChaChaPoly(s)
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
@@ -12,30 +11,24 @@ import (
|
||||
|
||||
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
roundtrip(t, NewCipherState(enc, CipherAESGCM), NewCipherState(dec, CipherAESGCM))
|
||||
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
roundtrip(t, NewCipherState(enc, noise.CipherChaChaPoly), NewCipherState(dec, noise.CipherChaChaPoly))
|
||||
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
}
|
||||
|
||||
func TestNewCipherStateDispatch(t *testing.T) {
|
||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
|
||||
if !boringEnabled && !fips140.Enabled() {
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
} else {
|
||||
// fips140
|
||||
assert.IsType(t, encA.Cipher(), NewCipherState(encA, CipherAESGCM))
|
||||
}
|
||||
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||
enc, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||
assert.Panics(t, func() {
|
||||
NewCipherState(enc, fakeCipher{})
|
||||
})
|
||||
|
||||
@@ -1,197 +0,0 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/cipher"
|
||||
"crypto/fips140"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"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) {
|
||||
if c == nil {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
if n >= RejectAfterMessages {
|
||||
return nil, ErrMessageCounterExhausted
|
||||
}
|
||||
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) {
|
||||
if c == nil {
|
||||
return []byte{}, nil
|
||||
}
|
||||
binary.BigEndian.PutUint64(nb[4:], n)
|
||||
return c.Open(out, nb, ciphertext, ad)
|
||||
}
|
||||
|
||||
func (c *aeadGCMFIPS140Cipher) Overhead() int {
|
||||
if c == nil {
|
||||
return 0
|
||||
}
|
||||
return c.AEAD.Overhead()
|
||||
}
|
||||
|
||||
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[:]
|
||||
}
|
||||
|
||||
func init() {
|
||||
if boringEnabled || fips140.Enabled() {
|
||||
initSelfTestAESGCMFIPS140()
|
||||
}
|
||||
}
|
||||
|
||||
// validates the go:linkname + reflection extraction and the nonce-reuse
|
||||
// protection at startup. 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).
|
||||
func initSelfTestAESGCMFIPS140() {
|
||||
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 defensive check that FIPS-140 requires, so
|
||||
// if the extraction ever silently yields an AEAD without it, refuse to start.
|
||||
if !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
|
||||
}
|
||||
@@ -1,48 +0,0 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,13 +0,0 @@
|
||||
//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")
|
||||
}
|
||||
}
|
||||
+4
-15
@@ -1,25 +1,14 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// EncryptLockNeeded indicates if calls to Encrypt need a lock
|
||||
var EncryptLockNeeded = fips140.Enabled()
|
||||
const EncryptLockNeeded = false
|
||||
|
||||
var CipherAESGCM noise.CipherFunc = initAESGCM()
|
||||
|
||||
func initAESGCM() noise.CipherFunc {
|
||||
if fips140.Enabled() {
|
||||
return CipherAESGCMFIPS140
|
||||
} else {
|
||||
return noise.CipherAESGCM
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
var boringEnabled = false
|
||||
// CipherAESGCM is the standard noise.CipherAESGCM when boringcrypto is not enabled
|
||||
var CipherAESGCM noise.CipherFunc = noise.CipherAESGCM
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
//go:build !boringcrypto
|
||||
// +build !boringcrypto
|
||||
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestEncryptLockNeeded(t *testing.T) {
|
||||
assert.False(t, EncryptLockNeeded)
|
||||
}
|
||||
+7
-6
@@ -8,6 +8,7 @@ import (
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
"github.com/google/gopacket/layers"
|
||||
"golang.org/x/net/ipv6"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
@@ -368,15 +369,15 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch proto {
|
||||
case iputil.IPProtocolICMPv6:
|
||||
switch layers.IPProtocol(proto) {
|
||||
case layers.IPProtocolICMPv6:
|
||||
// An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read.
|
||||
if dataLen < offset+4 {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||
switch data[offset] { //icmp type
|
||||
case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
|
||||
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
|
||||
if dataLen < offset+6 {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
@@ -385,7 +386,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
fp.RemotePort = 0
|
||||
}
|
||||
|
||||
case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
|
||||
case layers.IPProtocolTCP, layers.IPProtocolUDP:
|
||||
if dataLen < offset+4 {
|
||||
return ErrIPv6PacketTooShort
|
||||
}
|
||||
@@ -434,7 +435,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
|
||||
minLen := ihl
|
||||
if !fp.Fragment {
|
||||
if fp.Protocol == iputil.IPProtocolICMP {
|
||||
if fp.Protocol == firewall.ProtoICMP {
|
||||
minLen += minFwPacketLen + 2
|
||||
} else {
|
||||
minLen += minFwPacketLen
|
||||
@@ -456,7 +457,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
|
||||
if fp.Fragment {
|
||||
fp.RemotePort = 0
|
||||
fp.LocalPort = 0
|
||||
} else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
|
||||
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
|
||||
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
|
||||
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
|
||||
} else if incoming {
|
||||
|
||||
+19
-20
@@ -9,7 +9,6 @@ import (
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -59,7 +58,7 @@ func Test_newPacket(t *testing.T) {
|
||||
Src: net.IPv4(10, 0, 0, 1),
|
||||
Dst: net.IPv4(10, 0, 0, 2),
|
||||
Options: []byte{0, 1, 0, 2},
|
||||
Protocol: iputil.IPProtocolTCP,
|
||||
Protocol: firewall.ProtoTCP,
|
||||
}
|
||||
|
||||
b, _ = h.Marshal()
|
||||
@@ -67,7 +66,7 @@ func Test_newPacket(t *testing.T) {
|
||||
err = newPacket(b, true, p)
|
||||
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr)
|
||||
assert.Equal(t, uint16(3), p.RemotePort)
|
||||
@@ -240,7 +239,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// A good UDP packet
|
||||
ip = layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: iputil.IPProtocolUDP,
|
||||
NextHeader: firewall.ProtoUDP,
|
||||
HopLimit: 128,
|
||||
SrcIP: net.IPv6linklocalallrouters,
|
||||
DstIP: net.IPv6linklocalallnodes,
|
||||
@@ -263,7 +262,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// incoming
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
@@ -273,7 +272,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// outgoing
|
||||
err = newPacket(b, false, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||
@@ -290,7 +289,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// incoming
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
@@ -300,7 +299,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
// outgoing
|
||||
err = newPacket(b, false, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
||||
assert.Equal(t, uint16(36123), p.LocalPort)
|
||||
@@ -345,7 +344,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
|
||||
err = newPacket(b, true, p)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||
assert.Equal(t, uint16(36123), p.RemotePort)
|
||||
@@ -679,7 +678,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
pkt := make([]byte, realTCPAt+4)
|
||||
pkt[0] = 0x60 // version 6
|
||||
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
|
||||
pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
|
||||
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
|
||||
pkt[41] = 255 // HdrExtLen = 255
|
||||
|
||||
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
|
||||
@@ -688,7 +687,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
|
||||
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
|
||||
|
||||
require.NoError(t, newPacket(pkt, true, p))
|
||||
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||
// LocalPort is the destination port for incoming traffic. It must be the real port (22)
|
||||
// the host delivers to, not the forged 443 at the overflowed offset.
|
||||
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
|
||||
@@ -767,7 +766,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
|
||||
v4 := make([]byte, 28)
|
||||
v4[0] = 0x45
|
||||
v4[9] = iputil.IPProtocolTCP
|
||||
v4[9] = firewall.ProtoTCP
|
||||
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
|
||||
require.NoError(t, newPacket(v4, true, p))
|
||||
assert.Equal(t, 20, p.IPHdrLen)
|
||||
@@ -778,7 +777,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
// (Fragment false) but the coalescer must not touch it (FragAny true).
|
||||
ff := make([]byte, 28)
|
||||
ff[0] = 0x45
|
||||
ff[9] = iputil.IPProtocolUDP
|
||||
ff[9] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
|
||||
require.NoError(t, newPacket(ff, true, p))
|
||||
assert.False(t, p.Fragment)
|
||||
@@ -788,7 +787,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
// IPv4 non-first fragment (nonzero offset): both flags set.
|
||||
nf := make([]byte, 28)
|
||||
nf[0] = 0x45
|
||||
nf[9] = iputil.IPProtocolUDP
|
||||
nf[9] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
|
||||
require.NoError(t, newPacket(nf, true, p))
|
||||
assert.True(t, p.Fragment)
|
||||
@@ -797,7 +796,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
|
||||
opts := make([]byte, 32)
|
||||
opts[0] = 0x46
|
||||
opts[9] = iputil.IPProtocolTCP
|
||||
opts[9] = firewall.ProtoTCP
|
||||
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
|
||||
require.NoError(t, newPacket(opts, true, p))
|
||||
assert.Equal(t, 24, p.IPHdrLen)
|
||||
@@ -806,7 +805,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
// Plain IPv6 TCP: L4 at 40.
|
||||
v6 := make([]byte, 60)
|
||||
v6[0] = 0x60
|
||||
v6[6] = iputil.IPProtocolTCP
|
||||
v6[6] = firewall.ProtoTCP
|
||||
require.NoError(t, newPacket(v6, true, p))
|
||||
assert.Equal(t, 40, p.IPHdrLen)
|
||||
assert.False(t, p.FragAny)
|
||||
@@ -815,7 +814,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
hbh := make([]byte, 60)
|
||||
hbh[0] = 0x60
|
||||
hbh[6] = 0 // hop-by-hop
|
||||
hbh[40] = iputil.IPProtocolTCP
|
||||
hbh[40] = firewall.ProtoTCP
|
||||
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
|
||||
require.NoError(t, newPacket(hbh, true, p))
|
||||
assert.Equal(t, 48, p.IPHdrLen)
|
||||
@@ -825,17 +824,17 @@ func Test_newPacket_parsedFields(t *testing.T) {
|
||||
f6 := make([]byte, 60)
|
||||
f6[0] = 0x60
|
||||
f6[6] = 44 // fragment extension header
|
||||
f6[40] = iputil.IPProtocolUDP
|
||||
f6[40] = firewall.ProtoUDP
|
||||
require.NoError(t, newPacket(f6, true, p))
|
||||
assert.True(t, p.FragAny)
|
||||
assert.False(t, p.Fragment)
|
||||
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
|
||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||
|
||||
// IPv6 non-first fragment: both set, walk stops at the fragment header.
|
||||
f6n := make([]byte, 60)
|
||||
f6n[0] = 0x60
|
||||
f6n[6] = 44
|
||||
f6n[40] = iputil.IPProtocolUDP
|
||||
f6n[40] = firewall.ProtoUDP
|
||||
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
|
||||
require.NoError(t, newPacket(f6n, true, p))
|
||||
assert.True(t, p.Fragment)
|
||||
|
||||
@@ -133,7 +133,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||
}
|
||||
nameStr := c.GetString("tun.dev", "")
|
||||
useOffloads := c.GetBool("tun.use_offloads", false)
|
||||
useOffloads := c.GetBool("tun.use_offloads", true)
|
||||
|
||||
var fd int
|
||||
var name string
|
||||
|
||||
@@ -218,7 +218,8 @@ func (t *winTun) addRoutes(logErrors bool) error {
|
||||
return t.setMTU(luid, foundDefault4, carriesV6)
|
||||
}
|
||||
|
||||
// setMTU applies tun.mtu per address family. The default route metric rides along on the v4 handle.
|
||||
// setMTU applies tun.mtu per address family. The default route metric rides along on the v4 handle, DAD and
|
||||
// router discovery come off with the v6 one.
|
||||
func (t *winTun) setMTU(luid winipcfg.LUID, foundDefault4, carriesV6 bool) error {
|
||||
ipif, err := luid.IPInterface(windows.AF_INET)
|
||||
if err != nil {
|
||||
@@ -249,6 +250,11 @@ func (t *winTun) setMTU(luid winipcfg.LUID, foundDefault4, carriesV6 bool) error
|
||||
}
|
||||
|
||||
ipif6.NLMTU = uint32(t.MTU)
|
||||
// Nothing answers on the far side of this adapter but nebula, which drops the probes. DAD only holds the
|
||||
// address tentative for a round it can never lose, and solicitations only invite RAs we would drop anyway.
|
||||
// wireguard-windows turns both off on the same handle.
|
||||
ipif6.DadTransmits = 0
|
||||
ipif6.RouterDiscoveryBehavior = winipcfg.RouterDiscoveryDisabled
|
||||
if err := ipif6.Set(); err != nil {
|
||||
return fmt.Errorf("failed to set ipv6 interface: %w", err)
|
||||
}
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"crypto/fips140"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -238,17 +237,10 @@ func (cs *CertState) getCertificate(v cert.Version) cert.Certificate {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 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) {
|
||||
func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string) (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 {
|
||||
@@ -261,9 +253,6 @@ func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string, fips140E
|
||||
}
|
||||
|
||||
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
|
||||
@@ -337,10 +326,6 @@ 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 {
|
||||
@@ -420,7 +405,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: %w", err)
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -428,7 +413,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, fips140.Enforced())
|
||||
ncs, err := newCipherSuite(v1.Curve(), pkcs11backed, cipher)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -445,7 +430,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: %w", err)
|
||||
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -453,7 +438,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, fips140.Enforced())
|
||||
ncs, err := newCipherSuite(v2.Curve(), pkcs11backed, cipher)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
-94
@@ -1,94 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,19 +1,62 @@
|
||||
package nebula
|
||||
|
||||
// Configuration and lifecycle for the ssh debug console. The commands it serves are not
|
||||
// defined here; see commands.go, which registers them for every transport.
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"runtime/pprof"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
)
|
||||
|
||||
type sshListHostMapFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
ByIndex bool
|
||||
}
|
||||
|
||||
type sshPrintCertFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
Raw bool
|
||||
}
|
||||
|
||||
type sshPrintTunnelFlags struct {
|
||||
Pretty bool
|
||||
}
|
||||
|
||||
type sshChangeRemoteFlags struct {
|
||||
Address string
|
||||
}
|
||||
|
||||
type sshCloseTunnelFlags struct {
|
||||
LocalOnly bool
|
||||
}
|
||||
|
||||
type sshCreateTunnelFlags struct {
|
||||
Address string
|
||||
}
|
||||
|
||||
type sshDeviceInfoFlags struct {
|
||||
Json bool
|
||||
Pretty bool
|
||||
}
|
||||
|
||||
func wireSSHReload(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) {
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if c.GetBool("sshd.enabled", false) {
|
||||
@@ -154,3 +197,862 @@ func configSSH(l *slog.Logger, ssh *sshd.SSHServer, c *config.C) (func(), error)
|
||||
|
||||
return runner, nil
|
||||
}
|
||||
|
||||
func attachCommands(l *slog.Logger, c *config.C, ssh *sshd.SSHServer, f *Interface) {
|
||||
// sandboxDir defaults to a dir in temp. The intention is that end user will
|
||||
// create this dir as needed. Overriding this config value to "" allows
|
||||
// writing to anywhere in the system.
|
||||
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
|
||||
sandboxDir := c.GetString("sshd.sandbox_dir", defaultDir)
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "list-hostmap",
|
||||
ShortDescription: "List all known previously connected hosts",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshListHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshListHostMap(f.hostMap, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "list-pending-hostmap",
|
||||
ShortDescription: "List all handshaking hosts",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshListHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.ByIndex, "by-index", false, "gets all hosts in the hostmap from the index table")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshListHostMap(f.handshakeManager, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "list-lighthouse-addrmap",
|
||||
ShortDescription: "List all lighthouse map entries",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshListHostMapFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshListLighthouseMap(f.lightHouse, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "reload",
|
||||
ShortDescription: "Reloads configuration from disk, same as sending HUP to the process",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshReload(c, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "start-cpu-profile",
|
||||
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshStartCpuProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "stop-cpu-profile",
|
||||
ShortDescription: "Stops a cpu profile and writes output to the previously provided file",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
pprof.StopCPUProfile()
|
||||
return w.WriteLine("If a CPU profile was running it is now stopped")
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "save-heap-profile",
|
||||
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshGetHeapProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "mutex-profile-fraction",
|
||||
ShortDescription: "Gets or sets runtime.SetMutexProfileFraction",
|
||||
Callback: sshMutexProfileFraction,
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "save-mutex-profile",
|
||||
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshGetMutexProfile(sandboxDir, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "log-level",
|
||||
ShortDescription: "Gets or sets the current log level",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshLogLevel(l, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "log-format",
|
||||
ShortDescription: "Gets or sets the current log format",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshLogFormat(l, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "version",
|
||||
ShortDescription: "Prints the currently running version of nebula",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshVersion(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "device-info",
|
||||
ShortDescription: "Prints information about the network device.",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshDeviceInfoFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json with more information")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshDeviceInfo(f, fs, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "print-cert",
|
||||
ShortDescription: "Prints the current certificate being used or the certificate for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshPrintCertFlags{}
|
||||
fl.BoolVar(&s.Json, "json", false, "outputs as json")
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json, assumes -json")
|
||||
fl.BoolVar(&s.Raw, "raw", false, "raw prints the PEM encoded certificate, not compatible with -json or -pretty")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshPrintCert(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "print-tunnel",
|
||||
ShortDescription: "Prints json details about a tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshPrintTunnelFlags{}
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshPrintTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "print-relays",
|
||||
ShortDescription: "Prints json details about all relay info",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshPrintTunnelFlags{}
|
||||
fl.BoolVar(&s.Pretty, "pretty", false, "pretty prints json")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshPrintRelays(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "change-remote",
|
||||
ShortDescription: "Changes the remote address used in the tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshChangeRemoteFlags{}
|
||||
fl.StringVar(&s.Address, "address", "", "The new remote address, ip:port")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshChangeRemote(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "close-tunnel",
|
||||
ShortDescription: "Closes a tunnel for the provided vpn addr",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshCloseTunnelFlags{}
|
||||
fl.BoolVar(&s.LocalOnly, "local-only", false, "Disables notifying the remote that the tunnel is shutting down")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshCloseTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "create-tunnel",
|
||||
ShortDescription: "Creates a tunnel for the provided vpn address",
|
||||
Help: "The lighthouses will be queried for real addresses but you can provide one as well.",
|
||||
Flags: func() (*flag.FlagSet, any) {
|
||||
fl := flag.NewFlagSet("", flag.ContinueOnError)
|
||||
s := sshCreateTunnelFlags{}
|
||||
fl.StringVar(&s.Address, "address", "", "Optionally provide a real remote address, ip:port ")
|
||||
return fl, &s
|
||||
},
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshCreateTunnel(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
|
||||
ssh.RegisterCommand(&sshd.Command{
|
||||
Name: "query-lighthouse",
|
||||
ShortDescription: "Query the lighthouses for the provided vpn address",
|
||||
Help: "This command is asynchronous. Only currently known udp addresses will be printed.",
|
||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
||||
return sshQueryLighthouse(f, fs, a, w)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func sshListHostMap(hl controlHostLister, a any, w sshd.StringWriter) error {
|
||||
fs, ok := a.(*sshListHostMapFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
var hm []ControlHostInfo
|
||||
if fs.ByIndex {
|
||||
hm = listHostMapIndexes(hl)
|
||||
} else {
|
||||
hm = listHostMapHosts(hl)
|
||||
}
|
||||
|
||||
sort.Slice(hm, func(i, j int) bool {
|
||||
return hm[i].VpnAddrs[0].Compare(hm[j].VpnAddrs[0]) < 0
|
||||
})
|
||||
|
||||
if fs.Json || fs.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if fs.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
err := js.Encode(hm)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
} else {
|
||||
for _, v := range hm {
|
||||
err := w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddrs, v.RemoteAddrs))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func sshListLighthouseMap(lightHouse *LightHouse, a any, w sshd.StringWriter) error {
|
||||
fs, ok := a.(*sshListHostMapFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
type lighthouseInfo struct {
|
||||
VpnAddr string `json:"vpnAddr"`
|
||||
Addrs *CacheMap `json:"addrs"`
|
||||
}
|
||||
|
||||
lightHouse.RLock()
|
||||
addrMap := make([]lighthouseInfo, len(lightHouse.addrMap))
|
||||
x := 0
|
||||
for k, v := range lightHouse.addrMap {
|
||||
addrMap[x] = lighthouseInfo{
|
||||
VpnAddr: k.String(),
|
||||
Addrs: v.CopyCache(),
|
||||
}
|
||||
x++
|
||||
}
|
||||
lightHouse.RUnlock()
|
||||
|
||||
sort.Slice(addrMap, func(i, j int) bool {
|
||||
return strings.Compare(addrMap[i].VpnAddr, addrMap[j].VpnAddr) < 0
|
||||
})
|
||||
|
||||
if fs.Json || fs.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if fs.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
err := js.Encode(addrMap)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
} else {
|
||||
for _, v := range addrMap {
|
||||
b, err := json.Marshal(v.Addrs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = w.WriteLine(fmt.Sprintf("%s: %s", v.VpnAddr, string(b)))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// sshSanitizeFilePath validates that the given file path is within the sandbox directory.
|
||||
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
|
||||
func sshSanitizeFilePath(sandboxDir, filePath string) (string, error) {
|
||||
if sandboxDir == "" {
|
||||
return filePath, nil
|
||||
}
|
||||
|
||||
// Clean and resolve the path relative to the sandbox directory
|
||||
if !filepath.IsAbs(filePath) {
|
||||
filePath = filepath.Join(sandboxDir, filePath)
|
||||
}
|
||||
cleaned := filepath.Clean(filePath)
|
||||
|
||||
// Ensure the resolved path is within the sandbox directory
|
||||
cleanedSandbox := filepath.Clean(sandboxDir)
|
||||
if cleaned == cleanedSandbox {
|
||||
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
|
||||
}
|
||||
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
|
||||
}
|
||||
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func sshStartCpuProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
err := w.WriteLine("No path to write profile provided")
|
||||
return err
|
||||
}
|
||||
|
||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = pprof.StartCPUProfile(file)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to start cpu profile: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = w.WriteLine(fmt.Sprintf("Started cpu profile, issue stop-cpu-profile to write the output to %s", a))
|
||||
return err
|
||||
}
|
||||
|
||||
func sshVersion(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
return w.WriteLine(fmt.Sprintf("%s", ifce.version))
|
||||
}
|
||||
|
||||
func sshQueryLighthouse(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
var cm *CacheMap
|
||||
rl := ifce.lightHouse.Query(vpnAddr)
|
||||
if rl != nil {
|
||||
cm = rl.CopyCache()
|
||||
}
|
||||
return json.NewEncoder(w.GetWriter()).Encode(cm)
|
||||
}
|
||||
|
||||
func sshCloseTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
flags, ok := fs.(*sshCloseTunnelFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||
}
|
||||
|
||||
if !flags.LocalOnly {
|
||||
ifce.send(
|
||||
header.CloseTunnel,
|
||||
0,
|
||||
hostInfo.ConnectionState,
|
||||
hostInfo,
|
||||
[]byte{},
|
||||
make([]byte, 12, 12),
|
||||
make([]byte, mtu),
|
||||
)
|
||||
}
|
||||
|
||||
ifce.closeTunnel(hostInfo)
|
||||
return w.WriteLine("Closed")
|
||||
}
|
||||
|
||||
func sshCreateTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
flags, ok := fs.(*sshCreateTunnelFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Tunnel already exists"))
|
||||
}
|
||||
|
||||
hostInfo = ifce.handshakeManager.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Tunnel already handshaking"))
|
||||
}
|
||||
|
||||
var addr netip.AddrPort
|
||||
if flags.Address != "" {
|
||||
addr, err = netip.ParseAddrPort(flags.Address)
|
||||
if err != nil {
|
||||
return w.WriteLine("Address could not be parsed")
|
||||
}
|
||||
}
|
||||
|
||||
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
|
||||
if addr.IsValid() {
|
||||
hostInfo.SetRemote(addr)
|
||||
}
|
||||
|
||||
return w.WriteLine("Created")
|
||||
}
|
||||
|
||||
func sshChangeRemote(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
flags, ok := fs.(*sshChangeRemoteFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
if flags.Address == "" {
|
||||
return w.WriteLine("No address was provided")
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddrPort(flags.Address)
|
||||
if err != nil {
|
||||
return w.WriteLine("Address could not be parsed")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn address could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn address: %v", a[0]))
|
||||
}
|
||||
|
||||
hostInfo.SetRemote(addr)
|
||||
return w.WriteLine("Changed")
|
||||
}
|
||||
|
||||
func sshGetHeapProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No path to write profile provided")
|
||||
}
|
||||
|
||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = pprof.WriteHeapProfile(file)
|
||||
if err != nil {
|
||||
err = w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||
return err
|
||||
}
|
||||
|
||||
err = w.WriteLine(fmt.Sprintf("Mem profile created at %s", a))
|
||||
return err
|
||||
}
|
||||
|
||||
func sshMutexProfileFraction(fs any, a []string, w sshd.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
rate := runtime.SetMutexProfileFraction(-1)
|
||||
return w.WriteLine(fmt.Sprintf("Current value: %d", rate))
|
||||
}
|
||||
|
||||
newRate, err := strconv.Atoi(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Invalid argument: %s", a[0]))
|
||||
}
|
||||
|
||||
oldRate := runtime.SetMutexProfileFraction(newRate)
|
||||
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
|
||||
}
|
||||
|
||||
func sshGetMutexProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No path to write profile provided")
|
||||
}
|
||||
|
||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(err.Error())
|
||||
}
|
||||
|
||||
file, err := os.Create(filePath)
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
mutexProfile := pprof.Lookup("mutex")
|
||||
if mutexProfile == nil {
|
||||
return w.WriteLine("Unable to get pprof.Lookup(\"mutex\")")
|
||||
}
|
||||
|
||||
err = mutexProfile.WriteTo(file, 0)
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unable to write profile: %s", err))
|
||||
}
|
||||
|
||||
return w.WriteLine(fmt.Sprintf("Mutex profile created at %s", a))
|
||||
}
|
||||
|
||||
func sshLogLevel(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error {
|
||||
ctrl, ok := l.Handler().(interface {
|
||||
GetLevel() slog.Level
|
||||
SetLevel(slog.Level)
|
||||
})
|
||||
if !ok {
|
||||
return w.WriteLine("Log level is not reconfigurable on this logger")
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||
}
|
||||
|
||||
level, err := logging.ParseLevel(strings.ToLower(a[0]))
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("Unknown log level %s. Possible log levels: trace, debug, info, warn, error", a))
|
||||
}
|
||||
|
||||
ctrl.SetLevel(level)
|
||||
return w.WriteLine(fmt.Sprintf("Log level is: %s", logging.LevelName(ctrl.GetLevel())))
|
||||
}
|
||||
|
||||
func sshLogFormat(l *slog.Logger, fs any, a []string, w sshd.StringWriter) error {
|
||||
ctrl, ok := l.Handler().(interface {
|
||||
GetFormat() string
|
||||
SetFormat(string) error
|
||||
})
|
||||
if !ok {
|
||||
return w.WriteLine("Log format is not reconfigurable on this logger")
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||
}
|
||||
|
||||
if err := ctrl.SetFormat(strings.ToLower(a[0])); err != nil {
|
||||
return err
|
||||
}
|
||||
return w.WriteLine(fmt.Sprintf("Log format is: %s", ctrl.GetFormat()))
|
||||
}
|
||||
|
||||
func sshPrintCert(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
args, ok := fs.(*sshPrintCertFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
cert := ifce.pki.getCertState().GetDefaultCertificate()
|
||||
if len(a) > 0 {
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||
}
|
||||
|
||||
cert = hostInfo.GetCert().Certificate
|
||||
}
|
||||
|
||||
if args.Json || args.Pretty {
|
||||
b, err := cert.MarshalJSON()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if args.Pretty {
|
||||
buf := new(bytes.Buffer)
|
||||
err := json.Indent(buf, b, "", " ")
|
||||
b = buf.Bytes()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return w.WriteBytes(b)
|
||||
}
|
||||
|
||||
if args.Raw {
|
||||
b, err := cert.MarshalPEM()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return w.WriteBytes(b)
|
||||
}
|
||||
|
||||
return w.WriteLine(cert.String())
|
||||
}
|
||||
|
||||
func sshPrintRelays(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
args, ok := fs.(*sshPrintTunnelFlags)
|
||||
if !ok {
|
||||
w.WriteLine(fmt.Sprintf("sshPrintRelays failed to convert args type"))
|
||||
return nil
|
||||
}
|
||||
|
||||
relays := map[uint32]*HostInfo{}
|
||||
ifce.hostMap.Lock()
|
||||
maps.Copy(relays, ifce.hostMap.Relays)
|
||||
ifce.hostMap.Unlock()
|
||||
|
||||
type RelayFor struct {
|
||||
Error error
|
||||
Type string
|
||||
State string
|
||||
PeerAddr netip.Addr
|
||||
LocalIndex uint32
|
||||
RemoteIndex uint32
|
||||
RelayedThrough []netip.Addr
|
||||
}
|
||||
|
||||
type RelayOutput struct {
|
||||
NebulaAddr netip.Addr
|
||||
RelayForAddrs []RelayFor
|
||||
}
|
||||
|
||||
type CmdOutput struct {
|
||||
Relays []*RelayOutput
|
||||
}
|
||||
|
||||
co := CmdOutput{}
|
||||
|
||||
enc := json.NewEncoder(w.GetWriter())
|
||||
|
||||
if args.Pretty {
|
||||
enc.SetIndent("", " ")
|
||||
}
|
||||
|
||||
for k, v := range relays {
|
||||
ro := RelayOutput{NebulaAddr: v.vpnAddrs[0]}
|
||||
co.Relays = append(co.Relays, &ro)
|
||||
relayHI := ifce.hostMap.QueryVpnAddr(v.vpnAddrs[0])
|
||||
if relayHI == nil {
|
||||
ro.RelayForAddrs = append(ro.RelayForAddrs, RelayFor{Error: errors.New("could not find hostinfo")})
|
||||
continue
|
||||
}
|
||||
for _, vpnAddr := range relayHI.relayState.CopyRelayForIps() {
|
||||
rf := RelayFor{Error: nil}
|
||||
r, ok := relayHI.relayState.GetRelayForByAddr(vpnAddr)
|
||||
if ok {
|
||||
t := ""
|
||||
switch r.Type {
|
||||
case ForwardingType:
|
||||
t = "forwarding"
|
||||
case TerminalType:
|
||||
t = "terminal"
|
||||
default:
|
||||
t = "unknown"
|
||||
}
|
||||
|
||||
s := ""
|
||||
switch r.State {
|
||||
case Requested:
|
||||
s = "requested"
|
||||
case Established:
|
||||
s = "established"
|
||||
default:
|
||||
s = "unknown"
|
||||
}
|
||||
|
||||
rf.LocalIndex = r.LocalIndex
|
||||
rf.RemoteIndex = r.RemoteIndex
|
||||
rf.PeerAddr = r.PeerAddr
|
||||
rf.Type = t
|
||||
rf.State = s
|
||||
if rf.LocalIndex != k {
|
||||
rf.Error = fmt.Errorf("hostmap LocalIndex '%v' does not match RelayState LocalIndex", k)
|
||||
}
|
||||
}
|
||||
relayedHI := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if relayedHI != nil {
|
||||
rf.RelayedThrough = append(rf.RelayedThrough, relayedHI.relayState.CopyRelayIps()...)
|
||||
}
|
||||
|
||||
ro.RelayForAddrs = append(ro.RelayForAddrs, rf)
|
||||
}
|
||||
}
|
||||
err := enc.Encode(co)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func sshPrintTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) error {
|
||||
args, ok := fs.(*sshPrintTunnelFlags)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(a) == 0 {
|
||||
return w.WriteLine("No vpn address was provided")
|
||||
}
|
||||
|
||||
vpnAddr, err := netip.ParseAddr(a[0])
|
||||
if err != nil {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
if !vpnAddr.IsValid() {
|
||||
return w.WriteLine(fmt.Sprintf("The provided vpn addr could not be parsed: %s", a[0]))
|
||||
}
|
||||
|
||||
hostInfo := ifce.hostMap.QueryVpnAddr(vpnAddr)
|
||||
if hostInfo == nil {
|
||||
return w.WriteLine(fmt.Sprintf("Could not find tunnel for vpn addr: %v", a[0]))
|
||||
}
|
||||
|
||||
enc := json.NewEncoder(w.GetWriter())
|
||||
if args.Pretty {
|
||||
enc.SetIndent("", " ")
|
||||
}
|
||||
|
||||
return enc.Encode(copyHostInfo(hostInfo, ifce.hostMap.GetPreferredRanges()))
|
||||
}
|
||||
|
||||
func sshDeviceInfo(ifce *Interface, fs any, w sshd.StringWriter) error {
|
||||
|
||||
data := struct {
|
||||
Name string `json:"name"`
|
||||
Cidr []netip.Prefix `json:"cidr"`
|
||||
}{
|
||||
Name: ifce.inside.Name(),
|
||||
Cidr: make([]netip.Prefix, len(ifce.inside.Networks())),
|
||||
}
|
||||
|
||||
copy(data.Cidr, ifce.inside.Networks())
|
||||
|
||||
flags, ok := fs.(*sshDeviceInfoFlags)
|
||||
if !ok {
|
||||
return fmt.Errorf("internal error: expected flags to be sshDeviceInfoFlags but was %+v", fs)
|
||||
}
|
||||
|
||||
if flags.Json || flags.Pretty {
|
||||
js := json.NewEncoder(w.GetWriter())
|
||||
if flags.Pretty {
|
||||
js.SetIndent("", " ")
|
||||
}
|
||||
|
||||
return js.Encode(data)
|
||||
} else {
|
||||
return w.WriteLine(fmt.Sprintf("name=%v cidr=%v", data.Name, data.Cidr))
|
||||
}
|
||||
}
|
||||
|
||||
func sshReload(c *config.C, w sshd.StringWriter) error {
|
||||
err := w.WriteLine("Reloading config")
|
||||
c.ReloadConfig()
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
package diag
|
||||
package sshd
|
||||
|
||||
import (
|
||||
"errors"
|
||||
@@ -10,17 +10,6 @@ import (
|
||||
"github.com/armon/go-radix"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrUnknownCommand is returned by the Registry when the first argument names no
|
||||
// registered command. The user has already been told so on their writer.
|
||||
ErrUnknownCommand = errors.New("unknown command")
|
||||
|
||||
// ErrUsage wraps a flag parsing failure. The flag package has already written the
|
||||
// details to the caller's writer by the time this is returned, so a transport should
|
||||
// use it only to pick an exit status.
|
||||
ErrUsage = errors.New("usage")
|
||||
)
|
||||
|
||||
// CommandFlags is a function called before help or command execution to parse command line flags
|
||||
// It should return a flag.FlagSet instance and a pointer to the struct that will contain parsed flags
|
||||
type CommandFlags func() (*flag.FlagSet, any)
|
||||
@@ -55,10 +44,8 @@ func execCommand(c *Command, args []string, w StringWriter) error {
|
||||
fl.SetOutput(w.GetWriter())
|
||||
err := fl.Parse(args)
|
||||
if err != nil {
|
||||
// fl.Parse has dumped error information to the user via the w writer, so
|
||||
// the wrapper exists purely so a transport can tell a usage problem from a
|
||||
// command that ran and failed.
|
||||
return fmt.Errorf("%w: %w", ErrUsage, err)
|
||||
// fl.Parse has dumped error information to the user via the w writer.
|
||||
return err
|
||||
}
|
||||
args = fl.Args()
|
||||
}
|
||||
+20
-7
@@ -9,9 +9,8 @@ import (
|
||||
"net"
|
||||
"sync"
|
||||
|
||||
"github.com/armon/go-radix"
|
||||
"golang.org/x/crypto/ssh"
|
||||
|
||||
"github.com/slackhq/nebula/diag"
|
||||
)
|
||||
|
||||
type SSHServer struct {
|
||||
@@ -26,9 +25,10 @@ type SSHServer struct {
|
||||
trustedKeys map[string]map[string]bool
|
||||
trustedCAs []ssh.PublicKey
|
||||
|
||||
// The commands this server serves. Shared with every other transport, see diag.Registry.
|
||||
commands *diag.Registry
|
||||
listener net.Listener
|
||||
// List of available commands
|
||||
helpCommand *Command
|
||||
commands *radix.Tree
|
||||
listener net.Listener
|
||||
|
||||
// ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even
|
||||
// across reloads, since each Run derives a fresh child rather than reusing this one directly.
|
||||
@@ -38,11 +38,11 @@ type SSHServer struct {
|
||||
// NewSSHServer creates a new ssh server rigged with default commands and prepares to listen.
|
||||
// The ssh server's context is parented off the supplied ctx so cancelling it
|
||||
// (e.g. on Control.Stop) tears down active sessions and closes the listener.
|
||||
func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry) (*SSHServer, error) {
|
||||
func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
|
||||
s := &SSHServer{
|
||||
trustedKeys: make(map[string]map[string]bool),
|
||||
l: l,
|
||||
commands: commands,
|
||||
commands: radix.New(),
|
||||
ctx: ctx,
|
||||
}
|
||||
|
||||
@@ -90,6 +90,14 @@ func NewSSHServer(ctx context.Context, l *slog.Logger, commands *diag.Registry)
|
||||
ServerVersion: fmt.Sprintf("SSH-2.0-Nebula???"),
|
||||
}
|
||||
|
||||
s.RegisterCommand(&Command{
|
||||
Name: "help",
|
||||
ShortDescription: "prints available commands or help <command> for specific usage info",
|
||||
Callback: func(a any, args []string, w StringWriter) error {
|
||||
return helpCallback(s.commands, args, w)
|
||||
},
|
||||
})
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
@@ -152,6 +160,11 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegisterCommand adds a command that can be run by a user, by default only `help` is available
|
||||
func (s *SSHServer) RegisterCommand(c *Command) {
|
||||
s.commands.Insert(c.Name, c)
|
||||
}
|
||||
|
||||
// Run begins listening and accepting connections. Each invocation derives a fresh per-Run context
|
||||
// from the constructor-supplied ctx so a Stop+Run sequence (used by config reload) starts clean
|
||||
// rather than carrying a permanently-cancelled context across runs.
|
||||
|
||||
+45
-18
@@ -1,38 +1,37 @@
|
||||
package sshd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/anmitsu/go-shlex"
|
||||
"github.com/armon/go-radix"
|
||||
"golang.org/x/crypto/ssh"
|
||||
"golang.org/x/term"
|
||||
|
||||
"github.com/slackhq/nebula/diag"
|
||||
)
|
||||
|
||||
type session struct {
|
||||
l *slog.Logger
|
||||
c *ssh.ServerConn
|
||||
term *term.Terminal
|
||||
commands *diag.Registry
|
||||
commands *radix.Tree
|
||||
cancel func()
|
||||
}
|
||||
|
||||
func NewSession(commands *diag.Registry, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
|
||||
func NewSession(commands *radix.Tree, conn *ssh.ServerConn, chans <-chan ssh.NewChannel, cancel func(), l *slog.Logger) *session {
|
||||
s := &session{
|
||||
// A copy, so the logout command this session adds for itself stays invisible to every
|
||||
// other session and to `nebula ctl`.
|
||||
commands: commands.Clone(),
|
||||
commands: radix.NewFromMap(commands.ToMap()),
|
||||
l: l,
|
||||
c: conn,
|
||||
cancel: cancel,
|
||||
}
|
||||
|
||||
s.commands.RegisterCommand(&diag.Command{
|
||||
s.commands.Insert("logout", &Command{
|
||||
Name: "logout",
|
||||
ShortDescription: "Ends the current session",
|
||||
Callback: func(a any, args []string, w diag.StringWriter) error {
|
||||
Callback: func(a any, args []string, w StringWriter) error {
|
||||
s.Close()
|
||||
return nil
|
||||
},
|
||||
@@ -88,11 +87,9 @@ func (s *session) handleRequests(in <-chan *ssh.Request, channel ssh.Channel) {
|
||||
}
|
||||
|
||||
req.Reply(true, nil)
|
||||
dErr := s.commands.Dispatch(payload.Value, diag.NewWriter(channel))
|
||||
s.dispatchCommand(payload.Value, &stringWriter{channel})
|
||||
|
||||
// Report a real exit status rather than a hardcoded zero, so that
|
||||
// `ssh nebula-host list-hostmap` is scriptable the same way `nebula ctl` is.
|
||||
status := struct{ Status uint32 }{uint32(diag.StatusFor(dErr))}
|
||||
status := struct{ Status uint32 }{uint32(0)}
|
||||
channel.SendRequest("exit-status", false, ssh.Marshal(status))
|
||||
channel.Close()
|
||||
return
|
||||
@@ -114,7 +111,7 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
|
||||
term.AutoCompleteCallback = func(line string, pos int, key rune) (newLine string, newPos int, ok bool) {
|
||||
// key 9 is tab
|
||||
if key == 9 {
|
||||
cmds := s.commands.Match(line)
|
||||
cmds := matchCommand(s.commands, line)
|
||||
if len(cmds) == 1 {
|
||||
return cmds[0] + " ", len(cmds[0]) + 1, true
|
||||
}
|
||||
@@ -131,19 +128,49 @@ func (s *session) createTerm(channel ssh.Channel) *term.Terminal {
|
||||
}
|
||||
|
||||
func (s *session) handleInput() {
|
||||
w := diag.NewWriter(s.term)
|
||||
w := &stringWriter{w: s.term}
|
||||
for {
|
||||
line, err := s.term.ReadLine()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
|
||||
// The interactive console reports problems on the terminal the user is already
|
||||
// looking at, so the error is nothing extra to say here.
|
||||
_ = s.commands.Dispatch(line, w)
|
||||
s.dispatchCommand(line, w)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *session) dispatchCommand(line string, w StringWriter) {
|
||||
args, err := shlex.Split(line, true)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if len(args) == 0 {
|
||||
dumpCommands(s.commands, w)
|
||||
return
|
||||
}
|
||||
|
||||
c, err := lookupCommand(s.commands, args[0])
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if c == nil {
|
||||
err := w.WriteLine(fmt.Sprintf("did not understand: %s", line))
|
||||
_ = err
|
||||
|
||||
dumpCommands(s.commands, w)
|
||||
return
|
||||
}
|
||||
|
||||
if checkHelpArgs(args) {
|
||||
s.dispatchCommand(fmt.Sprintf("%s %s", "help", c.Name), w)
|
||||
return
|
||||
}
|
||||
|
||||
_ = execCommand(c, args[1:], w)
|
||||
}
|
||||
|
||||
func (s *session) Close() {
|
||||
s.c.Close()
|
||||
s.cancel()
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
package diag
|
||||
package sshd
|
||||
|
||||
import "io"
|
||||
|
||||
@@ -30,9 +30,3 @@ func (w *stringWriter) WriteBytes(b []byte) error {
|
||||
func (w *stringWriter) GetWriter() io.Writer {
|
||||
return w.w
|
||||
}
|
||||
|
||||
// NewWriter adapts an io.Writer to the StringWriter commands are handed. Transports
|
||||
// implement their own framing behind w; the commands never know the difference.
|
||||
func NewWriter(w io.Writer) StringWriter {
|
||||
return &stringWriter{w: w}
|
||||
}
|
||||
@@ -2,7 +2,6 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/fips140"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -288,12 +287,9 @@ 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()),
|
||||
"fips140Version": fips140.Version(),
|
||||
"fips140Enabled": strconv.FormatBool(fips140.Enabled()),
|
||||
"fips140Enforced": strconv.FormatBool(fips140.Enforced()),
|
||||
"version": s.buildVersion,
|
||||
"goversion": runtime.Version(),
|
||||
"boringcrypto": strconv.FormatBool(boringEnabled()),
|
||||
},
|
||||
})
|
||||
pr.MustRegister(g)
|
||||
|
||||
Reference in New Issue
Block a user