mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 14:16:58 +02:00
Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 7794e93762 | |||
| 3df60ae195 |
@@ -209,11 +209,10 @@ jobs:
|
|||||||
id: create_release
|
id: create_release
|
||||||
env:
|
env:
|
||||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
GITHUB_REF_NAME: ${{ github.ref_name }}
|
|
||||||
run: |
|
run: |
|
||||||
cd artifacts
|
cd artifacts
|
||||||
gh release create \
|
gh release create \
|
||||||
--verify-tag \
|
--verify-tag \
|
||||||
--title "Release ${GITHUB_REF_NAME}" \
|
--title "Release ${{ github.ref_name }}" \
|
||||||
"${GITHUB_REF_NAME}" \
|
"${{ github.ref_name }}" \
|
||||||
SHASUM256.txt *-latest/*.zip *-latest/*.tar.gz
|
SHASUM256.txt *-latest/*.zip *-latest/*.tar.gz
|
||||||
|
|||||||
@@ -18,8 +18,6 @@ jobs:
|
|||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||||
name: Run extra smoke tests
|
name: Run extra smoke tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
env:
|
|
||||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v6
|
||||||
@@ -32,13 +30,11 @@ jobs:
|
|||||||
- name: add hashicorp source
|
- name: add hashicorp source
|
||||||
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
||||||
|
|
||||||
- name: install vagrant and libvirt
|
- name: workaround AMD-V issue # https://github.com/cri-o/packaging/pull/306
|
||||||
run: |
|
run: sudo rmmod kvm_amd
|
||||||
sudo apt-get update && sudo apt-get install -y vagrant libvirt-daemon-system libvirt-dev
|
|
||||||
sudo chmod 666 /dev/kvm
|
- name: install vagrant
|
||||||
sudo usermod -aG libvirt $(whoami)
|
run: sudo apt-get update && sudo apt-get install -y vagrant virtualbox
|
||||||
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
|
||||||
vagrant plugin install vagrant-libvirt
|
|
||||||
|
|
||||||
- name: freebsd-amd64
|
- name: freebsd-amd64
|
||||||
run: make smoke-vagrant/freebsd-amd64
|
run: make smoke-vagrant/freebsd-amd64
|
||||||
@@ -49,19 +45,10 @@ jobs:
|
|||||||
- name: netbsd-amd64
|
- name: netbsd-amd64
|
||||||
run: make smoke-vagrant/netbsd-amd64
|
run: make smoke-vagrant/netbsd-amd64
|
||||||
|
|
||||||
|
- name: linux-386
|
||||||
|
run: make smoke-vagrant/linux-386
|
||||||
|
|
||||||
- name: linux-amd64-ipv6disable
|
- name: linux-amd64-ipv6disable
|
||||||
run: make smoke-vagrant/linux-amd64-ipv6disable
|
run: make smoke-vagrant/linux-amd64-ipv6disable
|
||||||
|
|
||||||
# linux-386 runs last because it requires disabling KVM to use VirtualBox,
|
|
||||||
# which prevents libvirt (used by the other tests) from working after this point.
|
|
||||||
- name: install virtualbox for i386 test
|
|
||||||
run: |
|
|
||||||
sudo apt-get install -y virtualbox
|
|
||||||
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
|
||||||
|
|
||||||
- name: linux-386
|
|
||||||
env:
|
|
||||||
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
|
||||||
run: make smoke-vagrant/linux-386
|
|
||||||
|
|
||||||
timeout-minutes: 30
|
timeout-minutes: 30
|
||||||
|
|||||||
@@ -16,10 +16,8 @@ relay:
|
|||||||
am_relay: true
|
am_relay: true
|
||||||
EOF
|
EOF
|
||||||
|
|
||||||
# TEST-NET-3 placeholder IPs; smoke-relay.sh seds them to real container IPs.
|
export LIGHTHOUSES="192.168.100.1 172.17.0.2:4242"
|
||||||
# Mapping: .2 lighthouse1, .3 host2, .4 host3, .5 host4.
|
export REMOTE_ALLOW_LIST='{"172.17.0.4/32": false, "172.17.0.5/32": false}'
|
||||||
export LIGHTHOUSES="192.168.100.1 203.0.113.2:4242"
|
|
||||||
export REMOTE_ALLOW_LIST='{"203.0.113.4/32": false, "203.0.113.5/32": false}'
|
|
||||||
|
|
||||||
HOST="host2" ../genconfig.sh >host2.yml <<EOF
|
HOST="host2" ../genconfig.sh >host2.yml <<EOF
|
||||||
relay:
|
relay:
|
||||||
@@ -27,7 +25,7 @@ relay:
|
|||||||
- 192.168.100.1
|
- 192.168.100.1
|
||||||
EOF
|
EOF
|
||||||
|
|
||||||
export REMOTE_ALLOW_LIST='{"203.0.113.3/32": false}'
|
export REMOTE_ALLOW_LIST='{"172.17.0.3/32": false}'
|
||||||
|
|
||||||
HOST="host3" ../genconfig.sh >host3.yml
|
HOST="host3" ../genconfig.sh >host3.yml
|
||||||
|
|
||||||
|
|||||||
@@ -5,15 +5,9 @@ set -e -x
|
|||||||
rm -rf ./build
|
rm -rf ./build
|
||||||
mkdir ./build
|
mkdir ./build
|
||||||
|
|
||||||
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
# TODO: Assumes your docker bridge network is a /24, and the first container that launches will be .1
|
||||||
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
# - We could make this better by launching the lighthouse first and then fetching what IP it is.
|
||||||
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
NET="$(docker network inspect bridge -f '{{ range .IPAM.Config }}{{ .Subnet }}{{ end }}' | cut -d. -f1-3)"
|
||||||
# sed the real container IPs in before starting nebula.
|
|
||||||
#
|
|
||||||
# Placeholder mapping (last octet == fixed container slot):
|
|
||||||
# 203.0.113.2 -> lighthouse1, 203.0.113.3 -> host2,
|
|
||||||
# 203.0.113.4 -> host3, 203.0.113.5 -> host4.
|
|
||||||
LIGHTHOUSE_IP="203.0.113.2"
|
|
||||||
|
|
||||||
(
|
(
|
||||||
cd build
|
cd build
|
||||||
@@ -31,16 +25,16 @@ LIGHTHOUSE_IP="203.0.113.2"
|
|||||||
../genconfig.sh >lighthouse1.yml
|
../genconfig.sh >lighthouse1.yml
|
||||||
|
|
||||||
HOST="host2" \
|
HOST="host2" \
|
||||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $NET.2:4242" \
|
||||||
../genconfig.sh >host2.yml
|
../genconfig.sh >host2.yml
|
||||||
|
|
||||||
HOST="host3" \
|
HOST="host3" \
|
||||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $NET.2: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="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
LIGHTHOUSES="192.168.100.1 $NET.2:4242" \
|
||||||
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host4.yml
|
../genconfig.sh >host4.yml
|
||||||
|
|
||||||
|
|||||||
@@ -6,8 +6,6 @@ set -o pipefail
|
|||||||
|
|
||||||
mkdir -p logs
|
mkdir -p logs
|
||||||
|
|
||||||
NETWORK="nebula-smoke-relay"
|
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
echo
|
echo
|
||||||
echo " *** cleanup"
|
echo " *** cleanup"
|
||||||
@@ -18,53 +16,22 @@ cleanup() {
|
|||||||
then
|
then
|
||||||
docker kill lighthouse1 host2 host3 host4
|
docker kill lighthouse1 host2 host3 host4
|
||||||
fi
|
fi
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
|
||||||
}
|
}
|
||||||
|
|
||||||
trap cleanup EXIT
|
trap cleanup EXIT
|
||||||
|
|
||||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
|
||||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
|
||||||
# fail the whole test — we only need one to be free.
|
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
|
||||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
|
||||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
break
|
|
||||||
fi
|
|
||||||
done
|
|
||||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
|
||||||
# .3 host2, .4 host3, .5 host4 — matches the placeholders in build-relay.sh.
|
|
||||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
|
||||||
PREFIX="${SUBNET%/*}"
|
|
||||||
PREFIX="${PREFIX%.*}"
|
|
||||||
LIGHTHOUSE_IP="$PREFIX.2"
|
|
||||||
HOST2_IP="$PREFIX.3"
|
|
||||||
HOST3_IP="$PREFIX.4"
|
|
||||||
HOST4_IP="$PREFIX.5"
|
|
||||||
|
|
||||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
|
||||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
|
||||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
|
||||||
mv "$f.tmp" "$f"
|
|
||||||
done
|
|
||||||
|
|
||||||
docker run --name lighthouse1 --rm nebula:smoke-relay -config lighthouse1.yml -test
|
docker run --name lighthouse1 --rm nebula:smoke-relay -config lighthouse1.yml -test
|
||||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" nebula:smoke-relay -config host2.yml -test
|
docker run --name host2 --rm nebula:smoke-relay -config host2.yml -test
|
||||||
docker run --name host3 --rm -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" nebula:smoke-relay -config host3.yml -test
|
docker run --name host3 --rm nebula:smoke-relay -config host3.yml -test
|
||||||
docker run --name host4 --rm -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" nebula:smoke-relay -config host4.yml -test
|
docker run --name host4 --rm nebula:smoke-relay -config host4.yml -test
|
||||||
|
|
||||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host3 --network "$NETWORK" --ip "$HOST3_IP" -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
docker run --name host3 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host4 --network "$NETWORK" --ip "$HOST4_IP" -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
docker run --name host4 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
@@ -109,13 +76,7 @@ docker exec host4 sh -c 'kill 1'
|
|||||||
docker exec host3 sh -c 'kill 1'
|
docker exec host3 sh -c 'kill 1'
|
||||||
docker exec host2 sh -c 'kill 1'
|
docker exec host2 sh -c 'kill 1'
|
||||||
docker exec lighthouse1 sh -c 'kill 1'
|
docker exec lighthouse1 sh -c 'kill 1'
|
||||||
|
sleep 5
|
||||||
# Wait up to 30s for all backgrounded jobs to exit rather than relying on a
|
|
||||||
# fixed sleep.
|
|
||||||
for _ in $(seq 1 30); do
|
|
||||||
[ -z "$(jobs -r)" ] && break
|
|
||||||
sleep 1
|
|
||||||
done
|
|
||||||
|
|
||||||
if [ "$(jobs -r)" ]
|
if [ "$(jobs -r)" ]
|
||||||
then
|
then
|
||||||
|
|||||||
@@ -8,8 +8,6 @@ export VAGRANT_CWD="$PWD/vagrant-$1"
|
|||||||
|
|
||||||
mkdir -p logs
|
mkdir -p logs
|
||||||
|
|
||||||
NETWORK="nebula-smoke"
|
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
echo
|
echo
|
||||||
echo " *** cleanup"
|
echo " *** cleanup"
|
||||||
@@ -21,51 +19,21 @@ cleanup() {
|
|||||||
docker kill lighthouse1 host2
|
docker kill lighthouse1 host2
|
||||||
fi
|
fi
|
||||||
vagrant destroy -f
|
vagrant destroy -f
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
|
||||||
}
|
}
|
||||||
|
|
||||||
trap cleanup EXIT
|
trap cleanup EXIT
|
||||||
|
|
||||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
|
||||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
|
||||||
# fail the whole test — we only need one to be free.
|
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
|
||||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
|
||||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
break
|
|
||||||
fi
|
|
||||||
done
|
|
||||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
|
||||||
# .3 host2 — matches the placeholders in build.sh.
|
|
||||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
|
||||||
PREFIX="${SUBNET%/*}"
|
|
||||||
PREFIX="${PREFIX%.*}"
|
|
||||||
LIGHTHOUSE_IP="$PREFIX.2"
|
|
||||||
HOST2_IP="$PREFIX.3"
|
|
||||||
|
|
||||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
|
||||||
# This must happen before `vagrant up` rsyncs build/ into the VM for host3.
|
|
||||||
for f in build/host2.yml build/host3.yml; do
|
|
||||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
|
||||||
mv "$f.tmp" "$f"
|
|
||||||
done
|
|
||||||
|
|
||||||
CONTAINER="nebula:${NAME:-smoke}"
|
CONTAINER="nebula:${NAME:-smoke}"
|
||||||
|
|
||||||
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
||||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" "$CONTAINER" -config host2.yml -test
|
docker run --name host2 --rm "$CONTAINER" -config host2.yml -test
|
||||||
|
|
||||||
vagrant up
|
vagrant up
|
||||||
vagrant ssh -c "cd /nebula && /nebula/$1-nebula -config host3.yml -test" -- -T
|
vagrant ssh -c "cd /nebula && /nebula/$1-nebula -config host3.yml -test" -- -T
|
||||||
|
|
||||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
vagrant ssh -c "cd /nebula && sudo sh -c 'echo \$\$ >/nebula/pid && exec /nebula/$1-nebula -config host3.yml'" 2>&1 -- -T | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
vagrant ssh -c "cd /nebula && sudo sh -c 'echo \$\$ >/nebula/pid && exec /nebula/$1-nebula -config host3.yml'" 2>&1 -- -T | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||||
sleep 15
|
sleep 15
|
||||||
@@ -128,14 +96,7 @@ vagrant ssh -c "ping -c1 192.168.100.2" -- -T
|
|||||||
vagrant ssh -c "sudo xargs kill </nebula/pid" -- -T
|
vagrant ssh -c "sudo xargs kill </nebula/pid" -- -T
|
||||||
docker exec host2 sh -c 'kill 1'
|
docker exec host2 sh -c 'kill 1'
|
||||||
docker exec lighthouse1 sh -c 'kill 1'
|
docker exec lighthouse1 sh -c 'kill 1'
|
||||||
|
sleep 1
|
||||||
# Wait up to 30s for all backgrounded jobs to exit. vagrant ssh in particular
|
|
||||||
# takes a beat to tear down after nebula exits on the VM, so a fixed sleep is
|
|
||||||
# racy.
|
|
||||||
for _ in $(seq 1 30); do
|
|
||||||
[ -z "$(jobs -r)" ] && break
|
|
||||||
sleep 1
|
|
||||||
done
|
|
||||||
|
|
||||||
if [ "$(jobs -r)" ]
|
if [ "$(jobs -r)" ]
|
||||||
then
|
then
|
||||||
|
|||||||
@@ -6,8 +6,6 @@ set -o pipefail
|
|||||||
|
|
||||||
mkdir -p logs
|
mkdir -p logs
|
||||||
|
|
||||||
NETWORK="nebula-smoke"
|
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
echo
|
echo
|
||||||
echo " *** cleanup"
|
echo " *** cleanup"
|
||||||
@@ -18,56 +16,24 @@ cleanup() {
|
|||||||
then
|
then
|
||||||
docker kill lighthouse1 host2 host3 host4
|
docker kill lighthouse1 host2 host3 host4
|
||||||
fi
|
fi
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
|
||||||
}
|
}
|
||||||
|
|
||||||
trap cleanup EXIT
|
trap cleanup EXIT
|
||||||
|
|
||||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
|
||||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
|
||||||
# fail the whole test — we only need one to be free.
|
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
|
||||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
|
||||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
break
|
|
||||||
fi
|
|
||||||
done
|
|
||||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
|
||||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
|
|
||||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
|
||||||
# .3 host2, .4 host3, .5 host4 — matches the placeholders in build.sh.
|
|
||||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
|
||||||
PREFIX="${SUBNET%/*}"
|
|
||||||
PREFIX="${PREFIX%.*}"
|
|
||||||
LIGHTHOUSE_IP="$PREFIX.2"
|
|
||||||
HOST2_IP="$PREFIX.3"
|
|
||||||
HOST3_IP="$PREFIX.4"
|
|
||||||
HOST4_IP="$PREFIX.5"
|
|
||||||
|
|
||||||
# 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.
|
|
||||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
|
||||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
|
||||||
mv "$f.tmp" "$f"
|
|
||||||
done
|
|
||||||
|
|
||||||
CONTAINER="nebula:${NAME:-smoke}"
|
CONTAINER="nebula:${NAME:-smoke}"
|
||||||
|
|
||||||
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
||||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" "$CONTAINER" -config host2.yml -test
|
docker run --name host2 --rm "$CONTAINER" -config host2.yml -test
|
||||||
docker run --name host3 --rm -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" "$CONTAINER" -config host3.yml -test
|
docker run --name host3 --rm "$CONTAINER" -config host3.yml -test
|
||||||
docker run --name host4 --rm -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" "$CONTAINER" -config host4.yml -test
|
docker run --name host4 --rm "$CONTAINER" -config host4.yml -test
|
||||||
|
|
||||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host3 --network "$NETWORK" --ip "$HOST3_IP" -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
docker run --name host3 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
docker run --name host4 --network "$NETWORK" --ip "$HOST4_IP" -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
docker run --name host4 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||||
sleep 1
|
sleep 1
|
||||||
|
|
||||||
# grab tcpdump pcaps for debugging
|
# grab tcpdump pcaps for debugging
|
||||||
@@ -158,20 +124,13 @@ set -x
|
|||||||
# host2 speaking to host4 on UDP 4000 should allow it to reply, when firewall rules would normally not permit this
|
# host2 speaking to host4 on UDP 4000 should allow it to reply, when firewall rules would normally not permit this
|
||||||
docker exec host2 sh -c "/usr/bin/echo host2 | ncat -nuv 192.168.100.4 4000"
|
docker exec host2 sh -c "/usr/bin/echo host2 | ncat -nuv 192.168.100.4 4000"
|
||||||
docker exec host2 ncat -e '/usr/bin/echo helloagainfromhost2' -nkluv 0.0.0.0 4000 &
|
docker exec host2 ncat -e '/usr/bin/echo helloagainfromhost2' -nkluv 0.0.0.0 4000 &
|
||||||
sleep 1
|
|
||||||
docker exec host4 sh -c "/usr/bin/echo host4 | ncat -nuv 192.168.100.2 4000"
|
docker exec host4 sh -c "/usr/bin/echo host4 | ncat -nuv 192.168.100.2 4000"
|
||||||
|
|
||||||
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'
|
||||||
docker exec host2 sh -c 'kill 1'
|
docker exec host2 sh -c 'kill 1'
|
||||||
docker exec lighthouse1 sh -c 'kill 1'
|
docker exec lighthouse1 sh -c 'kill 1'
|
||||||
|
sleep 5
|
||||||
# Wait up to 30s for all backgrounded jobs to exit rather than relying on a
|
|
||||||
# fixed sleep.
|
|
||||||
for _ in $(seq 1 30); do
|
|
||||||
[ -z "$(jobs -r)" ] && break
|
|
||||||
sleep 1
|
|
||||||
done
|
|
||||||
|
|
||||||
if [ "$(jobs -r)" ]
|
if [ "$(jobs -r)" ]
|
||||||
then
|
then
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# -*- mode: ruby -*-
|
# -*- mode: ruby -*-
|
||||||
# vi: set ft=ruby :
|
# vi: set ft=ruby :
|
||||||
Vagrant.configure("2") do |config|
|
Vagrant.configure("2") do |config|
|
||||||
config.vm.box = "bento/ubuntu-24.04"
|
config.vm.box = "ubuntu/jammy64"
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula"
|
config.vm.synced_folder "../build", "/nebula"
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# -*- mode: ruby -*-
|
# -*- mode: ruby -*-
|
||||||
# vi: set ft=ruby :
|
# vi: set ft=ruby :
|
||||||
Vagrant.configure("2") do |config|
|
Vagrant.configure("2") do |config|
|
||||||
config.vm.box = "DefinedNet/openbsd78"
|
config.vm.box = "generic/openbsd7"
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
||||||
end
|
end
|
||||||
|
|||||||
@@ -57,7 +57,7 @@ Check the [releases](https://github.com/slackhq/nebula/releases/latest) page for
|
|||||||
docker pull nebulaoss/nebula
|
docker pull nebulaoss/nebula
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Mobile ([source code](https://github.com/DefinedNet/mobile_nebula))
|
#### Mobile
|
||||||
|
|
||||||
- [iOS](https://apps.apple.com/us/app/mobile-nebula/id1509587936?itsct=apps_box&itscg=30200)
|
- [iOS](https://apps.apple.com/us/app/mobile-nebula/id1509587936?itsct=apps_box&itscg=30200)
|
||||||
- [Android](https://play.google.com/store/apps/details?id=net.defined.mobile_nebula&pcampaignid=pcampaignidMKT-Other-global-all-co-prtnr-py-PartBadge-Mar2515-1)
|
- [Android](https://play.google.com/store/apps/details?id=net.defined.mobile_nebula&pcampaignid=pcampaignidMKT-Other-global-all-co-prtnr-py-PartBadge-Mar2515-1)
|
||||||
@@ -76,8 +76,6 @@ Nebula was created to provide a mechanism for groups of hosts to communicate sec
|
|||||||
|
|
||||||
## Getting started (quickly)
|
## Getting started (quickly)
|
||||||
|
|
||||||
**Don't want to manage your own PKI and lighthouses?** [Managed Nebula](https://www.defined.net/) from Defined Networking handles all of this for you.
|
|
||||||
|
|
||||||
To set up a Nebula network, you'll need:
|
To set up a Nebula network, you'll need:
|
||||||
|
|
||||||
#### 1. The [Nebula binaries](https://github.com/slackhq/nebula/releases) or [Distribution Packages](https://github.com/slackhq/nebula#distribution-packages) for your specific platform. Specifically you'll need `nebula-cert` and the specific nebula binary for each platform you use.
|
#### 1. The [Nebula binaries](https://github.com/slackhq/nebula/releases) or [Distribution Packages](https://github.com/slackhq/nebula#distribution-packages) for your specific platform. Specifically you'll need `nebula-cert` and the specific nebula binary for each platform you use.
|
||||||
|
|||||||
@@ -1,70 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import "net/netip"
|
|
||||||
|
|
||||||
// sendBatchCap is the maximum number of encrypted packets accumulated before a
|
|
||||||
// flush is forced. TSO superpackets segment to at most ~45 packets on
|
|
||||||
// reasonable MTUs, so 128 leaves headroom without bloating the backing
|
|
||||||
// allocation.
|
|
||||||
const sendBatchCap = 128
|
|
||||||
|
|
||||||
// sendBatch accumulates encrypted UDP packets for a single sendmmsg flush.
|
|
||||||
// One sendBatch is owned by each listenIn goroutine; no locking is needed.
|
|
||||||
// The backing storage holds up to batchCap packets of slotCap bytes each;
|
|
||||||
// bufs and dsts are parallel slices of committed slots.
|
|
||||||
type sendBatch struct {
|
|
||||||
bufs [][]byte
|
|
||||||
dsts []netip.AddrPort
|
|
||||||
backing []byte
|
|
||||||
slotCap int
|
|
||||||
batchCap int
|
|
||||||
nextSlot int
|
|
||||||
}
|
|
||||||
|
|
||||||
func newSendBatch(batchCap, slotCap int) *sendBatch {
|
|
||||||
return &sendBatch{
|
|
||||||
bufs: make([][]byte, 0, batchCap),
|
|
||||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
|
||||||
backing: make([]byte, batchCap*slotCap),
|
|
||||||
slotCap: slotCap,
|
|
||||||
batchCap: batchCap,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next returns a zero-length slice with slotCap capacity over the next unused
|
|
||||||
// slot's backing bytes. The caller writes into the returned slice and then
|
|
||||||
// calls Commit with the final length and destination. Next returns nil when
|
|
||||||
// the batch is full.
|
|
||||||
func (b *sendBatch) Next() []byte {
|
|
||||||
if b.nextSlot >= b.batchCap {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
start := b.nextSlot * b.slotCap
|
|
||||||
return b.backing[start : start : start+b.slotCap]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Commit records the slot just returned by Next as a packet of length n
|
|
||||||
// destined for dst.
|
|
||||||
func (b *sendBatch) Commit(n int, dst netip.AddrPort) {
|
|
||||||
start := b.nextSlot * b.slotCap
|
|
||||||
b.bufs = append(b.bufs, b.backing[start:start+n])
|
|
||||||
b.dsts = append(b.dsts, dst)
|
|
||||||
b.nextSlot++
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reset clears committed slots; backing storage is retained for reuse.
|
|
||||||
func (b *sendBatch) Reset() {
|
|
||||||
b.bufs = b.bufs[:0]
|
|
||||||
b.dsts = b.dsts[:0]
|
|
||||||
b.nextSlot = 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// Len returns the number of committed packets.
|
|
||||||
func (b *sendBatch) Len() int {
|
|
||||||
return len(b.bufs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Cap returns the maximum number of slots in the batch.
|
|
||||||
func (b *sendBatch) Cap() int {
|
|
||||||
return b.batchCap
|
|
||||||
}
|
|
||||||
@@ -1,69 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSendBatchBookkeeping(t *testing.T) {
|
|
||||||
b := newSendBatch(4, 32)
|
|
||||||
if b.Len() != 0 || b.Cap() != 4 {
|
|
||||||
t.Fatalf("fresh batch: len=%d cap=%d", b.Len(), b.Cap())
|
|
||||||
}
|
|
||||||
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
|
||||||
for i := 0; i < 4; i++ {
|
|
||||||
slot := b.Next()
|
|
||||||
if slot == nil {
|
|
||||||
t.Fatalf("slot %d: Next returned nil before cap", i)
|
|
||||||
}
|
|
||||||
if cap(slot) != 32 || len(slot) != 0 {
|
|
||||||
t.Fatalf("slot %d: got len=%d cap=%d want len=0 cap=32", i, len(slot), cap(slot))
|
|
||||||
}
|
|
||||||
// Write a marker byte.
|
|
||||||
slot = append(slot, byte(i), byte(i+1), byte(i+2))
|
|
||||||
b.Commit(len(slot), ap)
|
|
||||||
}
|
|
||||||
if b.Next() != nil {
|
|
||||||
t.Fatalf("Next should return nil when full")
|
|
||||||
}
|
|
||||||
if b.Len() != 4 {
|
|
||||||
t.Fatalf("Len=%d want 4", b.Len())
|
|
||||||
}
|
|
||||||
for i, buf := range b.bufs {
|
|
||||||
if len(buf) != 3 || buf[0] != byte(i) {
|
|
||||||
t.Errorf("buf %d: %x", i, buf)
|
|
||||||
}
|
|
||||||
if b.dsts[i] != ap {
|
|
||||||
t.Errorf("dst %d: got %v want %v", i, b.dsts[i], ap)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reset returns empty and Next works again.
|
|
||||||
b.Reset()
|
|
||||||
if b.Len() != 0 {
|
|
||||||
t.Fatalf("after Reset Len=%d want 0", b.Len())
|
|
||||||
}
|
|
||||||
slot := b.Next()
|
|
||||||
if slot == nil || cap(slot) != 32 {
|
|
||||||
t.Fatalf("after Reset Next nil or wrong cap: %v cap=%d", slot == nil, cap(slot))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
|
||||||
b := newSendBatch(3, 8)
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
|
||||||
|
|
||||||
// Fill three slots, each with its own sentinel byte.
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
s := b.Next()
|
|
||||||
s = append(s, byte(0xA0+i), byte(0xB0+i))
|
|
||||||
b.Commit(len(s), ap)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, buf := range b.bufs {
|
|
||||||
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
|
||||||
t.Errorf("slot %d corrupted: %x", i, buf)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+9
-36
@@ -1,14 +1,11 @@
|
|||||||
package cert
|
package cert
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
|
||||||
"bytes"
|
|
||||||
"encoding/pem"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -32,46 +29,22 @@ func NewCAPool() *CAPool {
|
|||||||
// If the pool contains any expired certificates, an ErrExpired will be
|
// If the pool contains any expired certificates, an ErrExpired will be
|
||||||
// returned along with the pool. The caller must handle any such errors.
|
// returned along with the pool. The caller must handle any such errors.
|
||||||
func NewCAPoolFromPEM(caPEMs []byte) (*CAPool, error) {
|
func NewCAPoolFromPEM(caPEMs []byte) (*CAPool, error) {
|
||||||
return NewCAPoolFromPEMReader(bytes.NewReader(caPEMs))
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCAPoolFromPEMReader will create a new CA pool from the provided reader.
|
|
||||||
// The reader must contain a PEM-encoded set of nebula certificates.
|
|
||||||
func NewCAPoolFromPEMReader(r io.Reader) (*CAPool, error) {
|
|
||||||
pool := NewCAPool()
|
pool := NewCAPool()
|
||||||
|
var err error
|
||||||
var expired bool
|
var expired bool
|
||||||
|
for {
|
||||||
scanner := bufio.NewScanner(r)
|
caPEMs, err = pool.AddCAFromPEM(caPEMs)
|
||||||
scanner.Split(SplitPEM)
|
if errors.Is(err, ErrExpired) {
|
||||||
|
expired = true
|
||||||
for scanner.Scan() {
|
err = nil
|
||||||
pemBytes := scanner.Bytes()
|
|
||||||
|
|
||||||
block, rest := pem.Decode(pemBytes)
|
|
||||||
if len(bytes.TrimSpace(rest)) > 0 {
|
|
||||||
return nil, ErrInvalidPEMBlock
|
|
||||||
}
|
}
|
||||||
if block == nil {
|
|
||||||
return nil, ErrInvalidPEMBlock
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := unmarshalCertificateBlock(block)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if len(caPEMs) == 0 || strings.TrimSpace(string(caPEMs)) == "" {
|
||||||
err = pool.AddCA(c)
|
break
|
||||||
if errors.Is(err, ErrExpired) {
|
|
||||||
expired = true
|
|
||||||
continue
|
|
||||||
} else if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := scanner.Err(); err != nil {
|
|
||||||
return nil, ErrInvalidPEMBlock
|
|
||||||
}
|
|
||||||
|
|
||||||
if expired {
|
if expired {
|
||||||
return pool, ErrExpired
|
return pool, ErrExpired
|
||||||
|
|||||||
@@ -1,10 +1,7 @@
|
|||||||
package cert
|
package cert
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -115,60 +112,6 @@ k+coOv04r+zh33ISyhbsafnYduN17p2eD7CmHvHuerguXD9f32gcxo/KsFCKEjMe
|
|||||||
assert.Len(t, ppppp.CAs, 1)
|
assert.Len(t, ppppp.CAs, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
// oneByteReader wraps a reader to return at most 1 byte per Read call,
|
|
||||||
// exercising the streaming accumulation logic in NewCAPoolFromPEMReader.
|
|
||||||
type oneByteReader struct {
|
|
||||||
r io.Reader
|
|
||||||
}
|
|
||||||
|
|
||||||
func (o *oneByteReader) Read(p []byte) (int, error) {
|
|
||||||
if len(p) == 0 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
return o.r.Read(p[:1])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewCAPoolFromPEMReader_EmptyReader(t *testing.T) {
|
|
||||||
pool, err := NewCAPoolFromPEMReader(bytes.NewReader(nil))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, pool.CAs)
|
|
||||||
|
|
||||||
pool, err = NewCAPoolFromPEMReader(strings.NewReader(" \n\t\n "))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, pool.CAs)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewCAPoolFromPEMReader_OneByteReads(t *testing.T) {
|
|
||||||
ca1, _, _, pem1 := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(time.Hour), nil, nil, nil)
|
|
||||||
ca2, _, _, pem2 := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(time.Hour), nil, nil, nil)
|
|
||||||
|
|
||||||
bundle := append(pem1, pem2...)
|
|
||||||
pool, err := NewCAPoolFromPEMReader(&oneByteReader{r: bytes.NewReader(bundle)})
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Len(t, pool.CAs, 2)
|
|
||||||
|
|
||||||
fp1, err := ca1.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
fp2, err := ca2.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Contains(t, pool.CAs, fp1)
|
|
||||||
assert.Contains(t, pool.CAs, fp2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewCAPoolFromPEMReader_TruncatedPEM(t *testing.T) {
|
|
||||||
_, err := NewCAPoolFromPEMReader(strings.NewReader("-----BEGIN NEBULA CERTIFICATE-----\npartialdata"))
|
|
||||||
assert.ErrorIs(t, err, ErrInvalidPEMBlock)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewCAPoolFromPEMReader_TrailingGarbage(t *testing.T) {
|
|
||||||
_, _, _, pem1 := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(time.Hour), nil, nil, nil)
|
|
||||||
|
|
||||||
bundle := append(pem1, []byte("some trailing garbage")...)
|
|
||||||
_, err := NewCAPoolFromPEMReader(bytes.NewReader(bundle))
|
|
||||||
assert.ErrorIs(t, err, ErrInvalidPEMBlock)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Verify(t *testing.T) {
|
func TestCertificateV1_Verify(t *testing.T) {
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
||||||
c, _, _, _ := NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test cert", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
c, _, _, _ := NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test cert", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
||||||
|
|||||||
+1
-6
@@ -44,12 +44,7 @@ func swap(r, s []byte) ([]byte, []byte, error) {
|
|||||||
}
|
}
|
||||||
sNormalized := nMod.Nat().Sub(bigS, nMod)
|
sNormalized := nMod.Nat().Sub(bigS, nMod)
|
||||||
|
|
||||||
result := sNormalized.Bytes(nMod)
|
return r, sNormalized.Bytes(nMod), nil
|
||||||
for len(result) > 1 && result[0] == 0 {
|
|
||||||
result = result[1:]
|
|
||||||
}
|
|
||||||
|
|
||||||
return r, result, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func Normalize(sig []byte) ([]byte, error) {
|
func Normalize(sig []byte) ([]byte, error) {
|
||||||
|
|||||||
+13
-69
@@ -1,66 +1,12 @@
|
|||||||
package cert
|
package cert
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"encoding/pem"
|
"encoding/pem"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
var ErrTruncatedPEMBlock = errors.New("truncated PEM block")
|
|
||||||
|
|
||||||
// SplitPEM is a split function for bufio.Scanner that returns each PEM block.
|
|
||||||
func SplitPEM(data []byte, atEOF bool) (advance int, token []byte, err error) {
|
|
||||||
// Look for the start of a PEM block
|
|
||||||
start := bytes.Index(data, []byte("-----BEGIN "))
|
|
||||||
if start == -1 {
|
|
||||||
if atEOF && len(bytes.TrimSpace(data)) > 0 {
|
|
||||||
// Non-whitespace content with no PEM block
|
|
||||||
return 0, nil, ErrTruncatedPEMBlock
|
|
||||||
}
|
|
||||||
if atEOF {
|
|
||||||
return len(data), nil, nil
|
|
||||||
}
|
|
||||||
// Request more data
|
|
||||||
return 0, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Look for the end marker
|
|
||||||
endMarkerStart := bytes.Index(data[start:], []byte("-----END "))
|
|
||||||
if endMarkerStart == -1 {
|
|
||||||
if atEOF {
|
|
||||||
// Incomplete PEM block at EOF
|
|
||||||
return 0, nil, ErrTruncatedPEMBlock
|
|
||||||
}
|
|
||||||
// Need more data to find the end
|
|
||||||
return 0, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Find the actual end of the END line (after the newline)
|
|
||||||
endMarkerStart += start
|
|
||||||
endLineEnd := bytes.IndexByte(data[endMarkerStart:], '\n')
|
|
||||||
var end int
|
|
||||||
if endLineEnd == -1 {
|
|
||||||
if atEOF {
|
|
||||||
// END marker without newline at EOF - take it anyway
|
|
||||||
end = len(data)
|
|
||||||
} else {
|
|
||||||
// Need more data
|
|
||||||
return 0, nil, nil
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
end = endMarkerStart + endLineEnd + 1
|
|
||||||
}
|
|
||||||
|
|
||||||
// Extract the PEM block
|
|
||||||
pemBlock := data[start:end]
|
|
||||||
|
|
||||||
// Return the valid PEM block
|
|
||||||
return end, pemBlock, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
const ( //cert banners
|
const ( //cert banners
|
||||||
CertificateBanner = "NEBULA CERTIFICATE"
|
CertificateBanner = "NEBULA CERTIFICATE"
|
||||||
CertificateV2Banner = "NEBULA CERTIFICATE V2"
|
CertificateV2Banner = "NEBULA CERTIFICATE V2"
|
||||||
@@ -91,7 +37,19 @@ func UnmarshalCertificateFromPEM(b []byte) (Certificate, []byte, error) {
|
|||||||
return nil, r, ErrInvalidPEMBlock
|
return nil, r, ErrInvalidPEMBlock
|
||||||
}
|
}
|
||||||
|
|
||||||
c, err := unmarshalCertificateBlock(p)
|
var c Certificate
|
||||||
|
var err error
|
||||||
|
|
||||||
|
switch p.Type {
|
||||||
|
// Implementations must validate the resulting certificate contains valid information
|
||||||
|
case CertificateBanner:
|
||||||
|
c, err = unmarshalCertificateV1(p.Bytes, nil)
|
||||||
|
case CertificateV2Banner:
|
||||||
|
c, err = unmarshalCertificateV2(p.Bytes, nil, Curve_CURVE25519)
|
||||||
|
default:
|
||||||
|
return nil, r, ErrInvalidPEMCertificateBanner
|
||||||
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, r, err
|
return nil, r, err
|
||||||
}
|
}
|
||||||
@@ -100,20 +58,6 @@ func UnmarshalCertificateFromPEM(b []byte) (Certificate, []byte, error) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// unmarshalCertificateBlock decodes a single PEM block into a certificate.
|
|
||||||
// It expects a Nebula certificate banner and returns ErrInvalidPEMCertificateBanner otherwise.
|
|
||||||
func unmarshalCertificateBlock(block *pem.Block) (Certificate, error) {
|
|
||||||
switch block.Type {
|
|
||||||
// Implementations must validate the resulting certificate contains valid information
|
|
||||||
case CertificateBanner:
|
|
||||||
return unmarshalCertificateV1(block.Bytes, nil)
|
|
||||||
case CertificateV2Banner:
|
|
||||||
return unmarshalCertificateV2(block.Bytes, nil, Curve_CURVE25519)
|
|
||||||
default:
|
|
||||||
return nil, ErrInvalidPEMCertificateBanner
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func marshalCertPublicKeyToPEM(c Certificate) []byte {
|
func marshalCertPublicKeyToPEM(c Certificate) []byte {
|
||||||
if c.IsCA() {
|
if c.IsCA() {
|
||||||
return MarshalSigningPublicKeyToPEM(c.Curve(), c.PublicKey())
|
return MarshalSigningPublicKeyToPEM(c.Curve(), c.PublicKey())
|
||||||
|
|||||||
@@ -1,88 +1,12 @@
|
|||||||
package cert
|
package cert
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
func scanAll(t *testing.T, input string) ([]string, error) {
|
|
||||||
t.Helper()
|
|
||||||
scanner := bufio.NewScanner(strings.NewReader(input))
|
|
||||||
scanner.Split(SplitPEM)
|
|
||||||
var blocks []string
|
|
||||||
for scanner.Scan() {
|
|
||||||
blocks = append(blocks, scanner.Text())
|
|
||||||
}
|
|
||||||
return blocks, scanner.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_Single(t *testing.T) {
|
|
||||||
input := "-----BEGIN TEST-----\ndata\n-----END TEST-----\n"
|
|
||||||
blocks, err := scanAll(t, input)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, blocks, 1)
|
|
||||||
require.Equal(t, input, blocks[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_Multiple(t *testing.T) {
|
|
||||||
block1 := "-----BEGIN TEST-----\naaa\n-----END TEST-----\n"
|
|
||||||
block2 := "-----BEGIN TEST-----\nbbb\n-----END TEST-----\n"
|
|
||||||
blocks, err := scanAll(t, block1+block2)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, blocks, 2)
|
|
||||||
require.Equal(t, block1, blocks[0])
|
|
||||||
require.Equal(t, block2, blocks[1])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_CommentsAndWhitespaceBetweenBlocks(t *testing.T) {
|
|
||||||
input := "# comment\n\n-----BEGIN TEST-----\naaa\n-----END TEST-----\n\n# another comment\n\n-----BEGIN TEST-----\nbbb\n-----END TEST-----\n"
|
|
||||||
blocks, err := scanAll(t, input)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, blocks, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_Empty(t *testing.T) {
|
|
||||||
blocks, err := scanAll(t, "")
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Empty(t, blocks)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_WhitespaceOnly(t *testing.T) {
|
|
||||||
blocks, err := scanAll(t, " \n\t\n ")
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Empty(t, blocks)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_TrailingGarbage(t *testing.T) {
|
|
||||||
input := "-----BEGIN TEST-----\ndata\n-----END TEST-----\ngarbage"
|
|
||||||
blocks, err := scanAll(t, input)
|
|
||||||
require.ErrorIs(t, err, ErrTruncatedPEMBlock)
|
|
||||||
require.Len(t, blocks, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_TruncatedBlock(t *testing.T) {
|
|
||||||
input := "-----BEGIN TEST-----\npartial data with no end"
|
|
||||||
_, err := scanAll(t, input)
|
|
||||||
require.ErrorIs(t, err, ErrTruncatedPEMBlock)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_NoEndNewline(t *testing.T) {
|
|
||||||
input := "-----BEGIN TEST-----\ndata\n-----END TEST-----"
|
|
||||||
blocks, err := scanAll(t, input)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, blocks, 1)
|
|
||||||
require.Equal(t, input, blocks[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSplitPEM_GarbageOnly(t *testing.T) {
|
|
||||||
_, err := scanAll(t, "this is not PEM data")
|
|
||||||
require.ErrorIs(t, err, ErrTruncatedPEMBlock)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalCertificateFromPEM(t *testing.T) {
|
func TestUnmarshalCertificateFromPEM(t *testing.T) {
|
||||||
goodCert := []byte(`
|
goodCert := []byte(`
|
||||||
# A good cert
|
# A good cert
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
@@ -39,15 +40,21 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
caFile, err := os.Open(*vf.caPath)
|
rawCACert, err := os.ReadFile(*vf.caPath)
|
||||||
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 caFile.Close()
|
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caFile)
|
caPool := cert.NewCAPool()
|
||||||
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
for {
|
||||||
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
rawCACert, err = caPool.AddCAFromPEM(rawCACert)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if rawCACert == nil || len(rawCACert) == 0 || strings.TrimSpace(string(rawCACert)) == "" {
|
||||||
|
break
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := os.ReadFile(*vf.certPath)
|
rawCert, err := os.ReadFile(*vf.certPath)
|
||||||
|
|||||||
@@ -64,7 +64,7 @@ func Test_verify(t *testing.T) {
|
|||||||
err = verify([]string{"-ca", caFile.Name(), "-crt", "does_not_exist"}, ob, eb)
|
err = verify([]string{"-ca", caFile.Name(), "-crt", "does_not_exist"}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
require.ErrorIs(t, err, cert.ErrInvalidPEMBlock)
|
require.EqualError(t, err, "error while adding ca cert to pool: input did not contain a valid PEM encoded block")
|
||||||
|
|
||||||
// make a ca for later
|
// make a ca for later
|
||||||
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
|
|||||||
@@ -78,20 +78,8 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
ctrl.Start()
|
||||||
if err != nil {
|
ctrl.ShutdownBlock()
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
go ctrl.ShutdownBlock()
|
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
|
||||||
l.WithError(err).Error("Nebula stopped due to fatal error")
|
|
||||||
os.Exit(2)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.Info("Goodbye")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
|
|||||||
+2
-14
@@ -72,21 +72,9 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !*configTest {
|
if !*configTest {
|
||||||
wait, err := ctrl.Start()
|
ctrl.Start()
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
go ctrl.ShutdownBlock()
|
|
||||||
notifyReady(l)
|
notifyReady(l)
|
||||||
|
ctrl.ShutdownBlock()
|
||||||
if err := wait(); err != nil {
|
|
||||||
l.WithError(err).Error("Nebula stopped due to fatal error")
|
|
||||||
os.Exit(2)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.Info("Goodbye")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import (
|
|||||||
"github.com/flynn/noise"
|
"github.com/flynn/noise"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/overlay"
|
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
@@ -53,7 +52,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlay.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -136,7 +135,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlay.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -221,7 +220,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlay.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
@@ -348,7 +347,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlay.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
|
|||||||
+6
-72
@@ -2,11 +2,9 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"sync"
|
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
@@ -15,20 +13,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
)
|
)
|
||||||
|
|
||||||
type RunState int
|
|
||||||
|
|
||||||
const (
|
|
||||||
StateUnknown RunState = iota
|
|
||||||
StateReady
|
|
||||||
StateStarted
|
|
||||||
StateStopping
|
|
||||||
StateStopped
|
|
||||||
)
|
|
||||||
|
|
||||||
var ErrAlreadyStarted = errors.New("nebula is already started")
|
|
||||||
var ErrAlreadyStopped = errors.New("nebula cannot be restarted")
|
|
||||||
var ErrUnknownState = errors.New("nebula state is invalid")
|
|
||||||
|
|
||||||
// Every interaction here needs to take extra care to copy memory and not return or use arguments "as is" when touching
|
// Every interaction here needs to take extra care to copy memory and not return or use arguments "as is" when touching
|
||||||
// core. This means copying IP objects, slices, de-referencing pointers and taking the actual value, etc
|
// core. This means copying IP objects, slices, de-referencing pointers and taking the actual value, etc
|
||||||
|
|
||||||
@@ -42,9 +26,6 @@ type controlHostLister interface {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Control struct {
|
type Control struct {
|
||||||
stateLock sync.Mutex
|
|
||||||
state RunState
|
|
||||||
|
|
||||||
f *Interface
|
f *Interface
|
||||||
l *logrus.Logger
|
l *logrus.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
@@ -68,31 +49,10 @@ type ControlHostInfo struct {
|
|||||||
CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"`
|
CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Start actually runs nebula, this is a nonblocking call.
|
// Start actually runs nebula, this is a nonblocking call. To block use Control.ShutdownBlock()
|
||||||
// The returned function blocks until nebula has fully stopped and returns the
|
func (c *Control) Start() {
|
||||||
// first fatal reader error (if any). A nil error means nebula shut down
|
|
||||||
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
|
||||||
// triggered the shutdown.
|
|
||||||
func (c *Control) Start() (func() error, error) {
|
|
||||||
c.stateLock.Lock()
|
|
||||||
defer c.stateLock.Unlock()
|
|
||||||
switch c.state {
|
|
||||||
case StateReady:
|
|
||||||
//yay!
|
|
||||||
case StateStopped, StateStopping:
|
|
||||||
return nil, ErrAlreadyStopped
|
|
||||||
case StateStarted:
|
|
||||||
return nil, ErrAlreadyStarted
|
|
||||||
default:
|
|
||||||
return nil, ErrUnknownState
|
|
||||||
}
|
|
||||||
|
|
||||||
// Activate the interface
|
// Activate the interface
|
||||||
err := c.f.activate()
|
c.f.activate()
|
||||||
if err != nil {
|
|
||||||
c.state = StateStopped
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Call all the delayed funcs that waited patiently for the interface to be created.
|
// Call all the delayed funcs that waited patiently for the interface to be created.
|
||||||
if c.sshStart != nil {
|
if c.sshStart != nil {
|
||||||
@@ -111,40 +71,16 @@ func (c *Control) Start() (func() error, error) {
|
|||||||
c.lighthouseStart()
|
c.lighthouseStart()
|
||||||
}
|
}
|
||||||
|
|
||||||
c.f.triggerShutdown = c.Stop
|
|
||||||
|
|
||||||
// Start reading packets.
|
// Start reading packets.
|
||||||
out, err := c.f.run()
|
c.f.run()
|
||||||
if err != nil {
|
|
||||||
c.state = StateStopped
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
c.state = StateStarted
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Control) State() RunState {
|
|
||||||
c.stateLock.Lock()
|
|
||||||
defer c.stateLock.Unlock()
|
|
||||||
return c.state
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) Context() context.Context {
|
func (c *Control) Context() context.Context {
|
||||||
return c.ctx
|
return c.ctx
|
||||||
}
|
}
|
||||||
|
|
||||||
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
// Stop signals nebula to shutdown and close all tunnels, returns after the shutdown is complete
|
||||||
func (c *Control) Stop() {
|
func (c *Control) Stop() {
|
||||||
c.stateLock.Lock()
|
|
||||||
if c.state != StateStarted {
|
|
||||||
c.stateLock.Unlock()
|
|
||||||
// We are stopping or stopped already
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
c.state = StateStopping
|
|
||||||
c.stateLock.Unlock()
|
|
||||||
|
|
||||||
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
||||||
// being created while we're shutting them all down.
|
// being created while we're shutting them all down.
|
||||||
c.cancel()
|
c.cancel()
|
||||||
@@ -153,9 +89,7 @@ func (c *Control) Stop() {
|
|||||||
if err := c.f.Close(); err != nil {
|
if err := c.f.Close(); err != nil {
|
||||||
c.l.WithError(err).Error("Close interface failed")
|
c.l.WithError(err).Error("Close interface failed")
|
||||||
}
|
}
|
||||||
c.stateLock.Lock()
|
c.l.Info("Goodbye")
|
||||||
c.state = StateStopped
|
|
||||||
c.stateLock.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
||||||
|
|||||||
@@ -79,7 +79,6 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
}, &Interface{})
|
}, &Interface{})
|
||||||
|
|
||||||
c := Control{
|
c := Control{
|
||||||
state: StateReady,
|
|
||||||
f: &Interface{
|
f: &Interface{
|
||||||
hostMap: hm,
|
hostMap: hm,
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -1,565 +0,0 @@
|
|||||||
//go:build e2e_testing
|
|
||||||
// +build e2e_testing
|
|
||||||
|
|
||||||
package e2e
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// makeHandshakePacket creates a handshake packet with the given parameters.
|
|
||||||
func makeHandshakePacket(from, to netip.AddrPort, subtype header.MessageSubType, remoteIndex uint32, counter uint64) *udp.Packet {
|
|
||||||
data := make([]byte, 200)
|
|
||||||
header.Encode(data, header.Version, header.Handshake, subtype, remoteIndex, counter)
|
|
||||||
for i := header.Len; i < len(data); i++ {
|
|
||||||
data[i] = byte(i)
|
|
||||||
}
|
|
||||||
return &udp.Packet{To: to, From: from, Data: data}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|
||||||
// Verify the responder correctly handles receiving the same msg1 multiple times
|
|
||||||
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
|
||||||
// and the cached response is resent.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake from me to them")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
|
|
||||||
t.Log("Grab my msg1")
|
|
||||||
msg1 := myControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Inject msg1 into them, first time")
|
|
||||||
theirControl.InjectUDPPacket(msg1)
|
|
||||||
_ = theirControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Inject the SAME msg1 again, tests ErrAlreadySeen path")
|
|
||||||
theirControl.InjectUDPPacket(msg1)
|
|
||||||
resp2 := theirControl.GetFromUDP(true)
|
|
||||||
assert.NotNil(t, resp2, "should get cached response on duplicate msg1")
|
|
||||||
|
|
||||||
t.Log("Complete handshake with cached response")
|
|
||||||
myControl.InjectUDPPacket(resp2)
|
|
||||||
myControl.WaitForType(1, 0, theirControl)
|
|
||||||
|
|
||||||
t.Log("Drain cached packet and verify tunnel works")
|
|
||||||
cachedPacket := theirControl.GetFromTun(true)
|
|
||||||
assertUdpPacket(t, []byte("Hi"), cachedPacket, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
t.Log("Verify only one tunnel exists on each side")
|
|
||||||
assert.Len(t, myControl.ListHostmapHosts(false), 1)
|
|
||||||
assert.Len(t, theirControl.ListHostmapHosts(false), 1)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|
||||||
// Verify that a truncated handshake packet is ignored and the real
|
|
||||||
// packet can still complete the handshake.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
|
|
||||||
t.Log("Get msg1 and deliver to responder")
|
|
||||||
msg1 := myControl.GetFromUDP(true)
|
|
||||||
theirControl.InjectUDPPacket(msg1)
|
|
||||||
|
|
||||||
t.Log("Get the real response")
|
|
||||||
realResp := theirControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Truncate the response and inject, should be ignored")
|
|
||||||
truncResp := realResp.Copy()
|
|
||||||
truncResp.Data = truncResp.Data[:header.Len]
|
|
||||||
myControl.InjectUDPPacket(truncResp)
|
|
||||||
|
|
||||||
t.Log("Verify pending handshake survived the truncated packet")
|
|
||||||
assert.NotEmpty(t, myControl.ListHostmapHosts(true), "pending handshake should still exist")
|
|
||||||
|
|
||||||
t.Log("Inject real response, should complete handshake")
|
|
||||||
myControl.InjectUDPPacket(realResp)
|
|
||||||
myControl.WaitForType(1, 0, theirControl)
|
|
||||||
|
|
||||||
t.Log("Drain and verify tunnel")
|
|
||||||
cachedPacket := theirControl.GetFromTun(true)
|
|
||||||
assertUdpPacket(t, []byte("Hi"), cachedPacket, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|
||||||
// A msg2 arriving with no matching pending index should be silently dropped
|
|
||||||
// with no response sent and no state changes.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
t.Log("Complete a normal handshake")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
t.Log("Record hostmap state")
|
|
||||||
myIndexes := len(myControl.ListHostmapIndexes(false))
|
|
||||||
|
|
||||||
t.Log("Inject a fake msg2 with unknown RemoteIndex")
|
|
||||||
myControl.InjectUDPPacket(makeHandshakePacket(theirUdpAddr, myUdpAddr, header.HandshakeIXPSK0, 0xDEADBEEF, 2))
|
|
||||||
|
|
||||||
t.Log("Verify no new indexes created")
|
|
||||||
assert.Equal(t, myIndexes, len(myControl.ListHostmapIndexes(false)))
|
|
||||||
|
|
||||||
t.Log("Verify no UDP response was sent")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
assert.Nil(t, myControl.GetFromUDP(false), "should not send a response to orphaned msg2")
|
|
||||||
|
|
||||||
t.Log("Verify existing tunnel still works")
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
|
||||||
// A handshake packet with an unexpected message counter should be silently
|
|
||||||
// dropped with no side effects and no UDP response.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, _, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
t.Log("Inject handshake with MessageCounter=3")
|
|
||||||
myControl.InjectUDPPacket(makeHandshakePacket(theirUdpAddr, myUdpAddr, header.HandshakeIXPSK0, 0, 3))
|
|
||||||
|
|
||||||
t.Log("Inject handshake with MessageCounter=99")
|
|
||||||
myControl.InjectUDPPacket(makeHandshakePacket(theirUdpAddr, myUdpAddr, header.HandshakeIXPSK0, 0, 99))
|
|
||||||
|
|
||||||
t.Log("Verify no tunnels or pending handshakes")
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(false))
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
|
|
||||||
t.Log("Verify no UDP response was sent")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
assert.Nil(t, myControl.GetFromUDP(false))
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeUnknownSubtype(t *testing.T) {
|
|
||||||
// A handshake packet with an unknown subtype should be silently dropped.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, _, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, _, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
t.Log("Inject handshake with unknown subtype 99")
|
|
||||||
myControl.InjectUDPPacket(makeHandshakePacket(theirUdpAddr, myUdpAddr, header.MessageSubType(99), 0, 1))
|
|
||||||
|
|
||||||
t.Log("Verify no tunnels or pending handshakes")
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(false))
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
|
|
||||||
t.Log("Verify no UDP response was sent")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
assert.Nil(t, myControl.GetFromUDP(false))
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeLateResponse(t *testing.T) {
|
|
||||||
// After a handshake times out, a late response should be silently ignored
|
|
||||||
// with no new tunnels created.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{
|
|
||||||
"handshakes": m{
|
|
||||||
"try_interval": "200ms",
|
|
||||||
"retries": 2,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
|
|
||||||
t.Log("Grab msg1 but don't deliver")
|
|
||||||
msg1 := myControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Wait for handshake to time out")
|
|
||||||
for i := 0; i < 5; i++ {
|
|
||||||
time.Sleep(300 * time.Millisecond)
|
|
||||||
myControl.GetFromUDP(false)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Log("Confirm no pending handshakes remain")
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
|
|
||||||
t.Log("Deliver old msg1 to them, they create a tunnel")
|
|
||||||
theirControl.InjectUDPPacket(msg1)
|
|
||||||
resp := theirControl.GetFromUDP(true)
|
|
||||||
assert.NotNil(t, resp)
|
|
||||||
|
|
||||||
t.Log("Inject late response into me, should be ignored")
|
|
||||||
myControl.InjectUDPPacket(resp)
|
|
||||||
|
|
||||||
t.Log("No tunnel should exist on my side")
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(false))
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|
||||||
// Verify that a node rejects a handshake containing its own VPN IP in the
|
|
||||||
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
|
|
||||||
// Need a lighthouse entry to trigger a handshake
|
|
||||||
myControl.InjectLightHouseAddr(netip.MustParseAddr("10.128.0.2"), netip.MustParseAddrPort("10.0.0.2:4242"))
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
|
||||||
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
msg1 := myControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Drain any handshake retransmits before injecting")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
for myControl.GetFromUDP(false) != nil {
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Log("Feed my own msg1 back to me as if it came from someone else")
|
|
||||||
selfMsg := msg1.Copy()
|
|
||||||
selfMsg.From = netip.MustParseAddrPort("10.0.0.99:4242")
|
|
||||||
selfMsg.To = myUdpAddr
|
|
||||||
myControl.InjectUDPPacket(selfMsg)
|
|
||||||
|
|
||||||
t.Log("Verify no response was sent (self-connection rejected)")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
// Drain any further retransmits from the original handshake, then check
|
|
||||||
// that none of them are a handshake response (MessageCounter=2)
|
|
||||||
h := &header.H{}
|
|
||||||
for {
|
|
||||||
p := myControl.GetFromUDP(false)
|
|
||||||
if p == nil {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
_ = h.Parse(p.Data)
|
|
||||||
assert.NotEqual(t, uint64(2), h.MessageCounter,
|
|
||||||
"should not send a stage 2 response to self-connection")
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Log("Verify no tunnel to myself was created")
|
|
||||||
assert.Nil(t, myControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false))
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
|
||||||
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, _, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
_, _, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
|
|
||||||
t.Log("Inject handshake with MessageCounter=0")
|
|
||||||
myControl.InjectUDPPacket(makeHandshakePacket(theirUdpAddr, myUdpAddr, header.HandshakeIXPSK0, 0, 0))
|
|
||||||
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(false))
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
assert.Nil(t, myControl.GetFromUDP(false))
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeRemoteAllowList(t *testing.T) {
|
|
||||||
// Verify that a handshake from a blocked underlay IP is dropped with no
|
|
||||||
// response and no state changes. Then verify the same packet from an
|
|
||||||
// allowed IP succeeds.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{
|
|
||||||
"lighthouse": m{
|
|
||||||
"remote_allow_list": m{
|
|
||||||
"10.0.0.0/8": true,
|
|
||||||
"0.0.0.0/0": false,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake from them")
|
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
msg1 := theirControl.GetFromUDP(true)
|
|
||||||
|
|
||||||
t.Log("Rewrite the source to a blocked IP and inject")
|
|
||||||
blockedMsg := msg1.Copy()
|
|
||||||
blockedMsg.From = netip.MustParseAddrPort("192.168.1.1:4242")
|
|
||||||
myControl.InjectUDPPacket(blockedMsg)
|
|
||||||
|
|
||||||
t.Log("Verify no tunnel, no pending, no response from blocked source")
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(false))
|
|
||||||
assert.Empty(t, myControl.ListHostmapHosts(true))
|
|
||||||
assert.Nil(t, myControl.GetFromUDP(false), "should not respond to blocked source")
|
|
||||||
|
|
||||||
t.Log("Now inject the real packet from the allowed source")
|
|
||||||
myControl.InjectUDPPacket(msg1)
|
|
||||||
|
|
||||||
t.Log("Verify handshake completes from allowed source")
|
|
||||||
resp := myControl.GetFromUDP(true)
|
|
||||||
assert.NotNil(t, resp)
|
|
||||||
theirControl.InjectUDPPacket(resp)
|
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
|
||||||
|
|
||||||
t.Log("Drain cached packet and verify tunnel works")
|
|
||||||
cachedPacket := myControl.GetFromTun(true)
|
|
||||||
assertUdpPacket(t, []byte("Hi"), cachedPacket, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|
||||||
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
|
||||||
// remains functional and hostmap index count is stable.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
t.Log("Complete a normal handshake via the router")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
t.Log("Record hostmap state")
|
|
||||||
theirIndexes := len(theirControl.ListHostmapIndexes(false))
|
|
||||||
hi := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
|
||||||
assert.NotNil(t, hi)
|
|
||||||
originalRemote := hi.CurrentRemote
|
|
||||||
|
|
||||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
|
|
||||||
t.Log("Verify tunnel still works")
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
t.Log("Verify remote is still valid and index count is stable")
|
|
||||||
hi2 := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
|
||||||
assert.NotNil(t, hi2)
|
|
||||||
assert.Equal(t, originalRemote, hi2.CurrentRemote)
|
|
||||||
assert.Equal(t, theirIndexes, len(theirControl.ListHostmapIndexes(false)),
|
|
||||||
"no extra indexes should be created from ErrAlreadySeen")
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|
||||||
// Verify that when the wrong host responds, the cached packets are
|
|
||||||
// transferred to the new handshake, the evil tunnel is closed, evil's
|
|
||||||
// address is blocked, and the correct tunnel is eventually established.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
|
||||||
evilControl, evilVpnIpNet, evilUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "evil", "10.128.0.2/24", nil)
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), evilUdpAddr)
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, theirControl, evilControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
evilControl.Start()
|
|
||||||
|
|
||||||
t.Log("Send multiple packets to them (cached during handshake)")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
|
||||||
|
|
||||||
t.Log("Route until evil tunnel is closed")
|
|
||||||
h := &header.H{}
|
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
|
||||||
if err := h.Parse(p.Data); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
if h.Type == header.CloseTunnel && p.To == evilUdpAddr {
|
|
||||||
return router.RouteAndExit
|
|
||||||
}
|
|
||||||
return router.KeepRouting
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Log("Verify evil's address is blocked in the new pending handshake")
|
|
||||||
pendingHI := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), true)
|
|
||||||
if pendingHI != nil {
|
|
||||||
assert.NotContains(t, pendingHI.RemoteAddrs, evilUdpAddr,
|
|
||||||
"evil's address should be blocked")
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Log("Inject correct lighthouse addr for them")
|
|
||||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
t.Log("Route until cached packets arrive at the real them")
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
assert.NotNil(t, p, "a cached packet should be delivered to the correct host")
|
|
||||||
|
|
||||||
t.Log("Verify the correct host has a tunnel")
|
|
||||||
assertHostInfoPair(t, myUdpAddr, theirUdpAddr, myVpnIpNet, theirVpnIpNet, myControl, theirControl)
|
|
||||||
|
|
||||||
t.Log("Verify no hostinfo artifacts from evil remain")
|
|
||||||
assert.Nil(t, myControl.GetHostInfoByVpnAddr(evilVpnIpNet[0].Addr(), true),
|
|
||||||
"no pending hostinfo for evil")
|
|
||||||
assert.Nil(t, myControl.GetHostInfoByVpnAddr(evilVpnIpNet[0].Addr(), false),
|
|
||||||
"no main hostinfo for evil")
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
evilControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandshakeRelayComplete(t *testing.T) {
|
|
||||||
// Verify that a relay handshake completes correctly and relay state is
|
|
||||||
// properly maintained on all three nodes.
|
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
|
||||||
|
|
||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
|
||||||
|
|
||||||
r := router.NewR(t, myControl, relayControl, theirControl)
|
|
||||||
defer r.RenderFlow()
|
|
||||||
|
|
||||||
myControl.Start()
|
|
||||||
relayControl.Start()
|
|
||||||
theirControl.Start()
|
|
||||||
|
|
||||||
t.Log("Trigger handshake via relay")
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
|
||||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
|
||||||
|
|
||||||
t.Log("Verify bidirectional tunnel via relay")
|
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
|
||||||
|
|
||||||
t.Log("Verify relay state on my side shows relay-to-me")
|
|
||||||
myHI := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
|
|
||||||
assert.NotNil(t, myHI)
|
|
||||||
assert.NotEmpty(t, myHI.CurrentRelaysToMe, "should have relay-to-me for them")
|
|
||||||
|
|
||||||
t.Log("Verify relay state on their side shows relay-to-me")
|
|
||||||
theirHI := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
|
||||||
assert.NotNil(t, theirHI)
|
|
||||||
assert.NotEmpty(t, theirHI.CurrentRelaysToMe, "should have relay-to-me for me")
|
|
||||||
|
|
||||||
t.Log("Verify relay node shows through-me relays")
|
|
||||||
relayHI := relayControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false)
|
|
||||||
assert.NotNil(t, relayHI)
|
|
||||||
|
|
||||||
myControl.Stop()
|
|
||||||
relayControl.Stop()
|
|
||||||
theirControl.Stop()
|
|
||||||
}
|
|
||||||
|
|
||||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
|
||||||
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
|
||||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
|
||||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
|
||||||
// address is IPv6, the relay is skipped.
|
|
||||||
|
|
||||||
// NOTE: Relay reestablishment (Disestablished state transition) is covered
|
|
||||||
// by the existing TestReestablishRelays in handshakes_test.go.
|
|
||||||
+8
-6
@@ -204,12 +204,6 @@ punchy:
|
|||||||
# Trusted SSH CA public keys. These are the public keys of the CAs that are allowed to sign SSH keys for access.
|
# Trusted SSH CA public keys. These are the public keys of the CAs that are allowed to sign SSH keys for access.
|
||||||
#trusted_cas:
|
#trusted_cas:
|
||||||
#- "ssh public key string"
|
#- "ssh public key string"
|
||||||
# sandbox_dir restricts file paths for profiling commands (start-cpu-profile, save-heap-profile,
|
|
||||||
# save-mutex-profile) to the specified directory. Relative paths will be resolved within this directory,
|
|
||||||
# and absolute paths outside of it will be rejected. Default is $TMP/nebula-debug.
|
|
||||||
# The directory is NOT automatically created.
|
|
||||||
# Overriding this to "" is the same as "/" and will allow overwriting any path on the host.
|
|
||||||
#sandbox_dir: /var/tmp/nebula-debug
|
|
||||||
|
|
||||||
# EXPERIMENTAL: relay support for networks that can't establish direct connections.
|
# EXPERIMENTAL: relay support for networks that can't establish direct connections.
|
||||||
relay:
|
relay:
|
||||||
@@ -348,6 +342,14 @@ logging:
|
|||||||
# after receiving the response for lighthouse queries
|
# after receiving the response for lighthouse queries
|
||||||
#trigger_buffer: 64
|
#trigger_buffer: 64
|
||||||
|
|
||||||
|
# max_rate limits the number of new inbound handshakes per second. Once the limit is reached,
|
||||||
|
# new handshakes are dropped until the next second. A value of 0 means unlimited (default).
|
||||||
|
# This is useful for preventing DoS attacks that attempt to exhaust CPU with handshake crypto.
|
||||||
|
# Running `openssl speed ecdhp256` on your hardware can be a good rule of thumb for choosing
|
||||||
|
# a max, as each handshake performs similar DH operations. Note that this benchmarks a single
|
||||||
|
# core, so you may wish to scale the value by the number of `routines` configured.
|
||||||
|
#max_rate: 0
|
||||||
|
|
||||||
# Tunnel manager settings
|
# Tunnel manager settings
|
||||||
#tunnels:
|
#tunnels:
|
||||||
# drop_inactive controls whether inactive tunnels are maintained or dropped after the inactive_timeout period has
|
# drop_inactive controls whether inactive tunnels are maintained or dropped after the inactive_timeout period has
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
module github.com/slackhq/nebula
|
module github.com/slackhq/nebula
|
||||||
|
|
||||||
go 1.25.0
|
go 1.25
|
||||||
|
|
||||||
require (
|
require (
|
||||||
dario.cat/mergo v1.0.2
|
dario.cat/mergo v1.0.2
|
||||||
@@ -13,8 +13,8 @@ require (
|
|||||||
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
|
||||||
github.com/miekg/dns v1.1.72
|
github.com/miekg/dns v1.1.70
|
||||||
github.com/miekg/pkcs11 v1.1.2
|
github.com/miekg/pkcs11 v1.1.2-0.20231115102856-9078ad6b9d4b
|
||||||
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
||||||
github.com/prometheus/client_golang v1.23.2
|
github.com/prometheus/client_golang v1.23.2
|
||||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
||||||
@@ -24,15 +24,15 @@ require (
|
|||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.50.0
|
golang.org/x/crypto v0.47.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
golang.org/x/net v0.52.0
|
golang.org/x/net v0.49.0
|
||||||
golang.org/x/sync v0.20.0
|
golang.org/x/sync v0.19.0
|
||||||
golang.org/x/sys v0.43.0
|
golang.org/x/sys v0.40.0
|
||||||
golang.org/x/term v0.42.0
|
golang.org/x/term v0.39.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.5.3
|
||||||
google.golang.org/protobuf v1.36.11
|
google.golang.org/protobuf v1.36.11
|
||||||
gopkg.in/yaml.v3 v3.0.1
|
gopkg.in/yaml.v3 v3.0.1
|
||||||
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
|
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
|
||||||
@@ -50,7 +50,7 @@ require (
|
|||||||
github.com/prometheus/procfs v0.16.1 // indirect
|
github.com/prometheus/procfs v0.16.1 // indirect
|
||||||
github.com/vishvananda/netns v0.0.5 // indirect
|
github.com/vishvananda/netns v0.0.5 // indirect
|
||||||
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
||||||
golang.org/x/mod v0.34.0 // indirect
|
golang.org/x/mod v0.31.0 // indirect
|
||||||
golang.org/x/time v0.5.0 // indirect
|
golang.org/x/time v0.5.0 // indirect
|
||||||
golang.org/x/tools v0.43.0 // indirect
|
golang.org/x/tools v0.40.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -85,10 +85,10 @@ github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
|
|||||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||||
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||||
github.com/matttproud/golang_protobuf_extensions v1.0.1/go.mod h1:D8He9yQNgCq6Z5Ld7szi9bcBfOoFv/3dc6xSMkL2PC0=
|
github.com/matttproud/golang_protobuf_extensions v1.0.1/go.mod h1:D8He9yQNgCq6Z5Ld7szi9bcBfOoFv/3dc6xSMkL2PC0=
|
||||||
github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
|
github.com/miekg/dns v1.1.70 h1:DZ4u2AV35VJxdD9Fo9fIWm119BsQL5cZU1cQ9s0LkqA=
|
||||||
github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
|
github.com/miekg/dns v1.1.70/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
|
||||||
github.com/miekg/pkcs11 v1.1.2 h1:/VxmeAX5qU6Q3EwafypogwWbYryHFmF2RpkJmw3m4MQ=
|
github.com/miekg/pkcs11 v1.1.2-0.20231115102856-9078ad6b9d4b h1:J/AzCvg5z0Hn1rqZUJjpbzALUmkKX0Zwbc/i4fw7Sfk=
|
||||||
github.com/miekg/pkcs11 v1.1.2/go.mod h1:XsNlhZGX73bx86s2hdc/FuaLm2CPZJemRLMA+WTFxgs=
|
github.com/miekg/pkcs11 v1.1.2-0.20231115102856-9078ad6b9d4b/go.mod h1:XsNlhZGX73bx86s2hdc/FuaLm2CPZJemRLMA+WTFxgs=
|
||||||
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
||||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
||||||
github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0=
|
github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0=
|
||||||
@@ -164,16 +164,16 @@ 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.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
golang.org/x/crypto v0.47.0 h1:V6e3FRj+n4dbpw86FJ8Fv7XVOql7TEwpHapKoMJ/GO8=
|
||||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
golang.org/x/crypto v0.47.0/go.mod h1:ff3Y9VzzKbwSSEzWqJsJVBnWmRwRSHt/6Op5n9bQc4A=
|
||||||
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=
|
||||||
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
||||||
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||||
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
||||||
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
||||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
@@ -184,8 +184,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.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
|
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
|
||||||
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
|
||||||
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=
|
||||||
@@ -193,8 +193,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.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
||||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||||
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=
|
||||||
@@ -210,11 +210,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.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ=
|
||||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||||
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.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
golang.org/x/term v0.39.0 h1:RclSuaJf32jOqZz74CkPA9qFuVTX7vhLlpfj/IGWlqY=
|
||||||
golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
|
golang.org/x/term v0.39.0/go.mod h1:yxzUCTP/U+FzoxfdKmLaA0RV1WgE0VY7hXBwKtY/4ww=
|
||||||
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=
|
||||||
@@ -225,8 +225,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn
|
|||||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||||
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
||||||
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
||||||
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
||||||
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
|
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
@@ -235,8 +235,8 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU=
|
golang.zx2c4.com/wireguard/windows v0.5.3 h1:On6j2Rpn3OEMXqBq00QEDC7bWSZrPIHKIus8eIuExIE=
|
||||||
golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM=
|
golang.zx2c4.com/wireguard/windows v0.5.3/go.mod h1:9TEe8TJmtwyQebdFwAkEWOPr3prrtqm+REGFifP60hI=
|
||||||
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
||||||
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
||||||
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
||||||
|
|||||||
+57
-8
@@ -23,22 +23,25 @@ const (
|
|||||||
DefaultHandshakeRetries = 10
|
DefaultHandshakeRetries = 10
|
||||||
DefaultHandshakeTriggerBuffer = 64
|
DefaultHandshakeTriggerBuffer = 64
|
||||||
DefaultUseRelays = true
|
DefaultUseRelays = true
|
||||||
|
DefaultMaxHandshakeRate = 0 // 0 means unlimited
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
defaultHandshakeConfig = HandshakeConfig{
|
defaultHandshakeConfig = HandshakeConfig{
|
||||||
tryInterval: DefaultHandshakeTryInterval,
|
tryInterval: DefaultHandshakeTryInterval,
|
||||||
retries: DefaultHandshakeRetries,
|
retries: DefaultHandshakeRetries,
|
||||||
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
||||||
useRelays: DefaultUseRelays,
|
useRelays: DefaultUseRelays,
|
||||||
|
maxHandshakeRate: DefaultMaxHandshakeRate,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
type HandshakeConfig struct {
|
type HandshakeConfig struct {
|
||||||
tryInterval time.Duration
|
tryInterval time.Duration
|
||||||
retries int64
|
retries int64
|
||||||
triggerBuffer int
|
triggerBuffer int
|
||||||
useRelays bool
|
useRelays bool
|
||||||
|
maxHandshakeRate int
|
||||||
|
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
}
|
}
|
||||||
@@ -58,9 +61,15 @@ type HandshakeManager struct {
|
|||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
metricInitiated metrics.Counter
|
metricInitiated metrics.Counter
|
||||||
metricTimedOut metrics.Counter
|
metricTimedOut metrics.Counter
|
||||||
|
metricRateLimited metrics.Counter
|
||||||
f *Interface
|
f *Interface
|
||||||
l *logrus.Logger
|
l *logrus.Logger
|
||||||
|
|
||||||
|
// Rate limiting for new handshakes (token bucket)
|
||||||
|
rateBucket int // tokens currently available
|
||||||
|
rateMax int // max tokens (== max handshakes per second), 0 means unlimited
|
||||||
|
rateLastTick time.Time
|
||||||
|
|
||||||
// can be used to trigger outbound handshake for the given vpnIp
|
// can be used to trigger outbound handshake for the given vpnIp
|
||||||
trigger chan netip.Addr
|
trigger chan netip.Addr
|
||||||
}
|
}
|
||||||
@@ -116,10 +125,41 @@ func NewHandshakeManager(l *logrus.Logger, mainHostMap *HostMap, lightHouse *Lig
|
|||||||
messageMetrics: config.messageMetrics,
|
messageMetrics: config.messageMetrics,
|
||||||
metricInitiated: metrics.GetOrRegisterCounter("handshake_manager.initiated", nil),
|
metricInitiated: metrics.GetOrRegisterCounter("handshake_manager.initiated", nil),
|
||||||
metricTimedOut: metrics.GetOrRegisterCounter("handshake_manager.timed_out", nil),
|
metricTimedOut: metrics.GetOrRegisterCounter("handshake_manager.timed_out", nil),
|
||||||
|
metricRateLimited: metrics.GetOrRegisterCounter("handshake_manager.rate_limited", nil),
|
||||||
|
rateBucket: config.maxHandshakeRate,
|
||||||
|
rateMax: config.maxHandshakeRate,
|
||||||
|
rateLastTick: time.Now(),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// handshakeRateAllow checks the token bucket rate limiter and returns true if a
|
||||||
|
// new handshake is allowed. Must be called with hm.Lock held.
|
||||||
|
func (hm *HandshakeManager) handshakeRateAllow(now time.Time) bool {
|
||||||
|
if hm.rateMax == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refill tokens based on elapsed time
|
||||||
|
elapsed := now.Sub(hm.rateLastTick)
|
||||||
|
if elapsed >= time.Second {
|
||||||
|
// Add tokens for full seconds elapsed
|
||||||
|
tokens := int(elapsed/time.Second) * hm.rateMax
|
||||||
|
hm.rateBucket += tokens
|
||||||
|
if hm.rateBucket > hm.rateMax {
|
||||||
|
hm.rateBucket = hm.rateMax
|
||||||
|
}
|
||||||
|
hm.rateLastTick = now
|
||||||
|
}
|
||||||
|
|
||||||
|
if hm.rateBucket > 0 {
|
||||||
|
hm.rateBucket--
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) Run(ctx context.Context) {
|
func (hm *HandshakeManager) Run(ctx context.Context) {
|
||||||
clockSource := time.NewTicker(hm.config.tryInterval)
|
clockSource := time.NewTicker(hm.config.tryInterval)
|
||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
@@ -149,6 +189,15 @@ func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *head
|
|||||||
case header.HandshakeIXPSK0:
|
case header.HandshakeIXPSK0:
|
||||||
switch h.MessageCounter {
|
switch h.MessageCounter {
|
||||||
case 1:
|
case 1:
|
||||||
|
// Check rate limit for new incoming handshakes
|
||||||
|
hm.Lock()
|
||||||
|
allowed := hm.handshakeRateAllow(time.Now())
|
||||||
|
hm.Unlock()
|
||||||
|
if !allowed {
|
||||||
|
hm.metricRateLimited.Inc(1)
|
||||||
|
hm.l.WithField("from", via).Debug("Handshake rate limit reached, dropping incoming handshake")
|
||||||
|
return
|
||||||
|
}
|
||||||
ixHandshakeStage1(hm.f, via, packet, h)
|
ixHandshakeStage1(hm.f, via, packet, h)
|
||||||
|
|
||||||
case 2:
|
case 2:
|
||||||
|
|||||||
@@ -65,6 +65,68 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
|||||||
assert.NotContains(t, blah.vpnIps, ip)
|
assert.NotContains(t, blah.vpnIps, ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_HandshakeManagerRateLimit(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
||||||
|
preferredRanges := []netip.Prefix{localrange}
|
||||||
|
mainHM := newHostMap(l)
|
||||||
|
mainHM.preferredRanges.Store(&preferredRanges)
|
||||||
|
|
||||||
|
lh := newTestLighthouse()
|
||||||
|
|
||||||
|
config := defaultHandshakeConfig
|
||||||
|
config.maxHandshakeRate = 2
|
||||||
|
|
||||||
|
hm := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, config)
|
||||||
|
hm.f = &Interface{handshakeManager: hm, pki: &PKI{}, l: l}
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
// Should allow up to maxHandshakeRate handshakes
|
||||||
|
hm.Lock()
|
||||||
|
assert.True(t, hm.handshakeRateAllow(now), "first handshake should be allowed")
|
||||||
|
assert.True(t, hm.handshakeRateAllow(now), "second handshake should be allowed")
|
||||||
|
assert.False(t, hm.handshakeRateAllow(now), "third handshake should be rate limited")
|
||||||
|
hm.Unlock()
|
||||||
|
|
||||||
|
// After advancing time by 1 second, tokens should refill
|
||||||
|
hm.Lock()
|
||||||
|
assert.True(t, hm.handshakeRateAllow(now.Add(time.Second)), "handshake should be allowed after token refill")
|
||||||
|
hm.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_HandshakeManagerRateLimitUnlimited(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
||||||
|
preferredRanges := []netip.Prefix{localrange}
|
||||||
|
mainHM := newHostMap(l)
|
||||||
|
mainHM.preferredRanges.Store(&preferredRanges)
|
||||||
|
|
||||||
|
lh := newTestLighthouse()
|
||||||
|
|
||||||
|
cs := &CertState{
|
||||||
|
initiatingVersion: cert.Version1,
|
||||||
|
privateKey: []byte{},
|
||||||
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
|
v1HandshakeBytes: []byte{},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default config has maxHandshakeRate=0 (unlimited)
|
||||||
|
hm := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
|
hm.f = &Interface{handshakeManager: hm, pki: &PKI{}, l: l}
|
||||||
|
hm.f.pki.cs.Store(cs)
|
||||||
|
|
||||||
|
// Should allow many handshakes with no limit
|
||||||
|
// Limited to 10 due to test lighthouse query channel buffer
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
ip := netip.MustParseAddr("172.1.1.1").As16()
|
||||||
|
ip[15] = byte(i + 1)
|
||||||
|
addr := netip.AddrFrom16(ip)
|
||||||
|
h := hm.StartHandshake(addr, nil)
|
||||||
|
assert.NotNil(t, h, "handshake %d should be allowed with unlimited rate", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func testCountTimerWheelEntries(tw *LockingTimerWheel[netip.Addr]) (c int) {
|
func testCountTimerWheelEntries(tw *LockingTimerWheel[netip.Addr]) (c int) {
|
||||||
for _, i := range tw.t.wheel {
|
for _, i := range tw.t.wheel {
|
||||||
n := i.Head
|
n := i.Head
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb []byte, batch *sendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Level >= logrus.DebugLevel {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
@@ -33,7 +33,7 @@ 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].WriteReject(packet)
|
_, err := f.readers[q].Write(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.WithError(err).Error("Failed to forward to tun")
|
f.l.WithError(err).Error("Failed to forward to tun")
|
||||||
}
|
}
|
||||||
@@ -53,7 +53,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, out, q)
|
||||||
if f.l.Level >= logrus.DebugLevel {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
f.l.WithField("vpnAddr", fwPacket.RemoteAddr).
|
f.l.WithField("vpnAddr", fwPacket.RemoteAddr).
|
||||||
WithField("fwPacket", fwPacket).
|
WithField("fwPacket", fwPacket).
|
||||||
@@ -68,10 +68,10 @@ 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.sendInsideMessage(hostinfo, packet, nb, batch, rejectBuf, q)
|
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, out, q)
|
||||||
if f.l.Level >= logrus.DebugLevel {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(f.l).
|
hostinfo.logger(f.l).
|
||||||
WithField("fwPacket", fwPacket).
|
WithField("fwPacket", fwPacket).
|
||||||
@@ -81,63 +81,6 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendInsideMessage encrypts a firewall-approved inside packet into the
|
|
||||||
// caller's batch slot for later sendmmsg flush. When hostinfo.remote is not
|
|
||||||
// valid we fall through to the relay slow path via the unbatched sendNoMetrics
|
|
||||||
// so relay behavior is unchanged.
|
|
||||||
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, p, nb []byte, batch *sendBatch, rejectBuf []byte, q int) {
|
|
||||||
ci := hostinfo.ConnectionState
|
|
||||||
if ci.eKey == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if !hostinfo.remote.IsValid() {
|
|
||||||
// Slow path: relay fallback. Reuse rejectBuf as the ciphertext
|
|
||||||
// scratch; sendNoMetrics arranges header space for SendVia.
|
|
||||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
scratch := batch.Next()
|
|
||||||
if scratch == nil {
|
|
||||||
// Batch full: bypass batching and send this packet directly so we
|
|
||||||
// never drop traffic on over-subscribed iterations.
|
|
||||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
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.Level >= logrus.DebugLevel {
|
|
||||||
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).Debug("Lighthouse update triggered for punch due to rebind counter")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
out, err := ci.eKey.EncryptDanger(out, out, p, c, nb)
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
ci.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).WithError(err).
|
|
||||||
WithField("udpAddr", hostinfo.remote).WithField("counter", c).
|
|
||||||
Error("Failed to encrypt outgoing packet")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
batch.Commit(len(out), hostinfo.remote)
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
||||||
@@ -148,7 +91,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := f.readers[q].WriteReject(out)
|
_, err := f.readers[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.WithError(err).Error("Failed to write to tun")
|
f.l.WithError(err).Error("Failed to write to tun")
|
||||||
}
|
}
|
||||||
|
|||||||
+45
-120
@@ -4,8 +4,10 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"sync"
|
"os"
|
||||||
|
"runtime"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -16,8 +18,6 @@ import (
|
|||||||
"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/coalesce"
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -87,18 +87,7 @@ type Interface struct {
|
|||||||
conntrackCacheTimeout time.Duration
|
conntrackCacheTimeout time.Duration
|
||||||
|
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []tio.Queue
|
readers []io.ReadWriteCloser
|
||||||
// tunCoalescers is one tcpCoalescer per tun queue, wrapping readers[i].
|
|
||||||
// decryptToTun sends plaintext into the coalescer; listenOut calls its
|
|
||||||
// Flush at the end of each UDP recvmmsg batch.
|
|
||||||
tunCoalescers []*coalesce.TCPCoalescer
|
|
||||||
wg sync.WaitGroup
|
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
|
||||||
// nil means "no fatal error" (yet)
|
|
||||||
fatalErr atomic.Pointer[error]
|
|
||||||
// triggerShutdown is a function that will be run exactly once, when onFatal swaps something non-nil into fatalErr
|
|
||||||
triggerShutdown func()
|
|
||||||
|
|
||||||
metricHandshakes metrics.Histogram
|
metricHandshakes metrics.Histogram
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
@@ -189,8 +178,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]tio.Queue, c.routines),
|
readers: make([]io.ReadWriteCloser, c.routines),
|
||||||
tunCoalescers: make([]*coalesce.TCPCoalescer, c.routines),
|
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -222,7 +210,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
// activate creates the interface on the host. After the interface is created, any
|
// activate creates the interface on the host. After the interface is created, any
|
||||||
// other services that want to bind listeners to its IP may do so successfully. However,
|
// other services that want to bind listeners to its IP may do so successfully. However,
|
||||||
// the interface isn't going to process anything until run() is called.
|
// the interface isn't going to process anything until run() is called.
|
||||||
func (f *Interface) activate() error {
|
func (f *Interface) activate() {
|
||||||
// actually turn on tun dev
|
// actually turn on tun dev
|
||||||
|
|
||||||
addr, err := f.outside.LocalAddr()
|
addr, err := f.outside.LocalAddr()
|
||||||
@@ -245,65 +233,38 @@ 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 {
|
||||||
err = f.inside.NewMultiQueueReader()
|
reader, err = f.inside.NewMultiQueueReader()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
f.l.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
f.readers[i] = reader
|
||||||
f.readers = f.inside.Readers()
|
|
||||||
for i := range f.readers {
|
|
||||||
f.tunCoalescers[i] = coalesce.NewTCPCoalescer(f.readers[i]) //todo don't always do this
|
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
if err := f.inside.Activate(); err != nil {
|
||||||
if err = f.inside.Activate(); err != nil {
|
|
||||||
f.wg.Done()
|
|
||||||
f.inside.Close()
|
f.inside.Close()
|
||||||
return err
|
f.l.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) run() (func() error, error) {
|
func (f *Interface) run() {
|
||||||
// Launch n queues to read packets from udp
|
// Launch n queues to read packets from udp
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
go f.listenOut(i)
|
||||||
f.listenOut(i)
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Launch n queues to read packets from tun dev
|
// Launch n queues to read packets from tun dev
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
f.wg.Go(func() {
|
go f.listenIn(f.readers[i], i)
|
||||||
f.listenIn(f.readers[i], i)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
return func() error {
|
|
||||||
f.wg.Wait()
|
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
|
||||||
return *e
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
|
||||||
func (f *Interface) onFatal(err error) {
|
|
||||||
swapped := f.fatalErr.CompareAndSwap(nil, &err)
|
|
||||||
if !swapped {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if f.triggerShutdown != nil {
|
|
||||||
f.triggerShutdown()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenOut(i int) {
|
func (f *Interface) listenOut(i int) {
|
||||||
|
runtime.LockOSThread()
|
||||||
|
|
||||||
var li udp.Conn
|
var li udp.Conn
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
li = f.writers[i]
|
li = f.writers[i]
|
||||||
@@ -313,75 +274,39 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(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)
|
||||||
|
|
||||||
// plaintexts is a ring of decrypt scratches, one per packet in a UDP
|
li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
// recvmmsg batch. The coalescer borrows payload slices from here and
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l))
|
||||||
// requires they stay valid until Flush — so we rotate each packet and
|
|
||||||
// reset only in the batch-end flush callback.
|
|
||||||
var plaintexts [][]byte
|
|
||||||
idx := 0
|
|
||||||
coalescer := f.tunCoalescers[i]
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
|
||||||
if idx >= len(plaintexts) {
|
|
||||||
plaintexts = append(plaintexts, make([]byte, udp.MTU))
|
|
||||||
}
|
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintexts[idx][:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l))
|
|
||||||
idx++
|
|
||||||
}, func() {
|
|
||||||
if err := coalescer.Flush(); err != nil {
|
|
||||||
f.l.WithError(err).Error("Failed to flush tun coalescer")
|
|
||||||
}
|
|
||||||
idx = 0
|
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
|
||||||
f.l.WithError(err).Error("Error while reading inbound packet, closing")
|
|
||||||
f.onFatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
f.l.Debugf("underlay reader %v is done", i)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader tio.Queue, i int) {
|
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||||
rejectBuf := make([]byte, mtu)
|
runtime.LockOSThread()
|
||||||
batch := newSendBatch(sendBatchCap, udp.MTU+32)
|
|
||||||
|
packet := make([]byte, mtu)
|
||||||
|
out := make([]byte, mtu)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
pkts, err := reader.Read()
|
n, err := reader.Read(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !f.closed.Load() {
|
if errors.Is(err, os.ErrClosed) && f.closed.Load() {
|
||||||
f.l.WithError(err).WithField("reader", i).Error("Error while reading outbound packet, closing")
|
return
|
||||||
f.onFatal(err)
|
|
||||||
}
|
}
|
||||||
break
|
|
||||||
|
f.l.WithError(err).Error("Error while reading outbound packet")
|
||||||
|
// This only seems to happen when something fatal happens to the fd, so exit.
|
||||||
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
|
||||||
batch.Reset()
|
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get(f.l))
|
||||||
for _, pkt := range pkts {
|
|
||||||
if batch.Len() >= batch.Cap() {
|
|
||||||
f.flushBatch(batch, i)
|
|
||||||
batch.Reset()
|
|
||||||
}
|
|
||||||
f.consumeInsidePacket(pkt, fwPacket, nb, batch, rejectBuf, i, conntrackCache.Get(f.l))
|
|
||||||
}
|
|
||||||
if batch.Len() > 0 {
|
|
||||||
f.flushBatch(batch, i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
f.l.Debugf("overlay reader %v is done", i)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) flushBatch(batch *sendBatch, q int) {
|
|
||||||
if err := f.writers[q].WriteBatch(batch.bufs, batch.dsts); err != nil {
|
|
||||||
f.l.WithError(err).WithField("writer", q).Error("Failed to write outgoing batch")
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -557,23 +482,23 @@ func (f *Interface) GetCertState() *CertState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) Close() error {
|
func (f *Interface) Close() error {
|
||||||
var errs []error
|
|
||||||
f.closed.Store(true)
|
f.closed.Store(true)
|
||||||
|
|
||||||
// Release the udp readers
|
for _, u := range f.writers {
|
||||||
for i, u := range f.writers {
|
|
||||||
err := u.Close()
|
err := u.Close()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.WithError(err).WithField("writer", i).Error("Error while closing udp socket")
|
f.l.WithError(err).Error("Error while closing udp socket")
|
||||||
errs = append(errs, err)
|
}
|
||||||
|
}
|
||||||
|
for i, r := range f.readers {
|
||||||
|
if i == 0 {
|
||||||
|
continue // f.readers[0] is f.inside, which we want to save for last
|
||||||
|
}
|
||||||
|
if err := r.Close(); err != nil {
|
||||||
|
f.l.WithError(err).Error("Error while closing tun reader")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Release the tun device (closing the tun also closes all readers)
|
// Release the tun device
|
||||||
closeErr := f.inside.Close()
|
return f.inside.Close()
|
||||||
if closeErr != nil {
|
|
||||||
errs = append(errs, closeErr)
|
|
||||||
}
|
|
||||||
f.wg.Done()
|
|
||||||
return errors.Join(errs...)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,10 +3,7 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
|
||||||
_ "net/http/pprof"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -52,11 +49,6 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
|||||||
l.Println(string(b))
|
l.Println(string(b))
|
||||||
}
|
}
|
||||||
|
|
||||||
//todo!!!
|
|
||||||
go func() {
|
|
||||||
log.Println(http.ListenAndServe("0.0.0.0:6060", nil))
|
|
||||||
}()
|
|
||||||
|
|
||||||
err := configLogger(l, c)
|
err := configLogger(l, c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.ContextualizeIfNeeded("Failed to configure the logger", err)
|
return nil, util.ContextualizeIfNeeded("Failed to configure the logger", err)
|
||||||
@@ -212,10 +204,11 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
|||||||
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
||||||
|
|
||||||
handshakeConfig := HandshakeConfig{
|
handshakeConfig := HandshakeConfig{
|
||||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||||
useRelays: useRelays,
|
useRelays: useRelays,
|
||||||
|
maxHandshakeRate: c.GetInt("handshakes.max_rate", DefaultMaxHandshakeRate),
|
||||||
|
|
||||||
messageMetrics: messageMetrics,
|
messageMetrics: messageMetrics,
|
||||||
}
|
}
|
||||||
@@ -296,16 +289,15 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
|||||||
}
|
}
|
||||||
|
|
||||||
return &Control{
|
return &Control{
|
||||||
state: StateReady,
|
ifce,
|
||||||
f: ifce,
|
l,
|
||||||
l: l,
|
ctx,
|
||||||
ctx: ctx,
|
cancel,
|
||||||
cancel: cancel,
|
sshStart,
|
||||||
sshStart: sshStart,
|
statsStart,
|
||||||
statsStart: statsStart,
|
dnsStart,
|
||||||
dnsStart: dnsStart,
|
lightHouse.StartUpdateWorker,
|
||||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
connManager.Start,
|
||||||
connectionManagerStart: connManager.Start,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -15,14 +15,14 @@ type endianness interface {
|
|||||||
var noiseEndianness endianness = binary.BigEndian
|
var noiseEndianness endianness = binary.BigEndian
|
||||||
|
|
||||||
type NebulaCipherState struct {
|
type NebulaCipherState struct {
|
||||||
c cipher.AEAD
|
c noise.Cipher
|
||||||
//k [32]byte
|
//k [32]byte
|
||||||
//n uint64
|
//n uint64
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewNebulaCipherState(s *noise.CipherState) *NebulaCipherState {
|
func NewNebulaCipherState(s *noise.CipherState) *NebulaCipherState {
|
||||||
x := s.Cipher()
|
return &NebulaCipherState{c: s.Cipher()}
|
||||||
return &NebulaCipherState{c: x.(cipher.AEAD)}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// EncryptDanger encrypts and authenticates a given payload.
|
// EncryptDanger encrypts and authenticates a given payload.
|
||||||
@@ -46,7 +46,7 @@ func (s *NebulaCipherState) EncryptDanger(out, ad, plaintext []byte, n uint64, n
|
|||||||
nb[2] = 0
|
nb[2] = 0
|
||||||
nb[3] = 0
|
nb[3] = 0
|
||||||
noiseEndianness.PutUint64(nb[4:], n)
|
noiseEndianness.PutUint64(nb[4:], n)
|
||||||
out = s.c.Seal(out, nb, plaintext, ad)
|
out = s.c.(cipher.AEAD).Seal(out, nb, plaintext, ad)
|
||||||
//l.Debugf("Encryption: outlen: %d, nonce: %d, ad: %s, plainlen %d", len(out), n, ad, len(plaintext))
|
//l.Debugf("Encryption: outlen: %d, nonce: %d, ad: %s, plainlen %d", len(out), n, ad, len(plaintext))
|
||||||
return out, nil
|
return out, nil
|
||||||
} else {
|
} else {
|
||||||
@@ -61,7 +61,7 @@ func (s *NebulaCipherState) DecryptDanger(out, ad, ciphertext []byte, n uint64,
|
|||||||
nb[2] = 0
|
nb[2] = 0
|
||||||
nb[3] = 0
|
nb[3] = 0
|
||||||
noiseEndianness.PutUint64(nb[4:], n)
|
noiseEndianness.PutUint64(nb[4:], n)
|
||||||
return s.c.Open(out, nb, ciphertext, ad)
|
return s.c.(cipher.AEAD).Open(out, nb, ciphertext, ad)
|
||||||
} else {
|
} else {
|
||||||
return []byte{}, nil
|
return []byte{}, nil
|
||||||
}
|
}
|
||||||
@@ -69,7 +69,7 @@ func (s *NebulaCipherState) DecryptDanger(out, ad, ciphertext []byte, n uint64,
|
|||||||
|
|
||||||
func (s *NebulaCipherState) Overhead() int {
|
func (s *NebulaCipherState) Overhead() int {
|
||||||
if s != nil {
|
if s != nil {
|
||||||
return s.c.Overhead()
|
return s.c.(cipher.AEAD).Overhead()
|
||||||
}
|
}
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -535,7 +535,7 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out
|
|||||||
}
|
}
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
f.connectionManager.In(hostinfo)
|
||||||
err = f.tunCoalescers[q].Add(out)
|
_, err = f.readers[q].Write(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.WithError(err).Error("Failed to write to tun")
|
f.l.WithError(err).Error("Failed to write to tun")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,484 +0,0 @@
|
|||||||
package coalesce
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ipProtoTCP is the IANA protocol number for TCP. Hardcoded instead of
|
|
||||||
// reaching for golang.org/x/sys/unix — that package doesn't define the
|
|
||||||
// constant on Windows, which would break cross-compiles even though this
|
|
||||||
// file runs unchanged on every platform.
|
|
||||||
const ipProtoTCP = 6
|
|
||||||
|
|
||||||
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
|
||||||
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
|
||||||
const tcpCoalesceBufSize = 65535
|
|
||||||
|
|
||||||
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
|
||||||
// superpacket. Keeping this well below the kernel's TSO ceiling bounds
|
|
||||||
// latency.
|
|
||||||
const tcpCoalesceMaxSegs = 64
|
|
||||||
|
|
||||||
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
|
||||||
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
|
||||||
const tcpCoalesceHdrCap = 100
|
|
||||||
|
|
||||||
// initialSlots is the starting capacity of the slot pool. One flow per
|
|
||||||
// packet is the worst case so this matches a typical UDP recvmmsg batch.
|
|
||||||
const initialSlots = 64
|
|
||||||
|
|
||||||
// flowKey identifies a TCP flow by {src, dst, sport, dport, family}.
|
|
||||||
// Comparable, so linear scans over the slot list stay tight.
|
|
||||||
type flowKey struct {
|
|
||||||
src, dst [16]byte
|
|
||||||
sport, dport uint16
|
|
||||||
isV6 bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
|
||||||
// passthrough is true the slot holds a single borrowed packet that must be
|
|
||||||
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
|
||||||
// passthrough is false the slot is an in-progress coalesced superpacket:
|
|
||||||
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
|
||||||
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
|
||||||
// slices from the caller's plaintext buffers — no payload is ever copied.
|
|
||||||
// The caller (listenOut) must keep those buffers alive until Flush.
|
|
||||||
type coalesceSlot struct {
|
|
||||||
passthrough bool
|
|
||||||
rawPkt []byte // borrowed when passthrough
|
|
||||||
|
|
||||||
fk flowKey
|
|
||||||
hdrBuf [tcpCoalesceHdrCap]byte
|
|
||||||
hdrLen int
|
|
||||||
ipHdrLen int
|
|
||||||
isV6 bool
|
|
||||||
gsoSize int
|
|
||||||
numSeg int
|
|
||||||
totalPay int
|
|
||||||
nextSeq uint32
|
|
||||||
// psh closes the chain: set when the last-accepted segment had PSH or
|
|
||||||
// was sub-gsoSize. No further appends after that.
|
|
||||||
psh bool
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
|
||||||
// multiple concurrent flows and emits each flow's run as a single TSO
|
|
||||||
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
|
||||||
// deferred until Flush so arrival order is preserved on the wire. Owns
|
|
||||||
// no locks; one coalescer per TUN write queue.
|
|
||||||
type TCPCoalescer struct {
|
|
||||||
plainW io.Writer
|
|
||||||
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
|
||||||
|
|
||||||
// slots is the ordered event queue. Flush walks it once and emits each
|
|
||||||
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
|
||||||
slots []*coalesceSlot
|
|
||||||
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
|
||||||
// segments can extend an in-progress superpacket in O(1). Slots are
|
|
||||||
// removed from this map when they close (PSH or short-last-segment),
|
|
||||||
// when a non-admissible packet for that flow arrives, or in Flush.
|
|
||||||
openSlots map[flowKey]*coalesceSlot
|
|
||||||
pool []*coalesceSlot // free list for reuse
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewTCPCoalescer(w io.Writer) *TCPCoalescer {
|
|
||||||
c := &TCPCoalescer{
|
|
||||||
plainW: w,
|
|
||||||
slots: make([]*coalesceSlot, 0, initialSlots),
|
|
||||||
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
|
||||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
|
||||||
}
|
|
||||||
if gw, ok := w.(tio.GSOWriter); ok && gw.GSOSupported() {
|
|
||||||
c.gsoW = gw
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
// parsedTCP holds the fields extracted from a single parse so later steps
|
|
||||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
|
||||||
type parsedTCP struct {
|
|
||||||
fk flowKey
|
|
||||||
ipHdrLen int
|
|
||||||
tcpHdrLen int
|
|
||||||
hdrLen int
|
|
||||||
payLen int
|
|
||||||
seq uint32
|
|
||||||
flags byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
|
||||||
// regardless of whether it's admissible for coalescing. Returns ok=false
|
|
||||||
// for non-TCP or malformed input. Accepts IPv4 (no options, no fragmentation)
|
|
||||||
// and IPv6 (no extension headers).
|
|
||||||
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
|
||||||
var p parsedTCP
|
|
||||||
if len(pkt) < 20 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
v := pkt[0] >> 4
|
|
||||||
switch v {
|
|
||||||
case 4:
|
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
|
||||||
if ihl != 20 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if pkt[9] != ipProtoTCP {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
// Reject actual fragmentation (MF or non-zero frag offset).
|
|
||||||
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
|
||||||
if totalLen > len(pkt) || totalLen < ihl {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = 20
|
|
||||||
p.fk.isV6 = false
|
|
||||||
copy(p.fk.src[:4], pkt[12:16])
|
|
||||||
copy(p.fk.dst[:4], pkt[16:20])
|
|
||||||
pkt = pkt[:totalLen]
|
|
||||||
case 6:
|
|
||||||
if len(pkt) < 40 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if pkt[6] != ipProtoTCP {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
|
||||||
if 40+payloadLen > len(pkt) {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.ipHdrLen = 40
|
|
||||||
p.fk.isV6 = true
|
|
||||||
copy(p.fk.src[:], pkt[8:24])
|
|
||||||
copy(p.fk.dst[:], pkt[24:40])
|
|
||||||
pkt = pkt[:40+payloadLen]
|
|
||||||
default:
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(pkt) < p.ipHdrLen+20 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
|
||||||
if tcpOff < 20 || tcpOff > 60 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
if len(pkt) < p.ipHdrLen+tcpOff {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.tcpHdrLen = tcpOff
|
|
||||||
p.hdrLen = p.ipHdrLen + tcpOff
|
|
||||||
p.payLen = len(pkt) - p.hdrLen
|
|
||||||
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
|
||||||
p.flags = pkt[p.ipHdrLen+13]
|
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
|
||||||
return p, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// coalesceable reports whether a parsed TCP segment is eligible for
|
|
||||||
// coalescing. Accepts only ACK or ACK|PSH with a non-empty payload.
|
|
||||||
func (p parsedTCP) coalesceable() bool {
|
|
||||||
const ack = 0x10
|
|
||||||
const psh = 0x08
|
|
||||||
if p.flags&^(ack|psh) != 0 || p.flags&ack == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return p.payLen > 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add borrows pkt. The caller must keep pkt valid until the next Flush,
|
|
||||||
// whether or not the packet was coalesced — passthrough (non-admissible)
|
|
||||||
// packets are queued and written at Flush time, not synchronously.
|
|
||||||
func (c *TCPCoalescer) Add(pkt []byte) error {
|
|
||||||
if c.gsoW == nil {
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
info, ok := parseTCPBase(pkt)
|
|
||||||
if !ok {
|
|
||||||
// Non-TCP or malformed — can't possibly collide with an open flow.
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if !info.coalesceable() {
|
|
||||||
// TCP but not admissible (SYN/FIN/RST/URG/CWR/ECE or zero-payload).
|
|
||||||
// Seal this flow's open slot so later in-flow packets don't extend
|
|
||||||
// it and accidentally reorder past this passthrough.
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if open := c.openSlots[info.fk]; open != nil {
|
|
||||||
if c.canAppend(open, pkt, info) {
|
|
||||||
c.appendPayload(open, pkt, info)
|
|
||||||
if open.psh {
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
}
|
|
||||||
c.seed(pkt, info)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush emits every queued event in arrival order. Coalesced slots go out
|
|
||||||
// via WriteGSO; passthrough slots go out via plainW.Write. 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.
|
|
||||||
func (c *TCPCoalescer) Flush() error {
|
|
||||||
var first error
|
|
||||||
for _, s := range c.slots {
|
|
||||||
var err error
|
|
||||||
if s.passthrough {
|
|
||||||
_, err = c.plainW.Write(s.rawPkt)
|
|
||||||
} else {
|
|
||||||
err = c.flushSlot(s)
|
|
||||||
}
|
|
||||||
if err != nil && first == nil {
|
|
||||||
first = err
|
|
||||||
}
|
|
||||||
c.release(s)
|
|
||||||
}
|
|
||||||
for i := range c.slots {
|
|
||||||
c.slots[i] = nil
|
|
||||||
}
|
|
||||||
c.slots = c.slots[:0]
|
|
||||||
for k := range c.openSlots {
|
|
||||||
delete(c.openSlots, k)
|
|
||||||
}
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
|
||||||
s := c.take()
|
|
||||||
s.passthrough = true
|
|
||||||
s.rawPkt = pkt
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
|
||||||
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
|
||||||
// Pathological shape — can't fit our scratch, emit as-is.
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
s := c.take()
|
|
||||||
s.passthrough = false
|
|
||||||
s.rawPkt = nil
|
|
||||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
|
||||||
s.hdrLen = info.hdrLen
|
|
||||||
s.ipHdrLen = info.ipHdrLen
|
|
||||||
s.isV6 = info.fk.isV6
|
|
||||||
s.fk = info.fk
|
|
||||||
s.gsoSize = info.payLen
|
|
||||||
s.numSeg = 1
|
|
||||||
s.totalPay = info.payLen
|
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
|
||||||
s.psh = info.flags&0x08 != 0
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
if !s.psh {
|
|
||||||
c.openSlots[info.fk] = s
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed: same
|
|
||||||
// header shape and stable contents, adjacent seq, not oversized, chain not
|
|
||||||
// closed.
|
|
||||||
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
|
||||||
if s.psh {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.hdrLen != s.hdrLen {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.seq != s.nextSeq {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.numSeg >= tcpCoalesceMaxSegs {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.payLen > s.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
|
||||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
s.numSeg++
|
|
||||||
s.totalPay += info.payLen
|
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
|
||||||
if info.payLen < s.gsoSize || info.flags&0x08 != 0 {
|
|
||||||
s.psh = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) take() *coalesceSlot {
|
|
||||||
if n := len(c.pool); n > 0 {
|
|
||||||
s := c.pool[n-1]
|
|
||||||
c.pool[n-1] = nil
|
|
||||||
c.pool = c.pool[:n-1]
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &coalesceSlot{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
|
||||||
s.passthrough = false
|
|
||||||
s.rawPkt = nil
|
|
||||||
for i := range s.payIovs {
|
|
||||||
s.payIovs[i] = nil
|
|
||||||
}
|
|
||||||
s.payIovs = s.payIovs[:0]
|
|
||||||
s.numSeg = 0
|
|
||||||
s.totalPay = 0
|
|
||||||
s.psh = false
|
|
||||||
c.pool = append(c.pool, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// flushSlot patches the header and calls WriteGSO. Does not remove the
|
|
||||||
// slot from c.slots.
|
|
||||||
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
|
||||||
total := s.hdrLen + s.totalPay
|
|
||||||
l4Len := total - s.ipHdrLen
|
|
||||||
hdr := s.hdrBuf[:s.hdrLen]
|
|
||||||
|
|
||||||
if s.isV6 {
|
|
||||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
|
||||||
hdr[10] = 0
|
|
||||||
hdr[11] = 0
|
|
||||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
|
||||||
}
|
|
||||||
|
|
||||||
var psum uint32
|
|
||||||
if s.isV6 {
|
|
||||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
|
||||||
} else {
|
|
||||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
|
||||||
}
|
|
||||||
tcsum := s.ipHdrLen + 16
|
|
||||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
|
||||||
|
|
||||||
return c.gsoW.WriteGSO(hdr, s.payIovs, uint16(s.gsoSize), s.isV6, uint16(s.ipHdrLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
|
||||||
// equality on every field that must be identical across coalesced
|
|
||||||
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
|
||||||
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|
||||||
if len(a) != len(b) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if isV6 {
|
|
||||||
// IPv6: bytes [0:4] = version/TC/flow-label, [6:8] = next_hdr/hop,
|
|
||||||
// [8:40] = src+dst. Skip [4:6] payload length.
|
|
||||||
if !bytes.Equal(a[0:4], b[0:4]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[6:40], b[6:40]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// IPv4: [0:2] version/IHL/TOS, [6:10] flags/fragoff/TTL/proto,
|
|
||||||
// [12:20] src+dst. Skip [2:4] total len, [4:6] id, [10:12] csum.
|
|
||||||
if !bytes.Equal(a[0:2], b[0:2]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[6:10], b[6:10]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[12:20], b[12:20]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
|
||||||
// [18:tcpHdrLen] options (incl. urgent).
|
|
||||||
tcp := ipHdrLen
|
|
||||||
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
|
||||||
// already have its checksum field zeroed) and returns the folded/inverted
|
|
||||||
// 16-bit value to store.
|
|
||||||
func ipv4HdrChecksum(hdr []byte) uint16 {
|
|
||||||
var sum uint32
|
|
||||||
for i := 0; i+1 < len(hdr); i += 2 {
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
|
||||||
}
|
|
||||||
if len(hdr)%2 == 1 {
|
|
||||||
sum += uint32(hdr[len(hdr)-1]) << 8
|
|
||||||
}
|
|
||||||
for sum>>16 != 0 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
}
|
|
||||||
return ^uint16(sum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoSumIPv4 / pseudoSumIPv6 build the TCP pseudo-header partial sum
|
|
||||||
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
|
||||||
// before folding.
|
|
||||||
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
|
||||||
var sum uint32
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
|
||||||
sum += uint32(proto)
|
|
||||||
sum += uint32(l4Len)
|
|
||||||
return sum
|
|
||||||
}
|
|
||||||
|
|
||||||
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
|
||||||
var sum uint32
|
|
||||||
for i := 0; i < 16; i += 2 {
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
|
||||||
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
|
||||||
}
|
|
||||||
sum += uint32(l4Len >> 16)
|
|
||||||
sum += uint32(l4Len & 0xffff)
|
|
||||||
sum += uint32(proto)
|
|
||||||
return sum
|
|
||||||
}
|
|
||||||
|
|
||||||
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
|
||||||
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
|
||||||
// the L4 checksum field — the kernel will add the payload sum and invert.
|
|
||||||
func foldOnceNoInvert(sum uint32) uint16 {
|
|
||||||
for sum>>16 != 0 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
}
|
|
||||||
return uint16(sum)
|
|
||||||
}
|
|
||||||
@@ -1,576 +0,0 @@
|
|||||||
package coalesce
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// fakeTunWriter records plain Writes and WriteGSO calls without touching a
|
|
||||||
// real TUN fd. WriteGSO preserves the split between hdr and borrowed pays
|
|
||||||
// so tests can inspect each independently.
|
|
||||||
type fakeTunWriter struct {
|
|
||||||
gsoEnabled bool
|
|
||||||
writes [][]byte
|
|
||||||
gsoWrites []fakeGSOWrite
|
|
||||||
}
|
|
||||||
|
|
||||||
type fakeGSOWrite struct {
|
|
||||||
hdr []byte
|
|
||||||
pays [][]byte
|
|
||||||
gsoSize uint16
|
|
||||||
isV6 bool
|
|
||||||
csumStart uint16
|
|
||||||
}
|
|
||||||
|
|
||||||
// total returns hdrLen + sum of pay lens.
|
|
||||||
func (g fakeGSOWrite) total() int {
|
|
||||||
n := len(g.hdr)
|
|
||||||
for _, p := range g.pays {
|
|
||||||
n += len(p)
|
|
||||||
}
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
|
|
||||||
// payLen sums the pays.
|
|
||||||
func (g fakeGSOWrite) payLen() int {
|
|
||||||
var n int
|
|
||||||
for _, p := range g.pays {
|
|
||||||
n += len(p)
|
|
||||||
}
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *fakeTunWriter) Write(p []byte) (int, error) {
|
|
||||||
buf := make([]byte, len(p))
|
|
||||||
copy(buf, p)
|
|
||||||
w.writes = append(w.writes, buf)
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *fakeTunWriter) WriteGSO(hdr []byte, pays [][]byte, gsoSize uint16, isV6 bool, csumStart uint16) error {
|
|
||||||
hcopy := make([]byte, len(hdr))
|
|
||||||
copy(hcopy, hdr)
|
|
||||||
paysCopy := make([][]byte, len(pays))
|
|
||||||
for i, p := range pays {
|
|
||||||
pc := make([]byte, len(p))
|
|
||||||
copy(pc, p)
|
|
||||||
paysCopy[i] = pc
|
|
||||||
}
|
|
||||||
w.gsoWrites = append(w.gsoWrites, fakeGSOWrite{
|
|
||||||
hdr: hcopy,
|
|
||||||
pays: paysCopy,
|
|
||||||
gsoSize: gsoSize,
|
|
||||||
isV6: isV6,
|
|
||||||
csumStart: csumStart,
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *fakeTunWriter) GSOSupported() bool { return w.gsoEnabled }
|
|
||||||
|
|
||||||
// buildTCPv4 constructs a minimal IPv4+TCP packet with the given payload,
|
|
||||||
// seq, and flags. Assumes no IP options and a 20-byte TCP header.
|
|
||||||
func buildTCPv4(seq uint32, flags byte, payload []byte) []byte {
|
|
||||||
return buildTCPv4Ports(1000, 2000, seq, flags, payload)
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv4Ports is buildTCPv4 with caller-specified ports so tests can
|
|
||||||
// build distinct flows.
|
|
||||||
func buildTCPv4Ports(sport, dport uint16, seq uint32, flags byte, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 20
|
|
||||||
const tcpHdrLen = 20
|
|
||||||
total := ipHdrLen + tcpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x45
|
|
||||||
pkt[1] = 0x00
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = ipProtoTCP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
|
||||||
binary.BigEndian.PutUint32(pkt[24:28], seq)
|
|
||||||
binary.BigEndian.PutUint32(pkt[28:32], 12345)
|
|
||||||
pkt[32] = 0x50
|
|
||||||
pkt[33] = flags
|
|
||||||
binary.BigEndian.PutUint16(pkt[34:36], 0xffff)
|
|
||||||
|
|
||||||
copy(pkt[40:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
const (
|
|
||||||
tcpAck = 0x10
|
|
||||||
tcpPsh = 0x08
|
|
||||||
tcpSyn = 0x02
|
|
||||||
tcpFin = 0x01
|
|
||||||
tcpAckPsh = tcpAck | tcpPsh
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: false}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pkt := buildTCPv4(1000, tcpAck, []byte("hello"))
|
|
||||||
if err := c.Add(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// No sync write — passthrough is deferred to Flush.
|
|
||||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("no Add-time writes: got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerNonTCPPassthrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pkt := make([]byte, 28)
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
|
||||||
pkt[9] = 1
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
if err := c.Add(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ICMP should pass through unchanged")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerSeedThenFlushAlone(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pkt := buildTCPv4(1000, tcpAck, make([]byte, 1000))
|
|
||||||
if err := c.Add(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("unexpected output before flush")
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Single-segment flush now goes through WriteGSO with GSO_NONE
|
|
||||||
// (virtio NEEDS_CSUM lets the kernel fill in the L4 csum).
|
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
|
||||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if g.total() != 40+1000 {
|
|
||||||
t.Errorf("super total=%d want %d", g.total(), 40+1000)
|
|
||||||
}
|
|
||||||
if g.payLen() != 1000 {
|
|
||||||
t.Errorf("payLen=%d want 1000", g.payLen())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerCoalescesAdjacentACKs(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Add(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if g.gsoSize != 1200 {
|
|
||||||
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
|
||||||
}
|
|
||||||
if len(g.hdr) != 40 {
|
|
||||||
t.Errorf("hdrLen=%d want 40", len(g.hdr))
|
|
||||||
}
|
|
||||||
if g.csumStart != 20 {
|
|
||||||
t.Errorf("csumStart=%d want 20", g.csumStart)
|
|
||||||
}
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
if g.total() != 40+3*1200 {
|
|
||||||
t.Errorf("superpacket len=%d want %d", g.total(), 40+3*1200)
|
|
||||||
}
|
|
||||||
if tot := binary.BigEndian.Uint16(g.hdr[2:4]); int(tot) != g.total() {
|
|
||||||
t.Errorf("ip total_length=%d want %d", tot, g.total())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerRejectsSeqGap(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Add(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(3000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Each packet flushes as its own single-segment WriteGSO now.
|
|
||||||
if len(w.gsoWrites) != 2 || len(w.writes) != 0 {
|
|
||||||
t.Fatalf("seq gap: want 2 gso writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerRejectsFlagMismatch(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Add(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// SYN|ACK is non-admissible. Must flush matching flow's slot (gso)
|
|
||||||
// and then plain-write the SYN packet itself.
|
|
||||||
syn := buildTCPv4(2200, tcpSyn|tcpAck, pay)
|
|
||||||
if err := c.Add(syn); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("flag mismatch: want 1 plain + 1 gso, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerRejectsFIN(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
fin := buildTCPv4(1000, tcpAck|tcpFin, []byte("x"))
|
|
||||||
if err := c.Add(fin); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// FIN isn't admissible — passthrough as plain, no slot, no gso.
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("FIN should be passthrough, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerShortLastSegmentClosesChain(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
full := make([]byte, 1200)
|
|
||||||
half := make([]byte, 500)
|
|
||||||
if err := c.Add(buildTCPv4(1000, tcpAck, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(2200, tcpAck, half)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Chain now closed; next packet seeds a new slot on the same flow
|
|
||||||
// after flushing the old one.
|
|
||||||
if err := c.Add(buildTCPv4(2700, tcpAck, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Expect two gso writes: the first two packets coalesced, then the
|
|
||||||
// third flushed alone (single-seg via GSO_NONE).
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 0 {
|
|
||||||
t.Fatalf("want 0 plain writes got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
if w.gsoWrites[0].gsoSize != 1200 {
|
|
||||||
t.Errorf("gsoSize=%d want 1200", w.gsoWrites[0].gsoSize)
|
|
||||||
}
|
|
||||||
if got, want := w.gsoWrites[0].total(), 40+1200+500; got != want {
|
|
||||||
t.Errorf("super len=%d want %d", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerPSHFinalizesChain(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Add(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(2200, tcpAckPsh, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// First two coalesce; the third seeds a fresh slot that flushes alone.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 0 {
|
|
||||||
t.Fatalf("want 0 plain writes got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerRejectsDifferentFlow(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
p1 := buildTCPv4(1000, tcpAck, pay)
|
|
||||||
p2 := buildTCPv4(2200, tcpAck, pay)
|
|
||||||
binary.BigEndian.PutUint16(p2[20:22], 9999)
|
|
||||||
if err := c.Add(p1); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(p2); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Two independent flows, each flushes its own single-segment WriteGSO.
|
|
||||||
if len(w.gsoWrites) != 2 || len(w.writes) != 0 {
|
|
||||||
t.Fatalf("diff flow: want 2 gso writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerRejectsIPOptions(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 500)
|
|
||||||
pkt := buildTCPv4(1000, tcpAck, pay)
|
|
||||||
// Bump IHL to 6 to simulate 4 bytes of IP options. Don't actually add
|
|
||||||
// bytes — parser should bail before it matters.
|
|
||||||
pkt[0] = 0x46
|
|
||||||
if err := c.Add(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Non-admissible parse → passthrough as plain.
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("IP options should passthrough, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCoalescerCapBySegments(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 512)
|
|
||||||
seq := uint32(1000)
|
|
||||||
for i := 0; i < tcpCoalesceMaxSegs+5; i++ {
|
|
||||||
if err := c.Add(buildTCPv4(seq, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
seq += uint32(len(pay))
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
for _, g := range w.gsoWrites {
|
|
||||||
segs := len(g.pays)
|
|
||||||
if segs > tcpCoalesceMaxSegs {
|
|
||||||
t.Fatalf("super exceeded seg cap: %d > %d", segs, tcpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerMultipleFlowsInSameBatch proves two interleaved bulk TCP
|
|
||||||
// flows coalesce independently in a single Flush.
|
|
||||||
func TestCoalescerMultipleFlowsInSameBatch(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
// Flow A: sport 1000. Flow B: sport 3000.
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 2500, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 2900, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 0 {
|
|
||||||
t.Fatalf("want no plain writes, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
// Each superpacket should carry 3 segments.
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Errorf("gso[%d]: segs=%d want 3", i, len(g.pays))
|
|
||||||
}
|
|
||||||
if g.gsoSize != 1200 {
|
|
||||||
t.Errorf("gso[%d]: gsoSize=%d want 1200", i, g.gsoSize)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Verify each superpacket carries the source port it was seeded with.
|
|
||||||
seenSports := map[uint16]bool{}
|
|
||||||
for _, g := range w.gsoWrites {
|
|
||||||
sp := binary.BigEndian.Uint16(g.hdr[20:22])
|
|
||||||
seenSports[sp] = true
|
|
||||||
}
|
|
||||||
if !seenSports[1000] || !seenSports[3000] {
|
|
||||||
t.Errorf("expected superpackets for sports 1000 and 3000, got %v", seenSports)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerPreservesArrivalOrder confirms that with passthrough and
|
|
||||||
// coalesced events both queued, Flush emits them in Add order rather than
|
|
||||||
// writing passthrough packets synchronously.
|
|
||||||
func TestCoalescerPreservesArrivalOrder(t *testing.T) {
|
|
||||||
w := &orderedFakeWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
// Sequence: coalesceable TCP, ICMP (passthrough), coalesceable TCP on
|
|
||||||
// a different flow. Expected emit order: gso(X), plain(ICMP), gso(Y).
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
icmp := make([]byte, 28)
|
|
||||||
icmp[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(icmp[2:4], 28)
|
|
||||||
icmp[9] = 1
|
|
||||||
copy(icmp[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(icmp[16:20], []byte{10, 0, 0, 3})
|
|
||||||
if err := c.Add(icmp); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Nothing should have hit the writer synchronously.
|
|
||||||
if len(w.events) != 0 {
|
|
||||||
t.Fatalf("Add emitted events synchronously: %v", w.events)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if got, want := w.events, []string{"gso", "plain", "gso"}; !stringSliceEq(got, want) {
|
|
||||||
t.Fatalf("flush order=%v want %v", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// orderedFakeWriter records only the sequence of call types so tests can
|
|
||||||
// assert arrival order without inspecting bytes.
|
|
||||||
type orderedFakeWriter struct {
|
|
||||||
gsoEnabled bool
|
|
||||||
events []string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *orderedFakeWriter) Write(p []byte) (int, error) {
|
|
||||||
w.events = append(w.events, "plain")
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *orderedFakeWriter) WriteGSO(hdr []byte, pays [][]byte, gsoSize uint16, isV6 bool, csumStart uint16) error {
|
|
||||||
w.events = append(w.events, "gso")
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *orderedFakeWriter) GSOSupported() bool { return w.gsoEnabled }
|
|
||||||
|
|
||||||
func stringSliceEq(a, b []string) bool {
|
|
||||||
if len(a) != len(b) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
for i := range a {
|
|
||||||
if a[i] != b[i] {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerInterleavedFlowsPreserveOrdering checks that a non-admissible
|
|
||||||
// packet (SYN) mid-flow only flushes its own flow, not others.
|
|
||||||
func TestCoalescerInterleavedFlowsPreserveOrdering(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
|
|
||||||
// Flow A two segments.
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Flow B two segments.
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Flow A SYN (non-admissible) — must flush only flow A's slot.
|
|
||||||
syn := buildTCPv4Ports(1000, 2000, 9999, tcpSyn|tcpAck, pay)
|
|
||||||
if err := c.Add(syn); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Flow B continues — should still be coalesced with its seed.
|
|
||||||
if err := c.Add(buildTCPv4Ports(3000, 2000, 2900, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Expected:
|
|
||||||
// - 1 gso for flow A (first 2 segments)
|
|
||||||
// - 1 plain for flow A SYN
|
|
||||||
// - 1 gso for flow B (3 segments)
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want 1 plain write (SYN), got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
// Find the 3-segment gso (flow B) and the 2-segment gso (flow A).
|
|
||||||
var segCounts []int
|
|
||||||
for _, g := range w.gsoWrites {
|
|
||||||
segCounts = append(segCounts, len(g.pays))
|
|
||||||
}
|
|
||||||
if !(segCounts[0] == 2 && segCounts[1] == 3) && !(segCounts[0] == 3 && segCounts[1] == 2) {
|
|
||||||
t.Errorf("unexpected segment counts: %v (want 2 and 3)", segCounts)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+3
-9
@@ -4,21 +4,15 @@ 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.Closer
|
io.ReadWriteCloser
|
||||||
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 //todo remove?
|
SupportsMultiqueue() bool
|
||||||
NewMultiQueueReader() error
|
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
||||||
Readers() []tio.Queue
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,70 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
type offloadContainer struct {
|
|
||||||
pq []*Offload
|
|
||||||
// pqi is exactly the same as pq, but stored as the interface type
|
|
||||||
pqi []Queue
|
|
||||||
shutdownFd int
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewOffloadContainer() (Container, 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 := &offloadContainer{
|
|
||||||
pq: []*Offload{},
|
|
||||||
pqi: []Queue{},
|
|
||||||
shutdownFd: shutdownFd,
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *offloadContainer) Queues() []Queue {
|
|
||||||
return c.pqi
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *offloadContainer) Add(fd int) error {
|
|
||||||
x, err := newOffload(fd, c.shutdownFd)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
c.pq = append(c.pq, x)
|
|
||||||
c.pqi = append(c.pqi, x)
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *offloadContainer) wakeForShutdown() error {
|
|
||||||
var buf [8]byte
|
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
|
||||||
_, err := unix.Write(c.shutdownFd, buf[:])
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *offloadContainer) Close() error {
|
|
||||||
errs := []error{}
|
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return errors.Join(errs...)
|
|
||||||
}
|
|
||||||
@@ -1,69 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
type pollContainer struct {
|
|
||||||
pq []*Poll
|
|
||||||
// pqi is exactly the same as pq, but stored as the interface type
|
|
||||||
pqi []Queue
|
|
||||||
shutdownFd int
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewPollContainer() (Container, 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 := &pollContainer{
|
|
||||||
pq: []*Poll{},
|
|
||||||
pqi: []Queue{},
|
|
||||||
shutdownFd: shutdownFd,
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *pollContainer) Queues() []Queue {
|
|
||||||
return c.pqi
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *pollContainer) 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 *pollContainer) wakeForShutdown() error {
|
|
||||||
var buf [8]byte
|
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
|
||||||
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *pollContainer) Close() error {
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return errors.Join(errs...)
|
|
||||||
}
|
|
||||||
@@ -1,63 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import "io"
|
|
||||||
|
|
||||||
// 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 Container interface {
|
|
||||||
Queues() []Queue
|
|
||||||
Add(fd int) error
|
|
||||||
|
|
||||||
io.Closer
|
|
||||||
}
|
|
||||||
|
|
||||||
// Queue is a readable/writable Poll queue. One Queue is driven by a single
|
|
||||||
// read goroutine plus concurrent writers (see Write / WriteReject below).
|
|
||||||
type Queue interface {
|
|
||||||
io.Closer
|
|
||||||
|
|
||||||
// Read returns one or more packets. The returned slices are borrowed
|
|
||||||
// from the Queue's internal buffer and are only valid until the next
|
|
||||||
// Read or Close on this Queue — callers must encrypt or copy each
|
|
||||||
// slice before the next call. Not safe for concurrent Reads; exactly
|
|
||||||
// one goroutine per Queue reads.
|
|
||||||
Read() ([][]byte, error)
|
|
||||||
|
|
||||||
// Write emits a single packet on the plaintext (outside→inside)
|
|
||||||
// delivery path. May run concurrently with WriteReject on the same
|
|
||||||
// Queue, but not with itself.
|
|
||||||
Write(p []byte) (int, error)
|
|
||||||
|
|
||||||
// WriteReject writes a single packet that originated from the inside
|
|
||||||
// path (reject replies or self-forward) using scratch state distinct
|
|
||||||
// from Write, so it can run concurrently with Write on the same Queue
|
|
||||||
// without a data race. On backends without a shared-scratch Write, a
|
|
||||||
// trivial delegation to Write is acceptable.
|
|
||||||
WriteReject(p []byte) (int, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSOWriter is implemented by Queues that can emit a TCP TSO 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 return false from GSOSupported and coalescing is skipped.
|
|
||||||
//
|
|
||||||
// hdr contains the IPv4/IPv6 + TCP header prefix (mutable — callers will
|
|
||||||
// have filled in total length and pseudo-header partial). 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. gsoSize is the MSS:
|
|
||||||
// every segment except possibly the last is exactly that many bytes.
|
|
||||||
// csumStart is the byte offset where the TCP header begins within hdr.
|
|
||||||
//
|
|
||||||
// # TODO fold into Queue
|
|
||||||
//
|
|
||||||
// hdr's TCP checksum field must already hold the pseudo-header partial
|
|
||||||
// sum (single-fold, not inverted), per virtio NEEDS_CSUM semantics.
|
|
||||||
type GSOWriter interface {
|
|
||||||
WriteGSO(hdr []byte, pays [][]byte, gsoSize uint16, isV6 bool, csumStart uint16) error
|
|
||||||
GSOSupported() bool
|
|
||||||
}
|
|
||||||
@@ -1,434 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"sync/atomic"
|
|
||||||
"syscall"
|
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Space for segmented output. Worst case is many small segments, each paying
|
|
||||||
// an IP+TCP header. Should be a multiple of 64KiB.
|
|
||||||
// const tunSegBufSize = 0xffff * 8 TODO larger? config?
|
|
||||||
const tunSegBufSize = 131072
|
|
||||||
|
|
||||||
// tunSegBufCap is the total size we allocate for the per-reader segment
|
|
||||||
// buffer. It is sized as one worst-case TSO superpacket (tunSegBufSize) plus
|
|
||||||
// the same again as drain headroom so a Read wake can accumulate
|
|
||||||
// additional packets after an initial big read without overflowing.
|
|
||||||
const tunSegBufCap = tunSegBufSize * 2
|
|
||||||
|
|
||||||
// tunDrainCap caps how many packets a single Read will accumulate via
|
|
||||||
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
|
||||||
// bounding how much work a single caller holds before handing off.
|
|
||||||
const tunDrainCap = 64 //256
|
|
||||||
|
|
||||||
// gsoInitialPayIovs is the starting capacity (in payload fragments) of
|
|
||||||
// Offload.gsoIovs. Sized to cover the default coalesce segment cap without
|
|
||||||
// any reallocations.
|
|
||||||
const gsoInitialPayIovs = 66
|
|
||||||
|
|
||||||
// gsoWriteBufCap is the initial per-queue coalesce scratch capacity used by
|
|
||||||
// WriteGSO to assemble [virtio_hdr || IP/TCP hdr || pays...] into a single
|
|
||||||
// contiguous buffer so we can emit the superpacket via a single write()
|
|
||||||
// instead of writev(). One worst-case TSO superpacket is bounded by the
|
|
||||||
// virtio spec at 64KiB; 128KiB gives comfortable slack for the 10-byte
|
|
||||||
// virtio header, the IP/TCP header, and any future size bumps. Grown on
|
|
||||||
// demand if a superpacket exceeds this.
|
|
||||||
const gsoWriteBufCap = tunSegBufSize
|
|
||||||
|
|
||||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
|
||||||
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
|
||||||
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
|
|
||||||
// verification. All packets that reach the plain Write / WriteReject paths
|
|
||||||
// already carry a valid L4 checksum (either supplied by a remote peer whose
|
|
||||||
// ciphertext we AEAD-authenticated, or produced by finishChecksum during TSO
|
|
||||||
// segmentation, or built locally by CreateRejectPacket), so trusting them is
|
|
||||||
// safe.
|
|
||||||
var validVnetHdr = [virtioNetHdrLen]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
|
||||||
|
|
||||||
// Offload 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 Offload struct {
|
|
||||||
fd int
|
|
||||||
shutdownFd int
|
|
||||||
readPoll [2]unix.PollFd
|
|
||||||
writePoll [2]unix.PollFd
|
|
||||||
closed atomic.Bool
|
|
||||||
readBuf []byte // scratch for a single raw read (virtio hdr + superpacket)
|
|
||||||
segBuf []byte // backing store for segmented output
|
|
||||||
segOff int // cursor into segBuf for the current Read drain
|
|
||||||
pending [][]byte // segments returned from the most recent Read
|
|
||||||
writeIovs [2]unix.Iovec // preallocated iovecs for Write (coalescer passthrough); iovs[0] is fixed to validVnetHdr
|
|
||||||
// rejectIovs is a second preallocated iovec scratch used exclusively by
|
|
||||||
// WriteReject (reject + self-forward from the inside path). It mirrors
|
|
||||||
// writeIovs but lets listenIn goroutines emit reject packets without
|
|
||||||
// racing with the listenOut coalescer that owns writeIovs.
|
|
||||||
rejectIovs [2]unix.Iovec
|
|
||||||
|
|
||||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
|
||||||
// by WriteGSO. Separate from validVnetHdr so a concurrent non-GSO Write on
|
|
||||||
// another queue never observes a half-written header.
|
|
||||||
gsoHdrBuf [virtioNetHdrLen]byte
|
|
||||||
// gsoIovs is a legacy writev iovec scratch. No longer used by the
|
|
||||||
// WriteGSO path (which coalesces into gsoWriteBuf and uses a single
|
|
||||||
// write()) but retained for any other iovec-based path that may use it.
|
|
||||||
gsoIovs []unix.Iovec
|
|
||||||
|
|
||||||
// gsoWriteBuf is a per-queue scratch used by WriteGSO to coalesce the
|
|
||||||
// virtio_net_hdr + IP/TCP header + payload fragments into a single
|
|
||||||
// contiguous buffer, which is then written to the TUN fd with one
|
|
||||||
// write() syscall. This mirrors wireguard-go's approach and avoids
|
|
||||||
// triggering a kernel refcount use-after-free in skb_set_owner_w /
|
|
||||||
// sock_wfree observed on Linux 4.19 TUN when scatter-gather writev is
|
|
||||||
// combined with GSO-flagged virtio_net_hdr in the tun_chr_write_iter
|
|
||||||
// path. Grown on demand if a superpacket exceeds the initial cap.
|
|
||||||
gsoWriteBuf []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func newOffload(fd int, shutdownFd int) (*Offload, error) {
|
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
out := &Offload{
|
|
||||||
fd: fd,
|
|
||||||
shutdownFd: shutdownFd,
|
|
||||||
closed: atomic.Bool{},
|
|
||||||
readBuf: make([]byte, tunReadBufSize),
|
|
||||||
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},
|
|
||||||
},
|
|
||||||
|
|
||||||
segBuf: make([]byte, tunSegBufCap),
|
|
||||||
gsoIovs: make([]unix.Iovec, 2, 2+gsoInitialPayIovs),
|
|
||||||
gsoWriteBuf: make([]byte, 0, gsoWriteBufCap),
|
|
||||||
}
|
|
||||||
|
|
||||||
out.writeIovs[0].Base = &validVnetHdr[0]
|
|
||||||
out.writeIovs[0].SetLen(virtioNetHdrLen)
|
|
||||||
out.rejectIovs[0].Base = &validVnetHdr[0]
|
|
||||||
out.rejectIovs[0].SetLen(virtioNetHdrLen)
|
|
||||||
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
|
||||||
out.gsoIovs[0].SetLen(virtioNetHdrLen)
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) 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 *Offload) 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 *Offload) readRaw(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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read reads one or more superpackets from the tun and returns the
|
|
||||||
// resulting packets. The first read blocks via poll; once the fd is known
|
|
||||||
// readable we drain additional packets non-blocking until the kernel queue
|
|
||||||
// is empty (EAGAIN), we've collected tunDrainCap packets, or we're out of
|
|
||||||
// segBuf headroom. This amortizes the poll wake over bursts of small
|
|
||||||
// packets (e.g. TCP ACKs). Slices point into the Offload's internal buffers
|
|
||||||
// and are only valid until the next Read or Close on this Queue.
|
|
||||||
func (r *Offload) Read() ([][]byte, error) {
|
|
||||||
r.pending = r.pending[:0]
|
|
||||||
r.segOff = 0
|
|
||||||
|
|
||||||
// Initial (blocking) read. Retry on decode errors so a single bad
|
|
||||||
// packet does not stall the reader.
|
|
||||||
for {
|
|
||||||
n, err := r.readRaw(r.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if err := r.decodeRead(n); err != nil {
|
|
||||||
// Drop and read again — a bad packet should not kill the reader.
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
|
||||||
|
|
||||||
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
|
||||||
// cap is reached, or segBuf no longer has room for another worst-case
|
|
||||||
// superpacket.
|
|
||||||
for len(r.pending) < tunDrainCap && tunSegBufCap-r.segOff >= tunSegBufSize {
|
|
||||||
n, err := unix.Read(r.fd, r.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
// EAGAIN / EINTR / anything else: stop draining. We already
|
|
||||||
// have a valid batch from the first read.
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if n <= 0 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if err := r.decodeRead(n); err != nil {
|
|
||||||
// Drop this packet and stop the drain; we'd rather hand off
|
|
||||||
// what we have than keep spinning here.
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return r.pending, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// decodeRead decodes the virtio header plus payload in r.readBuf[:n], appends
|
|
||||||
// the segments to r.pending, and advances r.segOff by the total scratch used.
|
|
||||||
// Caller must have already ensured r.vnetHdr is true.
|
|
||||||
func (r *Offload) decodeRead(n int) error {
|
|
||||||
if n < virtioNetHdrLen {
|
|
||||||
return fmt.Errorf("short tun read: %d < %d", n, virtioNetHdrLen)
|
|
||||||
}
|
|
||||||
var hdr VirtioNetHdr
|
|
||||||
hdr.decode(r.readBuf[:virtioNetHdrLen])
|
|
||||||
before := len(r.pending)
|
|
||||||
if err := segmentInto(r.readBuf[virtioNetHdrLen:n], hdr, &r.pending, r.segBuf[r.segOff:]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
for k := before; k < len(r.pending); k++ {
|
|
||||||
r.segOff += len(r.pending[k])
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) Write(buf []byte) (int, error) {
|
|
||||||
return r.writeWithScratch(buf, &r.writeIovs)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WriteReject emits a packet using a dedicated iovec scratch (rejectIovs)
|
|
||||||
// distinct from the one used by the coalescer's Write path. This avoids a
|
|
||||||
// data race between the inside (listenIn) goroutine emitting reject or
|
|
||||||
// self-forward packets and the outside (listenOut) goroutine flushing TCP
|
|
||||||
// coalescer passthroughs on the same Offload.
|
|
||||||
func (r *Offload) WriteReject(buf []byte) (int, error) {
|
|
||||||
return r.writeWithScratch(buf, &r.rejectIovs)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error) {
|
|
||||||
if len(buf) == 0 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
// Point the payload iovec at the caller's buffer. iovs[0] is pre-wired
|
|
||||||
// to validVnetHdr during Offload construction so we don't rebuild it here.
|
|
||||||
iovs[1].Base = &buf[0]
|
|
||||||
iovs[1].SetLen(len(buf))
|
|
||||||
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
|
||||||
for {
|
|
||||||
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
|
||||||
if errno == 0 {
|
|
||||||
if int(n) < virtioNetHdrLen {
|
|
||||||
return 0, io.ErrShortWrite
|
|
||||||
}
|
|
||||||
return int(n) - virtioNetHdrLen, nil
|
|
||||||
}
|
|
||||||
if errno == unix.EAGAIN {
|
|
||||||
if err := r.blockOnWrite(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if errno == unix.EINTR {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if errno == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
}
|
|
||||||
return 0, errno
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// rawWriteSingle writes buf to the TUN fd with a single write() syscall.
|
|
||||||
// Unlike rawWrite (which uses writev), this avoids the kernel
|
|
||||||
// scatter-gather path that triggers a use-after-free in
|
|
||||||
// tun_chr_write_iter → sock_alloc_send_pskb → skb_set_owner_w on Linux
|
|
||||||
// 4.19 TUN when the virtio_net_hdr requests TSO segmentation. The caller
|
|
||||||
// is responsible for including the virtio_net_hdr prefix in buf.
|
|
||||||
func (r *Offload) rawWriteSingle(buf []byte) (int, error) {
|
|
||||||
for {
|
|
||||||
n, err := unix.Write(r.fd, buf)
|
|
||||||
if err == nil {
|
|
||||||
if n < virtioNetHdrLen {
|
|
||||||
return 0, io.ErrShortWrite
|
|
||||||
}
|
|
||||||
return n - virtioNetHdrLen, nil
|
|
||||||
}
|
|
||||||
if err == unix.EAGAIN {
|
|
||||||
if werr := r.blockOnWrite(); werr != nil {
|
|
||||||
return 0, werr
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if err == unix.EINTR {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if err == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
}
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSOSupported reports whether this queue was opened with IFF_VNET_HDR and
|
|
||||||
// can accept WriteGSO. When false, callers should fall back to per-segment
|
|
||||||
// Write calls.
|
|
||||||
func (r *Offload) GSOSupported() bool { return true }
|
|
||||||
|
|
||||||
// WriteGSO emits a TCP TSO superpacket. hdr is the IPv4/IPv6 + TCP header
|
|
||||||
// prefix (already finalized — total length, IP csum, and TCP pseudo-header
|
|
||||||
// partial set by the caller). pays are payload fragments whose concatenation
|
|
||||||
// forms the full coalesced payload. gsoSize is the MSS; every segment except
|
|
||||||
// possibly the last is exactly gsoSize bytes. csumStart is the byte offset
|
|
||||||
// where the TCP header begins within hdr.
|
|
||||||
//
|
|
||||||
// Implementation note: this path coalesces [virtio_hdr || hdr || pays...]
|
|
||||||
// into a single contiguous scratch buffer (r.gsoWriteBuf) and emits it via
|
|
||||||
// one write() syscall rather than writev() with a scatter-gather iovec.
|
|
||||||
// The scatter-gather path triggered a kernel-side use-after-free on Linux
|
|
||||||
// 4.19 TUN where tun_chr_write_iter → sock_alloc_send_pskb →
|
|
||||||
// skb_set_owner_w could be invoked with a zero sk_wmem_alloc, crashing
|
|
||||||
// the router. The single-write path mirrors wireguard-go's design (see
|
|
||||||
// golang.zx2c4.com/wireguard/tun/tun_linux.go Write — it always coalesces
|
|
||||||
// GRO-merged data into a single contiguous buffer before calling
|
|
||||||
// tunFile.Write) and has no equivalent failure mode.
|
|
||||||
func (r *Offload) WriteGSO(hdr []byte, pays [][]byte, gsoSize uint16, isV6 bool, csumStart uint16) error {
|
|
||||||
if len(hdr) == 0 || len(pays) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build the virtio_net_hdr. When pays total to <= gsoSize the kernel
|
|
||||||
// would produce a single segment; keep NEEDS_CSUM semantics but skip
|
|
||||||
// the GSO type so the kernel doesn't spuriously mark this as TSO.
|
|
||||||
vhdr := VirtioNetHdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
HdrLen: uint16(len(hdr)),
|
|
||||||
GSOSize: gsoSize,
|
|
||||||
CsumStart: csumStart,
|
|
||||||
CsumOffset: 16, // TCP checksum field lives 16 bytes into the TCP header
|
|
||||||
}
|
|
||||||
var totalPay int
|
|
||||||
for _, p := range pays {
|
|
||||||
totalPay += len(p)
|
|
||||||
}
|
|
||||||
if totalPay > int(gsoSize) {
|
|
||||||
if isV6 {
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
|
||||||
} else {
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
|
||||||
vhdr.GSOSize = 0
|
|
||||||
}
|
|
||||||
vhdr.encode(r.gsoHdrBuf[:])
|
|
||||||
|
|
||||||
// Coalesce [virtio_hdr || hdr || pays...] into a single contiguous
|
|
||||||
// buffer. This avoids the kernel scatter-gather write path entirely.
|
|
||||||
need := virtioNetHdrLen + len(hdr) + totalPay
|
|
||||||
if cap(r.gsoWriteBuf) < need {
|
|
||||||
// Grow geometrically to amortize reallocs.
|
|
||||||
newCap := cap(r.gsoWriteBuf) * 2
|
|
||||||
if newCap < need {
|
|
||||||
newCap = need
|
|
||||||
}
|
|
||||||
r.gsoWriteBuf = make([]byte, 0, newCap)
|
|
||||||
} else {
|
|
||||||
r.gsoWriteBuf = r.gsoWriteBuf[:0]
|
|
||||||
}
|
|
||||||
r.gsoWriteBuf = append(r.gsoWriteBuf, r.gsoHdrBuf[:]...)
|
|
||||||
r.gsoWriteBuf = append(r.gsoWriteBuf, hdr...)
|
|
||||||
for _, p := range pays {
|
|
||||||
r.gsoWriteBuf = append(r.gsoWriteBuf, p...)
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err := r.rawWriteSingle(r.gsoWriteBuf)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) Close() error {
|
|
||||||
if r.closed.Swap(true) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
//shutdownFd is owned by the container, so we should not close it
|
|
||||||
var err error
|
|
||||||
if r.fd >= 0 {
|
|
||||||
err = unix.Close(r.fd)
|
|
||||||
r.fd = -1
|
|
||||||
}
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
@@ -1,205 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"sync/atomic"
|
|
||||||
"syscall"
|
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Maximum size we accept for a single read from a TUN with IFF_VNET_HDR. A
|
|
||||||
// TSO superpacket can be up to 64KiB of payload plus a single L2/L3/L4 header
|
|
||||||
// prefix plus the virtio header.
|
|
||||||
const tunReadBufSize = 65535
|
|
||||||
|
|
||||||
type Poll struct {
|
|
||||||
fd int
|
|
||||||
|
|
||||||
readPoll [2]unix.PollFd
|
|
||||||
writePoll [2]unix.PollFd
|
|
||||||
closed atomic.Bool
|
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
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,
|
|
||||||
readBuf: make([]byte, tunReadBufSize),
|
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
tunEvents := t.writePoll[0].Revents
|
|
||||||
shutdownEvents := t.writePoll[1].Revents
|
|
||||||
t.writePoll[0].Revents = 0
|
|
||||||
t.writePoll[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) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.readOne(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *Poll) readOne(to []byte) (int, error) {
|
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
var head [4]byte
|
|
||||||
iovecs := [2]syscall.Iovec{ //todo plat-specific
|
|
||||||
{&head[0], 4},
|
|
||||||
{&to[0], uint64(len(to))},
|
|
||||||
}
|
|
||||||
for {
|
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_READV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
|
||||||
if errno == 0 {
|
|
||||||
bytesRead := int(n)
|
|
||||||
if bytesRead < 4 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
|
||||||
return bytesRead - 4, 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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
|
||||||
func (t *Poll) Write(from []byte) (int, error) {
|
|
||||||
if len(from) <= 1 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
var head [4]byte
|
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
switch ipVer {
|
|
||||||
case 4:
|
|
||||||
head[3] = syscall.AF_INET
|
|
||||||
case 6:
|
|
||||||
head[3] = syscall.AF_INET6
|
|
||||||
default:
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
iovecs := [2]syscall.Iovec{ //todo plat specific
|
|
||||||
{&head[0], 4},
|
|
||||||
{&from[0], uint64(len(from))},
|
|
||||||
}
|
|
||||||
for {
|
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_WRITEV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
|
||||||
if errno == 0 {
|
|
||||||
return int(n) - 4, 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) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
|
||||||
@@ -1,86 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
// +build linux,!android,!e2e_testing
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"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 to newOffload / 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 TestOffload_WakeForShutdown_WakesFriends(t *testing.T) {
|
|
||||||
pipe1 := newReadPipe(t)
|
|
||||||
pipe2 := newReadPipe(t)
|
|
||||||
parent, err := NewOffloadContainer()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newOffload: %v", err)
|
|
||||||
}
|
|
||||||
require.NoError(t, parent.Add(pipe1))
|
|
||||||
require.NoError(t, parent.Add(pipe2))
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = unix.Close(pipe1)
|
|
||||||
_ = unix.Close(pipe2)
|
|
||||||
})
|
|
||||||
|
|
||||||
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()
|
|
||||||
_, errs[i] = r.Read()
|
|
||||||
}(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 TestTunFile_Close_Idempotent(t *testing.T) {
|
|
||||||
tf, err := newOffload(newReadPipe(t), 1)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newOffload: %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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,281 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
// +build linux,!android,!e2e_testing
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
|
||||||
const (
|
|
||||||
ipv4HeaderMinLen = 20 // IHL=5, no options
|
|
||||||
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
|
||||||
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
|
||||||
tcpHeaderMinLen = 20 // data-offset=5, no options
|
|
||||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside an IPv4 header.
|
|
||||||
const (
|
|
||||||
ipv4TotalLenOff = 2
|
|
||||||
ipv4IDOff = 4
|
|
||||||
ipv4ChecksumOff = 10
|
|
||||||
ipv4SrcOff = 12
|
|
||||||
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside an IPv6 header.
|
|
||||||
const (
|
|
||||||
ipv6PayloadLenOff = 4
|
|
||||||
ipv6SrcOff = 8
|
|
||||||
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
|
||||||
const (
|
|
||||||
tcpSeqOff = 4
|
|
||||||
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
|
||||||
tcpFlagsOff = 13
|
|
||||||
tcpChecksumOff = 16
|
|
||||||
)
|
|
||||||
|
|
||||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
|
||||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
|
||||||
|
|
||||||
// segmentInto splits a TUN-side packet described by hdr into one or more
|
|
||||||
// IP packets, each appended to *out as a slice of scratch. scratch must be
|
|
||||||
// sized to hold every segment (including replicated headers).
|
|
||||||
func segmentInto(pkt []byte, hdr VirtioNetHdr, out *[][]byte, scratch []byte) error {
|
|
||||||
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
|
||||||
// carry coalescing info rather than checksum offsets. A TUN writing via
|
|
||||||
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
|
||||||
// miscompute the segment checksums — refuse the packet instead.
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
|
||||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
|
||||||
}
|
|
||||||
|
|
||||||
switch hdr.GSOType {
|
|
||||||
case unix.VIRTIO_NET_HDR_GSO_NONE:
|
|
||||||
if len(pkt) > len(scratch) {
|
|
||||||
return fmt.Errorf("packet larger than segment buffer: %d > %d", len(pkt), len(scratch))
|
|
||||||
}
|
|
||||||
copy(scratch, pkt)
|
|
||||||
seg := scratch[:len(pkt)]
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
|
||||||
if err := finishChecksum(seg, hdr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
*out = append(*out, seg)
|
|
||||||
return nil
|
|
||||||
|
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
|
||||||
return segmentTCP(pkt, hdr, out, scratch)
|
|
||||||
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("unsupported virtio gso type: %d", hdr.GSOType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// finishChecksum computes the L4 checksum for a non-GSO packet that the kernel
|
|
||||||
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
|
|
||||||
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
|
||||||
// the pseudo-header partial sum by the kernel), and store the result.
|
|
||||||
func finishChecksum(seg []byte, hdr VirtioNetHdr) error {
|
|
||||||
cs := int(hdr.CsumStart)
|
|
||||||
co := int(hdr.CsumOffset)
|
|
||||||
if cs+co+2 > len(seg) {
|
|
||||||
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
|
||||||
}
|
|
||||||
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
|
||||||
// L4 region starting at cs, folding the prior partial in as the seed.
|
|
||||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
|
||||||
seg[cs+co] = 0
|
|
||||||
seg[cs+co+1] = 0
|
|
||||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// segmentTCP software-segments a TSO superpacket into one IP packet per MSS
|
|
||||||
// chunk. The caller guarantees hdr.GSOType is TCPV4 or TCPV6.
|
|
||||||
//
|
|
||||||
// Hot-path shape: the per-segment loop only sums the payload chunk. The TCP
|
|
||||||
// header, the IPv4 header, and the pseudo-header src/dst/proto contributions
|
|
||||||
// are each summed once up front — every segment reuses those three pre-folded
|
|
||||||
// uint32 values and combines them with small per-segment deltas (seq, flags,
|
|
||||||
// tcpLen, ip_id, total_len) that are cheap to fold in.
|
|
||||||
func segmentTCP(pkt []byte, hdr VirtioNetHdr, out *[][]byte, scratch []byte) error {
|
|
||||||
if hdr.GSOSize == 0 {
|
|
||||||
return fmt.Errorf("gso_size is zero")
|
|
||||||
}
|
|
||||||
if hdr.CsumStart == 0 {
|
|
||||||
return fmt.Errorf("csum_start is zero")
|
|
||||||
}
|
|
||||||
|
|
||||||
isV4 := hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_TCPV4
|
|
||||||
csumStart := int(hdr.CsumStart)
|
|
||||||
|
|
||||||
if isV4 && csumStart < ipv4HeaderMinLen {
|
|
||||||
return fmt.Errorf("csum_start %d too small for IPv4", csumStart)
|
|
||||||
}
|
|
||||||
if !isV4 && csumStart < ipv6FixedLen {
|
|
||||||
return fmt.Errorf("csum_start %d too small for IPv6", csumStart)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Don't trust hdr.HdrLen from the kernel: on some paths it can be set
|
|
||||||
// to the full length of the first packet rather than the true L3+L4 header length.
|
|
||||||
// Instead, read the TCP data-offset field from the packet itself and derive
|
|
||||||
// headerLen = csum_start + tcpHdrLen. Matches wireguard-go's approach.
|
|
||||||
if csumStart+tcpFlagsOff+1 > len(pkt) {
|
|
||||||
return fmt.Errorf("packet too short for tcp header at csum_start=%d (pkt %d)", csumStart, len(pkt))
|
|
||||||
}
|
|
||||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
|
||||||
if tcpHdrLen < tcpHeaderMinLen || tcpHdrLen > tcpHeaderMaxLen {
|
|
||||||
return fmt.Errorf("tcp data-offset out of range: %d", tcpHdrLen)
|
|
||||||
}
|
|
||||||
headerLen := csumStart + tcpHdrLen
|
|
||||||
if headerLen > len(pkt) {
|
|
||||||
return fmt.Errorf("derived hdr_len %d > pkt %d", headerLen, len(pkt))
|
|
||||||
}
|
|
||||||
|
|
||||||
payload := pkt[headerLen:]
|
|
||||||
payLen := len(payload)
|
|
||||||
gso := int(hdr.GSOSize)
|
|
||||||
numSeg := (payLen + gso - 1) / gso
|
|
||||||
if numSeg == 0 {
|
|
||||||
numSeg = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
need := numSeg*headerLen + payLen
|
|
||||||
if need > len(scratch) {
|
|
||||||
return fmt.Errorf("scratch too small for %d segments: need %d have %d", numSeg, need, len(scratch))
|
|
||||||
}
|
|
||||||
|
|
||||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
|
||||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
|
||||||
|
|
||||||
// Precompute the TCP header sum with seq/flags/csum zeroed. Copy onto
|
|
||||||
// the stack, zero the per-segment-varying fields, sum once.
|
|
||||||
var tmp [tcpHeaderMaxLen]byte
|
|
||||||
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
|
||||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
|
||||||
tmp[tcpFlagsOff] = 0
|
|
||||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
|
||||||
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
|
||||||
|
|
||||||
// Pseudo-header src+dst+proto contribution (tcpLen varies per segment).
|
|
||||||
var baseProtoSum uint32
|
|
||||||
if isV4 {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
|
||||||
} else {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
|
||||||
}
|
|
||||||
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
|
||||||
|
|
||||||
// Precompute IPv4 header sum with total_len/id/csum zeroed.
|
|
||||||
var origIPID uint16
|
|
||||||
var ihl int
|
|
||||||
var baseIPHdrSum uint32
|
|
||||||
if isV4 {
|
|
||||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
|
||||||
ihl = int(pkt[0]&0x0f) * 4
|
|
||||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
|
||||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
|
||||||
}
|
|
||||||
var ipTmp [ipv4HeaderMaxLen]byte
|
|
||||||
copy(ipTmp[:ihl], pkt[:ihl])
|
|
||||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
|
||||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
|
||||||
}
|
|
||||||
|
|
||||||
off := 0
|
|
||||||
for i := 0; i < numSeg; i++ {
|
|
||||||
segStart := i * gso
|
|
||||||
segEnd := segStart + gso
|
|
||||||
if segEnd > payLen {
|
|
||||||
segEnd = payLen
|
|
||||||
}
|
|
||||||
segPayLen := segEnd - segStart
|
|
||||||
|
|
||||||
copy(scratch[off:], pkt[:headerLen])
|
|
||||||
copy(scratch[off+headerLen:], payload[segStart:segEnd])
|
|
||||||
seg := scratch[off : off+headerLen+segPayLen]
|
|
||||||
off += headerLen + segPayLen
|
|
||||||
|
|
||||||
segSeq := origSeq + uint32(segStart)
|
|
||||||
segFlags := origFlags
|
|
||||||
if i != numSeg-1 {
|
|
||||||
segFlags = origFlags &^ tcpFinPshMask
|
|
||||||
}
|
|
||||||
totalLen := headerLen + segPayLen
|
|
||||||
|
|
||||||
// Patch IP header and write the v4 header checksum from the precomputed base.
|
|
||||||
if isV4 {
|
|
||||||
segID := origIPID + uint16(i)
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
|
||||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
|
||||||
} else {
|
|
||||||
// IPv6 payload length excludes the fixed header but includes any
|
|
||||||
// extension headers between [ipv6FixedLen:csumStart].
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Patch TCP header.
|
|
||||||
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
|
||||||
seg[csumStart+tcpFlagsOff] = segFlags
|
|
||||||
// (csum is written below; its prior contents in `seg` don't affect the
|
|
||||||
// computation since we never sum over the segment's own header.)
|
|
||||||
|
|
||||||
tcpLen := tcpHdrLen + segPayLen
|
|
||||||
paySum := uint32(checksum.Checksum(payload[segStart:segEnd], 0))
|
|
||||||
|
|
||||||
// Combine pre-folded uint32s into a wider accumulator, then fold. Using
|
|
||||||
// uint64 guards against overflow when segSeq's high bits set.
|
|
||||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
|
||||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
|
||||||
|
|
||||||
*out = append(*out, seg)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
|
||||||
// complements it, yielding the on-wire Internet checksum value.
|
|
||||||
func foldComplement(sum uint32) uint16 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
return ^uint16(sum)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoHeaderIPv4 returns the folded pseudo-header sum used to verify a TCP
|
|
||||||
// segment's checksum in tests. src/dst are 4 bytes each.
|
|
||||||
func pseudoHeaderIPv4(src, dst []byte, proto byte, tcpLen int) uint16 {
|
|
||||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
|
||||||
s += uint32(proto) + uint32(tcpLen)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
return uint16(s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoHeaderIPv6 returns the folded pseudo-header sum used to verify a TCP
|
|
||||||
// segment's checksum in tests. src/dst are 16 bytes each.
|
|
||||||
func pseudoHeaderIPv6(src, dst []byte, proto byte, tcpLen int) uint16 {
|
|
||||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
|
||||||
s += uint32(tcpLen>>16) + uint32(tcpLen&0xffff) + uint32(proto)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
return uint16(s)
|
|
||||||
}
|
|
||||||
@@ -1,330 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
// +build linux,!android,!e2e_testing
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
|
||||||
// with a folded pseudo-header sum, equals all-ones (valid).
|
|
||||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
|
||||||
return checksum.Checksum(b, pseudo) == 0xffff
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTSOv4 builds a synthetic IPv4/TCP TSO superpacket with a payload of
|
|
||||||
// `payLen` bytes split at `mss`.
|
|
||||||
func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, VirtioNetHdr) {
|
|
||||||
t.Helper()
|
|
||||||
const ipLen = 20
|
|
||||||
const tcpLen = 20
|
|
||||||
pkt := make([]byte, ipLen+tcpLen+payLen)
|
|
||||||
|
|
||||||
// IPv4 header
|
|
||||||
pkt[0] = 0x45 // version 4, IHL 5
|
|
||||||
// total length is meaningless for TSO but set it anyway
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // original ID
|
|
||||||
pkt[8] = 64 // TTL
|
|
||||||
pkt[9] = unix.IPPROTO_TCP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
|
|
||||||
|
|
||||||
// TCP header
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
|
|
||||||
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
|
|
||||||
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
|
|
||||||
pkt[32] = 0x50 // data offset 5 words
|
|
||||||
pkt[33] = 0x18 // ACK | PSH
|
|
||||||
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
|
|
||||||
|
|
||||||
// payload
|
|
||||||
for i := 0; i < payLen; i++ {
|
|
||||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
|
||||||
}
|
|
||||||
|
|
||||||
return pkt, VirtioNetHdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
|
||||||
GSOSize: uint16(mss),
|
|
||||||
CsumStart: uint16(ipLen),
|
|
||||||
CsumOffset: 16,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentTCPv4(t *testing.T) {
|
|
||||||
const mss = 100
|
|
||||||
const numSeg = 3
|
|
||||||
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
|
||||||
|
|
||||||
scratch := make([]byte, tunSegBufSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentTCP: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != numSeg {
|
|
||||||
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg) != 40+mss {
|
|
||||||
t.Errorf("seg %d: unexpected len %d", i, len(seg))
|
|
||||||
}
|
|
||||||
totalLen := binary.BigEndian.Uint16(seg[2:4])
|
|
||||||
if totalLen != uint16(40+mss) {
|
|
||||||
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 40+mss)
|
|
||||||
}
|
|
||||||
id := binary.BigEndian.Uint16(seg[4:6])
|
|
||||||
if id != 0x4242+uint16(i) {
|
|
||||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242+uint16(i))
|
|
||||||
}
|
|
||||||
seq := binary.BigEndian.Uint32(seg[24:28])
|
|
||||||
wantSeq := uint32(10000 + i*mss)
|
|
||||||
if seq != wantSeq {
|
|
||||||
t.Errorf("seg %d: seq=%d want %d", i, seq, wantSeq)
|
|
||||||
}
|
|
||||||
flags := seg[33]
|
|
||||||
wantFlags := byte(0x10) // ACK only, PSH cleared
|
|
||||||
if i == numSeg-1 {
|
|
||||||
wantFlags = 0x18 // ACK | PSH preserved on last
|
|
||||||
}
|
|
||||||
if flags != wantFlags {
|
|
||||||
t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags)
|
|
||||||
}
|
|
||||||
// IPv4 header checksum must verify against itself.
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
// TCP checksum must verify against the pseudo-header.
|
|
||||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+mss)
|
|
||||||
if !verifyChecksum(seg[20:], psum) {
|
|
||||||
t.Errorf("seg %d: bad TCP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentTCPv4OddTail(t *testing.T) {
|
|
||||||
// Payload of 250 bytes with MSS 100 → segments of 100, 100, 50.
|
|
||||||
pkt, hdr := buildTSOv4(t, 250, 100)
|
|
||||||
scratch := make([]byte, tunSegBufSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentTCP: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != 3 {
|
|
||||||
t.Fatalf("want 3 segments, got %d", len(out))
|
|
||||||
}
|
|
||||||
wantPayLens := []int{100, 100, 50}
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg)-40 != wantPayLens[i] {
|
|
||||||
t.Errorf("seg %d: pay len %d want %d", i, len(seg)-40, wantPayLens[i])
|
|
||||||
}
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+wantPayLens[i])
|
|
||||||
if !verifyChecksum(seg[20:], psum) {
|
|
||||||
t.Errorf("seg %d: bad TCP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentTCPv6(t *testing.T) {
|
|
||||||
const ipLen = 40
|
|
||||||
const tcpLen = 20
|
|
||||||
const mss = 120
|
|
||||||
const numSeg = 2
|
|
||||||
payLen := mss * numSeg
|
|
||||||
pkt := make([]byte, ipLen+tcpLen+payLen)
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
pkt[0] = 0x60 // version 6
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen))
|
|
||||||
pkt[6] = unix.IPPROTO_TCP
|
|
||||||
pkt[7] = 64
|
|
||||||
// src/dst fe80::1 / fe80::2
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
// TCP header
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 80)
|
|
||||||
binary.BigEndian.PutUint32(pkt[44:48], 7)
|
|
||||||
binary.BigEndian.PutUint32(pkt[48:52], 99)
|
|
||||||
pkt[52] = 0x50
|
|
||||||
pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too
|
|
||||||
binary.BigEndian.PutUint16(pkt[54:56], 65535)
|
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
|
||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
|
||||||
}
|
|
||||||
|
|
||||||
hdr := VirtioNetHdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
|
||||||
GSOSize: uint16(mss),
|
|
||||||
CsumStart: uint16(ipLen),
|
|
||||||
CsumOffset: 16,
|
|
||||||
}
|
|
||||||
|
|
||||||
scratch := make([]byte, tunSegBufSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentTCP: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != numSeg {
|
|
||||||
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg) != ipLen+tcpLen+mss {
|
|
||||||
t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+tcpLen+mss)
|
|
||||||
}
|
|
||||||
pl := binary.BigEndian.Uint16(seg[4:6])
|
|
||||||
if pl != uint16(tcpLen+mss) {
|
|
||||||
t.Errorf("seg %d: payload_length=%d want %d", i, pl, tcpLen+mss)
|
|
||||||
}
|
|
||||||
seq := binary.BigEndian.Uint32(seg[44:48])
|
|
||||||
if seq != uint32(7+i*mss) {
|
|
||||||
t.Errorf("seg %d: seq=%d want %d", i, seq, 7+i*mss)
|
|
||||||
}
|
|
||||||
flags := seg[53]
|
|
||||||
// Original flags = 0x19 (FIN|ACK|PSH). FIN(0x01)+PSH(0x08) should be
|
|
||||||
// cleared on all but the last; ACK(0x10) always preserved.
|
|
||||||
wantFlags := byte(0x10)
|
|
||||||
if i == numSeg-1 {
|
|
||||||
wantFlags = 0x19
|
|
||||||
}
|
|
||||||
if flags != wantFlags {
|
|
||||||
t.Errorf("seg %d: flags=%#x want %#x", i, flags, wantFlags)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpLen+mss)
|
|
||||||
if !verifyChecksum(seg[ipLen:], psum) {
|
|
||||||
t.Errorf("seg %d: bad TCP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentGSONonePassesThrough(t *testing.T) {
|
|
||||||
pkt, hdr := buildTSOv4(t, 100, 100)
|
|
||||||
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
|
||||||
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
|
||||||
|
|
||||||
scratch := make([]byte, tunSegBufSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentInto(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentInto: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != 1 {
|
|
||||||
t.Fatalf("want 1 segment, got %d", len(out))
|
|
||||||
}
|
|
||||||
if len(out[0]) != len(pkt) {
|
|
||||||
t.Fatalf("unexpected length: %d vs %d", len(out[0]), len(pkt))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentRejectsUDP(t *testing.T) {
|
|
||||||
hdr := VirtioNetHdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentInto(nil, hdr, &out, nil); err == nil {
|
|
||||||
t.Fatalf("expected rejection for UDP GSO")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkSegmentTCPv4(b *testing.B) {
|
|
||||||
sizes := []struct {
|
|
||||||
name string
|
|
||||||
payLen int
|
|
||||||
mss int
|
|
||||||
}{
|
|
||||||
{"64KiB_MSS1460", 65000, 1460},
|
|
||||||
{"16KiB_MSS1460", 16384, 1460},
|
|
||||||
{"4KiB_MSS1460", 4096, 1460},
|
|
||||||
}
|
|
||||||
for _, sz := range sizes {
|
|
||||||
b.Run(sz.name, func(b *testing.B) {
|
|
||||||
const ipLen = 20
|
|
||||||
const tcpLen = 20
|
|
||||||
pkt := make([]byte, ipLen+tcpLen+sz.payLen)
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+sz.payLen))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = unix.IPPROTO_TCP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], 12345)
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], 80)
|
|
||||||
binary.BigEndian.PutUint32(pkt[24:28], 10000)
|
|
||||||
binary.BigEndian.PutUint32(pkt[28:32], 20000)
|
|
||||||
pkt[32] = 0x50
|
|
||||||
pkt[33] = 0x18
|
|
||||||
binary.BigEndian.PutUint16(pkt[34:36], 65535)
|
|
||||||
for i := 0; i < sz.payLen; i++ {
|
|
||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
|
||||||
}
|
|
||||||
hdr := VirtioNetHdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
|
||||||
GSOSize: uint16(sz.mss),
|
|
||||||
CsumStart: uint16(ipLen),
|
|
||||||
CsumOffset: 16,
|
|
||||||
}
|
|
||||||
|
|
||||||
scratch := make([]byte, tunSegBufSize)
|
|
||||||
out := make([][]byte, 0, 64)
|
|
||||||
|
|
||||||
b.SetBytes(int64(len(pkt)))
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
out = out[:0]
|
|
||||||
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTunFileWriteVnetHdrNoAlloc verifies the IFF_VNET_HDR fast-path write is
|
|
||||||
// allocation-free. We write to /dev/null so every call succeeds synchronously.
|
|
||||||
func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
|
|
||||||
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("open /dev/null: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() { _ = unix.Close(fd) })
|
|
||||||
|
|
||||||
tf := &Offload{fd: fd}
|
|
||||||
tf.writeIovs[0].Base = &validVnetHdr[0]
|
|
||||||
tf.writeIovs[0].SetLen(virtioNetHdrLen)
|
|
||||||
|
|
||||||
payload := make([]byte, 1400)
|
|
||||||
// Warm up (first call may trigger one-time internal allocations elsewhere).
|
|
||||||
if _, err := tf.Write(payload); err != nil {
|
|
||||||
t.Fatalf("Write: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
allocs := testing.AllocsPerRun(1000, func() {
|
|
||||||
if _, err := tf.Write(payload); err != nil {
|
|
||||||
t.Fatalf("Write: %v", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
if allocs != 0 {
|
|
||||||
t.Fatalf("Write allocated %.1f times per call, want 0", allocs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
package tio
|
|
||||||
|
|
||||||
import "encoding/binary"
|
|
||||||
|
|
||||||
// Size of the legacy struct virtio_net_hdr that the kernel prepends/expects on
|
|
||||||
// a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ not set).
|
|
||||||
const virtioNetHdrLen = 10
|
|
||||||
|
|
||||||
type VirtioNetHdr struct {
|
|
||||||
Flags uint8
|
|
||||||
GSOType uint8
|
|
||||||
HdrLen uint16
|
|
||||||
GSOSize uint16
|
|
||||||
CsumStart uint16
|
|
||||||
CsumOffset uint16
|
|
||||||
}
|
|
||||||
|
|
||||||
// decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
|
||||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
|
||||||
func (h *VirtioNetHdr) decode(b []byte) {
|
|
||||||
h.Flags = b[0]
|
|
||||||
h.GSOType = b[1]
|
|
||||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
|
||||||
h.GSOSize = binary.NativeEndian.Uint16(b[4:6])
|
|
||||||
h.CsumStart = binary.NativeEndian.Uint16(b[6:8])
|
|
||||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
|
||||||
}
|
|
||||||
|
|
||||||
// encode is the inverse of decode: writes the virtio_net_hdr fields into b
|
|
||||||
// (must be at least virtioNetHdrLen bytes). Used to emit a TSO superpacket
|
|
||||||
// on egress.
|
|
||||||
func (h *VirtioNetHdr) encode(b []byte) {
|
|
||||||
b[0] = h.Flags
|
|
||||||
b[1] = h.GSOType
|
|
||||||
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
|
||||||
binary.NativeEndian.PutUint16(b[4:6], h.GSOSize)
|
|
||||||
binary.NativeEndian.PutUint16(b[6:8], h.CsumStart)
|
|
||||||
binary.NativeEndian.PutUint16(b[8:10], h.CsumOffset)
|
|
||||||
}
|
|
||||||
+6
-34
@@ -13,45 +13,17 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
rwc io.ReadWriteCloser
|
io.ReadWriteCloser
|
||||||
fd int
|
fd int
|
||||||
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 *logrus.Logger
|
l *logrus.Logger
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.rwc.Read(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Write(p []byte) (int, error) {
|
|
||||||
return t.rwc.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.rwc.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
|
||||||
return t.rwc.Close()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
@@ -60,10 +32,10 @@ func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []net
|
|||||||
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
rwc: file,
|
ReadWriteCloser: file,
|
||||||
fd: deviceFd,
|
fd: deviceFd,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -127,6 +99,6 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||||
}
|
}
|
||||||
|
|||||||
+12
-32
@@ -16,7 +16,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -24,7 +23,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
rwc io.ReadWriteCloser
|
io.ReadWriteCloser
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
DefaultMTU int
|
DefaultMTU int
|
||||||
@@ -35,9 +34,6 @@ type tun struct {
|
|||||||
|
|
||||||
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
||||||
out []byte
|
out []byte
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ifReq struct {
|
type ifReq struct {
|
||||||
@@ -128,11 +124,11 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
rwc: os.NewFile(uintptr(fd), ""),
|
ReadWriteCloser: 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)
|
||||||
@@ -162,8 +158,8 @@ func newTunFromFd(_ *config.C, _ *logrus.Logger, _ int, _ []netip.Prefix) (*tun,
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.rwc != nil {
|
if t.ReadWriteCloser != nil {
|
||||||
return t.rwc.Close()
|
return t.ReadWriteCloser.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -507,31 +503,15 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) readOne(to []byte) (int, error) {
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
buf := make([]byte, len(to)+4)
|
buf := make([]byte, len(to)+4)
|
||||||
|
|
||||||
n, err := t.rwc.Read(buf)
|
n, err := t.ReadWriteCloser.Read(buf)
|
||||||
|
|
||||||
copy(to, buf[4:])
|
copy(to, buf[4:])
|
||||||
return n - 4, err
|
return n - 4, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.readOne(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
||||||
@@ -557,7 +537,7 @@ func (t *tun) Write(from []byte) (int, error) {
|
|||||||
|
|
||||||
copy(buf[4:], from)
|
copy(buf[4:], from)
|
||||||
|
|
||||||
n, err := t.rwc.Write(buf)
|
n, err := t.ReadWriteCloser.Write(buf)
|
||||||
return n - 4, err
|
return n - 4, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -573,6 +553,6 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||||
}
|
}
|
||||||
|
|||||||
+23
-38
@@ -9,7 +9,6 @@ import (
|
|||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -18,27 +17,9 @@ 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 *logrus.Logger
|
l *logrus.Logger
|
||||||
numReaders int
|
|
||||||
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *disabledTun) Read() ([][]byte, error) {
|
|
||||||
r, ok := <-t.read
|
|
||||||
if !ok {
|
|
||||||
return nil, io.EOF
|
|
||||||
}
|
|
||||||
|
|
||||||
t.tx.Inc(1)
|
|
||||||
if t.l.Level >= logrus.DebugLevel {
|
|
||||||
t.l.WithField("raw", prettyPacket(r)).Debugf("Write payload")
|
|
||||||
}
|
|
||||||
|
|
||||||
t.batchRet[0] = r
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *logrus.Logger) *disabledTun {
|
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *logrus.Logger) *disabledTun {
|
||||||
@@ -46,7 +27,6 @@ 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 {
|
||||||
@@ -76,6 +56,24 @@ func (*disabledTun) Name() string {
|
|||||||
return "disabled"
|
return "disabled"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Read(b []byte) (int, error) {
|
||||||
|
r, ok := <-t.read
|
||||||
|
if !ok {
|
||||||
|
return 0, io.EOF
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(r) > len(b) {
|
||||||
|
return 0, fmt.Errorf("packet larger than mtu: %d > %d bytes", len(r), len(b))
|
||||||
|
}
|
||||||
|
|
||||||
|
t.tx.Inc(1)
|
||||||
|
if t.l.Level >= logrus.DebugLevel {
|
||||||
|
t.l.WithField("raw", prettyPacket(r)).Debugf("Write payload")
|
||||||
|
}
|
||||||
|
|
||||||
|
return copy(b, r), 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)
|
||||||
@@ -107,25 +105,12 @@ func (t *disabledTun) Write(b []byte) (int, error) {
|
|||||||
return len(b), nil
|
return len(b), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) WriteReject(b []byte) (int, error) {
|
|
||||||
return t.Write(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *disabledTun) SupportsMultiqueue() bool {
|
func (t *disabledTun) SupportsMultiqueue() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) NewMultiQueueReader() error {
|
func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
t.numReaders++
|
return t, nil
|
||||||
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) Close() error {
|
func (t *disabledTun) Close() error {
|
||||||
|
|||||||
+78
-206
@@ -7,9 +7,9 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
@@ -18,7 +18,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -94,203 +93,107 @@ type tun struct {
|
|||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
linkAddr *netroute.LinkAddr
|
linkAddr *netroute.LinkAddr
|
||||||
l *logrus.Logger
|
l *logrus.Logger
|
||||||
|
devFd int
|
||||||
fd int
|
|
||||||
shutdownR int // read end of the shutdown pipe; closing the write end wakes blocked polls
|
|
||||||
shutdownW int // write end of the shutdown pipe; closing this signals shutdown to any blocked reader/writer
|
|
||||||
readPoll [2]unix.PollFd
|
|
||||||
writePoll [2]unix.PollFd
|
|
||||||
closed atomic.Bool
|
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// blockOnRead waits until the tun fd is readable or shutdown has been signaled.
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
// Returns os.ErrClosed if Close was called.
|
// use readv() to read from the tunnel device, to eliminate the need for copying the buffer
|
||||||
func (t *tun) blockOnRead() error {
|
if t.devFd < 0 {
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
return -1, syscall.EINVAL
|
||||||
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 *tun) 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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
tunEvents := t.writePoll[0].Revents
|
|
||||||
shutdownEvents := t.writePoll[1].Revents
|
|
||||||
t.writePoll[0].Revents = 0
|
|
||||||
t.writePoll[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 *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.readOne(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
head := make([]byte, 4)
|
||||||
iovecs := [2]syscall.Iovec{
|
|
||||||
|
iovecs := []syscall.Iovec{
|
||||||
{&head[0], 4},
|
{&head[0], 4},
|
||||||
{&to[0], uint64(len(to))},
|
{&to[0], uint64(len(to))},
|
||||||
}
|
}
|
||||||
for {
|
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_READV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
n, _, errno := syscall.Syscall(syscall.SYS_READV, uintptr(t.devFd), uintptr(unsafe.Pointer(&iovecs[0])), uintptr(2))
|
||||||
if errno == 0 {
|
|
||||||
bytesRead := int(n)
|
var err error
|
||||||
if bytesRead < 4 {
|
if errno != 0 {
|
||||||
return 0, nil
|
err = syscall.Errno(errno)
|
||||||
}
|
} else {
|
||||||
return bytesRead - 4, nil
|
err = nil
|
||||||
}
|
}
|
||||||
switch errno {
|
// fix bytes read number to exclude header
|
||||||
case unix.EAGAIN:
|
bytesRead := int(n)
|
||||||
if err := t.blockOnRead(); err != nil {
|
if bytesRead < 0 {
|
||||||
return 0, err
|
return bytesRead, err
|
||||||
}
|
} else if bytesRead < 4 {
|
||||||
case unix.EINTR:
|
return 0, err
|
||||||
// retry
|
} else {
|
||||||
case unix.EBADF:
|
return bytesRead - 4, err
|
||||||
return 0, os.ErrClosed
|
|
||||||
default:
|
|
||||||
return 0, errno
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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) {
|
||||||
|
// use writev() to write to the tunnel device, to eliminate the need for copying the buffer
|
||||||
|
if t.devFd < 0 {
|
||||||
|
return -1, syscall.EINVAL
|
||||||
|
}
|
||||||
|
|
||||||
if len(from) <= 1 {
|
if len(from) <= 1 {
|
||||||
return 0, syscall.EIO
|
return 0, syscall.EIO
|
||||||
}
|
}
|
||||||
|
|
||||||
ipVer := from[0] >> 4
|
ipVer := from[0] >> 4
|
||||||
var head [4]byte
|
var head []byte
|
||||||
// first 4 bytes is protocol family, in network byte order
|
// first 4 bytes is protocol family, in network byte order
|
||||||
switch ipVer {
|
if ipVer == 4 {
|
||||||
case 4:
|
head = []byte{0, 0, 0, syscall.AF_INET}
|
||||||
head[3] = syscall.AF_INET
|
} else if ipVer == 6 {
|
||||||
case 6:
|
head = []byte{0, 0, 0, syscall.AF_INET6}
|
||||||
head[3] = syscall.AF_INET6
|
} else {
|
||||||
default:
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
return 0, fmt.Errorf("unable to determine IP version from packet")
|
||||||
}
|
}
|
||||||
|
iovecs := []syscall.Iovec{
|
||||||
iovecs := [2]syscall.Iovec{
|
|
||||||
{&head[0], 4},
|
{&head[0], 4},
|
||||||
{&from[0], uint64(len(from))},
|
{&from[0], uint64(len(from))},
|
||||||
}
|
}
|
||||||
for {
|
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_WRITEV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
n, _, errno := syscall.Syscall(syscall.SYS_WRITEV, uintptr(t.devFd), uintptr(unsafe.Pointer(&iovecs[0])), uintptr(2))
|
||||||
if errno == 0 {
|
|
||||||
return int(n) - 4, nil
|
var err error
|
||||||
}
|
if errno != 0 {
|
||||||
switch errno {
|
err = syscall.Errno(errno)
|
||||||
case unix.EAGAIN:
|
} else {
|
||||||
if err := t.blockOnWrite(); err != nil {
|
err = nil
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
case unix.EINTR:
|
|
||||||
// retry
|
|
||||||
case unix.EBADF:
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
default:
|
|
||||||
return 0, errno
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return int(n) - 4, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.closed.Swap(true) {
|
if t.devFd >= 0 {
|
||||||
return nil
|
err := syscall.Close(t.devFd)
|
||||||
}
|
if err != nil {
|
||||||
|
|
||||||
// Closing the write end of the shutdown pipe causes any blocked Poll to
|
|
||||||
// return with POLLHUP on the shutdown fd, so readers/writers wake up and
|
|
||||||
// exit with os.ErrClosed.
|
|
||||||
if t.shutdownW >= 0 {
|
|
||||||
_ = unix.Close(t.shutdownW)
|
|
||||||
t.shutdownW = -1
|
|
||||||
}
|
|
||||||
|
|
||||||
if t.fd >= 0 {
|
|
||||||
if err := unix.Close(t.fd); err != nil {
|
|
||||||
t.l.WithError(err).Error("Error closing device")
|
t.l.WithError(err).Error("Error closing device")
|
||||||
}
|
}
|
||||||
t.fd = -1
|
t.devFd = -1
|
||||||
}
|
|
||||||
|
|
||||||
if t.shutdownR >= 0 {
|
c := make(chan struct{})
|
||||||
_ = unix.Close(t.shutdownR)
|
go func() {
|
||||||
t.shutdownR = -1
|
// destroying the interface can block if a read() is still pending. Do this asynchronously.
|
||||||
}
|
defer close(c)
|
||||||
|
s, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_DGRAM, syscall.IPPROTO_IP)
|
||||||
|
if err == nil {
|
||||||
|
defer syscall.Close(s)
|
||||||
|
ifreq := ifreqDestroy{Name: t.deviceBytes()}
|
||||||
|
err = ioctl(uintptr(s), syscall.SIOCIFDESTROY, uintptr(unsafe.Pointer(&ifreq)))
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.l.WithError(err).Error("Error destroying tunnel")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
c := make(chan struct{})
|
// wait up to 1 second so we start blocking at the ioctl
|
||||||
go func() {
|
select {
|
||||||
// destroying the interface can block if a read() is still pending. Do this asynchronously.
|
case <-c:
|
||||||
defer close(c)
|
case <-time.After(1 * time.Second):
|
||||||
s, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_DGRAM, syscall.IPPROTO_IP)
|
|
||||||
if err == nil {
|
|
||||||
defer syscall.Close(s)
|
|
||||||
ifreq := ifreqDestroy{Name: t.deviceBytes()}
|
|
||||||
err = ioctl(uintptr(s), syscall.SIOCIFDESTROY, uintptr(unsafe.Pointer(&ifreq)))
|
|
||||||
}
|
}
|
||||||
if err != nil {
|
|
||||||
t.l.WithError(err).Error("Error destroying tunnel")
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// wait up to 1 second so we start blocking at the ioctl
|
|
||||||
select {
|
|
||||||
case <-c:
|
|
||||||
case <-time.After(1 * time.Second):
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -306,38 +209,16 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
var err error
|
var err error
|
||||||
deviceName := c.GetString("tun.dev", "")
|
deviceName := c.GetString("tun.dev", "")
|
||||||
if deviceName != "" {
|
if deviceName != "" {
|
||||||
fd, err = unix.Open("/dev/"+deviceName, os.O_RDWR, 0)
|
fd, err = syscall.Open("/dev/"+deviceName, syscall.O_RDWR, 0)
|
||||||
}
|
}
|
||||||
if errors.Is(err, fs.ErrNotExist) || deviceName == "" {
|
if errors.Is(err, fs.ErrNotExist) || deviceName == "" {
|
||||||
// If the device doesn't already exist, request a new one and rename it
|
// If the device doesn't already exist, request a new one and rename it
|
||||||
fd, err = unix.Open("/dev/tun", os.O_RDWR, 0)
|
fd, err = syscall.Open("/dev/tun", syscall.O_RDWR, 0)
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = unix.SetNonblock(fd, true); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, fmt.Errorf("failed to set tun device as nonblocking: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Shutdown pipe lets Close wake any reader/writer blocked in Poll.
|
|
||||||
var pipeFds [2]int
|
|
||||||
if err = unix.Pipe2(pipeFds[:], unix.O_CLOEXEC|unix.O_NONBLOCK); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, fmt.Errorf("failed to create shutdown pipe: %w", err)
|
|
||||||
}
|
|
||||||
shutdownR, shutdownW := pipeFds[0], pipeFds[1]
|
|
||||||
|
|
||||||
closeOnErr := true
|
|
||||||
defer func() {
|
|
||||||
if closeOnErr {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
_ = unix.Close(shutdownR)
|
|
||||||
_ = unix.Close(shutdownW)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Read the name of the interface
|
// Read the name of the interface
|
||||||
var name [16]byte
|
var name [16]byte
|
||||||
arg := fiodgnameArg{length: 16, buf: unsafe.Pointer(&name)}
|
arg := fiodgnameArg{length: 16, buf: unsafe.Pointer(&name)}
|
||||||
@@ -356,7 +237,7 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
}
|
}
|
||||||
|
|
||||||
if ctrlErr != nil {
|
if ctrlErr != nil {
|
||||||
return nil, ctrlErr
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
ifName := string(bytes.TrimRight(name[:], "\x00"))
|
ifName := string(bytes.TrimRight(name[:], "\x00"))
|
||||||
@@ -372,6 +253,8 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
}
|
}
|
||||||
defer syscall.Close(s)
|
defer syscall.Close(s)
|
||||||
|
|
||||||
|
fd := uintptr(s)
|
||||||
|
|
||||||
var fromName [16]byte
|
var fromName [16]byte
|
||||||
var toName [16]byte
|
var toName [16]byte
|
||||||
copy(fromName[:], ifName)
|
copy(fromName[:], ifName)
|
||||||
@@ -383,7 +266,7 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Set the device name
|
// Set the device name
|
||||||
_ = ioctl(uintptr(s), syscall.SIOCSIFNAME, uintptr(unsafe.Pointer(&ifrr)))
|
_ = ioctl(fd, syscall.SIOCSIFNAME, uintptr(unsafe.Pointer(&ifrr)))
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
@@ -391,24 +274,13 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (
|
|||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
MTU: c.GetInt("tun.mtu", DefaultMTU),
|
MTU: c.GetInt("tun.mtu", DefaultMTU),
|
||||||
l: l,
|
l: l,
|
||||||
fd: fd,
|
devFd: fd,
|
||||||
shutdownR: shutdownR,
|
|
||||||
shutdownW: shutdownW,
|
|
||||||
readPoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
|
||||||
{Fd: int32(shutdownR), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
writePoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
|
||||||
{Fd: int32(shutdownR), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = t.reload(c, true)
|
err = t.reload(c, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
closeOnErr = false
|
|
||||||
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
err := t.reload(c, false)
|
err := t.reload(c, false)
|
||||||
@@ -582,7 +454,7 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+5
-33
@@ -16,44 +16,16 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
rwc io.ReadWriteCloser
|
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 *logrus.Logger
|
l *logrus.Logger
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.rwc.Read(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Write(p []byte) (int, error) {
|
|
||||||
return t.rwc.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.rwc.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
|
||||||
return t.rwc.Close()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTun(_ *config.C, _ *logrus.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
func newTun(_ *config.C, _ *logrus.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
||||||
@@ -63,9 +35,9 @@ func newTun(_ *config.C, _ *logrus.Logger, _ []netip.Prefix, _ bool) (*tun, erro
|
|||||||
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *logrus.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,
|
||||||
rwc: &tunReadCloser{f: file},
|
ReadWriteCloser: &tunReadCloser{f: file},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -183,6 +155,6 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||||
}
|
}
|
||||||
|
|||||||
+82
-137
@@ -5,6 +5,7 @@ package overlay
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,7 +18,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -25,8 +25,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
readers tio.Container
|
io.ReadWriteCloser
|
||||||
closeLock sync.Mutex
|
fd int
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
MaxMTU int
|
MaxMTU int
|
||||||
@@ -34,7 +34,6 @@ type tun struct {
|
|||||||
TXQueueLen int
|
TXQueueLen int
|
||||||
deviceIndex int
|
deviceIndex int
|
||||||
ioctlFd uintptr
|
ioctlFd uintptr
|
||||||
vnetHdr bool
|
|
||||||
|
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
@@ -73,9 +72,9 @@ type ifreqQLEN struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
// We don't know what flags the caller opened this fd with and can't turn
|
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
||||||
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
|
||||||
t, err := newTunGeneric(c, l, deviceFd, false, vpnNetworks)
|
t, err := newTunGeneric(c, l, file, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -85,83 +84,46 @@ func newTunFromFd(c *config.C, l *logrus.Logger, deviceFd int, vpnNetworks []net
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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)
|
|
||||||
if err == nil {
|
|
||||||
return fd, nil
|
|
||||||
}
|
|
||||||
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
|
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// tsoOffloadFlags are the TUN_F_* bits we ask the kernel to enable when a
|
|
||||||
// TSO-capable TUN is available. CSUM is required as a prerequisite for TSO.
|
|
||||||
const tsoOffloadFlags = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6
|
|
||||||
|
|
||||||
func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||||
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if multiqueue {
|
|
||||||
baseFlags |= unix.IFF_MULTI_QUEUE
|
|
||||||
}
|
|
||||||
nameStr := c.GetString("tun.dev", "")
|
|
||||||
|
|
||||||
// First try to open with IFF_VNET_HDR + TUNSETOFFLOAD so we can receive
|
|
||||||
// TSO superpackets. If either step fails (older kernel, unprivileged
|
|
||||||
// container, etc.) we close and fall back to a plain TUN.
|
|
||||||
fd, err := openTunDev()
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
||||||
}
|
if os.IsNotExist(err) {
|
||||||
vnetHdr := true
|
err = os.MkdirAll("/dev/net", 0755)
|
||||||
name, err := tunSetIff(fd, nameStr, baseFlags|unix.IFF_VNET_HDR|unix.IFF_NAPI)
|
if err != nil {
|
||||||
if err != nil {
|
return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||||
_ = unix.Close(fd)
|
}
|
||||||
vnetHdr = false
|
err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200)))
|
||||||
} else if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err != nil {
|
if err != nil {
|
||||||
l.WithError(err).Warn("Failed to enable TUN offload (TSO); proceeding without virtio headers")
|
return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||||
_ = unix.Close(fd)
|
}
|
||||||
vnetHdr = false
|
|
||||||
}
|
|
||||||
|
|
||||||
if !vnetHdr {
|
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
fd, err = openTunDev()
|
if err != nil {
|
||||||
if err != nil {
|
return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
return nil, err
|
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, vnetHdr, vpnNetworks)
|
var req ifReq
|
||||||
|
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||||
|
if multiqueue {
|
||||||
|
req.Flags |= unix.IFF_MULTI_QUEUE
|
||||||
|
}
|
||||||
|
nameStr := c.GetString("tun.dev", "")
|
||||||
|
copy(req.Name[:], nameStr)
|
||||||
|
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||||
|
return nil, &NameError{
|
||||||
|
Name: nameStr,
|
||||||
|
Underlying: err,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
name := strings.Trim(string(req.Name[:]), "\x00")
|
||||||
|
|
||||||
|
file := os.NewFile(uintptr(fd), "/dev/net/tun")
|
||||||
|
t, err := newTunGeneric(c, l, file, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -171,30 +133,10 @@ func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, multiqueu
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
func newTunGeneric(c *config.C, l *logrus.Logger, file *os.File, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
func newTunGeneric(c *config.C, l *logrus.Logger, fd int, vnetHdr bool, vpnNetworks []netip.Prefix) (*tun, error) {
|
|
||||||
var container tio.Container
|
|
||||||
var err error
|
|
||||||
if vnetHdr {
|
|
||||||
container, err = tio.NewOffloadContainer()
|
|
||||||
} else {
|
|
||||||
container, err = tio.NewPollContainer()
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
err = container.Add(fd)
|
|
||||||
if err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
readers: container,
|
ReadWriteCloser: file,
|
||||||
closeLock: sync.Mutex{},
|
fd: int(file.Fd()),
|
||||||
vnetHdr: vnetHdr,
|
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||||
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
useSystemRoutes: c.GetBool("tun.use_system_route_table", false),
|
||||||
@@ -203,8 +145,8 @@ func newTunGeneric(c *config.C, l *logrus.Logger, fd int, vnetHdr bool, vpnNetwo
|
|||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = t.reload(c, true); err != nil {
|
err := t.reload(c, true)
|
||||||
_ = t.Close()
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -296,38 +238,22 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() error {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
t.closeLock.Lock()
|
|
||||||
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 err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
var req ifReq
|
||||||
if t.vnetHdr {
|
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||||
flags |= unix.IFF_VNET_HDR | unix.IFF_NAPI
|
copy(req.Name[:], t.Device)
|
||||||
}
|
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||||
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
return nil, err
|
||||||
_ = unix.Close(fd)
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if t.vnetHdr {
|
file := os.NewFile(uintptr(fd), "/dev/net/tun")
|
||||||
if err = ioctl(uintptr(fd), unix.TUNSETOFFLOAD, uintptr(tsoOffloadFlags)); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return fmt.Errorf("failed to enable offload on multiqueue tun fd: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = t.readers.Add(fd)
|
return file, nil
|
||||||
if err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||||
@@ -335,6 +261,29 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Write(b []byte) (int, error) {
|
||||||
|
var nn int
|
||||||
|
maximum := len(b)
|
||||||
|
|
||||||
|
for {
|
||||||
|
n, err := unix.Write(t.fd, b[nn:maximum])
|
||||||
|
if n > 0 {
|
||||||
|
nn += n
|
||||||
|
}
|
||||||
|
if nn == len(b) {
|
||||||
|
return nn, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nn, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if n == 0 {
|
||||||
|
return nn, io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *tun) deviceBytes() (o [16]byte) {
|
func (t *tun) deviceBytes() (o [16]byte) {
|
||||||
for i, c := range t.Device {
|
for i, c := range t.Device {
|
||||||
o[i] = byte(c)
|
o[i] = byte(c)
|
||||||
@@ -757,23 +706,19 @@ 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()
|
|
||||||
defer t.closeLock.Unlock()
|
|
||||||
|
|
||||||
if t.routeChan != nil {
|
if t.routeChan != nil {
|
||||||
close(t.routeChan)
|
close(t.routeChan)
|
||||||
t.routeChan = nil
|
}
|
||||||
|
|
||||||
|
if t.ReadWriteCloser != nil {
|
||||||
|
_ = t.ReadWriteCloser.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
if t.ioctlFd > 0 {
|
if t.ioctlFd > 0 {
|
||||||
_ = unix.Close(int(t.ioctlFd))
|
_ = os.NewFile(t.ioctlFd, "ioctlFd").Close()
|
||||||
t.ioctlFd = 0
|
t.ioctlFd = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
return t.readers.Close()
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,9 +3,7 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import "testing"
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
+3
-22
@@ -6,6 +6,7 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"regexp"
|
"regexp"
|
||||||
@@ -16,7 +17,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -66,25 +66,6 @@ type tun struct {
|
|||||||
l *logrus.Logger
|
l *logrus.Logger
|
||||||
f *os.File
|
f *os.File
|
||||||
fd int
|
fd int
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.readOne(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
@@ -160,7 +141,7 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) readOne(to []byte) (int, error) {
|
func (t *tun) Read(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)
|
||||||
@@ -413,7 +394,7 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+3
-22
@@ -6,6 +6,7 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"regexp"
|
"regexp"
|
||||||
@@ -16,7 +17,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -59,25 +59,6 @@ type tun struct {
|
|||||||
fd int
|
fd int
|
||||||
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
||||||
out []byte
|
out []byte
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.readOne(t.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
@@ -143,7 +124,7 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) readOne(to []byte) (int, error) {
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
buf := make([]byte, len(to)+4)
|
buf := make([]byte, len(to)+4)
|
||||||
|
|
||||||
n, err := t.f.Read(buf)
|
n, err := t.f.Read(buf)
|
||||||
@@ -333,7 +314,7 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+10
-17
@@ -13,7 +13,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -27,17 +26,6 @@ type TestTun struct {
|
|||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
rxPackets chan []byte // Packets to receive into nebula
|
rxPackets chan []byte // Packets to receive into nebula
|
||||||
TxPackets chan []byte // Packets transmitted outside by nebula
|
TxPackets chan []byte // Packets transmitted outside by nebula
|
||||||
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *TestTun) Read() ([][]byte, error) {
|
|
||||||
p, ok := <-t.rxPackets
|
|
||||||
if !ok {
|
|
||||||
return nil, os.ErrClosed
|
|
||||||
}
|
|
||||||
t.batchRet[0] = p
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (*TestTun, error) {
|
func newTun(c *config.C, l *logrus.Logger, vpnNetworks []netip.Prefix, _ bool) (*TestTun, error) {
|
||||||
@@ -127,10 +115,6 @@ func (t *TestTun) Write(b []byte) (n int, err error) {
|
|||||||
return len(b), nil
|
return len(b), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) WriteReject(b []byte) (int, error) {
|
|
||||||
return t.Write(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *TestTun) Close() error {
|
func (t *TestTun) Close() error {
|
||||||
if t.closed.CompareAndSwap(false, true) {
|
if t.closed.CompareAndSwap(false, true) {
|
||||||
close(t.rxPackets)
|
close(t.rxPackets)
|
||||||
@@ -139,10 +123,19 @@ func (t *TestTun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) Read(b []byte) (int, error) {
|
||||||
|
p, ok := <-t.rxPackets
|
||||||
|
if !ok {
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
}
|
||||||
|
copy(b, p)
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
func (t *TestTun) SupportsMultiqueue() bool {
|
func (t *TestTun) SupportsMultiqueue() bool {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
||||||
}
|
}
|
||||||
|
|||||||
+6
-21
@@ -6,6 +6,7 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"crypto"
|
"crypto"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -17,7 +18,6 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
@@ -36,25 +36,6 @@ type winTun struct {
|
|||||||
l *logrus.Logger
|
l *logrus.Logger
|
||||||
|
|
||||||
tun *wintun.NativeTun
|
tun *wintun.NativeTun
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *winTun) Read() ([][]byte, error) {
|
|
||||||
if t.readBuf == nil {
|
|
||||||
t.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := t.tun.Read(t.readBuf, 0)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
t.batchRet[0] = t.readBuf[:n]
|
|
||||||
return t.batchRet[:], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *winTun) WriteReject(p []byte) (int, error) {
|
|
||||||
return t.Write(p)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *logrus.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
func newTunFromFd(_ *config.C, _ *logrus.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||||
@@ -248,6 +229,10 @@ 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)
|
||||||
}
|
}
|
||||||
@@ -256,7 +241,7 @@ func (t *winTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) NewMultiQueueReader() (tio.Queue, error) {
|
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+5
-32
@@ -6,7 +6,6 @@ import (
|
|||||||
|
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -24,34 +23,17 @@ 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
|
||||||
|
|
||||||
inboundReader *io.PipeReader
|
inboundReader *io.PipeReader
|
||||||
inboundWriter *io.PipeWriter
|
inboundWriter *io.PipeWriter
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
batchRet [1][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *UserDevice) Read() ([][]byte, error) {
|
|
||||||
if d.readBuf == nil {
|
|
||||||
d.readBuf = make([]byte, defaultBatchBufSize)
|
|
||||||
}
|
|
||||||
n, err := d.outboundReader.Read(d.readBuf)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
d.batchRet[0] = d.readBuf[:n]
|
|
||||||
return d.batchRet[:], nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Activate() error {
|
func (d *UserDevice) Activate() error {
|
||||||
@@ -68,29 +50,20 @@ func (d *UserDevice) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) NewMultiQueueReader() error {
|
func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
d.numReaders++
|
return d, nil
|
||||||
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) WriteReject(p []byte) (n int, err error) {
|
|
||||||
return d.Write(p)
|
|
||||||
}
|
|
||||||
func (d *UserDevice) Close() error {
|
func (d *UserDevice) Close() error {
|
||||||
d.inboundWriter.Close()
|
d.inboundWriter.Close()
|
||||||
d.outboundWriter.Close()
|
d.outboundWriter.Close()
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -488,25 +487,25 @@ func loadCertificate(b []byte) (cert.Certificate, []byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func loadCAPoolFromConfig(l *logrus.Logger, c *config.C) (*cert.CAPool, error) {
|
func loadCAPoolFromConfig(l *logrus.Logger, c *config.C) (*cert.CAPool, error) {
|
||||||
|
var rawCA []byte
|
||||||
|
var err error
|
||||||
|
|
||||||
caPathOrPEM := c.GetString("pki.ca", "")
|
caPathOrPEM := c.GetString("pki.ca", "")
|
||||||
if caPathOrPEM == "" {
|
if caPathOrPEM == "" {
|
||||||
return nil, errors.New("no pki.ca path or PEM data provided")
|
return nil, errors.New("no pki.ca path or PEM data provided")
|
||||||
}
|
}
|
||||||
|
|
||||||
var caReader io.ReadCloser
|
|
||||||
var err error
|
|
||||||
|
|
||||||
if strings.Contains(caPathOrPEM, "-----BEGIN") {
|
if strings.Contains(caPathOrPEM, "-----BEGIN") {
|
||||||
caReader = io.NopCloser(strings.NewReader(caPathOrPEM))
|
rawCA = []byte(caPathOrPEM)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
caReader, err = os.Open(caPathOrPEM)
|
rawCA, err = os.ReadFile(caPathOrPEM)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("unable to read pki.ca file %s: %s", caPathOrPEM, err)
|
return nil, fmt.Errorf("unable to read pki.ca file %s: %s", caPathOrPEM, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
defer caReader.Close()
|
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caReader)
|
caPool, err := cert.NewCAPoolFromPEM(rawCA)
|
||||||
if errors.Is(err, cert.ErrExpired) {
|
if errors.Is(err, cert.ErrExpired) {
|
||||||
var expired int
|
var expired int
|
||||||
for _, crt := range caPool.CAs {
|
for _, crt := range caPool.CAs {
|
||||||
|
|||||||
@@ -1,121 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
cert_test "github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/config"
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func BenchmarkReloadConfigWithCAs(b *testing.B) {
|
|
||||||
prevProcs := runtime.GOMAXPROCS(1)
|
|
||||||
b.Cleanup(func() { runtime.GOMAXPROCS(prevProcs) })
|
|
||||||
|
|
||||||
for _, size := range []int{100, 250, 500, 1000, 5000} {
|
|
||||||
b.Run(fmt.Sprintf("%dCAs", size), func(b *testing.B) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
dir := b.TempDir()
|
|
||||||
|
|
||||||
ca, caKey, caBundle := buildCABundle(b, size)
|
|
||||||
caPath, certPath, keyPath := writePKIFiles(b, dir, ca, caKey, caBundle)
|
|
||||||
|
|
||||||
configBody := fmt.Sprintf(`pki:
|
|
||||||
ca: %s
|
|
||||||
cert: %s
|
|
||||||
key: %s
|
|
||||||
`, caPath, certPath, keyPath)
|
|
||||||
|
|
||||||
configPath := filepath.Join(dir, "config.yml")
|
|
||||||
require.NoError(b, os.WriteFile(configPath, []byte(configBody), 0o600))
|
|
||||||
|
|
||||||
c := config.NewC(l)
|
|
||||||
require.NoError(b, c.Load(dir))
|
|
||||||
|
|
||||||
_, err := NewPKIFromConfig(l, c)
|
|
||||||
require.NoError(b, err)
|
|
||||||
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
|
|
||||||
for b.Loop() {
|
|
||||||
c.ReloadConfig()
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func buildCABundle(b *testing.B, count int) (cert.Certificate, []byte, []byte) {
|
|
||||||
b.Helper()
|
|
||||||
require.GreaterOrEqual(b, count, 1)
|
|
||||||
|
|
||||||
before := time.Now().Add(-24 * time.Hour)
|
|
||||||
after := time.Now().Add(24 * time.Hour)
|
|
||||||
|
|
||||||
ca, _, caKey, pem := cert_test.NewTestCaCert(
|
|
||||||
cert.Version2,
|
|
||||||
cert.Curve_CURVE25519,
|
|
||||||
before,
|
|
||||||
after,
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
|
|
||||||
buf := bytes.NewBuffer(pem)
|
|
||||||
buf.Write([]byte("\n# a comment!\n"))
|
|
||||||
|
|
||||||
for i := 1; i < count; i++ {
|
|
||||||
_, _, _, extraPEM := cert_test.NewTestCaCert(
|
|
||||||
cert.Version2,
|
|
||||||
cert.Curve_CURVE25519,
|
|
||||||
time.Now(),
|
|
||||||
time.Now().Add(time.Hour),
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
buf.Write([]byte("\n# a comment!\n"))
|
|
||||||
buf.Write(extraPEM)
|
|
||||||
}
|
|
||||||
|
|
||||||
return ca, caKey, buf.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
func writePKIFiles(b *testing.B, dir string, ca cert.Certificate, caKey []byte, caBundle []byte) (string, string, string) {
|
|
||||||
b.Helper()
|
|
||||||
|
|
||||||
networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}
|
|
||||||
|
|
||||||
_, _, keyPEM, certPEM := cert_test.NewTestCert(
|
|
||||||
cert.Version2,
|
|
||||||
cert.Curve_CURVE25519,
|
|
||||||
ca,
|
|
||||||
caKey,
|
|
||||||
"reload-benchmark",
|
|
||||||
time.Now(),
|
|
||||||
time.Now().Add(time.Hour),
|
|
||||||
networks,
|
|
||||||
nil,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
|
|
||||||
caPath := filepath.Join(dir, "ca.pem")
|
|
||||||
certPath := filepath.Join(dir, "cert.pem")
|
|
||||||
keyPath := filepath.Join(dir, "key.pem")
|
|
||||||
|
|
||||||
require.NoError(b, os.WriteFile(caPath, caBundle, 0o600))
|
|
||||||
require.NoError(b, os.WriteFile(certPath, certPEM, 0o600))
|
|
||||||
require.NoError(b, os.WriteFile(keyPath, keyPEM, 0o600))
|
|
||||||
|
|
||||||
return caPath, certPath, keyPath
|
|
||||||
}
|
|
||||||
+1
-10
@@ -44,10 +44,7 @@ type Service struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func New(control *nebula.Control) (*Service, error) {
|
func New(control *nebula.Control) (*Service, error) {
|
||||||
wait, err := control.Start()
|
control.Start()
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx := control.Context()
|
ctx := control.Context()
|
||||||
eg, ctx := errgroup.WithContext(ctx)
|
eg, ctx := errgroup.WithContext(ctx)
|
||||||
@@ -144,12 +141,6 @@ func New(control *nebula.Control) (*Service, error) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
// Add the nebula wait function to the group so a fatal reader error
|
|
||||||
// propagates out through errgroup.Wait().
|
|
||||||
eg.Go(func() error {
|
|
||||||
return wait()
|
|
||||||
})
|
|
||||||
|
|
||||||
return &s, nil
|
return &s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
|
||||||
"reflect"
|
"reflect"
|
||||||
"runtime"
|
"runtime"
|
||||||
"runtime/pprof"
|
"runtime/pprof"
|
||||||
@@ -189,12 +188,6 @@ func configSSH(l *logrus.Logger, ssh *sshd.SSHServer, c *config.C) (func(), erro
|
|||||||
}
|
}
|
||||||
|
|
||||||
func attachCommands(l *logrus.Logger, c *config.C, ssh *sshd.SSHServer, f *Interface) {
|
func attachCommands(l *logrus.Logger, c *config.C, ssh *sshd.SSHServer, f *Interface) {
|
||||||
// sandboxDir defaults to a dir in temp. The intention is that end user will
|
|
||||||
// create this dir as needed. Overriding this config value to "" allows
|
|
||||||
// writing to anywhere in the system.
|
|
||||||
defaultDir := filepath.Join(os.TempDir(), "nebula-debug")
|
|
||||||
sandboxDir := c.GetString("sshd.sandbox_dir", defaultDir)
|
|
||||||
|
|
||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
Name: "list-hostmap",
|
Name: "list-hostmap",
|
||||||
ShortDescription: "List all known previously connected hosts",
|
ShortDescription: "List all known previously connected hosts",
|
||||||
@@ -253,9 +246,7 @@ func attachCommands(l *logrus.Logger, c *config.C, ssh *sshd.SSHServer, f *Inter
|
|||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
Name: "start-cpu-profile",
|
Name: "start-cpu-profile",
|
||||||
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
|
ShortDescription: "Starts a cpu profile and write output to the provided file, ex: `cpu-profile.pb.gz`",
|
||||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
Callback: sshStartCpuProfile,
|
||||||
return sshStartCpuProfile(sandboxDir, fs, a, w)
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
|
|
||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
@@ -270,9 +261,7 @@ func attachCommands(l *logrus.Logger, c *config.C, ssh *sshd.SSHServer, f *Inter
|
|||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
Name: "save-heap-profile",
|
Name: "save-heap-profile",
|
||||||
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
|
ShortDescription: "Saves a heap profile to the provided path, ex: `heap-profile.pb.gz`",
|
||||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
Callback: sshGetHeapProfile,
|
||||||
return sshGetHeapProfile(sandboxDir, fs, a, w)
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
|
|
||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
@@ -284,9 +273,7 @@ func attachCommands(l *logrus.Logger, c *config.C, ssh *sshd.SSHServer, f *Inter
|
|||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
Name: "save-mutex-profile",
|
Name: "save-mutex-profile",
|
||||||
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
|
ShortDescription: "Saves a mutex profile to the provided path, ex: `mutex-profile.pb.gz`",
|
||||||
Callback: func(fs any, a []string, w sshd.StringWriter) error {
|
Callback: sshGetMutexProfile,
|
||||||
return sshGetMutexProfile(sandboxDir, fs, a, w)
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
|
|
||||||
ssh.RegisterCommand(&sshd.Command{
|
ssh.RegisterCommand(&sshd.Command{
|
||||||
@@ -519,43 +506,13 @@ func sshListLighthouseMap(lightHouse *LightHouse, a any, w sshd.StringWriter) er
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// sshSanitizeFilePath validates that the given file path is within the sandbox directory.
|
func sshStartCpuProfile(fs any, a []string, w sshd.StringWriter) error {
|
||||||
// If sandboxDir is empty, the path is returned as-is for backwards compatibility.
|
|
||||||
func sshSanitizeFilePath(sandboxDir, filePath string) (string, error) {
|
|
||||||
if sandboxDir == "" {
|
|
||||||
return filePath, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Clean and resolve the path relative to the sandbox directory
|
|
||||||
if !filepath.IsAbs(filePath) {
|
|
||||||
filePath = filepath.Join(sandboxDir, filePath)
|
|
||||||
}
|
|
||||||
cleaned := filepath.Clean(filePath)
|
|
||||||
|
|
||||||
// Ensure the resolved path is within the sandbox directory
|
|
||||||
cleanedSandbox := filepath.Clean(sandboxDir)
|
|
||||||
if cleaned == cleanedSandbox {
|
|
||||||
return "", fmt.Errorf("path %q resolves to the sandbox directory itself %q", filePath, sandboxDir)
|
|
||||||
}
|
|
||||||
if !strings.HasPrefix(cleaned, cleanedSandbox+string(filepath.Separator)) {
|
|
||||||
return "", fmt.Errorf("path %q is outside the sandbox directory %q", filePath, sandboxDir)
|
|
||||||
}
|
|
||||||
|
|
||||||
return cleaned, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func sshStartCpuProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
|
||||||
if len(a) == 0 {
|
if len(a) == 0 {
|
||||||
err := w.WriteLine("No path to write profile provided")
|
err := w.WriteLine("No path to write profile provided")
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
file, err := os.Create(a[0])
|
||||||
if err != nil {
|
|
||||||
return w.WriteLine(err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
file, err := os.Create(filePath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
return err
|
return err
|
||||||
@@ -675,6 +632,9 @@ func sshCreateTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) e
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
|
hostInfo = ifce.handshakeManager.StartHandshake(vpnAddr, nil)
|
||||||
|
if hostInfo == nil {
|
||||||
|
return w.WriteLine("Handshake rate limit reached")
|
||||||
|
}
|
||||||
if addr.IsValid() {
|
if addr.IsValid() {
|
||||||
hostInfo.SetRemote(addr)
|
hostInfo.SetRemote(addr)
|
||||||
}
|
}
|
||||||
@@ -719,17 +679,12 @@ func sshChangeRemote(ifce *Interface, fs any, a []string, w sshd.StringWriter) e
|
|||||||
return w.WriteLine("Changed")
|
return w.WriteLine("Changed")
|
||||||
}
|
}
|
||||||
|
|
||||||
func sshGetHeapProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
func sshGetHeapProfile(fs any, a []string, w sshd.StringWriter) error {
|
||||||
if len(a) == 0 {
|
if len(a) == 0 {
|
||||||
return w.WriteLine("No path to write profile provided")
|
return w.WriteLine("No path to write profile provided")
|
||||||
}
|
}
|
||||||
|
|
||||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
file, err := os.Create(a[0])
|
||||||
if err != nil {
|
|
||||||
return w.WriteLine(err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
file, err := os.Create(filePath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
err = w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
return err
|
return err
|
||||||
@@ -760,17 +715,12 @@ func sshMutexProfileFraction(fs any, a []string, w sshd.StringWriter) error {
|
|||||||
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
|
return w.WriteLine(fmt.Sprintf("New value: %d. Old value: %d", newRate, oldRate))
|
||||||
}
|
}
|
||||||
|
|
||||||
func sshGetMutexProfile(sandboxDir string, fs any, a []string, w sshd.StringWriter) error {
|
func sshGetMutexProfile(fs any, a []string, w sshd.StringWriter) error {
|
||||||
if len(a) == 0 {
|
if len(a) == 0 {
|
||||||
return w.WriteLine("No path to write profile provided")
|
return w.WriteLine("No path to write profile provided")
|
||||||
}
|
}
|
||||||
|
|
||||||
filePath, err := sshSanitizeFilePath(sandboxDir, a[0])
|
file, err := os.Create(a[0])
|
||||||
if err != nil {
|
|
||||||
return w.WriteLine(err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
file, err := os.Create(filePath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
return w.WriteLine(fmt.Sprintf("Unable to create profile file: %s", err))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
package overlay
|
package test
|
||||||
|
|
||||||
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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -26,15 +26,11 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read() ([][]byte, error) {
|
func (NoopTun) Read([]byte) (int, error) {
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (NoopTun) Write([]byte) (int, error) {
|
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) WriteReject(p []byte) (int, error) {
|
func (NoopTun) Write([]byte) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -42,12 +38,8 @@ func (NoopTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) NewMultiQueueReader() error {
|
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return errors.New("unsupported")
|
return nil, errors.New("unsupported")
|
||||||
}
|
|
||||||
|
|
||||||
func (NoopTun) Readers() []tio.Queue {
|
|
||||||
return []tio.Queue{NoopTun{}}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Close() error {
|
func (NoopTun) Close() error {
|
||||||
+3
-23
@@ -8,12 +8,6 @@ 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 a chunked retry.
|
|
||||||
const MaxWriteBatch = 128
|
|
||||||
|
|
||||||
type EncReader func(
|
type EncReader func(
|
||||||
addr netip.AddrPort,
|
addr netip.AddrPort,
|
||||||
payload []byte,
|
payload []byte,
|
||||||
@@ -22,19 +16,8 @@ type EncReader func(
|
|||||||
type Conn interface {
|
type Conn interface {
|
||||||
Rebind() error
|
Rebind() error
|
||||||
LocalAddr() (netip.AddrPort, error)
|
LocalAddr() (netip.AddrPort, error)
|
||||||
// ListenOut invokes r for each received packet. On batch-capable
|
ListenOut(r EncReader)
|
||||||
// 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. Linux uses
|
|
||||||
// sendmmsg(2) for a single syscall; other backends fall back to a
|
|
||||||
// WriteTo loop. 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) error
|
|
||||||
ReloadConfig(c *config.C)
|
ReloadConfig(c *config.C)
|
||||||
SupportsMultipleReaders() bool
|
SupportsMultipleReaders() bool
|
||||||
Close() error
|
Close() error
|
||||||
@@ -48,8 +31,8 @@ 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, _ func()) error {
|
func (NoopConn) ListenOut(_ EncReader) {
|
||||||
return nil
|
return
|
||||||
}
|
}
|
||||||
func (NoopConn) SupportsMultipleReaders() bool {
|
func (NoopConn) SupportsMultipleReaders() bool {
|
||||||
return false
|
return false
|
||||||
@@ -57,9 +40,6 @@ 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) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
+3
-23
@@ -140,26 +140,6 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) error {
|
|
||||||
for i, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addrs[i]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) WriteSegmented(bufs [][]byte, addr netip.AddrPort, _ int) error {
|
|
||||||
for _, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) SupportsGSO() bool { return false }
|
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -185,7 +165,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
|||||||
return func() {}
|
return func() {}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
func (u *StdConn) ListenOut(r EncReader) {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
@@ -193,14 +173,14 @@ func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
n, rua, err := u.ReadFromUDPAddrPort(buffer)
|
n, rua, err := u.ReadFromUDPAddrPort(buffer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, net.ErrClosed) {
|
if errors.Is(err, net.ErrClosed) {
|
||||||
return err
|
u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
u.l.WithError(err).Error("unexpected udp socket receive error")
|
u.l.WithError(err).Error("unexpected udp socket receive error")
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
||||||
flush()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+3
-23
@@ -44,26 +44,6 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) error {
|
|
||||||
for i, b := range bufs {
|
|
||||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *GenericConn) WriteSegmented(bufs [][]byte, addr netip.AddrPort, _ int) error {
|
|
||||||
for _, b := range bufs {
|
|
||||||
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *GenericConn) SupportsGSO() bool { return false }
|
|
||||||
|
|
||||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -93,7 +73,7 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
func (u *GenericConn) ListenOut(r EncReader) {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -103,7 +83,8 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
n, rua, err := u.ReadFromUDPAddrPort(buffer)
|
n, rua, err := u.ReadFromUDPAddrPort(buffer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, net.ErrClosed) {
|
if errors.Is(err, net.ErrClosed) {
|
||||||
return err
|
u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
|
||||||
|
return
|
||||||
}
|
}
|
||||||
// Dampen unexpected message warns to once per minute
|
// Dampen unexpected message warns to once per minute
|
||||||
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
|
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
|
||||||
@@ -114,7 +95,6 @@ func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
||||||
flush()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+148
-545
@@ -4,7 +4,6 @@
|
|||||||
package udp
|
package udp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
@@ -19,188 +18,58 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type StdConn struct {
|
type StdConn struct {
|
||||||
udpConn *net.UDPConn
|
sysFd int
|
||||||
rawConn syscall.RawConn
|
isV4 bool
|
||||||
isV4 bool
|
l *logrus.Logger
|
||||||
l *logrus.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
|
|
||||||
|
|
||||||
// Per-entry UDP_SEGMENT cmsg scratch. writeCmsg is one contiguous slab
|
|
||||||
// of MaxWriteBatch * writeCmsgSpace bytes; each entry's cmsg header is
|
|
||||||
// pre-filled once in prepareWriteMessages. WriteBatch only rewrites the
|
|
||||||
// 2-byte gso_size payload (and toggles Hdr.Control on/off) per call.
|
|
||||||
writeCmsg []byte
|
|
||||||
writeCmsgSpace int
|
|
||||||
|
|
||||||
// writeEntryEnd[e] is the bufs index *after* the last packet packed
|
|
||||||
// into mmsghdr entry e. Used to rewind `i` on partial sendmmsg success.
|
|
||||||
writeEntryEnd []int
|
|
||||||
|
|
||||||
// Preallocated closure + in/out slots for sendmmsg, so the hot path
|
|
||||||
// does not heap-allocate a fresh closure per call.
|
|
||||||
writeChunk int
|
|
||||||
writeSent int
|
|
||||||
writeErrno syscall.Errno
|
|
||||||
writeFunc func(fd uintptr) bool
|
|
||||||
|
|
||||||
// UDP GSO (sendmsg with UDP_SEGMENT cmsg) support. gsoSupported is
|
|
||||||
// probed once at socket creation. When true, WriteSegmented takes a
|
|
||||||
// single-syscall GSO path; otherwise it falls back to a WriteTo loop.
|
|
||||||
gsoSupported bool
|
|
||||||
|
|
||||||
// UDP GRO (recvmsg with UDP_GRO cmsg) support. groSupported is probed
|
|
||||||
// once at socket creation. When true, listenOutBatch allocates larger
|
|
||||||
// RX buffers and a per-entry cmsg slot so the kernel can coalesce
|
|
||||||
// consecutive same-flow datagrams into a single recvmmsg entry; the
|
|
||||||
// delivered cmsg carries the gso_size used to split them back apart.
|
|
||||||
groSupported bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func setReusePort(network, address string, c syscall.RawConn) error {
|
func maybeIPV4(ip net.IP) (net.IP, bool) {
|
||||||
var opErr error
|
ip4 := ip.To4()
|
||||||
err := c.Control(func(fd uintptr) {
|
if ip4 != nil {
|
||||||
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1)
|
return ip4, true
|
||||||
//CloseOnExec already set by the runtime
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
return opErr
|
return ip, false
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewListener(l *logrus.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
func NewListener(l *logrus.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) {
|
||||||
listen := netip.AddrPortFrom(ip, uint16(port))
|
af := unix.AF_INET6
|
||||||
lc := net.ListenConfig{}
|
if ip.Is4() {
|
||||||
if multi {
|
af = unix.AF_INET
|
||||||
lc.Control = setReusePort
|
|
||||||
}
|
}
|
||||||
//this context is only used during the bind operation, you can't cancel it to kill the socket
|
syscall.ForkLock.RLock()
|
||||||
pc, err := lc.ListenPacket(context.Background(), "udp", listen.String())
|
fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP)
|
||||||
|
if err == nil {
|
||||||
|
unix.CloseOnExec(fd)
|
||||||
|
}
|
||||||
|
syscall.ForkLock.RUnlock()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
unix.Close(fd)
|
||||||
return nil, fmt.Errorf("unable to open socket: %s", err)
|
return nil, fmt.Errorf("unable to open socket: %s", err)
|
||||||
}
|
}
|
||||||
udpConn := pc.(*net.UDPConn)
|
|
||||||
rawConn, err := udpConn.SyscallConn()
|
if multi {
|
||||||
if err != nil {
|
if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil {
|
||||||
_ = udpConn.Close()
|
return nil, fmt.Errorf("unable to set SO_REUSEPORT: %s", err)
|
||||||
return nil, err
|
}
|
||||||
}
|
|
||||||
//gotta find out if we got an AF_INET6 socket or not:
|
|
||||||
out := &StdConn{
|
|
||||||
udpConn: udpConn,
|
|
||||||
rawConn: rawConn,
|
|
||||||
l: l,
|
|
||||||
batch: batch,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
af, err := out.getSockOptInt(unix.SO_DOMAIN)
|
var sa unix.Sockaddr
|
||||||
if err != nil {
|
if ip.Is4() {
|
||||||
_ = out.Close()
|
sa4 := &unix.SockaddrInet4{Port: port}
|
||||||
return nil, err
|
sa4.Addr = ip.As4()
|
||||||
|
sa = sa4
|
||||||
|
} else {
|
||||||
|
sa6 := &unix.SockaddrInet6{Port: port}
|
||||||
|
sa6.Addr = ip.As16()
|
||||||
|
sa = sa6
|
||||||
}
|
}
|
||||||
out.isV4 = af == unix.AF_INET
|
if err = unix.Bind(fd, sa); err != nil {
|
||||||
|
return nil, fmt.Errorf("unable to bind to socket: %s", err)
|
||||||
out.prepareWriteMessages(MaxWriteBatch)
|
|
||||||
out.writeFunc = out.sendmmsgRawWrite
|
|
||||||
|
|
||||||
out.prepareGSO()
|
|
||||||
// GRO delivers coalesced superpackets that need a cmsg to split back
|
|
||||||
// into segments. The single-packet RX path uses ReadFromUDPAddrPort
|
|
||||||
// and cannot see that cmsg, so only enable GRO for the batch path.
|
|
||||||
if batch > 1 {
|
|
||||||
out.prepareGRO()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return out, nil
|
return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, err
|
||||||
}
|
|
||||||
|
|
||||||
// prepareWriteMessages allocates one mmsghdr/iovec/sockaddr/cmsg scratch
|
|
||||||
// slot per sendmmsg entry. The iovec slab is sized to the same n so a
|
|
||||||
// single entry can fan out to up to n iovecs (needed for UDP_SEGMENT runs
|
|
||||||
// that coalesce consecutive bufs into one entry). Hdr.Iov / Hdr.Iovlen /
|
|
||||||
// Hdr.Control / Hdr.Controllen are wired per call since each entry can
|
|
||||||
// span a variable number of iovecs and may or may not carry a cmsg.
|
|
||||||
func (u *StdConn) prepareWriteMessages(n int) {
|
|
||||||
u.writeMsgs = make([]rawMessage, n)
|
|
||||||
u.writeIovs = make([]iovec, n)
|
|
||||||
u.writeNames = make([][]byte, n)
|
|
||||||
u.writeEntryEnd = make([]int, n)
|
|
||||||
|
|
||||||
u.writeCmsgSpace = unix.CmsgSpace(2)
|
|
||||||
u.writeCmsg = make([]byte, n*u.writeCmsgSpace)
|
|
||||||
for k := 0; k < n; k++ {
|
|
||||||
off := k * u.writeCmsgSpace
|
|
||||||
h := (*unix.Cmsghdr)(unsafe.Pointer(&u.writeCmsg[off]))
|
|
||||||
h.Level = unix.SOL_UDP
|
|
||||||
h.Type = unix.UDP_SEGMENT
|
|
||||||
setCmsgLen(h, unix.CmsgLen(2))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range u.writeMsgs {
|
|
||||||
u.writeNames[i] = make([]byte, unix.SizeofSockaddrInet6)
|
|
||||||
u.writeMsgs[i].Hdr.Name = &u.writeNames[i][0]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// maxGSOSegments caps the per-sendmsg GSO fan-out. Linux kernels have
|
|
||||||
// historically capped UDP_MAX_SEGMENTS at 64; newer kernels raise it to 128
|
|
||||||
// but we stay conservative so the same code works everywhere.
|
|
||||||
const maxGSOSegments = 64
|
|
||||||
|
|
||||||
// maxGSOBytes bounds the total payload per sendmsg() when UDP_SEGMENT is
|
|
||||||
// set. The kernel stitches all iovecs into a single skb whose length the
|
|
||||||
// UDP length field can represent, and also enforces sk_gso_max_size (which
|
|
||||||
// on most devices is 65536). We use 65535 so ciphertext + headers always
|
|
||||||
// fits, avoiding EMSGSIZE on large TSO superpackets.
|
|
||||||
const maxGSOBytes = 65535
|
|
||||||
|
|
||||||
// prepareGSO probes UDP_SEGMENT support
|
|
||||||
func (u *StdConn) prepareGSO() {
|
|
||||||
var probeErr error
|
|
||||||
if err := u.rawConn.Control(func(fd uintptr) {
|
|
||||||
probeErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT, 0)
|
|
||||||
}); err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if probeErr != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
u.gsoSupported = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// udpGROBufferSize sizes the per-entry recvmmsg buffer when UDP_GRO is on.
|
|
||||||
// The kernel stitches a run of same-flow datagrams into a single skb whose
|
|
||||||
// length is bounded by sk_gso_max_size (typically 65535); anything larger
|
|
||||||
// would be MSG_TRUNCed. We use the maximum representable UDP length so a
|
|
||||||
// full superpacket always lands intact.
|
|
||||||
const udpGROBufferSize = 65535
|
|
||||||
|
|
||||||
// udpGROCmsgPayload is the size of the UDP_GRO cmsg data delivered by the
|
|
||||||
// kernel: a single int (gso_size in bytes). See udp_cmsg_recv() in
|
|
||||||
// net/ipv4/udp.c.
|
|
||||||
const udpGROCmsgPayload = 4
|
|
||||||
|
|
||||||
// prepareGRO turns on UDP_GRO so the kernel coalesces consecutive same-flow
|
|
||||||
// datagrams into one recvmmsg entry, with a cmsg carrying the gso_size used
|
|
||||||
// to split them back apart on the application side.
|
|
||||||
func (u *StdConn) prepareGRO() {
|
|
||||||
var probeErr error
|
|
||||||
if err := u.rawConn.Control(func(fd uintptr) {
|
|
||||||
probeErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
|
|
||||||
}); err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if probeErr != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
u.groSupported = true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||||
@@ -211,145 +80,62 @@ func (u *StdConn) Rebind() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) getSockOptInt(opt int) (int, error) {
|
|
||||||
if u.rawConn == nil {
|
|
||||||
return 0, fmt.Errorf("no UDP connection")
|
|
||||||
}
|
|
||||||
var out int
|
|
||||||
var opErr error
|
|
||||||
err := u.rawConn.Control(func(fd uintptr) {
|
|
||||||
out, opErr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt)
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
return out, opErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) setSockOptInt(opt int, n int) error {
|
|
||||||
if u.rawConn == nil {
|
|
||||||
return fmt.Errorf("no UDP connection")
|
|
||||||
}
|
|
||||||
var opErr error
|
|
||||||
err := u.rawConn.Control(func(fd uintptr) {
|
|
||||||
opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, opt, n)
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return opErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) SetRecvBuffer(n int) error {
|
func (u *StdConn) SetRecvBuffer(n int) error {
|
||||||
return u.setSockOptInt(unix.SO_RCVBUFFORCE, n)
|
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) SetSendBuffer(n int) error {
|
func (u *StdConn) SetSendBuffer(n int) error {
|
||||||
return u.setSockOptInt(unix.SO_SNDBUFFORCE, n)
|
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) SetSoMark(mark int) error {
|
func (u *StdConn) SetSoMark(mark int) error {
|
||||||
return u.setSockOptInt(unix.SO_MARK, mark)
|
return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK, mark)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) GetRecvBuffer() (int, error) {
|
func (u *StdConn) GetRecvBuffer() (int, error) {
|
||||||
return u.getSockOptInt(unix.SO_RCVBUF)
|
return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_RCVBUF)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) GetSendBuffer() (int, error) {
|
func (u *StdConn) GetSendBuffer() (int, error) {
|
||||||
return u.getSockOptInt(unix.SO_SNDBUF)
|
return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_SNDBUF)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) GetSoMark() (int, error) {
|
func (u *StdConn) GetSoMark() (int, error) {
|
||||||
return u.getSockOptInt(unix.SO_MARK)
|
return unix.GetsockoptInt(int(u.sysFd), unix.SOL_SOCKET, unix.SO_MARK)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.udpConn.LocalAddr()
|
sa, err := unix.Getsockname(u.sysFd)
|
||||||
|
if err != nil {
|
||||||
|
return netip.AddrPort{}, err
|
||||||
|
}
|
||||||
|
|
||||||
switch v := a.(type) {
|
switch sa := sa.(type) {
|
||||||
case *net.UDPAddr:
|
case *unix.SockaddrInet4:
|
||||||
addr, ok := netip.AddrFromSlice(v.IP)
|
return netip.AddrPortFrom(netip.AddrFrom4(sa.Addr), uint16(sa.Port)), nil
|
||||||
if !ok {
|
|
||||||
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned invalid IP address: %s", v.IP)
|
case *unix.SockaddrInet6:
|
||||||
}
|
return netip.AddrPortFrom(netip.AddrFrom16(sa.Addr), uint16(sa.Port)), nil
|
||||||
return netip.AddrPortFrom(addr, uint16(v.Port)), nil
|
|
||||||
|
|
||||||
default:
|
default:
|
||||||
return netip.AddrPort{}, fmt.Errorf("LocalAddr returned: %#v", a)
|
return netip.AddrPort{}, fmt.Errorf("unsupported sock type: %T", sa)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
|
func (u *StdConn) ListenOut(r EncReader) {
|
||||||
var errno syscall.Errno
|
|
||||||
n, _, errno := unix.Syscall6(
|
|
||||||
unix.SYS_RECVMMSG,
|
|
||||||
fd,
|
|
||||||
uintptr(unsafe.Pointer(&msgs[0])),
|
|
||||||
uintptr(len(msgs)),
|
|
||||||
unix.MSG_WAITFORONE,
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
)
|
|
||||||
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
|
||||||
// No data available, block for I/O and try again.
|
|
||||||
return int(n), false, nil
|
|
||||||
}
|
|
||||||
if errno != 0 {
|
|
||||||
return int(n), true, &net.OpError{Op: "recvmmsg", Err: errno}
|
|
||||||
}
|
|
||||||
return int(n), true, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) listenOutSingle(r EncReader, flush func()) error {
|
|
||||||
var err error
|
|
||||||
var n int
|
|
||||||
var from netip.AddrPort
|
|
||||||
buffer := make([]byte, MTU)
|
|
||||||
|
|
||||||
for {
|
|
||||||
n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
|
||||||
r(from, buffer[:n])
|
|
||||||
flush()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) listenOutBatch(r EncReader, flush func()) error {
|
|
||||||
var ip netip.Addr
|
var ip netip.Addr
|
||||||
var n int
|
|
||||||
var operr error
|
|
||||||
|
|
||||||
bufSize := MTU
|
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
||||||
cmsgSpace := 0
|
read := u.ReadMulti
|
||||||
if u.groSupported {
|
if u.batch == 1 {
|
||||||
bufSize = udpGROBufferSize
|
read = u.ReadSingle
|
||||||
cmsgSpace = unix.CmsgSpace(udpGROCmsgPayload)
|
|
||||||
}
|
|
||||||
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
|
|
||||||
//defining it outside the loop so it gets re-used
|
|
||||||
reader := func(fd uintptr) (done bool) {
|
|
||||||
n, done, operr = recvmmsg(fd, msgs)
|
|
||||||
return done
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for {
|
for {
|
||||||
if cmsgSpace > 0 {
|
n, err := read(msgs)
|
||||||
for i := range msgs {
|
|
||||||
setMsgControllen(&msgs[i].Hdr, cmsgSpace)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
err := u.rawConn.Read(reader)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
|
||||||
}
|
return
|
||||||
if operr != nil {
|
|
||||||
return operr
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := 0; i < n; i++ {
|
for i := 0; i < n; i++ {
|
||||||
@@ -359,281 +145,111 @@ func (u *StdConn) listenOutBatch(r EncReader, flush func()) error {
|
|||||||
} else {
|
} else {
|
||||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
||||||
}
|
}
|
||||||
from := netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4]))
|
r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len])
|
||||||
payload := buffers[i][:msgs[i].Len]
|
|
||||||
|
|
||||||
segSize := 0
|
|
||||||
if u.groSupported {
|
|
||||||
segSize = parseUDPGRO(&msgs[i].Hdr)
|
|
||||||
}
|
|
||||||
if segSize <= 0 || segSize >= len(payload) {
|
|
||||||
// No coalescing happened (or a lone datagram).
|
|
||||||
r(from, payload)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// GRO superpacket: the kernel guarantees every segment is
|
|
||||||
// exactly segSize bytes except for the final one, which may be
|
|
||||||
// short.
|
|
||||||
for off := 0; off < len(payload); off += segSize {
|
|
||||||
end := off + segSize
|
|
||||||
if end > len(payload) {
|
|
||||||
end = len(payload)
|
|
||||||
}
|
|
||||||
r(from, payload[off:end])
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
// End-of-batch: let callers (e.g. TUN write coalescer) flush any
|
|
||||||
// state they accumulated across this batch.
|
|
||||||
flush()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseUDPGRO walks the control buffer on hdr looking for a SOL_UDP/UDP_GRO
|
func (u *StdConn) ReadSingle(msgs []rawMessage) (int, error) {
|
||||||
// cmsg and returns the gso_size (bytes per coalesced segment) it carries.
|
for {
|
||||||
// Returns 0 when no UDP_GRO cmsg is present, which is the normal case for
|
n, _, err := unix.Syscall6(
|
||||||
// lone datagrams that the kernel did not coalesce.
|
unix.SYS_RECVMSG,
|
||||||
func parseUDPGRO(hdr *msghdr) int {
|
uintptr(u.sysFd),
|
||||||
controllen := int(hdr.Controllen)
|
uintptr(unsafe.Pointer(&(msgs[0].Hdr))),
|
||||||
if controllen < unix.SizeofCmsghdr || hdr.Control == nil {
|
0,
|
||||||
return 0
|
0,
|
||||||
}
|
0,
|
||||||
ctrl := unsafe.Slice(hdr.Control, controllen)
|
0,
|
||||||
off := 0
|
)
|
||||||
for off+unix.SizeofCmsghdr <= len(ctrl) {
|
|
||||||
ch := (*unix.Cmsghdr)(unsafe.Pointer(&ctrl[off]))
|
if err != 0 {
|
||||||
clen := int(ch.Len)
|
return 0, &net.OpError{Op: "recvmsg", Err: err}
|
||||||
if clen < unix.SizeofCmsghdr || off+clen > len(ctrl) {
|
|
||||||
return 0
|
|
||||||
}
|
}
|
||||||
if ch.Level == unix.SOL_UDP && ch.Type == unix.UDP_GRO {
|
|
||||||
dataOff := off + unix.CmsgLen(0)
|
msgs[0].Len = uint32(n)
|
||||||
if dataOff+udpGROCmsgPayload <= len(ctrl) {
|
return 1, nil
|
||||||
return int(int32(binary.NativeEndian.Uint32(ctrl[dataOff : dataOff+udpGROCmsgPayload])))
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
// Advance by the aligned cmsg space. CmsgSpace(n) is the stride
|
|
||||||
// from one header to the next (len aligned up to the platform's
|
|
||||||
// cmsg alignment).
|
|
||||||
off += unix.CmsgSpace(clen - unix.CmsgLen(0))
|
|
||||||
}
|
}
|
||||||
return 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
func (u *StdConn) ReadMulti(msgs []rawMessage) (int, error) {
|
||||||
if u.batch == 1 {
|
for {
|
||||||
return u.listenOutSingle(r, flush)
|
n, _, err := unix.Syscall6(
|
||||||
} else {
|
unix.SYS_RECVMMSG,
|
||||||
return u.listenOutBatch(r, flush)
|
uintptr(u.sysFd),
|
||||||
|
uintptr(unsafe.Pointer(&msgs[0])),
|
||||||
|
uintptr(len(msgs)),
|
||||||
|
unix.MSG_WAITFORONE,
|
||||||
|
0,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
|
||||||
|
if err != 0 {
|
||||||
|
return 0, &net.OpError{Op: "recvmmsg", Err: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
return int(n), nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
||||||
_, err := u.udpConn.WriteToUDPAddrPort(b, ip)
|
if u.isV4 {
|
||||||
return err
|
return u.writeTo4(b, ip)
|
||||||
|
}
|
||||||
|
return u.writeTo6(b, ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
// WriteBatch sends bufs via sendmmsg(2) using the preallocated scratch on
|
func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error {
|
||||||
// StdConn. Consecutive packets to the same destination with matching segment
|
var rsa unix.RawSockaddrInet6
|
||||||
// sizes (all but possibly the last) are coalesced into a single mmsghdr entry
|
rsa.Family = unix.AF_INET6
|
||||||
// carrying a UDP_SEGMENT cmsg, so one syscall can mix runs of GSO superpackets
|
rsa.Addr = ip.Addr().As16()
|
||||||
// with plain one-off datagrams. Without GSO support every packet is its own
|
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||||
// entry, matching the prior behaviour.
|
|
||||||
//
|
for {
|
||||||
// Chunks larger than the scratch are processed across multiple syscalls. If
|
_, _, err := unix.Syscall6(
|
||||||
// sendmmsg returns a fatal error before any entry is sent we fall back to
|
unix.SYS_SENDTO,
|
||||||
// per-packet WriteTo for that chunk so the caller still gets best-effort
|
uintptr(u.sysFd),
|
||||||
// delivery.
|
uintptr(unsafe.Pointer(&b[0])),
|
||||||
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) error {
|
uintptr(len(b)),
|
||||||
if len(bufs) != len(addrs) {
|
uintptr(0),
|
||||||
return fmt.Errorf("WriteBatch: len(bufs)=%d != len(addrs)=%d", len(bufs), len(addrs))
|
uintptr(unsafe.Pointer(&rsa)),
|
||||||
|
uintptr(unix.SizeofSockaddrInet6),
|
||||||
|
)
|
||||||
|
|
||||||
|
if err != 0 {
|
||||||
|
return &net.OpError{Op: "sendto", Err: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
i := 0
|
|
||||||
for i < len(bufs) {
|
|
||||||
baseI := i
|
|
||||||
entry := 0
|
|
||||||
iovIdx := 0
|
|
||||||
|
|
||||||
for entry < len(u.writeMsgs) && i < len(bufs) {
|
|
||||||
iovBudget := len(u.writeIovs) - iovIdx
|
|
||||||
if iovBudget < 1 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
runLen, segSize := u.planRun(bufs, addrs, i, iovBudget)
|
|
||||||
if runLen == 0 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
|
|
||||||
for k := 0; k < runLen; k++ {
|
|
||||||
b := bufs[i+k]
|
|
||||||
if len(b) == 0 {
|
|
||||||
u.writeIovs[iovIdx+k].Base = nil
|
|
||||||
setIovLen(&u.writeIovs[iovIdx+k], 0)
|
|
||||||
} else {
|
|
||||||
u.writeIovs[iovIdx+k].Base = &b[0]
|
|
||||||
setIovLen(&u.writeIovs[iovIdx+k], len(b))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
nlen, err := writeSockaddr(u.writeNames[entry], addrs[i], u.isV4)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
hdr := &u.writeMsgs[entry].Hdr
|
|
||||||
hdr.Iov = &u.writeIovs[iovIdx]
|
|
||||||
setMsgIovlen(hdr, runLen)
|
|
||||||
hdr.Namelen = uint32(nlen)
|
|
||||||
|
|
||||||
if runLen >= 2 {
|
|
||||||
off := entry * u.writeCmsgSpace
|
|
||||||
dataOff := off + unix.CmsgLen(0)
|
|
||||||
binary.NativeEndian.PutUint16(u.writeCmsg[dataOff:dataOff+2], uint16(segSize))
|
|
||||||
hdr.Control = &u.writeCmsg[off]
|
|
||||||
setMsgControllen(hdr, u.writeCmsgSpace)
|
|
||||||
} else {
|
|
||||||
hdr.Control = nil
|
|
||||||
setMsgControllen(hdr, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
i += runLen
|
|
||||||
iovIdx += runLen
|
|
||||||
u.writeEntryEnd[entry] = i
|
|
||||||
entry++
|
|
||||||
}
|
|
||||||
|
|
||||||
if entry == 0 {
|
|
||||||
return fmt.Errorf("sendmmsg: no progress")
|
|
||||||
}
|
|
||||||
|
|
||||||
sent, serr := u.sendmmsg(entry)
|
|
||||||
if serr != nil && sent <= 0 {
|
|
||||||
// Nothing went out for this chunk; fall back to WriteTo for each
|
|
||||||
// packet that was queued this iteration.
|
|
||||||
for k := baseI; k < i; k++ {
|
|
||||||
if werr := u.WriteTo(bufs[k], addrs[k]); werr != nil {
|
|
||||||
return werr
|
|
||||||
}
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if sent == 0 {
|
|
||||||
return fmt.Errorf("sendmmsg made no progress")
|
|
||||||
}
|
|
||||||
// Rewind i to the end of the last successfully sent entry. For a
|
|
||||||
// full-success send this leaves i unchanged; for a partial send it
|
|
||||||
// replays the remainder on the next outer-loop iteration.
|
|
||||||
i = u.writeEntryEnd[sent-1]
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// planRun groups consecutive packets starting at `start` that can be sent as
|
func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error {
|
||||||
// a single UDP GSO superpacket (one sendmmsg entry with UDP_SEGMENT cmsg).
|
if !ip.Addr().Is4() {
|
||||||
// A run of length 1 means the entry carries no cmsg and the kernel treats
|
return ErrInvalidIPv6RemoteForSocket
|
||||||
// it as a plain datagram. Returns the run length and the per-segment size
|
|
||||||
// (which equals len(bufs[start])). Without GSO support every call returns
|
|
||||||
// runLen=1.
|
|
||||||
func (u *StdConn) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovBudget int) (int, int) {
|
|
||||||
if start >= len(bufs) || iovBudget < 1 {
|
|
||||||
return 0, 0
|
|
||||||
}
|
}
|
||||||
segSize := len(bufs[start])
|
|
||||||
if !u.gsoSupported || segSize == 0 || segSize > maxGSOBytes {
|
|
||||||
return 1, segSize
|
|
||||||
}
|
|
||||||
dst := addrs[start]
|
|
||||||
maxLen := maxGSOSegments
|
|
||||||
if iovBudget < maxLen {
|
|
||||||
maxLen = iovBudget
|
|
||||||
}
|
|
||||||
runLen := 1
|
|
||||||
total := segSize
|
|
||||||
for runLen < maxLen && start+runLen < len(bufs) {
|
|
||||||
nextLen := len(bufs[start+runLen])
|
|
||||||
if nextLen == 0 || nextLen > segSize {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if addrs[start+runLen] != dst {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if total+nextLen > maxGSOBytes {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
total += nextLen
|
|
||||||
runLen++
|
|
||||||
if nextLen < segSize {
|
|
||||||
// A short packet must be the last in the run.
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return runLen, segSize
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendmmsgRawWrite is the preallocated callback passed to rawConn.Write. It
|
var rsa unix.RawSockaddrInet4
|
||||||
// reads its input (u.writeChunk) and writes its outputs (u.writeSent,
|
rsa.Family = unix.AF_INET
|
||||||
// u.writeErrno) through StdConn fields so the closure itself does not
|
rsa.Addr = ip.Addr().As4()
|
||||||
// capture per-call locals and therefore does not heap-allocate.
|
binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port())
|
||||||
func (u *StdConn) sendmmsgRawWrite(fd uintptr) bool {
|
|
||||||
r1, _, errno := unix.Syscall6(
|
|
||||||
unix.SYS_SENDMMSG,
|
|
||||||
fd,
|
|
||||||
uintptr(unsafe.Pointer(&u.writeMsgs[0])),
|
|
||||||
uintptr(u.writeChunk),
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
)
|
|
||||||
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
u.writeSent = int(r1)
|
|
||||||
u.writeErrno = errno
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *StdConn) sendmmsg(n int) (int, error) {
|
for {
|
||||||
u.writeChunk = n
|
_, _, err := unix.Syscall6(
|
||||||
u.writeSent = 0
|
unix.SYS_SENDTO,
|
||||||
u.writeErrno = 0
|
uintptr(u.sysFd),
|
||||||
if err := u.rawConn.Write(u.writeFunc); err != nil {
|
uintptr(unsafe.Pointer(&b[0])),
|
||||||
return u.writeSent, err
|
uintptr(len(b)),
|
||||||
}
|
uintptr(0),
|
||||||
if u.writeErrno != 0 {
|
uintptr(unsafe.Pointer(&rsa)),
|
||||||
return u.writeSent, &net.OpError{Op: "sendmmsg", Err: u.writeErrno}
|
uintptr(unix.SizeofSockaddrInet4),
|
||||||
}
|
)
|
||||||
return u.writeSent, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeSockaddr encodes addr into buf (which must be at least
|
if err != 0 {
|
||||||
// SizeofSockaddrInet6 bytes). Returns the number of bytes used. If isV4 is
|
return &net.OpError{Op: "sendto", Err: err}
|
||||||
// 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.
|
return nil
|
||||||
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[:])
|
|
||||||
for j := 8; j < 16; j++ {
|
|
||||||
buf[j] = 0
|
|
||||||
}
|
|
||||||
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) {
|
||||||
@@ -686,28 +302,15 @@ func (u *StdConn) ReloadConfig(c *config.C) {
|
|||||||
|
|
||||||
func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
|
func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error {
|
||||||
var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
|
var vallen uint32 = 4 * unix.SK_MEMINFO_VARS
|
||||||
|
_, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
|
||||||
if u.rawConn == nil {
|
if err != 0 {
|
||||||
return fmt.Errorf("no UDP connection")
|
|
||||||
}
|
|
||||||
var opErr error
|
|
||||||
err := u.rawConn.Control(func(fd uintptr) {
|
|
||||||
_, _, syserr := unix.Syscall6(unix.SYS_GETSOCKOPT, fd, uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0)
|
|
||||||
if syserr != 0 {
|
|
||||||
opErr = syserr
|
|
||||||
}
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return opErr
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) Close() error {
|
func (u *StdConn) Close() error {
|
||||||
if u.udpConn != nil {
|
return syscall.Close(u.sysFd)
|
||||||
return u.udpConn.Close()
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUDPStatsEmitter(udpConns []Conn) func() {
|
func NewUDPStatsEmitter(udpConns []Conn) func() {
|
||||||
|
|||||||
+3
-29
@@ -30,18 +30,13 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]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, bufSize)
|
buffers[i] = make([]byte, MTU)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -53,28 +48,7 @@ func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [
|
|||||||
|
|
||||||
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, cmsgs
|
return msgs, buffers, names
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
+3
-29
@@ -33,18 +33,13 @@ type rawMessage struct {
|
|||||||
Pad0 [4]byte
|
Pad0 [4]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]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, bufSize)
|
buffers[i] = make([]byte, MTU)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -56,28 +51,7 @@ func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [
|
|||||||
|
|
||||||
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, cmsgs
|
return msgs, buffers, names
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
+3
-23
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *logrus.Logger, sa windows.Sockaddr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
func (u *RIOConn) ListenOut(r EncReader) {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -151,7 +151,8 @@ func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, net.ErrClosed) {
|
if errors.Is(err, net.ErrClosed) {
|
||||||
return err
|
u.l.WithError(err).Debug("udp socket is closed, exiting read loop")
|
||||||
|
return
|
||||||
}
|
}
|
||||||
// Dampen unexpected message warns to once per minute
|
// Dampen unexpected message warns to once per minute
|
||||||
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
|
if lastRecvErr.IsZero() || time.Since(lastRecvErr) > time.Minute {
|
||||||
@@ -162,7 +163,6 @@ func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
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])
|
||||||
flush()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -317,26 +317,6 @@ 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) error {
|
|
||||||
for i, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addrs[i]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *RIOConn) WriteSegmented(bufs [][]byte, addr netip.AddrPort, _ int) error {
|
|
||||||
for _, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *RIOConn) SupportsGSO() bool { return false }
|
|
||||||
|
|
||||||
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 {
|
||||||
|
|||||||
+2
-24
@@ -6,7 +6,6 @@ package udp
|
|||||||
import (
|
import (
|
||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
@@ -107,34 +106,13 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) error {
|
func (u *TesterConn) ListenOut(r EncReader) {
|
||||||
for i, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addrs[i]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *TesterConn) WriteSegmented(bufs [][]byte, addr netip.AddrPort, _ int) error {
|
|
||||||
for _, b := range bufs {
|
|
||||||
if err := u.WriteTo(b, addr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (u *TesterConn) SupportsGSO() bool { return false }
|
|
||||||
|
|
||||||
func (u *TesterConn) ListenOut(r EncReader, flush func()) error {
|
|
||||||
for {
|
for {
|
||||||
p, ok := <-u.RxPackets
|
p, ok := <-u.RxPackets
|
||||||
if !ok {
|
if !ok {
|
||||||
return os.ErrClosed
|
return
|
||||||
}
|
}
|
||||||
r(p.From, p.Data)
|
r(p.From, p.Data)
|
||||||
flush()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user