mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 23:16:59 +02:00
Compare commits
15 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 60e07370de | |||
| 44530cb610 | |||
| c5abf30102 | |||
| 3a329ec217 | |||
| afcdf2163b | |||
| 9a30c5b6a1 | |||
| ad0c99f262 | |||
| c1ad7a3af2 | |||
| cef69465db | |||
| daef13e53a | |||
| 6e1df81473 | |||
| c61de54ec3 | |||
| 1b59636028 | |||
| 487bae4c2f | |||
| 2bdd284993 |
@@ -0,0 +1,34 @@
|
|||||||
|
name: gofmt
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- master
|
||||||
|
pull_request:
|
||||||
|
paths:
|
||||||
|
- '.github/workflows/gofmt.yml'
|
||||||
|
- '**.go'
|
||||||
|
jobs:
|
||||||
|
|
||||||
|
gofmt:
|
||||||
|
name: Run gofmt
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
|
||||||
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
|
- uses: actions/setup-go@v6
|
||||||
|
with:
|
||||||
|
go-version: '1.25'
|
||||||
|
check-latest: true
|
||||||
|
|
||||||
|
- name: Install goimports
|
||||||
|
run: |
|
||||||
|
go install golang.org/x/tools/cmd/goimports@latest
|
||||||
|
|
||||||
|
- name: gofmt
|
||||||
|
run: |
|
||||||
|
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
|
||||||
|
then
|
||||||
|
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
@@ -10,7 +10,7 @@ jobs:
|
|||||||
name: Build Linux/BSD All
|
name: Build Linux/BSD All
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -36,7 +36,7 @@ jobs:
|
|||||||
id-token: write
|
id-token: write
|
||||||
contents: read
|
contents: read
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -76,7 +76,7 @@ jobs:
|
|||||||
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
||||||
runs-on: macos-latest
|
runs-on: macos-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -134,7 +134,7 @@ jobs:
|
|||||||
# be overwritten
|
# be overwritten
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: actions/checkout@v7
|
uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
@@ -163,17 +163,14 @@ jobs:
|
|||||||
mkdir -p build/linux-{amd64,arm64}
|
mkdir -p build/linux-{amd64,arm64}
|
||||||
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
|
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
|
||||||
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
|
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
|
||||||
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 \
|
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}"
|
||||||
--build-arg VERSION="${GITHUB_REF#refs/tags/v}" \
|
|
||||||
--build-arg REVISION="${GITHUB_SHA}" \
|
|
||||||
--tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}"
|
|
||||||
|
|
||||||
release:
|
release:
|
||||||
name: Create and Upload Release
|
name: Create and Upload Release
|
||||||
needs: [build-linux, build-darwin, build-windows]
|
needs: [build-linux, build-darwin, build-windows]
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v8
|
uses: actions/download-artifact@v8
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ jobs:
|
|||||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
VAGRANT_DEFAULT_PROVIDER: libvirt
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -62,7 +62,7 @@ jobs:
|
|||||||
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -88,7 +88,7 @@ jobs:
|
|||||||
runs-on: windows-latest
|
runs-on: windows-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -36,14 +36,6 @@ jobs:
|
|||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./smoke.sh
|
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
|
- name: setup relay docker image
|
||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
run: ./build-relay.sh
|
run: ./build-relay.sh
|
||||||
|
|||||||
@@ -5,19 +5,6 @@ set -e -x
|
|||||||
rm -rf ./build
|
rm -rf ./build
|
||||||
mkdir ./build
|
mkdir ./build
|
||||||
|
|
||||||
if [ "$SMOKE_OVERLAY_IPV6" ]
|
|
||||||
then
|
|
||||||
LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401"
|
|
||||||
HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402"
|
|
||||||
HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403"
|
|
||||||
HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404"
|
|
||||||
else
|
|
||||||
LIGHTHOUSE_NIP="192.168.100.1"
|
|
||||||
HOST2_NIP="192.168.100.2"
|
|
||||||
HOST3_NIP="192.168.100.3"
|
|
||||||
HOST4_NIP="192.168.100.4"
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
||||||
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
||||||
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
||||||
@@ -44,24 +31,24 @@ LIGHTHOUSE_IP="203.0.113.2"
|
|||||||
../genconfig.sh >lighthouse1.yml
|
../genconfig.sh >lighthouse1.yml
|
||||||
|
|
||||||
HOST="host2" \
|
HOST="host2" \
|
||||||
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||||
../genconfig.sh >host2.yml
|
../genconfig.sh >host2.yml
|
||||||
|
|
||||||
HOST="host3" \
|
HOST="host3" \
|
||||||
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||||
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host3.yml
|
../genconfig.sh >host3.yml
|
||||||
|
|
||||||
HOST="host4" \
|
HOST="host4" \
|
||||||
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||||
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host4.yml
|
../genconfig.sh >host4.yml
|
||||||
|
|
||||||
../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test"
|
../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test"
|
||||||
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "$LIGHTHOUSE_NIP/24"
|
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "192.168.100.1/24"
|
||||||
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "$HOST2_NIP/24"
|
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "192.168.100.2/24"
|
||||||
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "$HOST3_NIP/24"
|
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "192.168.100.3/24"
|
||||||
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "$HOST4_NIP/24"
|
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
||||||
)
|
)
|
||||||
|
|
||||||
docker build -t "nebula:${NAME:-smoke}" .
|
docker build -t "nebula:${NAME:-smoke}" .
|
||||||
|
|||||||
@@ -47,19 +47,6 @@ HOST2_IP="$PREFIX.3"
|
|||||||
HOST3_IP="$PREFIX.4"
|
HOST3_IP="$PREFIX.4"
|
||||||
HOST4_IP="$PREFIX.5"
|
HOST4_IP="$PREFIX.5"
|
||||||
|
|
||||||
if [ "$SMOKE_OVERLAY_IPV6" ]
|
|
||||||
then
|
|
||||||
LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401"
|
|
||||||
HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402"
|
|
||||||
HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403"
|
|
||||||
HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404"
|
|
||||||
else
|
|
||||||
LIGHTHOUSE_NIP="192.168.100.1"
|
|
||||||
HOST2_NIP="192.168.100.2"
|
|
||||||
HOST3_NIP="192.168.100.3"
|
|
||||||
HOST4_NIP="192.168.100.4"
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
||||||
# build/lighthouse1.yml has no IPs to rewrite so it's skipped.
|
# build/lighthouse1.yml has no IPs to rewrite so it's skipped.
|
||||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
||||||
@@ -93,28 +80,28 @@ docker exec host3 tcpdump -i eth0 -q -w - -U 2>logs/host3.outside.log >logs/host
|
|||||||
docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
||||||
docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
||||||
|
|
||||||
docker exec host2 ncat -nklv 2000 &
|
docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host3 ncat -nklv 2000 &
|
docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 4000 &
|
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 0.0.0.0 4000 &
|
||||||
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 3000 &
|
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
||||||
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 3000 &
|
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from lighthouse1"
|
echo " *** Testing ping from lighthouse1"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec lighthouse1 ping -c1 $HOST2_NIP
|
docker exec lighthouse1 ping -c1 192.168.100.2
|
||||||
docker exec lighthouse1 ping -c1 $HOST3_NIP
|
docker exec lighthouse1 ping -c1 192.168.100.3
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host2"
|
echo " *** Testing ping from host2"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host2 ping -c1 $LIGHTHOUSE_NIP
|
docker exec host2 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ping -c1 $HOST3_NIP -w5 || exit 1
|
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -122,34 +109,34 @@ echo " *** Testing ncat from host2"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1
|
! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
||||||
! docker exec host2 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host3"
|
echo " *** Testing ping from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ping -c1 $LIGHTHOUSE_NIP
|
docker exec host3 ping -c1 192.168.100.1
|
||||||
docker exec host3 ping -c1 $HOST2_NIP
|
docker exec host3 ping -c1 192.168.100.2
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ncat from host3"
|
echo " *** Testing ncat from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ncat -nzv -w5 $HOST2_NIP 2000
|
docker exec host3 ncat -nzv -w5 192.168.100.2 2000
|
||||||
docker exec host3 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2
|
docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host4"
|
echo " *** Testing ping from host4"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host4 ping -c1 $LIGHTHOUSE_NIP
|
docker exec host4 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ping -c1 $HOST2_NIP -w5 || exit 1
|
! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
||||||
! docker exec host4 ping -c1 $HOST3_NIP -w5 || exit 1
|
! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -157,10 +144,10 @@ echo " *** Testing ncat from host4"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ncat -nzv -w5 $HOST2_NIP 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1
|
! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 || exit 1
|
! docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1
|
! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -172,7 +159,7 @@ set -x
|
|||||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
||||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
||||||
# the echo back from host4 never reaches host2.
|
# the echo back from host4 never reaches host2.
|
||||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4
|
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
|
||||||
|
|
||||||
docker exec host4 sh -c 'kill 1'
|
docker exec host4 sh -c 'kill 1'
|
||||||
docker exec host3 sh -c 'kill 1'
|
docker exec host3 sh -c 'kill 1'
|
||||||
|
|||||||
+72
-90
@@ -13,28 +13,20 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
static:
|
test-linux:
|
||||||
name: Static checks
|
name: Build all and test on ubuntu-linux
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Install goimports
|
- name: Build
|
||||||
run: go install golang.org/x/tools/cmd/goimports@latest
|
run: make all
|
||||||
|
|
||||||
- name: gofmt
|
|
||||||
run: |
|
|
||||||
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
|
|
||||||
then
|
|
||||||
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
|
|
||||||
- name: Vet
|
- name: Vet
|
||||||
run: make vet
|
run: make vet
|
||||||
@@ -44,41 +36,27 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
version: v2.5
|
version: v2.5
|
||||||
|
|
||||||
test:
|
- name: Test
|
||||||
name: Test ${{ matrix.name }}
|
run: make test
|
||||||
runs-on: ${{ matrix.os }}
|
|
||||||
strategy:
|
- name: End 2 end
|
||||||
fail-fast: false
|
run: make e2evv
|
||||||
matrix:
|
|
||||||
include:
|
- name: Build test mobile
|
||||||
- name: linux
|
run: make build-test-mobile
|
||||||
os: ubuntu-latest
|
|
||||||
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
- uses: actions/upload-artifact@v7
|
||||||
test-cmd: make test
|
with:
|
||||||
e2e-cmd: make e2evv
|
name: e2e packet flow linux-latest
|
||||||
- name: linux-boringcrypto
|
path: e2e/mermaid/linux-latest
|
||||||
os: ubuntu-latest
|
if-no-files-found: warn
|
||||||
build-cmd: make bin-boringcrypto
|
|
||||||
test-cmd: make test-boringcrypto
|
test-linux-boringcrypto:
|
||||||
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
name: Build and test on linux with boringcrypto
|
||||||
- name: linux-pkcs11
|
runs-on: ubuntu-latest
|
||||||
os: ubuntu-latest
|
|
||||||
build-cmd: make bin-pkcs11
|
|
||||||
test-cmd: make test-pkcs11
|
|
||||||
e2e-cmd: ''
|
|
||||||
- name: macos
|
|
||||||
os: macos-latest
|
|
||||||
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
|
||||||
test-cmd: make test
|
|
||||||
e2e-cmd: make e2evv
|
|
||||||
- name: windows
|
|
||||||
os: windows-latest
|
|
||||||
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
|
|
||||||
test-cmd: make test
|
|
||||||
e2e-cmd: make e2evv
|
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
@@ -86,65 +64,69 @@ jobs:
|
|||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
run: ${{ matrix.build-cmd }}
|
run: make bin-boringcrypto
|
||||||
|
|
||||||
- name: Cross-build darwin-amd64
|
|
||||||
if: matrix.name == 'macos'
|
|
||||||
run: GOARCH=amd64 go build -o /tmp/nebula-amd64 ./cmd/nebula && GOARCH=amd64 go build -o /tmp/nebula-cert-amd64 ./cmd/nebula-cert
|
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: ${{ matrix.test-cmd }}
|
run: make test-boringcrypto
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
if: matrix.e2e-cmd != ''
|
run: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
||||||
run: ${{ matrix.e2e-cmd }}
|
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v7
|
test-linux-pkcs11:
|
||||||
if: matrix.e2e-cmd != '' && always()
|
name: Build and test on linux with pkcs11
|
||||||
with:
|
|
||||||
name: e2e packet flow ${{ matrix.name }}
|
|
||||||
path: e2e/mermaid/
|
|
||||||
if-no-files-found: warn
|
|
||||||
|
|
||||||
cross-build:
|
|
||||||
name: Cross-build ${{ matrix.name }}
|
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
|
||||||
fail-fast: false
|
|
||||||
matrix:
|
|
||||||
include:
|
|
||||||
- {name: linux-arm, make-target: all-cross-linux-arm}
|
|
||||||
- {name: linux-mips, make-target: all-cross-linux-mips}
|
|
||||||
- {name: linux-other, make-target: all-cross-linux-other}
|
|
||||||
- {name: freebsd, make-target: all-freebsd}
|
|
||||||
- {name: openbsd, make-target: all-openbsd}
|
|
||||||
- {name: netbsd, make-target: all-netbsd}
|
|
||||||
- {name: windows, make-target: all-cross-windows}
|
|
||||||
- {name: mobile, make-target: build-test-mobile}
|
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v7
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version: '1.25'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build ${{ matrix.name }}
|
- name: Build
|
||||||
run: make -j"$(nproc)" ${{ matrix.make-target }}
|
run: make bin-pkcs11
|
||||||
|
|
||||||
finish:
|
- name: Test
|
||||||
name: CI status
|
run: make test-pkcs11
|
||||||
if: always()
|
|
||||||
needs: [static, test, cross-build]
|
test:
|
||||||
runs-on: ubuntu-latest
|
name: Build and test on ${{ matrix.os }}
|
||||||
|
runs-on: ${{ matrix.os }}
|
||||||
|
strategy:
|
||||||
|
matrix:
|
||||||
|
os: [windows-latest, macos-latest]
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- name: Fail if any upstream job failed
|
- uses: actions/checkout@v6
|
||||||
if: contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled')
|
|
||||||
run: |
|
|
||||||
echo "upstream results: ${{ toJSON(needs) }}"
|
|
||||||
exit 1
|
|
||||||
|
|
||||||
- name: All upstream jobs passed
|
- uses: actions/setup-go@v6
|
||||||
run: echo "ok"
|
with:
|
||||||
|
go-version: '1.25'
|
||||||
|
check-latest: true
|
||||||
|
|
||||||
|
- name: Build nebula
|
||||||
|
run: go build ./cmd/nebula
|
||||||
|
|
||||||
|
- name: Build nebula-cert
|
||||||
|
run: go build ./cmd/nebula-cert
|
||||||
|
|
||||||
|
- name: Vet
|
||||||
|
run: make vet
|
||||||
|
|
||||||
|
- name: golangci-lint
|
||||||
|
uses: golangci/golangci-lint-action@v9
|
||||||
|
with:
|
||||||
|
version: v2.5
|
||||||
|
|
||||||
|
- name: Test
|
||||||
|
run: make test
|
||||||
|
|
||||||
|
- name: End 2 end
|
||||||
|
run: make e2evv
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v7
|
||||||
|
with:
|
||||||
|
name: e2e packet flow ${{ matrix.os }}
|
||||||
|
path: e2e/mermaid/${{ matrix.os }}
|
||||||
|
if-no-files-found: warn
|
||||||
|
|||||||
@@ -60,18 +60,6 @@ ALL = $(ALL_LINUX) \
|
|||||||
windows-amd64 \
|
windows-amd64 \
|
||||||
windows-arm64
|
windows-arm64
|
||||||
|
|
||||||
# Cross-build shards used by .github/workflows/test.yml — same as ALL_*
|
|
||||||
# but with the arch that has a native CI runner removed, so the cross-build
|
|
||||||
# job is not duplicating coverage the native test jobs already give.
|
|
||||||
ALL_CROSS_LINUX = $(filter-out linux-amd64,$(ALL_LINUX))
|
|
||||||
|
|
||||||
# ALL_CROSS_LINUX further split into family sub-shards so each can run on
|
|
||||||
# its own CI runner in parallel. Union of the three must equal
|
|
||||||
# ALL_CROSS_LINUX; adding a new linux arch goes into the matching family.
|
|
||||||
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
|
|
||||||
|
|
||||||
e2e:
|
e2e:
|
||||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||||
|
|
||||||
@@ -94,35 +82,6 @@ DOCKER_BIN = build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
|||||||
|
|
||||||
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
||||||
|
|
||||||
all-linux: $(ALL_LINUX:%=build/%/nebula) $(ALL_LINUX:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-freebsd: $(ALL_FREEBSD:%=build/%/nebula) $(ALL_FREEBSD:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-openbsd: $(ALL_OPENBSD:%=build/%/nebula) $(ALL_OPENBSD:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-netbsd: $(ALL_NETBSD:%=build/%/nebula) $(ALL_NETBSD:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert build/darwin-arm64/nebula build/darwin-arm64/nebula-cert
|
|
||||||
|
|
||||||
all-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
|
|
||||||
|
|
||||||
# CI cross-build shards. darwin-arm64 is covered by the native macos-latest
|
|
||||||
# job; windows-amd64 is covered by the native windows-latest job; both are
|
|
||||||
# omitted here to avoid building them a second time. darwin-amd64 stays in
|
|
||||||
# all-cross-darwin because intel mac is only a labeled/master-time native
|
|
||||||
# job, so PRs still need cross-build coverage for it.
|
|
||||||
all-cross-linux: $(ALL_CROSS_LINUX:%=build/%/nebula) $(ALL_CROSS_LINUX:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-cross-linux-arm: $(ALL_CROSS_LINUX_ARM:%=build/%/nebula) $(ALL_CROSS_LINUX_ARM:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-cross-linux-mips: $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula) $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-cross-linux-other: $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula) $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula-cert)
|
|
||||||
|
|
||||||
all-cross-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert
|
|
||||||
|
|
||||||
all-cross-windows: build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
|
|
||||||
|
|
||||||
docker: docker/linux-$(shell go env GOARCH)
|
docker: docker/linux-$(shell go env GOARCH)
|
||||||
|
|
||||||
release: $(ALL:%=build/nebula-%.tar.gz)
|
release: $(ALL:%=build/nebula-%.tar.gz)
|
||||||
@@ -268,9 +227,6 @@ smoke-relay-docker: bin-docker
|
|||||||
cd .github/workflows/smoke/ && ./build-relay.sh
|
cd .github/workflows/smoke/ && ./build-relay.sh
|
||||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
|
|
||||||
smoke-docker-ipv6: smoke-docker
|
|
||||||
|
|
||||||
smoke-docker-race: BUILD_ARGS = -race
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
smoke-docker-race: CGO_ENABLED = 1
|
smoke-docker-race: CGO_ENABLED = 1
|
||||||
smoke-docker-race: smoke-docker
|
smoke-docker-race: smoke-docker
|
||||||
@@ -280,5 +236,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
|
|||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
.PHONY: bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
+4
-10
@@ -13,12 +13,6 @@ import (
|
|||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
// testCertNow is the reference "now" used to derive default before/after times
|
|
||||||
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
|
|
||||||
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
|
|
||||||
// signed with default times can never expire after its CA on a rounding race.
|
|
||||||
var testCertNow = time.Now().Round(time.Second)
|
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
// NewTestCaCert will create a new ca certificate
|
||||||
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
||||||
var err error
|
var err error
|
||||||
@@ -40,10 +34,10 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &TBSCertificate{
|
t := &TBSCertificate{
|
||||||
@@ -76,11 +70,11 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
|
|||||||
// Expiry times are defaulted if you do not pass them in
|
// Expiry times are defaulted if you do not pass them in
|
||||||
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(networks) == 0 {
|
if len(networks) == 0 {
|
||||||
|
|||||||
+4
-10
@@ -14,12 +14,6 @@ import (
|
|||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
// testCertNow is the reference "now" used to derive default before/after times
|
|
||||||
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
|
|
||||||
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
|
|
||||||
// signed with default times can never expire after its CA on a rounding race.
|
|
||||||
var testCertNow = time.Now().Round(time.Second)
|
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
// NewTestCaCert will create a new ca certificate
|
||||||
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
||||||
var err error
|
var err error
|
||||||
@@ -41,10 +35,10 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
|
|||||||
}
|
}
|
||||||
|
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
t := &cert.TBSCertificate{
|
||||||
@@ -77,11 +71,11 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
|
|||||||
// Expiry times are defaulted if you do not pass them in
|
// Expiry times are defaulted if you do not pass them in
|
||||||
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
||||||
if before.IsZero() {
|
if before.IsZero() {
|
||||||
before = testCertNow.Add(time.Second * -60)
|
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
if after.IsZero() {
|
if after.IsZero() {
|
||||||
after = testCertNow.Add(time.Second * 60)
|
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
||||||
}
|
}
|
||||||
|
|
||||||
var pub, priv []byte
|
var pub, priv []byte
|
||||||
|
|||||||
+7
-32
@@ -97,19 +97,6 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
// out-key is meaningless under PKCS#11 because the private key never
|
|
||||||
// leaves the HSM; reject it so we never silently accept or claim a
|
|
||||||
// stdout slot for it.
|
|
||||||
outKeySet := false
|
|
||||||
cf.set.Visit(func(f *flag.Flag) {
|
|
||||||
if f.Name == "out-key" {
|
|
||||||
outKeySet = true
|
|
||||||
}
|
|
||||||
})
|
|
||||||
if outKeySet {
|
|
||||||
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -184,21 +171,12 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
|
||||||
if err := reserveOutputs(&claims,
|
|
||||||
"out-key", *cf.outKeyPath,
|
|
||||||
"out-crt", *cf.outCertPath,
|
|
||||||
"out-qr", *cf.outQRPath,
|
|
||||||
); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
var passphrase []byte
|
var passphrase []byte
|
||||||
if !isP11 && *cf.encryption {
|
if !isP11 && *cf.encryption {
|
||||||
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
||||||
if len(passphrase) == 0 {
|
if len(passphrase) == 0 {
|
||||||
for i := 0; i < 5; i++ {
|
for i := 0; i < 5; i++ {
|
||||||
errOut.Write([]byte("Enter passphrase: "))
|
out.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if err == ErrNoTerminal {
|
if err == ErrNoTerminal {
|
||||||
@@ -283,16 +261,14 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
Curve: curve,
|
Curve: curve,
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 && !isStdio(*cf.outKeyPath) {
|
if !isP11 {
|
||||||
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isStdio(*cf.outCertPath) {
|
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
||||||
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
||||||
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var c cert.Certificate
|
var c cert.Certificate
|
||||||
@@ -318,7 +294,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
|
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*cf.outKeyPath, b, 0600, out)
|
err = os.WriteFile(*cf.outKeyPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -329,7 +305,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return fmt.Errorf("error while marshalling certificate: %s", err)
|
return fmt.Errorf("error while marshalling certificate: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*cf.outCertPath, b, 0600, out)
|
err = os.WriteFile(*cf.outCertPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -340,7 +316,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*cf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*cf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -356,7 +332,6 @@ func caSummary() string {
|
|||||||
func caHelp(out io.Writer) {
|
func caHelp(out io.Writer) {
|
||||||
cf := newCaFlags()
|
cf := newCaFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n"))
|
||||||
out.Write([]byte(stdioHelpText))
|
|
||||||
cf.set.SetOutput(out)
|
cf.set.SetOutput(out)
|
||||||
cf.set.PrintDefaults()
|
cf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ func Test_caHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
|
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
|
||||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
|
||||||
" -argon-iterations uint\n"+
|
" -argon-iterations uint\n"+
|
||||||
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
|
||||||
" -argon-memory uint\n"+
|
" -argon-memory uint\n"+
|
||||||
@@ -85,7 +84,7 @@ func Test_ca(t *testing.T) {
|
|||||||
err: nil,
|
err: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
pwPromptEB := "Enter passphrase: "
|
pwPromptOb := "Enter passphrase: "
|
||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, ca(
|
assertHelpError(t, ca(
|
||||||
@@ -169,8 +168,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, ca(args, ob, eb, testpw))
|
require.NoError(t, ca(args, ob, eb, testpw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Equal(t, pwPromptEB, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test encrypted key with passphrase environment variable
|
// test encrypted key with passphrase environment variable
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -208,8 +207,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.Error(t, ca(args, ob, eb, errpw))
|
require.Error(t, ca(args, ob, eb, errpw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Equal(t, pwPromptEB, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test when user fails to enter a password
|
// test when user fails to enter a password
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -218,8 +217,8 @@ func Test_ca(t *testing.T) {
|
|||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
|
||||||
assert.Equal(t, strings.Repeat(pwPromptEB, 5), eb.String()) // prompts 5 times before giving up
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
@@ -248,67 +247,3 @@ func Test_ca(t *testing.T) {
|
|||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_ca_stdio(t *testing.T) {
|
|
||||||
nopw := &StubPasswordReader{}
|
|
||||||
|
|
||||||
keyF, err := os.CreateTemp("", "ca.key")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
defer os.Remove(keyF.Name())
|
|
||||||
|
|
||||||
crtF, err := os.CreateTemp("", "ca.crt")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
defer os.Remove(crtF.Name())
|
|
||||||
|
|
||||||
// out-crt on stdout, out-key on disk
|
|
||||||
ob := &bytes.Buffer{}
|
|
||||||
eb := &bytes.Buffer{}
|
|
||||||
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", keyF.Name()}, ob, eb, nopw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
c, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.True(t, c.IsCA())
|
|
||||||
assert.Equal(t, "test-ca", c.Name())
|
|
||||||
|
|
||||||
// out-key on stdout, out-crt on disk
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", crtF.Name(), "-out-key", "-"}, ob, eb, nopw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
_, _, curve, err := cert.UnmarshalSigningPrivateKeyFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
|
||||||
|
|
||||||
// dual stdout is rejected up front
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.EqualError(t,
|
|
||||||
ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", "-"}, ob, eb, nopw),
|
|
||||||
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
|
|
||||||
// an output conflict combined with -encrypt must error BEFORE prompting
|
|
||||||
// for a passphrase; pr would record any read attempt
|
|
||||||
tracker := &trackingPasswordReader{}
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.EqualError(t,
|
|
||||||
ca([]string{"-name", "test-ca", "-duration", "1h", "-encrypt", "-out-crt", "-", "-out-key", "-"}, ob, eb, tracker),
|
|
||||||
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
assert.Zero(t, tracker.calls, "passphrase prompt should not have been called")
|
|
||||||
}
|
|
||||||
|
|
||||||
type trackingPasswordReader struct {
|
|
||||||
calls int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (pr *trackingPasswordReader) ReadPassword() ([]byte, error) {
|
|
||||||
pr.calls++
|
|
||||||
return []byte(""), nil
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -42,8 +42,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
} else if *cf.outKeyPath != "" {
|
|
||||||
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
|
||||||
}
|
}
|
||||||
if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
|
if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -71,14 +69,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
|
||||||
if err := reserveOutputs(&claims,
|
|
||||||
"out-key", *cf.outKeyPath,
|
|
||||||
"out-pub", *cf.outPubPath,
|
|
||||||
); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if isP11 {
|
if isP11 {
|
||||||
p11Client, err := pkclient.FromUrl(*cf.p11url)
|
p11Client, err := pkclient.FromUrl(*cf.p11url)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -92,12 +82,12 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while getting public key: %w", err)
|
return fmt.Errorf("error while getting public key: %w", err)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
err = writeOutput(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
err = os.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
err = writeOutput(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600, out)
|
err = os.WriteFile(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-pub: %s", err)
|
return fmt.Errorf("error while writing out-pub: %s", err)
|
||||||
}
|
}
|
||||||
@@ -112,7 +102,6 @@ func keygenSummary() string {
|
|||||||
func keygenHelp(out io.Writer) {
|
func keygenHelp(out io.Writer) {
|
||||||
cf := newKeygenFlags()
|
cf := newKeygenFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
||||||
_, _ = out.Write([]byte(stdioHelpText))
|
|
||||||
cf.set.SetOutput(out)
|
cf.set.SetOutput(out)
|
||||||
cf.set.PrintDefaults()
|
cf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ func Test_keygenHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+
|
"Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+
|
||||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
|
||||||
" -curve string\n"+
|
" -curve string\n"+
|
||||||
" \tECDH Curve (25519, P256) (default \"25519\")\n"+
|
" \tECDH Curve (25519, P256) (default \"25519\")\n"+
|
||||||
" -out-key string\n"+
|
" -out-key string\n"+
|
||||||
@@ -94,43 +93,3 @@ func Test_keygen(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Len(t, lPub, 32)
|
assert.Len(t, lPub, 32)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_keygen_stdio(t *testing.T) {
|
|
||||||
keyF, err := os.CreateTemp("", "test.key")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
defer os.Remove(keyF.Name())
|
|
||||||
|
|
||||||
pubF, err := os.CreateTemp("", "test.pub")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(pubF.Name())
|
|
||||||
defer os.Remove(pubF.Name())
|
|
||||||
|
|
||||||
// out-pub on stdout, out-key on disk
|
|
||||||
ob := &bytes.Buffer{}
|
|
||||||
eb := &bytes.Buffer{}
|
|
||||||
require.NoError(t, keygen([]string{"-out-pub", "-", "-out-key", keyF.Name()}, ob, eb))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
lPub, _, curve, err := cert.UnmarshalPublicKeyFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
|
||||||
assert.Len(t, lPub, 32)
|
|
||||||
|
|
||||||
// out-key on stdout, out-pub on disk
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.NoError(t, keygen([]string{"-out-pub", pubF.Name(), "-out-key", "-"}, ob, eb))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
lKey, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
|
||||||
assert.Len(t, lKey, 32)
|
|
||||||
|
|
||||||
// both on stdout is a conflict caught up front
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.EqualError(t, keygen([]string{"-out-pub", "-", "-out-key", "-"}, ob, eb),
|
|
||||||
`-out-key and -out-pub both set to "-", only one output may write to stdout`)
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -22,9 +22,7 @@ func (pr StdinPasswordReader) ReadPassword() ([]byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
password, err := term.ReadPassword(int(os.Stdin.Fd()))
|
||||||
// Terminal echo is off while reading, so the user's Enter key does not
|
fmt.Println()
|
||||||
// produce a visible newline. Emit one on stderr to match the prompt.
|
|
||||||
fmt.Fprintln(os.Stderr)
|
|
||||||
|
|
||||||
return password, err
|
return password, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -40,23 +40,11 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
rawCert, err := os.ReadFile(*pf.path)
|
||||||
if err := reserveInputs(&claims, "path", *pf.path); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if err := reserveOutputs(&claims, "out-qr", *pf.outQRPath); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
rawCert, err := readInput("path", *pf.path, &claims)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read cert; %s", err)
|
return fmt.Errorf("unable to read cert; %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// When the QR is going to stdout, suppress the human-readable text/json
|
|
||||||
// output so the binary stream is not contaminated.
|
|
||||||
qrToStdout := isStdio(*pf.outQRPath)
|
|
||||||
|
|
||||||
var c cert.Certificate
|
var c cert.Certificate
|
||||||
var qrBytes []byte
|
var qrBytes []byte
|
||||||
part := 0
|
part := 0
|
||||||
@@ -69,13 +57,11 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !qrToStdout {
|
if *pf.json {
|
||||||
if *pf.json {
|
jsonCerts = append(jsonCerts, c)
|
||||||
jsonCerts = append(jsonCerts, c)
|
} else {
|
||||||
} else {
|
_, _ = out.Write([]byte(c.String()))
|
||||||
_, _ = out.Write([]byte(c.String()))
|
_, _ = out.Write([]byte("\n"))
|
||||||
_, _ = out.Write([]byte("\n"))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
@@ -93,7 +79,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
part++
|
part++
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json && !qrToStdout {
|
if *pf.json {
|
||||||
b, _ := json.Marshal(jsonCerts)
|
b, _ := json.Marshal(jsonCerts)
|
||||||
_, _ = out.Write(b)
|
_, _ = out.Write(b)
|
||||||
_, _ = out.Write([]byte("\n"))
|
_, _ = out.Write([]byte("\n"))
|
||||||
@@ -105,7 +91,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*pf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*pf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -121,7 +107,6 @@ func printSummary() string {
|
|||||||
func printHelp(out io.Writer) {
|
func printHelp(out io.Writer) {
|
||||||
pf := newPrintFlags()
|
pf := newPrintFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n"))
|
||||||
out.Write([]byte(stdioHelpText))
|
|
||||||
pf.set.SetOutput(out)
|
pf.set.SetOutput(out)
|
||||||
pf.set.PrintDefaults()
|
pf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -25,7 +25,6 @@ func Test_printHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+
|
"Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+
|
||||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
|
||||||
" -json\n"+
|
" -json\n"+
|
||||||
" \tOptional: outputs certificates in json format\n"+
|
" \tOptional: outputs certificates in json format\n"+
|
||||||
" -out-qr string\n"+
|
" -out-qr string\n"+
|
||||||
@@ -179,44 +178,6 @@ func Test_printCert(t *testing.T) {
|
|||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// read cert from stdin
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
withStdin(t, bytes.NewReader(p))
|
|
||||||
err = printCert([]string{"-json", "-path", "-"}, ob, eb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(
|
|
||||||
t,
|
|
||||||
`[{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1}]
|
|
||||||
`,
|
|
||||||
ob.String(),
|
|
||||||
)
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
|
|
||||||
// -out-qr - sends only the PNG to stdout, suppressing the cert dump
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
withStdin(t, bytes.NewReader(p))
|
|
||||||
err = printCert([]string{"-path", "-", "-out-qr", "-"}, ob, eb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
stdout := ob.Bytes()
|
|
||||||
require.NotEmpty(t, stdout)
|
|
||||||
// PNG magic, no PEM/JSON noise prepended
|
|
||||||
assert.Equal(t, []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a}, stdout[:8])
|
|
||||||
assert.NotContains(t, string(stdout), "NebulaCertificate")
|
|
||||||
assert.NotContains(t, string(stdout), `"details"`)
|
|
||||||
|
|
||||||
// json + out-qr - still suppresses json
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
withStdin(t, bytes.NewReader(p))
|
|
||||||
err = printCert([]string{"-json", "-path", "-", "-out-qr", "-"}, ob, eb)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
assert.Equal(t, []byte{0x89, 'P', 'N', 'G'}, ob.Bytes()[:4])
|
|
||||||
assert.NotContains(t, ob.String(), `"details"`)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewTestCaCert will generate a CA cert
|
// NewTestCaCert will generate a CA cert
|
||||||
|
|||||||
+20
-42
@@ -85,9 +85,6 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
||||||
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
||||||
}
|
}
|
||||||
if isP11 && *sf.outKeyPath != "" {
|
|
||||||
return newHelpErrorf("cannot set -out-key with -pkcs11")
|
|
||||||
}
|
|
||||||
|
|
||||||
var v4Networks []netip.Prefix
|
var v4Networks []netip.Prefix
|
||||||
var v6Networks []netip.Prefix
|
var v6Networks []netip.Prefix
|
||||||
@@ -105,35 +102,13 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
||||||
}
|
}
|
||||||
|
|
||||||
if *sf.outKeyPath == "" {
|
|
||||||
*sf.outKeyPath = *sf.name + ".key"
|
|
||||||
}
|
|
||||||
if *sf.outCertPath == "" {
|
|
||||||
*sf.outCertPath = *sf.name + ".crt"
|
|
||||||
}
|
|
||||||
|
|
||||||
var claims ioClaims
|
|
||||||
if err := reserveInputs(&claims,
|
|
||||||
"ca-key", *sf.caKeyPath,
|
|
||||||
"ca-crt", *sf.caCertPath,
|
|
||||||
"in-pub", *sf.inPubPath,
|
|
||||||
); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if err := reserveOutputs(&claims,
|
|
||||||
"out-key", *sf.outKeyPath,
|
|
||||||
"out-crt", *sf.outCertPath,
|
|
||||||
"out-qr", *sf.outQRPath,
|
|
||||||
); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
var caKey []byte
|
var caKey []byte
|
||||||
|
|
||||||
if !isP11 {
|
if !isP11 {
|
||||||
var rawCAKey []byte
|
var rawCAKey []byte
|
||||||
rawCAKey, err = readInput("ca-key", *sf.caKeyPath, &claims)
|
rawCAKey, err := os.ReadFile(*sf.caKeyPath)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-key: %s", err)
|
return fmt.Errorf("error while reading ca-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -146,7 +121,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
if len(passphrase) == 0 {
|
if len(passphrase) == 0 {
|
||||||
// ask for a passphrase until we get one
|
// ask for a passphrase until we get one
|
||||||
for i := 0; i < 5; i++ {
|
for i := 0; i < 5; i++ {
|
||||||
errOut.Write([]byte("Enter passphrase: "))
|
out.Write([]byte("Enter passphrase: "))
|
||||||
passphrase, err = pr.ReadPassword()
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
if errors.Is(err, ErrNoTerminal) {
|
if errors.Is(err, ErrNoTerminal) {
|
||||||
@@ -172,7 +147,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCACert, err := readInput("ca-crt", *sf.caCertPath, &claims)
|
rawCACert, err := os.ReadFile(*sf.caCertPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-crt: %s", err)
|
return fmt.Errorf("error while reading ca-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -270,7 +245,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
|
|
||||||
if *sf.inPubPath != "" {
|
if *sf.inPubPath != "" {
|
||||||
var pubCurve cert.Curve
|
var pubCurve cert.Curve
|
||||||
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims)
|
rawPub, err := os.ReadFile(*sf.inPubPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading in-pub: %s", err)
|
return fmt.Errorf("error while reading in-pub: %s", err)
|
||||||
}
|
}
|
||||||
@@ -291,10 +266,16 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
pub, rawPriv = newKeypair(curve)
|
pub, rawPriv = newKeypair(curve)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isStdio(*sf.outCertPath) {
|
if *sf.outKeyPath == "" {
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
}
|
||||||
}
|
|
||||||
|
if *sf.outCertPath == "" {
|
||||||
|
*sf.outCertPath = *sf.name + ".crt"
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
|
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
var crts []cert.Certificate
|
var crts []cert.Certificate
|
||||||
@@ -379,13 +360,11 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 && *sf.inPubPath == "" {
|
if !isP11 && *sf.inPubPath == "" {
|
||||||
if !isStdio(*sf.outKeyPath) {
|
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
||||||
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
||||||
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out)
|
err = os.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
@@ -400,7 +379,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
b = append(b, sb...)
|
b = append(b, sb...)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*sf.outCertPath, b, 0600, out)
|
err = os.WriteFile(*sf.outCertPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -411,7 +390,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writeOutput(*sf.outQRPath, b, 0600, out)
|
err = os.WriteFile(*sf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
@@ -461,7 +440,6 @@ func signSummary() string {
|
|||||||
func signHelp(out io.Writer) {
|
func signHelp(out io.Writer) {
|
||||||
sf := newSignFlags()
|
sf := newSignFlags()
|
||||||
out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n"))
|
||||||
out.Write([]byte(stdioHelpText))
|
|
||||||
sf.set.SetOutput(out)
|
sf.set.SetOutput(out)
|
||||||
sf.set.PrintDefaults()
|
sf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ func Test_signHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+
|
"Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+
|
||||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
|
||||||
" -ca-crt string\n"+
|
" -ca-crt string\n"+
|
||||||
" \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+
|
" \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+
|
||||||
" -ca-key string\n"+
|
" -ca-key string\n"+
|
||||||
@@ -377,18 +376,15 @@ func Test_signCert(t *testing.T) {
|
|||||||
// test with the proper password
|
// test with the proper password
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
require.NoError(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test with the proper password in the environment
|
// test with the proper password in the environment
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
require.NoError(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
||||||
|
|
||||||
@@ -399,8 +395,8 @@ func Test_signCert(t *testing.T) {
|
|||||||
testpw.password = []byte("invalid password")
|
testpw.password = []byte("invalid password")
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, testpw))
|
require.Error(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test with the wrong password in environment
|
// test with the wrong password in environment
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -420,8 +416,8 @@ func Test_signCert(t *testing.T) {
|
|||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, nopw))
|
require.Error(t, signCert(args, ob, eb, nopw))
|
||||||
// normally the user hitting enter on the prompt would add newlines between these
|
// normally the user hitting enter on the prompt would add newlines between these
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test an error condition
|
// test an error condition
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
@@ -429,106 +425,6 @@ func Test_signCert(t *testing.T) {
|
|||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, errpw))
|
require.Error(t, signCert(args, ob, eb, errpw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Equal(t, "Enter passphrase: ", eb.String())
|
assert.Empty(t, eb.String())
|
||||||
}
|
|
||||||
|
|
||||||
func Test_signCert_stdio(t *testing.T) {
|
|
||||||
nopw := &StubPasswordReader{
|
|
||||||
password: []byte(""),
|
|
||||||
err: nil,
|
|
||||||
}
|
|
||||||
|
|
||||||
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
|
||||||
rawCAKey := cert.MarshalSigningPrivateKeyToPEM(cert.Curve_CURVE25519, caPriv)
|
|
||||||
|
|
||||||
ca, _ := NewTestCaCert("ca", caPub, caPriv, time.Now(), time.Now().Add(time.Minute*200), nil, nil, nil)
|
|
||||||
rawCACrt, _ := ca.MarshalPEM()
|
|
||||||
|
|
||||||
caCrtF, err := os.CreateTemp("", "sign-cert.crt")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer os.Remove(caCrtF.Name())
|
|
||||||
caCrtF.Write(rawCACrt)
|
|
||||||
|
|
||||||
caKeyF, err := os.CreateTemp("", "sign-cert.key")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer os.Remove(caKeyF.Name())
|
|
||||||
caKeyF.Write(rawCAKey)
|
|
||||||
|
|
||||||
keyF, err := os.CreateTemp("", "sign.key")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
defer os.Remove(keyF.Name())
|
|
||||||
|
|
||||||
// ca-key on stdin, cert to stdout
|
|
||||||
withStdin(t, bytes.NewReader(rawCAKey))
|
|
||||||
ob := &bytes.Buffer{}
|
|
||||||
eb := &bytes.Buffer{}
|
|
||||||
args := []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", keyF.Name(), "-duration", "100m"}
|
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
|
|
||||||
lCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, "stdin-test", lCrt.Name())
|
|
||||||
assert.True(t, lCrt.CheckSignature(caPub))
|
|
||||||
|
|
||||||
// two flags reading from stdin should error before any read attempt;
|
|
||||||
// otherwise an interactive shell would hang on io.ReadAll
|
|
||||||
stdinIn := bytes.NewReader(rawCAKey)
|
|
||||||
withStdin(t, stdinIn)
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-ca-crt", "-", "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw),
|
|
||||||
`-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
|
||||||
assert.Equal(t, len(rawCAKey), stdinIn.Len(), "stdin should be untouched when conflict is caught up front")
|
|
||||||
|
|
||||||
// two flags writing to stdout should error before any output is written
|
|
||||||
// AND before stdin is consumed
|
|
||||||
stdinR := bytes.NewReader(rawCAKey)
|
|
||||||
withStdin(t, stdinR)
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", "-", "-duration", "100m"}
|
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw),
|
|
||||||
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
// stdin should be untouched because the conflict was caught up front
|
|
||||||
assert.Equal(t, len(rawCAKey), stdinR.Len())
|
|
||||||
|
|
||||||
// out-key on stdout, cert on disk
|
|
||||||
keyF2, err := os.CreateTemp("", "sign.key")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(keyF2.Name())
|
|
||||||
defer os.Remove(keyF2.Name())
|
|
||||||
crtF, err := os.CreateTemp("", "sign.crt")
|
|
||||||
require.NoError(t, err)
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
defer os.Remove(crtF.Name())
|
|
||||||
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", "-", "-duration", "100m"}
|
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
_, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
|
||||||
|
|
||||||
// in-pub on stdin (caller already has a keypair, only the cert is generated)
|
|
||||||
inPub, _ := x25519Keypair()
|
|
||||||
rawInPub := cert.MarshalPublicKeyToPEM(cert.Curve_CURVE25519, inPub)
|
|
||||||
|
|
||||||
withStdin(t, bytes.NewReader(rawInPub))
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "in-pub-test", "-ip", "1.1.1.1/24", "-in-pub", "-", "-out-crt", "-", "-duration", "100m"}
|
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
stdinCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, "in-pub-test", stdinCrt.Name())
|
|
||||||
assert.Equal(t, inPub, stdinCrt.PublicKey())
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,117 +0,0 @@
|
|||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
)
|
|
||||||
|
|
||||||
// stdioPath is the special path value that selects stdin (for inputs) or
|
|
||||||
// stdout (for outputs) instead of a file on disk.
|
|
||||||
const stdioPath = "-"
|
|
||||||
|
|
||||||
// stdioHelpText is rendered just under the Usage line of each subcommand
|
|
||||||
// help so the - convention is documented once instead of on every flag.
|
|
||||||
const stdioHelpText = " Pass \"-\" to any path flag to read from stdin or write to stdout.\n"
|
|
||||||
|
|
||||||
// stdinReader is the source used when an input flag is set to "-".
|
|
||||||
// It is a package level var so tests can swap in a deterministic reader.
|
|
||||||
// Tests that mutate stdinReader cannot run with t.Parallel().
|
|
||||||
var stdinReader io.Reader = os.Stdin
|
|
||||||
|
|
||||||
// ioClaims tracks which flags have claimed stdin and stdout during a single
|
|
||||||
// command invocation so we can refuse a second flag asking for the same
|
|
||||||
// stream.
|
|
||||||
type ioClaims struct {
|
|
||||||
in string
|
|
||||||
out string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *ioClaims) claimIn(flagName string) error {
|
|
||||||
if c.in != "" && c.in != flagName {
|
|
||||||
return fmt.Errorf("-%s and -%s both set to %q, only one input may read from stdin", c.in, flagName, stdioPath)
|
|
||||||
}
|
|
||||||
c.in = flagName
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *ioClaims) claimOut(flagName string) error {
|
|
||||||
if c.out != "" && c.out != flagName {
|
|
||||||
return fmt.Errorf("-%s and -%s both set to %q, only one output may write to stdout", c.out, flagName, stdioPath)
|
|
||||||
}
|
|
||||||
c.out = flagName
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// reserveInputs walks alternating (flagName, path) pairs and claims stdin
|
|
||||||
// for any path equal to stdioPath. It must be called before any input is
|
|
||||||
// read so a conflict can be reported immediately instead of blocking on
|
|
||||||
// io.ReadAll while waiting for input that will never arrive.
|
|
||||||
func reserveInputs(claims *ioClaims, pairs ...string) error {
|
|
||||||
return reserveStdio(claims, "reserveInputs", (*ioClaims).claimIn, pairs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// reserveOutputs walks alternating (flagName, path) pairs and claims stdout
|
|
||||||
// for any path equal to stdioPath. It must be called before any output is
|
|
||||||
// written so a conflict cannot leave one stream half written before the
|
|
||||||
// second flag fails.
|
|
||||||
func reserveOutputs(claims *ioClaims, pairs ...string) error {
|
|
||||||
return reserveStdio(claims, "reserveOutputs", (*ioClaims).claimOut, pairs)
|
|
||||||
}
|
|
||||||
|
|
||||||
func reserveStdio(claims *ioClaims, who string, claim func(*ioClaims, string) error, pairs []string) error {
|
|
||||||
if len(pairs)%2 != 0 {
|
|
||||||
panic(who + " requires alternating name, path pairs")
|
|
||||||
}
|
|
||||||
for i := 0; i < len(pairs); i += 2 {
|
|
||||||
name, path := pairs[i], pairs[i+1]
|
|
||||||
if path != stdioPath {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if err := claim(claims, name); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// readInput returns the bytes referenced by path, reading from stdin when
|
|
||||||
// path is stdioPath.
|
|
||||||
func readInput(flagName, path string, claims *ioClaims) ([]byte, error) {
|
|
||||||
if path == stdioPath {
|
|
||||||
if err := claims.claimIn(flagName); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return io.ReadAll(stdinReader)
|
|
||||||
}
|
|
||||||
return os.ReadFile(path)
|
|
||||||
}
|
|
||||||
|
|
||||||
// openInput returns a reader for path. When path is stdioPath the returned
|
|
||||||
// reader wraps stdin and Close is a no-op.
|
|
||||||
func openInput(flagName, path string, claims *ioClaims) (io.ReadCloser, error) {
|
|
||||||
if path == stdioPath {
|
|
||||||
if err := claims.claimIn(flagName); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return io.NopCloser(stdinReader), nil
|
|
||||||
}
|
|
||||||
return os.Open(path)
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeOutput writes data to path, or to stdout when path is stdioPath. perm
|
|
||||||
// is only used for file output. The caller must have already claimed stdout
|
|
||||||
// via reserveOutputs before invoking with stdioPath.
|
|
||||||
func writeOutput(path string, data []byte, perm os.FileMode, stdout io.Writer) error {
|
|
||||||
if path == stdioPath {
|
|
||||||
_, err := stdout.Write(data)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return os.WriteFile(path, data, perm)
|
|
||||||
}
|
|
||||||
|
|
||||||
// isStdio reports whether path is the stdio sentinel and so should skip
|
|
||||||
// existence checks like "refuse to overwrite".
|
|
||||||
func isStdio(path string) bool {
|
|
||||||
return path == stdioPath
|
|
||||||
}
|
|
||||||
@@ -1,167 +0,0 @@
|
|||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// withStdin temporarily replaces stdinReader for the duration of t.
|
|
||||||
func withStdin(t *testing.T, r io.Reader) {
|
|
||||||
t.Helper()
|
|
||||||
prev := stdinReader
|
|
||||||
stdinReader = r
|
|
||||||
t.Cleanup(func() { stdinReader = prev })
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_readInput_stdin(t *testing.T) {
|
|
||||||
withStdin(t, bytes.NewBufferString("hello"))
|
|
||||||
var claims ioClaims
|
|
||||||
|
|
||||||
got, err := readInput("path", "-", &claims)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("hello"), got)
|
|
||||||
assert.Equal(t, "path", claims.in)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_readInput_file(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
p := filepath.Join(dir, "f")
|
|
||||||
require.NoError(t, os.WriteFile(p, []byte("file"), 0600))
|
|
||||||
var claims ioClaims
|
|
||||||
|
|
||||||
got, err := readInput("path", p, &claims)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("file"), got)
|
|
||||||
assert.Empty(t, claims.in)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_readInput_doubleStdinErrors(t *testing.T) {
|
|
||||||
withStdin(t, bytes.NewBufferString("hello"))
|
|
||||||
var claims ioClaims
|
|
||||||
|
|
||||||
_, err := readInput("ca-key", "-", &claims)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = readInput("ca-crt", "-", &claims)
|
|
||||||
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_openInput_stdin(t *testing.T) {
|
|
||||||
withStdin(t, bytes.NewBufferString("hi"))
|
|
||||||
var claims ioClaims
|
|
||||||
|
|
||||||
r, err := openInput("ca", "-", &claims)
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer r.Close()
|
|
||||||
b, err := io.ReadAll(r)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("hi"), b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_openInput_doubleStdinErrors(t *testing.T) {
|
|
||||||
withStdin(t, bytes.NewBufferString("hi"))
|
|
||||||
var claims ioClaims
|
|
||||||
|
|
||||||
r, err := openInput("ca", "-", &claims)
|
|
||||||
require.NoError(t, err)
|
|
||||||
r.Close()
|
|
||||||
|
|
||||||
_, err = openInput("crt", "-", &claims)
|
|
||||||
require.EqualError(t, err, `-ca and -crt both set to "-", only one input may read from stdin`)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_writeOutput_stdout(t *testing.T) {
|
|
||||||
out := &bytes.Buffer{}
|
|
||||||
|
|
||||||
err := writeOutput("-", []byte("payload"), 0600, out)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, "payload", out.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_writeOutput_file(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
p := filepath.Join(dir, "f")
|
|
||||||
out := &bytes.Buffer{}
|
|
||||||
|
|
||||||
err := writeOutput(p, []byte("payload"), 0600, out)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, out.String())
|
|
||||||
got, err := os.ReadFile(p)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, []byte("payload"), got)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_reserveOutputs_noConflict(t *testing.T) {
|
|
||||||
var claims ioClaims
|
|
||||||
require.NoError(t, reserveOutputs(&claims,
|
|
||||||
"out-key", "/tmp/key",
|
|
||||||
"out-crt", "-",
|
|
||||||
"out-qr", "",
|
|
||||||
))
|
|
||||||
assert.Equal(t, "out-crt", claims.out)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_reserveOutputs_conflict(t *testing.T) {
|
|
||||||
var claims ioClaims
|
|
||||||
err := reserveOutputs(&claims,
|
|
||||||
"out-key", "-",
|
|
||||||
"out-crt", "-",
|
|
||||||
)
|
|
||||||
require.EqualError(t, err, `-out-key and -out-crt both set to "-", only one output may write to stdout`)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_reserveOutputs_panicsOnOddPairs(t *testing.T) {
|
|
||||||
defer func() {
|
|
||||||
r := recover()
|
|
||||||
require.NotNil(t, r)
|
|
||||||
}()
|
|
||||||
var claims ioClaims
|
|
||||||
_ = reserveOutputs(&claims, "out-key")
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_reserveInputs_noConflict(t *testing.T) {
|
|
||||||
var claims ioClaims
|
|
||||||
require.NoError(t, reserveInputs(&claims,
|
|
||||||
"ca-key", "/tmp/ca.key",
|
|
||||||
"ca-crt", "-",
|
|
||||||
"in-pub", "",
|
|
||||||
))
|
|
||||||
assert.Equal(t, "ca-crt", claims.in)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_reserveInputs_conflict(t *testing.T) {
|
|
||||||
var claims ioClaims
|
|
||||||
err := reserveInputs(&claims,
|
|
||||||
"ca-key", "-",
|
|
||||||
"ca-crt", "-",
|
|
||||||
)
|
|
||||||
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_claimIn_idempotent(t *testing.T) {
|
|
||||||
// pre-claim then a lazy re-claim of the same flag should be a no-op
|
|
||||||
var claims ioClaims
|
|
||||||
require.NoError(t, claims.claimIn("ca-key"))
|
|
||||||
require.NoError(t, claims.claimIn("ca-key"))
|
|
||||||
assert.Equal(t, "ca-key", claims.in)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_claimOut_idempotent(t *testing.T) {
|
|
||||||
var claims ioClaims
|
|
||||||
require.NoError(t, claims.claimOut("out-crt"))
|
|
||||||
require.NoError(t, claims.claimOut("out-crt"))
|
|
||||||
assert.Equal(t, "out-crt", claims.out)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_isStdio(t *testing.T) {
|
|
||||||
assert.True(t, isStdio("-"))
|
|
||||||
assert.False(t, isStdio(""))
|
|
||||||
assert.False(t, isStdio("./-"))
|
|
||||||
assert.False(t, isStdio("foo"))
|
|
||||||
}
|
|
||||||
@@ -39,26 +39,18 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var claims ioClaims
|
caFile, err := os.Open(*vf.caPath)
|
||||||
if err := reserveInputs(&claims,
|
|
||||||
"ca", *vf.caPath,
|
|
||||||
"crt", *vf.certPath,
|
|
||||||
); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
caReader, err := openInput("ca", *vf.caPath, &claims)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca: %w", err)
|
return fmt.Errorf("error while reading ca: %w", err)
|
||||||
}
|
}
|
||||||
defer caReader.Close()
|
defer caFile.Close()
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caReader)
|
caPool, err := cert.NewCAPoolFromPEMReader(caFile)
|
||||||
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
||||||
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := readInput("crt", *vf.certPath, &claims)
|
rawCert, err := os.ReadFile(*vf.certPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read crt: %w", err)
|
return fmt.Errorf("unable to read crt: %w", err)
|
||||||
}
|
}
|
||||||
@@ -93,7 +85,6 @@ func verifySummary() string {
|
|||||||
func verifyHelp(out io.Writer) {
|
func verifyHelp(out io.Writer) {
|
||||||
vf := newVerifyFlags()
|
vf := newVerifyFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
||||||
_, _ = out.Write([]byte(stdioHelpText))
|
|
||||||
vf.set.SetOutput(out)
|
vf.set.SetOutput(out)
|
||||||
vf.set.PrintDefaults()
|
vf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ func Test_verifyHelp(t *testing.T) {
|
|||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
"Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+
|
"Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+
|
||||||
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
|
|
||||||
" -ca string\n"+
|
" -ca string\n"+
|
||||||
" \tRequired: path to a file containing one or more ca certificates\n"+
|
" \tRequired: path to a file containing one or more ca certificates\n"+
|
||||||
" -crt string\n"+
|
" -crt string\n"+
|
||||||
@@ -123,46 +122,3 @@ func Test_verify(t *testing.T) {
|
|||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_verify_stdio(t *testing.T) {
|
|
||||||
ob := &bytes.Buffer{}
|
|
||||||
eb := &bytes.Buffer{}
|
|
||||||
|
|
||||||
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
|
||||||
ca, _ := NewTestCaCert("test-ca", caPub, caPriv, time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour*2), nil, nil, nil)
|
|
||||||
caPEM, _ := ca.MarshalPEM()
|
|
||||||
|
|
||||||
crt, _ := NewTestCert(ca, caPriv, "test-cert", time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour), nil, nil, nil)
|
|
||||||
crtPEM, _ := crt.MarshalPEM()
|
|
||||||
|
|
||||||
caFile, err := os.CreateTemp("", "verify-ca")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer os.Remove(caFile.Name())
|
|
||||||
caFile.Write(caPEM)
|
|
||||||
|
|
||||||
// crt on stdin, ca on disk
|
|
||||||
withStdin(t, bytes.NewReader(crtPEM))
|
|
||||||
require.NoError(t, verify([]string{"-ca", caFile.Name(), "-crt", "-"}, ob, eb))
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
|
|
||||||
// ca on stdin, crt on disk
|
|
||||||
certFile, err := os.CreateTemp("", "verify-cert")
|
|
||||||
require.NoError(t, err)
|
|
||||||
defer os.Remove(certFile.Name())
|
|
||||||
certFile.Write(crtPEM)
|
|
||||||
|
|
||||||
withStdin(t, bytes.NewReader(caPEM))
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.NoError(t, verify([]string{"-ca", "-", "-crt", certFile.Name()}, ob, eb))
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
|
|
||||||
// both flags on stdin should error
|
|
||||||
withStdin(t, bytes.NewReader(caPEM))
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
require.EqualError(t, verify([]string{"-ca", "-", "-crt", "-"}, ob, eb),
|
|
||||||
`-ca and -crt both set to "-", only one input may read from stdin`)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -61,12 +61,9 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
fmt.Println("-config flag must be set")
|
||||||
if err != nil {
|
flag.Usage()
|
||||||
fmt.Println(err)
|
os.Exit(1)
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
*configPath = p
|
|
||||||
}
|
}
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ package main
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
@@ -55,13 +57,24 @@ func (p *program) Stop(s service.Service) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func fileExists(filename string) bool {
|
||||||
|
_, err := os.Stat(filename)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
ex, err := os.Executable()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
*configPath = p
|
*configPath = filepath.Dir(ex) + "/config.yaml"
|
||||||
|
if !fileExists(*configPath) {
|
||||||
|
*configPath = filepath.Dir(ex) + "/config.yml"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
svcConfig := &service.Config{
|
svcConfig := &service.Config{
|
||||||
|
|||||||
+3
-6
@@ -50,12 +50,9 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
p, err := config.DefaultPath()
|
fmt.Println("-config flag must be set")
|
||||||
if err != nil {
|
flag.Usage()
|
||||||
fmt.Println(err)
|
os.Exit(1)
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
*configPath = p
|
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
l := logging.NewLogger(os.Stdout)
|
||||||
|
|||||||
@@ -1,29 +0,0 @@
|
|||||||
package config
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
)
|
|
||||||
|
|
||||||
// DefaultPath returns a path to a config file alongside the running executable, preferring config.yaml over config.yml.
|
|
||||||
// If neither file exists an error is returned that names both paths checked.
|
|
||||||
func DefaultPath() (string, error) {
|
|
||||||
ex, err := os.Executable()
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
return defaultPathInDir(filepath.Dir(ex))
|
|
||||||
}
|
|
||||||
|
|
||||||
func defaultPathInDir(dir string) (string, error) {
|
|
||||||
yamlPath := filepath.Join(dir, "config.yaml")
|
|
||||||
if _, err := os.Stat(yamlPath); err == nil {
|
|
||||||
return yamlPath, nil
|
|
||||||
}
|
|
||||||
ymlPath := filepath.Join(dir, "config.yml")
|
|
||||||
if _, err := os.Stat(ymlPath); err == nil {
|
|
||||||
return ymlPath, nil
|
|
||||||
}
|
|
||||||
return "", fmt.Errorf("no default config found at %s or %s", yamlPath, ymlPath)
|
|
||||||
}
|
|
||||||
@@ -1,67 +0,0 @@
|
|||||||
package config
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestDefaultPathInDir(t *testing.T) {
|
|
||||||
t.Run("prefers config.yaml when both exist", func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
want := filepath.Join(dir, "config.yaml")
|
|
||||||
other := filepath.Join(dir, "config.yml")
|
|
||||||
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
|
||||||
require.NoError(t, os.WriteFile(other, []byte("a: 2"), 0644))
|
|
||||||
|
|
||||||
got, err := defaultPathInDir(dir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, want, got)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("returns config.yaml when only it exists", func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
want := filepath.Join(dir, "config.yaml")
|
|
||||||
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
|
||||||
|
|
||||||
got, err := defaultPathInDir(dir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, want, got)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("falls back to config.yml when only it exists", func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
want := filepath.Join(dir, "config.yml")
|
|
||||||
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
|
|
||||||
|
|
||||||
got, err := defaultPathInDir(dir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, want, got)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("errors when neither exists and names both paths", func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
got, err := defaultPathInDir(dir)
|
|
||||||
assert.Empty(t, got)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yaml"))
|
|
||||||
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yml"))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDefaultPath(t *testing.T) {
|
|
||||||
got, err := DefaultPath()
|
|
||||||
if err != nil {
|
|
||||||
ex, exErr := os.Executable()
|
|
||||||
require.NoError(t, exErr)
|
|
||||||
assert.Contains(t, err.Error(), filepath.Dir(ex))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ex, err := os.Executable()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, filepath.Dir(ex), filepath.Dir(got))
|
|
||||||
assert.Contains(t, []string{"config.yaml", "config.yml"}, filepath.Base(got))
|
|
||||||
}
|
|
||||||
+10
-2
@@ -136,6 +136,14 @@ func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time)
|
|||||||
return in, out
|
return in, out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddTrafficWatch must be called for every new HostInfo.
|
||||||
|
// We will continue to monitor the HostInfo until the tunnel is dropped.
|
||||||
|
func (cm *connectionManager) AddTrafficWatch(h *HostInfo) {
|
||||||
|
if h.out.Swap(true) == false {
|
||||||
|
cm.trafficTimer.Add(h.localIndexId, cm.checkInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) Start(ctx context.Context) {
|
func (cm *connectionManager) Start(ctx context.Context) {
|
||||||
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
@@ -298,8 +306,8 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
} else {
|
} else {
|
||||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
cm.l.Info("send CreateRelayRequest",
|
cm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", relayFrom,
|
"relayFrom", req.RelayFromAddr,
|
||||||
"relayTo", relayTo,
|
"relayTo", req.RelayToAddr,
|
||||||
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", req.ResponderRelayIndex,
|
"responderRelayIndex", req.ResponderRelayIndex,
|
||||||
"vpnAddrs", newhostinfo.vpnAddrs,
|
"vpnAddrs", newhostinfo.vpnAddrs,
|
||||||
|
|||||||
Vendored
+1
-1
@@ -62,7 +62,7 @@ function nebula.dissector(tvbuf, pktinfo, root)
|
|||||||
tree:add(pf_version, tvbuf:range(0,1))
|
tree:add(pf_version, tvbuf:range(0,1))
|
||||||
local type = tree:add(pf_type, tvbuf:range(0,1))
|
local type = tree:add(pf_type, tvbuf:range(0,1))
|
||||||
|
|
||||||
local nebula_type = bit.band(tvbuf:range(0,1):uint(), 0x0F)
|
local nebula_type = bit32.band(tvbuf:range(0,1):uint(), 0x0F)
|
||||||
if nebula_type == 0 then
|
if nebula_type == 0 then
|
||||||
local stage = tvbuf(8,8):uint64()
|
local stage = tvbuf(8,8):uint64()
|
||||||
tree:add(pf_subtype_handshake, tvbuf:range(1,1))
|
tree:add(pf_subtype_handshake, tvbuf:range(1,1))
|
||||||
|
|||||||
+26
-92
@@ -11,21 +11,19 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
type dnsServer struct {
|
type dnsServer struct {
|
||||||
sync.RWMutex
|
sync.RWMutex
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
dnsMap4 map[string]netip.Addr
|
dnsMap4 map[string]netip.Addr
|
||||||
dnsMap6 map[string]netip.Addr
|
dnsMap6 map[string]netip.Addr
|
||||||
hostMap *HostMap
|
hostMap *HostMap
|
||||||
pki *PKI
|
myVpnAddrsTable *bart.Lite
|
||||||
|
|
||||||
// selfHost is the cached FQDN we last seeded for ourselves
|
|
||||||
selfHost string
|
|
||||||
|
|
||||||
mux *dns.ServeMux
|
mux *dns.ServeMux
|
||||||
|
|
||||||
@@ -57,14 +55,14 @@ type dnsServer struct {
|
|||||||
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
||||||
// watcher that tears the listener down on nebula shutdown. The returned
|
// watcher that tears the listener down on nebula shutdown. The returned
|
||||||
// pointer is always non-nil, even on error.
|
// pointer is always non-nil, even on error.
|
||||||
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
||||||
ds := &dnsServer{
|
ds := &dnsServer{
|
||||||
l: l,
|
l: l,
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
dnsMap4: make(map[string]netip.Addr),
|
dnsMap4: make(map[string]netip.Addr),
|
||||||
dnsMap6: make(map[string]netip.Addr),
|
dnsMap6: make(map[string]netip.Addr),
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
pki: pki,
|
myVpnAddrsTable: cs.myVpnAddrsTable,
|
||||||
}
|
}
|
||||||
ds.mux = dns.NewServeMux()
|
ds.mux = dns.NewServeMux()
|
||||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||||
@@ -78,7 +76,6 @@ func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostM
|
|||||||
if err := ds.reload(c, true); err != nil {
|
if err := ds.reload(c, true); err != nil {
|
||||||
return ds, err
|
return ds, err
|
||||||
}
|
}
|
||||||
ds.seedSelf()
|
|
||||||
return ds, nil
|
return ds, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -116,7 +113,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
d.Stop()
|
d.Stop()
|
||||||
}
|
}
|
||||||
// Drop any records that accumulated while enabled; a later re-enable
|
// Drop any records that accumulated while enabled; a later re-enable
|
||||||
// will repopulate from fresh handshakes and a fresh seedSelf.
|
// will repopulate from fresh handshakes.
|
||||||
d.clearRecords()
|
d.clearRecords()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -124,14 +121,17 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|||||||
if running == nil {
|
if running == nil {
|
||||||
// Was disabled (or never started); bring it up now.
|
// Was disabled (or never started); bring it up now.
|
||||||
go d.Start()
|
go d.Start()
|
||||||
} else if !sameAddr {
|
return nil
|
||||||
d.shutdownServer(running, runningStarted, "reload")
|
|
||||||
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
|
|
||||||
go d.Start()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Refresh the self entry every enabled reload so cert renewals that change our name or VPN addresses are picked up.
|
if sameAddr {
|
||||||
d.seedSelf()
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
d.shutdownServer(running, runningStarted, "reload")
|
||||||
|
// Old Start goroutine has now exited; bring up a fresh listener on the
|
||||||
|
// new address.
|
||||||
|
go d.Start()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -249,20 +249,6 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// The hostmap only ever contains peers we have handshaked with, so it never carries an entry for ourselves.
|
|
||||||
// Answer self lookups straight from the local cert state.
|
|
||||||
if cs := d.certState(); cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
|
|
||||||
c := cs.GetDefaultCertificate()
|
|
||||||
if c == nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
b, err := c.MarshalJSON()
|
|
||||||
if err != nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return string(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
return ""
|
return ""
|
||||||
@@ -280,60 +266,12 @@ func (d *dnsServer) QueryCert(data string) string {
|
|||||||
return string(b)
|
return string(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearRecords drops all DNS records, including the self entry.
|
// clearRecords drops all DNS records.
|
||||||
func (d *dnsServer) clearRecords() {
|
func (d *dnsServer) clearRecords() {
|
||||||
d.Lock()
|
d.Lock()
|
||||||
defer d.Unlock()
|
defer d.Unlock()
|
||||||
clear(d.dnsMap4)
|
clear(d.dnsMap4)
|
||||||
clear(d.dnsMap6)
|
clear(d.dnsMap6)
|
||||||
d.selfHost = ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// seedSelf inserts (or refreshes) a record for our own cert name pointing at our VPN addresses,
|
|
||||||
// so a single-lighthouse network can resolve the lighthouse's own hostname without the two-process workaround.
|
|
||||||
func (d *dnsServer) seedSelf() {
|
|
||||||
if !d.enabled.Load() {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cs := d.certState()
|
|
||||||
if cs == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
c := cs.GetDefaultCertificate()
|
|
||||||
if c == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
newHost := strings.ToLower(c.Name()) + "."
|
|
||||||
|
|
||||||
d.Lock()
|
|
||||||
defer d.Unlock()
|
|
||||||
if d.selfHost != "" && d.selfHost != newHost {
|
|
||||||
delete(d.dnsMap4, d.selfHost)
|
|
||||||
delete(d.dnsMap6, d.selfHost)
|
|
||||||
}
|
|
||||||
d.selfHost = newHost
|
|
||||||
delete(d.dnsMap4, newHost)
|
|
||||||
delete(d.dnsMap6, newHost)
|
|
||||||
haveV4, haveV6 := false, false
|
|
||||||
for _, addr := range cs.myVpnAddrs {
|
|
||||||
if addr.Is4() && !haveV4 {
|
|
||||||
d.dnsMap4[newHost] = addr
|
|
||||||
haveV4 = true
|
|
||||||
} else if addr.Is6() && !haveV6 {
|
|
||||||
d.dnsMap6[newHost] = addr
|
|
||||||
haveV6 = true
|
|
||||||
}
|
|
||||||
if haveV4 && haveV6 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dnsServer) certState() *CertState {
|
|
||||||
if d.pki == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return d.pki.getCertState()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
||||||
@@ -371,12 +309,8 @@ func (d *dnsServer) isSelfNebulaOrLocalhost(addr string) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
cs := d.certState()
|
|
||||||
if cs == nil || cs.myVpnAddrsTable == nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
//if we found it in this table, it's good
|
//if we found it in this table, it's good
|
||||||
return cs.myVpnAddrsTable.Contains(b)
|
return d.myVpnAddrsTable.Contains(b)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||||
|
|||||||
@@ -9,10 +9,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -279,92 +276,6 @@ func TestDnsServer_Stop_beforeBind_doesNotHang(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTestPKI builds a minimal *PKI with a single v1 cert whose name and
|
|
||||||
// VPN addresses are caller-provided, suitable for exercising seedSelf and
|
|
||||||
// QueryCert self handling.
|
|
||||||
func newTestPKI(t *testing.T, name string, addrs []netip.Addr) *PKI {
|
|
||||||
t.Helper()
|
|
||||||
networks := make([]netip.Prefix, 0, len(addrs))
|
|
||||||
for _, a := range addrs {
|
|
||||||
bits := 32
|
|
||||||
if a.Is6() {
|
|
||||||
bits = 128
|
|
||||||
}
|
|
||||||
networks = append(networks, netip.PrefixFrom(a, bits))
|
|
||||||
}
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
c, _, _, _ := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, name, time.Time{}, time.Time{}, networks, nil, nil)
|
|
||||||
|
|
||||||
addrsTable := new(bart.Lite)
|
|
||||||
for _, a := range addrs {
|
|
||||||
addrsTable.Insert(netip.PrefixFrom(a, a.BitLen()))
|
|
||||||
}
|
|
||||||
|
|
||||||
cs := &CertState{
|
|
||||||
v2Cert: c,
|
|
||||||
initiatingVersion: cert.Version2,
|
|
||||||
myVpnAddrs: addrs,
|
|
||||||
myVpnAddrsTable: addrsTable,
|
|
||||||
}
|
|
||||||
pki := &PKI{}
|
|
||||||
pki.cs.Store(cs)
|
|
||||||
return pki
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDnsServer_seedSelf_addsOwnRecord(t *testing.T) {
|
|
||||||
ds, c := newTestDnsServer(t)
|
|
||||||
myV4 := netip.MustParseAddr("10.0.0.1")
|
|
||||||
myV6 := netip.MustParseAddr("fd00::1")
|
|
||||||
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4, myV6})
|
|
||||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
|
||||||
require.NoError(t, ds.reload(c, true))
|
|
||||||
|
|
||||||
ds.seedSelf()
|
|
||||||
got4, exists := ds.Query(dns.TypeA, "lighthouse.")
|
|
||||||
assert.True(t, exists)
|
|
||||||
assert.Equal(t, myV4, got4)
|
|
||||||
got6, exists := ds.Query(dns.TypeAAAA, "lighthouse.")
|
|
||||||
assert.True(t, exists)
|
|
||||||
assert.Equal(t, myV6, got6)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDnsServer_seedSelf_disabled_noOp(t *testing.T) {
|
|
||||||
ds, c := newTestDnsServer(t)
|
|
||||||
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
|
|
||||||
setDnsConfig(c, "127.0.0.1", "0", true, false)
|
|
||||||
require.NoError(t, ds.reload(c, true))
|
|
||||||
|
|
||||||
ds.seedSelf()
|
|
||||||
_, exists := ds.Query(dns.TypeA, "lighthouse.")
|
|
||||||
assert.False(t, exists)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDnsServer_clearRecords_dropsSelfHost(t *testing.T) {
|
|
||||||
ds, c := newTestDnsServer(t)
|
|
||||||
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
|
|
||||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
|
||||||
require.NoError(t, ds.reload(c, true))
|
|
||||||
ds.seedSelf()
|
|
||||||
require.NotEmpty(t, ds.selfHost)
|
|
||||||
|
|
||||||
ds.clearRecords()
|
|
||||||
assert.Empty(t, ds.selfHost)
|
|
||||||
_, exists := ds.Query(dns.TypeA, "lighthouse.")
|
|
||||||
assert.False(t, exists)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDnsServer_QueryCert_returnsOwnCert(t *testing.T) {
|
|
||||||
ds, _ := newTestDnsServer(t)
|
|
||||||
myV4 := netip.MustParseAddr("10.0.0.1")
|
|
||||||
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4})
|
|
||||||
|
|
||||||
got := ds.QueryCert(myV4.String() + ".")
|
|
||||||
assert.NotEmpty(t, got, "TXT lookup of our own VPN address should return our cert")
|
|
||||||
|
|
||||||
other := netip.MustParseAddr("10.0.0.99")
|
|
||||||
assert.Empty(t, ds.QueryCert(other.String()+"."), "unknown peer IP should return nothing")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
|
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
|
||||||
port := freeUDPPort(t)
|
port := freeUDPPort(t)
|
||||||
ds, c := newTestDnsServer(t)
|
ds, c := newTestDnsServer(t)
|
||||||
|
|||||||
@@ -1,16 +1,6 @@
|
|||||||
FROM gcr.io/distroless/static:latest
|
FROM gcr.io/distroless/static:latest
|
||||||
|
|
||||||
ARG TARGETOS TARGETARCH
|
ARG TARGETOS TARGETARCH
|
||||||
|
|
||||||
ARG VERSION=dev
|
|
||||||
ARG REVISION=unknown
|
|
||||||
LABEL org.opencontainers.image.title="nebula" \
|
|
||||||
org.opencontainers.image.description="A scalable overlay networking tool with a focus on performance, simplicity and security" \
|
|
||||||
org.opencontainers.image.vendor="Nebula OSS" \
|
|
||||||
org.opencontainers.image.source="https://github.com/slackhq/nebula" \
|
|
||||||
org.opencontainers.image.version="${VERSION}" \
|
|
||||||
org.opencontainers.image.revision="${REVISION}"
|
|
||||||
|
|
||||||
COPY build/$TARGETOS-$TARGETARCH/nebula /nebula
|
COPY build/$TARGETOS-$TARGETARCH/nebula /nebula
|
||||||
COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert
|
COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert
|
||||||
|
|
||||||
|
|||||||
+2
-4
@@ -4,15 +4,13 @@
|
|||||||
package e2e
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
@@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *slog.Logger {
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
return slog.New(slog.DiscardHandler)
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
level := slog.LevelInfo
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"gopkg.in/yaml.v3"
|
"gopkg.in/yaml.v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -374,100 +373,6 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
//relayControl.Stop()
|
//relayControl.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestRelayReplayProtection asserts that a relay (forwarding-type) node rejects
|
|
||||||
// replayed relay frames. A captured relay frame, re-injected with the same
|
|
||||||
// message counter, must be dropped by the replay window rather than re-forwarded
|
|
||||||
// to the relay target. Before the fix, handleOutsideRelayPacket authenticated the
|
|
||||||
// frame but never advanced the replay window, so every replay was re-forwarded.
|
|
||||||
func TestRelayReplayProtection(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{})
|
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
|
||||||
theirUdp := netip.MustParseAddrPort("10.0.0.2:4242")
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdp(cert.Version2, ca, caKey, "them ", "fc00::2/64", theirUdp, m{"relay": m{"use_relays": true}})
|
|
||||||
|
|
||||||
myVpnV6 := myVpnIpNet[1]
|
|
||||||
relayVpnV4 := relayVpnIpNet[0]
|
|
||||||
relayVpnV6 := relayVpnIpNet[1]
|
|
||||||
theirVpnV6 := theirVpnIpNet[0]
|
|
||||||
|
|
||||||
// Teach me how to reach the relay and that them is reachable via the relay
|
|
||||||
myControl.InjectLightHouseAddr(relayVpnV4.Addr(), relayUdpAddr)
|
|
||||||
myControl.InjectLightHouseAddr(relayVpnV6.Addr(), relayUdpAddr)
|
|
||||||
myControl.InjectRelays(theirVpnV6.Addr(), []netip.Addr{relayVpnV6.Addr()})
|
|
||||||
relayControl.InjectLightHouseAddr(theirVpnV6.Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
relayControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
// Establish the relayed tunnel in both directions so all handshakes complete.
|
|
||||||
t.Log("Establish the relayed tunnel")
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
|
||||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
|
||||||
|
|
||||||
// Drain anything still queued on me's UDP tx so the next packet we pull is the
|
|
||||||
// relay frame we are about to generate.
|
|
||||||
for myControl.GetFromUDP(false) != nil {
|
|
||||||
}
|
|
||||||
|
|
||||||
// Capture a single legitimate relay frame that me transmits toward the relay.
|
|
||||||
t.Log("Capture a relay frame from me -> relay")
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("replay me")))
|
|
||||||
relayFrame := myControl.GetFromUDP(true)
|
|
||||||
require.Equal(t, relayUdpAddr, relayFrame.To, "captured frame should be addressed to the relay")
|
|
||||||
var fh header.H
|
|
||||||
require.NoError(t, fh.Parse(relayFrame.Data))
|
|
||||||
require.Equal(t, header.Message, fh.Type)
|
|
||||||
require.Equal(t, header.MessageRelay, fh.Subtype)
|
|
||||||
|
|
||||||
// drainForwards counts relay frames the relay forwards toward them within the
|
|
||||||
// settle window. We match on destination + (Message, MessageRelay) so the
|
|
||||||
// relay's own direct traffic to them can't be miscounted.
|
|
||||||
drainForwards := func(settle time.Duration) int {
|
|
||||||
ch := relayControl.GetUDPTxChan()
|
|
||||||
count := 0
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case pkt := <-ch:
|
|
||||||
var ph header.H
|
|
||||||
if pkt.To == theirUdpAddr && ph.Parse(pkt.Data) == nil &&
|
|
||||||
ph.Type == header.Message && ph.Subtype == header.MessageRelay {
|
|
||||||
count++
|
|
||||||
}
|
|
||||||
pkt.Release()
|
|
||||||
case <-time.After(settle):
|
|
||||||
return count
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// First delivery of the captured frame: the relay should forward it once.
|
|
||||||
t.Log("Deliver the captured frame once; relay forwards it to them")
|
|
||||||
relayControl.InjectUDPPacket(relayFrame)
|
|
||||||
require.Equal(t, 1, drainForwards(200*time.Millisecond), "relay should forward the first, legitimate copy")
|
|
||||||
|
|
||||||
// Replay the exact same frame several times. A correct replay window rejects
|
|
||||||
// these duplicates so the relay forwards none of them.
|
|
||||||
t.Log("Replay the captured frame; relay must drop the duplicates")
|
|
||||||
const replays = 3
|
|
||||||
for i := 0; i < replays; i++ {
|
|
||||||
relayControl.InjectUDPPacket(relayFrame)
|
|
||||||
}
|
|
||||||
forwarded := drainForwards(200 * time.Millisecond)
|
|
||||||
assert.Equal(t, 0, forwarded, "relay re-forwarded %d/%d replayed relay frames; replay protection is ineffective on relay tunnels", forwarded, replays)
|
|
||||||
|
|
||||||
r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCloseTunnelAuthenticated(t *testing.T) {
|
func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|||||||
+1
-1
@@ -397,7 +397,7 @@ firewall:
|
|||||||
# `drop` (default): silently drop the packet.
|
# `drop` (default): silently drop the packet.
|
||||||
# `reject`: send a reject reply.
|
# `reject`: send a reject reply.
|
||||||
# - For TCP, this will be a RST "Connection Reset" packet.
|
# - For TCP, this will be a RST "Connection Reset" packet.
|
||||||
# - For other protocols, this will be an ICMP "Destination unreachable: Communication administratively prohibited" packet.
|
# - For other protocols, this will be an ICMP port unreachable packet.
|
||||||
outbound_action: drop
|
outbound_action: drop
|
||||||
inbound_action: drop
|
inbound_action: drop
|
||||||
|
|
||||||
|
|||||||
+26
-43
@@ -58,9 +58,8 @@ type Firewall struct {
|
|||||||
routableNetworks *bart.Lite
|
routableNetworks *bart.Lite
|
||||||
|
|
||||||
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
// assignedNetworks is a list of vpn networks assigned to us in the certificate.
|
||||||
assignedNetworks []netip.Prefix
|
assignedNetworks []netip.Prefix
|
||||||
// unsafeNetworks is the list of unsafe networks issued to us in the certificate
|
hasUnsafeNetworks bool
|
||||||
unsafeNetworks []netip.Prefix
|
|
||||||
|
|
||||||
rules string
|
rules string
|
||||||
rulesVersion uint16
|
rulesVersion uint16
|
||||||
@@ -159,9 +158,10 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
assignedNetworks = append(assignedNetworks, network)
|
assignedNetworks = append(assignedNetworks, network)
|
||||||
}
|
}
|
||||||
|
|
||||||
unsafeNetworks := c.UnsafeNetworks()
|
hasUnsafeNetworks := false
|
||||||
for _, n := range unsafeNetworks {
|
for _, n := range c.UnsafeNetworks() {
|
||||||
routableNetworks.Insert(n)
|
routableNetworks.Insert(n)
|
||||||
|
hasUnsafeNetworks = true
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Firewall{
|
return &Firewall{
|
||||||
@@ -169,15 +169,15 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
Conns: make(map[firewall.Packet]*conn),
|
Conns: make(map[firewall.Packet]*conn),
|
||||||
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
||||||
},
|
},
|
||||||
InRules: newFirewallTable(),
|
InRules: newFirewallTable(),
|
||||||
OutRules: newFirewallTable(),
|
OutRules: newFirewallTable(),
|
||||||
TCPTimeout: tcpTimeout,
|
TCPTimeout: tcpTimeout,
|
||||||
UDPTimeout: UDPTimeout,
|
UDPTimeout: UDPTimeout,
|
||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
unsafeNetworks: unsafeNetworks,
|
hasUnsafeNetworks: hasUnsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
||||||
@@ -897,7 +897,7 @@ func (flc *firewallLocalCIDR) addRule(f *Firewall, localCidr string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if localCidr == "" {
|
if localCidr == "" {
|
||||||
if len(f.unsafeNetworks) == 0 || f.defaultLocalCIDRAny {
|
if !f.hasUnsafeNetworks || f.defaultLocalCIDRAny {
|
||||||
flc.Any = true
|
flc.Any = true
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1055,6 +1055,7 @@ func (r *rule) sanity() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func parsePort(s string) (int32, int32, error) {
|
func parsePort(s string) (int32, int32, error) {
|
||||||
|
var err error
|
||||||
const notAPort int32 = -2
|
const notAPort int32 = -2
|
||||||
if s == "any" {
|
if s == "any" {
|
||||||
return firewall.PortAny, firewall.PortAny, nil
|
return firewall.PortAny, firewall.PortAny, nil
|
||||||
@@ -1063,11 +1064,11 @@ func parsePort(s string) (int32, int32, error) {
|
|||||||
return firewall.PortFragment, firewall.PortFragment, nil
|
return firewall.PortFragment, firewall.PortFragment, nil
|
||||||
}
|
}
|
||||||
if !strings.Contains(s, `-`) {
|
if !strings.Contains(s, `-`) {
|
||||||
rPort, err := parsePortValue("", s)
|
rPort, err := strconv.Atoi(s)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("was not a number; `%s`", s)
|
||||||
}
|
}
|
||||||
return rPort, rPort, nil
|
return int32(rPort), int32(rPort), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
sPorts := strings.SplitN(s, `-`, 2)
|
sPorts := strings.SplitN(s, `-`, 2)
|
||||||
@@ -1078,40 +1079,22 @@ func parsePort(s string) (int32, int32, error) {
|
|||||||
return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s)
|
return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s)
|
||||||
}
|
}
|
||||||
|
|
||||||
startPort, err := parsePortValue("beginning range ", sPorts[0])
|
rStartPort, err := strconv.Atoi(sPorts[0])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("beginning range was not a number; `%s`", sPorts[0])
|
||||||
}
|
}
|
||||||
|
|
||||||
endPort, err := parsePortValue("ending range ", sPorts[1])
|
rEndPort, err := strconv.Atoi(sPorts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return notAPort, notAPort, err
|
return notAPort, notAPort, fmt.Errorf("ending range was not a number; `%s`", sPorts[1])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
startPort := int32(rStartPort)
|
||||||
|
endPort := int32(rEndPort)
|
||||||
|
|
||||||
if startPort == firewall.PortAny {
|
if startPort == firewall.PortAny {
|
||||||
endPort = firewall.PortAny
|
endPort = firewall.PortAny
|
||||||
}
|
}
|
||||||
|
|
||||||
return startPort, endPort, nil
|
return startPort, endPort, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// parsePortValue accepts a base-10 decimal in [0, 65535] and returns it
|
|
||||||
// widened to int32. Using strconv.ParseUint with bitSize 16 rejects
|
|
||||||
// negative input, out-of-range input (>65535), and any non-decimal byte
|
|
||||||
// by construction, so the int32 widening that follows is provably safe
|
|
||||||
// and cannot collide with firewall.PortAny (0) or firewall.PortFragment
|
|
||||||
// (-1) via integer truncation.
|
|
||||||
//
|
|
||||||
// prefix is prepended to both error messages so callers can disambiguate
|
|
||||||
// the single-port path (prefix="") from the range bounds (prefix="beginning
|
|
||||||
// range " / "ending range "), preserving the historical error strings.
|
|
||||||
func parsePortValue(prefix, s string) (int32, error) {
|
|
||||||
n, err := strconv.ParseUint(s, 10, 16)
|
|
||||||
if err == nil {
|
|
||||||
return int32(n), nil
|
|
||||||
}
|
|
||||||
if errors.Is(err, strconv.ErrRange) {
|
|
||||||
return 0, fmt.Errorf("%sout of range [0,65535]; `%s`", prefix, s)
|
|
||||||
}
|
|
||||||
return 0, fmt.Errorf("%swas not a number; `%s`", prefix, s)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1029,75 +1029,6 @@ func Test_parsePort(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test_parsePort_invalid covers inputs that must error. The named bug is
|
|
||||||
// that int32(strconv.Atoi("4294967296")) truncates to 0 == firewall.PortAny,
|
|
||||||
// silently turning a typo into a match-all-ports rule; the rest are
|
|
||||||
// representative syntax/range probes.
|
|
||||||
func Test_parsePort_invalid(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
input string
|
|
||||||
wantErrContains string
|
|
||||||
}{
|
|
||||||
// Numeric overflow (the named bug + boundary).
|
|
||||||
{"named bug: 2^32 truncates to PortAny", "4294967296", "out of range"},
|
|
||||||
{"just above max real port", "65536", "out of range"},
|
|
||||||
|
|
||||||
// Negatives route through the range branch and hit the empty-half
|
|
||||||
// guard; included as defense in depth so a future refactor cannot
|
|
||||||
// accidentally reach the int32 cast.
|
|
||||||
{"negative", "-1", "could not be parsed"},
|
|
||||||
|
|
||||||
// Syntax probes.
|
|
||||||
{"NUL between digits", "4\x002", "was not a number"},
|
|
||||||
{"hex notation", "0x10", "was not a number"},
|
|
||||||
{"scientific notation", "1e3", "was not a number"},
|
|
||||||
{"leading whitespace", " 42", "was not a number"},
|
|
||||||
{"fullwidth digits", "42", "was not a number"},
|
|
||||||
|
|
||||||
// Range branch.
|
|
||||||
{"range upper out of range", "1-65536", "ending range out of range"},
|
|
||||||
{"range lower out of range", "65536-65537", "beginning range out of range"},
|
|
||||||
{"range with negative upper", "1--1", "ending range was not a number"},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
_, _, err := parsePort(tc.input)
|
|
||||||
require.Error(t, err, "input %q must error", tc.input)
|
|
||||||
require.ErrorContains(t, err, tc.wantErrContains)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test_parsePort_valid_boundaries locks in success cases at 0, 1, and 65535
|
|
||||||
// so a future refactor cannot regress the boundaries.
|
|
||||||
func Test_parsePort_valid_boundaries(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
input string
|
|
||||||
wantStart int32
|
|
||||||
wantEnd int32
|
|
||||||
}{
|
|
||||||
{"zero is PortAny", "0", 0, 0},
|
|
||||||
{"min real port", "1", 1, 1},
|
|
||||||
{"max real port", "65535", 65535, 65535},
|
|
||||||
{"range zero to max forces end to zero", "0-65535", 0, 0},
|
|
||||||
{"range max to max", "65535-65535", 65535, 65535},
|
|
||||||
{"range one to max", "1-65535", 1, 65535},
|
|
||||||
{"range with whitespace inside", " 1 - 2 ", 1, 2},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
s, e, err := parsePort(tc.input)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, tc.wantStart, s, "start port")
|
|
||||||
assert.Equal(t, tc.wantEnd, e, "end port")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewFirewallFromConfig(t *testing.T) {
|
func TestNewFirewallFromConfig(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Test a bad rule definition
|
// Test a bad rule definition
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ require (
|
|||||||
github.com/armon/go-radix v1.0.0
|
github.com/armon/go-radix v1.0.0
|
||||||
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
||||||
github.com/flynn/noise v1.1.0
|
github.com/flynn/noise v1.1.0
|
||||||
github.com/gaissmai/bart v0.28.0
|
github.com/gaissmai/bart v0.26.1
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
github.com/google/gopacket v1.1.19
|
github.com/google/gopacket v1.1.19
|
||||||
github.com/kardianos/service v1.2.4
|
github.com/kardianos/service v1.2.4
|
||||||
@@ -24,12 +24,12 @@ require (
|
|||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
go.uber.org/goleak v1.3.0
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.53.0
|
golang.org/x/crypto v0.50.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
golang.org/x/net v0.56.0
|
golang.org/x/net v0.53.0
|
||||||
golang.org/x/sync v0.21.0
|
golang.org/x/sync v0.20.0
|
||||||
golang.org/x/sys v0.46.0
|
golang.org/x/sys v0.43.0
|
||||||
golang.org/x/term v0.44.0
|
golang.org/x/term v0.42.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1
|
golang.zx2c4.com/wireguard/windows v0.6.1
|
||||||
|
|||||||
@@ -26,8 +26,8 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
|
|||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
||||||
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
||||||
github.com/gaissmai/bart v0.28.0 h1:89yZLo8NmyqD0RYgJ3QO9HhqqGGw+oWhf90cZm69Lko=
|
github.com/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo=
|
||||||
github.com/gaissmai/bart v0.28.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
github.com/gaissmai/bart v0.26.1/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||||
@@ -162,8 +162,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
|||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||||
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=
|
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||||
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
@@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
|
golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
||||||
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
|
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
|
|||||||
@@ -13,7 +13,6 @@ var (
|
|||||||
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
||||||
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
||||||
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
||||||
ErrInvalidRemoteIndex = errors.New("peer sent an invalid index in handshake payload")
|
|
||||||
ErrIndexAllocation = errors.New("failed to allocate local index")
|
ErrIndexAllocation = errors.New("failed to allocate local index")
|
||||||
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
||||||
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
||||||
|
|||||||
+2
-10
@@ -312,19 +312,11 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
|||||||
|
|
||||||
// Process payload
|
// Process payload
|
||||||
if flags.expectsPayload {
|
if flags.expectsPayload {
|
||||||
var remoteIndex uint32
|
|
||||||
if m.result.Initiator {
|
if m.result.Initiator {
|
||||||
remoteIndex = payload.ResponderIndex
|
m.result.RemoteIndex = payload.ResponderIndex
|
||||||
} else {
|
} else {
|
||||||
remoteIndex = payload.InitiatorIndex
|
m.result.RemoteIndex = payload.InitiatorIndex
|
||||||
}
|
}
|
||||||
// The payload presence check above can be satisfied by Time alone, so a payload
|
|
||||||
// could still carry a zero index here. We need to reject it.
|
|
||||||
if remoteIndex == 0 {
|
|
||||||
m.failed = true
|
|
||||||
return ErrInvalidRemoteIndex
|
|
||||||
}
|
|
||||||
m.result.RemoteIndex = remoteIndex
|
|
||||||
m.result.HandshakeTime = payload.Time
|
m.result.HandshakeTime = payload.Time
|
||||||
m.payloadSet = true
|
m.payloadSet = true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -229,24 +229,6 @@ func TestMachineProcessPayload(t *testing.T) {
|
|||||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||||
assert.True(t, m.Failed())
|
assert.True(t, m.Failed())
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("zero initiator index on responder is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, false, 100)
|
|
||||||
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 0, Time: 1})
|
|
||||||
err := m.processPayload(bytes, msgFlags{expectsPayload: true})
|
|
||||||
require.ErrorIs(t, err, ErrInvalidRemoteIndex)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
assert.Zero(t, m.result.RemoteIndex)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("zero responder index on initiator is fatal", func(t *testing.T) {
|
|
||||||
m := newTestMachine(t, cs, v, true, 100)
|
|
||||||
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 100, ResponderIndex: 0, Time: 1})
|
|
||||||
err := m.processPayload(bytes, msgFlags{expectsPayload: true})
|
|
||||||
require.ErrorIs(t, err, ErrInvalidRemoteIndex)
|
|
||||||
assert.True(t, m.Failed())
|
|
||||||
assert.Zero(t, m.result.RemoteIndex)
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
||||||
|
|||||||
@@ -83,7 +83,6 @@ type HandshakeHostInfo struct {
|
|||||||
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
|
||||||
counter int64 // How many attempts have we made so far
|
counter int64 // How many attempts have we made so far
|
||||||
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
|
||||||
lastRelays []netip.Addr // Relays we attempted to use during the previous attempt
|
|
||||||
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
||||||
|
|
||||||
hostinfo *HostInfo
|
hostinfo *HostInfo
|
||||||
@@ -218,6 +217,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
fields := []any{
|
fields := []any{
|
||||||
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
||||||
"initiatorIndex", hh.hostinfo.localIndexId,
|
"initiatorIndex", hh.hostinfo.localIndexId,
|
||||||
|
"remoteIndex", hh.hostinfo.remoteIndexId,
|
||||||
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
||||||
}
|
}
|
||||||
// hh.machine can be nil here if buildStage0Packet never succeeded
|
// hh.machine can be nil here if buildStage0Packet never succeeded
|
||||||
@@ -323,7 +323,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
hm.f.relayManager.StartRelays(hm.f, vpnIp, hh, stage0)
|
hm.f.relayManager.StartRelays(hm.f, vpnIp, hostinfo, stage0)
|
||||||
|
|
||||||
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
||||||
if !lighthouseTriggered {
|
if !lighthouseTriggered {
|
||||||
@@ -465,6 +465,7 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
|
"remoteIndex", hostinfo.remoteIndexId,
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -487,6 +488,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// We have a collision, but this can happen since we can't control
|
// We have a collision, but this can happen since we can't control
|
||||||
// the remote ID. Just log about the situation as a note.
|
// the remote ID. Just log about the situation as a note.
|
||||||
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
|
||||||
|
"remoteIndex", hostinfo.remoteIndexId,
|
||||||
"collision", existingRemoteIndex.vpnAddrs,
|
"collision", existingRemoteIndex.vpnAddrs,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -796,6 +798,7 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
|
|||||||
}
|
}
|
||||||
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
// Don't wait for UpdateWorker
|
||||||
@@ -962,6 +965,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
hm.Complete(hostinfo, f)
|
hm.Complete(hostinfo, f)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
if len(hh.packetStore) > 0 {
|
if len(hh.packetStore) > 0 {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -970,6 +974,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
for _, cp := range hh.packetStore {
|
||||||
|
//todo use a sendbatcher
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
|||||||
+1
-6
@@ -138,9 +138,9 @@ func (rs *RelayState) InsertRelayTo(ip netip.Addr) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (rs *RelayState) CopyRelayIps() []netip.Addr {
|
func (rs *RelayState) CopyRelayIps() []netip.Addr {
|
||||||
|
ret := make([]netip.Addr, len(rs.relays))
|
||||||
rs.RLock()
|
rs.RLock()
|
||||||
defer rs.RUnlock()
|
defer rs.RUnlock()
|
||||||
ret := make([]netip.Addr, len(rs.relays))
|
|
||||||
copy(ret, rs.relays)
|
copy(ret, rs.relays)
|
||||||
return ret
|
return ret
|
||||||
}
|
}
|
||||||
@@ -623,11 +623,6 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
|||||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
|
|
||||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hm.l.Debug("Hostmap vpnIp added",
|
hm.l.Debug("Hostmap vpnIp added",
|
||||||
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,10 +10,16 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(pkt wire.TunPacket, fwPacket *firewall.Packet, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
|
// only valid until the next Read on that queue. If you must keep
|
||||||
|
// the packet, use pkt.Clone() to detach it
|
||||||
|
packet := pkt.Bytes
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -37,7 +44,10 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
_, err := f.readers[q].Write(packet)
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
_, werr := f.readers[q].Write(seg)
|
||||||
|
return werr
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -53,11 +63,23 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
|
// so retaining segments past the loop is safe.
|
||||||
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -73,10 +95,9 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,6 +107,124 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
|
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
if encErr != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
|
"error", encErr,
|
||||||
|
"udpAddr", hostinfo.remote,
|
||||||
|
"counter", c,
|
||||||
|
)
|
||||||
|
// Skip this segment; the rest of the superpacket can still
|
||||||
|
// go out — TCP will retransmit anything we drop here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
||||||
|
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
||||||
|
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||||
|
// kernel-supplied superpacket bytes never get written into a separate
|
||||||
|
// scratch arena: PerSegment builds each segment's plaintext in
|
||||||
|
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
||||||
|
// SendBatch slot.
|
||||||
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt wire.TunPacket, nb []byte, sendBatch *batch.SendBatch) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
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 counter",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hostinfo.remote.IsValid() { //the relay path
|
||||||
|
//first, find our relay hostinfo:
|
||||||
|
var relayHostInfo *HostInfo
|
||||||
|
var relay *Relay
|
||||||
|
var err error
|
||||||
|
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||||
|
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.relayState.DeleteRelay(relayIP)
|
||||||
|
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||||
|
"relay", relayIP,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if relayHostInfo == nil || relay == nil {
|
||||||
|
//failure already logged
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err = pkt.PerSegment(func(seg []byte) error {
|
||||||
|
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
|
||||||
|
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
|
||||||
|
|
||||||
|
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
|
||||||
|
if innerPacket == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
//now we need to do a relay-encrypt:
|
||||||
|
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
|
||||||
|
if err != nil {
|
||||||
|
//already logged
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sendBatch.Commit(toSend, relayHostInfo.remote, 0)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
||||||
|
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
||||||
|
|
||||||
|
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
||||||
|
if out == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sendBatch.Commit(out, hostinfo.remote, 0)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.InSendReject {
|
if !f.firewall.InSendReject {
|
||||||
return
|
return
|
||||||
@@ -275,21 +414,13 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
|
||||||
relay *Relay,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) {
|
) ([]byte, error) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -311,7 +442,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return nil, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -331,13 +462,36 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
|
// handshake messages to peers through relay tunnels.
|
||||||
|
// via is the HostInfo through which the message is relayed.
|
||||||
|
// ad is the plaintext data to authenticate, but not encrypt
|
||||||
|
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||||
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
|
// q indicates which writer to use to send the packet.
|
||||||
|
func (f *Interface) SendVia(via *HostInfo,
|
||||||
|
relay *Relay,
|
||||||
|
ad,
|
||||||
|
nb,
|
||||||
|
out []byte,
|
||||||
|
nocopy bool,
|
||||||
|
) {
|
||||||
|
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||||
|
if err != nil {
|
||||||
|
via.logger(f.l).Info("Failed to prepareSendVia", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.remote)
|
err = f.writers[0].WriteTo(toSend, via.remote)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
@@ -391,6 +545,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
"counter", c,
|
"counter", c,
|
||||||
|
"attemptedCounter", c,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
+67
-56
@@ -4,22 +4,23 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -39,6 +40,7 @@ type InterfaceConfig struct {
|
|||||||
DropLocalBroadcast bool
|
DropLocalBroadcast bool
|
||||||
DropMulticast bool
|
DropMulticast bool
|
||||||
routines int
|
routines int
|
||||||
|
batchSize int
|
||||||
MessageMetrics *MessageMetrics
|
MessageMetrics *MessageMetrics
|
||||||
version string
|
version string
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
@@ -71,6 +73,7 @@ type Interface struct {
|
|||||||
dropLocalBroadcast bool
|
dropLocalBroadcast bool
|
||||||
dropMulticast bool
|
dropMulticast bool
|
||||||
routines int
|
routines int
|
||||||
|
batchSize int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
@@ -90,8 +93,12 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []io.ReadWriteCloser
|
readers []tio.Queue
|
||||||
wg sync.WaitGroup
|
// batchers is one per tun queue, wrapping readers[i].
|
||||||
|
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||||
|
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||||
|
batchers []batch.RxBatcher
|
||||||
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
// nil means "no fatal error" (yet)
|
// nil means "no fatal error" (yet)
|
||||||
@@ -187,9 +194,11 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
dropLocalBroadcast: c.DropLocalBroadcast,
|
dropLocalBroadcast: c.DropLocalBroadcast,
|
||||||
dropMulticast: c.DropMulticast,
|
dropMulticast: c.DropMulticast,
|
||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
|
batchSize: c.batchSize,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]io.ReadWriteCloser, c.routines),
|
readers: make([]tio.Queue, c.routines),
|
||||||
|
batchers: make([]batch.RxBatcher, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -247,15 +256,17 @@ func (f *Interface) activate() error {
|
|||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
// Prepare n tun queues
|
// Prepare n tun queues
|
||||||
var reader io.ReadWriteCloser = f.inside
|
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
reader, err = f.inside.NewMultiQueueReader()
|
if err = f.inside.NewMultiQueueReader(); err != nil {
|
||||||
if err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
f.readers[i] = reader
|
}
|
||||||
|
f.readers = f.inside.Readers()
|
||||||
|
for i := range f.readers {
|
||||||
|
arena := util.NewArena(max(f.batchSize, 1) * udp.MTU)
|
||||||
|
f.batchers[i] = batch.NewPassthrough(f.readers[i], f.batchSize, arena)
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
f.wg.Add(1) // for us to wait on Close() to return
|
||||||
@@ -313,14 +324,21 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
})
|
}
|
||||||
|
|
||||||
|
flusher := func() {
|
||||||
|
if err := f.batchers[i].Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err := li.ListenOut(listener, flusher)
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
if err != nil && !f.closed.Load() {
|
||||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||||
@@ -330,28 +348,38 @@ func (f *Interface) listenOut(i int) {
|
|||||||
f.l.Debug("underlay reader is done", "reader", i)
|
f.l.Debug("underlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
func (f *Interface) listenIn(reader tio.Queue, q int) {
|
||||||
packet := make([]byte, mtu)
|
packetMem := make([]byte, mtu+16) //MTU + some leading slack space for platforms that return "bonus info"
|
||||||
out := make([]byte, mtu)
|
// TODO get the amount of bonus info from the reader
|
||||||
|
packets := make([]wire.TunPacket, 1)
|
||||||
|
rejectBuf := make([]byte, mtu)
|
||||||
|
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||||
|
sb := batch.NewSendBatch(f.writers[q], batch.SendBatchCap, util.NewArena(arenaSize))
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
n, err := reader.Read(packets, packetMem)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !f.closed.Load() {
|
if !f.closed.Load() {
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", q)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
ctCache := conntrackCache.Get()
|
||||||
|
for i := range n {
|
||||||
|
f.consumeInsidePacket(packets[i], fwPacket, nb, sb, rejectBuf, q, ctCache)
|
||||||
|
}
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", q)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||||
@@ -377,22 +405,13 @@ func (f *Interface) reloadDisconnectInvalid(c *config.C) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) reloadFirewall(c *config.C) {
|
func (f *Interface) reloadFirewall(c *config.C) {
|
||||||
cs := f.pki.getCertState()
|
//TODO: need to trigger/detect if the certificate changed too
|
||||||
curCert := cs.getCertificate(cert.Version2)
|
if c.HasChanged("firewall") == false {
|
||||||
if curCert == nil {
|
|
||||||
curCert = cs.getCertificate(cert.Version1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The firewall builds its routableNetworks set from the certificate's UnsafeNetworks at construction.
|
|
||||||
// Check to see if that set has changed, and if so, rebuild the firewall.
|
|
||||||
certUnsafeChanged := curCert != nil && !slices.Equal(curCert.UnsafeNetworks(), f.firewall.unsafeNetworks)
|
|
||||||
|
|
||||||
if !c.HasChanged("firewall") && !certUnsafeChanged {
|
|
||||||
f.l.Debug("No firewall config change detected")
|
f.l.Debug("No firewall config change detected")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(f.l, cs, c)
|
fw, err := NewFirewallFromConfig(f.l, f.pki.getCertState(), c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Error while creating firewall during reload", "error", err)
|
f.l.Error("Error while creating firewall during reload", "error", err)
|
||||||
return
|
return
|
||||||
@@ -502,34 +521,26 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
|||||||
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
|
||||||
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
|
||||||
|
|
||||||
emit := func() {
|
|
||||||
f.firewall.EmitStats()
|
|
||||||
f.handshakeManager.EmitStats()
|
|
||||||
udpStats()
|
|
||||||
|
|
||||||
certState := f.pki.getCertState()
|
|
||||||
defaultCrt := certState.GetDefaultCertificate()
|
|
||||||
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
|
||||||
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
|
||||||
|
|
||||||
// Report the max certificate version we are capable of using
|
|
||||||
if certState.v2Cert != nil {
|
|
||||||
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
|
||||||
} else {
|
|
||||||
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Prime gauges so a Prometheus scrape that lands before the first tick
|
|
||||||
// sees real values instead of the zero defaults (issue #907).
|
|
||||||
emit()
|
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
emit()
|
f.firewall.EmitStats()
|
||||||
|
f.handshakeManager.EmitStats()
|
||||||
|
udpStats()
|
||||||
|
|
||||||
|
certState := f.pki.getCertState()
|
||||||
|
defaultCrt := certState.GetDefaultCertificate()
|
||||||
|
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
||||||
|
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
||||||
|
|
||||||
|
// Report the max certificate version we are capable of using
|
||||||
|
if certState.v2Cert != nil {
|
||||||
|
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
||||||
|
} else {
|
||||||
|
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,73 +0,0 @@
|
|||||||
//go:build linux || darwin
|
|
||||||
|
|
||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
|
||||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Test_emitStats_primesGauges covers issue #907: a Prometheus scrape that
|
|
||||||
// landed before the first ticker fire used to read 0 for the cert gauges.
|
|
||||||
// emitStats now primes the gauges before entering the ticker loop. We assert
|
|
||||||
// the gauge is zero before the first call and non-zero after.
|
|
||||||
func Test_emitStats_primesGauges(t *testing.T) {
|
|
||||||
defer metrics.DefaultRegistry.UnregisterAll()
|
|
||||||
|
|
||||||
l := test.NewLogger()
|
|
||||||
hostMap := newHostMap(l)
|
|
||||||
preferredRanges := []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}
|
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
notAfter := time.Now().Add(time.Hour)
|
|
||||||
cs := &CertState{
|
|
||||||
initiatingVersion: cert.Version1,
|
|
||||||
privateKey: []byte{},
|
|
||||||
v1Cert: &dummyCert{version: cert.Version1, notAfter: notAfter},
|
|
||||||
v1Credential: nil,
|
|
||||||
}
|
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
|
||||||
ifce := &Interface{
|
|
||||||
hostMap: hostMap,
|
|
||||||
inside: &overlaytest.NoopTun{},
|
|
||||||
outside: &udp.NoopConn{},
|
|
||||||
firewall: &Firewall{Conntrack: &FirewallConntrack{Conns: map[firewall.Packet]*conn{}}},
|
|
||||||
lightHouse: lh,
|
|
||||||
pki: &PKI{},
|
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
|
||||||
// On linux, udp.NewUDPStatsEmitter indexes writers[0] and asserts to
|
|
||||||
// *udp.StdConn. A zero value works: getMemInfo sees a nil rawConn,
|
|
||||||
// returns an error, and the emitter falls through to a no-op.
|
|
||||||
writers: []udp.Conn{&udp.StdConn{}},
|
|
||||||
}
|
|
||||||
ifce.pki.cs.Store(cs)
|
|
||||||
|
|
||||||
ttlGauge := metrics.GetOrRegisterGauge("certificate.ttl_seconds", nil)
|
|
||||||
require.Zero(t, ttlGauge.Value(), "gauge should be zero before emitStats runs")
|
|
||||||
|
|
||||||
// Pre-cancel the context so emitStats returns after priming the gauges
|
|
||||||
// without ever reading from ticker.C. The one hour interval is just a
|
|
||||||
// belt-and-suspenders, the test does not expect the ticker to fire.
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
ifce.emitStats(ctx, time.Hour)
|
|
||||||
|
|
||||||
ttl := ttlGauge.Value()
|
|
||||||
assert.Positive(t, ttl, "ttl gauge should be primed by emitStats before its first tick")
|
|
||||||
assert.LessOrEqual(t, ttl, int64(3600))
|
|
||||||
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.initiating_version", nil).Value())
|
|
||||||
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.max_version", nil).Value())
|
|
||||||
}
|
|
||||||
@@ -1,120 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/config"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestReloadFirewall_CertUnsafeNetworksChanged verifies that reloadFirewall
|
|
||||||
// rebuilds the firewall when only the certificate's UnsafeNetworks have changed,
|
|
||||||
// even if the firewall section of the YAML has not.
|
|
||||||
func TestReloadFirewall_CertUnsafeNetworksChanged(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
initialUnsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
|
||||||
|
|
||||||
// dummyCert avoids dragging the real signing pipeline into a unit test.
|
|
||||||
c1 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: initialUnsafe,
|
|
||||||
}
|
|
||||||
pki := &PKI{}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
rawYAML := `firewall:
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
`
|
|
||||||
cfg := config.NewC(l)
|
|
||||||
require.NoError(t, cfg.LoadString(rawYAML))
|
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, initialUnsafe, fw.unsafeNetworks)
|
|
||||||
|
|
||||||
f := &Interface{
|
|
||||||
pki: pki,
|
|
||||||
firewall: fw,
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Swap the cert with a different UnsafeNetworks set.
|
|
||||||
newUnsafe := []netip.Prefix{
|
|
||||||
netip.MustParsePrefix("198.51.100.0/24"),
|
|
||||||
netip.MustParsePrefix("203.0.113.0/24"),
|
|
||||||
}
|
|
||||||
c2 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: newUnsafe,
|
|
||||||
}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c2, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
// Reload with the same YAML so HasChanged("firewall") reports false.
|
|
||||||
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
|
||||||
require.False(t, cfg.HasChanged("firewall"))
|
|
||||||
|
|
||||||
f.reloadFirewall(cfg)
|
|
||||||
|
|
||||||
assert.NotSame(t, fw, f.firewall, "firewall pointer should have been replaced")
|
|
||||||
assert.Equal(t, newUnsafe, f.firewall.unsafeNetworks)
|
|
||||||
assert.True(t, f.firewall.routableNetworks.Contains(netip.MustParseAddr("203.0.113.5")))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestReloadFirewall_NoChange verifies that reloadFirewall is a no-op when
|
|
||||||
// neither the firewall config nor the cert's UnsafeNetworks have changed.
|
|
||||||
func TestReloadFirewall_NoChange(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
|
|
||||||
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
|
|
||||||
unsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
|
|
||||||
|
|
||||||
c1 := &dummyCert{
|
|
||||||
version: cert.Version2,
|
|
||||||
networks: []netip.Prefix{vpnNet},
|
|
||||||
unsafeNetworks: unsafe,
|
|
||||||
}
|
|
||||||
pki := &PKI{}
|
|
||||||
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
|
|
||||||
|
|
||||||
rawYAML := `firewall:
|
|
||||||
outbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
inbound:
|
|
||||||
- port: any
|
|
||||||
proto: any
|
|
||||||
host: any
|
|
||||||
`
|
|
||||||
cfg := config.NewC(l)
|
|
||||||
require.NoError(t, cfg.LoadString(rawYAML))
|
|
||||||
|
|
||||||
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
f := &Interface{
|
|
||||||
pki: pki,
|
|
||||||
firewall: fw,
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, cfg.ReloadConfigString(rawYAML))
|
|
||||||
f.reloadFirewall(cfg)
|
|
||||||
|
|
||||||
assert.Same(t, fw, f.firewall, "firewall should not have been replaced")
|
|
||||||
}
|
|
||||||
+19
-290
@@ -4,54 +4,26 @@ import (
|
|||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet:
|
// Need 96 bytes for the largest reject packet:
|
||||||
// - 20 byte ipv4 header
|
// - 20 byte ipv4 header
|
||||||
// - 8 byte icmpv4 header
|
// - 8 byte icmpv4 header
|
||||||
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
|
||||||
maxIPv4RejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
MaxRejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8
|
||||||
|
|
||||||
// MaxRejectPacketSize is sized for the largest possible reject packet (IPv6):
|
|
||||||
// - 40 byte ipv6 header
|
|
||||||
// - 8 byte icmpv6 header
|
|
||||||
// - up to 1000 byte body (original packet, possibly truncated. We want to stay
|
|
||||||
// under the MTU with Nebula overhead included)
|
|
||||||
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
|
|
||||||
|
|
||||||
MaxRejectPacketSize = maxIPv6RejectPacketSize
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
func CreateRejectPacket(packet []byte, out []byte) []byte {
|
||||||
if len(packet) < 1 {
|
if len(packet) < ipv4.HeaderLen || int(packet[0]>>4) != ipv4.Version {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
version := int(packet[0] >> 4)
|
switch packet[9] {
|
||||||
switch version {
|
case 6: // tcp
|
||||||
case ipv4.Version:
|
return ipv4CreateRejectTCPPacket(packet, out)
|
||||||
if len(packet) < ipv4.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Do not send reject packets for non-first fragments
|
|
||||||
if packet[6]&0x1f != 0 || packet[7] != 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
switch packet[9] {
|
|
||||||
case 6: // tcp
|
|
||||||
return ipv4CreateRejectTCPPacket(packet, out)
|
|
||||||
default:
|
|
||||||
return ipv4CreateRejectICMPPacket(packet, out)
|
|
||||||
}
|
|
||||||
case ipv6.Version:
|
|
||||||
if len(packet) < ipv6.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return ipv6CreateRejectPacket(packet, out)
|
|
||||||
default:
|
default:
|
||||||
return nil
|
return ipv4CreateRejectICMPPacket(packet, out)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -63,16 +35,11 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Do not generate ICMP errors in response to ICMP error packets
|
|
||||||
if packet[9] == 1 && len(packet) > ihl {
|
|
||||||
icmpType := packet[ihl]
|
|
||||||
if icmpType == 3 || icmpType == 4 || icmpType == 5 || icmpType == 11 || icmpType == 12 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMP reply includes original header and first 8 bytes of the packet
|
// ICMP reply includes original header and first 8 bytes of the packet
|
||||||
packetLen := min(len(packet), ihl+8)
|
packetLen := len(packet)
|
||||||
|
if packetLen > ihl+8 {
|
||||||
|
packetLen = ihl + 8
|
||||||
|
}
|
||||||
|
|
||||||
outLen := ipv4.HeaderLen + 8 + packetLen
|
outLen := ipv4.HeaderLen + 8 + packetLen
|
||||||
if outLen > cap(out) {
|
if outLen > cap(out) {
|
||||||
@@ -104,14 +71,14 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
|
|||||||
|
|
||||||
// ICMP Destination Unreachable
|
// ICMP Destination Unreachable
|
||||||
icmpOut := out[ipv4.HeaderLen:]
|
icmpOut := out[ipv4.HeaderLen:]
|
||||||
icmpOut[0] = 3 // type (Destination unreachable)
|
icmpOut[0] = 3 // type (Destination unreachable)
|
||||||
icmpOut[1] = 13 // code (Communication administratively prohibited)
|
icmpOut[1] = 3 // code (Port unreachable error)
|
||||||
icmpOut[2] = 0 // checksum
|
icmpOut[2] = 0 // checksum
|
||||||
icmpOut[3] = 0 // .
|
icmpOut[3] = 0 // .
|
||||||
icmpOut[4] = 0 // unused
|
icmpOut[4] = 0 // unused
|
||||||
icmpOut[5] = 0 // .
|
icmpOut[5] = 0 // .
|
||||||
icmpOut[6] = 0 // .
|
icmpOut[6] = 0 // .
|
||||||
icmpOut[7] = 0 // .
|
icmpOut[7] = 0 // .
|
||||||
|
|
||||||
// Copy original IP header and first 8 bytes as body
|
// Copy original IP header and first 8 bytes as body
|
||||||
copy(icmpOut[8:], packet[:packetLen])
|
copy(icmpOut[8:], packet[:packetLen])
|
||||||
@@ -198,193 +165,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func ipv6CreateRejectPacket(packet []byte, out []byte) []byte {
|
|
||||||
proto, offset, isFragment := ipv6FindUpperProtocol(packet)
|
|
||||||
if isFragment {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
switch proto {
|
|
||||||
case 6: // tcp
|
|
||||||
return ipv6CreateRejectTCPPacket(packet, out, offset)
|
|
||||||
default:
|
|
||||||
return ipv6CreateRejectICMPPacket(packet, out, proto, offset)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6CreateRejectICMPPacket(packet []byte, out []byte, proto uint8, offset int) []byte {
|
|
||||||
// Do not generate ICMPv6 errors in response to ICMPv6 error packets
|
|
||||||
if proto == 58 && len(packet) > offset {
|
|
||||||
icmpType := packet[offset]
|
|
||||||
if icmpType >= 1 && icmpType <= 4 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Include as much of the original packet as possible, up to 1000 bytes,
|
|
||||||
// so the response fits comfortably within any tunnel MTU.
|
|
||||||
packetLen := min(len(packet), 1000)
|
|
||||||
|
|
||||||
outLen := ipv6.HeaderLen + 8 + packetLen
|
|
||||||
if outLen > cap(out) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:outLen]
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
ipHdr := out[0:ipv6.HeaderLen]
|
|
||||||
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
|
||||||
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
|
||||||
ipHdr[2] = 0 // flow label
|
|
||||||
ipHdr[3] = 0 // flow label
|
|
||||||
|
|
||||||
payloadLen := uint16(outLen - ipv6.HeaderLen)
|
|
||||||
binary.BigEndian.PutUint16(ipHdr[4:], payloadLen) // payload length
|
|
||||||
ipHdr[6] = 58 // next header (ICMPv6)
|
|
||||||
ipHdr[7] = 64 // hop limit
|
|
||||||
|
|
||||||
// Swap dest / src IPs (each 16 bytes, src at 8, dst at 24)
|
|
||||||
copy(ipHdr[8:24], packet[24:40])
|
|
||||||
copy(ipHdr[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// ICMPv6 Destination Unreachable
|
|
||||||
icmpOut := out[ipv6.HeaderLen:]
|
|
||||||
icmpOut[0] = 1 // type (Destination Unreachable)
|
|
||||||
icmpOut[1] = 1 // code (Communication with destination administratively prohibited)
|
|
||||||
icmpOut[2] = 0 // checksum
|
|
||||||
icmpOut[3] = 0 // .
|
|
||||||
icmpOut[4] = 0 // unused
|
|
||||||
icmpOut[5] = 0 // .
|
|
||||||
icmpOut[6] = 0 // .
|
|
||||||
icmpOut[7] = 0 // .
|
|
||||||
|
|
||||||
copy(icmpOut[8:], packet[:packetLen])
|
|
||||||
|
|
||||||
// ICMPv6 checksum uses a pseudo-header
|
|
||||||
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 58, uint32(payloadLen))
|
|
||||||
binary.BigEndian.PutUint16(icmpOut[2:], tcpipChecksum(icmpOut, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte {
|
|
||||||
const tcpLen = 20
|
|
||||||
|
|
||||||
if len(packet) < offset+tcpLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
outLen := ipv6.HeaderLen + tcpLen
|
|
||||||
if outLen > cap(out) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:outLen]
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
ipHdr := out[0:ipv6.HeaderLen]
|
|
||||||
ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits)
|
|
||||||
ipHdr[1] = 0 // traffic class (low bits), flow label (high bits)
|
|
||||||
ipHdr[2] = 0 // flow label
|
|
||||||
ipHdr[3] = 0 // flow label
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(ipHdr[4:], tcpLen) // payload length
|
|
||||||
ipHdr[6] = 6 // next header (TCP)
|
|
||||||
ipHdr[7] = 64 // hop limit
|
|
||||||
|
|
||||||
// Swap dest / src IPs
|
|
||||||
copy(ipHdr[8:24], packet[24:40])
|
|
||||||
copy(ipHdr[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// TCP RST
|
|
||||||
tcpIn := packet[offset:]
|
|
||||||
var ackSeq, seq uint32
|
|
||||||
outFlags := byte(0b00000100) // RST
|
|
||||||
|
|
||||||
inAck := tcpIn[13]&0b00010000 != 0
|
|
||||||
if inAck {
|
|
||||||
seq = binary.BigEndian.Uint32(tcpIn[8:])
|
|
||||||
} else {
|
|
||||||
inSyn := uint32((tcpIn[13] & 0b00000010) >> 1)
|
|
||||||
inFin := uint32(tcpIn[13] & 0b00000001)
|
|
||||||
ackSeq = binary.BigEndian.Uint32(tcpIn[4:]) + inSyn + inFin + uint32(len(tcpIn)) - uint32(tcpIn[12]>>4)<<2
|
|
||||||
outFlags |= 0b00010000 // ACK
|
|
||||||
}
|
|
||||||
|
|
||||||
tcpOut := out[ipv6.HeaderLen:]
|
|
||||||
// Swap dest / src ports
|
|
||||||
copy(tcpOut[0:2], tcpIn[2:4])
|
|
||||||
copy(tcpOut[2:4], tcpIn[0:2])
|
|
||||||
binary.BigEndian.PutUint32(tcpOut[4:], seq)
|
|
||||||
binary.BigEndian.PutUint32(tcpOut[8:], ackSeq)
|
|
||||||
tcpOut[12] = (tcpLen >> 2) << 4 // data offset, reserved, NS
|
|
||||||
tcpOut[13] = outFlags // CWR, ECE, URG, ACK, PSH, RST, SYN, FIN
|
|
||||||
tcpOut[14] = 0 // window size
|
|
||||||
tcpOut[15] = 0 // .
|
|
||||||
tcpOut[16] = 0 // checksum
|
|
||||||
tcpOut[17] = 0 // .
|
|
||||||
tcpOut[18] = 0 // URG Pointer
|
|
||||||
tcpOut[19] = 0 // .
|
|
||||||
|
|
||||||
// Calculate checksum with IPv6 pseudo-header
|
|
||||||
csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 6, tcpLen)
|
|
||||||
binary.BigEndian.PutUint16(tcpOut[16:], tcpipChecksum(tcpOut, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) {
|
|
||||||
nextHeader = packet[6]
|
|
||||||
offset = ipv6.HeaderLen
|
|
||||||
|
|
||||||
for {
|
|
||||||
switch nextHeader {
|
|
||||||
case 0, 43, 60: // Hop-by-Hop, Routing, Destination
|
|
||||||
if len(packet) < offset+2 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += int(packet[offset+1]+1) << 3
|
|
||||||
|
|
||||||
case 44: // Fragment
|
|
||||||
if len(packet) < offset+8 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 {
|
|
||||||
isFragment = true
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += 8
|
|
||||||
|
|
||||||
case 51: // AH
|
|
||||||
if len(packet) < offset+2 {
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
nextHeader = packet[offset]
|
|
||||||
offset += int(packet[offset+1]+2) << 2
|
|
||||||
|
|
||||||
default:
|
|
||||||
return nextHeader, offset, isFragment
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
func CreateICMPEchoResponse(packet, out []byte) []byte {
|
||||||
if len(packet) < 1 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
switch packet[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
return createICMPv4EchoResponse(packet, out)
|
|
||||||
case 6:
|
|
||||||
return createICMPv6EchoResponse(packet, out)
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func createICMPv4EchoResponse(packet, out []byte) []byte {
|
|
||||||
// Return early if this is not a simple ICMP Echo Request
|
// Return early if this is not a simple ICMP Echo Request
|
||||||
//TODO: make constants out of these
|
//TODO: make constants out of these
|
||||||
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
|
||||||
@@ -418,43 +199,6 @@ func createICMPv4EchoResponse(packet, out []byte) []byte {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func createICMPv6EchoResponse(packet, out []byte) []byte {
|
|
||||||
// IPv6 header (40 bytes) + ICMPv6 header (8 bytes minimum)
|
|
||||||
if len(packet) < ipv6.HeaderLen+8 || len(packet) > 9001 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next Header must be ICMPv6 (58)
|
|
||||||
if packet[6] != 58 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMPv6 type must be Echo Request (128)
|
|
||||||
if packet[ipv6.HeaderLen] != 128 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
out = out[:len(packet)]
|
|
||||||
copy(out, packet)
|
|
||||||
|
|
||||||
// Swap src/dst addresses (bytes 8-23 and 24-39)
|
|
||||||
copy(out[8:24], packet[24:40])
|
|
||||||
copy(out[24:40], packet[8:24])
|
|
||||||
|
|
||||||
// Change ICMPv6 type to Echo Reply (129)
|
|
||||||
icmp := out[ipv6.HeaderLen:]
|
|
||||||
icmp[0] = 129
|
|
||||||
icmp[2] = 0
|
|
||||||
icmp[3] = 0
|
|
||||||
|
|
||||||
// ICMPv6 checksum uses a pseudo-header with src, dst, length, and next header
|
|
||||||
payloadLen := uint32(len(icmp))
|
|
||||||
csum := ipv6PseudoheaderChecksum(out[8:24], out[24:40], 58, payloadLen)
|
|
||||||
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
// calculates the TCP/IP checksum defined in rfc1071. The passed-in
|
||||||
// csum is any initial checksum data that's already been computed.
|
// csum is any initial checksum data that's already been computed.
|
||||||
//
|
//
|
||||||
@@ -492,18 +236,3 @@ func ipv4PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint3
|
|||||||
csum += length >> 16
|
csum += length >> 16
|
||||||
return csum
|
return csum
|
||||||
}
|
}
|
||||||
|
|
||||||
// based on:
|
|
||||||
// - https://github.com/google/gopacket/blob/v1.1.19/layers/tcpip.go#L37-L48
|
|
||||||
func ipv6PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint32) {
|
|
||||||
for i := 0; i < 16; i += 2 {
|
|
||||||
csum += uint32(src[i]) << 8
|
|
||||||
csum += uint32(src[i+1])
|
|
||||||
csum += uint32(dst[i]) << 8
|
|
||||||
csum += uint32(dst[i+1])
|
|
||||||
}
|
|
||||||
csum += proto
|
|
||||||
csum += length & 0xffff
|
|
||||||
csum += length >> 16
|
|
||||||
return csum
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-404
@@ -1,13 +1,11 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"net"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
"golang.org/x/net/ipv6"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_CreateRejectPacket(t *testing.T) {
|
func Test_CreateRejectPacket(t *testing.T) {
|
||||||
@@ -45,7 +43,7 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
|
||||||
|
|
||||||
expectedLen = maxIPv4RejectPacketSize
|
expectedLen = MaxRejectPacketSize
|
||||||
out = make([]byte, MaxRejectPacketSize)
|
out = make([]byte, MaxRejectPacketSize)
|
||||||
rejectPacket = CreateRejectPacket(b, out)
|
rejectPacket = CreateRejectPacket(b, out)
|
||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
@@ -73,404 +71,3 @@ func Test_CreateRejectPacket(t *testing.T) {
|
|||||||
assert.NotNil(t, rejectPacket)
|
assert.NotNil(t, rejectPacket)
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
assert.Len(t, rejectPacket, expectedLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_CreateRejectPacket_NoFragment(t *testing.T) {
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// IPv4: non-zero fragment offset should not generate reject packet
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 17, // UDP
|
|
||||||
}
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, make([]byte, 8)...)
|
|
||||||
// Set fragment offset to non-zero (byte 6-7, offset in 8-byte units)
|
|
||||||
b[6] = 0x00
|
|
||||||
b[7] = 0x01
|
|
||||||
assert.Nil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// MF flag with zero offset (first fragment) should still generate reject
|
|
||||||
b[6] = 0x20 // MF flag set
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// Non-fragment should still generate reject packet
|
|
||||||
b[6] = 0x00
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
|
|
||||||
// DF flag only (not a fragment) should still generate reject packet
|
|
||||||
b[6] = 0x40
|
|
||||||
b[7] = 0x00
|
|
||||||
assert.NotNil(t, CreateRejectPacket(b, out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_NoFragment(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// IPv6 with Fragment header and non-zero offset should not generate reject
|
|
||||||
fragHeader := []byte{
|
|
||||||
17, // next header: UDP
|
|
||||||
0, // reserved
|
|
||||||
0, 9, // fragment offset=1 (shifted left 3), M=1
|
|
||||||
0, 0, 0, 1, // identification
|
|
||||||
}
|
|
||||||
udpPayload := make([]byte, 8)
|
|
||||||
payload := append(fragHeader, udpPayload...)
|
|
||||||
packet := makeIPv6Packet(src, dst, 44, payload) // next header 44 = Fragment
|
|
||||||
assert.Nil(t, CreateRejectPacket(packet, out))
|
|
||||||
|
|
||||||
// Fragment header with zero offset (first fragment) should still generate reject
|
|
||||||
fragHeader[2] = 0
|
|
||||||
fragHeader[3] = 1 // offset=0, M=1
|
|
||||||
payload = append(fragHeader, udpPayload...)
|
|
||||||
packet = makeIPv6Packet(src, dst, 44, payload)
|
|
||||||
assert.NotNil(t, CreateRejectPacket(packet, out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacket_NoICMPError(t *testing.T) {
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// ICMP error types should not generate reject packets
|
|
||||||
icmpErrorTypes := []byte{3, 4, 5, 11, 12}
|
|
||||||
for _, icmpType := range icmpErrorTypes {
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 1, // ICMP
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(b, out)
|
|
||||||
assert.Nil(t, rejectPacket, "ICMP type %d should not generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMP non-error types should still generate reject packets
|
|
||||||
icmpNonErrorTypes := []byte{0, 8, 13, 14}
|
|
||||||
for _, icmpType := range icmpNonErrorTypes {
|
|
||||||
h := ipv4.Header{
|
|
||||||
Len: 20,
|
|
||||||
Src: net.IPv4(10, 0, 0, 1),
|
|
||||||
Dst: net.IPv4(10, 0, 0, 2),
|
|
||||||
Protocol: 1, // ICMP
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := h.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("h.Marshal: %v", err)
|
|
||||||
}
|
|
||||||
b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(b, out)
|
|
||||||
assert.NotNil(t, rejectPacket, "ICMP type %d should generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte {
|
|
||||||
b := make([]byte, ipv6.HeaderLen+len(payload))
|
|
||||||
b[0] = ipv6.Version << 4
|
|
||||||
binary.BigEndian.PutUint16(b[4:], uint16(len(payload)))
|
|
||||||
b[6] = nextHeader
|
|
||||||
b[7] = 64
|
|
||||||
copy(b[8:24], src.To16())
|
|
||||||
copy(b[24:40], dst.To16())
|
|
||||||
copy(b[ipv6.HeaderLen:], payload)
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_ICMP(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// Small UDP packet: entire original included in body
|
|
||||||
udpPayload := make([]byte, 20)
|
|
||||||
udpPayload[0] = 0x00 // src port high
|
|
||||||
udpPayload[1] = 0x50 // src port low (80)
|
|
||||||
udpPayload[2] = 0x01 // dst port high
|
|
||||||
udpPayload[3] = 0xBB // dst port low (443)
|
|
||||||
packet := makeIPv6Packet(src, dst, 17, udpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Small packet fits entirely: 40 (ipv6 hdr) + 8 (icmpv6 hdr) + 60 (original)
|
|
||||||
expectedLen := ipv6.HeaderLen + 8 + len(packet)
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
|
|
||||||
// Verify version
|
|
||||||
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
|
||||||
// Verify next header is ICMPv6 (58)
|
|
||||||
assert.Equal(t, byte(58), rejectPacket[6])
|
|
||||||
// Verify src/dst are swapped
|
|
||||||
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
|
||||||
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
|
||||||
// Verify ICMPv6 type=1 (Dest Unreachable), code=1 (Administratively prohibited)
|
|
||||||
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen])
|
|
||||||
assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen+1])
|
|
||||||
// Verify entire original packet is included in body
|
|
||||||
assert.Equal(t, packet, rejectPacket[ipv6.HeaderLen+8:])
|
|
||||||
|
|
||||||
// Large packet: body is truncated to 1000 bytes
|
|
||||||
largePkt := makeIPv6Packet(src, dst, 17, make([]byte, 1200))
|
|
||||||
rejectPacket = CreateRejectPacket(largePkt, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
assert.Len(t, rejectPacket, ipv6.HeaderLen+8+1000)
|
|
||||||
assert.Equal(t, largePkt[:1000], rejectPacket[ipv6.HeaderLen+8:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TCP(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// TCP SYN packet (next header 6)
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00 // src port high
|
|
||||||
tcpPayload[1] = 0x50 // src port low (80)
|
|
||||||
tcpPayload[2] = 0x01 // dst port high
|
|
||||||
tcpPayload[3] = 0xBB // dst port low (443)
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 0) // ack seq
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
|
||||||
tcpPayload[13] = 0b00000010 // SYN flag
|
|
||||||
|
|
||||||
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Expected: 40 (ipv6 hdr) + 20 (tcp RST)
|
|
||||||
expectedLen := ipv6.HeaderLen + 20
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
|
|
||||||
// Verify version
|
|
||||||
assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0)
|
|
||||||
// Verify next header is TCP (6)
|
|
||||||
assert.Equal(t, byte(6), rejectPacket[6])
|
|
||||||
// Verify src/dst are swapped
|
|
||||||
assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24]))
|
|
||||||
assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40]))
|
|
||||||
// Verify ports are swapped
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
assert.Equal(t, uint16(443), binary.BigEndian.Uint16(tcpOut[0:2]))
|
|
||||||
assert.Equal(t, uint16(80), binary.BigEndian.Uint16(tcpOut[2:4]))
|
|
||||||
// RST+ACK flags (since input was SYN without ACK)
|
|
||||||
assert.Equal(t, byte(0b00010100), tcpOut[13])
|
|
||||||
// ack_seq = original seq (1000) + SYN (1) + FIN (0) + segment data (0)
|
|
||||||
assert.Equal(t, uint32(1001), binary.BigEndian.Uint32(tcpOut[8:]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TCPWithACK(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// TCP packet with ACK set
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00
|
|
||||||
tcpPayload[1] = 0x50
|
|
||||||
tcpPayload[2] = 0x01
|
|
||||||
tcpPayload[3] = 0xBB
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 2000) // ack seq
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4 // data offset
|
|
||||||
tcpPayload[13] = 0b00010000 // ACK flag
|
|
||||||
|
|
||||||
packet := makeIPv6Packet(src, dst, 6, tcpPayload)
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
// RST only (no ACK) since input had ACK
|
|
||||||
assert.Equal(t, byte(0b00000100), tcpOut[13])
|
|
||||||
// seq = original ack_seq
|
|
||||||
assert.Equal(t, uint32(2000), binary.BigEndian.Uint32(tcpOut[4:]))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_NoICMPError(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
|
|
||||||
// ICMPv6 error types (1-4) should not generate reject packets
|
|
||||||
for icmpType := byte(1); icmpType <= 4; icmpType++ {
|
|
||||||
payload := make([]byte, 8)
|
|
||||||
payload[0] = icmpType
|
|
||||||
packet := makeIPv6Packet(src, dst, 58, payload)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.Nil(t, rejectPacket, "ICMPv6 type %d should not generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ICMPv6 non-error types should still generate reject packets
|
|
||||||
nonErrorTypes := []byte{128, 129, 133, 134}
|
|
||||||
for _, icmpType := range nonErrorTypes {
|
|
||||||
payload := make([]byte, 8)
|
|
||||||
payload[0] = icmpType
|
|
||||||
packet := makeIPv6Packet(src, dst, 58, payload)
|
|
||||||
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket, "ICMPv6 type %d should generate a reject packet", icmpType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_TooShort(t *testing.T) {
|
|
||||||
// Packet too short to be valid IPv6
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
assert.Nil(t, CreateRejectPacket([]byte{0x60}, out))
|
|
||||||
assert.Nil(t, CreateRejectPacket(make([]byte, 39), out))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_CreateRejectPacketIPv6_ExtensionHeaders(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1")
|
|
||||||
dst := net.ParseIP("fd00::2")
|
|
||||||
|
|
||||||
// IPv6 + Hop-by-Hop extension header + TCP
|
|
||||||
hopByHop := []byte{
|
|
||||||
6, // next header: TCP
|
|
||||||
0, // length (8 bytes total)
|
|
||||||
0, 0, // padding
|
|
||||||
0, 0, 0, 0,
|
|
||||||
}
|
|
||||||
tcpPayload := make([]byte, 20)
|
|
||||||
tcpPayload[0] = 0x00
|
|
||||||
tcpPayload[1] = 0x50
|
|
||||||
tcpPayload[2] = 0x01
|
|
||||||
tcpPayload[3] = 0xBB
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[4:], 1000)
|
|
||||||
binary.BigEndian.PutUint32(tcpPayload[8:], 2000)
|
|
||||||
tcpPayload[12] = (20 >> 2) << 4
|
|
||||||
tcpPayload[13] = 0b00010000 // ACK
|
|
||||||
|
|
||||||
payload := append(hopByHop, tcpPayload...)
|
|
||||||
packet := makeIPv6Packet(src, dst, 0, payload) // next header 0 = Hop-by-Hop
|
|
||||||
|
|
||||||
out := make([]byte, MaxRejectPacketSize)
|
|
||||||
rejectPacket := CreateRejectPacket(packet, out)
|
|
||||||
assert.NotNil(t, rejectPacket)
|
|
||||||
|
|
||||||
// Should produce TCP RST
|
|
||||||
expectedLen := ipv6.HeaderLen + 20
|
|
||||||
assert.Len(t, rejectPacket, expectedLen)
|
|
||||||
assert.Equal(t, byte(6), rejectPacket[6]) // next header is TCP
|
|
||||||
tcpOut := rejectPacket[ipv6.HeaderLen:]
|
|
||||||
assert.Equal(t, byte(0b00000100), tcpOut[13]) // RST only
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv4(t *testing.T) {
|
|
||||||
// Build a simple IPv4 ICMP Echo Request
|
|
||||||
packet := make([]byte, 28)
|
|
||||||
packet[0] = 0x45 // version 4, IHL 5
|
|
||||||
binary.BigEndian.PutUint16(packet[2:], uint16(28)) // total length
|
|
||||||
packet[8] = 64 // TTL
|
|
||||||
packet[9] = 1 // protocol ICMP
|
|
||||||
copy(packet[12:16], net.IPv4(10, 0, 0, 1).To4()) // src
|
|
||||||
copy(packet[16:20], net.IPv4(10, 0, 0, 2).To4()) // dst
|
|
||||||
packet[20] = 8 // ICMP Echo Request
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.NotNil(t, result)
|
|
||||||
assert.Equal(t, byte(0x45), result[0])
|
|
||||||
// src/dst swapped
|
|
||||||
assert.Equal(t, net.IPv4(10, 0, 0, 2).To4(), net.IP(result[12:16]))
|
|
||||||
assert.Equal(t, net.IPv4(10, 0, 0, 1).To4(), net.IP(result[16:20]))
|
|
||||||
// ICMP Echo Reply
|
|
||||||
assert.Equal(t, byte(0), result[20])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
// Build an IPv6 ICMPv6 Echo Request packet
|
|
||||||
// IPv6 header (40 bytes) + ICMPv6 (8 bytes)
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60 // version 6
|
|
||||||
payloadLen := uint16(8) // ICMPv6 header only
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], payloadLen)
|
|
||||||
packet[6] = 58 // Next Header: ICMPv6
|
|
||||||
packet[7] = 64 // Hop Limit
|
|
||||||
copy(packet[8:24], src) // src address
|
|
||||||
copy(packet[24:40], dst) // dst address
|
|
||||||
|
|
||||||
// ICMPv6 Echo Request
|
|
||||||
icmp := packet[40:]
|
|
||||||
icmp[0] = 128 // type: Echo Request
|
|
||||||
icmp[1] = 0 // code
|
|
||||||
binary.BigEndian.PutUint16(icmp[4:], 1) // identifier
|
|
||||||
binary.BigEndian.PutUint16(icmp[6:], 1) // sequence number
|
|
||||||
|
|
||||||
// Compute correct checksum for the request
|
|
||||||
csum := ipv6PseudoheaderChecksum(src, dst, 58, uint32(payloadLen))
|
|
||||||
binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum))
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.NotNil(t, result)
|
|
||||||
|
|
||||||
// Version should still be 6
|
|
||||||
assert.Equal(t, byte(6), result[0]>>4)
|
|
||||||
// src/dst swapped
|
|
||||||
assert.Equal(t, dst, net.IP(result[8:24]))
|
|
||||||
assert.Equal(t, src, net.IP(result[24:40]))
|
|
||||||
// ICMPv6 Echo Reply type
|
|
||||||
assert.Equal(t, byte(129), result[40])
|
|
||||||
|
|
||||||
// Verify checksum is valid (tcpipChecksum returns 0 when data+checksum is correct)
|
|
||||||
respIcmp := result[40:]
|
|
||||||
verifyCsum := ipv6PseudoheaderChecksum(result[8:24], result[24:40], 58, uint32(payloadLen))
|
|
||||||
assert.Equal(t, uint16(0), tcpipChecksum(respIcmp, verifyCsum))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6_NotEchoRequest(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], 8)
|
|
||||||
packet[6] = 58
|
|
||||||
packet[7] = 64
|
|
||||||
copy(packet[8:24], src)
|
|
||||||
copy(packet[24:40], dst)
|
|
||||||
|
|
||||||
// ICMPv6 type 1 (Destination Unreachable) - not Echo Request
|
|
||||||
packet[40] = 1
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.Nil(t, result)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) {
|
|
||||||
src := net.ParseIP("fd00::1").To16()
|
|
||||||
dst := net.ParseIP("fd00::2").To16()
|
|
||||||
|
|
||||||
packet := make([]byte, 48)
|
|
||||||
packet[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(packet[4:], 8)
|
|
||||||
packet[6] = 6 // TCP, not ICMPv6
|
|
||||||
packet[7] = 64
|
|
||||||
copy(packet[8:24], src)
|
|
||||||
copy(packet[24:40], dst)
|
|
||||||
|
|
||||||
out := make([]byte, len(packet))
|
|
||||||
result := CreateICMPEchoResponse(packet, out)
|
|
||||||
assert.Nil(t, result)
|
|
||||||
}
|
|
||||||
|
|||||||
+4
-21
@@ -272,18 +272,16 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
|
||||||
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
|
||||||
// Clean up. Entries still in the static_host_map will be re-built.
|
// Clean up. Entries still in the static_host_map will be re-built.
|
||||||
ourselves := lh.myVpnNetworks[0].Addr()
|
// Entries no longer present must have their (possible) background DNS goroutines stopped.
|
||||||
oldStaticList := lh.staticList.Load()
|
if existingStaticList := lh.staticList.Load(); existingStaticList != nil {
|
||||||
if oldStaticList != nil {
|
|
||||||
lh.RLock()
|
lh.RLock()
|
||||||
for staticVpnAddr := range *oldStaticList {
|
for staticVpnAddr := range *existingStaticList {
|
||||||
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
||||||
am.ResetForOwner(ourselves)
|
am.hr.Cancel()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
lh.RUnlock()
|
lh.RUnlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build a new list based on current config.
|
// Build a new list based on current config.
|
||||||
staticList := make(map[netip.Addr]struct{})
|
staticList := make(map[netip.Addr]struct{})
|
||||||
err := lh.loadStaticMap(c, staticList)
|
err := lh.loadStaticMap(c, staticList)
|
||||||
@@ -291,21 +289,6 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// For entries removed from static_host_map, stop the DNS goroutine and drop the cached addrs.
|
|
||||||
// All addrs must come from the lighthouses now that it's no longer a static host.
|
|
||||||
if oldStaticList != nil {
|
|
||||||
lh.RLock()
|
|
||||||
for staticVpnAddr := range *oldStaticList {
|
|
||||||
if _, stillStatic := staticList[staticVpnAddr]; stillStatic {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
|
|
||||||
am.ClearHostnameResults()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
lh.RUnlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
if !initial {
|
if !initial {
|
||||||
if c.HasChanged("static_host_map") {
|
if c.HasChanged("static_host_map") {
|
||||||
|
|||||||
@@ -303,132 +303,6 @@ func TestLighthouse_reload(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestLighthouse_reloadStaticHostMap verifies that reloading static_host_map applies the new
|
|
||||||
// config rather than appending to it. See issue #718.
|
|
||||||
func TestLighthouse_reloadStaticHostMap(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
c := config.NewC(l)
|
|
||||||
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
|
|
||||||
c.Settings["listen"] = map[string]any{"port": 4242}
|
|
||||||
c.Settings["static_host_map"] = map[string]any{
|
|
||||||
"10.128.0.2": []any{"1.1.1.1:4242"},
|
|
||||||
}
|
|
||||||
|
|
||||||
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
|
|
||||||
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)
|
|
||||||
|
|
||||||
staticHost := netip.MustParseAddr("10.128.0.2")
|
|
||||||
otherHost := netip.MustParseAddr("10.128.0.3")
|
|
||||||
|
|
||||||
// Capture the RemoteList pointer up front; an in-flight handshake would hold the same one
|
|
||||||
// on hostinfo.remotes, so it must reflect every reload below.
|
|
||||||
pinned := lh.Query(staticHost)
|
|
||||||
require.NotNil(t, pinned)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, pinned.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Replace the remote address. The new address should be the only entry.
|
|
||||||
nc := map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err := yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl := lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl, "RemoteList pointer must stay stable so in-flight handshakes pick up the change")
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Reload back to the original IP. Mirrors the round-trip in issue #718 step 6-8 where
|
|
||||||
// the buggy reload produced [1.1.1.1, 2.2.2.2, 1.1.1.1] instead of [1.1.1.1].
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"1.1.1.1:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Reload with the same config. An unchanged entry must not duplicate.
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Switch back to 2.2.2.2 so the rest of the test continues against a known address.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
// Add a second host alongside the first. Both should be present, neither duplicated.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.2": []any{"2.2.2.2:4242"},
|
|
||||||
"10.128.0.3": []any{"3.3.3.3:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl, "adding a sibling entry must not displace the existing RemoteList")
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
rl = lh.Query(otherHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
// Drop the first host entirely. The vpnAddr is no longer marked static, our owner
|
|
||||||
// contribution is cleared, but the addrMap entry stays in place so non-static cache
|
|
||||||
// data (from lighthouse queries) on the same RemoteList isn't lost. In-flight handshakes
|
|
||||||
// that already had the pointer see an empty address list rather than retrying stale ones.
|
|
||||||
nc = map[string]any{
|
|
||||||
"static_host_map": map[string]any{
|
|
||||||
"10.128.0.3": []any{"3.3.3.3:4242"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
rc, err = yaml.Marshal(nc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, c.ReloadConfigString(string(rc)))
|
|
||||||
|
|
||||||
_, isStatic := lh.GetStaticHostList()[staticHost]
|
|
||||||
assert.False(t, isStatic)
|
|
||||||
|
|
||||||
rl = lh.Query(staticHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Same(t, pinned, rl)
|
|
||||||
assert.Empty(t, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
|
|
||||||
rl = lh.Query(otherHost)
|
|
||||||
require.NotNil(t, rl)
|
|
||||||
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
|
||||||
req := &NebulaMeta{
|
req := &NebulaMeta{
|
||||||
Type: NebulaMeta_HostQuery,
|
Type: NebulaMeta_HostQuery,
|
||||||
|
|||||||
@@ -194,7 +194,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
||||||
lightHouse.handshakeTrigger = handshakeManager.trigger
|
lightHouse.handshakeTrigger = handshakeManager.trigger
|
||||||
|
|
||||||
ds, err := newDnsServerFromConfig(ctx, l, pki, hostMap, c)
|
ds, err := newDnsServerFromConfig(ctx, l, pki.getCertState(), hostMap, c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
l.Warn("Failed to start DNS responder", "error", err)
|
l.Warn("Failed to start DNS responder", "error", err)
|
||||||
}
|
}
|
||||||
@@ -215,6 +215,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
DropLocalBroadcast: c.GetBool("tun.drop_local_broadcast", false),
|
DropLocalBroadcast: c.GetBool("tun.drop_local_broadcast", false),
|
||||||
DropMulticast: c.GetBool("tun.drop_multicast", false),
|
DropMulticast: c.GetBool("tun.drop_multicast", false),
|
||||||
routines: routines,
|
routines: routines,
|
||||||
|
batchSize: c.GetInt("listen.batch", 64),
|
||||||
MessageMetrics: messageMetrics,
|
MessageMetrics: messageMetrics,
|
||||||
version: buildVersion,
|
version: buildVersion,
|
||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
|
|||||||
@@ -147,6 +147,62 @@ func buildCipherStatesB(b *testing.B, c noise.CipherFunc) (*noise.CipherState, *
|
|||||||
return eI, dR
|
return eI, dR
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestDecryptDangerRelayShapeNoAlloc covers the AD-only relay path used in
|
||||||
|
// outside.go's handleOutsideRelayPacket: the body is AD, the trailing 16 bytes
|
||||||
|
// are the AEAD tag, the plaintext is empty, and the caller passes nil as the
|
||||||
|
// destination because it only needs the auth side-effect. The call must
|
||||||
|
// succeed, return an empty plaintext, and not allocate on the hot path.
|
||||||
|
func TestDecryptDangerRelayShapeNoAlloc(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
c noise.CipherFunc
|
||||||
|
wrap func(*noise.CipherState) CipherState
|
||||||
|
}{
|
||||||
|
{"AESGCM", CipherAESGCM, func(cs *noise.CipherState) CipherState { return NewCipherStateAESGCM(cs) }},
|
||||||
|
{"ChaChaPoly", noise.CipherChaChaPoly, func(cs *noise.CipherState) CipherState { return NewCipherStateChaChaPoly(cs) }},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
encCS, decCS := buildCipherStates(t, tc.c)
|
||||||
|
enc, dec := tc.wrap(encCS), tc.wrap(decCS)
|
||||||
|
|
||||||
|
ad := make([]byte, 1200) // typical relay packet body size
|
||||||
|
for i := range ad {
|
||||||
|
ad[i] = byte(i)
|
||||||
|
}
|
||||||
|
nb := make([]byte, 12)
|
||||||
|
|
||||||
|
// Build the "signature value" the way handleOutsideRelayPacket sees it:
|
||||||
|
// empty plaintext encrypted with the body as AD yields just the 16-byte tag.
|
||||||
|
tag, err := enc.EncryptDanger(nil, ad, nil, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, tag, dec.Overhead())
|
||||||
|
|
||||||
|
// Sanity: the relay-shaped call returns empty plaintext, no error.
|
||||||
|
out, err := dec.DecryptDanger(nil, ad, tag, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, out)
|
||||||
|
|
||||||
|
// Tampering with the AD must fail authentication.
|
||||||
|
ad[0] ^= 0xff
|
||||||
|
_, err = dec.DecryptDanger(nil, ad, tag, 1, nb)
|
||||||
|
require.Error(t, err)
|
||||||
|
ad[0] ^= 0xff
|
||||||
|
|
||||||
|
// The hot path must not allocate. AllocsPerRun does a warm-up run, so any
|
||||||
|
// one-time setup is excluded. Counter has to advance so the AEAD nonce is
|
||||||
|
// unique per call, but we don't care whether the auth succeeds — we only
|
||||||
|
// care about whether the call path allocates.
|
||||||
|
var counter uint64 = 2
|
||||||
|
allocs := testing.AllocsPerRun(100, func() {
|
||||||
|
_, _ = dec.DecryptDanger(nil, ad, tag, counter, nb)
|
||||||
|
counter++
|
||||||
|
})
|
||||||
|
assert.Equal(t, 0.0, allocs, "DecryptDanger(nil, ...) must not allocate")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCipherStateNilSafety(t *testing.T) {
|
func TestCipherStateNilSafety(t *testing.T) {
|
||||||
var aes *CipherStateAESGCM
|
var aes *CipherStateAESGCM
|
||||||
_, err := aes.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
_, err := aes.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||||
|
|||||||
+15
-19
@@ -22,7 +22,7 @@ const (
|
|||||||
|
|
||||||
var ErrOutOfWindow = errors.New("out of window packet")
|
var ErrOutOfWindow = errors.New("out of window packet")
|
||||||
|
|
||||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) readOutsidePackets(via ViaSender, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
@@ -110,11 +110,11 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
|
|
||||||
// Relay packets are special
|
// Relay packets are special
|
||||||
if isMessageRelay {
|
if isMessageRelay {
|
||||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache)
|
f.handleOutsideRelayPacket(hostinfo, via, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
out := f.batchers[q].Reserve(len(packet))[:0]
|
||||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -168,7 +168,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
// The entire body is sent as AD, not encrypted.
|
// The entire body is sent as AD, not encrypted.
|
||||||
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
|
||||||
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||||
@@ -176,16 +176,10 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
// which will gracefully fail in the DecryptDanger call.
|
// which will gracefully fail in the DecryptDanger call.
|
||||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||||
var err error
|
// The decrypted output is empty (relay packets carry their payload as AD) and unused.
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
// The recursive readOutsidePackets call below operates on signedPayload. Passing
|
||||||
if err != nil {
|
// nil avoids reserving an arena slot.
|
||||||
return
|
if _, err := hostinfo.ConnectionState.dKey.DecryptDanger(nil, signedPayload, signatureValue, h.MessageCounter, nb); err != nil {
|
||||||
}
|
|
||||||
// Advance the replay window now that the frame is authenticated
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h)
|
|
||||||
}
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
// Successfully validated the thing. Get rid of the Relay header.
|
// Successfully validated the thing. Get rid of the Relay header.
|
||||||
@@ -201,7 +195,8 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
|
||||||
// its internal mapping. This should never happen.
|
// its internal mapping. This should never happen.
|
||||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||||
"relayRemoteIndex", h.RemoteIndex,
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -217,15 +212,15 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
f.readOutsidePackets(via, signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
||||||
"relayTo", relay.PeerAddr,
|
"relayTo", relay.PeerAddr,
|
||||||
"relayFrom", hostinfo.vpnAddrs[0],
|
|
||||||
"error", err,
|
"error", err,
|
||||||
|
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -235,7 +230,8 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
switch targetRelay.Type {
|
switch targetRelay.Type {
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel
|
||||||
// Find the target HostInfo
|
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||||
|
out := f.batchers[q].Reserve(len(packet) + header.Len + hostinfo.ConnectionState.dKey.Overhead())[:0]
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
@@ -542,7 +538,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = f.readers[q].Write(out)
|
err = f.batchers[q].Commit(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||||
|
// slices via Reserve and releases them in bulk via Reset. Coalescers take
|
||||||
|
// an *Arena at construction so the caller controls the slab lifetime and
|
||||||
|
// can share one slab across multiple coalescers (MultiCoalescer hands the
|
||||||
|
// same *Arena to every lane so the lanes don't carry their own backings).
|
||||||
|
//
|
||||||
|
// Reserve borrows; the slice is valid until the next Reset. The slab grows
|
||||||
|
// (by allocating a fresh, larger backing array) if a Reserve doesn't fit;
|
||||||
|
// pre-size the arena via NewArena to avoid that path on the hot path.
|
||||||
|
type Arena struct {
|
||||||
|
buf []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewArena returns an Arena with a pre-allocated backing of the given
|
||||||
|
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
||||||
|
// only feeds the coalescer pre-made []byte packets via Commit).
|
||||||
|
func NewArena(capacity int) *Arena {
|
||||||
|
return &Arena{buf: make([]byte, 0, capacity)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the
|
||||||
|
// request doesn't fit the current backing, a fresh, larger backing is
|
||||||
|
// allocated; already-borrowed slices reference the old backing and remain
|
||||||
|
// valid until Reset.
|
||||||
|
func (a *Arena) Reserve(sz int) []byte {
|
||||||
|
if len(a.buf)+sz > cap(a.buf) {
|
||||||
|
newCap := max(cap(a.buf)*2, sz)
|
||||||
|
a.buf = make([]byte, 0, newCap)
|
||||||
|
}
|
||||||
|
start := len(a.buf)
|
||||||
|
a.buf = a.buf[:start+sz]
|
||||||
|
return a.buf[start : start+sz : start+sz]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset releases every slice handed out since the last Reset. Callers must
|
||||||
|
// not use any previously-borrowed slice after this returns. The underlying
|
||||||
|
// backing array is retained so subsequent Reserves don't re-allocate.
|
||||||
|
func (a *Arena) Reset() {
|
||||||
|
a.buf = a.buf[:0]
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||||
|
type Passthrough struct {
|
||||||
|
out io.Writer
|
||||||
|
slots [][]byte
|
||||||
|
arena *util.Arena
|
||||||
|
cursor int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPassthrough(w io.Writer, slots int, arena *util.Arena) *Passthrough {
|
||||||
|
return &Passthrough{
|
||||||
|
out: w,
|
||||||
|
slots: make([][]byte, 0, slots),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Reserve(sz int) []byte {
|
||||||
|
return p.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Commit(pkt []byte) error {
|
||||||
|
p.slots = append(p.slots, pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Flush() error {
|
||||||
|
var firstErr error
|
||||||
|
for _, s := range p.slots {
|
||||||
|
_, err := p.out.Write(s)
|
||||||
|
if err != nil && firstErr == nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
clear(p.slots)
|
||||||
|
p.slots = p.slots[:0]
|
||||||
|
p.arena.Reset()
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
type RxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||||
|
Commit(pkt []byte) error
|
||||||
|
// Flush emits every queued packet in arrival order. Returns the
|
||||||
|
// first error observed; keeps draining so one bad packet doesn't hold up
|
||||||
|
// the rest. After Flush returns, borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
const SendBatchCap = 128
|
||||||
|
|
||||||
|
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||||
|
type batchWriter interface {
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||||
|
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||||
|
// Slot bytes are borrowed from the injected Arena and remain valid until
|
||||||
|
// Flush, which Resets the arena.
|
||||||
|
type SendBatch struct {
|
||||||
|
out batchWriter
|
||||||
|
bufs [][]byte
|
||||||
|
dsts []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
arena *util.Arena
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSendBatch makes a SendBatch with batchCap slots backed by arena.
|
||||||
|
func NewSendBatch(out batchWriter, batchCap int, arena *util.Arena) *SendBatch {
|
||||||
|
return &SendBatch{
|
||||||
|
out: out,
|
||||||
|
bufs: make([][]byte, 0, batchCap),
|
||||||
|
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||||
|
ecns: make([]byte, 0, batchCap),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Reserve(sz int) []byte {
|
||||||
|
return b.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
||||||
|
b.bufs = append(b.bufs, pkt)
|
||||||
|
b.dsts = append(b.dsts, dst)
|
||||||
|
b.ecns = append(b.ecns, outerECN)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Flush() error {
|
||||||
|
var err error
|
||||||
|
if len(b.bufs) > 0 {
|
||||||
|
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||||
|
}
|
||||||
|
clear(b.bufs)
|
||||||
|
b.bufs = b.bufs[:0]
|
||||||
|
b.dsts = b.dsts[:0]
|
||||||
|
b.ecns = b.ecns[:0]
|
||||||
|
b.arena.Reset()
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,126 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeBatchWriter struct {
|
||||||
|
bufs [][]byte
|
||||||
|
addrs []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||||
|
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||||
|
// returns, so tests must capture data before that happens.
|
||||||
|
w.bufs = make([][]byte, len(bufs))
|
||||||
|
for i, b := range bufs {
|
||||||
|
cp := make([]byte, len(b))
|
||||||
|
copy(cp, b)
|
||||||
|
w.bufs[i] = cp
|
||||||
|
}
|
||||||
|
w.addrs = append(w.addrs[:0], addrs...)
|
||||||
|
w.ecns = append(w.ecns[:0], ecns...)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 4, util.NewArena(32))
|
||||||
|
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||||
|
}
|
||||||
|
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 4 {
|
||||||
|
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
||||||
|
}
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if len(buf) != 3 || buf[0] != byte(i) {
|
||||||
|
t.Errorf("buf %d: %x", i, buf)
|
||||||
|
}
|
||||||
|
if fw.addrs[i] != ap {
|
||||||
|
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush again with nothing committed — should be a no-op.
|
||||||
|
fw.bufs = nil
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("empty Flush: %v", err)
|
||||||
|
}
|
||||||
|
if fw.bufs != nil {
|
||||||
|
t.Fatalf("empty Flush triggered WriteBatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reuse after Flush.
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 3, util.NewArena(8))
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
s := b.Reserve(8)
|
||||||
|
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||||
|
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
// Tiny initial backing forces a grow on the second Reserve.
|
||||||
|
b := NewSendBatch(fw, 1, util.NewArena(4))
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
s1 := b.Reserve(4)
|
||||||
|
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||||
|
b.Commit(pkt1, ap, 0)
|
||||||
|
|
||||||
|
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||||
|
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||||
|
b.Commit(pkt2, ap, 0)
|
||||||
|
|
||||||
|
// pkt1 must still be intact even though backing reallocated.
|
||||||
|
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||||
|
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 2 {
|
||||||
|
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
||||||
|
}
|
||||||
|
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
||||||
|
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
||||||
|
}
|
||||||
|
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
||||||
|
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
+8
-2
@@ -4,15 +4,21 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||||
|
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||||
|
const defaultBatchBufSize = 65535
|
||||||
|
|
||||||
type Device interface {
|
type Device interface {
|
||||||
io.ReadWriteCloser
|
io.Closer
|
||||||
Activate() error
|
Activate() error
|
||||||
Networks() []netip.Prefix
|
Networks() []netip.Prefix
|
||||||
Name() string
|
Name() string
|
||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
SupportsMultiqueue() bool
|
SupportsMultiqueue() bool
|
||||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
NewMultiQueueReader() error
|
||||||
|
Readers() []tio.Queue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,10 +4,11 @@ package overlaytest
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
// NoopTun is an overlay.Device that silently discards every read and write.
|
// NoopTun is an overlay.Device that silently discards every read and write.
|
||||||
@@ -15,6 +16,10 @@ import (
|
|||||||
// exercise the datapath.
|
// exercise the datapath.
|
||||||
type NoopTun struct{}
|
type NoopTun struct{}
|
||||||
|
|
||||||
|
func (NoopTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (NoopTun) RoutesFor(addr netip.Addr) routing.Gateways {
|
func (NoopTun) RoutesFor(addr netip.Addr) routing.Gateways {
|
||||||
return routing.Gateways{}
|
return routing.Gateways{}
|
||||||
}
|
}
|
||||||
@@ -31,7 +36,7 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read([]byte) (int, error) {
|
func (NoopTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -43,8 +48,12 @@ func (NoopTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (NoopTun) NewMultiQueueReader() error {
|
||||||
return nil, errors.New("unsupported")
|
return errors.New("unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (NoopTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{NoopTun{}}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Close() error {
|
func (NoopTun) Close() error {
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
type pollQueueSet struct {
|
||||||
|
pq []*Poll
|
||||||
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
|
pqi []Queue
|
||||||
|
shutdownFd int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPollQueueSet() (QueueSet, error) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &pollQueueSet{
|
||||||
|
pq: []*Poll{},
|
||||||
|
pqi: []Queue{},
|
||||||
|
shutdownFd: shutdownFd,
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Queues() []Queue {
|
||||||
|
return c.pqi
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Add(fd int) error {
|
||||||
|
x, err := newPoll(fd, c.shutdownFd)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.pq = append(c.pq, x)
|
||||||
|
c.pqi = append(c.pqi, x)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) wakeForShutdown() error {
|
||||||
|
var buf [8]byte
|
||||||
|
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||||
|
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Close() error {
|
||||||
|
if c.shutdownFd < 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
errs := []error{}
|
||||||
|
|
||||||
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, x := range c.pq {
|
||||||
|
if err := x.Close(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// All Polls reference shutdownFd in their pollfd arrays, so close it
|
||||||
|
// only after every Poll.Close has returned.
|
||||||
|
if err := unix.Close(c.shutdownFd); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
c.shutdownFd = -1
|
||||||
|
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
)
|
||||||
|
|
||||||
|
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
||||||
|
type QueueSet interface {
|
||||||
|
io.Closer
|
||||||
|
Queues() []Queue
|
||||||
|
|
||||||
|
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
||||||
|
Add(fd int) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||||
|
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||||
|
type Capabilities struct {
|
||||||
|
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed
|
||||||
|
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe.
|
||||||
|
TSO bool
|
||||||
|
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so
|
||||||
|
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
||||||
|
USO bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// Queue is a readable/writable Poll queue. One Queue is driven by a single
|
||||||
|
// read goroutine plus a single writer (see Write below).
|
||||||
|
type Queue interface {
|
||||||
|
io.Closer
|
||||||
|
|
||||||
|
// Read will read at least 1 packet from the tun (up to len(p)).
|
||||||
|
// mem will be used to provide the backing for each of p[n].Bytes.
|
||||||
|
// Callers should size mem and p to avoid exhausting mem before p.
|
||||||
|
// Returns the number of packets actually read, or error.
|
||||||
|
Read(p []wire.TunPacket, mem []byte) (int, error)
|
||||||
|
|
||||||
|
// Write emits a single packet on the plaintext (outside→inside)
|
||||||
|
// delivery path.
|
||||||
|
Write(p []byte) (int, error)
|
||||||
|
|
||||||
|
// Capabilities returns the Queue's negotiated offload capabilities,
|
||||||
|
// or the zero value when q does not advertise any.
|
||||||
|
Capabilities() Capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
||||||
|
// The zero value means "not a superpacket" — Bytes is one regular IP
|
||||||
|
// datagram and no segmentation is required.
|
||||||
|
type GSOInfo struct {
|
||||||
|
// Size is the GSO segment size: max payload bytes per segment
|
||||||
|
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means
|
||||||
|
// not a superpacket.
|
||||||
|
Size uint16
|
||||||
|
// HdrLen is the total L3+L4 header length within Bytes (already
|
||||||
|
// corrected via correctHdrLen, so safe to slice on).
|
||||||
|
HdrLen uint16
|
||||||
|
// CsumStart is the L4 header offset inside Bytes (== L3 header
|
||||||
|
// length).
|
||||||
|
CsumStart uint16
|
||||||
|
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows
|
||||||
|
// which checksum/header layout to apply.
|
||||||
|
Proto GSOProto
|
||||||
|
}
|
||||||
|
|
||||||
|
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
|
||||||
|
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
||||||
|
// inside the transport header virtio NEEDS_CSUM expects.
|
||||||
|
type GSOProto uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
GSOProtoNone GSOProto = iota
|
||||||
|
GSOProtoTCP
|
||||||
|
GSOProtoUDP
|
||||||
|
)
|
||||||
|
|
||||||
|
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
||||||
|
// assembled from a header prefix plus one or more borrowed payload
|
||||||
|
// fragments, in a single vectored write (writev with a leading
|
||||||
|
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
||||||
|
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
||||||
|
// support do not implement this interface and coalescing is skipped.
|
||||||
|
//
|
||||||
|
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
||||||
|
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
||||||
|
// header (mutable - the L4 checksum field must hold the pseudo-header
|
||||||
|
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||||
|
// pays are non-overlapping payload fragments whose concatenation is the
|
||||||
|
// full superpacket payload; they are read-only from the writer's
|
||||||
|
// perspective and must remain valid until the call returns. Every segment
|
||||||
|
// in pays except possibly the last is exactly the same size. proto picks
|
||||||
|
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
|
||||||
|
//
|
||||||
|
// Callers should also consult CapsProvider (via SupportsGSO or
|
||||||
|
// QueueCapabilities) for the per-protocol negotiated capability; an
|
||||||
|
// implementation of GSOWriter is necessary but not sufficient since USO
|
||||||
|
// may not have been negotiated even when TSO was.
|
||||||
|
type GSOWriter interface {
|
||||||
|
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||||
|
// queue advertises the negotiated capability for `want`. A writer that
|
||||||
|
// implements GSOWriter but not CapsProvider is treated as permissive
|
||||||
|
// (used by tests and fakes that don't negotiate).
|
||||||
|
func SupportsGSO(w Queue, want GSOProto) (GSOWriter, bool) {
|
||||||
|
gw, ok := w.(GSOWriter)
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
caps := w.Capabilities()
|
||||||
|
switch want {
|
||||||
|
case GSOProtoTCP:
|
||||||
|
return gw, caps.TSO
|
||||||
|
case GSOProtoUDP:
|
||||||
|
return gw, caps.USO
|
||||||
|
default:
|
||||||
|
return gw, false
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Poll struct {
|
||||||
|
fd int
|
||||||
|
|
||||||
|
readPoll [2]unix.PollFd
|
||||||
|
writePoll [2]unix.PollFd
|
||||||
|
writeLock sync.Mutex
|
||||||
|
closed atomic.Bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||||
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &Poll{
|
||||||
|
fd: fd,
|
||||||
|
readPoll: [2]unix.PollFd{
|
||||||
|
{Fd: int32(fd), Events: unix.POLLIN},
|
||||||
|
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||||
|
},
|
||||||
|
writePoll: [2]unix.PollFd{
|
||||||
|
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||||
|
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||||
|
},
|
||||||
|
writeLock: sync.Mutex{},
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
|
||||||
|
// Returns os.ErrClosed if Close was called.
|
||||||
|
func (t *Poll) blockOnRead() error {
|
||||||
|
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||||
|
var err error
|
||||||
|
for {
|
||||||
|
_, err = unix.Poll(t.readPoll[:], -1)
|
||||||
|
if err != unix.EINTR {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tunEvents := t.readPoll[0].Revents
|
||||||
|
shutdownEvents := t.readPoll[1].Revents
|
||||||
|
t.readPoll[0].Revents = 0
|
||||||
|
t.readPoll[1].Revents = 0
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
if tunEvents&problemFlags != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) blockOnWrite() error {
|
||||||
|
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||||
|
var err error
|
||||||
|
for {
|
||||||
|
_, err = unix.Poll(t.writePoll[:], -1)
|
||||||
|
if err != unix.EINTR {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.writeLock.Lock()
|
||||||
|
tunEvents := t.writePoll[0].Revents
|
||||||
|
shutdownEvents := t.writePoll[1].Revents
|
||||||
|
t.writePoll[0].Revents = 0
|
||||||
|
t.writePoll[1].Revents = 0
|
||||||
|
t.writeLock.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
if tunEvents&problemFlags != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) readOne(to []byte) (int, error) {
|
||||||
|
for {
|
||||||
|
n, errno := unix.Read(t.fd, to)
|
||||||
|
if errno == nil {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
switch errno {
|
||||||
|
case unix.EAGAIN:
|
||||||
|
if err := t.blockOnRead(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
case unix.EINTR:
|
||||||
|
// retry
|
||||||
|
case unix.EBADF:
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
default:
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Write(from []byte) (int, error) {
|
||||||
|
for {
|
||||||
|
n, errno := unix.Write(t.fd, from)
|
||||||
|
if errno == nil {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
switch errno {
|
||||||
|
case unix.EAGAIN:
|
||||||
|
if err := t.blockOnWrite(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
case unix.EINTR:
|
||||||
|
// retry
|
||||||
|
case unix.EBADF:
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
default:
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Close() error {
|
||||||
|
if t.closed.Swap(true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
//shutdownFd is owned by the container, so we should not close it
|
||||||
|
var err error
|
||||||
|
if t.fd >= 0 {
|
||||||
|
err = unix.Close(t.fd)
|
||||||
|
t.fd = -1
|
||||||
|
}
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Capabilities() Capabilities {
|
||||||
|
return Capabilities{}
|
||||||
|
}
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
// +build linux,!android,!e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||||
|
// The caller takes ownership of the read fd (pass it into a QueueSet).
|
||||||
|
func newReadPipe(t *testing.T) int {
|
||||||
|
t.Helper()
|
||||||
|
var fds [2]int
|
||||||
|
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||||
|
t.Fatalf("pipe2: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||||
|
return fds[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||||
|
parent, err := NewPollQueueSet()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, parent.Add(newReadPipe(t)))
|
||||||
|
require.NoError(t, parent.Add(newReadPipe(t)))
|
||||||
|
// QueueSet.Close owns the read fds we Added — don't register a separate
|
||||||
|
// Cleanup to close them or we'll double-close whatever fd the kernel
|
||||||
|
// has since reused.
|
||||||
|
|
||||||
|
readers := parent.Queues()
|
||||||
|
errs := make([]error, len(readers))
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i, r := range readers {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(i int, r Queue) {
|
||||||
|
defer wg.Done()
|
||||||
|
pkts := make([]wire.TunPacket, 1)
|
||||||
|
_, errs[i] = r.Read(pkts, make([]byte, 64))
|
||||||
|
}(i, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
if err := parent.Close(); err != nil {
|
||||||
|
t.Fatalf("Close: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() { wg.Wait(); close(done) }()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("readers did not wake")
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, err := range errs {
|
||||||
|
if !errors.Is(err, os.ErrClosed) {
|
||||||
|
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPoll_Close_Idempotent(t *testing.T) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
|
||||||
|
|
||||||
|
tf, err := newPoll(newReadPipe(t), shutdownFd)
|
||||||
|
require.NoError(t, err)
|
||||||
|
if err := tf.Close(); err != nil {
|
||||||
|
t.Fatalf("first Close: %v", err)
|
||||||
|
}
|
||||||
|
if err := tf.Close(); err != nil {
|
||||||
|
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPollQueueSet_Close_ClosesEventfd(t *testing.T) {
|
||||||
|
qs, err := NewPollQueueSet()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||||
|
|
||||||
|
fd := qs.(*pollQueueSet).shutdownFd
|
||||||
|
require.NoError(t, qs.Close())
|
||||||
|
|
||||||
|
// Closing the eventfd again should fail with EBADF, proving Close
|
||||||
|
// actually released it.
|
||||||
|
if err := unix.Close(fd); err == nil {
|
||||||
|
t.Fatalf("eventfd %d still open after QueueSet.Close", fd)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second Close must be a no-op (and must not double-close the eventfd
|
||||||
|
// in case the kernel handed it out to another caller in the meantime).
|
||||||
|
if err := qs.Close(); err != nil {
|
||||||
|
t.Fatalf("second Close: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+39
-8
@@ -13,12 +13,14 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
fd int
|
fd int
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
@@ -26,16 +28,37 @@ type tun struct {
|
|||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.rwc.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Write(p []byte) (int, error) {
|
||||||
|
return t.rwc.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Close() error {
|
||||||
|
return t.rwc.Close()
|
||||||
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
// XXX Android returns an fd in non-blocking mode which is necessary for shutdown to work properly.
|
// XXX Android returns an fd in non-blocking mode which is necessary for shutdown to work properly.
|
||||||
// Be sure not to call file.Fd() as it will set the fd to blocking mode.
|
// Be sure not to call file.Fd() as it will set the fd to blocking mode.
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
ReadWriteCloser: file,
|
rwc: file,
|
||||||
fd: deviceFd,
|
fd: deviceFd,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -62,7 +85,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t tun) Activate() error {
|
func (t *tun) Activate() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -99,6 +122,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
return fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
+32
-18
@@ -16,14 +16,16 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
DefaultMTU int
|
DefaultMTU int
|
||||||
@@ -124,11 +126,11 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
ReadWriteCloser: os.NewFile(uintptr(fd), ""),
|
rwc: os.NewFile(uintptr(fd), ""),
|
||||||
Device: name,
|
Device: name,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err = t.reload(c, true)
|
err = t.reload(c, true)
|
||||||
@@ -158,8 +160,8 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, e
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.ReadWriteCloser != nil {
|
if t.rwc != nil {
|
||||||
return t.ReadWriteCloser.Close()
|
return t.rwc.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -502,13 +504,17 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
buf := make([]byte, len(to)+4)
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
n, err := t.ReadWriteCloser.Read(buf)
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
copy(to, buf[4:])
|
n, err := t.rwc.Read(mem)
|
||||||
return n - 4, err
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
// Write is only valid for single threaded use
|
||||||
@@ -536,7 +542,7 @@ func (t *tun) Write(from []byte) (int, error) {
|
|||||||
|
|
||||||
copy(buf[4:], from)
|
copy(buf[4:], from)
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Write(buf)
|
n, err := t.rwc.Write(buf)
|
||||||
return n - 4, err
|
return n - 4, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -552,6 +558,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
return fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-6
@@ -10,7 +10,9 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type disabledTun struct {
|
type disabledTun struct {
|
||||||
@@ -18,9 +20,10 @@ type disabledTun struct {
|
|||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
|
|
||||||
// Track these metrics since we don't have the tun device to do it for us
|
// Track these metrics since we don't have the tun device to do it for us
|
||||||
tx metrics.Counter
|
tx metrics.Counter
|
||||||
rx metrics.Counter
|
rx metrics.Counter
|
||||||
l *slog.Logger
|
numReaders int
|
||||||
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
||||||
@@ -28,6 +31,7 @@ func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled boo
|
|||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
read: make(chan []byte, queueLen),
|
read: make(chan []byte, queueLen),
|
||||||
l: l,
|
l: l,
|
||||||
|
numReaders: 1,
|
||||||
}
|
}
|
||||||
|
|
||||||
if metricsEnabled {
|
if metricsEnabled {
|
||||||
@@ -57,7 +61,7 @@ func (*disabledTun) Name() string {
|
|||||||
return "disabled"
|
return "disabled"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
func (t *disabledTun) readOne(b []byte) (int, error) {
|
||||||
r, ok := <-t.read
|
r, ok := <-t.read
|
||||||
if !ok {
|
if !ok {
|
||||||
return 0, io.EOF
|
return 0, io.EOF
|
||||||
@@ -75,6 +79,19 @@ func (t *disabledTun) Read(b []byte) (int, error) {
|
|||||||
return copy(b, r), nil
|
return copy(b, r), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
||||||
out := make([]byte, len(b))
|
out := make([]byte, len(b))
|
||||||
out = iputil.CreateICMPEchoResponse(b, out)
|
out = iputil.CreateICMPEchoResponse(b, out)
|
||||||
@@ -110,8 +127,21 @@ func (t *disabledTun) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *disabledTun) NewMultiQueueReader() error {
|
||||||
return t, nil
|
t.numReaders++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Readers() []tio.Queue {
|
||||||
|
out := make([]tio.Queue, t.numReaders)
|
||||||
|
for i := range t.numReaders {
|
||||||
|
out[i] = t
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) Close() error {
|
func (t *disabledTun) Close() error {
|
||||||
|
|||||||
@@ -1,120 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
// +build linux,!android,!e2e_testing
|
|
||||||
|
|
||||||
package overlay
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
|
||||||
// The caller takes ownership of the read fd (pass it to newTunFd / newFriend).
|
|
||||||
func newReadPipe(t *testing.T) int {
|
|
||||||
t.Helper()
|
|
||||||
var fds [2]int
|
|
||||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
|
||||||
t.Fatalf("pipe2: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
|
||||||
return fds[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_WakeForShutdown_UnblocksRead(t *testing.T) {
|
|
||||||
tf, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() { _ = tf.Close() })
|
|
||||||
|
|
||||||
done := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
_, err := tf.Read(make([]byte, 64))
|
|
||||||
done <- err
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Verify Read is actually blocked in poll.
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
t.Fatalf("Read returned before shutdown signal: %v", err)
|
|
||||||
case <-time.After(50 * time.Millisecond):
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := tf.wakeForShutdown(); err != nil {
|
|
||||||
t.Fatalf("wakeForShutdown: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
if !errors.Is(err, os.ErrClosed) {
|
|
||||||
t.Fatalf("expected os.ErrClosed, got %v", err)
|
|
||||||
}
|
|
||||||
case <-time.After(2 * time.Second):
|
|
||||||
t.Fatal("Read did not wake on shutdown")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_WakeForShutdown_WakesFriends(t *testing.T) {
|
|
||||||
parent, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
friend, err := parent.newFriend(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
_ = parent.Close()
|
|
||||||
t.Fatalf("newFriend: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = friend.Close()
|
|
||||||
_ = parent.Close()
|
|
||||||
})
|
|
||||||
|
|
||||||
readers := []*tunFile{parent, friend}
|
|
||||||
errs := make([]error, len(readers))
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for i, r := range readers {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int, r *tunFile) {
|
|
||||||
defer wg.Done()
|
|
||||||
_, errs[i] = r.Read(make([]byte, 64))
|
|
||||||
}(i, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
time.Sleep(50 * time.Millisecond)
|
|
||||||
|
|
||||||
if err := parent.wakeForShutdown(); err != nil {
|
|
||||||
t.Fatalf("wakeForShutdown: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() { wg.Wait(); close(done) }()
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(2 * time.Second):
|
|
||||||
t.Fatal("readers did not wake")
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, err := range errs {
|
|
||||||
if !errors.Is(err, os.ErrClosed) {
|
|
||||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_Close_Idempotent(t *testing.T) {
|
|
||||||
tf, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
if err := tf.Close(); err != nil {
|
|
||||||
t.Fatalf("first Close: %v", err)
|
|
||||||
}
|
|
||||||
if err := tf.Close(); err != nil {
|
|
||||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+26
-5
@@ -7,7 +7,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -18,9 +17,10 @@ import (
|
|||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -157,7 +157,20 @@ func (t *tun) blockOnWrite() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) readOne(to []byte) (int, error) {
|
||||||
// first 4 bytes is protocol family, in network byte order
|
// first 4 bytes is protocol family, in network byte order
|
||||||
var head [4]byte
|
var head [4]byte
|
||||||
iovecs := [2]syscall.Iovec{
|
iovecs := [2]syscall.Iovec{
|
||||||
@@ -565,8 +578,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
@@ -593,6 +606,14 @@ func (t *tun) addRoutes(logErrors bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *tun) removeRoutes(routes []Route) error {
|
func (t *tun) removeRoutes(routes []Route) error {
|
||||||
for _, r := range routes {
|
for _, r := range routes {
|
||||||
if !r.Install {
|
if !r.Install {
|
||||||
|
|||||||
+39
-17
@@ -16,18 +16,41 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.rwc.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Write(p []byte) (int, error) {
|
||||||
|
return t.rwc.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Close() error {
|
||||||
|
return t.rwc.Close()
|
||||||
|
}
|
||||||
|
|
||||||
func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
||||||
return nil, fmt.Errorf("newTun not supported in iOS")
|
return nil, fmt.Errorf("newTun not supported in iOS")
|
||||||
}
|
}
|
||||||
@@ -35,9 +58,9 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error)
|
|||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||||
t := &tun{
|
t := &tun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
ReadWriteCloser: &tunReadCloser{f: file},
|
rwc: &tunReadCloser{f: file},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -96,18 +119,9 @@ type tunReadCloser struct {
|
|||||||
wBuf []byte
|
wBuf []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Read returns a packet with the BSD 4-byte header, watch out!
|
||||||
func (tr *tunReadCloser) Read(to []byte) (int, error) {
|
func (tr *tunReadCloser) Read(to []byte) (int, error) {
|
||||||
tr.rMu.Lock()
|
return tr.f.Read(to)
|
||||||
defer tr.rMu.Unlock()
|
|
||||||
|
|
||||||
if cap(tr.rBuf) < len(to)+4 {
|
|
||||||
tr.rBuf = make([]byte, len(to)+4)
|
|
||||||
}
|
|
||||||
tr.rBuf = tr.rBuf[:len(to)+4]
|
|
||||||
|
|
||||||
n, err := tr.f.Read(tr.rBuf)
|
|
||||||
copy(to, tr.rBuf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tr *tunReadCloser) Write(from []byte) (int, error) {
|
func (tr *tunReadCloser) Write(from []byte) (int, error) {
|
||||||
@@ -155,6 +169,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
return fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
+84
-241
@@ -4,9 +4,7 @@
|
|||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -19,180 +17,15 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/vishvananda/netlink"
|
"github.com/vishvananda/netlink"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
// tunFile wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
|
||||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
|
||||||
type tunFile struct {
|
|
||||||
fd int
|
|
||||||
shutdownFd int
|
|
||||||
lastOne bool
|
|
||||||
readPoll [2]unix.PollFd
|
|
||||||
writePoll [2]unix.PollFd
|
|
||||||
closed bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// newFriend makes a tunFile for a MultiQueueReader that copies the shutdown eventfd from the parent tun
|
|
||||||
func (r *tunFile) newFriend(fd int) (*tunFile, error) {
|
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
|
||||||
}
|
|
||||||
return &tunFile{
|
|
||||||
fd: fd,
|
|
||||||
shutdownFd: r.shutdownFd,
|
|
||||||
readPoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
|
||||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
writePoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
|
||||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTunFd(fd int) (*tunFile, error) {
|
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
out := &tunFile{
|
|
||||||
fd: fd,
|
|
||||||
shutdownFd: shutdownFd,
|
|
||||||
lastOne: true,
|
|
||||||
readPoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
|
||||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
writePoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
|
||||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) blockOnRead() error {
|
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
|
||||||
var err error
|
|
||||||
for {
|
|
||||||
_, err = unix.Poll(r.readPoll[:], -1)
|
|
||||||
if err != unix.EINTR {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
//always reset these!
|
|
||||||
tunEvents := r.readPoll[0].Revents
|
|
||||||
shutdownEvents := r.readPoll[1].Revents
|
|
||||||
r.readPoll[0].Revents = 0
|
|
||||||
r.readPoll[1].Revents = 0
|
|
||||||
//do the err check before trusting the potentially bogus bits we just got
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
} else if tunEvents&problemFlags != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) blockOnWrite() error {
|
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
|
||||||
var err error
|
|
||||||
for {
|
|
||||||
_, err = unix.Poll(r.writePoll[:], -1)
|
|
||||||
if err != unix.EINTR {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
//always reset these!
|
|
||||||
tunEvents := r.writePoll[0].Revents
|
|
||||||
shutdownEvents := r.writePoll[1].Revents
|
|
||||||
r.writePoll[0].Revents = 0
|
|
||||||
r.writePoll[1].Revents = 0
|
|
||||||
//do the err check before trusting the potentially bogus bits we just got
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
} else if tunEvents&problemFlags != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Read(buf []byte) (int, error) {
|
|
||||||
for {
|
|
||||||
if n, err := unix.Read(r.fd, buf); err == nil {
|
|
||||||
return n, nil
|
|
||||||
} else if err == unix.EAGAIN {
|
|
||||||
if err = r.blockOnRead(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
} else if err == unix.EINTR {
|
|
||||||
continue
|
|
||||||
} else if err == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
} else {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Write(buf []byte) (int, error) {
|
|
||||||
for {
|
|
||||||
if n, err := unix.Write(r.fd, buf); err == nil {
|
|
||||||
return n, nil
|
|
||||||
} else if err == unix.EAGAIN {
|
|
||||||
if err = r.blockOnWrite(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
} else if err == unix.EINTR {
|
|
||||||
continue
|
|
||||||
} else if err == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
} else {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) wakeForShutdown() error {
|
|
||||||
var buf [8]byte
|
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
|
||||||
_, err := unix.Write(int(r.readPoll[1].Fd), buf[:])
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Close() error {
|
|
||||||
if r.closed { // avoid closing more than once. Technically a fd could get re-used, which would be a problem
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
r.closed = true
|
|
||||||
if r.lastOne {
|
|
||||||
_ = unix.Close(r.shutdownFd)
|
|
||||||
}
|
|
||||||
return unix.Close(r.fd)
|
|
||||||
}
|
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
*tunFile
|
readers tio.QueueSet
|
||||||
readers []*tunFile
|
|
||||||
closeLock sync.Mutex
|
closeLock sync.Mutex
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
@@ -239,7 +72,9 @@ type ifreqQLEN struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
// We don't know what flags the caller opened this fd with and can't turn
|
||||||
|
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||||
|
t, err := newTunGeneric(c, l, deviceFd, false, false, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -249,46 +84,60 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||||
|
// missing (docker containers occasionally omit it).
|
||||||
|
func openTunDev() (int, error) {
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err == nil {
|
||||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
return fd, nil
|
||||||
if os.IsNotExist(err) {
|
|
||||||
err = os.MkdirAll("/dev/net", 0755)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
|
||||||
}
|
|
||||||
err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200)))
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
return -1, err
|
||||||
|
}
|
||||||
|
if err = os.MkdirAll("/dev/net", 0755); err != nil {
|
||||||
|
return -1, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||||
|
}
|
||||||
|
if err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))); err != nil {
|
||||||
|
return -1, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||||
|
}
|
||||||
|
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
|
if err != nil {
|
||||||
|
return -1, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||||
|
}
|
||||||
|
return fd, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen
|
||||||
|
// device name on success.
|
||||||
|
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||||
var req ifReq
|
var req ifReq
|
||||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
req.Flags = flags
|
||||||
|
copy(req.Name[:], name)
|
||||||
|
if err := ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||||
|
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||||
if multiqueue {
|
if multiqueue {
|
||||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||||
}
|
}
|
||||||
nameStr := c.GetString("tun.dev", "")
|
nameStr := c.GetString("tun.dev", "")
|
||||||
copy(req.Name[:], nameStr)
|
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, &NameError{
|
|
||||||
Name: nameStr,
|
|
||||||
Underlying: err,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
name := strings.Trim(string(req.Name[:]), "\x00")
|
|
||||||
|
|
||||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
fd, err := openTunDev()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
name, err := tunSetIff(fd, nameStr, baseFlags)
|
||||||
|
if err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, &NameError{Name: nameStr, Underlying: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
t, err := newTunGeneric(c, l, fd, false, false, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -299,15 +148,21 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
}
|
}
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr, usoEnabled bool, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
tfd, err := newTunFd(fd)
|
qs, err := tio.NewPollQueueSet()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
err = qs.Add(fd)
|
||||||
|
if err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
tunFile: tfd,
|
readers: qs,
|
||||||
readers: []*tunFile{tfd},
|
|
||||||
closeLock: sync.Mutex{},
|
closeLock: sync.Mutex{},
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||||
@@ -410,32 +265,29 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
t.closeLock.Lock()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
|
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var req ifReq
|
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
|
||||||
copy(req.Name[:], t.Device)
|
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
out, err := t.tunFile.newFriend(fd)
|
err = t.readers.Add(fd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
t.readers = append(t.readers, out)
|
return nil
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||||
@@ -603,6 +455,15 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
|||||||
Table: unix.RT_TABLE_MAIN,
|
Table: unix.RT_TABLE_MAIN,
|
||||||
Type: unix.RTN_UNICAST,
|
Type: unix.RTN_UNICAST,
|
||||||
}
|
}
|
||||||
|
// Match the metric the kernel uses for its auto-installed connected
|
||||||
|
// route, so RouteReplace overwrites it in place instead of adding a
|
||||||
|
// second route at a worse metric. IPv6 connected routes are installed
|
||||||
|
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the
|
||||||
|
// kernel route wins lookups and our MTU / AdvMSS / Features never
|
||||||
|
// apply on v6.
|
||||||
|
if cidr.Addr().Is6() {
|
||||||
|
nr.Priority = 256
|
||||||
|
}
|
||||||
err := netlink.RouteReplace(&nr)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||||
@@ -869,6 +730,10 @@ func (t *tun) updateRoutes(r netlink.RouteUpdate) {
|
|||||||
t.routeTree.Store(newTree)
|
t.routeTree.Store(newTree)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return t.readers.Queues()
|
||||||
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
t.closeLock.Lock()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
@@ -878,32 +743,10 @@ func (t *tun) Close() error {
|
|||||||
t.routeChan = nil
|
t.routeChan = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit
|
|
||||||
_ = t.tunFile.wakeForShutdown()
|
|
||||||
|
|
||||||
if t.ioctlFd > 0 {
|
if t.ioctlFd > 0 {
|
||||||
_ = unix.Close(int(t.ioctlFd))
|
_ = unix.Close(int(t.ioctlFd))
|
||||||
t.ioctlFd = 0
|
t.ioctlFd = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := range t.readers {
|
return t.readers.Close()
|
||||||
if i == 0 {
|
|
||||||
continue //we want to close the zeroth reader last
|
|
||||||
}
|
|
||||||
err := t.readers[i].Close()
|
|
||||||
if err != nil {
|
|
||||||
t.l.Error("error closing tun reader", "reader", i, "error", err)
|
|
||||||
} else {
|
|
||||||
t.l.Info("closed tun reader", "reader", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
//this is t.readers[0] too
|
|
||||||
err := t.tunFile.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.l.Error("error closing tun reader", "reader", 0, "error", err)
|
|
||||||
} else {
|
|
||||||
t.l.Info("closed tun reader", "reader", 0)
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,7 +3,9 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
+26
-4
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,8 +16,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
@@ -68,6 +69,27 @@ type tun struct {
|
|||||||
fd int
|
fd int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
||||||
@@ -141,7 +163,7 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) readOne(to []byte) (int, error) {
|
||||||
rc, err := t.f.SyscallConn()
|
rc, err := t.f.SyscallConn()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to get syscall conn for tun: %w", err)
|
return 0, fmt.Errorf("failed to get syscall conn for tun: %w", err)
|
||||||
@@ -394,8 +416,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
|
|||||||
+25
-12
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,8 +16,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
@@ -61,6 +62,19 @@ type tun struct {
|
|||||||
out []byte
|
out []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.f.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
||||||
@@ -124,15 +138,6 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
|
||||||
buf := make([]byte, len(to)+4)
|
|
||||||
|
|
||||||
n, err := t.f.Read(buf)
|
|
||||||
|
|
||||||
copy(to, buf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
// Write is only valid for single threaded use
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
buf := t.out
|
buf := t.out
|
||||||
@@ -314,8 +319,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
@@ -366,6 +371,14 @@ func (t *tun) deviceBytes() (o [16]byte) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func addRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
func addRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
||||||
sock, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
sock, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+26
-3
@@ -14,8 +14,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type TestTun struct {
|
type TestTun struct {
|
||||||
@@ -162,7 +164,20 @@ func (t *TestTun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) Read(b []byte) (int, error) {
|
func (t *TestTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) read(b []byte) (int, error) {
|
||||||
p, ok := <-t.rxPackets
|
p, ok := <-t.rxPackets
|
||||||
if !ok {
|
if !ok {
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
@@ -177,10 +192,18 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
|||||||
return n, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *TestTun) SupportsMultiqueue() bool {
|
func (t *TestTun) SupportsMultiqueue() bool {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *TestTun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
return fmt.Errorf("TODO: multiqueue not implemented")
|
||||||
}
|
}
|
||||||
|
|||||||
+25
-7
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"crypto"
|
"crypto"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -18,9 +17,11 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/slackhq/nebula/wintun"
|
"github.com/slackhq/nebula/wintun"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
"golang.org/x/sys/windows"
|
"golang.org/x/sys/windows"
|
||||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
)
|
)
|
||||||
@@ -47,6 +48,19 @@ type winTun struct {
|
|||||||
tun *wintun.NativeTun
|
tun *wintun.NativeTun
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.tun.Read(mem, 0)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||||
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
||||||
}
|
}
|
||||||
@@ -255,10 +269,6 @@ func (t *winTun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) Read(b []byte) (int, error) {
|
|
||||||
return t.tun.Read(b, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *winTun) Write(b []byte) (int, error) {
|
func (t *winTun) Write(b []byte) (int, error) {
|
||||||
return t.tun.Write(b, 0)
|
return t.tun.Write(b, 0)
|
||||||
}
|
}
|
||||||
@@ -267,8 +277,16 @@ func (t *winTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *winTun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
return fmt.Errorf("TODO: multiqueue not implemented for windows")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) Close() error {
|
func (t *winTun) Close() error {
|
||||||
|
|||||||
+33
-5
@@ -6,7 +6,9 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewUserDeviceFromConfig(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, routines int) (Device, error) {
|
func NewUserDeviceFromConfig(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, routines int) (Device, error) {
|
||||||
@@ -23,11 +25,13 @@ func NewUserDevice(vpnNetworks []netip.Prefix) (Device, error) {
|
|||||||
outboundWriter: ow,
|
outboundWriter: ow,
|
||||||
inboundReader: ir,
|
inboundReader: ir,
|
||||||
inboundWriter: iw,
|
inboundWriter: iw,
|
||||||
|
numReaders: 1,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type UserDevice struct {
|
type UserDevice struct {
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
|
numReaders int
|
||||||
|
|
||||||
outboundReader *io.PipeReader
|
outboundReader *io.PipeReader
|
||||||
outboundWriter *io.PipeWriter
|
outboundWriter *io.PipeWriter
|
||||||
@@ -36,6 +40,23 @@ type UserDevice struct {
|
|||||||
inboundWriter *io.PipeWriter
|
inboundWriter *io.PipeWriter
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := d.outboundReader.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Activate() error {
|
func (d *UserDevice) Activate() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -50,20 +71,27 @@ func (d *UserDevice) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (d *UserDevice) NewMultiQueueReader() error {
|
||||||
return d, nil
|
d.numReaders++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Readers() []tio.Queue {
|
||||||
|
out := make([]tio.Queue, d.numReaders)
|
||||||
|
for i := range d.numReaders {
|
||||||
|
out[i] = d
|
||||||
|
}
|
||||||
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
||||||
return d.inboundReader, d.outboundWriter
|
return d.inboundReader, d.outboundWriter
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Read(p []byte) (n int, err error) {
|
|
||||||
return d.outboundReader.Read(p)
|
|
||||||
}
|
|
||||||
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
||||||
return d.inboundWriter.Write(p)
|
return d.inboundWriter.Write(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Close() error {
|
func (d *UserDevice) Close() error {
|
||||||
d.inboundWriter.Close()
|
d.inboundWriter.Close()
|
||||||
d.outboundWriter.Close()
|
d.outboundWriter.Close()
|
||||||
|
|||||||
+39
-57
@@ -7,7 +7,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
@@ -58,25 +57,14 @@ func (rm *relayManager) GetUseRelays() bool {
|
|||||||
// For each candidate relay it either kicks off a handshake to the relay, sends a CreateRelayRequest, retransmits
|
// For each candidate relay it either kicks off a handshake to the relay, sends a CreateRelayRequest, retransmits
|
||||||
// one that may have been lost, or, once the relay is Established, forwards the in-progress
|
// one that may have been lost, or, once the relay is Established, forwards the in-progress
|
||||||
// stage 0 handshake packet for vpnIp through it.
|
// stage 0 handshake packet for vpnIp through it.
|
||||||
func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *HandshakeHostInfo, stage0 []byte) {
|
func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hostinfo *HostInfo, stage0 []byte) {
|
||||||
hostinfo := hh.hostinfo
|
|
||||||
if !rm.GetUseRelays() || len(hostinfo.remotes.relays) == 0 {
|
if !rm.GetUseRelays() || len(hostinfo.remotes.relays) == 0 {
|
||||||
hh.lastRelays = nil
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
relays := hostinfo.remotes.relays
|
hostinfo.logger(rm.l).Info("Attempt to relay through hosts", "relays", hostinfo.remotes.relays)
|
||||||
listLevel := slog.LevelDebug
|
|
||||||
prior := hh.lastRelays
|
|
||||||
if !slices.Equal(relays, prior) {
|
|
||||||
listLevel = slog.LevelInfo
|
|
||||||
hh.lastRelays = slices.Clone(relays)
|
|
||||||
}
|
|
||||||
hl := hostinfo.logger(rm.l)
|
|
||||||
hl.Log(context.Background(), listLevel, "Attempt to relay through hosts", "relays", relays)
|
|
||||||
|
|
||||||
// Send a RelayRequest to all known Relay IP's
|
// Send a RelayRequest to all known Relay IP's
|
||||||
for _, relay := range relays {
|
for _, relay := range hostinfo.remotes.relays {
|
||||||
// Don't relay through the host I'm trying to connect to
|
// Don't relay through the host I'm trying to connect to
|
||||||
if relay == vpnIp {
|
if relay == vpnIp {
|
||||||
continue
|
continue
|
||||||
@@ -87,19 +75,12 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Each relay's per-attempt log fires at Info on the first time we hit it and Debug after that.
|
|
||||||
level := slog.LevelInfo
|
|
||||||
if slices.Contains(prior, relay) {
|
|
||||||
level = slog.LevelDebug
|
|
||||||
}
|
|
||||||
|
|
||||||
relayHostInfo := rm.hostmap.QueryVpnAddr(relay)
|
relayHostInfo := rm.hostmap.QueryVpnAddr(relay)
|
||||||
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
||||||
hl.Log(context.Background(), level, "Establish tunnel to relay target", "relay", relay.String())
|
hostinfo.logger(rm.l).Info("Establish tunnel to relay target", "relay", relay.String())
|
||||||
f.Handshake(relay)
|
f.Handshake(relay)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check the relay HostInfo to see if we already established a relay through
|
// Check the relay HostInfo to see if we already established a relay through
|
||||||
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
||||||
if !ok {
|
if !ok {
|
||||||
@@ -107,7 +88,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
if relayHostInfo.remote.IsValid() {
|
if relayHostInfo.remote.IsValid() {
|
||||||
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
hostinfo.logger(rm.l).Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
m := NebulaControl{
|
m := NebulaControl{
|
||||||
@@ -118,12 +99,12 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
switch relayHostInfo.GetCert().Certificate.Version() {
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
case cert.Version1:
|
case cert.Version1:
|
||||||
if !f.myVpnAddrs[0].Is4() {
|
if !f.myVpnAddrs[0].Is4() {
|
||||||
hl.Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if !vpnIp.Is4() {
|
if !vpnIp.Is4() {
|
||||||
hl.Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -135,16 +116,16 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
||||||
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
default:
|
default:
|
||||||
hl.Error("Unknown certificate version found while creating relay")
|
hostinfo.logger(rm.l).Error("Unknown certificate version found while creating relay")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
msg, err := m.Marshal()
|
msg, err := m.Marshal()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hl.Error("Failed to marshal Control message to create relay", "error", err)
|
hostinfo.logger(rm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
rm.l.Log(context.Background(), level, "send CreateRelayRequest",
|
rm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", f.myVpnAddrs[0],
|
"relayFrom", f.myVpnAddrs[0],
|
||||||
"relayTo", vpnIp,
|
"relayTo", vpnIp,
|
||||||
"initiatorRelayIndex", idx,
|
"initiatorRelayIndex", idx,
|
||||||
@@ -157,14 +138,14 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
|
|
||||||
switch existingRelay.State {
|
switch existingRelay.State {
|
||||||
case Established:
|
case Established:
|
||||||
hl.Log(context.Background(), level, "Send handshake via relay", "relay", relay.String())
|
hostinfo.logger(rm.l).Info("Send handshake via relay", "relay", relay.String())
|
||||||
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false)
|
f.SendVia(relayHostInfo, existingRelay, stage0, make([]byte, 12), make([]byte, mtu), false)
|
||||||
case Disestablished:
|
case Disestablished:
|
||||||
// Mark this relay as 'requested'
|
// Mark this relay as 'requested'
|
||||||
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||||
fallthrough
|
fallthrough
|
||||||
case Requested:
|
case Requested:
|
||||||
hl.Log(context.Background(), level, "Re-send CreateRelay request", "relay", relay.String())
|
hostinfo.logger(rm.l).Info("Re-send CreateRelay request", "relay", relay.String())
|
||||||
// Re-send the CreateRelay request, in case the previous one was lost.
|
// Re-send the CreateRelay request, in case the previous one was lost.
|
||||||
m := NebulaControl{
|
m := NebulaControl{
|
||||||
Type: NebulaControl_CreateRelayRequest,
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
@@ -174,12 +155,12 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
switch relayHostInfo.GetCert().Certificate.Version() {
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
case cert.Version1:
|
case cert.Version1:
|
||||||
if !f.myVpnAddrs[0].Is4() {
|
if !f.myVpnAddrs[0].Is4() {
|
||||||
hl.Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if !vpnIp.Is4() {
|
if !vpnIp.Is4() {
|
||||||
hl.Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -191,16 +172,16 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
||||||
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
default:
|
default:
|
||||||
hl.Error("Unknown certificate version found while creating relay")
|
hostinfo.logger(rm.l).Error("Unknown certificate version found while creating relay")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
msg, err := m.Marshal()
|
msg, err := m.Marshal()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hl.Error("Failed to marshal Control message to create relay", "error", err)
|
hostinfo.logger(rm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
// This must send over the hostinfo, not over hm.Hosts[ip]
|
// This must send over the hostinfo, not over hm.Hosts[ip]
|
||||||
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
rm.l.Log(context.Background(), level, "send CreateRelayRequest",
|
rm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", f.myVpnAddrs[0],
|
"relayFrom", f.myVpnAddrs[0],
|
||||||
"relayTo", vpnIp,
|
"relayTo", vpnIp,
|
||||||
"initiatorRelayIndex", existingRelay.LocalIndex,
|
"initiatorRelayIndex", existingRelay.LocalIndex,
|
||||||
@@ -211,7 +192,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak
|
|||||||
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
||||||
fallthrough
|
fallthrough
|
||||||
default:
|
default:
|
||||||
hl.Error("Relay unexpected state",
|
hostinfo.logger(rm.l).Error("Relay unexpected state",
|
||||||
"vpnIp", vpnIp,
|
"vpnIp", vpnIp,
|
||||||
"state", existingRelay.State,
|
"state", existingRelay.State,
|
||||||
"relay", relay,
|
"relay", relay,
|
||||||
@@ -318,16 +299,17 @@ func (rm *relayManager) HandleControlMsg(h *HostInfo, d []byte, f *Interface) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f *Interface, m *NebulaControl) {
|
func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f *Interface, m *NebulaControl) {
|
||||||
relayFrom := protoAddrToNetAddr(m.RelayFromAddr)
|
|
||||||
relayTo := protoAddrToNetAddr(m.RelayToAddr)
|
|
||||||
rm.l.Info("handleCreateRelayResponse",
|
rm.l.Info("handleCreateRelayResponse",
|
||||||
"relayFrom", relayFrom,
|
"relayFrom", protoAddrToNetAddr(m.RelayFromAddr),
|
||||||
"relayTo", relayTo,
|
"relayTo", protoAddrToNetAddr(m.RelayToAddr),
|
||||||
"initiatorRelayIndex", m.InitiatorRelayIndex,
|
"initiatorRelayIndex", m.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", m.ResponderRelayIndex,
|
"responderRelayIndex", m.ResponderRelayIndex,
|
||||||
"vpnAddrs", h.vpnAddrs,
|
"vpnAddrs", h.vpnAddrs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
target := m.RelayToAddr
|
||||||
|
targetAddr := protoAddrToNetAddr(target)
|
||||||
|
|
||||||
relay, err := rm.EstablishRelay(h, m)
|
relay, err := rm.EstablishRelay(h, m)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
rm.l.Error("Failed to update relay for relayTo", "error", err)
|
rm.l.Error("Failed to update relay for relayTo", "error", err)
|
||||||
@@ -343,7 +325,7 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f
|
|||||||
rm.l.Error("Can't find a HostInfo for peer", "relayTo", relay.PeerAddr)
|
rm.l.Error("Can't find a HostInfo for peer", "relayTo", relay.PeerAddr)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
peerRelay, ok := peerHostInfo.relayState.QueryRelayForByIp(relayTo)
|
peerRelay, ok := peerHostInfo.relayState.QueryRelayForByIp(targetAddr)
|
||||||
if !ok {
|
if !ok {
|
||||||
rm.l.Error("peerRelay does not have Relay state for relayTo", "relayTo", peerHostInfo.vpnAddrs[0])
|
rm.l.Error("peerRelay does not have Relay state for relayTo", "relayTo", peerHostInfo.vpnAddrs[0])
|
||||||
return
|
return
|
||||||
@@ -353,19 +335,19 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f
|
|||||||
// I initiated the request to this peer, but haven't heard back from the peer yet. I must wait for this peer
|
// I initiated the request to this peer, but haven't heard back from the peer yet. I must wait for this peer
|
||||||
// to respond to complete the connection.
|
// to respond to complete the connection.
|
||||||
case PeerRequested, Disestablished, Established:
|
case PeerRequested, Disestablished, Established:
|
||||||
peerHostInfo.relayState.UpdateRelayForByIpState(relayTo, Established)
|
peerHostInfo.relayState.UpdateRelayForByIpState(targetAddr, Established)
|
||||||
resp := NebulaControl{
|
resp := NebulaControl{
|
||||||
Type: NebulaControl_CreateRelayResponse,
|
Type: NebulaControl_CreateRelayResponse,
|
||||||
ResponderRelayIndex: peerRelay.LocalIndex,
|
ResponderRelayIndex: peerRelay.LocalIndex,
|
||||||
InitiatorRelayIndex: peerRelay.RemoteIndex,
|
InitiatorRelayIndex: peerRelay.RemoteIndex,
|
||||||
}
|
}
|
||||||
|
|
||||||
peer := peerHostInfo.vpnAddrs[0]
|
|
||||||
if v == cert.Version1 {
|
if v == cert.Version1 {
|
||||||
|
peer := peerHostInfo.vpnAddrs[0]
|
||||||
if !peer.Is4() {
|
if !peer.Is4() {
|
||||||
rm.l.Error("Refusing to CreateRelayResponse for a v1 relay with an ipv6 address",
|
rm.l.Error("Refusing to CreateRelayResponse for a v1 relay with an ipv6 address",
|
||||||
"relayFrom", peer,
|
"relayFrom", peer,
|
||||||
"relayTo", relayTo,
|
"relayTo", target,
|
||||||
"initiatorRelayIndex", resp.InitiatorRelayIndex,
|
"initiatorRelayIndex", resp.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", resp.ResponderRelayIndex,
|
"responderRelayIndex", resp.ResponderRelayIndex,
|
||||||
"vpnAddrs", peerHostInfo.vpnAddrs,
|
"vpnAddrs", peerHostInfo.vpnAddrs,
|
||||||
@@ -375,26 +357,26 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f
|
|||||||
|
|
||||||
b := peer.As4()
|
b := peer.As4()
|
||||||
resp.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
resp.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
b = relayTo.As4()
|
b = targetAddr.As4()
|
||||||
resp.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
resp.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
} else {
|
} else {
|
||||||
resp.RelayFromAddr = netAddrToProtoAddr(peer)
|
resp.RelayFromAddr = netAddrToProtoAddr(peerHostInfo.vpnAddrs[0])
|
||||||
resp.RelayToAddr = m.RelayToAddr
|
resp.RelayToAddr = target
|
||||||
}
|
}
|
||||||
|
|
||||||
msg, err := resp.Marshal()
|
msg, err := resp.Marshal()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
rm.l.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
rm.l.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
||||||
return
|
} else {
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
|
rm.l.Info("send CreateRelayResponse",
|
||||||
|
"relayFrom", resp.RelayFromAddr,
|
||||||
|
"relayTo", resp.RelayToAddr,
|
||||||
|
"initiatorRelayIndex", resp.InitiatorRelayIndex,
|
||||||
|
"responderRelayIndex", resp.ResponderRelayIndex,
|
||||||
|
"vpnAddrs", peerHostInfo.vpnAddrs,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
|
||||||
rm.l.Info("send CreateRelayResponse",
|
|
||||||
"relayFrom", peer,
|
|
||||||
"relayTo", relayTo,
|
|
||||||
"initiatorRelayIndex", resp.InitiatorRelayIndex,
|
|
||||||
"responderRelayIndex", resp.ResponderRelayIndex,
|
|
||||||
"vpnAddrs", peerHostInfo.vpnAddrs,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,97 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestStartRelaysLogDedupe verifies that repeated attempts with the same relay set drop the log
|
|
||||||
// chatter to Debug, mirroring how the normal handshake retry loop quiets down once it's already
|
|
||||||
// announced its targets.
|
|
||||||
func TestStartRelaysLogDedupe(t *testing.T) {
|
|
||||||
vpnIp := netip.MustParseAddr("100.64.99.4")
|
|
||||||
otherRelay := netip.MustParseAddr("100.64.99.5")
|
|
||||||
|
|
||||||
newHH := func() *HandshakeHostInfo {
|
|
||||||
// Use the target's own vpnIp as the "relay" so the loop body skips it without
|
|
||||||
// touching any sender-side state. That isolates the test to the level-selection
|
|
||||||
// behavior of the top-level "Attempt to relay through hosts" log.
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
|
||||||
localIndexId: 1,
|
|
||||||
remotes: NewRemoteList([]netip.Addr{vpnIp}, nil),
|
|
||||||
}
|
|
||||||
hostinfo.remotes.relays = []netip.Addr{vpnIp}
|
|
||||||
return &HandshakeHostInfo{hostinfo: hostinfo}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Park any extra relay addresses we'll introduce mid-test in myVpnAddrsTable so the loop
|
|
||||||
// body always skips before touching f.Handshake (which would need a real handshakeManager).
|
|
||||||
addrTable := new(bart.Lite)
|
|
||||||
addrTable.Insert(netip.PrefixFrom(otherRelay, otherRelay.BitLen()))
|
|
||||||
f := &Interface{myVpnAddrsTable: addrTable}
|
|
||||||
|
|
||||||
newRM := func(buf *bytes.Buffer) *relayManager {
|
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
|
||||||
rm := &relayManager{l: l, hostmap: newHostMap(l)}
|
|
||||||
rm.useRelays.Store(true)
|
|
||||||
return rm
|
|
||||||
}
|
|
||||||
|
|
||||||
const msg = `msg="Attempt to relay through hosts"`
|
|
||||||
|
|
||||||
t.Run("first attempt logs at Info", func(t *testing.T) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rm := newRM(&buf)
|
|
||||||
hh := newHH()
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
assert.Equal(t, []netip.Addr{vpnIp}, hh.lastRelays, "lastRelays should record the relay set we just attempted")
|
|
||||||
assert.Contains(t, buf.String(), "level=INFO "+msg, "expected Info level on first attempt")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("repeat attempt with same relays drops to Debug", func(t *testing.T) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rm := newRM(&buf)
|
|
||||||
hh := newHH()
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
first := append([]netip.Addr(nil), hh.lastRelays...)
|
|
||||||
buf.Reset()
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
assert.Equal(t, first, hh.lastRelays)
|
|
||||||
assert.Contains(t, buf.String(), "level=DEBUG "+msg, "expected Debug level on identical retry")
|
|
||||||
assert.NotContains(t, buf.String(), "level=INFO "+msg, "Info should not fire on identical retry")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("changed relay list bumps back to Info", func(t *testing.T) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rm := newRM(&buf)
|
|
||||||
hh := newHH()
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
buf.Reset()
|
|
||||||
|
|
||||||
// The lighthouse handed us a new set this round.
|
|
||||||
hh.hostinfo.remotes.relays = []netip.Addr{vpnIp, otherRelay}
|
|
||||||
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
assert.Equal(t, []netip.Addr{vpnIp, otherRelay}, hh.lastRelays)
|
|
||||||
assert.Contains(t, buf.String(), "level=INFO "+msg, "expected Info when the relay list changes")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("disabled relays clears lastRelays and emits no Attempt log", func(t *testing.T) {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
rm := newRM(&buf)
|
|
||||||
rm.useRelays.Store(false)
|
|
||||||
hh := newHH()
|
|
||||||
hh.lastRelays = []netip.Addr{vpnIp}
|
|
||||||
|
|
||||||
rm.StartRelays(f, vpnIp, hh, nil)
|
|
||||||
assert.Nil(t, hh.lastRelays, "with relays disabled lastRelays should be cleared")
|
|
||||||
assert.NotContains(t, buf.String(), msg, "should not log when we shortcut out")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -239,31 +239,6 @@ func (r *RemoteList) unlockedSetHostnamesResults(hr *hostnamesResults) {
|
|||||||
r.hr = hr
|
r.hr = hr
|
||||||
}
|
}
|
||||||
|
|
||||||
// ResetForOwner zeros the reported address slices for the given owner and marks the addrs list dirty.
|
|
||||||
// Any pending hostname resolution will be canceled.
|
|
||||||
func (r *RemoteList) ResetForOwner(ownerVpnAddr netip.Addr) {
|
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
r.hr.Cancel()
|
|
||||||
if c, ok := r.cache[ownerVpnAddr]; ok {
|
|
||||||
if c.v4 != nil {
|
|
||||||
c.v4.reported = c.v4.reported[:0]
|
|
||||||
}
|
|
||||||
if c.v6 != nil {
|
|
||||||
c.v6.reported = c.v6.reported[:0]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
r.shouldRebuild = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// ClearHostnameResults cancels the in-flight DNS resolver goroutine (if any) and drops the resolved IP cache.
|
|
||||||
func (r *RemoteList) ClearHostnameResults() {
|
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
r.unlockedSetHostnamesResults(nil)
|
|
||||||
r.shouldRebuild = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Len locks and reports the size of the deduplicated address list
|
// Len locks and reports the size of the deduplicated address list
|
||||||
// The deduplication work may need to occur here, so you must pass preferredRanges
|
// The deduplication work may need to occur here, so you must pass preferredRanges
|
||||||
func (r *RemoteList) Len(preferredRanges []netip.Prefix) int {
|
func (r *RemoteList) Len(preferredRanges []netip.Prefix) int {
|
||||||
|
|||||||
@@ -6,22 +6,8 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// trackedHostnameResults builds a *hostnamesResults with a known cancel function and a
|
|
||||||
// pre-populated ips map so tests can assert cancellation and verify previously-resolved
|
|
||||||
// IPs survive a cancel without spinning up a real DNS resolver.
|
|
||||||
func trackedHostnameResults(cancelFn func(), addrs ...string) *hostnamesResults {
|
|
||||||
hr := &hostnamesResults{cancelFn: cancelFn}
|
|
||||||
ips := map[netip.AddrPort]struct{}{}
|
|
||||||
for _, a := range addrs {
|
|
||||||
ips[netip.MustParseAddrPort(a)] = struct{}{}
|
|
||||||
}
|
|
||||||
hr.ips.Store(&ips)
|
|
||||||
return hr
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteList_Rebuild(t *testing.T) {
|
func TestRemoteList_Rebuild(t *testing.T) {
|
||||||
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("0.0.0.0")}, nil)
|
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("0.0.0.0")}, nil)
|
||||||
rl.unlockedSetV4(
|
rl.unlockedSetV4(
|
||||||
@@ -126,81 +112,6 @@ func TestRemoteList_Rebuild(t *testing.T) {
|
|||||||
assert.Equal(t, "172.31.0.1:10101", rl.addrs[9].String())
|
assert.Equal(t, "172.31.0.1:10101", rl.addrs[9].String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRemoteList_ResetForOwner(t *testing.T) {
|
|
||||||
ourselves := netip.MustParseAddr("10.0.0.1")
|
|
||||||
otherOwner := netip.MustParseAddr("10.0.0.2")
|
|
||||||
vpnAddr := netip.MustParseAddr("10.0.0.99")
|
|
||||||
|
|
||||||
rl := NewRemoteList([]netip.Addr{vpnAddr}, nil)
|
|
||||||
rl.unlockedSetV4(ourselves, vpnAddr,
|
|
||||||
[]*V4AddrPort{newIp4AndPortFromString("1.1.1.1:4242")},
|
|
||||||
func(netip.Addr, *V4AddrPort) bool { return true },
|
|
||||||
)
|
|
||||||
rl.unlockedSetV6(ourselves, vpnAddr,
|
|
||||||
[]*V6AddrPort{newIp6AndPortFromString("[1::1]:4242")},
|
|
||||||
func(netip.Addr, *V6AddrPort) bool { return true },
|
|
||||||
)
|
|
||||||
rl.unlockedSetV4(otherOwner, vpnAddr,
|
|
||||||
[]*V4AddrPort{newIp4AndPortFromString("2.2.2.2:4242")},
|
|
||||||
func(netip.Addr, *V4AddrPort) bool { return true },
|
|
||||||
)
|
|
||||||
|
|
||||||
canceled := 0
|
|
||||||
hr := trackedHostnameResults(func() { canceled++ }, "3.3.3.3:4242")
|
|
||||||
rl.Lock()
|
|
||||||
rl.unlockedSetHostnamesResults(hr)
|
|
||||||
rl.Unlock()
|
|
||||||
|
|
||||||
rl.ResetForOwner(ourselves)
|
|
||||||
|
|
||||||
rl.RLock()
|
|
||||||
defer rl.RUnlock()
|
|
||||||
assert.Empty(t, rl.cache[ourselves].v4.reported, "our v4 reported should be cleared")
|
|
||||||
assert.Empty(t, rl.cache[ourselves].v6.reported, "our v6 reported should be cleared")
|
|
||||||
assert.Len(t, rl.cache[otherOwner].v4.reported, 1, "other owner's contribution must be preserved")
|
|
||||||
assert.Equal(t, "2.2.2.2:4242", protoV4AddrPortToNetAddrPort(rl.cache[otherOwner].v4.reported[0]).String())
|
|
||||||
assert.Equal(t, 1, canceled, "DNS resolution goroutine should be canceled")
|
|
||||||
assert.Same(t, hr, rl.hr, "hostnamesResults must be preserved so DNS-resolved IPs keep feeding addrs until replaced")
|
|
||||||
assert.NotEmpty(t, rl.hr.GetAddrs(), "previously-resolved IPs should still be readable after cancel")
|
|
||||||
assert.True(t, rl.shouldRebuild, "shouldRebuild must be set so the next Rebuild recomputes addrs")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteList_ResetForOwner_NoEntry(t *testing.T) {
|
|
||||||
// An owner with no cache entry must not panic; shouldRebuild is still set and any
|
|
||||||
// existing hostnamesResults is canceled.
|
|
||||||
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("10.0.0.99")}, nil)
|
|
||||||
canceled := 0
|
|
||||||
rl.Lock()
|
|
||||||
rl.unlockedSetHostnamesResults(trackedHostnameResults(func() { canceled++ }, "3.3.3.3:4242"))
|
|
||||||
rl.Unlock()
|
|
||||||
|
|
||||||
rl.ResetForOwner(netip.MustParseAddr("10.0.0.1"))
|
|
||||||
|
|
||||||
rl.RLock()
|
|
||||||
defer rl.RUnlock()
|
|
||||||
assert.Equal(t, 1, canceled)
|
|
||||||
assert.True(t, rl.shouldRebuild)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteList_ClearHostnameResults(t *testing.T) {
|
|
||||||
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("10.0.0.99")}, nil)
|
|
||||||
|
|
||||||
canceled := 0
|
|
||||||
hr := trackedHostnameResults(func() { canceled++ }, "3.3.3.3:4242")
|
|
||||||
rl.Lock()
|
|
||||||
rl.unlockedSetHostnamesResults(hr)
|
|
||||||
rl.Unlock()
|
|
||||||
require.NotEmpty(t, hr.GetAddrs(), "hostnamesResults should have its fastrack IPs populated")
|
|
||||||
|
|
||||||
rl.ClearHostnameResults()
|
|
||||||
|
|
||||||
rl.RLock()
|
|
||||||
defer rl.RUnlock()
|
|
||||||
assert.Equal(t, 1, canceled, "DNS resolution goroutine should be canceled")
|
|
||||||
assert.Nil(t, rl.hr, "hostnamesResults should be dropped")
|
|
||||||
assert.True(t, rl.shouldRebuild)
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkFullRebuild(b *testing.B) {
|
func BenchmarkFullRebuild(b *testing.B) {
|
||||||
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("0.0.0.0")}, nil)
|
rl := NewRemoteList([]netip.Addr{netip.MustParseAddr("0.0.0.0")}, nil)
|
||||||
rl.unlockedSetV4(
|
rl.unlockedSetV4(
|
||||||
|
|||||||
@@ -331,7 +331,7 @@ func loadStatsConfig(c *config.C) (statsConfig, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
cfg.interval = c.GetDuration("stats.interval", 0)
|
cfg.interval = c.GetDuration("stats.interval", 0)
|
||||||
if cfg.interval <= 0 {
|
if cfg.interval == 0 {
|
||||||
return cfg, fmt.Errorf("stats.interval was an invalid duration: %s", c.GetString("stats.interval", ""))
|
return cfg, fmt.Errorf("stats.interval was an invalid duration: %s", c.GetString("stats.interval", ""))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+38
-2
@@ -8,16 +8,49 @@ import (
|
|||||||
|
|
||||||
const MTU = 9001
|
const MTU = 9001
|
||||||
|
|
||||||
|
// MaxWriteBatch is the largest batch any Conn.WriteBatch implementation is
|
||||||
|
// required to accept. Callers SHOULD NOT pass more than this per call; Linux
|
||||||
|
// backends preallocate sendmmsg scratch sized to this value, so exceeding it
|
||||||
|
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||||
|
const MaxWriteBatch = 128
|
||||||
|
|
||||||
|
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
|
||||||
|
// data, kernel offload state, etc.) and passed to EncReader callbacks.
|
||||||
|
// Backends that do not produce a particular signal leave its zero value.
|
||||||
|
//
|
||||||
|
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
|
||||||
|
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
|
||||||
|
// means Not-ECT, which is also the value backends without ECN RX support
|
||||||
|
// supply on every packet.
|
||||||
|
type RxMeta struct {
|
||||||
|
OuterECN byte
|
||||||
|
}
|
||||||
|
|
||||||
type EncReader func(
|
type EncReader func(
|
||||||
addr netip.AddrPort,
|
addr netip.AddrPort,
|
||||||
payload []byte,
|
payload []byte,
|
||||||
|
meta RxMeta,
|
||||||
)
|
)
|
||||||
|
|
||||||
type Conn interface {
|
type Conn interface {
|
||||||
Rebind() error
|
Rebind() error
|
||||||
LocalAddr() (netip.AddrPort, error)
|
LocalAddr() (netip.AddrPort, error)
|
||||||
ListenOut(r EncReader) error
|
// ListenOut invokes r for each received packet. On batch-capable
|
||||||
|
// backends (recvmmsg), flush is called after each batch is fully
|
||||||
|
// delivered — callers use it to flush per-batch accumulators such as
|
||||||
|
// TUN write coalescers. Single-packet backends call flush after each
|
||||||
|
// packet. flush must not be nil.
|
||||||
|
ListenOut(r EncReader, flush func()) error
|
||||||
WriteTo(b []byte, addr netip.AddrPort) error
|
WriteTo(b []byte, addr netip.AddrPort) error
|
||||||
|
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||||
|
// destination. bufs and addrs must have the same length. outerECNs may
|
||||||
|
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
|
||||||
|
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
|
||||||
|
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
|
||||||
|
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
|
||||||
|
// cmsg; other backends ignore it. Returns on the first error; callers
|
||||||
|
// may observe a partial send if some packets went out before the error.
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
ReloadConfig(c *config.C)
|
ReloadConfig(c *config.C)
|
||||||
SupportsMultipleReaders() bool
|
SupportsMultipleReaders() bool
|
||||||
Close() error
|
Close() error
|
||||||
@@ -31,7 +64,7 @@ func (NoopConn) Rebind() error {
|
|||||||
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
return netip.AddrPort{}, nil
|
return netip.AddrPort{}, nil
|
||||||
}
|
}
|
||||||
func (NoopConn) ListenOut(_ EncReader) error {
|
func (NoopConn) ListenOut(_ EncReader, _ func()) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (NoopConn) SupportsMultipleReaders() bool {
|
func (NoopConn) SupportsMultipleReaders() bool {
|
||||||
@@ -40,6 +73,9 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
|||||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-3
@@ -140,6 +140,15 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -165,7 +174,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
|||||||
return func() {}
|
return func() {}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
@@ -175,11 +184,12 @@ func (u *StdConn) ListenOut(r EncReader) error {
|
|||||||
if errors.Is(err, net.ErrClosed) {
|
if errors.Is(err, net.ErrClosed) {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
u.l.Error("unexpected udp socket receive error", "error", err)
|
u.l.Error("unexpected udp socket receive error", "error", err)
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -44,6 +44,15 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -73,7 +82,7 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) ListenOut(r EncReader) error {
|
func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -93,7 +102,8 @@ func (u *GenericConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+171
-14
@@ -24,6 +24,22 @@ type StdConn struct {
|
|||||||
isV4 bool
|
isV4 bool
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
batch int
|
batch int
|
||||||
|
|
||||||
|
// sendmmsg scratch. Each queue has its own StdConn, so no locking is
|
||||||
|
// needed. Sized to MaxWriteBatch at construction; WriteBatch chunks
|
||||||
|
// larger inputs.
|
||||||
|
writeMsgs []rawMessage
|
||||||
|
writeIovs []iovec
|
||||||
|
writeNames [][]byte
|
||||||
|
|
||||||
|
// sendmmsg(2) callback state. sendmmsgCB is bound once in NewListener
|
||||||
|
// to the sendmmsgRun method value so passing it to rawConn.Write does
|
||||||
|
// not allocate a fresh closure per send; sendmmsgN/Sent/Errno carry
|
||||||
|
// the inputs and outputs across the call without escaping locals.
|
||||||
|
sendmmsgCB func(fd uintptr) bool
|
||||||
|
sendmmsgN int
|
||||||
|
sendmmsgSent int
|
||||||
|
sendmmsgErrno syscall.Errno
|
||||||
}
|
}
|
||||||
|
|
||||||
func setReusePort(network, address string, c syscall.RawConn) error {
|
func setReusePort(network, address string, c syscall.RawConn) error {
|
||||||
@@ -70,9 +86,23 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
|||||||
}
|
}
|
||||||
out.isV4 = af == unix.AF_INET
|
out.isV4 = af == unix.AF_INET
|
||||||
|
|
||||||
|
out.prepareWriteMessages(MaxWriteBatch)
|
||||||
|
out.sendmmsgCB = out.sendmmsgRun
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) prepareWriteMessages(n int) {
|
||||||
|
u.writeMsgs = make([]rawMessage, n)
|
||||||
|
u.writeIovs = make([]iovec, n)
|
||||||
|
u.writeNames = make([][]byte, n)
|
||||||
|
|
||||||
|
for i := range u.writeMsgs {
|
||||||
|
u.writeNames[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
u.writeMsgs[i].Hdr.Name = &u.writeNames[i][0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -171,7 +201,7 @@ func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
|
|||||||
return int(n), true, nil
|
return int(n), true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutSingle(r EncReader) error {
|
func (u *StdConn) listenOutSingle(r EncReader, flush func()) error {
|
||||||
var err error
|
var err error
|
||||||
var n int
|
var n int
|
||||||
var from netip.AddrPort
|
var from netip.AddrPort
|
||||||
@@ -183,16 +213,33 @@ func (u *StdConn) listenOutSingle(r EncReader) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
||||||
r(from, buffer[:n])
|
// listenOutSingle uses ReadFromUDPAddrPort which discards cmsgs,
|
||||||
|
// so the outer ECN field is not visible on this path. Zero RxMeta
|
||||||
|
// (Not-ECT) means RFC 6040 combine is a no-op.
|
||||||
|
r(from, buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutBatch(r EncReader) error {
|
// readSockaddr decodes the source address out of a recvmmsg name buffer
|
||||||
|
func (u *StdConn) readSockaddr(name []byte) netip.AddrPort {
|
||||||
var ip netip.Addr
|
var ip netip.Addr
|
||||||
|
// It's ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||||
|
if u.isV4 {
|
||||||
|
ip, _ = netip.AddrFromSlice(name[4:8])
|
||||||
|
} else {
|
||||||
|
ip, _ = netip.AddrFromSlice(name[8:24])
|
||||||
|
}
|
||||||
|
return netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(name[2:4]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) listenOutBatch(r EncReader, flush func()) error {
|
||||||
var n int
|
var n int
|
||||||
var operr error
|
var operr error
|
||||||
|
|
||||||
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
bufSize := MTU
|
||||||
|
cmsgSpace := 0
|
||||||
|
msgs, buffers, names, _ := u.PrepareRawMessages(u.batch, bufSize, cmsgSpace)
|
||||||
|
|
||||||
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
||||||
//defining it outside the loop so it gets re-used
|
//defining it outside the loop so it gets re-used
|
||||||
@@ -211,22 +258,18 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for i := 0; i < n; i++ {
|
for i := 0; i < n; i++ {
|
||||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
r(u.readSockaddr(names[i]), buffers[i][:msgs[i].Len], RxMeta{})
|
||||||
if u.isV4 {
|
|
||||||
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
|
||||||
} else {
|
|
||||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
|
||||||
}
|
|
||||||
r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
if u.batch == 1 {
|
if u.batch == 1 {
|
||||||
return u.listenOutSingle(r)
|
return u.listenOutSingle(r, flush)
|
||||||
} else {
|
} else {
|
||||||
return u.listenOutBatch(r)
|
return u.listenOutBatch(r, flush)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,6 +278,120 @@ func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WriteBatch sends bufs via sendmmsg(2) using the preallocated scratch on
|
||||||
|
// StdConn. If supported, consecutive packets to the same destination with
|
||||||
|
// matching segment sizes (all but possibly the last) are coalesced into a
|
||||||
|
// single mmsghdr entry
|
||||||
|
//
|
||||||
|
// If sendmmsg returns an error and zero entries went out, we fall back to
|
||||||
|
// per-packet WriteTo for that chunk so the caller still gets best-effort
|
||||||
|
// delivery. On a partial send we resume at the first un-acked entry on
|
||||||
|
// the next iteration.
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i := 0; i < len(bufs); {
|
||||||
|
chunk := min(len(bufs)-i, len(u.writeMsgs))
|
||||||
|
|
||||||
|
for k := 0; k < chunk; k++ {
|
||||||
|
u.writeIovs[k].Base = &bufs[i+k][0]
|
||||||
|
setIovLen(&u.writeIovs[k], len(bufs[i+k]))
|
||||||
|
|
||||||
|
nlen, err := writeSockaddr(u.writeNames[k], addrs[i+k], u.isV4)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
hdr := &u.writeMsgs[k].Hdr
|
||||||
|
hdr.Iov = &u.writeIovs[k]
|
||||||
|
setMsgIovlen(hdr, 1)
|
||||||
|
hdr.Namelen = uint32(nlen)
|
||||||
|
}
|
||||||
|
|
||||||
|
sent, serr := u.sendmmsg(chunk)
|
||||||
|
if serr != nil && sent <= 0 {
|
||||||
|
// sendmmsg returns -1 / sent=0 when entry 0 itself failed; log
|
||||||
|
// that entry's destination and fall back to per-packet WriteTo
|
||||||
|
// for the whole chunk so the caller still gets best-effort
|
||||||
|
// delivery without duplicating packets the kernel accepted.
|
||||||
|
u.l.Warn("sendmmsg failed, falling back to per-packet WriteTo",
|
||||||
|
"err", serr,
|
||||||
|
"entries", chunk,
|
||||||
|
"entry0_dst", addrs[i],
|
||||||
|
"isV4", u.isV4,
|
||||||
|
)
|
||||||
|
for k := 0; k < chunk; k++ {
|
||||||
|
if werr := u.WriteTo(bufs[i+k], addrs[i+k]); werr != nil {
|
||||||
|
return werr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
i += chunk
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i += sent
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendmmsg issues sendmmsg(2) against the first n entries of u.writeMsgs.
|
||||||
|
// The bound u.sendmmsgCB is passed to rawConn.Write so no closure is
|
||||||
|
// allocated per call; inputs and outputs ride on the StdConn fields.
|
||||||
|
func (u *StdConn) sendmmsg(n int) (int, error) {
|
||||||
|
u.sendmmsgN = n
|
||||||
|
u.sendmmsgSent = 0
|
||||||
|
u.sendmmsgErrno = 0
|
||||||
|
if err := u.rawConn.Write(u.sendmmsgCB); err != nil {
|
||||||
|
return u.sendmmsgSent, err
|
||||||
|
}
|
||||||
|
if u.sendmmsgErrno != 0 {
|
||||||
|
return u.sendmmsgSent, &net.OpError{Op: "sendmmsg", Err: u.sendmmsgErrno}
|
||||||
|
}
|
||||||
|
return u.sendmmsgSent, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendmmsgRun is the rawConn.Write callback. It is bound once into
|
||||||
|
// u.sendmmsgCB at construction so it stays alloc-free in the hot path;
|
||||||
|
// inputs (sendmmsgN) and outputs (sendmmsgSent, sendmmsgErrno) ride on
|
||||||
|
// the receiver rather than escaping locals.
|
||||||
|
func (u *StdConn) sendmmsgRun(fd uintptr) bool {
|
||||||
|
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
|
||||||
|
uintptr(unsafe.Pointer(&u.writeMsgs[0])), uintptr(u.sendmmsgN),
|
||||||
|
0, 0, 0,
|
||||||
|
)
|
||||||
|
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
u.sendmmsgSent = int(r1)
|
||||||
|
u.sendmmsgErrno = errno
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeSockaddr encodes addr into buf (which must be at least
|
||||||
|
// SizeofSockaddrInet6 bytes). Returns the number of bytes used. If isV4 is
|
||||||
|
// true and addr is not a v4 (or v4-in-v6) address, returns an error.
|
||||||
|
func writeSockaddr(buf []byte, addr netip.AddrPort, isV4 bool) (int, error) {
|
||||||
|
ap := addr.Addr().Unmap()
|
||||||
|
if isV4 {
|
||||||
|
if !ap.Is4() {
|
||||||
|
return 0, ErrInvalidIPv6RemoteForSocket
|
||||||
|
}
|
||||||
|
// struct sockaddr_in: { sa_family_t(2), in_port_t(2, BE), in_addr(4), zero(8) }
|
||||||
|
// sa_family is host endian.
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
ip4 := ap.As4()
|
||||||
|
copy(buf[4:8], ip4[:])
|
||||||
|
clear(buf[8:16])
|
||||||
|
return unix.SizeofSockaddrInet4, nil
|
||||||
|
}
|
||||||
|
// struct sockaddr_in6: { sa_family_t(2), in_port_t(2, BE), flowinfo(4), in6_addr(16), scope_id(4) }
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET6)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
binary.NativeEndian.PutUint32(buf[4:8], 0)
|
||||||
|
ip6 := addr.Addr().As16()
|
||||||
|
copy(buf[8:24], ip6[:])
|
||||||
|
binary.NativeEndian.PutUint32(buf[24:28], 0)
|
||||||
|
return unix.SizeofSockaddrInet6, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) ReloadConfig(c *config.C) {
|
func (u *StdConn) ReloadConfig(c *config.C) {
|
||||||
b := c.GetInt("listen.read_buffer", 0)
|
b := c.GetInt("listen.read_buffer", 0)
|
||||||
if b > 0 {
|
if b > 0 {
|
||||||
|
|||||||
+29
-3
@@ -30,13 +30,18 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -48,7 +53,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint32(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+29
-3
@@ -33,13 +33,18 @@ type rawMessage struct {
|
|||||||
Pad0 [4]byte
|
Pad0 [4]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -51,7 +56,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint64(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+12
-2
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *slog.Logger, sa windows.Sockaddr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *RIOConn) ListenOut(r EncReader) error {
|
func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -161,7 +161,8 @@ func (u *RIOConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n])
|
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -316,6 +317,15 @@ func (u *RIOConn) WriteTo(buf []byte, ip netip.AddrPort) error {
|
|||||||
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *RIOConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
sa, err := windows.Getsockname(u.sock)
|
sa, err := windows.Getsockname(u.sock)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+11
-2
@@ -157,15 +157,24 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *TesterConn) ListenOut(r EncReader) error {
|
func (u *TesterConn) ListenOut(r EncReader, flush func()) error {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-u.done:
|
case <-u.done:
|
||||||
return os.ErrClosed
|
return os.ErrClosed
|
||||||
case p := <-u.RxPackets:
|
case p := <-u.RxPackets:
|
||||||
r(p.From, p.Data)
|
r(p.From, p.Data, RxMeta{})
|
||||||
p.Release()
|
p.Release()
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
package wire
|
||||||
|
|
||||||
|
// TunPacket is the unit a read from a tun device returns.
|
||||||
|
// On supported platforms, it may be a superpacket, but a single TunPacket will never have more than one destination.
|
||||||
|
type TunPacket struct {
|
||||||
|
// Bytes contains the actual packet
|
||||||
|
Bytes []byte
|
||||||
|
// Meta contains other information to help process the packet correctly, such as offsets for segmentation offloads
|
||||||
|
// Fields in Meta should be as portable/platform-agnostic as possible.
|
||||||
|
Meta struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// PerSegment invokes fn once per segment of pkt.
|
||||||
|
// This is a stub implementation that does not actually support segmentation
|
||||||
|
func (t *TunPacket) PerSegment(fn func(seg []byte) error) error {
|
||||||
|
return fn(t.Bytes)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user