mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 23:37:01 +02:00
Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0af7e6a1dd | |||
| c14ad0f27b | |||
| af958676b8 | |||
| 70779b26ad | |||
| 90aa8efb29 | |||
| dd7f12f7c0 |
@@ -8,13 +8,13 @@ body:
|
|||||||
attributes:
|
attributes:
|
||||||
value: |
|
value: |
|
||||||
### Thank you for taking the time to file a bug report!
|
### Thank you for taking the time to file a bug report!
|
||||||
|
|
||||||
Please fill out this form as completely as possible.
|
Please fill out this form as completely as possible.
|
||||||
|
|
||||||
- type: input
|
- type: input
|
||||||
id: version
|
id: version
|
||||||
attributes:
|
attributes:
|
||||||
label: What version of `nebula` are you using? (`nebula -version`)
|
label: What version of `nebula` are you using?
|
||||||
placeholder: 0.0.0
|
placeholder: 0.0.0
|
||||||
validations:
|
validations:
|
||||||
required: true
|
required: true
|
||||||
@@ -41,17 +41,10 @@ body:
|
|||||||
attributes:
|
attributes:
|
||||||
label: Logs from affected hosts
|
label: Logs from affected hosts
|
||||||
description: |
|
description: |
|
||||||
Please provide logs from ALL affected hosts during the time of the issue. If you do not provide logs we will be unable to assist you!
|
Provide logs from all affected hosts during the time of the issue.
|
||||||
|
|
||||||
[Learn how to find Nebula logs here.](https://nebula.defined.net/docs/guides/viewing-nebula-logs/)
|
|
||||||
|
|
||||||
Improve formatting by using <code>```</code> at the beginning and end of each log block.
|
Improve formatting by using <code>```</code> at the beginning and end of each log block.
|
||||||
value: |
|
|
||||||
```
|
|
||||||
|
|
||||||
```
|
|
||||||
validations:
|
validations:
|
||||||
required: true
|
required: false
|
||||||
|
|
||||||
- type: textarea
|
- type: textarea
|
||||||
id: configs
|
id: configs
|
||||||
@@ -59,11 +52,6 @@ body:
|
|||||||
label: Config files from affected hosts
|
label: Config files from affected hosts
|
||||||
description: |
|
description: |
|
||||||
Provide config files for all affected hosts.
|
Provide config files for all affected hosts.
|
||||||
|
|
||||||
Improve formatting by using <code>```</code> at the beginning and end of each config file.
|
Improve formatting by using <code>```</code> at the beginning and end of each config file.
|
||||||
value: |
|
|
||||||
```
|
|
||||||
|
|
||||||
```
|
|
||||||
validations:
|
validations:
|
||||||
required: true
|
required: false
|
||||||
|
|||||||
@@ -1,21 +1,13 @@
|
|||||||
blank_issues_enabled: true
|
blank_issues_enabled: true
|
||||||
contact_links:
|
contact_links:
|
||||||
- name: 💨 Performance Issues
|
|
||||||
url: https://github.com/slackhq/nebula/discussions/new/choose
|
|
||||||
about: 'We ask that you create a discussion instead of an issue for performance-related questions. This allows us to have a more open conversation about the issue and helps us to better understand the problem.'
|
|
||||||
|
|
||||||
- name: 📄 Documentation Issues
|
|
||||||
url: https://github.com/definednet/nebula-docs
|
|
||||||
about: "If you've found an issue with the website documentation, please file it in the nebula-docs repository."
|
|
||||||
|
|
||||||
- name: 📱 Mobile Nebula Issues
|
|
||||||
url: https://github.com/definednet/mobile_nebula
|
|
||||||
about: "If you're using the mobile Nebula app and have found an issue, please file it in the mobile_nebula repository."
|
|
||||||
|
|
||||||
- name: 📘 Documentation
|
- name: 📘 Documentation
|
||||||
url: https://nebula.defined.net/docs/
|
url: https://nebula.defined.net/docs/
|
||||||
about: 'The documentation is the best place to start if you are new to Nebula.'
|
about: Review documentation.
|
||||||
|
|
||||||
- name: 💁 Support/Chat
|
- name: 💁 Support/Chat
|
||||||
url: https://join.slack.com/t/nebulaoss/shared_invite/zt-39pk4xopc-CUKlGcb5Z39dQ0cK1v7ehA
|
url: https://join.slack.com/t/nebulaoss/shared_invite/enQtOTA5MDI4NDg3MTg4LTkwY2EwNTI4NzQyMzc0M2ZlODBjNWI3NTY1MzhiOThiMmZlZjVkMTI0NGY4YTMyNjUwMWEyNzNkZTJmYzQxOGU
|
||||||
about: 'For faster support, join us on Slack for assistance!'
|
about: 'This issue tracker is not for support questions. Join us on Slack for assistance!'
|
||||||
|
|
||||||
|
- name: 📱 Mobile Nebula
|
||||||
|
url: https://github.com/definednet/mobile_nebula
|
||||||
|
about: 'This issue tracker is not for mobile support. Try the Mobile Nebula repo instead!'
|
||||||
|
|||||||
@@ -1,22 +0,0 @@
|
|||||||
version: 2
|
|
||||||
updates:
|
|
||||||
- package-ecosystem: "github-actions"
|
|
||||||
directory: "/"
|
|
||||||
schedule:
|
|
||||||
interval: "weekly"
|
|
||||||
|
|
||||||
- package-ecosystem: "gomod"
|
|
||||||
directory: "/"
|
|
||||||
schedule:
|
|
||||||
interval: "weekly"
|
|
||||||
groups:
|
|
||||||
golang-x-dependencies:
|
|
||||||
patterns:
|
|
||||||
- "golang.org/x/*"
|
|
||||||
zx2c4-dependencies:
|
|
||||||
patterns:
|
|
||||||
- "golang.zx2c4.com/*"
|
|
||||||
protobuf-dependencies:
|
|
||||||
patterns:
|
|
||||||
- "github.com/golang/protobuf"
|
|
||||||
- "google.golang.org/protobuf"
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
<!--
|
|
||||||
Thank you for taking the time to submit a pull request!
|
|
||||||
|
|
||||||
Please be sure to provide a clear description of what you're trying to achieve with the change.
|
|
||||||
|
|
||||||
- If you're submitting a new feature, please explain how to use it and document any new config options in the example config.
|
|
||||||
- If you're submitting a bugfix, please link the related issue or describe the circumstances surrounding the issue.
|
|
||||||
- If you're changing a default, explain why you believe the new default is appropriate for most users.
|
|
||||||
|
|
||||||
P.S. If you're only updating the README or other docs, please file a pull request here instead: https://github.com/DefinedNet/nebula-docs
|
|
||||||
-->
|
|
||||||
@@ -14,11 +14,11 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Install goimports
|
- name: Install goimports
|
||||||
|
|||||||
@@ -7,24 +7,24 @@ name: Create release and upload binaries
|
|||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
build-linux:
|
build-linux:
|
||||||
name: Build Linux/BSD All
|
name: Build Linux All
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
run: |
|
run: |
|
||||||
make BUILD_NUMBER="${GITHUB_REF#refs/tags/v}" release-linux release-freebsd release-openbsd release-netbsd
|
make BUILD_NUMBER="${GITHUB_REF#refs/tags/v}" release-linux release-freebsd
|
||||||
mkdir release
|
mkdir release
|
||||||
mv build/*.tar.gz release
|
mv build/*.tar.gz release
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v6
|
uses: actions/upload-artifact@v3
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: release
|
path: release
|
||||||
@@ -33,11 +33,11 @@ jobs:
|
|||||||
name: Build Windows
|
name: Build Windows
|
||||||
runs-on: windows-latest
|
runs-on: windows-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -55,7 +55,7 @@ jobs:
|
|||||||
mv dist\windows\wintun build\dist\windows\
|
mv dist\windows\wintun build\dist\windows\
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v6
|
uses: actions/upload-artifact@v3
|
||||||
with:
|
with:
|
||||||
name: windows-latest
|
name: windows-latest
|
||||||
path: build
|
path: build
|
||||||
@@ -64,18 +64,18 @@ jobs:
|
|||||||
name: Build Universal Darwin
|
name: Build Universal Darwin
|
||||||
env:
|
env:
|
||||||
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
|
||||||
runs-on: macos-latest
|
runs-on: macos-11
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Import certificates
|
- name: Import certificates
|
||||||
if: env.HAS_SIGNING_CREDS == 'true'
|
if: env.HAS_SIGNING_CREDS == 'true'
|
||||||
uses: Apple-Actions/import-codesign-certs@v6
|
uses: Apple-Actions/import-codesign-certs@v1
|
||||||
with:
|
with:
|
||||||
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
||||||
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
||||||
@@ -104,66 +104,20 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v6
|
uses: actions/upload-artifact@v3
|
||||||
with:
|
with:
|
||||||
name: darwin-latest
|
name: darwin-latest
|
||||||
path: ./release/*
|
path: ./release/*
|
||||||
|
|
||||||
build-docker:
|
|
||||||
name: Create and Upload Docker Images
|
|
||||||
# Technically we only need build-linux to succeed, but if any platforms fail we'll
|
|
||||||
# want to investigate and restart the build
|
|
||||||
needs: [build-linux, build-darwin, build-windows]
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
env:
|
|
||||||
HAS_DOCKER_CREDS: ${{ vars.DOCKERHUB_USERNAME != '' && secrets.DOCKERHUB_TOKEN != '' }}
|
|
||||||
# XXX It's not possible to write a conditional here, so instead we do it on every step
|
|
||||||
#if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
steps:
|
|
||||||
# Be sure to checkout the code before downloading artifacts, or they will
|
|
||||||
# be overwritten
|
|
||||||
- name: Checkout code
|
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- name: Download artifacts
|
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
uses: actions/download-artifact@v7
|
|
||||||
with:
|
|
||||||
name: linux-latest
|
|
||||||
path: artifacts
|
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
uses: docker/login-action@v3
|
|
||||||
with:
|
|
||||||
username: ${{ vars.DOCKERHUB_USERNAME }}
|
|
||||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
uses: docker/setup-buildx-action@v3
|
|
||||||
|
|
||||||
- name: Build and push images
|
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
|
||||||
env:
|
|
||||||
DOCKER_IMAGE_REPO: ${{ vars.DOCKER_IMAGE_REPO || 'nebulaoss/nebula' }}
|
|
||||||
DOCKER_IMAGE_TAG: ${{ vars.DOCKER_IMAGE_TAG || 'latest' }}
|
|
||||||
run: |
|
|
||||||
mkdir -p build/linux-{amd64,arm64}
|
|
||||||
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
|
|
||||||
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
|
|
||||||
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}"
|
|
||||||
|
|
||||||
release:
|
release:
|
||||||
name: Create and Upload Release
|
name: Create and Upload Release
|
||||||
needs: [build-linux, build-darwin, build-windows]
|
needs: [build-linux, build-darwin, build-windows]
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v7
|
uses: actions/download-artifact@v3
|
||||||
with:
|
with:
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
@@ -209,11 +163,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
|
||||||
|
|||||||
@@ -1,67 +0,0 @@
|
|||||||
name: smoke-extra
|
|
||||||
on:
|
|
||||||
push:
|
|
||||||
branches:
|
|
||||||
- master
|
|
||||||
pull_request:
|
|
||||||
types: [opened, synchronize, labeled, reopened]
|
|
||||||
paths:
|
|
||||||
- '.github/workflows/smoke**'
|
|
||||||
- '**Makefile'
|
|
||||||
- '**.go'
|
|
||||||
- '**.proto'
|
|
||||||
- 'go.mod'
|
|
||||||
- 'go.sum'
|
|
||||||
jobs:
|
|
||||||
|
|
||||||
smoke-extra:
|
|
||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
|
||||||
name: Run extra smoke tests
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
env:
|
|
||||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
|
||||||
with:
|
|
||||||
go-version: '1.25'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- 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
|
|
||||||
|
|
||||||
- name: install vagrant and libvirt
|
|
||||||
run: |
|
|
||||||
sudo apt-get update && sudo apt-get install -y vagrant libvirt-daemon-system libvirt-dev
|
|
||||||
sudo chmod 666 /dev/kvm
|
|
||||||
sudo usermod -aG libvirt $(whoami)
|
|
||||||
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
|
||||||
vagrant plugin install vagrant-libvirt
|
|
||||||
|
|
||||||
- name: freebsd-amd64
|
|
||||||
run: make smoke-vagrant/freebsd-amd64
|
|
||||||
|
|
||||||
- name: openbsd-amd64
|
|
||||||
run: make smoke-vagrant/openbsd-amd64
|
|
||||||
|
|
||||||
- name: netbsd-amd64
|
|
||||||
run: make smoke-vagrant/netbsd-amd64
|
|
||||||
|
|
||||||
- name: 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
|
|
||||||
@@ -18,15 +18,15 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: build
|
- name: build
|
||||||
run: make bin-docker CGO_ENABLED=1 BUILD_ARGS=-race
|
run: make bin-docker
|
||||||
|
|
||||||
- name: setup docker image
|
- name: setup docker image
|
||||||
working-directory: ./.github/workflows/smoke
|
working-directory: ./.github/workflows/smoke
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
@@ -43,4 +41,4 @@ EOF
|
|||||||
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
||||||
)
|
)
|
||||||
|
|
||||||
docker build -t nebula:smoke-relay .
|
sudo docker build -t nebula:smoke-relay .
|
||||||
|
|||||||
@@ -5,42 +5,27 @@ 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
|
|
||||||
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
|
||||||
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
|
||||||
# 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
|
||||||
|
|
||||||
cp ../../../../build/linux-amd64/nebula .
|
cp ../../../../build/linux-amd64/nebula .
|
||||||
cp ../../../../build/linux-amd64/nebula-cert .
|
cp ../../../../build/linux-amd64/nebula-cert .
|
||||||
|
|
||||||
if [ "$1" ]
|
|
||||||
then
|
|
||||||
cp "../../../../build/$1/nebula" "$1-nebula"
|
|
||||||
fi
|
|
||||||
|
|
||||||
HOST="lighthouse1" \
|
HOST="lighthouse1" \
|
||||||
AM_LIGHTHOUSE=true \
|
AM_LIGHTHOUSE=true \
|
||||||
../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 172.17.0.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 172.17.0.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 172.17.0.2:4242" \
|
||||||
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||||
../genconfig.sh >host4.yml
|
../genconfig.sh >host4.yml
|
||||||
|
|
||||||
@@ -51,4 +36,4 @@ LIGHTHOUSE_IP="203.0.113.2"
|
|||||||
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
|
||||||
)
|
)
|
||||||
|
|
||||||
docker build -t "nebula:${NAME:-smoke}" .
|
sudo docker build -t "nebula:${NAME:-smoke}" .
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ listen:
|
|||||||
port: ${LISTEN_PORT:-4242}
|
port: ${LISTEN_PORT:-4242}
|
||||||
|
|
||||||
tun:
|
tun:
|
||||||
dev: ${TUN_DEV:-tun0}
|
dev: ${TUN_DEV:-nebula1}
|
||||||
|
|
||||||
firewall:
|
firewall:
|
||||||
inbound_action: reject
|
inbound_action: reject
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -16,55 +14,24 @@ cleanup() {
|
|||||||
set +e
|
set +e
|
||||||
if [ "$(jobs -r)" ]
|
if [ "$(jobs -r)" ]
|
||||||
then
|
then
|
||||||
docker kill lighthouse1 host2 host3 host4
|
sudo 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
|
sudo docker run --name lighthouse1 --rm nebula:smoke-relay -config lighthouse1.yml -test
|
||||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
sudo docker run --name host2 --rm nebula:smoke-relay -config host2.yml -test
|
||||||
# fail the whole test — we only need one to be free.
|
sudo docker run --name host3 --rm nebula:smoke-relay -config host3.yml -test
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
sudo docker run --name host4 --rm nebula:smoke-relay -config host4.yml -test
|
||||||
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,
|
sudo 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] /' &
|
||||||
# .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 host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" 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 host4 --rm -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" 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] /' &
|
|
||||||
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] /' &
|
sudo 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] /' &
|
sudo 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] /' &
|
sudo 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
|
||||||
@@ -72,50 +39,44 @@ echo
|
|||||||
echo " *** Testing ping from lighthouse1"
|
echo " *** Testing ping from lighthouse1"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.2
|
sudo docker exec lighthouse1 ping -c1 192.168.100.2
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.3
|
sudo docker exec lighthouse1 ping -c1 192.168.100.3
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.4
|
sudo docker exec lighthouse1 ping -c1 192.168.100.4
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host2"
|
echo " *** Testing ping from host2"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host2 ping -c1 192.168.100.1
|
sudo docker exec host2 ping -c1 192.168.100.1
|
||||||
# Should fail because no relay configured in this direction
|
# Should fail because no relay configured in this direction
|
||||||
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
! sudo docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
! docker exec host2 ping -c1 192.168.100.4 -w5 || exit 1
|
! sudo docker exec host2 ping -c1 192.168.100.4 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host3"
|
echo " *** Testing ping from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ping -c1 192.168.100.1
|
sudo docker exec host3 ping -c1 192.168.100.1
|
||||||
docker exec host3 ping -c1 192.168.100.2
|
sudo docker exec host3 ping -c1 192.168.100.2
|
||||||
docker exec host3 ping -c1 192.168.100.4
|
sudo docker exec host3 ping -c1 192.168.100.4
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host4"
|
echo " *** Testing ping from host4"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host4 ping -c1 192.168.100.1
|
sudo docker exec host4 ping -c1 192.168.100.1
|
||||||
# Should fail because relays not allowed
|
# Should fail because relays not allowed
|
||||||
! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
! sudo docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
||||||
docker exec host4 ping -c1 192.168.100.3
|
sudo docker exec host4 ping -c1 192.168.100.3
|
||||||
|
|
||||||
docker exec host4 sh -c 'kill 1'
|
sudo docker exec host4 sh -c 'kill 1'
|
||||||
docker exec host3 sh -c 'kill 1'
|
sudo docker exec host3 sh -c 'kill 1'
|
||||||
docker exec host2 sh -c 'kill 1'
|
sudo docker exec host2 sh -c 'kill 1'
|
||||||
docker exec lighthouse1 sh -c 'kill 1'
|
sudo docker exec lighthouse1 sh -c 'kill 1'
|
||||||
|
sleep 1
|
||||||
# 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,144 +0,0 @@
|
|||||||
#!/bin/bash
|
|
||||||
|
|
||||||
set -e -x
|
|
||||||
|
|
||||||
set -o pipefail
|
|
||||||
|
|
||||||
export VAGRANT_CWD="$PWD/vagrant-$1"
|
|
||||||
|
|
||||||
mkdir -p logs
|
|
||||||
|
|
||||||
NETWORK="nebula-smoke"
|
|
||||||
|
|
||||||
cleanup() {
|
|
||||||
echo
|
|
||||||
echo " *** cleanup"
|
|
||||||
echo
|
|
||||||
|
|
||||||
set +e
|
|
||||||
if [ "$(jobs -r)" ]
|
|
||||||
then
|
|
||||||
docker kill lighthouse1 host2
|
|
||||||
fi
|
|
||||||
vagrant destroy -f
|
|
||||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
|
||||||
}
|
|
||||||
|
|
||||||
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}"
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
vagrant up
|
|
||||||
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] /' &
|
|
||||||
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] /' &
|
|
||||||
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] /' &
|
|
||||||
sleep 15
|
|
||||||
|
|
||||||
# grab tcpdump pcaps for debugging
|
|
||||||
docker exec lighthouse1 tcpdump -i nebula1 -q -w - -U 2>logs/lighthouse1.inside.log >logs/lighthouse1.inside.pcap &
|
|
||||||
docker exec lighthouse1 tcpdump -i eth0 -q -w - -U 2>logs/lighthouse1.outside.log >logs/lighthouse1.outside.pcap &
|
|
||||||
docker exec host2 tcpdump -i nebula1 -q -w - -U 2>logs/host2.inside.log >logs/host2.inside.pcap &
|
|
||||||
docker exec host2 tcpdump -i eth0 -q -w - -U 2>logs/host2.outside.log >logs/host2.outside.pcap &
|
|
||||||
# vagrant ssh -c "tcpdump -i nebula1 -q -w - -U" 2>logs/host3.inside.log >logs/host3.inside.pcap &
|
|
||||||
# vagrant ssh -c "tcpdump -i eth0 -q -w - -U" 2>logs/host3.outside.log >logs/host3.outside.pcap &
|
|
||||||
|
|
||||||
#docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
|
||||||
#vagrant ssh -c "ncat -nklv 0.0.0.0 2000" &
|
|
||||||
#docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
|
||||||
#vagrant ssh -c "ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000" &
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** Testing ping from lighthouse1"
|
|
||||||
echo
|
|
||||||
set -x
|
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.2
|
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.3
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** Testing ping from host2"
|
|
||||||
echo
|
|
||||||
set -x
|
|
||||||
docker exec host2 ping -c1 192.168.100.1
|
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
|
||||||
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
|
||||||
|
|
||||||
#set +x
|
|
||||||
#echo
|
|
||||||
#echo " *** Testing ncat from host2"
|
|
||||||
#echo
|
|
||||||
#set -x
|
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
|
||||||
#! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
|
||||||
#! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
|
||||||
|
|
||||||
set +x
|
|
||||||
echo
|
|
||||||
echo " *** Testing ping from host3"
|
|
||||||
echo
|
|
||||||
set -x
|
|
||||||
vagrant ssh -c "ping -c1 192.168.100.1" -- -T
|
|
||||||
vagrant ssh -c "ping -c1 192.168.100.2" -- -T
|
|
||||||
|
|
||||||
#set +x
|
|
||||||
#echo
|
|
||||||
#echo " *** Testing ncat from host3"
|
|
||||||
#echo
|
|
||||||
#set -x
|
|
||||||
#vagrant ssh -c "ncat -nzv -w5 192.168.100.2 2000"
|
|
||||||
#vagrant ssh -c "ncat -nzuv -w5 192.168.100.2 3000" | grep -q host2
|
|
||||||
|
|
||||||
vagrant ssh -c "sudo xargs kill </nebula/pid" -- -T
|
|
||||||
docker exec host2 sh -c 'kill 1'
|
|
||||||
docker exec lighthouse1 sh -c 'kill 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)" ]
|
|
||||||
then
|
|
||||||
echo "nebula still running after SIGTERM sent" >&2
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
@@ -6,8 +6,6 @@ set -o pipefail
|
|||||||
|
|
||||||
mkdir -p logs
|
mkdir -p logs
|
||||||
|
|
||||||
NETWORK="nebula-smoke"
|
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
echo
|
echo
|
||||||
echo " *** cleanup"
|
echo " *** cleanup"
|
||||||
@@ -16,92 +14,59 @@ cleanup() {
|
|||||||
set +e
|
set +e
|
||||||
if [ "$(jobs -r)" ]
|
if [ "$(jobs -r)" ]
|
||||||
then
|
then
|
||||||
docker kill lighthouse1 host2 host3 host4
|
sudo 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
|
sudo 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
|
sudo 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
|
sudo 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
|
sudo 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] /' &
|
sudo 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] /' &
|
sudo 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] /' &
|
sudo 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] /' &
|
sudo 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
|
||||||
docker exec lighthouse1 tcpdump -i tun0 -q -w - -U 2>logs/lighthouse1.inside.log >logs/lighthouse1.inside.pcap &
|
sudo docker exec lighthouse1 tcpdump -i nebula1 -q -w - -U 2>logs/lighthouse1.inside.log >logs/lighthouse1.inside.pcap &
|
||||||
docker exec lighthouse1 tcpdump -i eth0 -q -w - -U 2>logs/lighthouse1.outside.log >logs/lighthouse1.outside.pcap &
|
sudo docker exec lighthouse1 tcpdump -i eth0 -q -w - -U 2>logs/lighthouse1.outside.log >logs/lighthouse1.outside.pcap &
|
||||||
docker exec host2 tcpdump -i tun0 -q -w - -U 2>logs/host2.inside.log >logs/host2.inside.pcap &
|
sudo docker exec host2 tcpdump -i nebula1 -q -w - -U 2>logs/host2.inside.log >logs/host2.inside.pcap &
|
||||||
docker exec host2 tcpdump -i eth0 -q -w - -U 2>logs/host2.outside.log >logs/host2.outside.pcap &
|
sudo docker exec host2 tcpdump -i eth0 -q -w - -U 2>logs/host2.outside.log >logs/host2.outside.pcap &
|
||||||
docker exec host3 tcpdump -i tun0 -q -w - -U 2>logs/host3.inside.log >logs/host3.inside.pcap &
|
sudo docker exec host3 tcpdump -i nebula1 -q -w - -U 2>logs/host3.inside.log >logs/host3.inside.pcap &
|
||||||
docker exec host3 tcpdump -i eth0 -q -w - -U 2>logs/host3.outside.log >logs/host3.outside.pcap &
|
sudo docker exec host3 tcpdump -i eth0 -q -w - -U 2>logs/host3.outside.log >logs/host3.outside.pcap &
|
||||||
docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
sudo docker exec host4 tcpdump -i nebula1 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap &
|
||||||
docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
sudo docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
|
||||||
|
|
||||||
docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
sudo docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
sudo docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
||||||
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 0.0.0.0 4000 &
|
sudo docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
||||||
docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 &
|
sudo docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
|
||||||
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
|
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from lighthouse1"
|
echo " *** Testing ping from lighthouse1"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.2
|
sudo docker exec lighthouse1 ping -c1 192.168.100.2
|
||||||
docker exec lighthouse1 ping -c1 192.168.100.3
|
sudo docker exec lighthouse1 ping -c1 192.168.100.3
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host2"
|
echo " *** Testing ping from host2"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host2 ping -c1 192.168.100.1
|
sudo docker exec host2 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
! sudo docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -109,34 +74,34 @@ echo " *** Testing ncat from host2"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host3 inbound firewall
|
# Should fail because not allowed by host3 inbound firewall
|
||||||
! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
! sudo docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
||||||
! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
! sudo docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host3"
|
echo " *** Testing ping from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ping -c1 192.168.100.1
|
sudo docker exec host3 ping -c1 192.168.100.1
|
||||||
docker exec host3 ping -c1 192.168.100.2
|
sudo docker exec host3 ping -c1 192.168.100.2
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ncat from host3"
|
echo " *** Testing ncat from host3"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host3 ncat -nzv -w5 192.168.100.2 2000
|
sudo docker exec host3 ncat -nzv -w5 192.168.100.2 2000
|
||||||
docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2
|
sudo docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing ping from host4"
|
echo " *** Testing ping from host4"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
docker exec host4 ping -c1 192.168.100.1
|
sudo docker exec host4 ping -c1 192.168.100.1
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
! sudo docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
|
||||||
! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
|
! sudo docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
@@ -144,34 +109,27 @@ echo " *** Testing ncat from host4"
|
|||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
# Should fail because not allowed by host4 outbound firewall
|
# Should fail because not allowed by host4 outbound firewall
|
||||||
! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
|
! sudo docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
! sudo docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1
|
! sudo docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1
|
||||||
! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
! sudo docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
|
||||||
|
|
||||||
set +x
|
set +x
|
||||||
echo
|
echo
|
||||||
echo " *** Testing conntrack"
|
echo " *** Testing conntrack"
|
||||||
echo
|
echo
|
||||||
set -x
|
set -x
|
||||||
|
# host2 can ping host3 now that host3 pinged it first
|
||||||
|
sudo docker exec host2 ping -c1 192.168.100.3
|
||||||
|
# host4 can ping host2 once conntrack established
|
||||||
|
sudo docker exec host2 ping -c1 192.168.100.4
|
||||||
|
sudo docker exec host4 ping -c1 192.168.100.2
|
||||||
|
|
||||||
# host4's outbound firewall only allows ICMP to the lighthouse, so host4
|
sudo docker exec host4 sh -c 'kill 1'
|
||||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
sudo docker exec host3 sh -c 'kill 1'
|
||||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
sudo docker exec host2 sh -c 'kill 1'
|
||||||
# the echo back from host4 never reaches host2.
|
sudo docker exec lighthouse1 sh -c 'kill 1'
|
||||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
|
sleep 1
|
||||||
|
|
||||||
docker exec host4 sh -c 'kill 1'
|
|
||||||
docker exec host3 sh -c 'kill 1'
|
|
||||||
docker exec host2 sh -c 'kill 1'
|
|
||||||
docker exec lighthouse1 sh -c 'kill 1'
|
|
||||||
|
|
||||||
# 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 +0,0 @@
|
|||||||
# -*- mode: ruby -*-
|
|
||||||
# vi: set ft=ruby :
|
|
||||||
Vagrant.configure("2") do |config|
|
|
||||||
config.vm.box = "generic/freebsd14"
|
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
|
||||||
end
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
# -*- mode: ruby -*-
|
|
||||||
# vi: set ft=ruby :
|
|
||||||
Vagrant.configure("2") do |config|
|
|
||||||
config.vm.box = "ubuntu/xenial32"
|
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula"
|
|
||||||
end
|
|
||||||
@@ -1,16 +0,0 @@
|
|||||||
# -*- mode: ruby -*-
|
|
||||||
# vi: set ft=ruby :
|
|
||||||
Vagrant.configure("2") do |config|
|
|
||||||
config.vm.box = "bento/ubuntu-24.04"
|
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula"
|
|
||||||
|
|
||||||
config.vm.provision :shell do |shell|
|
|
||||||
shell.inline = <<-EOF
|
|
||||||
sed -i 's/GRUB_CMDLINE_LINUX=""/GRUB_CMDLINE_LINUX="ipv6.disable=1"/' /etc/default/grub
|
|
||||||
update-grub
|
|
||||||
EOF
|
|
||||||
shell.privileged = true
|
|
||||||
shell.reboot = true
|
|
||||||
end
|
|
||||||
end
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
# -*- mode: ruby -*-
|
|
||||||
# vi: set ft=ruby :
|
|
||||||
Vagrant.configure("2") do |config|
|
|
||||||
config.vm.box = "generic/netbsd9"
|
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
|
||||||
end
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
# -*- mode: ruby -*-
|
|
||||||
# vi: set ft=ruby :
|
|
||||||
Vagrant.configure("2") do |config|
|
|
||||||
config.vm.box = "DefinedNet/openbsd78"
|
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
|
||||||
end
|
|
||||||
+17
-48
@@ -18,11 +18,11 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -31,24 +31,16 @@ jobs:
|
|||||||
- name: Vet
|
- name: Vet
|
||||||
run: make vet
|
run: make vet
|
||||||
|
|
||||||
- name: golangci-lint
|
|
||||||
uses: golangci/golangci-lint-action@v9
|
|
||||||
with:
|
|
||||||
version: v2.5
|
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: make test
|
run: make test
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
run: make e2evv
|
run: make e2evv
|
||||||
|
|
||||||
- name: Build test mobile
|
- uses: actions/upload-artifact@v3
|
||||||
run: make build-test-mobile
|
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v6
|
|
||||||
with:
|
with:
|
||||||
name: e2e packet flow linux-latest
|
name: e2e packet flow
|
||||||
path: e2e/mermaid/linux-latest
|
path: e2e/mermaid/
|
||||||
if-no-files-found: warn
|
if-no-files-found: warn
|
||||||
|
|
||||||
test-linux-boringcrypto:
|
test-linux-boringcrypto:
|
||||||
@@ -56,11 +48,11 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build
|
- name: Build
|
||||||
@@ -70,39 +62,21 @@ jobs:
|
|||||||
run: make test-boringcrypto
|
run: make test-boringcrypto
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
run: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
|
run: make e2evv GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
|
|
||||||
test-linux-pkcs11:
|
|
||||||
name: Build and test on linux with pkcs11
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
|
||||||
with:
|
|
||||||
go-version: '1.25'
|
|
||||||
check-latest: true
|
|
||||||
|
|
||||||
- name: Build
|
|
||||||
run: make bin-pkcs11
|
|
||||||
|
|
||||||
- name: Test
|
|
||||||
run: make test-pkcs11
|
|
||||||
|
|
||||||
test:
|
test:
|
||||||
name: Build and test on ${{ matrix.os }}
|
name: Build and test on ${{ matrix.os }}
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
strategy:
|
strategy:
|
||||||
matrix:
|
matrix:
|
||||||
os: [windows-latest, macos-latest]
|
os: [windows-latest, macos-11]
|
||||||
steps:
|
steps:
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v3
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
- uses: actions/setup-go@v4
|
||||||
with:
|
with:
|
||||||
go-version: '1.25'
|
go-version-file: 'go.mod'
|
||||||
check-latest: true
|
check-latest: true
|
||||||
|
|
||||||
- name: Build nebula
|
- name: Build nebula
|
||||||
@@ -114,19 +88,14 @@ jobs:
|
|||||||
- name: Vet
|
- name: Vet
|
||||||
run: make vet
|
run: make vet
|
||||||
|
|
||||||
- name: golangci-lint
|
|
||||||
uses: golangci/golangci-lint-action@v9
|
|
||||||
with:
|
|
||||||
version: v2.5
|
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: make test
|
run: make test
|
||||||
|
|
||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
run: make e2evv
|
run: make e2evv
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v6
|
- uses: actions/upload-artifact@v3
|
||||||
with:
|
with:
|
||||||
name: e2e packet flow ${{ matrix.os }}
|
name: e2e packet flow
|
||||||
path: e2e/mermaid/${{ matrix.os }}
|
path: e2e/mermaid/
|
||||||
if-no-files-found: warn
|
if-no-files-found: warn
|
||||||
|
|||||||
+1
-3
@@ -5,8 +5,7 @@
|
|||||||
/nebula-darwin
|
/nebula-darwin
|
||||||
/nebula.exe
|
/nebula.exe
|
||||||
/nebula-cert.exe
|
/nebula-cert.exe
|
||||||
**/coverage.out
|
/coverage.out
|
||||||
**/cover.out
|
|
||||||
/cpu.pprof
|
/cpu.pprof
|
||||||
/build
|
/build
|
||||||
/*.tar.gz
|
/*.tar.gz
|
||||||
@@ -14,6 +13,5 @@
|
|||||||
**.crt
|
**.crt
|
||||||
**.key
|
**.key
|
||||||
**.pem
|
**.pem
|
||||||
**.pub
|
|
||||||
!/examples/quickstart-vagrant/ansible/roles/nebula/files/vagrant-test-ca.key
|
!/examples/quickstart-vagrant/ansible/roles/nebula/files/vagrant-test-ca.key
|
||||||
!/examples/quickstart-vagrant/ansible/roles/nebula/files/vagrant-test-ca.crt
|
!/examples/quickstart-vagrant/ansible/roles/nebula/files/vagrant-test-ca.crt
|
||||||
|
|||||||
@@ -1,37 +0,0 @@
|
|||||||
version: "2"
|
|
||||||
linters:
|
|
||||||
default: none
|
|
||||||
enable:
|
|
||||||
- sloglint
|
|
||||||
- testifylint
|
|
||||||
settings:
|
|
||||||
sloglint:
|
|
||||||
# Enforce key-value pair form for Info/Debug/Warn/Error/Log/With and
|
|
||||||
# the package-level slog equivalents. Use l.Log(ctx, level, ...) for
|
|
||||||
# custom levels instead of LogAttrs when you can.
|
|
||||||
#
|
|
||||||
# LogAttrs is also flagged by this rule because it takes ...slog.Attr;
|
|
||||||
# the few legitimate sites (where attrs is built up as a []slog.Attr)
|
|
||||||
# carry a //nolint:sloglint with rationale.
|
|
||||||
kv-only: true
|
|
||||||
# no-mixed-args is on by default: forbids mixing kv and attrs in one call.
|
|
||||||
# discard-handler is on by default (since Go 1.24): suggests
|
|
||||||
# slog.DiscardHandler over slog.NewTextHandler(io.Discard, nil).
|
|
||||||
exclusions:
|
|
||||||
generated: lax
|
|
||||||
presets:
|
|
||||||
- comments
|
|
||||||
- common-false-positives
|
|
||||||
- legacy
|
|
||||||
- std-error-handling
|
|
||||||
paths:
|
|
||||||
- third_party$
|
|
||||||
- builtin$
|
|
||||||
- examples$
|
|
||||||
formatters:
|
|
||||||
exclusions:
|
|
||||||
generated: lax
|
|
||||||
paths:
|
|
||||||
- third_party$
|
|
||||||
- builtin$
|
|
||||||
- examples$
|
|
||||||
+1
-316
@@ -7,306 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
## [1.10.3] - 2026-02-06
|
|
||||||
|
|
||||||
### Security
|
|
||||||
|
|
||||||
- Fix an issue where blocklist bypass is possible when using curve P256 since the signature can have 2 valid representations.
|
|
||||||
Both fingerprint representations will be tested against the blocklist.
|
|
||||||
Any newly issued P256 based certificates will have their signature clamped to the low-s form.
|
|
||||||
Nebula will assert the low-s signature form when validating certificates in a future version. [GHSA-69x3-g4r3-p962](https://github.com/slackhq/nebula/security/advisories/GHSA-69x3-g4r3-p962)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Improve error reporting if nebula fails to start due to a tun device naming issue. (#1588)
|
|
||||||
|
|
||||||
## [1.10.2] - 2026-01-21
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix panic when using `use_system_route_table` that was introduced in v1.10.1. (#1580)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Fix some typos in comments. (#1582)
|
|
||||||
- Dependency updates. (#1581)
|
|
||||||
|
|
||||||
## [1.10.1] - 2026-01-16
|
|
||||||
|
|
||||||
See the [v1.10.1](https://github.com/slackhq/nebula/milestone/26?closed=1) milestone for a complete list of changes.
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix a bug where an unsafe route derived from the system route table could be lost on a config reload. (#1573)
|
|
||||||
- Fix the PEM banner for ECDSA P256 public keys. (#1552)
|
|
||||||
- Fix a regression on Windows from 1.9.x where nebula could fall back to a less performant UDP listener if
|
|
||||||
non-critical ioctls failed. (#1568)
|
|
||||||
- Fix a bug in handshake processing when a peer sends an unexpected public key. (#1566)
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Add a config option to control accepting `recv_error` packets which defaults to `always`. (#1569)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Various dependency updates. (#1541, #1549, #1550, #1557, #1558, #1560, #1561, #1570, #1571)
|
|
||||||
|
|
||||||
## [1.10.0] - 2025-12-04
|
|
||||||
|
|
||||||
See the [v1.10.0](https://github.com/slackhq/nebula/milestone/16?closed=1) milestone for a complete list of changes.
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Support for ipv6 and multiple ipv4/6 addresses in the overlay.
|
|
||||||
A new v2 ASN.1 based certificate format.
|
|
||||||
Certificates now have a unified interface for external implementations.
|
|
||||||
(#1212, #1216, #1345, #1359, #1381, #1419, #1464, #1466, #1451, #1476, #1467, #1481, #1399, #1488, #1492, #1495, #1468, #1521, #1535, #1538)
|
|
||||||
- Add the ability to mark packets on linux to better target nebula packets in iptables/nftables. (#1331)
|
|
||||||
- Add ECMP support for `unsafe_routes`. (#1332)
|
|
||||||
- PKCS11 support for P256 keys when built with `pkcs11` tag (#1153, #1482)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- **NOTE**: `default_local_cidr_any` now defaults to false, meaning that any firewall rule
|
|
||||||
intended to target an `unsafe_routes` entry must explicitly declare it via the
|
|
||||||
`local_cidr` field. This is almost always the intended behavior. This flag is
|
|
||||||
deprecated and will be removed in a future release. (#1373)
|
|
||||||
- Improve logging when a relay is in use on an inbound packet. (#1533)
|
|
||||||
- Avoid fatal errors if `rountines` is > 1 on systems that don't support more than 1 routine. (#1531)
|
|
||||||
- Log a warning if a firewall rule contains an `any` that negates a more restrictive filter. (#1513)
|
|
||||||
- Accept encrypted CA passphrase from an environment variable. (#1421)
|
|
||||||
- Allow handshaking with any trusted remote. (#1509)
|
|
||||||
- Log only the count of blocklisted certificate fingerprints instead of the entire list. (#1525)
|
|
||||||
- Don't fatal when the ssh server is unable to be configured successfully. (#1520)
|
|
||||||
- Update to build against go v1.25. (#1483)
|
|
||||||
- Allow projects using `nebula` as a library with userspace networking to configure the `logger` and build version. (#1239)
|
|
||||||
- Upgrade to `yaml.v3`. (#1148, #1371, #1438, #1478)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix a potential bug with udp ipv4 only on darwin. (#1532)
|
|
||||||
- Improve lost packet statistics. (#1441, #1537)
|
|
||||||
- Honor `remote_allow_list` in hole punch response. (#1186)
|
|
||||||
- Fix a panic when `tun.use_system_route_table` is `true` and a route lacks a destination. (#1437)
|
|
||||||
- Fix an issue when `tun.use_system_route_table: true` could result in heavy CPU utilization when many thousands of routes
|
|
||||||
are present. (#1326)
|
|
||||||
- Fix tests for 32 bit machines. (#1394)
|
|
||||||
- Fix a possible 32bit integer underflow in config handling. (#1353)
|
|
||||||
- Fix moving a udp address from one vpn address to another in the `static_host_map`
|
|
||||||
which could cause rapid re-handshaking with an incorrect remote. (#1259)
|
|
||||||
- Improve smoke tests in environments where the docker network is not the default. (#1347)
|
|
||||||
|
|
||||||
## [1.9.7] - 2025-10-10
|
|
||||||
|
|
||||||
### Security
|
|
||||||
|
|
||||||
- Fix an issue where Nebula could incorrectly accept and process a packet from an erroneous source IP when the sender's
|
|
||||||
certificate is configured with unsafe_routes (cert v1/v2) or multiple IPs (cert v2). (#1494)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Disable sending `recv_error` messages when a packet is received outside the allowable counter window. (#1459)
|
|
||||||
- Improve error messages and remove some unnecessary fatal conditions in the Windows and generic udp listener. (#1453)
|
|
||||||
|
|
||||||
## [1.9.6] - 2025-7-15
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Support dropping inactive tunnels. This is disabled by default in this release but can be enabled with `tunnels.drop_inactive`. See example config for more details. (#1413)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix Darwin freeze due to presence of some Network Extensions (#1426)
|
|
||||||
- Ensure the same relay tunnel is always used when multiple relay tunnels are present (#1422)
|
|
||||||
- Fix Windows freeze due to ICMP error handling (#1412)
|
|
||||||
- Fix relay migration panic (#1403)
|
|
||||||
|
|
||||||
## [1.9.5] - 2024-12-05
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Gracefully ignore v2 certificates. (#1282)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix relays that refuse to re-establish after one of the remote tunnel pairs breaks. (#1277)
|
|
||||||
|
|
||||||
## [1.9.4] - 2024-09-09
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Support UDP dialing with gVisor. (#1181)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Make some Nebula state programmatically available via control object. (#1188)
|
|
||||||
- Switch internal representation of IPs to netip, to prepare for IPv6 support
|
|
||||||
in the overlay. (#1173)
|
|
||||||
- Minor build and cleanup changes. (#1171, #1164, #1162)
|
|
||||||
- Various dependency updates. (#1195, #1190, #1174, #1168, #1167, #1161, #1147, #1146)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix a bug on big endian hosts, like mips. (#1194)
|
|
||||||
- Fix a rare panic if a local index collision happens. (#1191)
|
|
||||||
- Fix integer wraparound in the calculation of handshake timeouts on 32-bit targets. (#1185)
|
|
||||||
|
|
||||||
## [1.9.3] - 2024-06-06
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Initialize messageCounter to 2 instead of verifying later. (#1156)
|
|
||||||
|
|
||||||
## [1.9.2] - 2024-06-03
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Ensure messageCounter is set before handshake is complete. (#1154)
|
|
||||||
|
|
||||||
## [1.9.1] - 2024-05-29
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fixed a potential deadlock in GetOrHandshake. (#1151)
|
|
||||||
|
|
||||||
## [1.9.0] - 2024-05-07
|
|
||||||
|
|
||||||
### Deprecated
|
|
||||||
|
|
||||||
- This release adds a new setting `default_local_cidr_any` that defaults to
|
|
||||||
true to match previous behavior, but will default to false in the next
|
|
||||||
release (1.10). When set to false, `local_cidr` is matched correctly for
|
|
||||||
firewall rules on hosts acting as unsafe routers, and should be set for any
|
|
||||||
firewall rules you want to allow unsafe route hosts to access. See the issue
|
|
||||||
and example config for more details. (#1071, #1099)
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Nebula now has an official Docker image `nebulaoss/nebula` that is
|
|
||||||
distroless and contains just the `nebula` and `nebula-cert` binaries. You
|
|
||||||
can find it here: https://hub.docker.com/r/nebulaoss/nebula (#1037)
|
|
||||||
|
|
||||||
- Experimental binaries for `loong64` are now provided. (#1003)
|
|
||||||
|
|
||||||
- Added example service script for OpenRC. (#711)
|
|
||||||
|
|
||||||
- The SSH daemon now supports inlined host keys. (#1054)
|
|
||||||
|
|
||||||
- The SSH daemon now supports certificates with `sshd.trusted_cas`. (#1098)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Config setting `tun.unsafe_routes` is now reloadable. (#1083)
|
|
||||||
|
|
||||||
- Small documentation and internal improvements. (#1065, #1067, #1069, #1108,
|
|
||||||
#1109, #1111, #1135)
|
|
||||||
|
|
||||||
- Various dependency updates. (#1139, #1138, #1134, #1133, #1126, #1123, #1110,
|
|
||||||
#1094, #1092, #1087, #1086, #1085, #1072, #1063, #1059, #1055, #1053, #1047,
|
|
||||||
#1046, #1034, #1022)
|
|
||||||
|
|
||||||
### Removed
|
|
||||||
|
|
||||||
- Support for the deprecated `local_range` option has been removed. Please
|
|
||||||
change to `preferred_ranges` (which is also now reloadable). (#1043)
|
|
||||||
|
|
||||||
- We are now building with go1.22, which means that for Windows you need at
|
|
||||||
least Windows 10 or Windows Server 2016. This is because support for earlier
|
|
||||||
versions was removed in Go 1.21. See https://go.dev/doc/go1.21#windows (#981)
|
|
||||||
|
|
||||||
- Removed vagrant example, as it was unmaintained. (#1129)
|
|
||||||
|
|
||||||
- Removed Fedora and Arch nebula.service files, as they are maintained in the
|
|
||||||
upstream repos. (#1128, #1132)
|
|
||||||
|
|
||||||
- Remove the TCP round trip tracking metrics, as they never had correct data
|
|
||||||
and were an experiment to begin with. (#1114)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fixed a potential deadlock introduced in 1.8.1. (#1112)
|
|
||||||
|
|
||||||
- Fixed support for Linux when IPv6 has been disabled at the OS level. (#787)
|
|
||||||
|
|
||||||
- DNS will return NXDOMAIN now when there are no results. (#845)
|
|
||||||
|
|
||||||
- Allow `::` in `lighthouse.dns.host`. (#1115)
|
|
||||||
|
|
||||||
- Capitalization of `NotAfter` fixed in DNS TXT response. (#1127)
|
|
||||||
|
|
||||||
- Don't log invalid certificates. It is untrusted data and can cause a large
|
|
||||||
volume of logs. (#1116)
|
|
||||||
|
|
||||||
## [1.8.2] - 2024-01-08
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix multiple routines when listen.port is zero. This was a regression
|
|
||||||
introduced in v1.6.0. (#1057)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- Small dependency update for Noise. (#1038)
|
|
||||||
|
|
||||||
## [1.8.1] - 2023-12-19
|
|
||||||
|
|
||||||
### Security
|
|
||||||
|
|
||||||
- Update `golang.org/x/crypto`, which includes a fix for CVE-2023-48795. (#1048)
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Fix a deadlock introduced in v1.8.0 that could occur during handshakes. (#1044)
|
|
||||||
|
|
||||||
- Fix mobile builds. (#1035)
|
|
||||||
|
|
||||||
## [1.8.0] - 2023-12-06
|
|
||||||
|
|
||||||
### Deprecated
|
|
||||||
|
|
||||||
- The next minor release of Nebula, 1.9.0, will require at least Windows 10 or
|
|
||||||
Windows Server 2016. This is because support for earlier versions was removed
|
|
||||||
in Go 1.21. See https://go.dev/doc/go1.21#windows
|
|
||||||
|
|
||||||
### Added
|
|
||||||
|
|
||||||
- Linux: Notify systemd of service readiness. This should resolve timing issues
|
|
||||||
with services that depend on Nebula being active. For an example of how to
|
|
||||||
enable this, see: `examples/service_scripts/nebula.service`. (#929)
|
|
||||||
|
|
||||||
- Windows: Use Registered IO (RIO) when possible. Testing on a Windows 11
|
|
||||||
machine shows ~50x improvement in throughput. (#905)
|
|
||||||
|
|
||||||
- NetBSD, OpenBSD: Added rudimentary support. (#916, #812)
|
|
||||||
|
|
||||||
- FreeBSD: Add support for naming tun devices. (#903)
|
|
||||||
|
|
||||||
### Changed
|
|
||||||
|
|
||||||
- `pki.disconnect_invalid` will now default to true. This means that once a
|
|
||||||
certificate expires, the tunnel will be disconnected. If you use SIGHUP to
|
|
||||||
reload certificates without restarting Nebula, you should ensure all of your
|
|
||||||
clients are on 1.7.0 or newer before you enable this feature. (#859)
|
|
||||||
|
|
||||||
- Limit how often a busy tunnel can requery the lighthouse. The new config
|
|
||||||
option `timers.requery_wait_duration` defaults to `60s`. (#940)
|
|
||||||
|
|
||||||
- The internal structures for hostmaps were refactored to reduce memory usage
|
|
||||||
and the potential for subtle bugs. (#843, #938, #953, #954, #955)
|
|
||||||
|
|
||||||
- Lots of dependency updates.
|
|
||||||
|
|
||||||
### Fixed
|
|
||||||
|
|
||||||
- Windows: Retry wintun device creation if it fails the first time. (#985)
|
|
||||||
|
|
||||||
- Fix issues with firewall reject packets that could cause panics. (#957)
|
|
||||||
|
|
||||||
- Fix relay migration during re-handshakes. (#964)
|
|
||||||
|
|
||||||
- Various other refactors and fixes. (#935, #952, #972, #961, #996, #1002,
|
|
||||||
#987, #1004, #1030, #1032, ...)
|
|
||||||
|
|
||||||
## [1.7.2] - 2023-06-01
|
## [1.7.2] - 2023-06-01
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
@@ -788,22 +488,7 @@ created.)
|
|||||||
|
|
||||||
- Initial public release.
|
- Initial public release.
|
||||||
|
|
||||||
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.10.3...HEAD
|
[Unreleased]: https://github.com/slackhq/nebula/compare/v1.7.2...HEAD
|
||||||
[1.10.3]: https://github.com/slackhq/nebula/releases/tag/v1.10.3
|
|
||||||
[1.10.2]: https://github.com/slackhq/nebula/releases/tag/v1.10.2
|
|
||||||
[1.10.1]: https://github.com/slackhq/nebula/releases/tag/v1.10.1
|
|
||||||
[1.10.0]: https://github.com/slackhq/nebula/releases/tag/v1.10.0
|
|
||||||
[1.9.7]: https://github.com/slackhq/nebula/releases/tag/v1.9.7
|
|
||||||
[1.9.6]: https://github.com/slackhq/nebula/releases/tag/v1.9.6
|
|
||||||
[1.9.5]: https://github.com/slackhq/nebula/releases/tag/v1.9.5
|
|
||||||
[1.9.4]: https://github.com/slackhq/nebula/releases/tag/v1.9.4
|
|
||||||
[1.9.3]: https://github.com/slackhq/nebula/releases/tag/v1.9.3
|
|
||||||
[1.9.2]: https://github.com/slackhq/nebula/releases/tag/v1.9.2
|
|
||||||
[1.9.1]: https://github.com/slackhq/nebula/releases/tag/v1.9.1
|
|
||||||
[1.9.0]: https://github.com/slackhq/nebula/releases/tag/v1.9.0
|
|
||||||
[1.8.2]: https://github.com/slackhq/nebula/releases/tag/v1.8.2
|
|
||||||
[1.8.1]: https://github.com/slackhq/nebula/releases/tag/v1.8.1
|
|
||||||
[1.8.0]: https://github.com/slackhq/nebula/releases/tag/v1.8.0
|
|
||||||
[1.7.2]: https://github.com/slackhq/nebula/releases/tag/v1.7.2
|
[1.7.2]: https://github.com/slackhq/nebula/releases/tag/v1.7.2
|
||||||
[1.7.1]: https://github.com/slackhq/nebula/releases/tag/v1.7.1
|
[1.7.1]: https://github.com/slackhq/nebula/releases/tag/v1.7.1
|
||||||
[1.7.0]: https://github.com/slackhq/nebula/releases/tag/v1.7.0
|
[1.7.0]: https://github.com/slackhq/nebula/releases/tag/v1.7.0
|
||||||
|
|||||||
@@ -1 +0,0 @@
|
|||||||
#ECCN:Open Source
|
|
||||||
@@ -33,5 +33,6 @@ l.WithError(err).
|
|||||||
WithField("vpnIp", IntIp(hostinfo.hostId)).
|
WithField("vpnIp", IntIp(hostinfo.hostId)).
|
||||||
WithField("udpAddr", addr).
|
WithField("udpAddr", addr).
|
||||||
WithField("handshake", m{"stage": 1, "style": "ix"}).
|
WithField("handshake", m{"stage": 1, "style": "ix"}).
|
||||||
|
WithField("cert", remoteCert).
|
||||||
Info("Invalid certificate from host")
|
Info("Invalid certificate from host")
|
||||||
```
|
```
|
||||||
@@ -1,14 +1,20 @@
|
|||||||
|
GOMINVERSION = 1.20
|
||||||
NEBULA_CMD_PATH = "./cmd/nebula"
|
NEBULA_CMD_PATH = "./cmd/nebula"
|
||||||
|
GO111MODULE = on
|
||||||
|
export GO111MODULE
|
||||||
CGO_ENABLED = 0
|
CGO_ENABLED = 0
|
||||||
export CGO_ENABLED
|
export CGO_ENABLED
|
||||||
|
|
||||||
# Set up OS specific bits
|
# Set up OS specific bits
|
||||||
ifeq ($(OS),Windows_NT)
|
ifeq ($(OS),Windows_NT)
|
||||||
|
#TODO: we should be able to ditch awk as well
|
||||||
|
GOVERSION := $(shell go version | awk "{print substr($$3, 3)}")
|
||||||
|
GOISMIN := $(shell IF "$(GOVERSION)" GEQ "$(GOMINVERSION)" ECHO 1)
|
||||||
NEBULA_CMD_SUFFIX = .exe
|
NEBULA_CMD_SUFFIX = .exe
|
||||||
NULL_FILE = nul
|
NULL_FILE = nul
|
||||||
# RIO on windows does pointer stuff that makes go vet angry
|
|
||||||
VET_FLAGS = -unsafeptr=false
|
|
||||||
else
|
else
|
||||||
|
GOVERSION := $(shell go version | awk '{print substr($$3, 3)}')
|
||||||
|
GOISMIN := $(shell expr "$(GOVERSION)" ">=" "$(GOMINVERSION)")
|
||||||
NEBULA_CMD_SUFFIX =
|
NEBULA_CMD_SUFFIX =
|
||||||
NULL_FILE = /dev/null
|
NULL_FILE = /dev/null
|
||||||
endif
|
endif
|
||||||
@@ -22,9 +28,6 @@ ifndef BUILD_NUMBER
|
|||||||
endif
|
endif
|
||||||
endif
|
endif
|
||||||
|
|
||||||
DOCKER_IMAGE_REPO ?= nebulaoss/nebula
|
|
||||||
DOCKER_IMAGE_TAG ?= latest
|
|
||||||
|
|
||||||
LDFLAGS = -X main.Build=$(BUILD_NUMBER)
|
LDFLAGS = -X main.Build=$(BUILD_NUMBER)
|
||||||
|
|
||||||
ALL_LINUX = linux-amd64 \
|
ALL_LINUX = linux-amd64 \
|
||||||
@@ -39,22 +42,13 @@ ALL_LINUX = linux-amd64 \
|
|||||||
linux-mips64 \
|
linux-mips64 \
|
||||||
linux-mips64le \
|
linux-mips64le \
|
||||||
linux-mips-softfloat \
|
linux-mips-softfloat \
|
||||||
linux-riscv64 \
|
linux-riscv64
|
||||||
linux-loong64
|
|
||||||
|
|
||||||
ALL_FREEBSD = freebsd-amd64 \
|
ALL_FREEBSD = freebsd-amd64 \
|
||||||
freebsd-arm64
|
freebsd-arm64
|
||||||
|
|
||||||
ALL_OPENBSD = openbsd-amd64 \
|
|
||||||
openbsd-arm64
|
|
||||||
|
|
||||||
ALL_NETBSD = netbsd-amd64 \
|
|
||||||
netbsd-arm64
|
|
||||||
|
|
||||||
ALL = $(ALL_LINUX) \
|
ALL = $(ALL_LINUX) \
|
||||||
$(ALL_FREEBSD) \
|
$(ALL_FREEBSD) \
|
||||||
$(ALL_OPENBSD) \
|
|
||||||
$(ALL_NETBSD) \
|
|
||||||
darwin-amd64 \
|
darwin-amd64 \
|
||||||
darwin-arm64 \
|
darwin-arm64 \
|
||||||
windows-amd64 \
|
windows-amd64 \
|
||||||
@@ -63,7 +57,7 @@ ALL = $(ALL_LINUX) \
|
|||||||
e2e:
|
e2e:
|
||||||
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
|
||||||
|
|
||||||
e2ev: TEST_FLAGS += -v
|
e2ev: TEST_FLAGS = -v
|
||||||
e2ev: e2e
|
e2ev: e2e
|
||||||
|
|
||||||
e2evv: TEST_ENV += TEST_LOGS=1
|
e2evv: TEST_ENV += TEST_LOGS=1
|
||||||
@@ -78,25 +72,17 @@ e2evvvv: e2ev
|
|||||||
e2e-bench: TEST_FLAGS = -bench=. -benchmem -run=^$
|
e2e-bench: TEST_FLAGS = -bench=. -benchmem -run=^$
|
||||||
e2e-bench: e2e
|
e2e-bench: e2e
|
||||||
|
|
||||||
DOCKER_BIN = build/linux-amd64/nebula build/linux-amd64/nebula-cert
|
|
||||||
|
|
||||||
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
|
||||||
|
|
||||||
docker: docker/linux-$(shell go env GOARCH)
|
|
||||||
|
|
||||||
release: $(ALL:%=build/nebula-%.tar.gz)
|
release: $(ALL:%=build/nebula-%.tar.gz)
|
||||||
|
|
||||||
release-linux: $(ALL_LINUX:%=build/nebula-%.tar.gz)
|
release-linux: $(ALL_LINUX:%=build/nebula-%.tar.gz)
|
||||||
|
|
||||||
release-freebsd: $(ALL_FREEBSD:%=build/nebula-%.tar.gz)
|
release-freebsd: $(ALL_FREEBSD:%=build/nebula-%.tar.gz)
|
||||||
|
|
||||||
release-openbsd: $(ALL_OPENBSD:%=build/nebula-%.tar.gz)
|
|
||||||
|
|
||||||
release-netbsd: $(ALL_NETBSD:%=build/nebula-%.tar.gz)
|
|
||||||
|
|
||||||
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
release-boringcrypto: build/nebula-linux-$(shell go env GOARCH)-boringcrypto.tar.gz
|
||||||
|
|
||||||
BUILD_ARGS += -trimpath
|
BUILD_ARGS = -trimpath
|
||||||
|
|
||||||
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
bin-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe
|
||||||
mv $? .
|
mv $? .
|
||||||
@@ -116,10 +102,6 @@ bin-freebsd-arm64: build/freebsd-arm64/nebula build/freebsd-arm64/nebula-cert
|
|||||||
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
bin-boringcrypto: build/linux-$(shell go env GOARCH)-boringcrypto/nebula build/linux-$(shell go env GOARCH)-boringcrypto/nebula-cert
|
||||||
mv $? .
|
mv $? .
|
||||||
|
|
||||||
bin-pkcs11: BUILD_ARGS += -tags pkcs11
|
|
||||||
bin-pkcs11: CGO_ENABLED = 1
|
|
||||||
bin-pkcs11: bin
|
|
||||||
|
|
||||||
bin:
|
bin:
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
|
||||||
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
|
||||||
@@ -137,8 +119,6 @@ build/linux-mips-softfloat/%: LDFLAGS += -s -w
|
|||||||
# boringcrypto
|
# boringcrypto
|
||||||
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-amd64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
build/linux-arm64-boringcrypto/%: GOENV += GOEXPERIMENT=boringcrypto CGO_ENABLED=1
|
||||||
build/linux-amd64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
|
||||||
build/linux-arm64-boringcrypto/%: LDFLAGS += -checklinkname=0
|
|
||||||
|
|
||||||
build/%/nebula: .FORCE
|
build/%/nebula: .FORCE
|
||||||
GOOS=$(firstword $(subst -, , $*)) \
|
GOOS=$(firstword $(subst -, , $*)) \
|
||||||
@@ -162,31 +142,19 @@ build/nebula-%.tar.gz: build/%/nebula build/%/nebula-cert
|
|||||||
build/nebula-%.zip: build/%/nebula.exe build/%/nebula-cert.exe
|
build/nebula-%.zip: build/%/nebula.exe build/%/nebula-cert.exe
|
||||||
cd build/$* && zip ../nebula-$*.zip nebula.exe nebula-cert.exe
|
cd build/$* && zip ../nebula-$*.zip nebula.exe nebula-cert.exe
|
||||||
|
|
||||||
docker/%: build/%/nebula build/%/nebula-cert
|
|
||||||
docker build . $(DOCKER_BUILD_ARGS) -f docker/Dockerfile --platform "$(subst -,/,$*)" --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:$(BUILD_NUMBER)"
|
|
||||||
|
|
||||||
vet:
|
vet:
|
||||||
go vet $(VET_FLAGS) -v ./...
|
go vet -v ./...
|
||||||
|
|
||||||
test:
|
test:
|
||||||
go test -v ./...
|
go test -v ./...
|
||||||
|
|
||||||
test-boringcrypto:
|
test-boringcrypto:
|
||||||
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -ldflags "-checklinkname=0" -v ./...
|
GOEXPERIMENT=boringcrypto CGO_ENABLED=1 go test -v ./...
|
||||||
|
|
||||||
test-pkcs11:
|
|
||||||
CGO_ENABLED=1 go test -v -tags pkcs11 ./...
|
|
||||||
|
|
||||||
test-cov-html:
|
test-cov-html:
|
||||||
go test -coverprofile=coverage.out
|
go test -coverprofile=coverage.out
|
||||||
go tool cover -html=coverage.out
|
go tool cover -html=coverage.out
|
||||||
|
|
||||||
build-test-mobile:
|
|
||||||
GOARCH=amd64 GOOS=ios go build $(shell go list ./... | grep -v '/cmd/\|/examples/')
|
|
||||||
GOARCH=arm64 GOOS=ios go build $(shell go list ./... | grep -v '/cmd/\|/examples/')
|
|
||||||
GOARCH=amd64 GOOS=android go build $(shell go list ./... | grep -v '/cmd/\|/examples/')
|
|
||||||
GOARCH=arm64 GOOS=android go build $(shell go list ./... | grep -v '/cmd/\|/examples/')
|
|
||||||
|
|
||||||
bench:
|
bench:
|
||||||
go test -bench=.
|
go test -bench=.
|
||||||
|
|
||||||
@@ -198,7 +166,7 @@ bench-cpu-long:
|
|||||||
go test -bench=. -benchtime=60s -cpuprofile=cpu.pprof
|
go test -bench=. -benchtime=60s -cpuprofile=cpu.pprof
|
||||||
go tool pprof go-audit.test cpu.pprof
|
go tool pprof go-audit.test cpu.pprof
|
||||||
|
|
||||||
proto: nebula.pb.go cert/cert_v1.pb.go
|
proto: nebula.pb.go cert/cert.pb.go
|
||||||
|
|
||||||
nebula.pb.go: nebula.proto .FORCE
|
nebula.pb.go: nebula.proto .FORCE
|
||||||
go build github.com/gogo/protobuf/protoc-gen-gogofaster
|
go build github.com/gogo/protobuf/protoc-gen-gogofaster
|
||||||
@@ -228,13 +196,8 @@ smoke-relay-docker: bin-docker
|
|||||||
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
cd .github/workflows/smoke/ && ./smoke-relay.sh
|
||||||
|
|
||||||
smoke-docker-race: BUILD_ARGS = -race
|
smoke-docker-race: BUILD_ARGS = -race
|
||||||
smoke-docker-race: CGO_ENABLED = 1
|
|
||||||
smoke-docker-race: smoke-docker
|
smoke-docker-race: smoke-docker
|
||||||
|
|
||||||
smoke-vagrant/%: bin-docker build/%/nebula
|
|
||||||
cd .github/workflows/smoke/ && ./build.sh $*
|
|
||||||
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
|
|
||||||
|
|
||||||
.FORCE:
|
.FORCE:
|
||||||
.PHONY: bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
|
.PHONY: e2e e2ev e2evv e2evvv e2evvvv test test-cov-html bench bench-cpu bench-cpu-long bin proto release service smoke-docker smoke-docker-race
|
||||||
.DEFAULT_GOAL := bin
|
.DEFAULT_GOAL := bin
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ It lets you seamlessly connect computers anywhere in the world. Nebula is portab
|
|||||||
It can be used to connect a small number of computers, but is also able to connect tens of thousands of computers.
|
It can be used to connect a small number of computers, but is also able to connect tens of thousands of computers.
|
||||||
|
|
||||||
Nebula incorporates a number of existing concepts like encryption, security groups, certificates,
|
Nebula incorporates a number of existing concepts like encryption, security groups, certificates,
|
||||||
and tunneling.
|
and tunneling, and each of those individual pieces existed before Nebula in various forms.
|
||||||
What makes Nebula different to existing offerings is that it brings all of these ideas together,
|
What makes Nebula different to existing offerings is that it brings all of these ideas together,
|
||||||
resulting in a sum that is greater than its individual parts.
|
resulting in a sum that is greater than its individual parts.
|
||||||
|
|
||||||
@@ -12,7 +12,7 @@ Further documentation can be found [here](https://nebula.defined.net/docs/).
|
|||||||
|
|
||||||
You can read more about Nebula [here](https://medium.com/p/884110a5579).
|
You can read more about Nebula [here](https://medium.com/p/884110a5579).
|
||||||
|
|
||||||
You can also join the NebulaOSS Slack group [here](https://join.slack.com/t/nebulaoss/shared_invite/zt-39pk4xopc-CUKlGcb5Z39dQ0cK1v7ehA).
|
You can also join the NebulaOSS Slack group [here](https://join.slack.com/t/nebulaoss/shared_invite/enQtOTA5MDI4NDg3MTg4LTkwY2EwNTI4NzQyMzc0M2ZlODBjNWI3NTY1MzhiOThiMmZlZjVkMTI0NGY4YTMyNjUwMWEyNzNkZTJmYzQxOGU).
|
||||||
|
|
||||||
## Supported Platforms
|
## Supported Platforms
|
||||||
|
|
||||||
@@ -27,47 +27,31 @@ Check the [releases](https://github.com/slackhq/nebula/releases/latest) page for
|
|||||||
|
|
||||||
#### Distribution Packages
|
#### Distribution Packages
|
||||||
|
|
||||||
- [Arch Linux](https://archlinux.org/packages/extra/x86_64/nebula/)
|
- [Arch Linux](https://archlinux.org/packages/community/x86_64/nebula/)
|
||||||
```sh
|
```
|
||||||
sudo pacman -S nebula
|
$ sudo pacman -S nebula
|
||||||
```
|
```
|
||||||
|
|
||||||
- [Fedora Linux](https://src.fedoraproject.org/rpms/nebula)
|
- [Fedora Linux](https://src.fedoraproject.org/rpms/nebula)
|
||||||
```sh
|
```
|
||||||
sudo dnf install nebula
|
$ sudo dnf install nebula
|
||||||
```
|
```
|
||||||
|
|
||||||
- [Debian Linux](https://packages.debian.org/source/stable/nebula)
|
- [macOS Homebrew](https://github.com/Homebrew/homebrew-core/blob/HEAD/Formula/nebula.rb)
|
||||||
```sh
|
```
|
||||||
sudo apt install nebula
|
$ brew install nebula
|
||||||
```
|
```
|
||||||
|
|
||||||
- [Alpine Linux](https://pkgs.alpinelinux.org/packages?name=nebula)
|
#### Mobile
|
||||||
```sh
|
|
||||||
sudo apk add nebula
|
|
||||||
```
|
|
||||||
|
|
||||||
- [macOS Homebrew](https://github.com/Homebrew/homebrew-core/blob/HEAD/Formula/n/nebula.rb)
|
|
||||||
```sh
|
|
||||||
brew install nebula
|
|
||||||
```
|
|
||||||
|
|
||||||
- [Docker](https://hub.docker.com/r/nebulaoss/nebula)
|
|
||||||
```sh
|
|
||||||
docker pull nebulaoss/nebula
|
|
||||||
```
|
|
||||||
|
|
||||||
#### Mobile ([source code](https://github.com/DefinedNet/mobile_nebula))
|
|
||||||
|
|
||||||
- [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)
|
||||||
|
|
||||||
## Technical Overview
|
## Technical Overview
|
||||||
|
|
||||||
Nebula is a mutually authenticated peer-to-peer software-defined network based on the [Noise Protocol Framework](https://noiseprotocol.org/).
|
Nebula is a mutually authenticated peer-to-peer software defined network based on the [Noise Protocol Framework](https://noiseprotocol.org/).
|
||||||
Nebula uses certificates to assert a node's IP address, name, and membership within user-defined groups.
|
Nebula uses certificates to assert a node's IP address, name, and membership within user-defined groups.
|
||||||
Nebula's user-defined groups allow for provider agnostic traffic filtering between nodes.
|
Nebula's user-defined groups allow for provider agnostic traffic filtering between nodes.
|
||||||
Discovery nodes (aka lighthouses) allow individual peers to find each other and optionally use UDP hole punching to establish connections from behind most firewalls or NATs.
|
Discovery nodes allow individual peers to find each other and optionally use UDP hole punching to establish connections from behind most firewalls or NATs.
|
||||||
Users can move data between nodes in any number of cloud service providers, datacenters, and endpoints, without needing to maintain a particular addressing scheme.
|
Users can move data between nodes in any number of cloud service providers, datacenters, and endpoints, without needing to maintain a particular addressing scheme.
|
||||||
|
|
||||||
Nebula uses Elliptic-curve Diffie-Hellman (`ECDH`) key exchange and `AES-256-GCM` in its default configuration.
|
Nebula uses Elliptic-curve Diffie-Hellman (`ECDH`) key exchange and `AES-256-GCM` in its default configuration.
|
||||||
@@ -76,42 +60,34 @@ 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.
|
||||||
|
|
||||||
#### 2. (Optional, but you really should..) At least one discovery node with a routable IP address, which we call a lighthouse.
|
#### 2. (Optional, but you really should..) At least one discovery node with a routable IP address, which we call a lighthouse.
|
||||||
|
|
||||||
Nebula lighthouses allow nodes to find each other, anywhere in the world. A lighthouse is the only node in a Nebula network whose IP should not change. Running a lighthouse requires very few compute resources, and you can easily use the least expensive option from a cloud hosting provider. If you're not sure which provider to use, a number of us have used $6/mo [DigitalOcean](https://digitalocean.com) droplets as lighthouses.
|
Nebula lighthouses allow nodes to find each other, anywhere in the world. A lighthouse is the only node in a Nebula network whose IP should not change. Running a lighthouse requires very few compute resources, and you can easily use the least expensive option from a cloud hosting provider. If you're not sure which provider to use, a number of us have used $5/mo [DigitalOcean](https://digitalocean.com) droplets as lighthouses.
|
||||||
|
|
||||||
|
Once you have launched an instance, ensure that Nebula udp traffic (default port udp/4242) can reach it over the internet.
|
||||||
|
|
||||||
Once you have launched an instance, ensure that Nebula udp traffic (default port udp/4242) can reach it over the internet.
|
|
||||||
|
|
||||||
#### 3. A Nebula certificate authority, which will be the root of trust for a particular Nebula network.
|
#### 3. A Nebula certificate authority, which will be the root of trust for a particular Nebula network.
|
||||||
|
|
||||||
```sh
|
```
|
||||||
./nebula-cert ca -name "Myorganization, Inc"
|
./nebula-cert ca -name "Myorganization, Inc"
|
||||||
```
|
```
|
||||||
|
This will create files named `ca.key` and `ca.cert` in the current directory. The `ca.key` file is the most sensitive file you'll create, because it is the key used to sign the certificates for individual nebula nodes/hosts. Please store this file somewhere safe, preferably with strong encryption.
|
||||||
This will create files named `ca.key` and `ca.cert` in the current directory. The `ca.key` file is the most sensitive file you'll create, because it is the key used to sign the certificates for individual nebula nodes/hosts. Please store this file somewhere safe, preferably with strong encryption.
|
|
||||||
|
|
||||||
**Be aware!** By default, certificate authorities have a 1-year lifetime before expiration. See [this guide](https://nebula.defined.net/docs/guides/rotating-certificate-authority/) for details on rotating a CA.
|
|
||||||
|
|
||||||
#### 4. Nebula host keys and certificates generated from that certificate authority
|
#### 4. Nebula host keys and certificates generated from that certificate authority
|
||||||
|
|
||||||
This assumes you have four nodes, named lighthouse1, laptop, server1, host3. You can name the nodes any way you'd like, including FQDN. You'll also need to choose IP addresses and the associated subnet. In this example, we are creating a nebula network that will use 192.168.100.x/24 as its network range. This example also demonstrates nebula groups, which can later be used to define traffic rules in a nebula network.
|
This assumes you have four nodes, named lighthouse1, laptop, server1, host3. You can name the nodes any way you'd like, including FQDN. You'll also need to choose IP addresses and the associated subnet. In this example, we are creating a nebula network that will use 192.168.100.x/24 as its network range. This example also demonstrates nebula groups, which can later be used to define traffic rules in a nebula network.
|
||||||
```sh
|
```
|
||||||
./nebula-cert sign -name "lighthouse1" -ip "192.168.100.1/24"
|
./nebula-cert sign -name "lighthouse1" -ip "192.168.100.1/24"
|
||||||
./nebula-cert sign -name "laptop" -ip "192.168.100.2/24" -groups "laptop,home,ssh"
|
./nebula-cert sign -name "laptop" -ip "192.168.100.2/24" -groups "laptop,home,ssh"
|
||||||
./nebula-cert sign -name "server1" -ip "192.168.100.9/24" -groups "servers"
|
./nebula-cert sign -name "server1" -ip "192.168.100.9/24" -groups "servers"
|
||||||
./nebula-cert sign -name "host3" -ip "192.168.100.10/24"
|
./nebula-cert sign -name "host3" -ip "192.168.100.10/24"
|
||||||
```
|
```
|
||||||
|
|
||||||
By default, host certificates will expire 1 second before the CA expires. Use the `-duration` flag to specify a shorter lifetime.
|
|
||||||
|
|
||||||
#### 5. Configuration files for each host
|
#### 5. Configuration files for each host
|
||||||
|
|
||||||
Download a copy of the nebula [example configuration](https://github.com/slackhq/nebula/blob/master/examples/config.yml).
|
Download a copy of the nebula [example configuration](https://github.com/slackhq/nebula/blob/master/examples/config.yml).
|
||||||
|
|
||||||
* On the lighthouse node, you'll need to ensure `am_lighthouse: true` is set.
|
* On the lighthouse node, you'll need to ensure `am_lighthouse: true` is set.
|
||||||
@@ -126,16 +102,13 @@ For each host, copy the nebula binary to the host, along with `config.yml` from
|
|||||||
**DO NOT COPY `ca.key` TO INDIVIDUAL NODES.**
|
**DO NOT COPY `ca.key` TO INDIVIDUAL NODES.**
|
||||||
|
|
||||||
#### 7. Run nebula on each host
|
#### 7. Run nebula on each host
|
||||||
|
```
|
||||||
```sh
|
|
||||||
./nebula -config /path/to/config.yml
|
./nebula -config /path/to/config.yml
|
||||||
```
|
```
|
||||||
|
|
||||||
For more detailed instructions, [find the full documentation here](https://nebula.defined.net/docs/).
|
|
||||||
|
|
||||||
## Building Nebula from source
|
## Building Nebula from source
|
||||||
|
|
||||||
Make sure you have [go](https://go.dev/doc/install) installed and clone this repo. Change to the nebula directory.
|
Download go and clone this repo. Change to the nebula directory.
|
||||||
|
|
||||||
To build nebula for all platforms:
|
To build nebula for all platforms:
|
||||||
`make all`
|
`make all`
|
||||||
@@ -151,10 +124,8 @@ The default curve used for cryptographic handshakes and signatures is Curve25519
|
|||||||
|
|
||||||
In addition, Nebula can be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets:
|
In addition, Nebula can be built using the [BoringCrypto GOEXPERIMENT](https://github.com/golang/go/blob/go1.20/src/crypto/internal/boring/README.md) by running either of the following make targets:
|
||||||
|
|
||||||
```sh
|
make bin-boringcrypto
|
||||||
make bin-boringcrypto
|
make release-boringcrypto
|
||||||
make release-boringcrypto
|
|
||||||
```
|
|
||||||
|
|
||||||
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
This is not the recommended default deployment, but may be useful based on your compliance requirements.
|
||||||
|
|
||||||
@@ -162,3 +133,5 @@ This is not the recommended default deployment, but may be useful based on your
|
|||||||
|
|
||||||
Nebula was created at Slack Technologies, Inc by Nate Brown and Ryan Huber, with contributions from Oliver Fross, Alan Lam, Wade Simmons, and Lining Wang.
|
Nebula was created at Slack Technologies, Inc by Nate Brown and Ryan Huber, with contributions from Oliver Fross, Alan Lam, Wade Simmons, and Lining Wang.
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+110
-55
@@ -2,16 +2,17 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/netip"
|
"net"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/slackhq/nebula/cidr"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
)
|
)
|
||||||
|
|
||||||
type AllowList struct {
|
type AllowList struct {
|
||||||
// The values of this cidrTree are `bool`, signifying allow/deny
|
// The values of this cidrTree are `bool`, signifying allow/deny
|
||||||
cidrTree *bart.Table[bool]
|
cidrTree *cidr.Tree6
|
||||||
}
|
}
|
||||||
|
|
||||||
type RemoteAllowList struct {
|
type RemoteAllowList struct {
|
||||||
@@ -19,7 +20,7 @@ type RemoteAllowList struct {
|
|||||||
|
|
||||||
// Inside Range Specific, keys of this tree are inside CIDRs and values
|
// Inside Range Specific, keys of this tree are inside CIDRs and values
|
||||||
// are *AllowList
|
// are *AllowList
|
||||||
insideAllowLists *bart.Table[*AllowList]
|
insideAllowLists *cidr.Tree6
|
||||||
}
|
}
|
||||||
|
|
||||||
type LocalAllowList struct {
|
type LocalAllowList struct {
|
||||||
@@ -36,7 +37,7 @@ type AllowListNameRule struct {
|
|||||||
|
|
||||||
func NewLocalAllowListFromConfig(c *config.C, k string) (*LocalAllowList, error) {
|
func NewLocalAllowListFromConfig(c *config.C, k string) (*LocalAllowList, error) {
|
||||||
var nameRules []AllowListNameRule
|
var nameRules []AllowListNameRule
|
||||||
handleKey := func(key string, value any) (bool, error) {
|
handleKey := func(key string, value interface{}) (bool, error) {
|
||||||
if key == "interfaces" {
|
if key == "interfaces" {
|
||||||
var err error
|
var err error
|
||||||
nameRules, err = getAllowListInterfaces(k, value)
|
nameRules, err = getAllowListInterfaces(k, value)
|
||||||
@@ -70,7 +71,7 @@ func NewRemoteAllowListFromConfig(c *config.C, k, rangesKey string) (*RemoteAllo
|
|||||||
|
|
||||||
// If the handleKey func returns true, the rest of the parsing is skipped
|
// If the handleKey func returns true, the rest of the parsing is skipped
|
||||||
// for this key. This allows parsing of special values like `interfaces`.
|
// for this key. This allows parsing of special values like `interfaces`.
|
||||||
func newAllowListFromConfig(c *config.C, k string, handleKey func(key string, value any) (bool, error)) (*AllowList, error) {
|
func newAllowListFromConfig(c *config.C, k string, handleKey func(key string, value interface{}) (bool, error)) (*AllowList, error) {
|
||||||
r := c.Get(k)
|
r := c.Get(k)
|
||||||
if r == nil {
|
if r == nil {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
@@ -81,13 +82,13 @@ func newAllowListFromConfig(c *config.C, k string, handleKey func(key string, va
|
|||||||
|
|
||||||
// If the handleKey func returns true, the rest of the parsing is skipped
|
// If the handleKey func returns true, the rest of the parsing is skipped
|
||||||
// for this key. This allows parsing of special values like `interfaces`.
|
// for this key. This allows parsing of special values like `interfaces`.
|
||||||
func newAllowList(k string, raw any, handleKey func(key string, value any) (bool, error)) (*AllowList, error) {
|
func newAllowList(k string, raw interface{}, handleKey func(key string, value interface{}) (bool, error)) (*AllowList, error) {
|
||||||
rawMap, ok := raw.(map[string]any)
|
rawMap, ok := raw.(map[interface{}]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, raw)
|
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, raw)
|
||||||
}
|
}
|
||||||
|
|
||||||
tree := new(bart.Table[bool])
|
tree := cidr.NewTree6()
|
||||||
|
|
||||||
// Keep track of the rules we have added for both ipv4 and ipv6
|
// Keep track of the rules we have added for both ipv4 and ipv6
|
||||||
type allowListRules struct {
|
type allowListRules struct {
|
||||||
@@ -100,7 +101,12 @@ func newAllowList(k string, raw any, handleKey func(key string, value any) (bool
|
|||||||
rules4 := allowListRules{firstValue: true, allValuesMatch: true, defaultSet: false}
|
rules4 := allowListRules{firstValue: true, allValuesMatch: true, defaultSet: false}
|
||||||
rules6 := allowListRules{firstValue: true, allValuesMatch: true, defaultSet: false}
|
rules6 := allowListRules{firstValue: true, allValuesMatch: true, defaultSet: false}
|
||||||
|
|
||||||
for rawCIDR, rawValue := range rawMap {
|
for rawKey, rawValue := range rawMap {
|
||||||
|
rawCIDR, ok := rawKey.(string)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("config `%s` has invalid key (type %T): %v", k, rawKey, rawKey)
|
||||||
|
}
|
||||||
|
|
||||||
if handleKey != nil {
|
if handleKey != nil {
|
||||||
handled, err := handleKey(rawCIDR, rawValue)
|
handled, err := handleKey(rawCIDR, rawValue)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -111,24 +117,23 @@ func newAllowList(k string, raw any, handleKey func(key string, value any) (bool
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
value, ok := config.AsBool(rawValue)
|
value, ok := rawValue.(bool)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid value (type %T): %v", k, rawValue, rawValue)
|
return nil, fmt.Errorf("config `%s` has invalid value (type %T): %v", k, rawValue, rawValue)
|
||||||
}
|
}
|
||||||
|
|
||||||
ipNet, err := netip.ParsePrefix(rawCIDR)
|
_, ipNet, err := net.ParseCIDR(rawCIDR)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s. %w", k, rawCIDR, err)
|
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s", k, rawCIDR)
|
||||||
}
|
}
|
||||||
|
|
||||||
ipNet = netip.PrefixFrom(ipNet.Addr().Unmap(), ipNet.Bits())
|
// TODO: should we error on duplicate CIDRs in the config?
|
||||||
|
tree.AddCIDR(ipNet, value)
|
||||||
|
|
||||||
tree.Insert(ipNet, value)
|
maskBits, maskSize := ipNet.Mask.Size()
|
||||||
|
|
||||||
maskBits := ipNet.Bits()
|
|
||||||
|
|
||||||
var rules *allowListRules
|
var rules *allowListRules
|
||||||
if ipNet.Addr().Is4() {
|
if maskSize == 32 {
|
||||||
rules = &rules4
|
rules = &rules4
|
||||||
} else {
|
} else {
|
||||||
rules = &rules6
|
rules = &rules6
|
||||||
@@ -151,7 +156,8 @@ func newAllowList(k string, raw any, handleKey func(key string, value any) (bool
|
|||||||
|
|
||||||
if !rules4.defaultSet {
|
if !rules4.defaultSet {
|
||||||
if rules4.allValuesMatch {
|
if rules4.allValuesMatch {
|
||||||
tree.Insert(netip.PrefixFrom(netip.IPv4Unspecified(), 0), !rules4.allValues)
|
_, zeroCIDR, _ := net.ParseCIDR("0.0.0.0/0")
|
||||||
|
tree.AddCIDR(zeroCIDR, !rules4.allValues)
|
||||||
} else {
|
} else {
|
||||||
return nil, fmt.Errorf("config `%s` contains both true and false rules, but no default set for 0.0.0.0/0", k)
|
return nil, fmt.Errorf("config `%s` contains both true and false rules, but no default set for 0.0.0.0/0", k)
|
||||||
}
|
}
|
||||||
@@ -159,7 +165,8 @@ func newAllowList(k string, raw any, handleKey func(key string, value any) (bool
|
|||||||
|
|
||||||
if !rules6.defaultSet {
|
if !rules6.defaultSet {
|
||||||
if rules6.allValuesMatch {
|
if rules6.allValuesMatch {
|
||||||
tree.Insert(netip.PrefixFrom(netip.IPv6Unspecified(), 0), !rules6.allValues)
|
_, zeroCIDR, _ := net.ParseCIDR("::/0")
|
||||||
|
tree.AddCIDR(zeroCIDR, !rules6.allValues)
|
||||||
} else {
|
} else {
|
||||||
return nil, fmt.Errorf("config `%s` contains both true and false rules, but no default set for ::/0", k)
|
return nil, fmt.Errorf("config `%s` contains both true and false rules, but no default set for ::/0", k)
|
||||||
}
|
}
|
||||||
@@ -168,18 +175,22 @@ func newAllowList(k string, raw any, handleKey func(key string, value any) (bool
|
|||||||
return &AllowList{cidrTree: tree}, nil
|
return &AllowList{cidrTree: tree}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func getAllowListInterfaces(k string, v any) ([]AllowListNameRule, error) {
|
func getAllowListInterfaces(k string, v interface{}) ([]AllowListNameRule, error) {
|
||||||
var nameRules []AllowListNameRule
|
var nameRules []AllowListNameRule
|
||||||
|
|
||||||
rawRules, ok := v.(map[string]any)
|
rawRules, ok := v.(map[interface{}]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s.interfaces` is invalid (type %T): %v", k, v, v)
|
return nil, fmt.Errorf("config `%s.interfaces` is invalid (type %T): %v", k, v, v)
|
||||||
}
|
}
|
||||||
|
|
||||||
firstEntry := true
|
firstEntry := true
|
||||||
var allValues bool
|
var allValues bool
|
||||||
for name, rawAllow := range rawRules {
|
for rawName, rawAllow := range rawRules {
|
||||||
allow, ok := config.AsBool(rawAllow)
|
name, ok := rawName.(string)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("config `%s.interfaces` has invalid key (type %T): %v", k, rawName, rawName)
|
||||||
|
}
|
||||||
|
allow, ok := rawAllow.(bool)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s.interfaces` has invalid value (type %T): %v", k, rawAllow, rawAllow)
|
return nil, fmt.Errorf("config `%s.interfaces` has invalid value (type %T): %v", k, rawAllow, rawAllow)
|
||||||
}
|
}
|
||||||
@@ -207,49 +218,87 @@ func getAllowListInterfaces(k string, v any) ([]AllowListNameRule, error) {
|
|||||||
return nameRules, nil
|
return nameRules, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func getRemoteAllowRanges(c *config.C, k string) (*bart.Table[*AllowList], error) {
|
func getRemoteAllowRanges(c *config.C, k string) (*cidr.Tree6, error) {
|
||||||
value := c.Get(k)
|
value := c.Get(k)
|
||||||
if value == nil {
|
if value == nil {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
remoteAllowRanges := new(bart.Table[*AllowList])
|
remoteAllowRanges := cidr.NewTree6()
|
||||||
|
|
||||||
rawMap, ok := value.(map[string]any)
|
rawMap, ok := value.(map[interface{}]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, value)
|
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, value)
|
||||||
}
|
}
|
||||||
for rawCIDR, rawValue := range rawMap {
|
for rawKey, rawValue := range rawMap {
|
||||||
|
rawCIDR, ok := rawKey.(string)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("config `%s` has invalid key (type %T): %v", k, rawKey, rawKey)
|
||||||
|
}
|
||||||
|
|
||||||
allowList, err := newAllowList(fmt.Sprintf("%s.%s", k, rawCIDR), rawValue, nil)
|
allowList, err := newAllowList(fmt.Sprintf("%s.%s", k, rawCIDR), rawValue, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
ipNet, err := netip.ParsePrefix(rawCIDR)
|
_, ipNet, err := net.ParseCIDR(rawCIDR)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s. %w", k, rawCIDR, err)
|
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s", k, rawCIDR)
|
||||||
}
|
}
|
||||||
|
|
||||||
remoteAllowRanges.Insert(netip.PrefixFrom(ipNet.Addr().Unmap(), ipNet.Bits()), allowList)
|
remoteAllowRanges.AddCIDR(ipNet, allowList)
|
||||||
}
|
}
|
||||||
|
|
||||||
return remoteAllowRanges, nil
|
return remoteAllowRanges, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *AllowList) Allow(addr netip.Addr) bool {
|
func (al *AllowList) Allow(ip net.IP) bool {
|
||||||
if al == nil {
|
if al == nil {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
result, _ := al.cidrTree.Lookup(addr)
|
result := al.cidrTree.MostSpecificContains(ip)
|
||||||
return result
|
switch v := result.(type) {
|
||||||
|
case bool:
|
||||||
|
return v
|
||||||
|
default:
|
||||||
|
panic(fmt.Errorf("invalid state, allowlist returned: %T %v", result, result))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *LocalAllowList) Allow(udpAddr netip.Addr) bool {
|
func (al *AllowList) AllowIpV4(ip iputil.VpnIp) bool {
|
||||||
if al == nil {
|
if al == nil {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
return al.AllowList.Allow(udpAddr)
|
|
||||||
|
result := al.cidrTree.MostSpecificContainsIpV4(ip)
|
||||||
|
switch v := result.(type) {
|
||||||
|
case bool:
|
||||||
|
return v
|
||||||
|
default:
|
||||||
|
panic(fmt.Errorf("invalid state, allowlist returned: %T %v", result, result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *AllowList) AllowIpV6(hi, lo uint64) bool {
|
||||||
|
if al == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
result := al.cidrTree.MostSpecificContainsIpV6(hi, lo)
|
||||||
|
switch v := result.(type) {
|
||||||
|
case bool:
|
||||||
|
return v
|
||||||
|
default:
|
||||||
|
panic(fmt.Errorf("invalid state, allowlist returned: %T %v", result, result))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *LocalAllowList) Allow(ip net.IP) bool {
|
||||||
|
if al == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return al.AllowList.Allow(ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *LocalAllowList) AllowName(name string) bool {
|
func (al *LocalAllowList) AllowName(name string) bool {
|
||||||
@@ -267,39 +316,45 @@ func (al *LocalAllowList) AllowName(name string) bool {
|
|||||||
return !al.nameRules[0].Allow
|
return !al.nameRules[0].Allow
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *RemoteAllowList) AllowUnknownVpnAddr(vpnAddr netip.Addr) bool {
|
func (al *RemoteAllowList) AllowUnknownVpnIp(ip net.IP) bool {
|
||||||
if al == nil {
|
if al == nil {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
return al.AllowList.Allow(vpnAddr)
|
return al.AllowList.Allow(ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *RemoteAllowList) Allow(vpnAddr netip.Addr, udpAddr netip.Addr) bool {
|
func (al *RemoteAllowList) Allow(vpnIp iputil.VpnIp, ip net.IP) bool {
|
||||||
if !al.getInsideAllowList(vpnAddr).Allow(udpAddr) {
|
if !al.getInsideAllowList(vpnIp).Allow(ip) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
return al.AllowList.Allow(udpAddr)
|
return al.AllowList.Allow(ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *RemoteAllowList) AllowAll(vpnAddrs []netip.Addr, udpAddr netip.Addr) bool {
|
func (al *RemoteAllowList) AllowIpV4(vpnIp iputil.VpnIp, ip iputil.VpnIp) bool {
|
||||||
if !al.AllowList.Allow(udpAddr) {
|
if al == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if !al.getInsideAllowList(vpnIp).AllowIpV4(ip) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
return al.AllowList.AllowIpV4(ip)
|
||||||
for _, vpnAddr := range vpnAddrs {
|
|
||||||
if !al.getInsideAllowList(vpnAddr).Allow(udpAddr) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (al *RemoteAllowList) getInsideAllowList(vpnAddr netip.Addr) *AllowList {
|
func (al *RemoteAllowList) AllowIpV6(vpnIp iputil.VpnIp, hi, lo uint64) bool {
|
||||||
|
if al == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if !al.getInsideAllowList(vpnIp).AllowIpV6(hi, lo) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return al.AllowList.AllowIpV6(hi, lo)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (al *RemoteAllowList) getInsideAllowList(vpnIp iputil.VpnIp) *AllowList {
|
||||||
if al.insideAllowLists != nil {
|
if al.insideAllowLists != nil {
|
||||||
inside, ok := al.insideAllowLists.Lookup(vpnAddr)
|
inside := al.insideAllowLists.MostSpecificContainsIpV4(vpnIp)
|
||||||
if ok {
|
if inside != nil {
|
||||||
return inside
|
return inside.(*AllowList)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
+44
-45
@@ -1,41 +1,40 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net"
|
||||||
"regexp"
|
"regexp"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/slackhq/nebula/cidr"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewAllowListFromConfig(t *testing.T) {
|
func TestNewAllowListFromConfig(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"192.168.0.0": true,
|
"192.168.0.0": true,
|
||||||
}
|
}
|
||||||
r, err := newAllowListFromConfig(c, "allowlist", nil)
|
r, err := newAllowListFromConfig(c, "allowlist", nil)
|
||||||
require.EqualError(t, err, "config `allowlist` has invalid CIDR: 192.168.0.0. netip.ParsePrefix(\"192.168.0.0\"): no '/'")
|
assert.EqualError(t, err, "config `allowlist` has invalid CIDR: 192.168.0.0")
|
||||||
assert.Nil(t, r)
|
assert.Nil(t, r)
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"192.168.0.0/16": "abc",
|
"192.168.0.0/16": "abc",
|
||||||
}
|
}
|
||||||
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
||||||
require.EqualError(t, err, "config `allowlist` has invalid value (type string): abc")
|
assert.EqualError(t, err, "config `allowlist` has invalid value (type string): abc")
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"192.168.0.0/16": true,
|
"192.168.0.0/16": true,
|
||||||
"10.0.0.0/8": false,
|
"10.0.0.0/8": false,
|
||||||
}
|
}
|
||||||
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
||||||
require.EqualError(t, err, "config `allowlist` contains both true and false rules, but no default set for 0.0.0.0/0")
|
assert.EqualError(t, err, "config `allowlist` contains both true and false rules, but no default set for 0.0.0.0/0")
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"0.0.0.0/0": true,
|
"0.0.0.0/0": true,
|
||||||
"10.0.0.0/8": false,
|
"10.0.0.0/8": false,
|
||||||
"10.42.42.0/24": true,
|
"10.42.42.0/24": true,
|
||||||
@@ -43,9 +42,9 @@ func TestNewAllowListFromConfig(t *testing.T) {
|
|||||||
"fd00:fd00::/16": false,
|
"fd00:fd00::/16": false,
|
||||||
}
|
}
|
||||||
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
r, err = newAllowListFromConfig(c, "allowlist", nil)
|
||||||
require.EqualError(t, err, "config `allowlist` contains both true and false rules, but no default set for ::/0")
|
assert.EqualError(t, err, "config `allowlist` contains both true and false rules, but no default set for ::/0")
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"0.0.0.0/0": true,
|
"0.0.0.0/0": true,
|
||||||
"10.0.0.0/8": false,
|
"10.0.0.0/8": false,
|
||||||
"10.42.42.0/24": true,
|
"10.42.42.0/24": true,
|
||||||
@@ -55,7 +54,7 @@ func TestNewAllowListFromConfig(t *testing.T) {
|
|||||||
assert.NotNil(t, r)
|
assert.NotNil(t, r)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"0.0.0.0/0": true,
|
"0.0.0.0/0": true,
|
||||||
"10.0.0.0/8": false,
|
"10.0.0.0/8": false,
|
||||||
"10.42.42.0/24": true,
|
"10.42.42.0/24": true,
|
||||||
@@ -70,25 +69,25 @@ func TestNewAllowListFromConfig(t *testing.T) {
|
|||||||
|
|
||||||
// Test interface names
|
// Test interface names
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"interfaces": map[string]any{
|
"interfaces": map[interface{}]interface{}{
|
||||||
`docker.*`: "foo",
|
`docker.*`: "foo",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
lr, err := NewLocalAllowListFromConfig(c, "allowlist")
|
lr, err := NewLocalAllowListFromConfig(c, "allowlist")
|
||||||
require.EqualError(t, err, "config `allowlist.interfaces` has invalid value (type string): foo")
|
assert.EqualError(t, err, "config `allowlist.interfaces` has invalid value (type string): foo")
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"interfaces": map[string]any{
|
"interfaces": map[interface{}]interface{}{
|
||||||
`docker.*`: false,
|
`docker.*`: false,
|
||||||
`eth.*`: true,
|
`eth.*`: true,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
lr, err = NewLocalAllowListFromConfig(c, "allowlist")
|
lr, err = NewLocalAllowListFromConfig(c, "allowlist")
|
||||||
require.EqualError(t, err, "config `allowlist.interfaces` values must all be the same true/false value")
|
assert.EqualError(t, err, "config `allowlist.interfaces` values must all be the same true/false value")
|
||||||
|
|
||||||
c.Settings["allowlist"] = map[string]any{
|
c.Settings["allowlist"] = map[interface{}]interface{}{
|
||||||
"interfaces": map[string]any{
|
"interfaces": map[interface{}]interface{}{
|
||||||
`docker.*`: false,
|
`docker.*`: false,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -99,30 +98,30 @@ func TestNewAllowListFromConfig(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowList_Allow(t *testing.T) {
|
func TestAllowList_Allow(t *testing.T) {
|
||||||
assert.True(t, ((*AllowList)(nil)).Allow(netip.MustParseAddr("1.1.1.1")))
|
assert.Equal(t, true, ((*AllowList)(nil)).Allow(net.ParseIP("1.1.1.1")))
|
||||||
|
|
||||||
tree := new(bart.Table[bool])
|
tree := cidr.NewTree6()
|
||||||
tree.Insert(netip.MustParsePrefix("0.0.0.0/0"), true)
|
tree.AddCIDR(cidr.Parse("0.0.0.0/0"), true)
|
||||||
tree.Insert(netip.MustParsePrefix("10.0.0.0/8"), false)
|
tree.AddCIDR(cidr.Parse("10.0.0.0/8"), false)
|
||||||
tree.Insert(netip.MustParsePrefix("10.42.42.42/32"), true)
|
tree.AddCIDR(cidr.Parse("10.42.42.42/32"), true)
|
||||||
tree.Insert(netip.MustParsePrefix("10.42.0.0/16"), true)
|
tree.AddCIDR(cidr.Parse("10.42.0.0/16"), true)
|
||||||
tree.Insert(netip.MustParsePrefix("10.42.42.0/24"), true)
|
tree.AddCIDR(cidr.Parse("10.42.42.0/24"), true)
|
||||||
tree.Insert(netip.MustParsePrefix("10.42.42.0/24"), false)
|
tree.AddCIDR(cidr.Parse("10.42.42.0/24"), false)
|
||||||
tree.Insert(netip.MustParsePrefix("::1/128"), true)
|
tree.AddCIDR(cidr.Parse("::1/128"), true)
|
||||||
tree.Insert(netip.MustParsePrefix("::2/128"), false)
|
tree.AddCIDR(cidr.Parse("::2/128"), false)
|
||||||
al := &AllowList{cidrTree: tree}
|
al := &AllowList{cidrTree: tree}
|
||||||
|
|
||||||
assert.True(t, al.Allow(netip.MustParseAddr("1.1.1.1")))
|
assert.Equal(t, true, al.Allow(net.ParseIP("1.1.1.1")))
|
||||||
assert.False(t, al.Allow(netip.MustParseAddr("10.0.0.4")))
|
assert.Equal(t, false, al.Allow(net.ParseIP("10.0.0.4")))
|
||||||
assert.True(t, al.Allow(netip.MustParseAddr("10.42.42.42")))
|
assert.Equal(t, true, al.Allow(net.ParseIP("10.42.42.42")))
|
||||||
assert.False(t, al.Allow(netip.MustParseAddr("10.42.42.41")))
|
assert.Equal(t, false, al.Allow(net.ParseIP("10.42.42.41")))
|
||||||
assert.True(t, al.Allow(netip.MustParseAddr("10.42.0.1")))
|
assert.Equal(t, true, al.Allow(net.ParseIP("10.42.0.1")))
|
||||||
assert.True(t, al.Allow(netip.MustParseAddr("::1")))
|
assert.Equal(t, true, al.Allow(net.ParseIP("::1")))
|
||||||
assert.False(t, al.Allow(netip.MustParseAddr("::2")))
|
assert.Equal(t, false, al.Allow(net.ParseIP("::2")))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLocalAllowList_AllowName(t *testing.T) {
|
func TestLocalAllowList_AllowName(t *testing.T) {
|
||||||
assert.True(t, ((*LocalAllowList)(nil)).AllowName("docker0"))
|
assert.Equal(t, true, ((*LocalAllowList)(nil)).AllowName("docker0"))
|
||||||
|
|
||||||
rules := []AllowListNameRule{
|
rules := []AllowListNameRule{
|
||||||
{Name: regexp.MustCompile("^docker.*$"), Allow: false},
|
{Name: regexp.MustCompile("^docker.*$"), Allow: false},
|
||||||
@@ -130,9 +129,9 @@ func TestLocalAllowList_AllowName(t *testing.T) {
|
|||||||
}
|
}
|
||||||
al := &LocalAllowList{nameRules: rules}
|
al := &LocalAllowList{nameRules: rules}
|
||||||
|
|
||||||
assert.False(t, al.AllowName("docker0"))
|
assert.Equal(t, false, al.AllowName("docker0"))
|
||||||
assert.False(t, al.AllowName("tun0"))
|
assert.Equal(t, false, al.AllowName("tun0"))
|
||||||
assert.True(t, al.AllowName("eth0"))
|
assert.Equal(t, true, al.AllowName("eth0"))
|
||||||
|
|
||||||
rules = []AllowListNameRule{
|
rules = []AllowListNameRule{
|
||||||
{Name: regexp.MustCompile("^eth.*$"), Allow: true},
|
{Name: regexp.MustCompile("^eth.*$"), Allow: true},
|
||||||
@@ -140,7 +139,7 @@ func TestLocalAllowList_AllowName(t *testing.T) {
|
|||||||
}
|
}
|
||||||
al = &LocalAllowList{nameRules: rules}
|
al = &LocalAllowList{nameRules: rules}
|
||||||
|
|
||||||
assert.False(t, al.AllowName("docker0"))
|
assert.Equal(t, false, al.AllowName("docker0"))
|
||||||
assert.True(t, al.AllowName("eth0"))
|
assert.Equal(t, true, al.AllowName("eth0"))
|
||||||
assert.True(t, al.AllowName("ens5"))
|
assert.Equal(t, true, al.AllowName("ens5"))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,23 +1,22 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Bits struct {
|
type Bits struct {
|
||||||
length uint64
|
length uint64
|
||||||
current uint64
|
current uint64
|
||||||
bits []bool
|
bits []bool
|
||||||
|
firstSeen bool
|
||||||
lostCounter metrics.Counter
|
lostCounter metrics.Counter
|
||||||
dupeCounter metrics.Counter
|
dupeCounter metrics.Counter
|
||||||
outOfWindowCounter metrics.Counter
|
outOfWindowCounter metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewBits(bits uint64) *Bits {
|
func NewBits(bits uint64) *Bits {
|
||||||
b := &Bits{
|
return &Bits{
|
||||||
length: bits,
|
length: bits,
|
||||||
bits: make([]bool, bits, bits),
|
bits: make([]bool, bits, bits),
|
||||||
current: 0,
|
current: 0,
|
||||||
@@ -25,40 +24,34 @@ func NewBits(bits uint64) *Bits {
|
|||||||
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
||||||
outOfWindowCounter: metrics.GetOrRegisterCounter("network.packets.out_of_window", nil),
|
outOfWindowCounter: metrics.GetOrRegisterCounter("network.packets.out_of_window", nil),
|
||||||
}
|
}
|
||||||
|
|
||||||
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
|
||||||
b.bits[0] = true
|
|
||||||
b.current = 0
|
|
||||||
return b
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
func (b *Bits) Check(l logrus.FieldLogger, i uint64) bool {
|
||||||
// If i is the next number, return true.
|
// If i is the next number, return true.
|
||||||
if i > b.current {
|
if i > b.current || (i == 0 && b.firstSeen == false && b.current < b.length) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the window, check if it's been set already.
|
// If i is within the window, check if it's been set already. The first window will fail this check
|
||||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
if i > b.current-b.length {
|
||||||
|
return !b.bits[i%b.length]
|
||||||
|
}
|
||||||
|
|
||||||
|
// If i is within the first window
|
||||||
|
if i < b.length {
|
||||||
return !b.bits[i%b.length]
|
return !b.bits[i%b.length]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Not within the window
|
// Not within the window
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
l.Debugf("rejected a packet (top) %d %d\n", b.current, i)
|
||||||
l.Debug("rejected a packet (top)",
|
|
||||||
"current", b.current,
|
|
||||||
"incoming", i,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
func (b *Bits) Update(l *logrus.Logger, i uint64) bool {
|
||||||
// If i is the next number, return true and update current.
|
// If i is the next number, return true and update current.
|
||||||
if i == b.current+1 {
|
if i == b.current+1 {
|
||||||
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
// Report missed packets, we can only understand what was missed after the first window has been gone through
|
||||||
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
if i > b.length && b.bits[i%b.length] == false {
|
||||||
if b.bits[i%b.length] == false && i > b.length {
|
|
||||||
b.lostCounter.Inc(1)
|
b.lostCounter.Inc(1)
|
||||||
}
|
}
|
||||||
b.bits[i%b.length] = true
|
b.bits[i%b.length] = true
|
||||||
@@ -66,39 +59,73 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is a jump, adjust the window, record lost, update current, and return true
|
// If i packet is greater than current but less than the maximum length of our bitmap,
|
||||||
if i > b.current {
|
// flip everything in between to false and move ahead.
|
||||||
lost := int64(0)
|
if i > b.current && i < b.current+b.length {
|
||||||
// Zero out the bits between the current and the new counter value, limited by the window size,
|
// In between current and i need to be zero'd to allow those packets to come in later
|
||||||
// since the window is shifting
|
for n := b.current + 1; n < i; n++ {
|
||||||
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
|
||||||
if b.bits[n%b.length] == false && n > b.length {
|
|
||||||
lost++
|
|
||||||
}
|
|
||||||
b.bits[n%b.length] = false
|
b.bits[n%b.length] = false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only record any skipped packets as a result of the window moving further than the window length
|
b.bits[i%b.length] = true
|
||||||
// Any loss within the new window will be accounted for in future calls
|
b.current = i
|
||||||
lost += max(0, int64(i-b.current-b.length))
|
//l.Debugf("missed %d packets between %d and %d\n", i-b.current, i, b.current)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// If i is greater than the delta between current and the total length of our bitmap,
|
||||||
|
// just flip everything in the map and move ahead.
|
||||||
|
if i >= b.current+b.length {
|
||||||
|
// The current window loss will be accounted for later, only record the jump as loss up until then
|
||||||
|
lost := maxInt64(0, int64(i-b.current-b.length))
|
||||||
|
//TODO: explain this
|
||||||
|
if b.current == 0 {
|
||||||
|
lost++
|
||||||
|
}
|
||||||
|
|
||||||
|
for n := range b.bits {
|
||||||
|
// Don't want to count the first window as a loss
|
||||||
|
//TODO: this is likely wrong, we are wanting to track only the bit slots that we aren't going to track anymore and this is marking everything as missed
|
||||||
|
//if b.bits[n] == false {
|
||||||
|
// lost++
|
||||||
|
//}
|
||||||
|
b.bits[n] = false
|
||||||
|
}
|
||||||
|
|
||||||
b.lostCounter.Inc(lost)
|
b.lostCounter.Inc(lost)
|
||||||
|
|
||||||
|
if l.Level >= logrus.DebugLevel {
|
||||||
|
l.WithField("receiveWindow", m{"accepted": true, "currentCounter": b.current, "incomingCounter": i, "reason": "window shifting"}).
|
||||||
|
Debug("Receive window")
|
||||||
|
}
|
||||||
b.bits[i%b.length] = true
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the current window but below the current counter,
|
// Allow for the 0 packet to come in within the first window
|
||||||
// Check to see if it's a duplicate
|
if i == 0 && b.firstSeen == false && b.current < b.length {
|
||||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
b.firstSeen = true
|
||||||
if b.current == i || b.bits[i%b.length] == true {
|
b.bits[i%b.length] = true
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
return true
|
||||||
l.Debug("Receive window",
|
}
|
||||||
"accepted", false,
|
|
||||||
"currentCounter", b.current,
|
// If i is within the window of current minus length (the total pat window size),
|
||||||
"incomingCounter", i,
|
// allow it and flip to true but to NOT change current. We also have to account for the first window
|
||||||
"reason", "duplicate",
|
if ((b.current >= b.length && i > b.current-b.length) || (b.current < b.length && i < b.length)) && i <= b.current {
|
||||||
)
|
if b.current == i {
|
||||||
|
if l.Level >= logrus.DebugLevel {
|
||||||
|
l.WithField("receiveWindow", m{"accepted": false, "currentCounter": b.current, "incomingCounter": i, "reason": "duplicate"}).
|
||||||
|
Debug("Receive window")
|
||||||
|
}
|
||||||
|
b.dupeCounter.Inc(1)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if b.bits[i%b.length] == true {
|
||||||
|
if l.Level >= logrus.DebugLevel {
|
||||||
|
l.WithField("receiveWindow", m{"accepted": false, "currentCounter": b.current, "incomingCounter": i, "reason": "old duplicate"}).
|
||||||
|
Debug("Receive window")
|
||||||
}
|
}
|
||||||
b.dupeCounter.Inc(1)
|
b.dupeCounter.Inc(1)
|
||||||
return false
|
return false
|
||||||
@@ -106,17 +133,25 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
|||||||
|
|
||||||
b.bits[i%b.length] = true
|
b.bits[i%b.length] = true
|
||||||
return true
|
return true
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// In all other cases, fail and don't change current.
|
// In all other cases, fail and don't change current.
|
||||||
b.outOfWindowCounter.Inc(1)
|
b.outOfWindowCounter.Inc(1)
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Level >= logrus.DebugLevel {
|
||||||
l.Debug("Receive window",
|
l.WithField("accepted", false).
|
||||||
"accepted", false,
|
WithField("currentCounter", b.current).
|
||||||
"currentCounter", b.current,
|
WithField("incomingCounter", i).
|
||||||
"incomingCounter", i,
|
WithField("reason", "nonsense").
|
||||||
"reason", "nonsense",
|
Debug("Receive window")
|
||||||
)
|
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func maxInt64(a, b int64) int64 {
|
||||||
|
if a > b {
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|||||||
+23
-86
@@ -15,41 +15,48 @@ func TestBits(t *testing.T) {
|
|||||||
assert.Len(t, b.bits, 10)
|
assert.Len(t, b.bits, 10)
|
||||||
|
|
||||||
// This is initialized to zero - receive one. This should work.
|
// This is initialized to zero - receive one. This should work.
|
||||||
|
|
||||||
assert.True(t, b.Check(l, 1))
|
assert.True(t, b.Check(l, 1))
|
||||||
assert.True(t, b.Update(l, 1))
|
u := b.Update(l, 1)
|
||||||
|
assert.True(t, u)
|
||||||
assert.EqualValues(t, 1, b.current)
|
assert.EqualValues(t, 1, b.current)
|
||||||
g := []bool{true, true, false, false, false, false, false, false, false, false}
|
g := []bool{false, true, false, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two
|
// Receive two
|
||||||
assert.True(t, b.Check(l, 2))
|
assert.True(t, b.Check(l, 2))
|
||||||
assert.True(t, b.Update(l, 2))
|
u = b.Update(l, 2)
|
||||||
|
assert.True(t, u)
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
g = []bool{true, true, true, false, false, false, false, false, false, false}
|
g = []bool{false, true, true, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two again - it will fail
|
// Receive two again - it will fail
|
||||||
assert.False(t, b.Check(l, 2))
|
assert.False(t, b.Check(l, 2))
|
||||||
assert.False(t, b.Update(l, 2))
|
u = b.Update(l, 2)
|
||||||
|
assert.False(t, u)
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
|
|
||||||
// Jump ahead to 15, which should clear everything and set the 6th element
|
// Jump ahead to 15, which should clear everything and set the 6th element
|
||||||
assert.True(t, b.Check(l, 15))
|
assert.True(t, b.Check(l, 15))
|
||||||
assert.True(t, b.Update(l, 15))
|
u = b.Update(l, 15)
|
||||||
|
assert.True(t, u)
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 14, which is allowed because it is in the window
|
// Mark 14, which is allowed because it is in the window
|
||||||
assert.True(t, b.Check(l, 14))
|
assert.True(t, b.Check(l, 14))
|
||||||
assert.True(t, b.Update(l, 14))
|
u = b.Update(l, 14)
|
||||||
|
assert.True(t, u)
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 5, which is not allowed because it is not in the window
|
// Mark 5, which is not allowed because it is not in the window
|
||||||
assert.False(t, b.Check(l, 5))
|
assert.False(t, b.Check(l, 5))
|
||||||
assert.False(t, b.Update(l, 5))
|
u = b.Update(l, 5)
|
||||||
|
assert.False(t, u)
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.bits)
|
||||||
@@ -62,29 +69,10 @@ func TestBits(t *testing.T) {
|
|||||||
|
|
||||||
// Walk through a few windows in order
|
// Walk through a few windows in order
|
||||||
b = NewBits(10)
|
b = NewBits(10)
|
||||||
for i := uint64(1); i <= 100; i++ {
|
for i := uint64(0); i <= 100; i++ {
|
||||||
assert.True(t, b.Check(l, i), "Error while checking %v", i)
|
assert.True(t, b.Check(l, i), "Error while checking %v", i)
|
||||||
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
assert.False(t, b.Check(l, 1), "Out of window check")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBitsLargeJumps(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
b := NewBits(10)
|
|
||||||
b.lostCounter.Clear()
|
|
||||||
|
|
||||||
b = NewBits(10)
|
|
||||||
b.lostCounter.Clear()
|
|
||||||
assert.True(t, b.Update(l, 55)) // We saw packet 55 and can still track 45,46,47,48,49,50,51,52,53,54
|
|
||||||
assert.Equal(t, int64(45), b.lostCounter.Count())
|
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 100)) // We saw packet 55 and 100 and can still track 90,91,92,93,94,95,96,97,98,99
|
|
||||||
assert.Equal(t, int64(89), b.lostCounter.Count())
|
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 200)) // We saw packet 55, 100, and 200 and can still track 190,191,192,193,194,195,196,197,198,199
|
|
||||||
assert.Equal(t, int64(188), b.lostCounter.Count())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsDupeCounter(t *testing.T) {
|
func TestBitsDupeCounter(t *testing.T) {
|
||||||
@@ -136,7 +124,8 @@ func TestBitsOutOfWindowCounter(t *testing.T) {
|
|||||||
assert.False(t, b.Update(l, 0))
|
assert.False(t, b.Update(l, 0))
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
//tODO: make sure lostcounter doesn't increase in orderly increment
|
||||||
|
assert.Equal(t, int64(20), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
@@ -148,6 +137,8 @@ func TestBitsLostCounter(t *testing.T) {
|
|||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
|
//assert.True(t, b.Update(0))
|
||||||
|
assert.True(t, b.Update(l, 0))
|
||||||
assert.True(t, b.Update(l, 20))
|
assert.True(t, b.Update(l, 20))
|
||||||
assert.True(t, b.Update(l, 21))
|
assert.True(t, b.Update(l, 21))
|
||||||
assert.True(t, b.Update(l, 22))
|
assert.True(t, b.Update(l, 22))
|
||||||
@@ -158,7 +149,7 @@ func TestBitsLostCounter(t *testing.T) {
|
|||||||
assert.True(t, b.Update(l, 27))
|
assert.True(t, b.Update(l, 27))
|
||||||
assert.True(t, b.Update(l, 28))
|
assert.True(t, b.Update(l, 28))
|
||||||
assert.True(t, b.Update(l, 29))
|
assert.True(t, b.Update(l, 29))
|
||||||
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
assert.Equal(t, int64(20), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
@@ -167,6 +158,8 @@ func TestBitsLostCounter(t *testing.T) {
|
|||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
|
assert.True(t, b.Update(l, 0))
|
||||||
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 9))
|
assert.True(t, b.Update(l, 9))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
// 10 will set 0 index, 0 was already set, no lost packets
|
// 10 will set 0 index, 0 was already set, no lost packets
|
||||||
@@ -221,62 +214,6 @@ func TestBitsLostCounter(t *testing.T) {
|
|||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsLostCounterIssue1(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
b := NewBits(10)
|
|
||||||
b.lostCounter.Clear()
|
|
||||||
b.dupeCounter.Clear()
|
|
||||||
b.outOfWindowCounter.Clear()
|
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 4))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 1))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 9))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 2))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 3))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 5))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 6))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 7))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
// assert.True(t, b.Update(l, 8))
|
|
||||||
assert.True(t, b.Update(l, 10))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 11))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 14))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
// Issue seems to be here, we reset missing packet 8 to false here and don't increment the lost counter
|
|
||||||
assert.True(t, b.Update(l, 19))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 12))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 13))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 15))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 16))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 17))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 18))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 20))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.True(t, b.Update(l, 21))
|
|
||||||
|
|
||||||
// We missed packet 8 above
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkBits(b *testing.B) {
|
func BenchmarkBits(b *testing.B) {
|
||||||
z := NewBits(10)
|
z := NewBits(10)
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
//go:build boringcrypto
|
//go:build boringcrypto
|
||||||
|
// +build boringcrypto
|
||||||
|
|
||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
|
|||||||
+38
-58
@@ -1,40 +1,41 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/slackhq/nebula/cidr"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
)
|
)
|
||||||
|
|
||||||
// This allows us to "guess" what the remote might be for a host while we wait
|
// This allows us to "guess" what the remote might be for a host while we wait
|
||||||
// for the lighthouse response. See "lighthouse.calculated_remotes" in the
|
// for the lighthouse response. See "lighthouse.calculated_remotes" in the
|
||||||
// example config file.
|
// example config file.
|
||||||
type calculatedRemote struct {
|
type calculatedRemote struct {
|
||||||
ipNet netip.Prefix
|
ipNet net.IPNet
|
||||||
mask netip.Prefix
|
maskIP iputil.VpnIp
|
||||||
port uint32
|
mask iputil.VpnIp
|
||||||
|
port uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCalculatedRemote(cidr, maskCidr netip.Prefix, port int) (*calculatedRemote, error) {
|
func newCalculatedRemote(ipNet *net.IPNet, port int) (*calculatedRemote, error) {
|
||||||
if maskCidr.Addr().BitLen() != cidr.Addr().BitLen() {
|
// Ensure this is an IPv4 mask that we expect
|
||||||
return nil, fmt.Errorf("invalid mask: %s for cidr: %s", maskCidr, cidr)
|
ones, bits := ipNet.Mask.Size()
|
||||||
|
if ones == 0 || bits != 32 {
|
||||||
|
return nil, fmt.Errorf("invalid mask: %v", ipNet)
|
||||||
}
|
}
|
||||||
|
|
||||||
masked := maskCidr.Masked()
|
|
||||||
if port < 0 || port > math.MaxUint16 {
|
if port < 0 || port > math.MaxUint16 {
|
||||||
return nil, fmt.Errorf("invalid port: %d", port)
|
return nil, fmt.Errorf("invalid port: %d", port)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &calculatedRemote{
|
return &calculatedRemote{
|
||||||
ipNet: maskCidr,
|
ipNet: *ipNet,
|
||||||
mask: masked,
|
maskIP: iputil.Ip2VpnIp(ipNet.IP),
|
||||||
port: uint32(port),
|
mask: iputil.Ip2VpnIp(ipNet.Mask),
|
||||||
|
port: uint32(port),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -42,70 +43,49 @@ func (c *calculatedRemote) String() string {
|
|||||||
return fmt.Sprintf("CalculatedRemote(mask=%v port=%d)", c.ipNet, c.port)
|
return fmt.Sprintf("CalculatedRemote(mask=%v port=%d)", c.ipNet, c.port)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *calculatedRemote) ApplyV4(addr netip.Addr) *V4AddrPort {
|
func (c *calculatedRemote) Apply(ip iputil.VpnIp) *Ip4AndPort {
|
||||||
// Combine the masked bytes of the "mask" IP with the unmasked bytes of the overlay IP
|
// Combine the masked bytes of the "mask" IP with the unmasked bytes
|
||||||
maskb := net.CIDRMask(c.mask.Bits(), c.mask.Addr().BitLen())
|
// of the overlay IP
|
||||||
mask := binary.BigEndian.Uint32(maskb[:])
|
masked := (c.maskIP & c.mask) | (ip & ^c.mask)
|
||||||
|
|
||||||
b := c.mask.Addr().As4()
|
return &Ip4AndPort{Ip: uint32(masked), Port: c.port}
|
||||||
maskAddr := binary.BigEndian.Uint32(b[:])
|
|
||||||
|
|
||||||
b = addr.As4()
|
|
||||||
intAddr := binary.BigEndian.Uint32(b[:])
|
|
||||||
|
|
||||||
return &V4AddrPort{(maskAddr & mask) | (intAddr & ^mask), c.port}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *calculatedRemote) ApplyV6(addr netip.Addr) *V6AddrPort {
|
func NewCalculatedRemotesFromConfig(c *config.C, k string) (*cidr.Tree4, error) {
|
||||||
mask := net.CIDRMask(c.mask.Bits(), c.mask.Addr().BitLen())
|
|
||||||
maskAddr := c.mask.Addr().As16()
|
|
||||||
calcAddr := addr.As16()
|
|
||||||
|
|
||||||
ap := V6AddrPort{Port: c.port}
|
|
||||||
|
|
||||||
maskb := binary.BigEndian.Uint64(mask[:8])
|
|
||||||
maskAddrb := binary.BigEndian.Uint64(maskAddr[:8])
|
|
||||||
calcAddrb := binary.BigEndian.Uint64(calcAddr[:8])
|
|
||||||
ap.Hi = (maskAddrb & maskb) | (calcAddrb & ^maskb)
|
|
||||||
|
|
||||||
maskb = binary.BigEndian.Uint64(mask[8:])
|
|
||||||
maskAddrb = binary.BigEndian.Uint64(maskAddr[8:])
|
|
||||||
calcAddrb = binary.BigEndian.Uint64(calcAddr[8:])
|
|
||||||
ap.Lo = (maskAddrb & maskb) | (calcAddrb & ^maskb)
|
|
||||||
|
|
||||||
return &ap
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewCalculatedRemotesFromConfig(c *config.C, k string) (*bart.Table[[]*calculatedRemote], error) {
|
|
||||||
value := c.Get(k)
|
value := c.Get(k)
|
||||||
if value == nil {
|
if value == nil {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
calculatedRemotes := new(bart.Table[[]*calculatedRemote])
|
calculatedRemotes := cidr.NewTree4()
|
||||||
|
|
||||||
rawMap, ok := value.(map[string]any)
|
rawMap, ok := value.(map[any]any)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, value)
|
return nil, fmt.Errorf("config `%s` has invalid type: %T", k, value)
|
||||||
}
|
}
|
||||||
for rawCIDR, rawValue := range rawMap {
|
for rawKey, rawValue := range rawMap {
|
||||||
cidr, err := netip.ParsePrefix(rawCIDR)
|
rawCIDR, ok := rawKey.(string)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("config `%s` has invalid key (type %T): %v", k, rawKey, rawKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, ipNet, err := net.ParseCIDR(rawCIDR)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s", k, rawCIDR)
|
return nil, fmt.Errorf("config `%s` has invalid CIDR: %s", k, rawCIDR)
|
||||||
}
|
}
|
||||||
|
|
||||||
entry, err := newCalculatedRemotesListFromConfig(cidr, rawValue)
|
entry, err := newCalculatedRemotesListFromConfig(rawValue)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("config '%s.%s': %w", k, rawCIDR, err)
|
return nil, fmt.Errorf("config '%s.%s': %w", k, rawCIDR, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
calculatedRemotes.Insert(cidr, entry)
|
calculatedRemotes.AddCIDR(ipNet, entry)
|
||||||
}
|
}
|
||||||
|
|
||||||
return calculatedRemotes, nil
|
return calculatedRemotes, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCalculatedRemotesListFromConfig(cidr netip.Prefix, raw any) ([]*calculatedRemote, error) {
|
func newCalculatedRemotesListFromConfig(raw any) ([]*calculatedRemote, error) {
|
||||||
rawList, ok := raw.([]any)
|
rawList, ok := raw.([]any)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("calculated_remotes entry has invalid type: %T", raw)
|
return nil, fmt.Errorf("calculated_remotes entry has invalid type: %T", raw)
|
||||||
@@ -113,7 +93,7 @@ func newCalculatedRemotesListFromConfig(cidr netip.Prefix, raw any) ([]*calculat
|
|||||||
|
|
||||||
var l []*calculatedRemote
|
var l []*calculatedRemote
|
||||||
for _, e := range rawList {
|
for _, e := range rawList {
|
||||||
c, err := newCalculatedRemotesEntryFromConfig(cidr, e)
|
c, err := newCalculatedRemotesEntryFromConfig(e)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("calculated_remotes entry: %w", err)
|
return nil, fmt.Errorf("calculated_remotes entry: %w", err)
|
||||||
}
|
}
|
||||||
@@ -123,8 +103,8 @@ func newCalculatedRemotesListFromConfig(cidr netip.Prefix, raw any) ([]*calculat
|
|||||||
return l, nil
|
return l, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCalculatedRemotesEntryFromConfig(cidr netip.Prefix, raw any) (*calculatedRemote, error) {
|
func newCalculatedRemotesEntryFromConfig(raw any) (*calculatedRemote, error) {
|
||||||
rawMap, ok := raw.(map[string]any)
|
rawMap, ok := raw.(map[any]any)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("invalid type: %T", raw)
|
return nil, fmt.Errorf("invalid type: %T", raw)
|
||||||
}
|
}
|
||||||
@@ -137,7 +117,7 @@ func newCalculatedRemotesEntryFromConfig(cidr netip.Prefix, raw any) (*calculate
|
|||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("invalid mask (type %T): %v", rawValue, rawValue)
|
return nil, fmt.Errorf("invalid mask (type %T): %v", rawValue, rawValue)
|
||||||
}
|
}
|
||||||
maskCidr, err := netip.ParsePrefix(rawMask)
|
_, ipNet, err := net.ParseCIDR(rawMask)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("invalid mask: %s", rawMask)
|
return nil, fmt.Errorf("invalid mask: %s", rawMask)
|
||||||
}
|
}
|
||||||
@@ -159,5 +139,5 @@ func newCalculatedRemotesEntryFromConfig(cidr netip.Prefix, raw any) (*calculate
|
|||||||
return nil, fmt.Errorf("invalid port (type %T): %v", rawValue, rawValue)
|
return nil, fmt.Errorf("invalid port (type %T): %v", rawValue, rawValue)
|
||||||
}
|
}
|
||||||
|
|
||||||
return newCalculatedRemote(cidr, maskCidr, port)
|
return newCalculatedRemote(ipNet, port)
|
||||||
}
|
}
|
||||||
|
|||||||
+10
-64
@@ -1,81 +1,27 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestCalculatedRemoteApply(t *testing.T) {
|
func TestCalculatedRemoteApply(t *testing.T) {
|
||||||
// Test v4 addresses
|
_, ipNet, err := net.ParseCIDR("192.168.1.0/24")
|
||||||
ipNet := netip.MustParsePrefix("192.168.1.0/24")
|
|
||||||
c, err := newCalculatedRemote(ipNet, ipNet, 4242)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
input, err := netip.ParseAddr("10.0.10.182")
|
c, err := newCalculatedRemote(ipNet, 4242)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
expected, err := netip.ParseAddr("192.168.1.182")
|
input := iputil.Ip2VpnIp([]byte{10, 0, 10, 182})
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, netAddrToProtoV4AddrPort(expected, 4242), c.ApplyV4(input))
|
expected := &Ip4AndPort{
|
||||||
|
Ip: uint32(iputil.Ip2VpnIp([]byte{192, 168, 1, 182})),
|
||||||
|
Port: 4242,
|
||||||
|
}
|
||||||
|
|
||||||
// Test v6 addresses
|
assert.Equal(t, expected, c.Apply(input))
|
||||||
ipNet = netip.MustParsePrefix("ffff:ffff:ffff:ffff::0/64")
|
|
||||||
c, err = newCalculatedRemote(ipNet, ipNet, 4242)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
input, err = netip.ParseAddr("beef:beef:beef:beef:beef:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
expected, err = netip.ParseAddr("ffff:ffff:ffff:ffff:beef:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, netAddrToProtoV6AddrPort(expected, 4242), c.ApplyV6(input))
|
|
||||||
|
|
||||||
// Test v6 addresses part 2
|
|
||||||
ipNet = netip.MustParsePrefix("ffff:ffff:ffff:ffff:ffff::0/80")
|
|
||||||
c, err = newCalculatedRemote(ipNet, ipNet, 4242)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
input, err = netip.ParseAddr("beef:beef:beef:beef:beef:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
expected, err = netip.ParseAddr("ffff:ffff:ffff:ffff:ffff:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, netAddrToProtoV6AddrPort(expected, 4242), c.ApplyV6(input))
|
|
||||||
|
|
||||||
// Test v6 addresses part 2
|
|
||||||
ipNet = netip.MustParsePrefix("ffff:ffff:ffff::0/48")
|
|
||||||
c, err = newCalculatedRemote(ipNet, ipNet, 4242)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
input, err = netip.ParseAddr("beef:beef:beef:beef:beef:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
expected, err = netip.ParseAddr("ffff:ffff:ffff:beef:beef:beef:beef:beef")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, netAddrToProtoV6AddrPort(expected, 4242), c.ApplyV6(input))
|
|
||||||
}
|
|
||||||
|
|
||||||
func Test_newCalculatedRemote(t *testing.T) {
|
|
||||||
c, err := newCalculatedRemote(netip.MustParsePrefix("1::1/128"), netip.MustParsePrefix("1.0.0.0/32"), 4242)
|
|
||||||
require.EqualError(t, err, "invalid mask: 1.0.0.0/32 for cidr: 1::1/128")
|
|
||||||
require.Nil(t, c)
|
|
||||||
|
|
||||||
c, err = newCalculatedRemote(netip.MustParsePrefix("1.0.0.0/32"), netip.MustParsePrefix("1::1/128"), 4242)
|
|
||||||
require.EqualError(t, err, "invalid mask: 1::1/128 for cidr: 1.0.0.0/32")
|
|
||||||
require.Nil(t, c)
|
|
||||||
|
|
||||||
c, err = newCalculatedRemote(netip.MustParsePrefix("1.0.0.0/32"), netip.MustParsePrefix("1.0.0.0/32"), 4242)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, c)
|
|
||||||
|
|
||||||
c, err = newCalculatedRemote(netip.MustParsePrefix("1::1/128"), netip.MustParsePrefix("1::1/128"), 4242)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, c)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,163 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/ioutil"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
type CertState struct {
|
||||||
|
certificate *cert.NebulaCertificate
|
||||||
|
rawCertificate []byte
|
||||||
|
rawCertificateNoKey []byte
|
||||||
|
publicKey []byte
|
||||||
|
privateKey []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewCertState(certificate *cert.NebulaCertificate, privateKey []byte) (*CertState, error) {
|
||||||
|
// Marshal the certificate to ensure it is valid
|
||||||
|
rawCertificate, err := certificate.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid nebula certificate on interface: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
publicKey := certificate.Details.PublicKey
|
||||||
|
cs := &CertState{
|
||||||
|
rawCertificate: rawCertificate,
|
||||||
|
certificate: certificate, // PublicKey has been set to nil above
|
||||||
|
privateKey: privateKey,
|
||||||
|
publicKey: publicKey,
|
||||||
|
}
|
||||||
|
|
||||||
|
cs.certificate.Details.PublicKey = nil
|
||||||
|
rawCertNoKey, err := cs.certificate.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("error marshalling certificate no key: %s", err)
|
||||||
|
}
|
||||||
|
cs.rawCertificateNoKey = rawCertNoKey
|
||||||
|
// put public key back
|
||||||
|
cs.certificate.Details.PublicKey = cs.publicKey
|
||||||
|
return cs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewCertStateFromConfig(c *config.C) (*CertState, error) {
|
||||||
|
var pemPrivateKey []byte
|
||||||
|
var err error
|
||||||
|
|
||||||
|
privPathOrPEM := c.GetString("pki.key", "")
|
||||||
|
|
||||||
|
if privPathOrPEM == "" {
|
||||||
|
return nil, errors.New("no pki.key path or PEM data provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(privPathOrPEM, "-----BEGIN") {
|
||||||
|
pemPrivateKey = []byte(privPathOrPEM)
|
||||||
|
privPathOrPEM = "<inline>"
|
||||||
|
} else {
|
||||||
|
pemPrivateKey, err = ioutil.ReadFile(privPathOrPEM)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("unable to read pki.key file %s: %s", privPathOrPEM, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
rawKey, _, curve, err := cert.UnmarshalPrivateKey(pemPrivateKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("error while unmarshaling pki.key %s: %s", privPathOrPEM, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var rawCert []byte
|
||||||
|
|
||||||
|
pubPathOrPEM := c.GetString("pki.cert", "")
|
||||||
|
|
||||||
|
if pubPathOrPEM == "" {
|
||||||
|
return nil, errors.New("no pki.cert path or PEM data provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(pubPathOrPEM, "-----BEGIN") {
|
||||||
|
rawCert = []byte(pubPathOrPEM)
|
||||||
|
pubPathOrPEM = "<inline>"
|
||||||
|
} else {
|
||||||
|
rawCert, err = ioutil.ReadFile(pubPathOrPEM)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("unable to read pki.cert file %s: %s", pubPathOrPEM, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
nebulaCert, _, err := cert.UnmarshalNebulaCertificateFromPEM(rawCert)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("error while unmarshaling pki.cert %s: %s", pubPathOrPEM, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if nebulaCert.Expired(time.Now()) {
|
||||||
|
return nil, fmt.Errorf("nebula certificate for this host is expired")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(nebulaCert.Details.Ips) == 0 {
|
||||||
|
return nil, fmt.Errorf("no IPs encoded in certificate")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = nebulaCert.VerifyPrivateKey(curve, rawKey); err != nil {
|
||||||
|
return nil, fmt.Errorf("private key is not a pair with public key in nebula cert")
|
||||||
|
}
|
||||||
|
|
||||||
|
return NewCertState(nebulaCert, rawKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadCAFromConfig(l *logrus.Logger, c *config.C) (*cert.NebulaCAPool, error) {
|
||||||
|
var rawCA []byte
|
||||||
|
var err error
|
||||||
|
|
||||||
|
caPathOrPEM := c.GetString("pki.ca", "")
|
||||||
|
if caPathOrPEM == "" {
|
||||||
|
return nil, errors.New("no pki.ca path or PEM data provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(caPathOrPEM, "-----BEGIN") {
|
||||||
|
rawCA = []byte(caPathOrPEM)
|
||||||
|
|
||||||
|
} else {
|
||||||
|
rawCA, err = ioutil.ReadFile(caPathOrPEM)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("unable to read pki.ca file %s: %s", caPathOrPEM, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
CAs, err := cert.NewCAPoolFromBytes(rawCA)
|
||||||
|
if errors.Is(err, cert.ErrExpired) {
|
||||||
|
var expired int
|
||||||
|
for _, cert := range CAs.CAs {
|
||||||
|
if cert.Expired(time.Now()) {
|
||||||
|
expired++
|
||||||
|
l.WithField("cert", cert).Warn("expired certificate present in CA pool")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if expired >= len(CAs.CAs) {
|
||||||
|
return nil, errors.New("no valid CA certificates present")
|
||||||
|
}
|
||||||
|
|
||||||
|
} else if err != nil {
|
||||||
|
return nil, fmt.Errorf("error while adding CA certificate to CA trust store: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, fp := range c.GetStringSlice("pki.blocklist", []string{}) {
|
||||||
|
l.WithField("fingerprint", fp).Info("Blocklisting cert")
|
||||||
|
CAs.BlocklistFingerprint(fp)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Support deprecated config for at least one minor release to allow for migrations
|
||||||
|
//TODO: remove in 2022 or later
|
||||||
|
for _, fp := range c.GetStringSlice("pki.blacklist", []string{}) {
|
||||||
|
l.WithField("fingerprint", fp).Info("Blocklisting cert")
|
||||||
|
l.Warn("pki.blacklist is deprecated and will not be supported in a future release. Please migrate your config to use pki.blocklist")
|
||||||
|
CAs.BlocklistFingerprint(fp)
|
||||||
|
}
|
||||||
|
|
||||||
|
return CAs, nil
|
||||||
|
}
|
||||||
+1
-1
@@ -1,7 +1,7 @@
|
|||||||
GO111MODULE = on
|
GO111MODULE = on
|
||||||
export GO111MODULE
|
export GO111MODULE
|
||||||
|
|
||||||
cert_v1.pb.go: cert_v1.proto .FORCE
|
cert.pb.go: cert.proto .FORCE
|
||||||
go build google.golang.org/protobuf/cmd/protoc-gen-go
|
go build google.golang.org/protobuf/cmd/protoc-gen-go
|
||||||
PATH="$(CURDIR):$(PATH)" protoc --go_out=. --go_opt=paths=source_relative $<
|
PATH="$(CURDIR):$(PATH)" protoc --go_out=. --go_opt=paths=source_relative $<
|
||||||
rm protoc-gen-go
|
rm protoc-gen-go
|
||||||
|
|||||||
+4
-15
@@ -2,25 +2,14 @@
|
|||||||
|
|
||||||
This is a library for interacting with `nebula` style certificates and authorities.
|
This is a library for interacting with `nebula` style certificates and authorities.
|
||||||
|
|
||||||
There are now 2 versions of `nebula` certificates:
|
A `protobuf` definition of the certificate format is also included
|
||||||
|
|
||||||
## v1
|
### Compiling the protobuf definition
|
||||||
|
|
||||||
This version is deprecated.
|
Make sure you have `protoc` installed.
|
||||||
|
|
||||||
A `protobuf` definition of the certificate format is included at `cert_v1.proto`
|
|
||||||
|
|
||||||
To compile the definition you will need `protoc` installed.
|
|
||||||
|
|
||||||
To compile for `go` with the same version of protobuf specified in go.mod:
|
To compile for `go` with the same version of protobuf specified in go.mod:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
make proto
|
make
|
||||||
```
|
```
|
||||||
|
|
||||||
## v2
|
|
||||||
|
|
||||||
This is the latest version which uses asn.1 DER encoding. It can support ipv4 and ipv6 and tolerate
|
|
||||||
future certificate changes better than v1.
|
|
||||||
|
|
||||||
`cert_v2.asn1` defines the wire format and can be used to compile marshalers.
|
|
||||||
@@ -1,52 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"golang.org/x/crypto/cryptobyte"
|
|
||||||
"golang.org/x/crypto/cryptobyte/asn1"
|
|
||||||
)
|
|
||||||
|
|
||||||
// readOptionalASN1Boolean reads an asn.1 boolean with a specific tag instead of a asn.1 tag wrapping a boolean with a value
|
|
||||||
// https://github.com/golang/go/issues/64811#issuecomment-1944446920
|
|
||||||
func readOptionalASN1Boolean(b *cryptobyte.String, out *bool, tag asn1.Tag, defaultValue bool) bool {
|
|
||||||
var present bool
|
|
||||||
var child cryptobyte.String
|
|
||||||
if !b.ReadOptionalASN1(&child, &present, tag) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if !present {
|
|
||||||
*out = defaultValue
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ensure we have 1 byte
|
|
||||||
if len(child) == 1 {
|
|
||||||
*out = child[0] > 0
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// readOptionalASN1Byte reads an asn.1 uint8 with a specific tag instead of a asn.1 tag wrapping a uint8 with a value
|
|
||||||
// Similar issue as with readOptionalASN1Boolean
|
|
||||||
func readOptionalASN1Byte(b *cryptobyte.String, out *byte, tag asn1.Tag, defaultValue byte) bool {
|
|
||||||
var present bool
|
|
||||||
var child cryptobyte.String
|
|
||||||
if !b.ReadOptionalASN1(&child, &present, tag) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if !present {
|
|
||||||
*out = defaultValue
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ensure we have 1 byte
|
|
||||||
if len(child) == 1 {
|
|
||||||
*out = child[0]
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
+140
@@ -0,0 +1,140 @@
|
|||||||
|
package cert
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type NebulaCAPool struct {
|
||||||
|
CAs map[string]*NebulaCertificate
|
||||||
|
certBlocklist map[string]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCAPool creates a CAPool
|
||||||
|
func NewCAPool() *NebulaCAPool {
|
||||||
|
ca := NebulaCAPool{
|
||||||
|
CAs: make(map[string]*NebulaCertificate),
|
||||||
|
certBlocklist: make(map[string]struct{}),
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ca
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCAPoolFromBytes will create a new CA pool from the provided
|
||||||
|
// input bytes, which must be a PEM-encoded set of nebula certificates.
|
||||||
|
// If the pool contains any expired certificates, an ErrExpired will be
|
||||||
|
// returned along with the pool. The caller must handle any such errors.
|
||||||
|
func NewCAPoolFromBytes(caPEMs []byte) (*NebulaCAPool, error) {
|
||||||
|
pool := NewCAPool()
|
||||||
|
var err error
|
||||||
|
var expired bool
|
||||||
|
for {
|
||||||
|
caPEMs, err = pool.AddCACertificate(caPEMs)
|
||||||
|
if errors.Is(err, ErrExpired) {
|
||||||
|
expired = true
|
||||||
|
err = nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(caPEMs) == 0 || strings.TrimSpace(string(caPEMs)) == "" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if expired {
|
||||||
|
return pool, ErrExpired
|
||||||
|
}
|
||||||
|
|
||||||
|
return pool, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddCACertificate verifies a Nebula CA certificate and adds it to the pool
|
||||||
|
// Only the first pem encoded object will be consumed, any remaining bytes are returned.
|
||||||
|
// Parsed certificates will be verified and must be a CA
|
||||||
|
func (ncp *NebulaCAPool) AddCACertificate(pemBytes []byte) ([]byte, error) {
|
||||||
|
c, pemBytes, err := UnmarshalNebulaCertificateFromPEM(pemBytes)
|
||||||
|
if err != nil {
|
||||||
|
return pemBytes, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if !c.Details.IsCA {
|
||||||
|
return pemBytes, fmt.Errorf("%s: %w", c.Details.Name, ErrNotCA)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !c.CheckSignature(c.Details.PublicKey) {
|
||||||
|
return pemBytes, fmt.Errorf("%s: %w", c.Details.Name, ErrNotSelfSigned)
|
||||||
|
}
|
||||||
|
|
||||||
|
sum, err := c.Sha256Sum()
|
||||||
|
if err != nil {
|
||||||
|
return pemBytes, fmt.Errorf("could not calculate shasum for provided CA; error: %s; %s", err, c.Details.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
ncp.CAs[sum] = c
|
||||||
|
if c.Expired(time.Now()) {
|
||||||
|
return pemBytes, fmt.Errorf("%s: %w", c.Details.Name, ErrExpired)
|
||||||
|
}
|
||||||
|
|
||||||
|
return pemBytes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// BlocklistFingerprint adds a cert fingerprint to the blocklist
|
||||||
|
func (ncp *NebulaCAPool) BlocklistFingerprint(f string) {
|
||||||
|
ncp.certBlocklist[f] = struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetCertBlocklist removes all previously blocklisted cert fingerprints
|
||||||
|
func (ncp *NebulaCAPool) ResetCertBlocklist() {
|
||||||
|
ncp.certBlocklist = make(map[string]struct{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NOTE: This uses an internal cache for Sha256Sum() that will not be invalidated
|
||||||
|
// automatically if you manually change any fields in the NebulaCertificate.
|
||||||
|
func (ncp *NebulaCAPool) IsBlocklisted(c *NebulaCertificate) bool {
|
||||||
|
return ncp.isBlocklistedWithCache(c, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsBlocklisted returns true if the fingerprint fails to generate or has been explicitly blocklisted
|
||||||
|
func (ncp *NebulaCAPool) isBlocklistedWithCache(c *NebulaCertificate, useCache bool) bool {
|
||||||
|
h, err := c.sha256SumWithCache(useCache)
|
||||||
|
if err != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, ok := ncp.certBlocklist[h]; ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCAForCert attempts to return the signing certificate for the provided certificate.
|
||||||
|
// No signature validation is performed
|
||||||
|
func (ncp *NebulaCAPool) GetCAForCert(c *NebulaCertificate) (*NebulaCertificate, error) {
|
||||||
|
if c.Details.Issuer == "" {
|
||||||
|
return nil, fmt.Errorf("no issuer in certificate")
|
||||||
|
}
|
||||||
|
|
||||||
|
signer, ok := ncp.CAs[c.Details.Issuer]
|
||||||
|
if ok {
|
||||||
|
return signer, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("could not find ca for the certificate")
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetFingerprints returns an array of trusted CA fingerprints
|
||||||
|
func (ncp *NebulaCAPool) GetFingerprints() []string {
|
||||||
|
fp := make([]string, len(ncp.CAs))
|
||||||
|
|
||||||
|
i := 0
|
||||||
|
for k := range ncp.CAs {
|
||||||
|
fp[i] = k
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
|
||||||
|
return fp
|
||||||
|
}
|
||||||
-341
@@ -1,341 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"bytes"
|
|
||||||
"encoding/pem"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
type CAPool struct {
|
|
||||||
CAs map[string]*CachedCertificate
|
|
||||||
certBlocklist map[string]struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCAPool creates an empty CAPool
|
|
||||||
func NewCAPool() *CAPool {
|
|
||||||
ca := CAPool{
|
|
||||||
CAs: make(map[string]*CachedCertificate),
|
|
||||||
certBlocklist: make(map[string]struct{}),
|
|
||||||
}
|
|
||||||
|
|
||||||
return &ca
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewCAPoolFromPEM will create a new CA pool from the provided
|
|
||||||
// input bytes, which must be a PEM-encoded set of nebula certificates.
|
|
||||||
// If the pool contains any expired certificates, an ErrExpired will be
|
|
||||||
// returned along with the pool. The caller must handle any such errors.
|
|
||||||
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()
|
|
||||||
|
|
||||||
var expired bool
|
|
||||||
|
|
||||||
scanner := bufio.NewScanner(r)
|
|
||||||
scanner.Split(SplitPEM)
|
|
||||||
|
|
||||||
for scanner.Scan() {
|
|
||||||
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 {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
err = pool.AddCA(c)
|
|
||||||
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 {
|
|
||||||
return pool, ErrExpired
|
|
||||||
}
|
|
||||||
|
|
||||||
return pool, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddCAFromPEM verifies a Nebula CA certificate and adds it to the pool.
|
|
||||||
// Only the first pem encoded object will be consumed, any remaining bytes are returned.
|
|
||||||
// Parsed certificates will be verified and must be a CA
|
|
||||||
func (ncp *CAPool) AddCAFromPEM(pemBytes []byte) ([]byte, error) {
|
|
||||||
c, pemBytes, err := UnmarshalCertificateFromPEM(pemBytes)
|
|
||||||
if err != nil {
|
|
||||||
return pemBytes, err
|
|
||||||
}
|
|
||||||
|
|
||||||
err = ncp.AddCA(c)
|
|
||||||
if err != nil {
|
|
||||||
return pemBytes, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return pemBytes, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddCA verifies a Nebula CA certificate and adds it to the pool.
|
|
||||||
func (ncp *CAPool) AddCA(c Certificate) error {
|
|
||||||
if !c.IsCA() {
|
|
||||||
return fmt.Errorf("%s: %w", c.Name(), ErrNotCA)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !c.CheckSignature(c.PublicKey()) {
|
|
||||||
return fmt.Errorf("%s: %w", c.Name(), ErrNotSelfSigned)
|
|
||||||
}
|
|
||||||
|
|
||||||
sum, err := c.Fingerprint()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("could not calculate fingerprint for provided CA; error: %w; %s", err, c.Name())
|
|
||||||
}
|
|
||||||
|
|
||||||
cc := &CachedCertificate{
|
|
||||||
Certificate: c,
|
|
||||||
Fingerprint: sum,
|
|
||||||
InvertedGroups: make(map[string]struct{}),
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, g := range c.Groups() {
|
|
||||||
cc.InvertedGroups[g] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
ncp.CAs[sum] = cc
|
|
||||||
|
|
||||||
if c.Expired(time.Now()) {
|
|
||||||
return fmt.Errorf("%s: %w", c.Name(), ErrExpired)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// BlocklistFingerprint adds a cert fingerprint to the blocklist
|
|
||||||
func (ncp *CAPool) BlocklistFingerprint(f string) {
|
|
||||||
ncp.certBlocklist[f] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ResetCertBlocklist removes all previously blocklisted cert fingerprints
|
|
||||||
func (ncp *CAPool) ResetCertBlocklist() {
|
|
||||||
ncp.certBlocklist = make(map[string]struct{})
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsBlocklisted tests the provided fingerprint against the pools blocklist.
|
|
||||||
// Returns true if the fingerprint is blocked.
|
|
||||||
func (ncp *CAPool) IsBlocklisted(fingerprint string) bool {
|
|
||||||
if _, ok := ncp.certBlocklist[fingerprint]; ok {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// VerifyCertificate verifies the certificate is valid and is signed by a trusted CA in the pool.
|
|
||||||
// If the certificate is valid then the returned CachedCertificate can be used in subsequent verification attempts
|
|
||||||
// to increase performance.
|
|
||||||
func (ncp *CAPool) VerifyCertificate(now time.Time, c Certificate) (*CachedCertificate, error) {
|
|
||||||
if c == nil {
|
|
||||||
return nil, fmt.Errorf("no certificate")
|
|
||||||
}
|
|
||||||
fp, err := c.Fingerprint()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("could not calculate fingerprint to verify: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
signer, err := ncp.verify(c, now, fp, "")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Pre nebula v1.10.3 could generate signatures in either high or low s form and validation
|
|
||||||
// of signatures allowed for either. Nebula v1.10.3 and beyond clamps signature generation to low-s form
|
|
||||||
// but validation still allows for either. Since a change in the signature bytes affects the fingerprint, we
|
|
||||||
// need to test both forms until such a time comes that we enforce low-s form on signature validation.
|
|
||||||
fp2, err := CalculateAlternateFingerprint(c)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("could not calculate alternate fingerprint to verify: %w", err)
|
|
||||||
}
|
|
||||||
if fp2 != "" && ncp.IsBlocklisted(fp2) {
|
|
||||||
return nil, ErrBlockListed
|
|
||||||
}
|
|
||||||
|
|
||||||
cc := CachedCertificate{
|
|
||||||
Certificate: c,
|
|
||||||
InvertedGroups: make(map[string]struct{}),
|
|
||||||
Fingerprint: fp,
|
|
||||||
fingerprint2: fp2,
|
|
||||||
signerFingerprint: signer.Fingerprint,
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, g := range c.Groups() {
|
|
||||||
cc.InvertedGroups[g] = struct{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
return &cc, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// VerifyCachedCertificate is the same as VerifyCertificate other than it operates on a pre-verified structure and
|
|
||||||
// is a cheaper operation to perform as a result.
|
|
||||||
func (ncp *CAPool) VerifyCachedCertificate(now time.Time, c *CachedCertificate) error {
|
|
||||||
// Check any available alternate fingerprint forms for this certificate, re P256 high-s/low-s
|
|
||||||
if c.fingerprint2 != "" && ncp.IsBlocklisted(c.fingerprint2) {
|
|
||||||
return ErrBlockListed
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err := ncp.verify(c.Certificate, now, c.Fingerprint, c.signerFingerprint)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (ncp *CAPool) verify(c Certificate, now time.Time, certFp string, signerFp string) (*CachedCertificate, error) {
|
|
||||||
if ncp.IsBlocklisted(certFp) {
|
|
||||||
return nil, ErrBlockListed
|
|
||||||
}
|
|
||||||
|
|
||||||
signer, err := ncp.GetCAForCert(c)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if signer.Certificate.Expired(now) {
|
|
||||||
return nil, ErrRootExpired
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.Expired(now) {
|
|
||||||
return nil, ErrExpired
|
|
||||||
}
|
|
||||||
|
|
||||||
// If we are checking a cached certificate then we can bail early here
|
|
||||||
// Either the root is no longer trusted or everything is fine
|
|
||||||
if len(signerFp) > 0 {
|
|
||||||
if signerFp != signer.Fingerprint {
|
|
||||||
return nil, ErrFingerprintMismatch
|
|
||||||
}
|
|
||||||
return signer, nil
|
|
||||||
}
|
|
||||||
if !c.CheckSignature(signer.Certificate.PublicKey()) {
|
|
||||||
return nil, ErrSignatureMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
err = CheckCAConstraints(signer.Certificate, c)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return signer, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetCAForCert attempts to return the signing certificate for the provided certificate.
|
|
||||||
// No signature validation is performed
|
|
||||||
func (ncp *CAPool) GetCAForCert(c Certificate) (*CachedCertificate, error) {
|
|
||||||
issuer := c.Issuer()
|
|
||||||
if issuer == "" {
|
|
||||||
return nil, fmt.Errorf("no issuer in certificate")
|
|
||||||
}
|
|
||||||
|
|
||||||
signer, ok := ncp.CAs[issuer]
|
|
||||||
if ok {
|
|
||||||
return signer, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil, ErrCaNotFound
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetFingerprints returns an array of trusted CA fingerprints
|
|
||||||
func (ncp *CAPool) GetFingerprints() []string {
|
|
||||||
fp := make([]string, len(ncp.CAs))
|
|
||||||
|
|
||||||
i := 0
|
|
||||||
for k := range ncp.CAs {
|
|
||||||
fp[i] = k
|
|
||||||
i++
|
|
||||||
}
|
|
||||||
|
|
||||||
return fp
|
|
||||||
}
|
|
||||||
|
|
||||||
// CheckCAConstraints returns an error if the sub certificate violates constraints present in the signer certificate.
|
|
||||||
func CheckCAConstraints(signer Certificate, sub Certificate) error {
|
|
||||||
return checkCAConstraints(signer, sub.NotBefore(), sub.NotAfter(), sub.Groups(), sub.Networks(), sub.UnsafeNetworks())
|
|
||||||
}
|
|
||||||
|
|
||||||
// checkCAConstraints is a very generic function allowing both Certificates and TBSCertificates to be tested.
|
|
||||||
func checkCAConstraints(signer Certificate, notBefore, notAfter time.Time, groups []string, networks, unsafeNetworks []netip.Prefix) error {
|
|
||||||
// Make sure this cert isn't valid after the root
|
|
||||||
if notAfter.After(signer.NotAfter()) {
|
|
||||||
return fmt.Errorf("certificate expires after signing certificate")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Make sure this cert wasn't valid before the root
|
|
||||||
if notBefore.Before(signer.NotBefore()) {
|
|
||||||
return fmt.Errorf("certificate is valid before the signing certificate")
|
|
||||||
}
|
|
||||||
|
|
||||||
// If the signer has a limited set of groups make sure the cert only contains a subset
|
|
||||||
signerGroups := signer.Groups()
|
|
||||||
if len(signerGroups) > 0 {
|
|
||||||
for _, g := range groups {
|
|
||||||
if !slices.Contains(signerGroups, g) {
|
|
||||||
return fmt.Errorf("certificate contained a group not present on the signing ca: %s", g)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// If the signer has a limited set of ip ranges to issue from make sure the cert only contains a subset
|
|
||||||
signingNetworks := signer.Networks()
|
|
||||||
if len(signingNetworks) > 0 {
|
|
||||||
for _, certNetwork := range networks {
|
|
||||||
found := false
|
|
||||||
for _, signingNetwork := range signingNetworks {
|
|
||||||
if signingNetwork.Contains(certNetwork.Addr()) && signingNetwork.Bits() <= certNetwork.Bits() {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !found {
|
|
||||||
return fmt.Errorf("certificate contained a network assignment outside the limitations of the signing ca: %s", certNetwork.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// If the signer has a limited set of subnet ranges to issue from make sure the cert only contains a subset
|
|
||||||
signingUnsafeNetworks := signer.UnsafeNetworks()
|
|
||||||
if len(signingUnsafeNetworks) > 0 {
|
|
||||||
for _, certUnsafeNetwork := range unsafeNetworks {
|
|
||||||
found := false
|
|
||||||
for _, caNetwork := range signingUnsafeNetworks {
|
|
||||||
if caNetwork.Contains(certUnsafeNetwork.Addr()) && caNetwork.Bits() <= certUnsafeNetwork.Bits() {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !found {
|
|
||||||
return fmt.Errorf("certificate contained an unsafe network assignment outside the limitations of the signing ca: %s", certUnsafeNetwork.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,656 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert/p256"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewCAPoolFromBytes(t *testing.T) {
|
|
||||||
noNewLines := `
|
|
||||||
# Current provisional, Remove once everything moves over to the real root.
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
Cj4KDm5lYnVsYSByb290IGNhKM0cMM24zPCvBzogV24YEw5YiqeI/oYo8XXFsoo+
|
|
||||||
PBmiOafNJhLacf9rsspAARJAz9OAnh8TKAUKix1kKVMyQU4iM3LsFfZRf6ODWXIf
|
|
||||||
2qWMpB6fpd3PSoVYziPoOt2bIHIFLlgRLPJz3I3xBEdBCQ==
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
# root-ca01
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CkEKEW5lYnVsYSByb290IGNhIDAxKM0cMM24zPCvBzogPzbWTxt8ZgXPQEwup7Br
|
|
||||||
BrtIt1O0q5AuTRT3+t2x1VJAARJAZ+2ib23qBXjdy49oU1YysrwuKkWWKrtJ7Jye
|
|
||||||
rFBQpDXikOukhQD/mfkloFwJ+Yjsfru7IpTN4ZfjXL+kN/2sCA==
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
`
|
|
||||||
|
|
||||||
withNewLines := `
|
|
||||||
# Current provisional, Remove once everything moves over to the real root.
|
|
||||||
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
Cj4KDm5lYnVsYSByb290IGNhKM0cMM24zPCvBzogV24YEw5YiqeI/oYo8XXFsoo+
|
|
||||||
PBmiOafNJhLacf9rsspAARJAz9OAnh8TKAUKix1kKVMyQU4iM3LsFfZRf6ODWXIf
|
|
||||||
2qWMpB6fpd3PSoVYziPoOt2bIHIFLlgRLPJz3I3xBEdBCQ==
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
|
|
||||||
# root-ca01
|
|
||||||
|
|
||||||
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CkEKEW5lYnVsYSByb290IGNhIDAxKM0cMM24zPCvBzogPzbWTxt8ZgXPQEwup7Br
|
|
||||||
BrtIt1O0q5AuTRT3+t2x1VJAARJAZ+2ib23qBXjdy49oU1YysrwuKkWWKrtJ7Jye
|
|
||||||
rFBQpDXikOukhQD/mfkloFwJ+Yjsfru7IpTN4ZfjXL+kN/2sCA==
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
|
|
||||||
`
|
|
||||||
|
|
||||||
expired := `
|
|
||||||
# expired certificate
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CjMKB2V4cGlyZWQozRwwzRw6ICJSG94CqX8wn5I65Pwn25V6HftVfWeIySVtp2DA
|
|
||||||
7TY/QAESQMaAk5iJT5EnQwK524ZaaHGEJLUqqbh5yyOHhboIGiVTWkFeH3HccTW8
|
|
||||||
Tq5a8AyWDQdfXbtEZ1FwabeHfH5Asw0=
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
`
|
|
||||||
|
|
||||||
p256 := `
|
|
||||||
# p256 certificate
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CmQKEG5lYnVsYSBQMjU2IHRlc3QozRwwzbjM8K8HOkEEdrmmg40zQp44AkMq6DZp
|
|
||||||
k+coOv04r+zh33ISyhbsafnYduN17p2eD7CmHvHuerguXD9f32gcxo/KsFCKEjMe
|
|
||||||
+0ABoAYBEkcwRQIgVoTg38L7uWku9xQgsr06kxZ/viQLOO/w1Qj1vFUEnhcCIQCq
|
|
||||||
75SjTiV92kv/1GcbT3wWpAZQQDBiUHVMVmh1822szA==
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
`
|
|
||||||
|
|
||||||
rootCA := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "nebula root ca",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
rootCA01 := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "nebula root ca 01",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
rootCAP256 := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "nebula P256 test",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
p, err := NewCAPoolFromPEM([]byte(noNewLines))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, p.CAs["ce4e6c7a596996eb0d82a8875f0f0137a4b53ce22d2421c9fd7150e7a26f6300"].Certificate.Name(), rootCA.details.name)
|
|
||||||
assert.Equal(t, p.CAs["04c585fcd9a49b276df956a22b7ebea3bf23f1fca5a17c0b56ce2e626631969e"].Certificate.Name(), rootCA01.details.name)
|
|
||||||
|
|
||||||
pp, err := NewCAPoolFromPEM([]byte(withNewLines))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, pp.CAs["ce4e6c7a596996eb0d82a8875f0f0137a4b53ce22d2421c9fd7150e7a26f6300"].Certificate.Name(), rootCA.details.name)
|
|
||||||
assert.Equal(t, pp.CAs["04c585fcd9a49b276df956a22b7ebea3bf23f1fca5a17c0b56ce2e626631969e"].Certificate.Name(), rootCA01.details.name)
|
|
||||||
|
|
||||||
// expired cert, no valid certs
|
|
||||||
ppp, err := NewCAPoolFromPEM([]byte(expired))
|
|
||||||
assert.Equal(t, ErrExpired, err)
|
|
||||||
assert.Equal(t, "expired", ppp.CAs["c39b35a0e8f246203fe4f32b9aa8bfd155f1ae6a6be9d78370641e43397f48f5"].Certificate.Name())
|
|
||||||
|
|
||||||
// expired cert, with valid certs
|
|
||||||
pppp, err := NewCAPoolFromPEM(append([]byte(expired), noNewLines...))
|
|
||||||
assert.Equal(t, ErrExpired, err)
|
|
||||||
assert.Equal(t, pppp.CAs["ce4e6c7a596996eb0d82a8875f0f0137a4b53ce22d2421c9fd7150e7a26f6300"].Certificate.Name(), rootCA.details.name)
|
|
||||||
assert.Equal(t, pppp.CAs["04c585fcd9a49b276df956a22b7ebea3bf23f1fca5a17c0b56ce2e626631969e"].Certificate.Name(), rootCA01.details.name)
|
|
||||||
assert.Equal(t, "expired", pppp.CAs["c39b35a0e8f246203fe4f32b9aa8bfd155f1ae6a6be9d78370641e43397f48f5"].Certificate.Name())
|
|
||||||
assert.Len(t, pppp.CAs, 3)
|
|
||||||
|
|
||||||
ppppp, err := NewCAPoolFromPEM([]byte(p256))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, ppppp.CAs["552bf7d99bec1fc775a0e4c324bf6d8f789b3078f1919c7960d2e5e0c351ee97"].Certificate.Name(), rootCAP256.details.name)
|
|
||||||
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) {
|
|
||||||
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)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
require.NoError(t, caPool.AddCA(ca))
|
|
||||||
|
|
||||||
f, err := c.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
caPool.BlocklistFingerprint(f)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now().Add(time.Hour*1000), c)
|
|
||||||
require.EqualError(t, err, "root certificate is expired")
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate is valid before the signing certificate", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test cert2", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
})
|
|
||||||
|
|
||||||
// Test group assertion
|
|
||||||
ca, _, caKey, _ = NewTestCaCert(Version1, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{"test1", "test2"})
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool = NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate contained a group not present on the signing ca: bad", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1", "bad"})
|
|
||||||
})
|
|
||||||
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test2", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_VerifyP256(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
|
||||||
c, _, _, _ := NewTestCert(Version1, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
require.NoError(t, caPool.AddCA(ca))
|
|
||||||
|
|
||||||
f, err := c.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
caPool.BlocklistFingerprint(f)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
// Create a copy of the cert and swap to the alternate form for the signature
|
|
||||||
nc := c.Copy()
|
|
||||||
b, err := p256.Swap(c.Signature())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, nc.(*certificateV1).setSignature(b))
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), nc)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now().Add(time.Hour*1000), c)
|
|
||||||
require.EqualError(t, err, "root certificate is expired")
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate is valid before the signing certificate", func() {
|
|
||||||
NewTestCert(Version1, Curve_P256, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
})
|
|
||||||
|
|
||||||
// Test group assertion
|
|
||||||
ca, _, caKey, _ = NewTestCaCert(Version1, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{"test1", "test2"})
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool = NewCAPool()
|
|
||||||
b, err = caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate contained a group not present on the signing ca: bad", func() {
|
|
||||||
NewTestCert(Version1, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1", "bad"})
|
|
||||||
})
|
|
||||||
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1"})
|
|
||||||
cc, err := caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Reset the blocklist and block the alternate form fingerprint
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
caPool.BlocklistFingerprint(cc.fingerprint2)
|
|
||||||
err = caPool.VerifyCachedCertificate(time.Now(), cc)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
err = caPool.VerifyCachedCertificate(time.Now(), cc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Verify_IPs(t *testing.T) {
|
|
||||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
|
||||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
|
||||||
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
// ip is outside the network
|
|
||||||
cIp1 := mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
cIp2 := mustParsePrefixUnmapped("192.168.0.1/16")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is outside the network reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip and mask are within the network
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/16")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/25")
|
|
||||||
c, _, _, _ := NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp2, caIp1}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed with just 1
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp1}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Verify_Subnets(t *testing.T) {
|
|
||||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
|
||||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, []netip.Prefix{caIp1, caIp2}, []string{"test"})
|
|
||||||
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
// ip is outside the network
|
|
||||||
cIp1 := mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
cIp2 := mustParsePrefixUnmapped("192.168.0.1/16")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is outside the network reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip and mask are within the network
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/16")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/25")
|
|
||||||
c, _, _, _ := NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp1, caIp2}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp2, caIp1}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed with just 1
|
|
||||||
c, _, _, _ = NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp1}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Verify(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test cert", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
require.NoError(t, caPool.AddCA(ca))
|
|
||||||
|
|
||||||
f, err := c.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
caPool.BlocklistFingerprint(f)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now().Add(time.Hour*1000), c)
|
|
||||||
require.EqualError(t, err, "root certificate is expired")
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate is valid before the signing certificate", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test cert2", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
})
|
|
||||||
|
|
||||||
// Test group assertion
|
|
||||||
ca, _, caKey, _ = NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{"test1", "test2"})
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool = NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate contained a group not present on the signing ca: bad", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1", "bad"})
|
|
||||||
})
|
|
||||||
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test2", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_VerifyP256(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
require.NoError(t, caPool.AddCA(ca))
|
|
||||||
|
|
||||||
f, err := c.Fingerprint()
|
|
||||||
require.NoError(t, err)
|
|
||||||
caPool.BlocklistFingerprint(f)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
// Create a copy of the cert and swap to the alternate form for the signature
|
|
||||||
nc := c.Copy()
|
|
||||||
b, err := p256.Swap(c.Signature())
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, nc.(*certificateV2).setSignature(b))
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), nc)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now().Add(time.Hour*1000), c)
|
|
||||||
require.EqualError(t, err, "root certificate is expired")
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate is valid before the signing certificate", func() {
|
|
||||||
NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
})
|
|
||||||
|
|
||||||
// Test group assertion
|
|
||||||
ca, _, caKey, _ = NewTestCaCert(Version2, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{"test1", "test2"})
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool = NewCAPool()
|
|
||||||
b, err = caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
assert.PanicsWithError(t, "certificate contained a group not present on the signing ca: bad", func() {
|
|
||||||
NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1", "bad"})
|
|
||||||
})
|
|
||||||
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, []string{"test1"})
|
|
||||||
cc, err := caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Reset the blocklist and block the alternate form fingerprint
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
caPool.BlocklistFingerprint(cc.fingerprint2)
|
|
||||||
err = caPool.VerifyCachedCertificate(time.Now(), cc)
|
|
||||||
require.EqualError(t, err, "certificate is in the block list")
|
|
||||||
|
|
||||||
caPool.ResetCertBlocklist()
|
|
||||||
err = caPool.VerifyCachedCertificate(time.Now(), cc)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Verify_IPs(t *testing.T) {
|
|
||||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
|
||||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
|
||||||
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
// ip is outside the network
|
|
||||||
cIp1 := mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
cIp2 := mustParsePrefixUnmapped("192.168.0.1/16")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is outside the network reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
assert.PanicsWithError(t, "certificate contained a network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip and mask are within the network
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/16")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/25")
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1, cIp2}, nil, []string{"test"})
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp2, caIp1}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed with just 1
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{caIp1}, nil, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Verify_Subnets(t *testing.T) {
|
|
||||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
|
||||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, []netip.Prefix{caIp1, caIp2}, []string{"test"})
|
|
||||||
|
|
||||||
caPem, err := ca.MarshalPEM()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
caPool := NewCAPool()
|
|
||||||
b, err := caPool.AddCAFromPEM(caPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
|
|
||||||
// ip is outside the network
|
|
||||||
cIp1 := mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
cIp2 := mustParsePrefixUnmapped("192.168.0.1/16")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is outside the network reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.1.0.0/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.1.0.0/24", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip is within the network but mask is outside reversed order of above
|
|
||||||
cIp1 = mustParsePrefixUnmapped("192.168.0.1/24")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("10.0.1.0/15")
|
|
||||||
assert.PanicsWithError(t, "certificate contained an unsafe network assignment outside the limitations of the signing ca: 10.0.1.0/15", func() {
|
|
||||||
NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
})
|
|
||||||
|
|
||||||
// ip and mask are within the network
|
|
||||||
cIp1 = mustParsePrefixUnmapped("10.0.1.0/16")
|
|
||||||
cIp2 = mustParsePrefixUnmapped("192.168.0.1/25")
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{cIp1, cIp2}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp1, caIp2}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp2, caIp1}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Exact matches reversed with just 1
|
|
||||||
c, _, _, _ = NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, []netip.Prefix{caIp1}, []string{"test"})
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
+986
-147
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,8 @@
|
|||||||
// Code generated by protoc-gen-go. DO NOT EDIT.
|
// Code generated by protoc-gen-go. DO NOT EDIT.
|
||||||
// versions:
|
// versions:
|
||||||
// protoc-gen-go v1.34.2
|
// protoc-gen-go v1.30.0
|
||||||
// protoc v3.21.5
|
// protoc v3.21.5
|
||||||
// source: cert_v1.proto
|
// source: cert.proto
|
||||||
|
|
||||||
package cert
|
package cert
|
||||||
|
|
||||||
@@ -50,11 +50,11 @@ func (x Curve) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (Curve) Descriptor() protoreflect.EnumDescriptor {
|
func (Curve) Descriptor() protoreflect.EnumDescriptor {
|
||||||
return file_cert_v1_proto_enumTypes[0].Descriptor()
|
return file_cert_proto_enumTypes[0].Descriptor()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (Curve) Type() protoreflect.EnumType {
|
func (Curve) Type() protoreflect.EnumType {
|
||||||
return &file_cert_v1_proto_enumTypes[0]
|
return &file_cert_proto_enumTypes[0]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x Curve) Number() protoreflect.EnumNumber {
|
func (x Curve) Number() protoreflect.EnumNumber {
|
||||||
@@ -63,7 +63,7 @@ func (x Curve) Number() protoreflect.EnumNumber {
|
|||||||
|
|
||||||
// Deprecated: Use Curve.Descriptor instead.
|
// Deprecated: Use Curve.Descriptor instead.
|
||||||
func (Curve) EnumDescriptor() ([]byte, []int) {
|
func (Curve) EnumDescriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{0}
|
return file_cert_proto_rawDescGZIP(), []int{0}
|
||||||
}
|
}
|
||||||
|
|
||||||
type RawNebulaCertificate struct {
|
type RawNebulaCertificate struct {
|
||||||
@@ -78,7 +78,7 @@ type RawNebulaCertificate struct {
|
|||||||
func (x *RawNebulaCertificate) Reset() {
|
func (x *RawNebulaCertificate) Reset() {
|
||||||
*x = RawNebulaCertificate{}
|
*x = RawNebulaCertificate{}
|
||||||
if protoimpl.UnsafeEnabled {
|
if protoimpl.UnsafeEnabled {
|
||||||
mi := &file_cert_v1_proto_msgTypes[0]
|
mi := &file_cert_proto_msgTypes[0]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -91,7 +91,7 @@ func (x *RawNebulaCertificate) String() string {
|
|||||||
func (*RawNebulaCertificate) ProtoMessage() {}
|
func (*RawNebulaCertificate) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *RawNebulaCertificate) ProtoReflect() protoreflect.Message {
|
func (x *RawNebulaCertificate) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_cert_v1_proto_msgTypes[0]
|
mi := &file_cert_proto_msgTypes[0]
|
||||||
if protoimpl.UnsafeEnabled && x != nil {
|
if protoimpl.UnsafeEnabled && x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -104,7 +104,7 @@ func (x *RawNebulaCertificate) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use RawNebulaCertificate.ProtoReflect.Descriptor instead.
|
// Deprecated: Use RawNebulaCertificate.ProtoReflect.Descriptor instead.
|
||||||
func (*RawNebulaCertificate) Descriptor() ([]byte, []int) {
|
func (*RawNebulaCertificate) Descriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{0}
|
return file_cert_proto_rawDescGZIP(), []int{0}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RawNebulaCertificate) GetDetails() *RawNebulaCertificateDetails {
|
func (x *RawNebulaCertificate) GetDetails() *RawNebulaCertificateDetails {
|
||||||
@@ -143,7 +143,7 @@ type RawNebulaCertificateDetails struct {
|
|||||||
func (x *RawNebulaCertificateDetails) Reset() {
|
func (x *RawNebulaCertificateDetails) Reset() {
|
||||||
*x = RawNebulaCertificateDetails{}
|
*x = RawNebulaCertificateDetails{}
|
||||||
if protoimpl.UnsafeEnabled {
|
if protoimpl.UnsafeEnabled {
|
||||||
mi := &file_cert_v1_proto_msgTypes[1]
|
mi := &file_cert_proto_msgTypes[1]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -156,7 +156,7 @@ func (x *RawNebulaCertificateDetails) String() string {
|
|||||||
func (*RawNebulaCertificateDetails) ProtoMessage() {}
|
func (*RawNebulaCertificateDetails) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *RawNebulaCertificateDetails) ProtoReflect() protoreflect.Message {
|
func (x *RawNebulaCertificateDetails) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_cert_v1_proto_msgTypes[1]
|
mi := &file_cert_proto_msgTypes[1]
|
||||||
if protoimpl.UnsafeEnabled && x != nil {
|
if protoimpl.UnsafeEnabled && x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -169,7 +169,7 @@ func (x *RawNebulaCertificateDetails) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use RawNebulaCertificateDetails.ProtoReflect.Descriptor instead.
|
// Deprecated: Use RawNebulaCertificateDetails.ProtoReflect.Descriptor instead.
|
||||||
func (*RawNebulaCertificateDetails) Descriptor() ([]byte, []int) {
|
func (*RawNebulaCertificateDetails) Descriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{1}
|
return file_cert_proto_rawDescGZIP(), []int{1}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RawNebulaCertificateDetails) GetName() string {
|
func (x *RawNebulaCertificateDetails) GetName() string {
|
||||||
@@ -254,7 +254,7 @@ type RawNebulaEncryptedData struct {
|
|||||||
func (x *RawNebulaEncryptedData) Reset() {
|
func (x *RawNebulaEncryptedData) Reset() {
|
||||||
*x = RawNebulaEncryptedData{}
|
*x = RawNebulaEncryptedData{}
|
||||||
if protoimpl.UnsafeEnabled {
|
if protoimpl.UnsafeEnabled {
|
||||||
mi := &file_cert_v1_proto_msgTypes[2]
|
mi := &file_cert_proto_msgTypes[2]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -267,7 +267,7 @@ func (x *RawNebulaEncryptedData) String() string {
|
|||||||
func (*RawNebulaEncryptedData) ProtoMessage() {}
|
func (*RawNebulaEncryptedData) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *RawNebulaEncryptedData) ProtoReflect() protoreflect.Message {
|
func (x *RawNebulaEncryptedData) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_cert_v1_proto_msgTypes[2]
|
mi := &file_cert_proto_msgTypes[2]
|
||||||
if protoimpl.UnsafeEnabled && x != nil {
|
if protoimpl.UnsafeEnabled && x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -280,7 +280,7 @@ func (x *RawNebulaEncryptedData) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use RawNebulaEncryptedData.ProtoReflect.Descriptor instead.
|
// Deprecated: Use RawNebulaEncryptedData.ProtoReflect.Descriptor instead.
|
||||||
func (*RawNebulaEncryptedData) Descriptor() ([]byte, []int) {
|
func (*RawNebulaEncryptedData) Descriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{2}
|
return file_cert_proto_rawDescGZIP(), []int{2}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RawNebulaEncryptedData) GetEncryptionMetadata() *RawNebulaEncryptionMetadata {
|
func (x *RawNebulaEncryptedData) GetEncryptionMetadata() *RawNebulaEncryptionMetadata {
|
||||||
@@ -309,7 +309,7 @@ type RawNebulaEncryptionMetadata struct {
|
|||||||
func (x *RawNebulaEncryptionMetadata) Reset() {
|
func (x *RawNebulaEncryptionMetadata) Reset() {
|
||||||
*x = RawNebulaEncryptionMetadata{}
|
*x = RawNebulaEncryptionMetadata{}
|
||||||
if protoimpl.UnsafeEnabled {
|
if protoimpl.UnsafeEnabled {
|
||||||
mi := &file_cert_v1_proto_msgTypes[3]
|
mi := &file_cert_proto_msgTypes[3]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -322,7 +322,7 @@ func (x *RawNebulaEncryptionMetadata) String() string {
|
|||||||
func (*RawNebulaEncryptionMetadata) ProtoMessage() {}
|
func (*RawNebulaEncryptionMetadata) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *RawNebulaEncryptionMetadata) ProtoReflect() protoreflect.Message {
|
func (x *RawNebulaEncryptionMetadata) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_cert_v1_proto_msgTypes[3]
|
mi := &file_cert_proto_msgTypes[3]
|
||||||
if protoimpl.UnsafeEnabled && x != nil {
|
if protoimpl.UnsafeEnabled && x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -335,7 +335,7 @@ func (x *RawNebulaEncryptionMetadata) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use RawNebulaEncryptionMetadata.ProtoReflect.Descriptor instead.
|
// Deprecated: Use RawNebulaEncryptionMetadata.ProtoReflect.Descriptor instead.
|
||||||
func (*RawNebulaEncryptionMetadata) Descriptor() ([]byte, []int) {
|
func (*RawNebulaEncryptionMetadata) Descriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{3}
|
return file_cert_proto_rawDescGZIP(), []int{3}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RawNebulaEncryptionMetadata) GetEncryptionAlgorithm() string {
|
func (x *RawNebulaEncryptionMetadata) GetEncryptionAlgorithm() string {
|
||||||
@@ -367,7 +367,7 @@ type RawNebulaArgon2Parameters struct {
|
|||||||
func (x *RawNebulaArgon2Parameters) Reset() {
|
func (x *RawNebulaArgon2Parameters) Reset() {
|
||||||
*x = RawNebulaArgon2Parameters{}
|
*x = RawNebulaArgon2Parameters{}
|
||||||
if protoimpl.UnsafeEnabled {
|
if protoimpl.UnsafeEnabled {
|
||||||
mi := &file_cert_v1_proto_msgTypes[4]
|
mi := &file_cert_proto_msgTypes[4]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -380,7 +380,7 @@ func (x *RawNebulaArgon2Parameters) String() string {
|
|||||||
func (*RawNebulaArgon2Parameters) ProtoMessage() {}
|
func (*RawNebulaArgon2Parameters) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *RawNebulaArgon2Parameters) ProtoReflect() protoreflect.Message {
|
func (x *RawNebulaArgon2Parameters) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_cert_v1_proto_msgTypes[4]
|
mi := &file_cert_proto_msgTypes[4]
|
||||||
if protoimpl.UnsafeEnabled && x != nil {
|
if protoimpl.UnsafeEnabled && x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -393,7 +393,7 @@ func (x *RawNebulaArgon2Parameters) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use RawNebulaArgon2Parameters.ProtoReflect.Descriptor instead.
|
// Deprecated: Use RawNebulaArgon2Parameters.ProtoReflect.Descriptor instead.
|
||||||
func (*RawNebulaArgon2Parameters) Descriptor() ([]byte, []int) {
|
func (*RawNebulaArgon2Parameters) Descriptor() ([]byte, []int) {
|
||||||
return file_cert_v1_proto_rawDescGZIP(), []int{4}
|
return file_cert_proto_rawDescGZIP(), []int{4}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *RawNebulaArgon2Parameters) GetVersion() int32 {
|
func (x *RawNebulaArgon2Parameters) GetVersion() int32 {
|
||||||
@@ -431,87 +431,87 @@ func (x *RawNebulaArgon2Parameters) GetSalt() []byte {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var File_cert_v1_proto protoreflect.FileDescriptor
|
var File_cert_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
var file_cert_v1_proto_rawDesc = []byte{
|
var file_cert_proto_rawDesc = []byte{
|
||||||
0x0a, 0x0d, 0x63, 0x65, 0x72, 0x74, 0x5f, 0x76, 0x31, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x12,
|
0x0a, 0x0a, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x12, 0x04, 0x63, 0x65,
|
||||||
0x04, 0x63, 0x65, 0x72, 0x74, 0x22, 0x71, 0x0a, 0x14, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75,
|
0x72, 0x74, 0x22, 0x71, 0x0a, 0x14, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x43,
|
||||||
0x6c, 0x61, 0x43, 0x65, 0x72, 0x74, 0x69, 0x66, 0x69, 0x63, 0x61, 0x74, 0x65, 0x12, 0x3b, 0x0a,
|
0x65, 0x72, 0x74, 0x69, 0x66, 0x69, 0x63, 0x61, 0x74, 0x65, 0x12, 0x3b, 0x0a, 0x07, 0x44, 0x65,
|
||||||
0x07, 0x44, 0x65, 0x74, 0x61, 0x69, 0x6c, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x21,
|
0x74, 0x61, 0x69, 0x6c, 0x73, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x21, 0x2e, 0x63, 0x65,
|
||||||
0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x43,
|
0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x43, 0x65, 0x72, 0x74,
|
||||||
0x65, 0x72, 0x74, 0x69, 0x66, 0x69, 0x63, 0x61, 0x74, 0x65, 0x44, 0x65, 0x74, 0x61, 0x69, 0x6c,
|
0x69, 0x66, 0x69, 0x63, 0x61, 0x74, 0x65, 0x44, 0x65, 0x74, 0x61, 0x69, 0x6c, 0x73, 0x52, 0x07,
|
||||||
0x73, 0x52, 0x07, 0x44, 0x65, 0x74, 0x61, 0x69, 0x6c, 0x73, 0x12, 0x1c, 0x0a, 0x09, 0x53, 0x69,
|
0x44, 0x65, 0x74, 0x61, 0x69, 0x6c, 0x73, 0x12, 0x1c, 0x0a, 0x09, 0x53, 0x69, 0x67, 0x6e, 0x61,
|
||||||
0x67, 0x6e, 0x61, 0x74, 0x75, 0x72, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x09, 0x53,
|
0x74, 0x75, 0x72, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x09, 0x53, 0x69, 0x67, 0x6e,
|
||||||
0x69, 0x67, 0x6e, 0x61, 0x74, 0x75, 0x72, 0x65, 0x22, 0x9c, 0x02, 0x0a, 0x1b, 0x52, 0x61, 0x77,
|
0x61, 0x74, 0x75, 0x72, 0x65, 0x22, 0x9c, 0x02, 0x0a, 0x1b, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62,
|
||||||
0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x43, 0x65, 0x72, 0x74, 0x69, 0x66, 0x69, 0x63, 0x61, 0x74,
|
0x75, 0x6c, 0x61, 0x43, 0x65, 0x72, 0x74, 0x69, 0x66, 0x69, 0x63, 0x61, 0x74, 0x65, 0x44, 0x65,
|
||||||
0x65, 0x44, 0x65, 0x74, 0x61, 0x69, 0x6c, 0x73, 0x12, 0x12, 0x0a, 0x04, 0x4e, 0x61, 0x6d, 0x65,
|
0x74, 0x61, 0x69, 0x6c, 0x73, 0x12, 0x12, 0x0a, 0x04, 0x4e, 0x61, 0x6d, 0x65, 0x18, 0x01, 0x20,
|
||||||
0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x04, 0x4e, 0x61, 0x6d, 0x65, 0x12, 0x10, 0x0a, 0x03,
|
0x01, 0x28, 0x09, 0x52, 0x04, 0x4e, 0x61, 0x6d, 0x65, 0x12, 0x10, 0x0a, 0x03, 0x49, 0x70, 0x73,
|
||||||
0x49, 0x70, 0x73, 0x18, 0x02, 0x20, 0x03, 0x28, 0x0d, 0x52, 0x03, 0x49, 0x70, 0x73, 0x12, 0x18,
|
0x18, 0x02, 0x20, 0x03, 0x28, 0x0d, 0x52, 0x03, 0x49, 0x70, 0x73, 0x12, 0x18, 0x0a, 0x07, 0x53,
|
||||||
0x0a, 0x07, 0x53, 0x75, 0x62, 0x6e, 0x65, 0x74, 0x73, 0x18, 0x03, 0x20, 0x03, 0x28, 0x0d, 0x52,
|
0x75, 0x62, 0x6e, 0x65, 0x74, 0x73, 0x18, 0x03, 0x20, 0x03, 0x28, 0x0d, 0x52, 0x07, 0x53, 0x75,
|
||||||
0x07, 0x53, 0x75, 0x62, 0x6e, 0x65, 0x74, 0x73, 0x12, 0x16, 0x0a, 0x06, 0x47, 0x72, 0x6f, 0x75,
|
0x62, 0x6e, 0x65, 0x74, 0x73, 0x12, 0x16, 0x0a, 0x06, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x73, 0x18,
|
||||||
0x70, 0x73, 0x18, 0x04, 0x20, 0x03, 0x28, 0x09, 0x52, 0x06, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x73,
|
0x04, 0x20, 0x03, 0x28, 0x09, 0x52, 0x06, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x73, 0x12, 0x1c, 0x0a,
|
||||||
0x12, 0x1c, 0x0a, 0x09, 0x4e, 0x6f, 0x74, 0x42, 0x65, 0x66, 0x6f, 0x72, 0x65, 0x18, 0x05, 0x20,
|
0x09, 0x4e, 0x6f, 0x74, 0x42, 0x65, 0x66, 0x6f, 0x72, 0x65, 0x18, 0x05, 0x20, 0x01, 0x28, 0x03,
|
||||||
0x01, 0x28, 0x03, 0x52, 0x09, 0x4e, 0x6f, 0x74, 0x42, 0x65, 0x66, 0x6f, 0x72, 0x65, 0x12, 0x1a,
|
0x52, 0x09, 0x4e, 0x6f, 0x74, 0x42, 0x65, 0x66, 0x6f, 0x72, 0x65, 0x12, 0x1a, 0x0a, 0x08, 0x4e,
|
||||||
0x0a, 0x08, 0x4e, 0x6f, 0x74, 0x41, 0x66, 0x74, 0x65, 0x72, 0x18, 0x06, 0x20, 0x01, 0x28, 0x03,
|
0x6f, 0x74, 0x41, 0x66, 0x74, 0x65, 0x72, 0x18, 0x06, 0x20, 0x01, 0x28, 0x03, 0x52, 0x08, 0x4e,
|
||||||
0x52, 0x08, 0x4e, 0x6f, 0x74, 0x41, 0x66, 0x74, 0x65, 0x72, 0x12, 0x1c, 0x0a, 0x09, 0x50, 0x75,
|
0x6f, 0x74, 0x41, 0x66, 0x74, 0x65, 0x72, 0x12, 0x1c, 0x0a, 0x09, 0x50, 0x75, 0x62, 0x6c, 0x69,
|
||||||
0x62, 0x6c, 0x69, 0x63, 0x4b, 0x65, 0x79, 0x18, 0x07, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x09, 0x50,
|
0x63, 0x4b, 0x65, 0x79, 0x18, 0x07, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x09, 0x50, 0x75, 0x62, 0x6c,
|
||||||
0x75, 0x62, 0x6c, 0x69, 0x63, 0x4b, 0x65, 0x79, 0x12, 0x12, 0x0a, 0x04, 0x49, 0x73, 0x43, 0x41,
|
0x69, 0x63, 0x4b, 0x65, 0x79, 0x12, 0x12, 0x0a, 0x04, 0x49, 0x73, 0x43, 0x41, 0x18, 0x08, 0x20,
|
||||||
0x18, 0x08, 0x20, 0x01, 0x28, 0x08, 0x52, 0x04, 0x49, 0x73, 0x43, 0x41, 0x12, 0x16, 0x0a, 0x06,
|
0x01, 0x28, 0x08, 0x52, 0x04, 0x49, 0x73, 0x43, 0x41, 0x12, 0x16, 0x0a, 0x06, 0x49, 0x73, 0x73,
|
||||||
0x49, 0x73, 0x73, 0x75, 0x65, 0x72, 0x18, 0x09, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x06, 0x49, 0x73,
|
0x75, 0x65, 0x72, 0x18, 0x09, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x06, 0x49, 0x73, 0x73, 0x75, 0x65,
|
||||||
0x73, 0x75, 0x65, 0x72, 0x12, 0x21, 0x0a, 0x05, 0x63, 0x75, 0x72, 0x76, 0x65, 0x18, 0x64, 0x20,
|
0x72, 0x12, 0x21, 0x0a, 0x05, 0x63, 0x75, 0x72, 0x76, 0x65, 0x18, 0x64, 0x20, 0x01, 0x28, 0x0e,
|
||||||
0x01, 0x28, 0x0e, 0x32, 0x0b, 0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x43, 0x75, 0x72, 0x76, 0x65,
|
0x32, 0x0b, 0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x43, 0x75, 0x72, 0x76, 0x65, 0x52, 0x05, 0x63,
|
||||||
0x52, 0x05, 0x63, 0x75, 0x72, 0x76, 0x65, 0x22, 0x8b, 0x01, 0x0a, 0x16, 0x52, 0x61, 0x77, 0x4e,
|
0x75, 0x72, 0x76, 0x65, 0x22, 0x8b, 0x01, 0x0a, 0x16, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75,
|
||||||
0x65, 0x62, 0x75, 0x6c, 0x61, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x65, 0x64, 0x44, 0x61,
|
0x6c, 0x61, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x65, 0x64, 0x44, 0x61, 0x74, 0x61, 0x12,
|
||||||
0x74, 0x61, 0x12, 0x51, 0x0a, 0x12, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e,
|
0x51, 0x0a, 0x12, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74,
|
||||||
0x4d, 0x65, 0x74, 0x61, 0x64, 0x61, 0x74, 0x61, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x21,
|
0x61, 0x64, 0x61, 0x74, 0x61, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x21, 0x2e, 0x63, 0x65,
|
||||||
0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x45,
|
0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x45, 0x6e, 0x63, 0x72,
|
||||||
0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74, 0x61, 0x64, 0x61, 0x74,
|
0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74, 0x61, 0x64, 0x61, 0x74, 0x61, 0x52, 0x12,
|
||||||
0x61, 0x52, 0x12, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74,
|
0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74, 0x61, 0x64, 0x61,
|
||||||
0x61, 0x64, 0x61, 0x74, 0x61, 0x12, 0x1e, 0x0a, 0x0a, 0x43, 0x69, 0x70, 0x68, 0x65, 0x72, 0x74,
|
0x74, 0x61, 0x12, 0x1e, 0x0a, 0x0a, 0x43, 0x69, 0x70, 0x68, 0x65, 0x72, 0x74, 0x65, 0x78, 0x74,
|
||||||
0x65, 0x78, 0x74, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x0a, 0x43, 0x69, 0x70, 0x68, 0x65,
|
0x18, 0x02, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x0a, 0x43, 0x69, 0x70, 0x68, 0x65, 0x72, 0x74, 0x65,
|
||||||
0x72, 0x74, 0x65, 0x78, 0x74, 0x22, 0x9c, 0x01, 0x0a, 0x1b, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62,
|
0x78, 0x74, 0x22, 0x9c, 0x01, 0x0a, 0x1b, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61,
|
||||||
0x75, 0x6c, 0x61, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74,
|
0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x74, 0x61, 0x64, 0x61,
|
||||||
0x61, 0x64, 0x61, 0x74, 0x61, 0x12, 0x30, 0x0a, 0x13, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74,
|
0x74, 0x61, 0x12, 0x30, 0x0a, 0x13, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e,
|
||||||
0x69, 0x6f, 0x6e, 0x41, 0x6c, 0x67, 0x6f, 0x72, 0x69, 0x74, 0x68, 0x6d, 0x18, 0x01, 0x20, 0x01,
|
0x41, 0x6c, 0x67, 0x6f, 0x72, 0x69, 0x74, 0x68, 0x6d, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52,
|
||||||
0x28, 0x09, 0x52, 0x13, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x41, 0x6c,
|
0x13, 0x45, 0x6e, 0x63, 0x72, 0x79, 0x70, 0x74, 0x69, 0x6f, 0x6e, 0x41, 0x6c, 0x67, 0x6f, 0x72,
|
||||||
0x67, 0x6f, 0x72, 0x69, 0x74, 0x68, 0x6d, 0x12, 0x4b, 0x0a, 0x10, 0x41, 0x72, 0x67, 0x6f, 0x6e,
|
0x69, 0x74, 0x68, 0x6d, 0x12, 0x4b, 0x0a, 0x10, 0x41, 0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61,
|
||||||
0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65, 0x72, 0x73, 0x18, 0x02, 0x20, 0x01, 0x28,
|
0x72, 0x61, 0x6d, 0x65, 0x74, 0x65, 0x72, 0x73, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1f,
|
||||||
0x0b, 0x32, 0x1f, 0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75,
|
0x2e, 0x63, 0x65, 0x72, 0x74, 0x2e, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x41,
|
||||||
0x6c, 0x61, 0x41, 0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65,
|
0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65, 0x72, 0x73, 0x52,
|
||||||
0x72, 0x73, 0x52, 0x10, 0x41, 0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65,
|
0x10, 0x41, 0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65, 0x72,
|
||||||
0x74, 0x65, 0x72, 0x73, 0x22, 0xa3, 0x01, 0x0a, 0x19, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75,
|
0x73, 0x22, 0xa3, 0x01, 0x0a, 0x19, 0x52, 0x61, 0x77, 0x4e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x41,
|
||||||
0x6c, 0x61, 0x41, 0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65,
|
0x72, 0x67, 0x6f, 0x6e, 0x32, 0x50, 0x61, 0x72, 0x61, 0x6d, 0x65, 0x74, 0x65, 0x72, 0x73, 0x12,
|
||||||
0x72, 0x73, 0x12, 0x18, 0x0a, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, 0x18, 0x01, 0x20,
|
0x18, 0x0a, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, 0x18, 0x01, 0x20, 0x01, 0x28, 0x05,
|
||||||
0x01, 0x28, 0x05, 0x52, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, 0x12, 0x16, 0x0a, 0x06,
|
0x52, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, 0x12, 0x16, 0x0a, 0x06, 0x6d, 0x65, 0x6d,
|
||||||
0x6d, 0x65, 0x6d, 0x6f, 0x72, 0x79, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x06, 0x6d, 0x65,
|
0x6f, 0x72, 0x79, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x06, 0x6d, 0x65, 0x6d, 0x6f, 0x72,
|
||||||
0x6d, 0x6f, 0x72, 0x79, 0x12, 0x20, 0x0a, 0x0b, 0x70, 0x61, 0x72, 0x61, 0x6c, 0x6c, 0x65, 0x6c,
|
0x79, 0x12, 0x20, 0x0a, 0x0b, 0x70, 0x61, 0x72, 0x61, 0x6c, 0x6c, 0x65, 0x6c, 0x69, 0x73, 0x6d,
|
||||||
0x69, 0x73, 0x6d, 0x18, 0x04, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x0b, 0x70, 0x61, 0x72, 0x61, 0x6c,
|
0x18, 0x04, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x0b, 0x70, 0x61, 0x72, 0x61, 0x6c, 0x6c, 0x65, 0x6c,
|
||||||
0x6c, 0x65, 0x6c, 0x69, 0x73, 0x6d, 0x12, 0x1e, 0x0a, 0x0a, 0x69, 0x74, 0x65, 0x72, 0x61, 0x74,
|
0x69, 0x73, 0x6d, 0x12, 0x1e, 0x0a, 0x0a, 0x69, 0x74, 0x65, 0x72, 0x61, 0x74, 0x69, 0x6f, 0x6e,
|
||||||
0x69, 0x6f, 0x6e, 0x73, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x0a, 0x69, 0x74, 0x65, 0x72,
|
0x73, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0d, 0x52, 0x0a, 0x69, 0x74, 0x65, 0x72, 0x61, 0x74, 0x69,
|
||||||
0x61, 0x74, 0x69, 0x6f, 0x6e, 0x73, 0x12, 0x12, 0x0a, 0x04, 0x73, 0x61, 0x6c, 0x74, 0x18, 0x05,
|
0x6f, 0x6e, 0x73, 0x12, 0x12, 0x0a, 0x04, 0x73, 0x61, 0x6c, 0x74, 0x18, 0x05, 0x20, 0x01, 0x28,
|
||||||
0x20, 0x01, 0x28, 0x0c, 0x52, 0x04, 0x73, 0x61, 0x6c, 0x74, 0x2a, 0x21, 0x0a, 0x05, 0x43, 0x75,
|
0x0c, 0x52, 0x04, 0x73, 0x61, 0x6c, 0x74, 0x2a, 0x21, 0x0a, 0x05, 0x43, 0x75, 0x72, 0x76, 0x65,
|
||||||
0x72, 0x76, 0x65, 0x12, 0x0e, 0x0a, 0x0a, 0x43, 0x55, 0x52, 0x56, 0x45, 0x32, 0x35, 0x35, 0x31,
|
0x12, 0x0e, 0x0a, 0x0a, 0x43, 0x55, 0x52, 0x56, 0x45, 0x32, 0x35, 0x35, 0x31, 0x39, 0x10, 0x00,
|
||||||
0x39, 0x10, 0x00, 0x12, 0x08, 0x0a, 0x04, 0x50, 0x32, 0x35, 0x36, 0x10, 0x01, 0x42, 0x20, 0x5a,
|
0x12, 0x08, 0x0a, 0x04, 0x50, 0x32, 0x35, 0x36, 0x10, 0x01, 0x42, 0x20, 0x5a, 0x1e, 0x67, 0x69,
|
||||||
0x1e, 0x67, 0x69, 0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f, 0x6d, 0x2f, 0x73, 0x6c, 0x61, 0x63,
|
0x74, 0x68, 0x75, 0x62, 0x2e, 0x63, 0x6f, 0x6d, 0x2f, 0x73, 0x6c, 0x61, 0x63, 0x6b, 0x68, 0x71,
|
||||||
0x6b, 0x68, 0x71, 0x2f, 0x6e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x2f, 0x63, 0x65, 0x72, 0x74, 0x62,
|
0x2f, 0x6e, 0x65, 0x62, 0x75, 0x6c, 0x61, 0x2f, 0x63, 0x65, 0x72, 0x74, 0x62, 0x06, 0x70, 0x72,
|
||||||
0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33,
|
0x6f, 0x74, 0x6f, 0x33,
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
file_cert_v1_proto_rawDescOnce sync.Once
|
file_cert_proto_rawDescOnce sync.Once
|
||||||
file_cert_v1_proto_rawDescData = file_cert_v1_proto_rawDesc
|
file_cert_proto_rawDescData = file_cert_proto_rawDesc
|
||||||
)
|
)
|
||||||
|
|
||||||
func file_cert_v1_proto_rawDescGZIP() []byte {
|
func file_cert_proto_rawDescGZIP() []byte {
|
||||||
file_cert_v1_proto_rawDescOnce.Do(func() {
|
file_cert_proto_rawDescOnce.Do(func() {
|
||||||
file_cert_v1_proto_rawDescData = protoimpl.X.CompressGZIP(file_cert_v1_proto_rawDescData)
|
file_cert_proto_rawDescData = protoimpl.X.CompressGZIP(file_cert_proto_rawDescData)
|
||||||
})
|
})
|
||||||
return file_cert_v1_proto_rawDescData
|
return file_cert_proto_rawDescData
|
||||||
}
|
}
|
||||||
|
|
||||||
var file_cert_v1_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
|
var file_cert_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
|
||||||
var file_cert_v1_proto_msgTypes = make([]protoimpl.MessageInfo, 5)
|
var file_cert_proto_msgTypes = make([]protoimpl.MessageInfo, 5)
|
||||||
var file_cert_v1_proto_goTypes = []any{
|
var file_cert_proto_goTypes = []interface{}{
|
||||||
(Curve)(0), // 0: cert.Curve
|
(Curve)(0), // 0: cert.Curve
|
||||||
(*RawNebulaCertificate)(nil), // 1: cert.RawNebulaCertificate
|
(*RawNebulaCertificate)(nil), // 1: cert.RawNebulaCertificate
|
||||||
(*RawNebulaCertificateDetails)(nil), // 2: cert.RawNebulaCertificateDetails
|
(*RawNebulaCertificateDetails)(nil), // 2: cert.RawNebulaCertificateDetails
|
||||||
@@ -519,7 +519,7 @@ var file_cert_v1_proto_goTypes = []any{
|
|||||||
(*RawNebulaEncryptionMetadata)(nil), // 4: cert.RawNebulaEncryptionMetadata
|
(*RawNebulaEncryptionMetadata)(nil), // 4: cert.RawNebulaEncryptionMetadata
|
||||||
(*RawNebulaArgon2Parameters)(nil), // 5: cert.RawNebulaArgon2Parameters
|
(*RawNebulaArgon2Parameters)(nil), // 5: cert.RawNebulaArgon2Parameters
|
||||||
}
|
}
|
||||||
var file_cert_v1_proto_depIdxs = []int32{
|
var file_cert_proto_depIdxs = []int32{
|
||||||
2, // 0: cert.RawNebulaCertificate.Details:type_name -> cert.RawNebulaCertificateDetails
|
2, // 0: cert.RawNebulaCertificate.Details:type_name -> cert.RawNebulaCertificateDetails
|
||||||
0, // 1: cert.RawNebulaCertificateDetails.curve:type_name -> cert.Curve
|
0, // 1: cert.RawNebulaCertificateDetails.curve:type_name -> cert.Curve
|
||||||
4, // 2: cert.RawNebulaEncryptedData.EncryptionMetadata:type_name -> cert.RawNebulaEncryptionMetadata
|
4, // 2: cert.RawNebulaEncryptedData.EncryptionMetadata:type_name -> cert.RawNebulaEncryptionMetadata
|
||||||
@@ -531,13 +531,13 @@ var file_cert_v1_proto_depIdxs = []int32{
|
|||||||
0, // [0:4] is the sub-list for field type_name
|
0, // [0:4] is the sub-list for field type_name
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { file_cert_v1_proto_init() }
|
func init() { file_cert_proto_init() }
|
||||||
func file_cert_v1_proto_init() {
|
func file_cert_proto_init() {
|
||||||
if File_cert_v1_proto != nil {
|
if File_cert_proto != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !protoimpl.UnsafeEnabled {
|
if !protoimpl.UnsafeEnabled {
|
||||||
file_cert_v1_proto_msgTypes[0].Exporter = func(v any, i int) any {
|
file_cert_proto_msgTypes[0].Exporter = func(v interface{}, i int) interface{} {
|
||||||
switch v := v.(*RawNebulaCertificate); i {
|
switch v := v.(*RawNebulaCertificate); i {
|
||||||
case 0:
|
case 0:
|
||||||
return &v.state
|
return &v.state
|
||||||
@@ -549,7 +549,7 @@ func file_cert_v1_proto_init() {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
file_cert_v1_proto_msgTypes[1].Exporter = func(v any, i int) any {
|
file_cert_proto_msgTypes[1].Exporter = func(v interface{}, i int) interface{} {
|
||||||
switch v := v.(*RawNebulaCertificateDetails); i {
|
switch v := v.(*RawNebulaCertificateDetails); i {
|
||||||
case 0:
|
case 0:
|
||||||
return &v.state
|
return &v.state
|
||||||
@@ -561,7 +561,7 @@ func file_cert_v1_proto_init() {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
file_cert_v1_proto_msgTypes[2].Exporter = func(v any, i int) any {
|
file_cert_proto_msgTypes[2].Exporter = func(v interface{}, i int) interface{} {
|
||||||
switch v := v.(*RawNebulaEncryptedData); i {
|
switch v := v.(*RawNebulaEncryptedData); i {
|
||||||
case 0:
|
case 0:
|
||||||
return &v.state
|
return &v.state
|
||||||
@@ -573,7 +573,7 @@ func file_cert_v1_proto_init() {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
file_cert_v1_proto_msgTypes[3].Exporter = func(v any, i int) any {
|
file_cert_proto_msgTypes[3].Exporter = func(v interface{}, i int) interface{} {
|
||||||
switch v := v.(*RawNebulaEncryptionMetadata); i {
|
switch v := v.(*RawNebulaEncryptionMetadata); i {
|
||||||
case 0:
|
case 0:
|
||||||
return &v.state
|
return &v.state
|
||||||
@@ -585,7 +585,7 @@ func file_cert_v1_proto_init() {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
file_cert_v1_proto_msgTypes[4].Exporter = func(v any, i int) any {
|
file_cert_proto_msgTypes[4].Exporter = func(v interface{}, i int) interface{} {
|
||||||
switch v := v.(*RawNebulaArgon2Parameters); i {
|
switch v := v.(*RawNebulaArgon2Parameters); i {
|
||||||
case 0:
|
case 0:
|
||||||
return &v.state
|
return &v.state
|
||||||
@@ -602,19 +602,19 @@ func file_cert_v1_proto_init() {
|
|||||||
out := protoimpl.TypeBuilder{
|
out := protoimpl.TypeBuilder{
|
||||||
File: protoimpl.DescBuilder{
|
File: protoimpl.DescBuilder{
|
||||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||||
RawDescriptor: file_cert_v1_proto_rawDesc,
|
RawDescriptor: file_cert_proto_rawDesc,
|
||||||
NumEnums: 1,
|
NumEnums: 1,
|
||||||
NumMessages: 5,
|
NumMessages: 5,
|
||||||
NumExtensions: 0,
|
NumExtensions: 0,
|
||||||
NumServices: 0,
|
NumServices: 0,
|
||||||
},
|
},
|
||||||
GoTypes: file_cert_v1_proto_goTypes,
|
GoTypes: file_cert_proto_goTypes,
|
||||||
DependencyIndexes: file_cert_v1_proto_depIdxs,
|
DependencyIndexes: file_cert_proto_depIdxs,
|
||||||
EnumInfos: file_cert_v1_proto_enumTypes,
|
EnumInfos: file_cert_proto_enumTypes,
|
||||||
MessageInfos: file_cert_v1_proto_msgTypes,
|
MessageInfos: file_cert_proto_msgTypes,
|
||||||
}.Build()
|
}.Build()
|
||||||
File_cert_v1_proto = out.File
|
File_cert_proto = out.File
|
||||||
file_cert_v1_proto_rawDesc = nil
|
file_cert_proto_rawDesc = nil
|
||||||
file_cert_v1_proto_goTypes = nil
|
file_cert_proto_goTypes = nil
|
||||||
file_cert_v1_proto_depIdxs = nil
|
file_cert_proto_depIdxs = nil
|
||||||
}
|
}
|
||||||
+1230
File diff suppressed because it is too large
Load Diff
-502
@@ -1,502 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/ecdh"
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/binary"
|
|
||||||
"encoding/hex"
|
|
||||||
"encoding/json"
|
|
||||||
"encoding/pem"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"golang.org/x/crypto/curve25519"
|
|
||||||
"google.golang.org/protobuf/proto"
|
|
||||||
)
|
|
||||||
|
|
||||||
const publicKeyLen = 32
|
|
||||||
|
|
||||||
type certificateV1 struct {
|
|
||||||
details detailsV1
|
|
||||||
signature []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type detailsV1 struct {
|
|
||||||
name string
|
|
||||||
networks []netip.Prefix
|
|
||||||
unsafeNetworks []netip.Prefix
|
|
||||||
groups []string
|
|
||||||
notBefore time.Time
|
|
||||||
notAfter time.Time
|
|
||||||
publicKey []byte
|
|
||||||
isCA bool
|
|
||||||
issuer string
|
|
||||||
|
|
||||||
curve Curve
|
|
||||||
}
|
|
||||||
|
|
||||||
type m = map[string]any
|
|
||||||
|
|
||||||
func (c *certificateV1) Version() Version {
|
|
||||||
return Version1
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Curve() Curve {
|
|
||||||
return c.details.curve
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Groups() []string {
|
|
||||||
return c.details.groups
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) IsCA() bool {
|
|
||||||
return c.details.isCA
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Issuer() string {
|
|
||||||
return c.details.issuer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Name() string {
|
|
||||||
return c.details.name
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Networks() []netip.Prefix {
|
|
||||||
return c.details.networks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) NotAfter() time.Time {
|
|
||||||
return c.details.notAfter
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) NotBefore() time.Time {
|
|
||||||
return c.details.notBefore
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) PublicKey() []byte {
|
|
||||||
return c.details.publicKey
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) MarshalPublicKeyPEM() []byte {
|
|
||||||
return marshalCertPublicKeyToPEM(c)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Signature() []byte {
|
|
||||||
return c.signature
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) UnsafeNetworks() []netip.Prefix {
|
|
||||||
return c.details.unsafeNetworks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Fingerprint() (string, error) {
|
|
||||||
b, err := c.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
sum := sha256.Sum256(b)
|
|
||||||
return hex.EncodeToString(sum[:]), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) CheckSignature(key []byte) bool {
|
|
||||||
b, err := proto.Marshal(c.getRawDetails())
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
switch c.details.curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
|
||||||
case Curve_P256:
|
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
hashed := sha256.Sum256(b)
|
|
||||||
return ecdsa.VerifyASN1(pubKey, hashed[:], c.signature)
|
|
||||||
default:
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Expired(t time.Time) bool {
|
|
||||||
return c.details.notBefore.After(t) || c.details.notAfter.Before(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) VerifyPrivateKey(curve Curve, key []byte) error {
|
|
||||||
if curve != c.details.curve {
|
|
||||||
return fmt.Errorf("curve in cert and private key supplied don't match")
|
|
||||||
}
|
|
||||||
if c.details.isCA {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
// the call to PublicKey below will panic slice bounds out of range otherwise
|
|
||||||
if len(key) != ed25519.PrivateKeySize {
|
|
||||||
return fmt.Errorf("key was not 64 bytes, is invalid ed25519 private key")
|
|
||||||
}
|
|
||||||
|
|
||||||
if !ed25519.PublicKey(c.details.publicKey).Equal(ed25519.PrivateKey(key).Public()) {
|
|
||||||
return fmt.Errorf("public key in cert and private key supplied don't match")
|
|
||||||
}
|
|
||||||
case Curve_P256:
|
|
||||||
privkey, err := ecdh.P256().NewPrivateKey(key)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("cannot parse private key as P256: %w", err)
|
|
||||||
}
|
|
||||||
pub := privkey.PublicKey().Bytes()
|
|
||||||
if !bytes.Equal(pub, c.details.publicKey) {
|
|
||||||
return fmt.Errorf("public key in cert and private key supplied don't match")
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", curve)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var pub []byte
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
var err error
|
|
||||||
pub, err = curve25519.X25519(key, curve25519.Basepoint)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
case Curve_P256:
|
|
||||||
privkey, err := ecdh.P256().NewPrivateKey(key)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
pub = privkey.PublicKey().Bytes()
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", curve)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(pub, c.details.publicKey) {
|
|
||||||
return fmt.Errorf("public key in cert and private key supplied don't match")
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// getRawDetails marshals the raw details into protobuf ready struct
|
|
||||||
func (c *certificateV1) getRawDetails() *RawNebulaCertificateDetails {
|
|
||||||
rd := &RawNebulaCertificateDetails{
|
|
||||||
Name: c.details.name,
|
|
||||||
Groups: c.details.groups,
|
|
||||||
NotBefore: c.details.notBefore.Unix(),
|
|
||||||
NotAfter: c.details.notAfter.Unix(),
|
|
||||||
PublicKey: make([]byte, len(c.details.publicKey)),
|
|
||||||
IsCA: c.details.isCA,
|
|
||||||
Curve: c.details.curve,
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, ipNet := range c.details.networks {
|
|
||||||
mask := net.CIDRMask(ipNet.Bits(), ipNet.Addr().BitLen())
|
|
||||||
rd.Ips = append(rd.Ips, addr2int(ipNet.Addr()), ip2int(mask))
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, ipNet := range c.details.unsafeNetworks {
|
|
||||||
mask := net.CIDRMask(ipNet.Bits(), ipNet.Addr().BitLen())
|
|
||||||
rd.Subnets = append(rd.Subnets, addr2int(ipNet.Addr()), ip2int(mask))
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(rd.PublicKey, c.details.publicKey[:])
|
|
||||||
|
|
||||||
// I know, this is terrible
|
|
||||||
rd.Issuer, _ = hex.DecodeString(c.details.issuer)
|
|
||||||
|
|
||||||
return rd
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) String() string {
|
|
||||||
b, err := json.MarshalIndent(c.marshalJSON(), "", "\t")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("<error marshalling certificate: %v>", err)
|
|
||||||
}
|
|
||||||
return string(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) MarshalForHandshakes() ([]byte, error) {
|
|
||||||
pubKey := c.details.publicKey
|
|
||||||
c.details.publicKey = nil
|
|
||||||
rawCertNoKey, err := c.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
c.details.publicKey = pubKey
|
|
||||||
return rawCertNoKey, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Marshal() ([]byte, error) {
|
|
||||||
rc := RawNebulaCertificate{
|
|
||||||
Details: c.getRawDetails(),
|
|
||||||
Signature: c.signature,
|
|
||||||
}
|
|
||||||
|
|
||||||
return proto.Marshal(&rc)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) MarshalPEM() ([]byte, error) {
|
|
||||||
b, err := c.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: CertificateBanner, Bytes: b}), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) MarshalJSON() ([]byte, error) {
|
|
||||||
return json.Marshal(c.marshalJSON())
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) marshalJSON() m {
|
|
||||||
fp, _ := c.Fingerprint()
|
|
||||||
return m{
|
|
||||||
"version": Version1,
|
|
||||||
"details": m{
|
|
||||||
"name": c.details.name,
|
|
||||||
"networks": c.details.networks,
|
|
||||||
"unsafeNetworks": c.details.unsafeNetworks,
|
|
||||||
"groups": c.details.groups,
|
|
||||||
"notBefore": c.details.notBefore,
|
|
||||||
"notAfter": c.details.notAfter,
|
|
||||||
"publicKey": fmt.Sprintf("%x", c.details.publicKey),
|
|
||||||
"isCa": c.details.isCA,
|
|
||||||
"issuer": c.details.issuer,
|
|
||||||
"curve": c.details.curve.String(),
|
|
||||||
},
|
|
||||||
"fingerprint": fp,
|
|
||||||
"signature": fmt.Sprintf("%x", c.Signature()),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) Copy() Certificate {
|
|
||||||
nc := &certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: c.details.name,
|
|
||||||
notBefore: c.details.notBefore,
|
|
||||||
notAfter: c.details.notAfter,
|
|
||||||
publicKey: make([]byte, len(c.details.publicKey)),
|
|
||||||
isCA: c.details.isCA,
|
|
||||||
issuer: c.details.issuer,
|
|
||||||
curve: c.details.curve,
|
|
||||||
},
|
|
||||||
signature: make([]byte, len(c.signature)),
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.groups != nil {
|
|
||||||
nc.details.groups = make([]string, len(c.details.groups))
|
|
||||||
copy(nc.details.groups, c.details.groups)
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.networks != nil {
|
|
||||||
nc.details.networks = make([]netip.Prefix, len(c.details.networks))
|
|
||||||
copy(nc.details.networks, c.details.networks)
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.unsafeNetworks != nil {
|
|
||||||
nc.details.unsafeNetworks = make([]netip.Prefix, len(c.details.unsafeNetworks))
|
|
||||||
copy(nc.details.unsafeNetworks, c.details.unsafeNetworks)
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(nc.signature, c.signature)
|
|
||||||
copy(nc.details.publicKey, c.details.publicKey)
|
|
||||||
|
|
||||||
return nc
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) fromTBSCertificate(t *TBSCertificate) error {
|
|
||||||
c.details = detailsV1{
|
|
||||||
name: t.Name,
|
|
||||||
networks: t.Networks,
|
|
||||||
unsafeNetworks: t.UnsafeNetworks,
|
|
||||||
groups: t.Groups,
|
|
||||||
notBefore: t.NotBefore,
|
|
||||||
notAfter: t.NotAfter,
|
|
||||||
publicKey: t.PublicKey,
|
|
||||||
isCA: t.IsCA,
|
|
||||||
curve: t.Curve,
|
|
||||||
issuer: t.issuer,
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.validate()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) validate() error {
|
|
||||||
// Empty names are allowed
|
|
||||||
|
|
||||||
if len(c.details.publicKey) == 0 {
|
|
||||||
return ErrInvalidPublicKey
|
|
||||||
}
|
|
||||||
|
|
||||||
// Original v1 rules allowed multiple networks to be present but ignored all but the first one.
|
|
||||||
// Continue to allow this behavior
|
|
||||||
if !c.details.isCA && len(c.details.networks) == 0 {
|
|
||||||
return NewErrInvalidCertificateProperties("non-CA certificates must contain exactly one network")
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, network := range c.details.networks {
|
|
||||||
if !network.IsValid() || !network.Addr().IsValid() {
|
|
||||||
return NewErrInvalidCertificateProperties("invalid network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Is6() {
|
|
||||||
return NewErrInvalidCertificateProperties("certificate may not contain IPv6 networks: %v", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().IsUnspecified() {
|
|
||||||
return NewErrInvalidCertificateProperties("non-CA certificates must not use the zero address as a network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Zone() != "" {
|
|
||||||
return NewErrInvalidCertificateProperties("networks may not contain zones: %s", network)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, network := range c.details.unsafeNetworks {
|
|
||||||
if !network.IsValid() || !network.Addr().IsValid() {
|
|
||||||
return NewErrInvalidCertificateProperties("invalid unsafe network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Is6() {
|
|
||||||
return NewErrInvalidCertificateProperties("certificate may not contain IPv6 unsafe networks: %v", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Zone() != "" {
|
|
||||||
return NewErrInvalidCertificateProperties("unsafe networks may not contain zones: %s", network)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// v1 doesn't bother with sort order or uniqueness of networks or unsafe networks.
|
|
||||||
// We can't modify the unmarshalled data because verification requires re-marshalling and a re-ordered
|
|
||||||
// unsafe networks would result in a different signature.
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) marshalForSigning() ([]byte, error) {
|
|
||||||
b, err := proto.Marshal(c.getRawDetails())
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return b, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV1) setSignature(b []byte) error {
|
|
||||||
if len(b) == 0 {
|
|
||||||
return ErrEmptySignature
|
|
||||||
}
|
|
||||||
c.signature = b
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// unmarshalCertificateV1 will unmarshal a protobuf byte representation of a nebula cert
|
|
||||||
// if the publicKey is provided here then it is not required to be present in `b`
|
|
||||||
func unmarshalCertificateV1(b []byte, publicKey []byte) (*certificateV1, error) {
|
|
||||||
if len(b) == 0 {
|
|
||||||
return nil, fmt.Errorf("nil byte array")
|
|
||||||
}
|
|
||||||
var rc RawNebulaCertificate
|
|
||||||
err := proto.Unmarshal(b, &rc)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if rc.Details == nil {
|
|
||||||
return nil, fmt.Errorf("encoded Details was nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(rc.Details.Ips)%2 != 0 {
|
|
||||||
return nil, fmt.Errorf("encoded IPs should be in pairs, an odd number was found")
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(rc.Details.Subnets)%2 != 0 {
|
|
||||||
return nil, fmt.Errorf("encoded Subnets should be in pairs, an odd number was found")
|
|
||||||
}
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: rc.Details.Name,
|
|
||||||
groups: make([]string, len(rc.Details.Groups)),
|
|
||||||
networks: make([]netip.Prefix, len(rc.Details.Ips)/2),
|
|
||||||
unsafeNetworks: make([]netip.Prefix, len(rc.Details.Subnets)/2),
|
|
||||||
notBefore: time.Unix(rc.Details.NotBefore, 0),
|
|
||||||
notAfter: time.Unix(rc.Details.NotAfter, 0),
|
|
||||||
publicKey: nil,
|
|
||||||
isCA: rc.Details.IsCA,
|
|
||||||
curve: rc.Details.Curve,
|
|
||||||
},
|
|
||||||
signature: make([]byte, len(rc.Signature)),
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(nc.signature, rc.Signature)
|
|
||||||
copy(nc.details.groups, rc.Details.Groups)
|
|
||||||
nc.details.issuer = hex.EncodeToString(rc.Details.Issuer)
|
|
||||||
|
|
||||||
// If a public key is passed in as an argument, the certificate pubkey must be empty
|
|
||||||
// and the passed-in pubkey copied into the cert.
|
|
||||||
if len(publicKey) > 0 {
|
|
||||||
if len(rc.Details.PublicKey) != 0 {
|
|
||||||
return nil, ErrCertPubkeyPresent
|
|
||||||
}
|
|
||||||
nc.details.publicKey = make([]byte, len(publicKey))
|
|
||||||
copy(nc.details.publicKey, publicKey)
|
|
||||||
} else {
|
|
||||||
nc.details.publicKey = make([]byte, len(rc.Details.PublicKey))
|
|
||||||
copy(nc.details.publicKey, rc.Details.PublicKey)
|
|
||||||
}
|
|
||||||
|
|
||||||
var ip netip.Addr
|
|
||||||
for i, rawIp := range rc.Details.Ips {
|
|
||||||
if i%2 == 0 {
|
|
||||||
ip = int2addr(rawIp)
|
|
||||||
} else {
|
|
||||||
ones, _ := net.IPMask(int2ip(rawIp)).Size()
|
|
||||||
nc.details.networks[i/2] = netip.PrefixFrom(ip, ones)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, rawIp := range rc.Details.Subnets {
|
|
||||||
if i%2 == 0 {
|
|
||||||
ip = int2addr(rawIp)
|
|
||||||
} else {
|
|
||||||
ones, _ := net.IPMask(int2ip(rawIp)).Size()
|
|
||||||
nc.details.unsafeNetworks[i/2] = netip.PrefixFrom(ip, ones)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = nc.validate()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &nc, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func ip2int(ip []byte) uint32 {
|
|
||||||
if len(ip) == 16 {
|
|
||||||
return binary.BigEndian.Uint32(ip[12:16])
|
|
||||||
}
|
|
||||||
return binary.BigEndian.Uint32(ip)
|
|
||||||
}
|
|
||||||
|
|
||||||
func int2ip(nn uint32) net.IP {
|
|
||||||
ip := make(net.IP, net.IPv4len)
|
|
||||||
binary.BigEndian.PutUint32(ip, nn)
|
|
||||||
return ip
|
|
||||||
}
|
|
||||||
|
|
||||||
func addr2int(addr netip.Addr) uint32 {
|
|
||||||
b := addr.Unmap().As4()
|
|
||||||
return binary.BigEndian.Uint32(b[:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func int2addr(nn uint32) netip.Addr {
|
|
||||||
ip := [4]byte{}
|
|
||||||
binary.BigEndian.PutUint32(ip[:], nn)
|
|
||||||
return netip.AddrFrom4(ip).Unmap()
|
|
||||||
}
|
|
||||||
@@ -1,334 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ed25519"
|
|
||||||
"fmt"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"google.golang.org/protobuf/proto"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestCertificateV1_Marshal(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
publicKey: pubKey,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := nc.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
//t.Log("Cert size:", len(b))
|
|
||||||
|
|
||||||
nc2, err := unmarshalCertificateV1(b, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, Version1, nc.Version())
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
assert.Equal(t, nc.Signature(), nc2.Signature())
|
|
||||||
assert.Equal(t, nc.Name(), nc2.Name())
|
|
||||||
assert.Equal(t, nc.NotBefore(), nc2.NotBefore())
|
|
||||||
assert.Equal(t, nc.NotAfter(), nc2.NotAfter())
|
|
||||||
assert.Equal(t, nc.PublicKey(), nc2.PublicKey())
|
|
||||||
assert.Equal(t, nc.IsCA(), nc2.IsCA())
|
|
||||||
|
|
||||||
assert.Equal(t, nc.Networks(), nc2.Networks())
|
|
||||||
assert.Equal(t, nc.UnsafeNetworks(), nc2.UnsafeNetworks())
|
|
||||||
|
|
||||||
assert.Equal(t, nc.Groups(), nc2.Groups())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Unmarshal(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
invalidPubkey := []byte("00000000000000000000000000000000")
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
publicKey: pubKey,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
// This certificate has a pubkey included
|
|
||||||
certWithPubkey, err := nc.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// This certificate is missing the pubkey section
|
|
||||||
certWithoutPubkey, err := nc.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Cert has no pubkey and no pubkey passed in must fail to validate
|
|
||||||
isNil, err := unmarshalCertificateV1(certWithoutPubkey, nil)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
// Cert has different pubkey than one passed in must fail
|
|
||||||
isNil, err = unmarshalCertificateV1(certWithPubkey, invalidPubkey)
|
|
||||||
require.Nil(t, isNil)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
// Cert has pubkey and no pubkey argument works ok
|
|
||||||
_, err = unmarshalCertificateV1(certWithPubkey, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Cert has no pubkey and valid, correctly signed pubkey passed in
|
|
||||||
nc2, err := unmarshalCertificateV1(certWithoutPubkey, pubKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, pubKey, nc2.PublicKey())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_PublicKeyPem(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := ed25519.PublicKey("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{},
|
|
||||||
unsafeNetworks: []netip.Prefix{},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
publicKey: pubKey,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Equal(t, Version1, nc.Version())
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
pubPem := "-----BEGIN NEBULA X25519 PUBLIC KEY-----\nMTIzNDU2Nzg5MGFiY2VkZmdoaWoxMjM0NTY3ODkwYWI=\n-----END NEBULA X25519 PUBLIC KEY-----\n"
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), pubPem)
|
|
||||||
assert.False(t, nc.IsCA())
|
|
||||||
|
|
||||||
nc.details.isCA = true
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
pubPem = "-----BEGIN NEBULA ED25519 PUBLIC KEY-----\nMTIzNDU2Nzg5MGFiY2VkZmdoaWoxMjM0NTY3ODkwYWI=\n-----END NEBULA ED25519 PUBLIC KEY-----\n"
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), pubPem)
|
|
||||||
assert.True(t, nc.IsCA())
|
|
||||||
|
|
||||||
pubP256KeyPem := []byte(`-----BEGIN NEBULA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
|
|
||||||
pubP256KeyPemCA := []byte(`-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
pubP256Key, _, _, err := UnmarshalPublicKeyFromPEM(pubP256KeyPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
nc.details.curve = Curve_P256
|
|
||||||
nc.details.publicKey = pubP256Key
|
|
||||||
assert.Equal(t, Curve_P256, nc.Curve())
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), string(pubP256KeyPemCA))
|
|
||||||
assert.True(t, nc.IsCA())
|
|
||||||
|
|
||||||
nc.details.isCA = false
|
|
||||||
assert.Equal(t, Curve_P256, nc.Curve())
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), string(pubP256KeyPem))
|
|
||||||
assert.False(t, nc.IsCA())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Expired(t *testing.T) {
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
notBefore: time.Now().Add(time.Second * -60).Round(time.Second),
|
|
||||||
notAfter: time.Now().Add(time.Second * 60).Round(time.Second),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.True(t, nc.Expired(time.Now().Add(time.Hour)))
|
|
||||||
assert.True(t, nc.Expired(time.Now().Add(-time.Hour)))
|
|
||||||
assert.False(t, nc.Expired(time.Now()))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_MarshalJSON(t *testing.T) {
|
|
||||||
time.Local = time.UTC
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: time.Date(1, 0, 0, 1, 0, 0, 0, time.UTC),
|
|
||||||
notAfter: time.Date(1, 0, 0, 2, 0, 0, 0, time.UTC),
|
|
||||||
publicKey: pubKey,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := nc.MarshalJSON()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.JSONEq(
|
|
||||||
t,
|
|
||||||
"{\"details\":{\"curve\":\"CURVE25519\",\"groups\":[\"test-group1\",\"test-group2\",\"test-group3\"],\"isCa\":false,\"issuer\":\"1234567890abcedfghij1234567890ab\",\"name\":\"testing\",\"networks\":[\"10.1.1.1/24\",\"10.1.1.2/16\"],\"notAfter\":\"0000-11-30T02:00:00Z\",\"notBefore\":\"0000-11-30T01:00:00Z\",\"publicKey\":\"313233343536373839306162636564666768696a313233343536373839306162\",\"unsafeNetworks\":[\"9.1.1.2/24\",\"9.1.1.3/16\"]},\"fingerprint\":\"3944c53d4267a229295b56cb2d27d459164c010ac97d655063ba421e0670f4ba\",\"signature\":\"313233343536373839306162636564666768696a313233343536373839306162\",\"version\":1}",
|
|
||||||
string(b),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_VerifyPrivateKey(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
err := ca.VerifyPrivateKey(Curve_CURVE25519, caKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, _, caKey2, _ := NewTestCaCert(Version1, Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
err = ca.VerifyPrivateKey(Curve_CURVE25519, caKey2)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
c, _, priv, _ := NewTestCert(Version1, Curve_CURVE25519, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
rawPriv, b, curve, err := UnmarshalPrivateKeyFromPEM(priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
err = c.VerifyPrivateKey(Curve_CURVE25519, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, priv2 := X25519Keypair()
|
|
||||||
err = c.VerifyPrivateKey(Curve_CURVE25519, priv2)
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_VerifyPrivateKeyP256(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version1, Curve_P256, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
err := ca.VerifyPrivateKey(Curve_P256, caKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, _, caKey2, _ := NewTestCaCert(Version1, Curve_P256, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
err = ca.VerifyPrivateKey(Curve_P256, caKey2)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
c, _, priv, _ := NewTestCert(Version1, Curve_P256, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
rawPriv, b, curve, err := UnmarshalPrivateKeyFromPEM(priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, priv2 := P256Keypair()
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, priv2)
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ensure that upgrading the protobuf library does not change how certificates
|
|
||||||
// are marshalled, since this would break signature verification
|
|
||||||
func TestMarshalingCertificateV1Consistency(t *testing.T) {
|
|
||||||
before := time.Date(1970, time.January, 1, 1, 1, 1, 1, time.UTC)
|
|
||||||
after := time.Date(9999, time.January, 1, 1, 1, 1, 1, time.UTC)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV1{
|
|
||||||
details: detailsV1{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
publicKey: pubKey,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := nc.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, "0a8e010a0774657374696e671212828284508080fcff0f8182845080feffff0f1a12838284488080fcff0f8282844880feffff0f220b746573742d67726f757031220b746573742d67726f757032220b746573742d67726f75703328cd1c30cdb8ccf0af073a20313233343536373839306162636564666768696a3132333435363738393061624a081234567890abcedf1220313233343536373839306162636564666768696a313233343536373839306162", fmt.Sprintf("%x", b))
|
|
||||||
|
|
||||||
b, err = proto.Marshal(nc.getRawDetails())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, "0a0774657374696e671212828284508080fcff0f8182845080feffff0f1a12838284488080fcff0f8282844880feffff0f220b746573742d67726f757031220b746573742d67726f757032220b746573742d67726f75703328cd1c30cdb8ccf0af073a20313233343536373839306162636564666768696a3132333435363738393061624a081234567890abcedf", fmt.Sprintf("%x", b))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_Copy(t *testing.T) {
|
|
||||||
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", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
|
||||||
cc := c.Copy()
|
|
||||||
test.AssertDeepCopyEqual(t, c, cc)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalCertificateV1(t *testing.T) {
|
|
||||||
// Test that we don't panic with an invalid certificate (#332)
|
|
||||||
data := []byte("\x98\x00\x00")
|
|
||||||
_, err := unmarshalCertificateV1(data, nil)
|
|
||||||
require.EqualError(t, err, "encoded Details was nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
func appendByteSlices(b ...[]byte) []byte {
|
|
||||||
retSlice := []byte{}
|
|
||||||
for _, v := range b {
|
|
||||||
retSlice = append(retSlice, v...)
|
|
||||||
}
|
|
||||||
return retSlice
|
|
||||||
}
|
|
||||||
|
|
||||||
func mustParsePrefixUnmapped(s string) netip.Prefix {
|
|
||||||
prefix := netip.MustParsePrefix(s)
|
|
||||||
return netip.PrefixFrom(prefix.Addr().Unmap(), prefix.Bits())
|
|
||||||
}
|
|
||||||
@@ -1,37 +0,0 @@
|
|||||||
Nebula DEFINITIONS AUTOMATIC TAGS ::= BEGIN
|
|
||||||
|
|
||||||
Name ::= UTF8String (SIZE (1..253))
|
|
||||||
Time ::= INTEGER (0..18446744073709551615) -- Seconds since unix epoch, uint64 maximum
|
|
||||||
Network ::= OCTET STRING (SIZE (5,17)) -- IP addresses are 4 or 16 bytes + 1 byte for the prefix length
|
|
||||||
Curve ::= ENUMERATED {
|
|
||||||
curve25519 (0),
|
|
||||||
p256 (1)
|
|
||||||
}
|
|
||||||
|
|
||||||
-- The maximum size of a certificate must not exceed 65536 bytes
|
|
||||||
Certificate ::= SEQUENCE {
|
|
||||||
details OCTET STRING,
|
|
||||||
curve Curve DEFAULT curve25519,
|
|
||||||
publicKey OCTET STRING,
|
|
||||||
-- signature(details + curve + publicKey) using the appropriate method for curve
|
|
||||||
signature OCTET STRING
|
|
||||||
}
|
|
||||||
|
|
||||||
Details ::= SEQUENCE {
|
|
||||||
name Name,
|
|
||||||
|
|
||||||
-- At least 1 ipv4 or ipv6 address must be present if isCA is false
|
|
||||||
networks SEQUENCE OF Network OPTIONAL,
|
|
||||||
unsafeNetworks SEQUENCE OF Network OPTIONAL,
|
|
||||||
groups SEQUENCE OF Name OPTIONAL,
|
|
||||||
isCA BOOLEAN DEFAULT false,
|
|
||||||
notBefore Time,
|
|
||||||
notAfter Time,
|
|
||||||
|
|
||||||
-- issuer is only required if isCA is false, if isCA is true then it must not be present
|
|
||||||
issuer OCTET STRING OPTIONAL,
|
|
||||||
...
|
|
||||||
-- New fields can be added below here
|
|
||||||
}
|
|
||||||
|
|
||||||
END
|
|
||||||
-742
@@ -1,742 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/ecdh"
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
|
||||||
"encoding/json"
|
|
||||||
"encoding/pem"
|
|
||||||
"fmt"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"golang.org/x/crypto/cryptobyte"
|
|
||||||
"golang.org/x/crypto/cryptobyte/asn1"
|
|
||||||
"golang.org/x/crypto/curve25519"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
classConstructed = 0x20
|
|
||||||
classContextSpecific = 0x80
|
|
||||||
|
|
||||||
TagCertDetails = 0 | classConstructed | classContextSpecific
|
|
||||||
TagCertCurve = 1 | classContextSpecific
|
|
||||||
TagCertPublicKey = 2 | classContextSpecific
|
|
||||||
TagCertSignature = 3 | classContextSpecific
|
|
||||||
|
|
||||||
TagDetailsName = 0 | classContextSpecific
|
|
||||||
TagDetailsNetworks = 1 | classConstructed | classContextSpecific
|
|
||||||
TagDetailsUnsafeNetworks = 2 | classConstructed | classContextSpecific
|
|
||||||
TagDetailsGroups = 3 | classConstructed | classContextSpecific
|
|
||||||
TagDetailsIsCA = 4 | classContextSpecific
|
|
||||||
TagDetailsNotBefore = 5 | classContextSpecific
|
|
||||||
TagDetailsNotAfter = 6 | classContextSpecific
|
|
||||||
TagDetailsIssuer = 7 | classContextSpecific
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// MaxCertificateSize is the maximum length a valid certificate can be
|
|
||||||
MaxCertificateSize = 65536
|
|
||||||
|
|
||||||
// MaxNameLength is limited to a maximum realistic DNS domain name to help facilitate DNS systems
|
|
||||||
MaxNameLength = 253
|
|
||||||
|
|
||||||
// MaxNetworkLength is the maximum length a network value can be.
|
|
||||||
// 16 bytes for an ipv6 address + 1 byte for the prefix length
|
|
||||||
MaxNetworkLength = 17
|
|
||||||
)
|
|
||||||
|
|
||||||
type certificateV2 struct {
|
|
||||||
details detailsV2
|
|
||||||
|
|
||||||
// RawDetails contains the entire asn.1 DER encoded Details struct
|
|
||||||
// This is to benefit forwards compatibility in signature checking.
|
|
||||||
// signature(RawDetails + Curve + PublicKey) == Signature
|
|
||||||
rawDetails []byte
|
|
||||||
curve Curve
|
|
||||||
publicKey []byte
|
|
||||||
signature []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type detailsV2 struct {
|
|
||||||
name string
|
|
||||||
networks []netip.Prefix // MUST BE SORTED
|
|
||||||
unsafeNetworks []netip.Prefix // MUST BE SORTED
|
|
||||||
groups []string
|
|
||||||
isCA bool
|
|
||||||
notBefore time.Time
|
|
||||||
notAfter time.Time
|
|
||||||
issuer string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Version() Version {
|
|
||||||
return Version2
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Curve() Curve {
|
|
||||||
return c.curve
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Groups() []string {
|
|
||||||
return c.details.groups
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) IsCA() bool {
|
|
||||||
return c.details.isCA
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Issuer() string {
|
|
||||||
return c.details.issuer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Name() string {
|
|
||||||
return c.details.name
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Networks() []netip.Prefix {
|
|
||||||
return c.details.networks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) NotAfter() time.Time {
|
|
||||||
return c.details.notAfter
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) NotBefore() time.Time {
|
|
||||||
return c.details.notBefore
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) PublicKey() []byte {
|
|
||||||
return c.publicKey
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) MarshalPublicKeyPEM() []byte {
|
|
||||||
return marshalCertPublicKeyToPEM(c)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Signature() []byte {
|
|
||||||
return c.signature
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) UnsafeNetworks() []netip.Prefix {
|
|
||||||
return c.details.unsafeNetworks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Fingerprint() (string, error) {
|
|
||||||
if len(c.rawDetails) == 0 {
|
|
||||||
return "", ErrMissingDetails
|
|
||||||
}
|
|
||||||
|
|
||||||
b := make([]byte, len(c.rawDetails)+1+len(c.publicKey)+len(c.signature))
|
|
||||||
copy(b, c.rawDetails)
|
|
||||||
b[len(c.rawDetails)] = byte(c.curve)
|
|
||||||
copy(b[len(c.rawDetails)+1:], c.publicKey)
|
|
||||||
copy(b[len(c.rawDetails)+1+len(c.publicKey):], c.signature)
|
|
||||||
sum := sha256.Sum256(b)
|
|
||||||
return hex.EncodeToString(sum[:]), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) CheckSignature(key []byte) bool {
|
|
||||||
if len(c.rawDetails) == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
b := make([]byte, len(c.rawDetails)+1+len(c.publicKey))
|
|
||||||
copy(b, c.rawDetails)
|
|
||||||
b[len(c.rawDetails)] = byte(c.curve)
|
|
||||||
copy(b[len(c.rawDetails)+1:], c.publicKey)
|
|
||||||
|
|
||||||
switch c.curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
|
||||||
case Curve_P256:
|
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
hashed := sha256.Sum256(b)
|
|
||||||
return ecdsa.VerifyASN1(pubKey, hashed[:], c.signature)
|
|
||||||
default:
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Expired(t time.Time) bool {
|
|
||||||
return c.details.notBefore.After(t) || c.details.notAfter.Before(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) VerifyPrivateKey(curve Curve, key []byte) error {
|
|
||||||
if curve != c.curve {
|
|
||||||
return ErrPublicPrivateCurveMismatch
|
|
||||||
}
|
|
||||||
if c.details.isCA {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
// the call to PublicKey below will panic slice bounds out of range otherwise
|
|
||||||
if len(key) != ed25519.PrivateKeySize {
|
|
||||||
return ErrInvalidPrivateKey
|
|
||||||
}
|
|
||||||
|
|
||||||
if !ed25519.PublicKey(c.publicKey).Equal(ed25519.PrivateKey(key).Public()) {
|
|
||||||
return ErrPublicPrivateKeyMismatch
|
|
||||||
}
|
|
||||||
case Curve_P256:
|
|
||||||
privkey, err := ecdh.P256().NewPrivateKey(key)
|
|
||||||
if err != nil {
|
|
||||||
return ErrInvalidPrivateKey
|
|
||||||
}
|
|
||||||
pub := privkey.PublicKey().Bytes()
|
|
||||||
if !bytes.Equal(pub, c.publicKey) {
|
|
||||||
return ErrPublicPrivateKeyMismatch
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", curve)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var pub []byte
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
var err error
|
|
||||||
pub, err = curve25519.X25519(key, curve25519.Basepoint)
|
|
||||||
if err != nil {
|
|
||||||
return ErrInvalidPrivateKey
|
|
||||||
}
|
|
||||||
case Curve_P256:
|
|
||||||
privkey, err := ecdh.P256().NewPrivateKey(key)
|
|
||||||
if err != nil {
|
|
||||||
return ErrInvalidPrivateKey
|
|
||||||
}
|
|
||||||
pub = privkey.PublicKey().Bytes()
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", curve)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(pub, c.publicKey) {
|
|
||||||
return ErrPublicPrivateKeyMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) String() string {
|
|
||||||
mb, err := c.marshalJSON()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("<error marshalling certificate: %v>", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := json.MarshalIndent(mb, "", "\t")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Sprintf("<error marshalling certificate: %v>", err)
|
|
||||||
}
|
|
||||||
return string(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) MarshalForHandshakes() ([]byte, error) {
|
|
||||||
if c.rawDetails == nil {
|
|
||||||
return nil, ErrEmptyRawDetails
|
|
||||||
}
|
|
||||||
var b cryptobyte.Builder
|
|
||||||
// Outermost certificate
|
|
||||||
b.AddASN1(asn1.SEQUENCE, func(b *cryptobyte.Builder) {
|
|
||||||
|
|
||||||
// Add the cert details which is already marshalled
|
|
||||||
b.AddBytes(c.rawDetails)
|
|
||||||
|
|
||||||
// Skipping the curve and public key since those come across in a different part of the handshake
|
|
||||||
|
|
||||||
// Add the signature
|
|
||||||
b.AddASN1(TagCertSignature, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes(c.signature)
|
|
||||||
})
|
|
||||||
})
|
|
||||||
|
|
||||||
return b.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Marshal() ([]byte, error) {
|
|
||||||
if c.rawDetails == nil {
|
|
||||||
return nil, ErrEmptyRawDetails
|
|
||||||
}
|
|
||||||
var b cryptobyte.Builder
|
|
||||||
// Outermost certificate
|
|
||||||
b.AddASN1(asn1.SEQUENCE, func(b *cryptobyte.Builder) {
|
|
||||||
|
|
||||||
// Add the cert details which is already marshalled
|
|
||||||
b.AddBytes(c.rawDetails)
|
|
||||||
|
|
||||||
// Add the curve only if its not the default value
|
|
||||||
if c.curve != Curve_CURVE25519 {
|
|
||||||
b.AddASN1(TagCertCurve, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes([]byte{byte(c.curve)})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add the public key if it is not empty
|
|
||||||
if c.publicKey != nil {
|
|
||||||
b.AddASN1(TagCertPublicKey, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes(c.publicKey)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add the signature
|
|
||||||
b.AddASN1(TagCertSignature, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes(c.signature)
|
|
||||||
})
|
|
||||||
})
|
|
||||||
|
|
||||||
return b.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) MarshalPEM() ([]byte, error) {
|
|
||||||
b, err := c.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: CertificateV2Banner, Bytes: b}), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) MarshalJSON() ([]byte, error) {
|
|
||||||
b, err := c.marshalJSON()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return json.Marshal(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) marshalJSON() (m, error) {
|
|
||||||
fp, err := c.Fingerprint()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return m{
|
|
||||||
"details": m{
|
|
||||||
"name": c.details.name,
|
|
||||||
"networks": c.details.networks,
|
|
||||||
"unsafeNetworks": c.details.unsafeNetworks,
|
|
||||||
"groups": c.details.groups,
|
|
||||||
"notBefore": c.details.notBefore,
|
|
||||||
"notAfter": c.details.notAfter,
|
|
||||||
"isCa": c.details.isCA,
|
|
||||||
"issuer": c.details.issuer,
|
|
||||||
},
|
|
||||||
"version": Version2,
|
|
||||||
"publicKey": fmt.Sprintf("%x", c.publicKey),
|
|
||||||
"curve": c.curve.String(),
|
|
||||||
"fingerprint": fp,
|
|
||||||
"signature": fmt.Sprintf("%x", c.Signature()),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) Copy() Certificate {
|
|
||||||
nc := &certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: c.details.name,
|
|
||||||
notBefore: c.details.notBefore,
|
|
||||||
notAfter: c.details.notAfter,
|
|
||||||
isCA: c.details.isCA,
|
|
||||||
issuer: c.details.issuer,
|
|
||||||
},
|
|
||||||
curve: c.curve,
|
|
||||||
publicKey: make([]byte, len(c.publicKey)),
|
|
||||||
signature: make([]byte, len(c.signature)),
|
|
||||||
rawDetails: make([]byte, len(c.rawDetails)),
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.groups != nil {
|
|
||||||
nc.details.groups = make([]string, len(c.details.groups))
|
|
||||||
copy(nc.details.groups, c.details.groups)
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.networks != nil {
|
|
||||||
nc.details.networks = make([]netip.Prefix, len(c.details.networks))
|
|
||||||
copy(nc.details.networks, c.details.networks)
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.details.unsafeNetworks != nil {
|
|
||||||
nc.details.unsafeNetworks = make([]netip.Prefix, len(c.details.unsafeNetworks))
|
|
||||||
copy(nc.details.unsafeNetworks, c.details.unsafeNetworks)
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(nc.rawDetails, c.rawDetails)
|
|
||||||
copy(nc.signature, c.signature)
|
|
||||||
copy(nc.publicKey, c.publicKey)
|
|
||||||
|
|
||||||
return nc
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) fromTBSCertificate(t *TBSCertificate) error {
|
|
||||||
c.details = detailsV2{
|
|
||||||
name: t.Name,
|
|
||||||
networks: t.Networks,
|
|
||||||
unsafeNetworks: t.UnsafeNetworks,
|
|
||||||
groups: t.Groups,
|
|
||||||
isCA: t.IsCA,
|
|
||||||
notBefore: t.NotBefore,
|
|
||||||
notAfter: t.NotAfter,
|
|
||||||
issuer: t.issuer,
|
|
||||||
}
|
|
||||||
c.curve = t.Curve
|
|
||||||
c.publicKey = t.PublicKey
|
|
||||||
return c.validate()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) validate() error {
|
|
||||||
// Empty names are allowed
|
|
||||||
|
|
||||||
if len(c.publicKey) == 0 {
|
|
||||||
return ErrInvalidPublicKey
|
|
||||||
}
|
|
||||||
|
|
||||||
if !c.details.isCA && len(c.details.networks) == 0 {
|
|
||||||
return NewErrInvalidCertificateProperties("non-CA certificate must contain at least 1 network")
|
|
||||||
}
|
|
||||||
|
|
||||||
hasV4Networks := false
|
|
||||||
hasV6Networks := false
|
|
||||||
for _, network := range c.details.networks {
|
|
||||||
if !network.IsValid() || !network.Addr().IsValid() {
|
|
||||||
return NewErrInvalidCertificateProperties("invalid network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().IsUnspecified() {
|
|
||||||
return NewErrInvalidCertificateProperties("non-CA certificates must not use the zero address as a network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Zone() != "" {
|
|
||||||
return NewErrInvalidCertificateProperties("networks may not contain zones: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Is4In6() {
|
|
||||||
return NewErrInvalidCertificateProperties("4in6 networks are not allowed: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
hasV4Networks = hasV4Networks || network.Addr().Is4()
|
|
||||||
hasV6Networks = hasV6Networks || network.Addr().Is6()
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortFunc(c.details.networks, comparePrefix)
|
|
||||||
err := findDuplicatePrefix(c.details.networks)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, network := range c.details.unsafeNetworks {
|
|
||||||
if !network.IsValid() || !network.Addr().IsValid() {
|
|
||||||
return NewErrInvalidCertificateProperties("invalid unsafe network: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if network.Addr().Zone() != "" {
|
|
||||||
return NewErrInvalidCertificateProperties("unsafe networks may not contain zones: %s", network)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !c.details.isCA {
|
|
||||||
if network.Addr().Is6() {
|
|
||||||
if !hasV6Networks {
|
|
||||||
return NewErrInvalidCertificateProperties("IPv6 unsafe networks require an IPv6 address assignment: %s", network)
|
|
||||||
}
|
|
||||||
} else if network.Addr().Is4() {
|
|
||||||
if !hasV4Networks {
|
|
||||||
return NewErrInvalidCertificateProperties("IPv4 unsafe networks require an IPv4 address assignment: %s", network)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortFunc(c.details.unsafeNetworks, comparePrefix)
|
|
||||||
err = findDuplicatePrefix(c.details.unsafeNetworks)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) marshalForSigning() ([]byte, error) {
|
|
||||||
d, err := c.details.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("marshalling certificate details failed: %w", err)
|
|
||||||
}
|
|
||||||
c.rawDetails = d
|
|
||||||
|
|
||||||
b := make([]byte, len(c.rawDetails)+1+len(c.publicKey))
|
|
||||||
copy(b, c.rawDetails)
|
|
||||||
b[len(c.rawDetails)] = byte(c.curve)
|
|
||||||
copy(b[len(c.rawDetails)+1:], c.publicKey)
|
|
||||||
return b, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *certificateV2) setSignature(b []byte) error {
|
|
||||||
if len(b) == 0 {
|
|
||||||
return ErrEmptySignature
|
|
||||||
}
|
|
||||||
c.signature = b
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *detailsV2) Marshal() ([]byte, error) {
|
|
||||||
var b cryptobyte.Builder
|
|
||||||
var err error
|
|
||||||
|
|
||||||
// Details are a structure
|
|
||||||
b.AddASN1(TagCertDetails, func(b *cryptobyte.Builder) {
|
|
||||||
|
|
||||||
// Add the name
|
|
||||||
b.AddASN1(TagDetailsName, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes([]byte(d.name))
|
|
||||||
})
|
|
||||||
|
|
||||||
// Add the networks if any exist
|
|
||||||
if len(d.networks) > 0 {
|
|
||||||
b.AddASN1(TagDetailsNetworks, func(b *cryptobyte.Builder) {
|
|
||||||
for _, n := range d.networks {
|
|
||||||
sb, innerErr := n.MarshalBinary()
|
|
||||||
if innerErr != nil {
|
|
||||||
// MarshalBinary never returns an error
|
|
||||||
err = fmt.Errorf("unable to marshal network: %w", innerErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
b.AddASN1OctetString(sb)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add the unsafe networks if any exist
|
|
||||||
if len(d.unsafeNetworks) > 0 {
|
|
||||||
b.AddASN1(TagDetailsUnsafeNetworks, func(b *cryptobyte.Builder) {
|
|
||||||
for _, n := range d.unsafeNetworks {
|
|
||||||
sb, innerErr := n.MarshalBinary()
|
|
||||||
if innerErr != nil {
|
|
||||||
// MarshalBinary never returns an error
|
|
||||||
err = fmt.Errorf("unable to marshal unsafe network: %w", innerErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
b.AddASN1OctetString(sb)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add groups if any exist
|
|
||||||
if len(d.groups) > 0 {
|
|
||||||
b.AddASN1(TagDetailsGroups, func(b *cryptobyte.Builder) {
|
|
||||||
for _, group := range d.groups {
|
|
||||||
b.AddASN1(asn1.UTF8String, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes([]byte(group))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add IsCA only if true
|
|
||||||
if d.isCA {
|
|
||||||
b.AddASN1(TagDetailsIsCA, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddUint8(0xff)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add not before
|
|
||||||
b.AddASN1Int64WithTag(d.notBefore.Unix(), TagDetailsNotBefore)
|
|
||||||
|
|
||||||
// Add not after
|
|
||||||
b.AddASN1Int64WithTag(d.notAfter.Unix(), TagDetailsNotAfter)
|
|
||||||
|
|
||||||
// Add the issuer if present
|
|
||||||
if d.issuer != "" {
|
|
||||||
issuerBytes, innerErr := hex.DecodeString(d.issuer)
|
|
||||||
if innerErr != nil {
|
|
||||||
err = fmt.Errorf("failed to decode issuer: %w", innerErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
b.AddASN1(TagDetailsIssuer, func(b *cryptobyte.Builder) {
|
|
||||||
b.AddBytes(issuerBytes)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return b.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
func unmarshalCertificateV2(b []byte, publicKey []byte, curve Curve) (*certificateV2, error) {
|
|
||||||
l := len(b)
|
|
||||||
if l == 0 || l > MaxCertificateSize {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
input := cryptobyte.String(b)
|
|
||||||
// Open the envelope
|
|
||||||
if !input.ReadASN1(&input, asn1.SEQUENCE) || input.Empty() {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Grab the cert details, we need to preserve the tag and length
|
|
||||||
var rawDetails cryptobyte.String
|
|
||||||
if !input.ReadASN1Element(&rawDetails, TagCertDetails) || rawDetails.Empty() {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
//Maybe grab the curve
|
|
||||||
var rawCurve byte
|
|
||||||
if !readOptionalASN1Byte(&input, &rawCurve, TagCertCurve, byte(curve)) {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
curve = Curve(rawCurve)
|
|
||||||
|
|
||||||
// Maybe grab the public key
|
|
||||||
var rawPublicKey cryptobyte.String
|
|
||||||
if len(publicKey) > 0 {
|
|
||||||
// If a public key is passed in, then the handshake certificate must
|
|
||||||
// not have a public key present
|
|
||||||
if input.PeekASN1Tag(TagCertPublicKey) {
|
|
||||||
return nil, ErrCertPubkeyPresent
|
|
||||||
}
|
|
||||||
rawPublicKey = make(cryptobyte.String, len(publicKey))
|
|
||||||
copy(rawPublicKey, publicKey)
|
|
||||||
} else if !input.ReadOptionalASN1(&rawPublicKey, nil, TagCertPublicKey) {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(rawPublicKey) == 0 {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Grab the signature
|
|
||||||
var rawSignature cryptobyte.String
|
|
||||||
if !input.ReadASN1(&rawSignature, TagCertSignature) || rawSignature.Empty() {
|
|
||||||
return nil, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Finally unmarshal the details
|
|
||||||
details, err := unmarshalDetails(rawDetails)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
c := &certificateV2{
|
|
||||||
details: details,
|
|
||||||
rawDetails: rawDetails,
|
|
||||||
curve: curve,
|
|
||||||
publicKey: rawPublicKey,
|
|
||||||
signature: rawSignature,
|
|
||||||
}
|
|
||||||
|
|
||||||
err = c.validate()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func unmarshalDetails(b cryptobyte.String) (detailsV2, error) {
|
|
||||||
// Open the envelope
|
|
||||||
if !b.ReadASN1(&b, TagCertDetails) || b.Empty() {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read the name
|
|
||||||
var name cryptobyte.String
|
|
||||||
if !b.ReadASN1(&name, TagDetailsName) || name.Empty() || len(name) > MaxNameLength {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read the network addresses
|
|
||||||
var subString cryptobyte.String
|
|
||||||
var found bool
|
|
||||||
|
|
||||||
if !b.ReadOptionalASN1(&subString, &found, TagDetailsNetworks) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var networks []netip.Prefix
|
|
||||||
var val cryptobyte.String
|
|
||||||
if found {
|
|
||||||
for !subString.Empty() {
|
|
||||||
if !subString.ReadASN1(&val, asn1.OCTET_STRING) || val.Empty() || len(val) > MaxNetworkLength {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var n netip.Prefix
|
|
||||||
if err := n.UnmarshalBinary(val); err != nil {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
networks = append(networks, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read out any unsafe networks
|
|
||||||
if !b.ReadOptionalASN1(&subString, &found, TagDetailsUnsafeNetworks) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var unsafeNetworks []netip.Prefix
|
|
||||||
if found {
|
|
||||||
for !subString.Empty() {
|
|
||||||
if !subString.ReadASN1(&val, asn1.OCTET_STRING) || val.Empty() || len(val) > MaxNetworkLength {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var n netip.Prefix
|
|
||||||
if err := n.UnmarshalBinary(val); err != nil {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
unsafeNetworks = append(unsafeNetworks, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read out any groups
|
|
||||||
if !b.ReadOptionalASN1(&subString, &found, TagDetailsGroups) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var groups []string
|
|
||||||
if found {
|
|
||||||
for !subString.Empty() {
|
|
||||||
if !subString.ReadASN1(&val, asn1.UTF8String) || val.Empty() {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
groups = append(groups, string(val))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read out IsCA
|
|
||||||
var isCa bool
|
|
||||||
if !readOptionalASN1Boolean(&b, &isCa, TagDetailsIsCA, false) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read not before and not after
|
|
||||||
var notBefore int64
|
|
||||||
if !b.ReadASN1Int64WithTag(¬Before, TagDetailsNotBefore) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
var notAfter int64
|
|
||||||
if !b.ReadASN1Int64WithTag(¬After, TagDetailsNotAfter) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read issuer
|
|
||||||
var issuer cryptobyte.String
|
|
||||||
if !b.ReadOptionalASN1(&issuer, nil, TagDetailsIssuer) {
|
|
||||||
return detailsV2{}, ErrBadFormat
|
|
||||||
}
|
|
||||||
|
|
||||||
return detailsV2{
|
|
||||||
name: string(name),
|
|
||||||
networks: networks,
|
|
||||||
unsafeNetworks: unsafeNetworks,
|
|
||||||
groups: groups,
|
|
||||||
isCA: isCa,
|
|
||||||
notBefore: time.Unix(notBefore, 0),
|
|
||||||
notAfter: time.Unix(notAfter, 0),
|
|
||||||
issuer: hex.EncodeToString(issuer),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
@@ -1,379 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/rand"
|
|
||||||
"encoding/hex"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/test"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestCertificateV2_Marshal(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcdef1234567890abcdef",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcdef1234567890abcdef"),
|
|
||||||
publicKey: pubKey,
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := nc.details.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
nc.rawDetails = db
|
|
||||||
|
|
||||||
b, err := nc.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
//t.Log("Cert size:", len(b))
|
|
||||||
|
|
||||||
nc2, err := unmarshalCertificateV2(b, nil, Curve_CURVE25519)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, Version2, nc.Version())
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
assert.Equal(t, nc.Signature(), nc2.Signature())
|
|
||||||
assert.Equal(t, nc.Name(), nc2.Name())
|
|
||||||
assert.Equal(t, nc.NotBefore(), nc2.NotBefore())
|
|
||||||
assert.Equal(t, nc.NotAfter(), nc2.NotAfter())
|
|
||||||
assert.Equal(t, nc.PublicKey(), nc2.PublicKey())
|
|
||||||
assert.Equal(t, nc.IsCA(), nc2.IsCA())
|
|
||||||
assert.Equal(t, nc.Issuer(), nc2.Issuer())
|
|
||||||
|
|
||||||
// unmarshalling will sort networks and unsafeNetworks, we need to do the same
|
|
||||||
// but first make sure it fails
|
|
||||||
assert.NotEqual(t, nc.Networks(), nc2.Networks())
|
|
||||||
assert.NotEqual(t, nc.UnsafeNetworks(), nc2.UnsafeNetworks())
|
|
||||||
|
|
||||||
slices.SortFunc(nc.details.networks, comparePrefix)
|
|
||||||
slices.SortFunc(nc.details.unsafeNetworks, comparePrefix)
|
|
||||||
|
|
||||||
assert.Equal(t, nc.Networks(), nc2.Networks())
|
|
||||||
assert.Equal(t, nc.UnsafeNetworks(), nc2.UnsafeNetworks())
|
|
||||||
|
|
||||||
assert.Equal(t, nc.Groups(), nc2.Groups())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Unmarshal(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcdef1234567890abcdef",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcdef1234567890abcdef"),
|
|
||||||
publicKey: pubKey,
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := nc.details.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
nc.rawDetails = db
|
|
||||||
|
|
||||||
certWithPubkey, err := nc.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
//t.Log("Cert size:", len(b))
|
|
||||||
certWithoutPubkey, err := nc.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Cert must not have a pubkey if one is passed in as an argument
|
|
||||||
_, err = unmarshalCertificateV2(certWithPubkey, pubKey, Curve_CURVE25519)
|
|
||||||
require.ErrorIs(t, err, ErrCertPubkeyPresent)
|
|
||||||
|
|
||||||
// Certs must have pubkeys
|
|
||||||
_, err = unmarshalCertificateV2(certWithoutPubkey, nil, Curve_CURVE25519)
|
|
||||||
require.ErrorIs(t, err, ErrBadFormat)
|
|
||||||
|
|
||||||
// Ensure proper unmarshal if a pubkey is passed in
|
|
||||||
nc2, err := unmarshalCertificateV2(certWithoutPubkey, pubKey, Curve_CURVE25519)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
assert.Equal(t, nc.PublicKey(), nc2.PublicKey())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_PublicKeyPem(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := ed25519.PublicKey("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{},
|
|
||||||
unsafeNetworks: []netip.Prefix{},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedfghij1234567890ab",
|
|
||||||
},
|
|
||||||
publicKey: pubKey,
|
|
||||||
signature: []byte("1234567890abcedfghij1234567890ab"),
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Equal(t, Version2, nc.Version())
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
pubPem := "-----BEGIN NEBULA X25519 PUBLIC KEY-----\nMTIzNDU2Nzg5MGFiY2VkZmdoaWoxMjM0NTY3ODkwYWI=\n-----END NEBULA X25519 PUBLIC KEY-----\n"
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), pubPem)
|
|
||||||
assert.False(t, nc.IsCA())
|
|
||||||
|
|
||||||
nc.details.isCA = true
|
|
||||||
assert.Equal(t, Curve_CURVE25519, nc.Curve())
|
|
||||||
pubPem = "-----BEGIN NEBULA ED25519 PUBLIC KEY-----\nMTIzNDU2Nzg5MGFiY2VkZmdoaWoxMjM0NTY3ODkwYWI=\n-----END NEBULA ED25519 PUBLIC KEY-----\n"
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), pubPem)
|
|
||||||
assert.True(t, nc.IsCA())
|
|
||||||
|
|
||||||
pubP256KeyPem := []byte(`-----BEGIN NEBULA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
|
|
||||||
pubP256KeyPemCA := []byte(`-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
|
|
||||||
pubP256Key, _, _, err := UnmarshalPublicKeyFromPEM(pubP256KeyPem)
|
|
||||||
require.NoError(t, err)
|
|
||||||
nc.curve = Curve_P256
|
|
||||||
nc.publicKey = pubP256Key
|
|
||||||
assert.Equal(t, Curve_P256, nc.Curve())
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), string(pubP256KeyPemCA))
|
|
||||||
assert.True(t, nc.IsCA())
|
|
||||||
|
|
||||||
nc.details.isCA = false
|
|
||||||
assert.Equal(t, Curve_P256, nc.Curve())
|
|
||||||
assert.Equal(t, string(nc.MarshalPublicKeyPEM()), string(pubP256KeyPem))
|
|
||||||
assert.False(t, nc.IsCA())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Expired(t *testing.T) {
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
notBefore: time.Now().Add(time.Second * -60).Round(time.Second),
|
|
||||||
notAfter: time.Now().Add(time.Second * 60).Round(time.Second),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.True(t, nc.Expired(time.Now().Add(time.Hour)))
|
|
||||||
assert.True(t, nc.Expired(time.Now().Add(-time.Hour)))
|
|
||||||
assert.False(t, nc.Expired(time.Now()))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_MarshalJSON(t *testing.T) {
|
|
||||||
time.Local = time.UTC
|
|
||||||
pubKey := []byte("1234567890abcedf1234567890abcedf")
|
|
||||||
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: time.Date(1, 0, 0, 1, 0, 0, 0, time.UTC),
|
|
||||||
notAfter: time.Date(1, 0, 0, 2, 0, 0, 0, time.UTC),
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcedf1234567890abcedf",
|
|
||||||
},
|
|
||||||
publicKey: pubKey,
|
|
||||||
signature: []byte("1234567890abcedf1234567890abcedf1234567890abcedf1234567890abcedf"),
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err := nc.MarshalJSON()
|
|
||||||
require.ErrorIs(t, err, ErrMissingDetails)
|
|
||||||
|
|
||||||
rd, err := nc.details.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
nc.rawDetails = rd
|
|
||||||
b, err = nc.MarshalJSON()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.JSONEq(
|
|
||||||
t,
|
|
||||||
"{\"curve\":\"CURVE25519\",\"details\":{\"groups\":[\"test-group1\",\"test-group2\",\"test-group3\"],\"isCa\":false,\"issuer\":\"1234567890abcedf1234567890abcedf\",\"name\":\"testing\",\"networks\":[\"10.1.1.1/24\",\"10.1.1.2/16\"],\"notAfter\":\"0000-11-30T02:00:00Z\",\"notBefore\":\"0000-11-30T01:00:00Z\",\"unsafeNetworks\":[\"9.1.1.2/24\",\"9.1.1.3/16\"]},\"fingerprint\":\"152d9a7400c1e001cb76cffd035215ebb351f69eeb797f7f847dd086e15e56dd\",\"publicKey\":\"3132333435363738393061626365646631323334353637383930616263656466\",\"signature\":\"31323334353637383930616263656466313233343536373839306162636564663132333435363738393061626365646631323334353637383930616263656466\",\"version\":2}",
|
|
||||||
string(b),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_VerifyPrivateKey(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
err := ca.VerifyPrivateKey(Curve_CURVE25519, caKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = ca.VerifyPrivateKey(Curve_CURVE25519, caKey[:16])
|
|
||||||
require.ErrorIs(t, err, ErrInvalidPrivateKey)
|
|
||||||
|
|
||||||
_, caKey2, err := ed25519.GenerateKey(rand.Reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
err = ca.VerifyPrivateKey(Curve_CURVE25519, caKey2)
|
|
||||||
require.ErrorIs(t, err, ErrPublicPrivateKeyMismatch)
|
|
||||||
|
|
||||||
c, _, priv, _ := NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
rawPriv, b, curve, err := UnmarshalPrivateKeyFromPEM(priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
err = c.VerifyPrivateKey(Curve_CURVE25519, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, priv2 := X25519Keypair()
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, priv2)
|
|
||||||
require.ErrorIs(t, err, ErrPublicPrivateCurveMismatch)
|
|
||||||
|
|
||||||
err = c.VerifyPrivateKey(Curve_CURVE25519, priv2)
|
|
||||||
require.ErrorIs(t, err, ErrPublicPrivateKeyMismatch)
|
|
||||||
|
|
||||||
err = c.VerifyPrivateKey(Curve_CURVE25519, priv2[:16])
|
|
||||||
require.ErrorIs(t, err, ErrInvalidPrivateKey)
|
|
||||||
|
|
||||||
ac, ok := c.(*certificateV2)
|
|
||||||
require.True(t, ok)
|
|
||||||
ac.curve = Curve(99)
|
|
||||||
err = c.VerifyPrivateKey(Curve(99), priv2)
|
|
||||||
require.EqualError(t, err, "invalid curve: 99")
|
|
||||||
|
|
||||||
ca2, _, caKey2, _ := NewTestCaCert(Version2, Curve_P256, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
err = ca.VerifyPrivateKey(Curve_CURVE25519, caKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = ca2.VerifyPrivateKey(Curve_P256, caKey2[:16])
|
|
||||||
require.ErrorIs(t, err, ErrInvalidPrivateKey)
|
|
||||||
|
|
||||||
c, _, priv, _ = NewTestCert(Version2, Curve_P256, ca2, caKey2, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
rawPriv, b, curve, err = UnmarshalPrivateKeyFromPEM(priv)
|
|
||||||
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, priv[:16])
|
|
||||||
require.ErrorIs(t, err, ErrInvalidPrivateKey)
|
|
||||||
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, priv)
|
|
||||||
require.ErrorIs(t, err, ErrInvalidPrivateKey)
|
|
||||||
|
|
||||||
aCa, ok := ca2.(*certificateV2)
|
|
||||||
require.True(t, ok)
|
|
||||||
aCa.curve = Curve(99)
|
|
||||||
err = aCa.VerifyPrivateKey(Curve(99), priv2)
|
|
||||||
require.EqualError(t, err, "invalid curve: 99")
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_VerifyPrivateKeyP256(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_P256, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
err := ca.VerifyPrivateKey(Curve_P256, caKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, _, caKey2, _ := NewTestCaCert(Version2, Curve_P256, time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
err = ca.VerifyPrivateKey(Curve_P256, caKey2)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
c, _, priv, _ := NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Time{}, time.Time{}, nil, nil, nil)
|
|
||||||
rawPriv, b, curve, err := UnmarshalPrivateKeyFromPEM(priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, b)
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, priv2 := P256Keypair()
|
|
||||||
err = c.VerifyPrivateKey(Curve_P256, priv2)
|
|
||||||
require.Error(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_Copy(t *testing.T) {
|
|
||||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, nil)
|
|
||||||
c, _, _, _ := NewTestCert(Version2, Curve_CURVE25519, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), nil, nil, nil)
|
|
||||||
cc := c.Copy()
|
|
||||||
test.AssertDeepCopyEqual(t, c, cc)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalCertificateV2(t *testing.T) {
|
|
||||||
data := []byte("\x98\x00\x00")
|
|
||||||
_, err := unmarshalCertificateV2(data, nil, Curve_CURVE25519)
|
|
||||||
require.EqualError(t, err, "bad wire format")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV2_marshalForSigningStability(t *testing.T) {
|
|
||||||
before := time.Date(1996, time.May, 5, 0, 0, 0, 0, time.UTC)
|
|
||||||
after := before.Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
nc := certificateV2{
|
|
||||||
details: detailsV2{
|
|
||||||
name: "testing",
|
|
||||||
networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
},
|
|
||||||
unsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
},
|
|
||||||
groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
notBefore: before,
|
|
||||||
notAfter: after,
|
|
||||||
isCA: false,
|
|
||||||
issuer: "1234567890abcdef1234567890abcdef",
|
|
||||||
},
|
|
||||||
signature: []byte("1234567890abcdef1234567890abcdef"),
|
|
||||||
publicKey: pubKey,
|
|
||||||
}
|
|
||||||
|
|
||||||
const expectedRawDetailsStr = "a070800774657374696e67a10e04050a0101021004050a01010118a20e0405090101031004050901010218a3270c0b746573742d67726f7570310c0b746573742d67726f7570320c0b746573742d67726f7570338504318bef808604318befbc87101234567890abcdef1234567890abcdef"
|
|
||||||
expectedRawDetails, err := hex.DecodeString(expectedRawDetailsStr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := nc.details.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, expectedRawDetails, db)
|
|
||||||
|
|
||||||
expectedForSigning, err := hex.DecodeString(expectedRawDetailsStr + "00313233343536373839306162636564666768696a313233343536373839306162")
|
|
||||||
b, err := nc.marshalForSigning()
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, expectedForSigning, b)
|
|
||||||
}
|
|
||||||
+2
-162
@@ -3,28 +3,14 @@ package cert
|
|||||||
import (
|
import (
|
||||||
"crypto/aes"
|
"crypto/aes"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/pem"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math"
|
|
||||||
|
|
||||||
"golang.org/x/crypto/argon2"
|
"golang.org/x/crypto/argon2"
|
||||||
"google.golang.org/protobuf/proto"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type NebulaEncryptedData struct {
|
// KDF factors
|
||||||
EncryptionMetadata NebulaEncryptionMetadata
|
|
||||||
Ciphertext []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type NebulaEncryptionMetadata struct {
|
|
||||||
EncryptionAlgorithm string
|
|
||||||
Argon2Parameters Argon2Parameters
|
|
||||||
}
|
|
||||||
|
|
||||||
// Argon2Parameters KDF factors
|
|
||||||
type Argon2Parameters struct {
|
type Argon2Parameters struct {
|
||||||
version rune
|
version rune
|
||||||
Memory uint32 // KiB
|
Memory uint32 // KiB
|
||||||
@@ -33,7 +19,7 @@ type Argon2Parameters struct {
|
|||||||
salt []byte
|
salt []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewArgon2Parameters Returns a new Argon2Parameters object with current version set
|
// Returns a new Argon2Parameters object with current version set
|
||||||
func NewArgon2Parameters(memory uint32, parallelism uint8, iterations uint32) *Argon2Parameters {
|
func NewArgon2Parameters(memory uint32, parallelism uint8, iterations uint32) *Argon2Parameters {
|
||||||
return &Argon2Parameters{
|
return &Argon2Parameters{
|
||||||
version: argon2.Version,
|
version: argon2.Version,
|
||||||
@@ -91,9 +77,6 @@ func aes256Decrypt(passphrase []byte, kdfParams *Argon2Parameters, data []byte)
|
|||||||
}
|
}
|
||||||
|
|
||||||
gcm, err := cipher.NewGCM(block)
|
gcm, err := cipher.NewGCM(block)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
nonce, ciphertext, err := splitNonceCiphertext(data, gcm.NonceSize())
|
nonce, ciphertext, err := splitNonceCiphertext(data, gcm.NonceSize())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -155,146 +138,3 @@ func splitNonceCiphertext(blob []byte, nonceSize int) ([]byte, []byte, error) {
|
|||||||
|
|
||||||
return blob[:nonceSize], blob[nonceSize:], nil
|
return blob[:nonceSize], blob[nonceSize:], nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// EncryptAndMarshalSigningPrivateKey is a simple helper to encrypt and PEM encode a private key
|
|
||||||
func EncryptAndMarshalSigningPrivateKey(curve Curve, b []byte, passphrase []byte, kdfParams *Argon2Parameters) ([]byte, error) {
|
|
||||||
ciphertext, err := aes256Encrypt(passphrase, kdfParams, b)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
b, err = proto.Marshal(&RawNebulaEncryptedData{
|
|
||||||
EncryptionMetadata: &RawNebulaEncryptionMetadata{
|
|
||||||
EncryptionAlgorithm: "AES-256-GCM",
|
|
||||||
Argon2Parameters: &RawNebulaArgon2Parameters{
|
|
||||||
Version: kdfParams.version,
|
|
||||||
Memory: kdfParams.Memory,
|
|
||||||
Parallelism: uint32(kdfParams.Parallelism),
|
|
||||||
Iterations: kdfParams.Iterations,
|
|
||||||
Salt: kdfParams.salt,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Ciphertext: ciphertext,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: EncryptedEd25519PrivateKeyBanner, Bytes: b}), nil
|
|
||||||
case Curve_P256:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: EncryptedECDSAP256PrivateKeyBanner, Bytes: b}), nil
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("invalid curve: %v", curve)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnmarshalNebulaEncryptedData will unmarshal a protobuf byte representation of a nebula cert into its
|
|
||||||
// protobuf-generated struct.
|
|
||||||
func UnmarshalNebulaEncryptedData(b []byte) (*NebulaEncryptedData, error) {
|
|
||||||
if len(b) == 0 {
|
|
||||||
return nil, fmt.Errorf("nil byte array")
|
|
||||||
}
|
|
||||||
var rned RawNebulaEncryptedData
|
|
||||||
err := proto.Unmarshal(b, &rned)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if rned.EncryptionMetadata == nil {
|
|
||||||
return nil, fmt.Errorf("encoded EncryptionMetadata was nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
if rned.EncryptionMetadata.Argon2Parameters == nil {
|
|
||||||
return nil, fmt.Errorf("encoded Argon2Parameters was nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
params, err := unmarshalArgon2Parameters(rned.EncryptionMetadata.Argon2Parameters)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
ned := NebulaEncryptedData{
|
|
||||||
EncryptionMetadata: NebulaEncryptionMetadata{
|
|
||||||
EncryptionAlgorithm: rned.EncryptionMetadata.EncryptionAlgorithm,
|
|
||||||
Argon2Parameters: *params,
|
|
||||||
},
|
|
||||||
Ciphertext: rned.Ciphertext,
|
|
||||||
}
|
|
||||||
|
|
||||||
return &ned, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func unmarshalArgon2Parameters(params *RawNebulaArgon2Parameters) (*Argon2Parameters, error) {
|
|
||||||
if params.Version < math.MinInt32 || params.Version > math.MaxInt32 {
|
|
||||||
return nil, fmt.Errorf("Argon2Parameters Version must be at least %d and no more than %d", math.MinInt32, math.MaxInt32)
|
|
||||||
}
|
|
||||||
if params.Memory <= 0 || params.Memory > math.MaxUint32 {
|
|
||||||
return nil, fmt.Errorf("Argon2Parameters Memory must be be greater than 0 and no more than %d KiB", uint32(math.MaxUint32))
|
|
||||||
}
|
|
||||||
if params.Parallelism <= 0 || params.Parallelism > math.MaxUint8 {
|
|
||||||
return nil, fmt.Errorf("Argon2Parameters Parallelism must be be greater than 0 and no more than %d", math.MaxUint8)
|
|
||||||
}
|
|
||||||
if params.Iterations <= 0 || params.Iterations > math.MaxUint32 {
|
|
||||||
return nil, fmt.Errorf("-argon-iterations must be be greater than 0 and no more than %d", uint32(math.MaxUint32))
|
|
||||||
}
|
|
||||||
|
|
||||||
return &Argon2Parameters{
|
|
||||||
version: params.Version,
|
|
||||||
Memory: params.Memory,
|
|
||||||
Parallelism: uint8(params.Parallelism),
|
|
||||||
Iterations: params.Iterations,
|
|
||||||
salt: params.Salt,
|
|
||||||
}, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// DecryptAndUnmarshalSigningPrivateKey will try to pem decode and decrypt an Ed25519/ECDSA private key with
|
|
||||||
// the given passphrase, returning any other bytes b or an error on failure
|
|
||||||
func DecryptAndUnmarshalSigningPrivateKey(passphrase, b []byte) (Curve, []byte, []byte, error) {
|
|
||||||
var curve Curve
|
|
||||||
|
|
||||||
k, r := pem.Decode(b)
|
|
||||||
if k == nil {
|
|
||||||
return curve, nil, r, fmt.Errorf("input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
switch k.Type {
|
|
||||||
case EncryptedEd25519PrivateKeyBanner:
|
|
||||||
curve = Curve_CURVE25519
|
|
||||||
case EncryptedECDSAP256PrivateKeyBanner:
|
|
||||||
curve = Curve_P256
|
|
||||||
default:
|
|
||||||
return curve, nil, r, fmt.Errorf("bytes did not contain a proper nebula encrypted Ed25519/ECDSA private key banner")
|
|
||||||
}
|
|
||||||
|
|
||||||
ned, err := UnmarshalNebulaEncryptedData(k.Bytes)
|
|
||||||
if err != nil {
|
|
||||||
return curve, nil, r, err
|
|
||||||
}
|
|
||||||
|
|
||||||
var bytes []byte
|
|
||||||
switch ned.EncryptionMetadata.EncryptionAlgorithm {
|
|
||||||
case "AES-256-GCM":
|
|
||||||
bytes, err = aes256Decrypt(passphrase, &ned.EncryptionMetadata.Argon2Parameters, ned.Ciphertext)
|
|
||||||
if err != nil {
|
|
||||||
return curve, nil, r, err
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return curve, nil, r, fmt.Errorf("unsupported encryption algorithm: %s", ned.EncryptionMetadata.EncryptionAlgorithm)
|
|
||||||
}
|
|
||||||
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
if len(bytes) != ed25519.PrivateKeySize {
|
|
||||||
return curve, nil, r, fmt.Errorf("key was not %d bytes, is invalid ed25519 private key", ed25519.PrivateKeySize)
|
|
||||||
}
|
|
||||||
case Curve_P256:
|
|
||||||
if len(bytes) != 32 {
|
|
||||||
return curve, nil, r, fmt.Errorf("key was not 32 bytes, is invalid ECDSA P256 private key")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return curve, bytes, r, nil
|
|
||||||
}
|
|
||||||
|
|||||||
+2
-90
@@ -4,110 +4,22 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/crypto/argon2"
|
"golang.org/x/crypto/argon2"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewArgon2Parameters(t *testing.T) {
|
func TestNewArgon2Parameters(t *testing.T) {
|
||||||
p := NewArgon2Parameters(64*1024, 4, 3)
|
p := NewArgon2Parameters(64*1024, 4, 3)
|
||||||
assert.Equal(t, &Argon2Parameters{
|
assert.EqualValues(t, &Argon2Parameters{
|
||||||
version: argon2.Version,
|
version: argon2.Version,
|
||||||
Memory: 64 * 1024,
|
Memory: 64 * 1024,
|
||||||
Parallelism: 4,
|
Parallelism: 4,
|
||||||
Iterations: 3,
|
Iterations: 3,
|
||||||
}, p)
|
}, p)
|
||||||
p = NewArgon2Parameters(2*1024*1024, 2, 1)
|
p = NewArgon2Parameters(2*1024*1024, 2, 1)
|
||||||
assert.Equal(t, &Argon2Parameters{
|
assert.EqualValues(t, &Argon2Parameters{
|
||||||
version: argon2.Version,
|
version: argon2.Version,
|
||||||
Memory: 2 * 1024 * 1024,
|
Memory: 2 * 1024 * 1024,
|
||||||
Parallelism: 2,
|
Parallelism: 2,
|
||||||
Iterations: 1,
|
Iterations: 1,
|
||||||
}, p)
|
}, p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDecryptAndUnmarshalSigningPrivateKey(t *testing.T) {
|
|
||||||
passphrase := []byte("DO NOT USE")
|
|
||||||
privKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
CjsKC0FFUy0yNTYtR0NNEiwIExCAgAQYAyAEKiCPoDfGQiosxNPTbPn5EsMlc2MI
|
|
||||||
c0Bt4oz6gTrFQhX3aBJcimhHKeAuhyTGvllD0Z19fe+DFPcLH3h5VrdjVfIAajg0
|
|
||||||
KrbV3n9UHif/Au5skWmquNJzoW1E4MTdRbvpti6o+WdQ49DxjBFhx0YH8LBqrbPU
|
|
||||||
0BGkUHmIO7daP24=
|
|
||||||
-----END NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A key which, once decrypted, is too short
|
|
||||||
-----BEGIN NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
CjsKC0FFUy0yNTYtR0NNEiwIExCAgAQYAyAEKiAVJwdfl3r+eqi/vF6S7OMdpjfo
|
|
||||||
hAzmTCRnr58Su4AqmBJbCv3zleYCEKYJP6UI3S8ekLMGISsgO4hm5leukCCyqT0Z
|
|
||||||
cQ76yrberpzkJKoPLGisX8f+xdy4aXSZl7oEYWQte1+vqbtl/eY9PGZhxUQdcyq7
|
|
||||||
hqzIyrRqfUgVuA==
|
|
||||||
-----END NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner (not encrypted)
|
|
||||||
-----BEGIN NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
bWRp2CTVFhW9HD/qCd28ltDgK3w8VXSeaEYczDWos8sMUBqDb9jP3+NYwcS4lURG
|
|
||||||
XgLvodMXZJuaFPssp+WwtA==
|
|
||||||
-----END NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
CjwKC0FFUy0yNTYtR0NNEi0IExCAgIABGAEgBCognnjujd67Vsv99p22wfAjQaDT
|
|
||||||
oCMW1mdjkU3gACKNW4MSXOWR9Sts4C81yk1RUku2gvGKs3TB9LYoklLsIizSYOLl
|
|
||||||
+Vs//O1T0I1Xbml2XBAROsb/VSoDln/6LMqR4B6fn6B3GOsLBBqRI8daDl9lRMPB
|
|
||||||
qrlJ69wer3ZUHFXA
|
|
||||||
-END NEBULA ED25519 ENCRYPTED PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(privKey, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
curve, k, rest, err := DecryptAndUnmarshalSigningPrivateKey(passphrase, keyBundle)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Len(t, k, 64)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
curve, k, rest, err = DecryptAndUnmarshalSigningPrivateKey(passphrase, rest)
|
|
||||||
require.EqualError(t, err, "key was not 64 bytes, is invalid ed25519 private key")
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
curve, k, rest, err = DecryptAndUnmarshalSigningPrivateKey(passphrase, rest)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper nebula encrypted Ed25519/ECDSA private key banner")
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
curve, k, rest, err = DecryptAndUnmarshalSigningPrivateKey(passphrase, rest)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
|
|
||||||
// Fail due to invalid passphrase
|
|
||||||
curve, k, rest, err = DecryptAndUnmarshalSigningPrivateKey([]byte("invalid passphrase"), privKey)
|
|
||||||
require.EqualError(t, err, "invalid passphrase or corrupt private key")
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, []byte{}, rest)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEncryptAndMarshalSigningPrivateKey(t *testing.T) {
|
|
||||||
// Having proved that decryption works correctly above, we can test the
|
|
||||||
// encryption function produces a value which can be decrypted
|
|
||||||
passphrase := []byte("passphrase")
|
|
||||||
bytes := []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA")
|
|
||||||
kdfParams := NewArgon2Parameters(64*1024, 4, 3)
|
|
||||||
key, err := EncryptAndMarshalSigningPrivateKey(Curve_CURVE25519, bytes, passphrase, kdfParams)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Verify the "key" can be decrypted successfully
|
|
||||||
curve, k, rest, err := DecryptAndUnmarshalSigningPrivateKey(passphrase, key)
|
|
||||||
assert.Len(t, k, 64)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Equal(t, []byte{}, rest)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// EncryptAndMarshalEd25519PrivateKey does not create any errors itself
|
|
||||||
}
|
|
||||||
|
|||||||
+6
-43
@@ -2,50 +2,13 @@ package cert
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
ErrBadFormat = errors.New("bad wire format")
|
ErrRootExpired = errors.New("root certificate is expired")
|
||||||
ErrRootExpired = errors.New("root certificate is expired")
|
ErrExpired = errors.New("certificate is expired")
|
||||||
ErrExpired = errors.New("certificate is expired")
|
ErrNotCA = errors.New("certificate is not a CA")
|
||||||
ErrNotCA = errors.New("certificate is not a CA")
|
ErrNotSelfSigned = errors.New("certificate is not self-signed")
|
||||||
ErrNotSelfSigned = errors.New("certificate is not self-signed")
|
ErrBlockListed = errors.New("certificate is in the block list")
|
||||||
ErrBlockListed = errors.New("certificate is in the block list")
|
ErrSignatureMismatch = errors.New("certificate signature did not match")
|
||||||
ErrFingerprintMismatch = errors.New("certificate fingerprint did not match")
|
|
||||||
ErrSignatureMismatch = errors.New("certificate signature did not match")
|
|
||||||
ErrInvalidPublicKey = errors.New("invalid public key")
|
|
||||||
ErrInvalidPrivateKey = errors.New("invalid private key")
|
|
||||||
ErrPublicPrivateCurveMismatch = errors.New("public key does not match private key curve")
|
|
||||||
ErrPublicPrivateKeyMismatch = errors.New("public key and private key are not a pair")
|
|
||||||
ErrPrivateKeyEncrypted = errors.New("private key must be decrypted")
|
|
||||||
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
|
||||||
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
|
||||||
ErrCertPubkeyPresent = errors.New("certificate has unexpected pubkey present")
|
|
||||||
|
|
||||||
ErrInvalidPEMBlock = errors.New("input did not contain a valid PEM encoded block")
|
|
||||||
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
|
||||||
ErrInvalidPEMX25519PublicKeyBanner = errors.New("bytes did not contain a proper X25519 public key banner")
|
|
||||||
ErrInvalidPEMX25519PrivateKeyBanner = errors.New("bytes did not contain a proper X25519 private key banner")
|
|
||||||
ErrInvalidPEMEd25519PublicKeyBanner = errors.New("bytes did not contain a proper Ed25519 public key banner")
|
|
||||||
ErrInvalidPEMEd25519PrivateKeyBanner = errors.New("bytes did not contain a proper Ed25519 private key banner")
|
|
||||||
|
|
||||||
ErrNoPeerStaticKey = errors.New("no peer static key was present")
|
|
||||||
ErrNoPayload = errors.New("provided payload was empty")
|
|
||||||
|
|
||||||
ErrMissingDetails = errors.New("certificate did not contain details")
|
|
||||||
ErrEmptySignature = errors.New("empty signature")
|
|
||||||
ErrEmptyRawDetails = errors.New("empty rawDetails not allowed")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type ErrInvalidCertificateProperties struct {
|
|
||||||
str string
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewErrInvalidCertificateProperties(format string, a ...any) error {
|
|
||||||
return &ErrInvalidCertificateProperties{fmt.Sprintf(format, a...)}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *ErrInvalidCertificateProperties) Error() string {
|
|
||||||
return e.str
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,141 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ecdh"
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"golang.org/x/crypto/curve25519"
|
|
||||||
"golang.org/x/crypto/ed25519"
|
|
||||||
)
|
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
|
||||||
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
|
||||||
var err error
|
|
||||||
var pub, priv []byte
|
|
||||||
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
pub, priv, err = ed25519.GenerateKey(rand.Reader)
|
|
||||||
case Curve_P256:
|
|
||||||
privk, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub = elliptic.Marshal(elliptic.P256(), privk.PublicKey.X, privk.PublicKey.Y)
|
|
||||||
priv = privk.D.FillBytes(make([]byte, 32))
|
|
||||||
default:
|
|
||||||
// There is no default to allow the underlying lib to respond with an error
|
|
||||||
}
|
|
||||||
|
|
||||||
if before.IsZero() {
|
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
}
|
|
||||||
if after.IsZero() {
|
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
t := &TBSCertificate{
|
|
||||||
Curve: curve,
|
|
||||||
Version: version,
|
|
||||||
Name: "test ca",
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pub,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
IsCA: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := t.Sign(nil, curve, priv)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pem, err := c.MarshalPEM()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, pub, priv, pem
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewTestCert will generate a signed certificate with the provided details.
|
|
||||||
// Expiry times are defaulted if you do not pass them in
|
|
||||||
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
|
|
||||||
if before.IsZero() {
|
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
if after.IsZero() {
|
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(networks) == 0 {
|
|
||||||
networks = []netip.Prefix{netip.MustParsePrefix("10.0.0.123/8")}
|
|
||||||
}
|
|
||||||
|
|
||||||
var pub, priv []byte
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
pub, priv = X25519Keypair()
|
|
||||||
case Curve_P256:
|
|
||||||
pub, priv = P256Keypair()
|
|
||||||
default:
|
|
||||||
panic("unknown curve")
|
|
||||||
}
|
|
||||||
|
|
||||||
nc := &TBSCertificate{
|
|
||||||
Version: v,
|
|
||||||
Curve: curve,
|
|
||||||
Name: name,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pub,
|
|
||||||
IsCA: false,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := nc.Sign(ca, ca.Curve(), key)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pem, err := c.MarshalPEM()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, pub, MarshalPrivateKeyToPEM(curve, priv), pem
|
|
||||||
}
|
|
||||||
|
|
||||||
func X25519Keypair() ([]byte, []byte) {
|
|
||||||
privkey := make([]byte, 32)
|
|
||||||
if _, err := io.ReadFull(rand.Reader, privkey); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pubkey, err := curve25519.X25519(privkey, curve25519.Basepoint)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return pubkey, privkey
|
|
||||||
}
|
|
||||||
|
|
||||||
func P256Keypair() ([]byte, []byte) {
|
|
||||||
privkey, err := ecdh.P256().GenerateKey(rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
pubkey := privkey.PublicKey()
|
|
||||||
return pubkey.Bytes(), privkey.Bytes()
|
|
||||||
}
|
|
||||||
@@ -1,127 +0,0 @@
|
|||||||
package p256
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/elliptic"
|
|
||||||
"errors"
|
|
||||||
"math/big"
|
|
||||||
|
|
||||||
"filippo.io/bigmod"
|
|
||||||
|
|
||||||
"golang.org/x/crypto/cryptobyte"
|
|
||||||
"golang.org/x/crypto/cryptobyte/asn1"
|
|
||||||
)
|
|
||||||
|
|
||||||
var halfN = new(big.Int).Rsh(elliptic.P256().Params().N, 1)
|
|
||||||
var nMod *bigmod.Modulus
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
n, err := bigmod.NewModulus(elliptic.P256().Params().N.Bytes())
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
nMod = n
|
|
||||||
}
|
|
||||||
|
|
||||||
func IsNormalized(sig []byte) (bool, error) {
|
|
||||||
r, s, err := parseSignature(sig)
|
|
||||||
if err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
return checkLowS(r, s), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func checkLowS(_, s []byte) bool {
|
|
||||||
bigS := new(big.Int).SetBytes(s)
|
|
||||||
// Check if S <= (N/2), because we want to include the midpoint in the set of low-s
|
|
||||||
return bigS.Cmp(halfN) <= 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func swap(r, s []byte) ([]byte, []byte, error) {
|
|
||||||
var err error
|
|
||||||
bigS, err := bigmod.NewNat().SetBytes(s, nMod)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
sNormalized := nMod.Nat().Sub(bigS, nMod)
|
|
||||||
|
|
||||||
result := sNormalized.Bytes(nMod)
|
|
||||||
for len(result) > 1 && result[0] == 0 {
|
|
||||||
result = result[1:]
|
|
||||||
}
|
|
||||||
|
|
||||||
return r, result, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func Normalize(sig []byte) ([]byte, error) {
|
|
||||||
r, s, err := parseSignature(sig)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if checkLowS(r, s) {
|
|
||||||
return sig, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
newR, newS, err := swap(r, s)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return encodeSignature(newR, newS)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Swap will change sig between its current form to the opposite high or low form.
|
|
||||||
func Swap(sig []byte) ([]byte, error) {
|
|
||||||
r, s, err := parseSignature(sig)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
newR, newS, err := swap(r, s)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return encodeSignature(newR, newS)
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseSignature taken exactly from crypto/ecdsa/ecdsa.go
|
|
||||||
func parseSignature(sig []byte) (r, s []byte, err error) {
|
|
||||||
var inner cryptobyte.String
|
|
||||||
input := cryptobyte.String(sig)
|
|
||||||
if !input.ReadASN1(&inner, asn1.SEQUENCE) ||
|
|
||||||
!input.Empty() ||
|
|
||||||
!inner.ReadASN1Integer(&r) ||
|
|
||||||
!inner.ReadASN1Integer(&s) ||
|
|
||||||
!inner.Empty() {
|
|
||||||
return nil, nil, errors.New("invalid ASN.1")
|
|
||||||
}
|
|
||||||
return r, s, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func encodeSignature(r, s []byte) ([]byte, error) {
|
|
||||||
var b cryptobyte.Builder
|
|
||||||
b.AddASN1(asn1.SEQUENCE, func(b *cryptobyte.Builder) {
|
|
||||||
addASN1IntBytes(b, r)
|
|
||||||
addASN1IntBytes(b, s)
|
|
||||||
})
|
|
||||||
return b.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
// addASN1IntBytes encodes in ASN.1 a positive integer represented as
|
|
||||||
// a big-endian byte slice with zero or more leading zeroes.
|
|
||||||
func addASN1IntBytes(b *cryptobyte.Builder, bytes []byte) {
|
|
||||||
for len(bytes) > 0 && bytes[0] == 0 {
|
|
||||||
bytes = bytes[1:]
|
|
||||||
}
|
|
||||||
if len(bytes) == 0 {
|
|
||||||
b.SetError(errors.New("invalid integer"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
b.AddASN1(asn1.INTEGER, func(c *cryptobyte.Builder) {
|
|
||||||
if bytes[0]&0x80 != 0 {
|
|
||||||
c.AddUint8(0)
|
|
||||||
}
|
|
||||||
c.AddBytes(bytes)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -1,28 +0,0 @@
|
|||||||
package p256
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestFlipping(t *testing.T) {
|
|
||||||
priv, err1 := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
require.NoError(t, err1)
|
|
||||||
|
|
||||||
out, err := ecdsa.SignASN1(rand.Reader, priv, []byte("big chungus"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
r, s, err := parseSignature(out)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
r, s1, err := swap(r, s)
|
|
||||||
require.NoError(t, err)
|
|
||||||
r, s2, err := swap(r, s1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, s, s2)
|
|
||||||
require.NotEqual(t, s, s1)
|
|
||||||
}
|
|
||||||
-247
@@ -1,247 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/pem"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"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
|
|
||||||
CertificateBanner = "NEBULA CERTIFICATE"
|
|
||||||
CertificateV2Banner = "NEBULA CERTIFICATE V2"
|
|
||||||
)
|
|
||||||
|
|
||||||
const ( //key-agreement-key banners
|
|
||||||
X25519PrivateKeyBanner = "NEBULA X25519 PRIVATE KEY"
|
|
||||||
X25519PublicKeyBanner = "NEBULA X25519 PUBLIC KEY"
|
|
||||||
P256PrivateKeyBanner = "NEBULA P256 PRIVATE KEY"
|
|
||||||
P256PublicKeyBanner = "NEBULA P256 PUBLIC KEY"
|
|
||||||
)
|
|
||||||
|
|
||||||
/* including "ECDSA" in the P256 banners is a clue that these keys should be used only for signing */
|
|
||||||
const ( //signing key banners
|
|
||||||
EncryptedECDSAP256PrivateKeyBanner = "NEBULA ECDSA P256 ENCRYPTED PRIVATE KEY"
|
|
||||||
ECDSAP256PrivateKeyBanner = "NEBULA ECDSA P256 PRIVATE KEY"
|
|
||||||
ECDSAP256PublicKeyBanner = "NEBULA ECDSA P256 PUBLIC KEY"
|
|
||||||
EncryptedEd25519PrivateKeyBanner = "NEBULA ED25519 ENCRYPTED PRIVATE KEY"
|
|
||||||
Ed25519PrivateKeyBanner = "NEBULA ED25519 PRIVATE KEY"
|
|
||||||
Ed25519PublicKeyBanner = "NEBULA ED25519 PUBLIC KEY"
|
|
||||||
)
|
|
||||||
|
|
||||||
// UnmarshalCertificateFromPEM will try to unmarshal the first pem block in a byte array, returning any non consumed
|
|
||||||
// data or an error on failure
|
|
||||||
func UnmarshalCertificateFromPEM(b []byte) (Certificate, []byte, error) {
|
|
||||||
p, r := pem.Decode(b)
|
|
||||||
if p == nil {
|
|
||||||
return nil, r, ErrInvalidPEMBlock
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := unmarshalCertificateBlock(p)
|
|
||||||
if err != nil {
|
|
||||||
return nil, r, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, r, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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 {
|
|
||||||
if c.IsCA() {
|
|
||||||
return MarshalSigningPublicKeyToPEM(c.Curve(), c.PublicKey())
|
|
||||||
} else {
|
|
||||||
return MarshalPublicKeyToPEM(c.Curve(), c.PublicKey())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarshalPublicKeyToPEM returns a PEM representation of a public key used for ECDH.
|
|
||||||
// if your public key came from a certificate, prefer Certificate.PublicKeyPEM() if possible, to avoid mistakes!
|
|
||||||
func MarshalPublicKeyToPEM(curve Curve, b []byte) []byte {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: X25519PublicKeyBanner, Bytes: b})
|
|
||||||
case Curve_P256:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: P256PublicKeyBanner, Bytes: b})
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// MarshalSigningPublicKeyToPEM returns a PEM representation of a public key used for signing.
|
|
||||||
// if your public key came from a certificate, prefer Certificate.PublicKeyPEM() if possible, to avoid mistakes!
|
|
||||||
func MarshalSigningPublicKeyToPEM(curve Curve, b []byte) []byte {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: Ed25519PublicKeyBanner, Bytes: b})
|
|
||||||
case Curve_P256:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: ECDSAP256PublicKeyBanner, Bytes: b})
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
|
||||||
k, r := pem.Decode(b)
|
|
||||||
if k == nil {
|
|
||||||
return nil, r, 0, fmt.Errorf("input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
var expectedLen int
|
|
||||||
var curve Curve
|
|
||||||
switch k.Type {
|
|
||||||
case X25519PublicKeyBanner, Ed25519PublicKeyBanner:
|
|
||||||
expectedLen = 32
|
|
||||||
curve = Curve_CURVE25519
|
|
||||||
case P256PublicKeyBanner, ECDSAP256PublicKeyBanner:
|
|
||||||
// Uncompressed
|
|
||||||
expectedLen = 65
|
|
||||||
curve = Curve_P256
|
|
||||||
default:
|
|
||||||
return nil, r, 0, fmt.Errorf("bytes did not contain a proper public key banner")
|
|
||||||
}
|
|
||||||
if len(k.Bytes) != expectedLen {
|
|
||||||
return nil, r, 0, fmt.Errorf("key was not %d bytes, is invalid %s public key", expectedLen, curve)
|
|
||||||
}
|
|
||||||
return k.Bytes, r, curve, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func MarshalPrivateKeyToPEM(curve Curve, b []byte) []byte {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: X25519PrivateKeyBanner, Bytes: b})
|
|
||||||
case Curve_P256:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: P256PrivateKeyBanner, Bytes: b})
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func MarshalSigningPrivateKeyToPEM(curve Curve, b []byte) []byte {
|
|
||||||
switch curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: Ed25519PrivateKeyBanner, Bytes: b})
|
|
||||||
case Curve_P256:
|
|
||||||
return pem.EncodeToMemory(&pem.Block{Type: ECDSAP256PrivateKeyBanner, Bytes: b})
|
|
||||||
default:
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// UnmarshalPrivateKeyFromPEM will try to unmarshal the first pem block in a byte array, returning any non
|
|
||||||
// consumed data or an error on failure
|
|
||||||
func UnmarshalPrivateKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
|
||||||
k, r := pem.Decode(b)
|
|
||||||
if k == nil {
|
|
||||||
return nil, r, 0, fmt.Errorf("input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
var expectedLen int
|
|
||||||
var curve Curve
|
|
||||||
switch k.Type {
|
|
||||||
case X25519PrivateKeyBanner:
|
|
||||||
expectedLen = 32
|
|
||||||
curve = Curve_CURVE25519
|
|
||||||
case P256PrivateKeyBanner:
|
|
||||||
expectedLen = 32
|
|
||||||
curve = Curve_P256
|
|
||||||
default:
|
|
||||||
return nil, r, 0, fmt.Errorf("bytes did not contain a proper private key banner")
|
|
||||||
}
|
|
||||||
if len(k.Bytes) != expectedLen {
|
|
||||||
return nil, r, 0, fmt.Errorf("key was not %d bytes, is invalid %s private key", expectedLen, curve)
|
|
||||||
}
|
|
||||||
return k.Bytes, r, curve, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func UnmarshalSigningPrivateKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
|
|
||||||
k, r := pem.Decode(b)
|
|
||||||
if k == nil {
|
|
||||||
return nil, r, 0, fmt.Errorf("input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
var curve Curve
|
|
||||||
switch k.Type {
|
|
||||||
case EncryptedEd25519PrivateKeyBanner:
|
|
||||||
return nil, nil, Curve_CURVE25519, ErrPrivateKeyEncrypted
|
|
||||||
case EncryptedECDSAP256PrivateKeyBanner:
|
|
||||||
return nil, nil, Curve_P256, ErrPrivateKeyEncrypted
|
|
||||||
case Ed25519PrivateKeyBanner:
|
|
||||||
curve = Curve_CURVE25519
|
|
||||||
if len(k.Bytes) != ed25519.PrivateKeySize {
|
|
||||||
return nil, r, 0, fmt.Errorf("key was not %d bytes, is invalid Ed25519 private key", ed25519.PrivateKeySize)
|
|
||||||
}
|
|
||||||
case ECDSAP256PrivateKeyBanner:
|
|
||||||
curve = Curve_P256
|
|
||||||
if len(k.Bytes) != 32 {
|
|
||||||
return nil, r, 0, fmt.Errorf("key was not 32 bytes, is invalid ECDSA P256 private key")
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return nil, r, 0, fmt.Errorf("bytes did not contain a proper Ed25519/ECDSA private key banner")
|
|
||||||
}
|
|
||||||
return k.Bytes, r, curve, nil
|
|
||||||
}
|
|
||||||
@@ -1,384 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"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) {
|
|
||||||
goodCert := []byte(`
|
|
||||||
# A good cert
|
|
||||||
-----BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CkAKDm5lYnVsYSByb290IGNhKJfap9AFMJfg1+YGOiCUQGByMuNRhIlQBOyzXWbL
|
|
||||||
vcKBwDhov900phEfJ5DN3kABEkDCq5R8qBiu8sl54yVfgRcQXEDt3cHr8UTSLszv
|
|
||||||
bzBEr00kERQxxTzTsH8cpYEgRoipvmExvg8WP8NdAJEYJosB
|
|
||||||
-----END NEBULA CERTIFICATE-----
|
|
||||||
`)
|
|
||||||
badBanner := []byte(`# A bad banner
|
|
||||||
-----BEGIN NOT A NEBULA CERTIFICATE-----
|
|
||||||
CkAKDm5lYnVsYSByb290IGNhKJfap9AFMJfg1+YGOiCUQGByMuNRhIlQBOyzXWbL
|
|
||||||
vcKBwDhov900phEfJ5DN3kABEkDCq5R8qBiu8sl54yVfgRcQXEDt3cHr8UTSLszv
|
|
||||||
bzBEr00kERQxxTzTsH8cpYEgRoipvmExvg8WP8NdAJEYJosB
|
|
||||||
-----END NOT A NEBULA CERTIFICATE-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA CERTIFICATE-----
|
|
||||||
CkAKDm5lYnVsYSByb290IGNhKJfap9AFMJfg1+YGOiCUQGByMuNRhIlQBOyzXWbL
|
|
||||||
vcKBwDhov900phEfJ5DN3kABEkDCq5R8qBiu8sl54yVfgRcQXEDt3cHr8UTSLszv
|
|
||||||
bzBEr00kERQxxTzTsH8cpYEgRoipvmExvg8WP8NdAJEYJosB
|
|
||||||
-END NEBULA CERTIFICATE----`)
|
|
||||||
|
|
||||||
certBundle := appendByteSlices(goodCert, badBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
cert, rest, err := UnmarshalCertificateFromPEM(certBundle)
|
|
||||||
assert.NotNil(t, cert)
|
|
||||||
assert.Equal(t, rest, append(badBanner, invalidPem...))
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Fail due to invalid banner.
|
|
||||||
cert, rest, err = UnmarshalCertificateFromPEM(rest)
|
|
||||||
assert.Nil(t, cert)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper certificate banner")
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
cert, rest, err = UnmarshalCertificateFromPEM(rest)
|
|
||||||
assert.Nil(t, cert)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalSigningPrivateKeyFromPEM(t *testing.T) {
|
|
||||||
privKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
privP256Key := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA ECDSA P256 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ECDSA P256 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A short key
|
|
||||||
-----BEGIN NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
-----END NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner
|
|
||||||
-----BEGIN NOT A NEBULA PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NOT A NEBULA PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA ED25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-END NEBULA ED25519 PRIVATE KEY-----`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(privKey, privP256Key, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err := UnmarshalSigningPrivateKeyFromPEM(keyBundle)
|
|
||||||
assert.Len(t, k, 64)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(privP256Key, shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err = UnmarshalSigningPrivateKeyFromPEM(rest)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, curve, err = UnmarshalSigningPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "key was not 64 bytes, is invalid Ed25519 private key")
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
k, rest, curve, err = UnmarshalSigningPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA private key banner")
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
k, rest, curve, err = UnmarshalSigningPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalPrivateKeyFromPEM(t *testing.T) {
|
|
||||||
privKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA X25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA X25519 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
privP256Key := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA P256 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA P256 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A short key
|
|
||||||
-----BEGIN NEBULA X25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NEBULA X25519 PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner
|
|
||||||
-----BEGIN NOT A NEBULA PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NOT A NEBULA PRIVATE KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA X25519 PRIVATE KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-END NEBULA X25519 PRIVATE KEY-----`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(privKey, privP256Key, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err := UnmarshalPrivateKeyFromPEM(keyBundle)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(privP256Key, shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err = UnmarshalPrivateKeyFromPEM(rest)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, curve, err = UnmarshalPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 private key")
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
k, rest, curve, err = UnmarshalPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper private key banner")
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
k, rest, curve, err = UnmarshalPrivateKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalPublicKeyFromPEM(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
pubKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A short key
|
|
||||||
-----BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner
|
|
||||||
-----BEGIN NOT A NEBULA PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NOT A NEBULA PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA ED25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-END NEBULA ED25519 PUBLIC KEY-----`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(pubKey, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUnmarshalX25519PublicKey(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
pubKey := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA X25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA X25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
pubP256Key := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
oldPubP256Key := []byte(`# A good key
|
|
||||||
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NEBULA ECDSA P256 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
shortKey := []byte(`# A short key
|
|
||||||
-----BEGIN NEBULA X25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==
|
|
||||||
-----END NEBULA X25519 PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidBanner := []byte(`# Invalid banner
|
|
||||||
-----BEGIN NOT A NEBULA PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-----END NOT A NEBULA PUBLIC KEY-----
|
|
||||||
`)
|
|
||||||
invalidPem := []byte(`# Not a valid PEM format
|
|
||||||
-BEGIN NEBULA X25519 PUBLIC KEY-----
|
|
||||||
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
|
|
||||||
-END NEBULA X25519 PUBLIC KEY-----`)
|
|
||||||
|
|
||||||
keyBundle := appendByteSlices(pubKey, pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
|
|
||||||
assert.Len(t, k, 32)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_CURVE25519, curve)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Len(t, k, 65)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(oldPubP256Key, shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
|
|
||||||
// Success test case
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Len(t, k, 65)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
|
|
||||||
assert.Equal(t, Curve_P256, curve)
|
|
||||||
|
|
||||||
// Fail due to short key
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
|
|
||||||
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
|
|
||||||
|
|
||||||
// Fail due to invalid banner
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
require.EqualError(t, err, "bytes did not contain a proper public key banner")
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
|
|
||||||
// Fail due to invalid PEM format, because
|
|
||||||
// it's missing the requisite pre-encapsulation boundary.
|
|
||||||
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
|
|
||||||
assert.Nil(t, k)
|
|
||||||
assert.Equal(t, rest, invalidPem)
|
|
||||||
require.EqualError(t, err, "input did not contain a valid PEM encoded block")
|
|
||||||
}
|
|
||||||
-170
@@ -1,170 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"crypto/sha256"
|
|
||||||
"fmt"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert/p256"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TBSCertificate represents a certificate intended to be signed.
|
|
||||||
// It is invalid to use this structure as a Certificate.
|
|
||||||
type TBSCertificate struct {
|
|
||||||
Version Version
|
|
||||||
Name string
|
|
||||||
Networks []netip.Prefix
|
|
||||||
UnsafeNetworks []netip.Prefix
|
|
||||||
Groups []string
|
|
||||||
IsCA bool
|
|
||||||
NotBefore time.Time
|
|
||||||
NotAfter time.Time
|
|
||||||
PublicKey []byte
|
|
||||||
Curve Curve
|
|
||||||
issuer string
|
|
||||||
}
|
|
||||||
|
|
||||||
type beingSignedCertificate interface {
|
|
||||||
// fromTBSCertificate copies the values from the TBSCertificate to this versions internal representation
|
|
||||||
// Implementations must validate the resulting certificate contains valid information
|
|
||||||
fromTBSCertificate(*TBSCertificate) error
|
|
||||||
|
|
||||||
// marshalForSigning returns the bytes that should be signed
|
|
||||||
marshalForSigning() ([]byte, error)
|
|
||||||
|
|
||||||
// setSignature sets the signature for the certificate that has just been signed. The signature must not be blank.
|
|
||||||
setSignature([]byte) error
|
|
||||||
}
|
|
||||||
|
|
||||||
type SignerLambda func(certBytes []byte) ([]byte, error)
|
|
||||||
|
|
||||||
// Sign will create a sealed certificate using details provided by the TBSCertificate as long as those
|
|
||||||
// details do not violate constraints of the signing certificate.
|
|
||||||
// If the TBSCertificate is a CA then signer must be nil.
|
|
||||||
func (t *TBSCertificate) Sign(signer Certificate, curve Curve, key []byte) (Certificate, error) {
|
|
||||||
switch t.Curve {
|
|
||||||
case Curve_CURVE25519:
|
|
||||||
pk := ed25519.PrivateKey(key)
|
|
||||||
sp := func(certBytes []byte) ([]byte, error) {
|
|
||||||
sig := ed25519.Sign(pk, certBytes)
|
|
||||||
return sig, nil
|
|
||||||
}
|
|
||||||
return t.SignWith(signer, curve, sp)
|
|
||||||
case Curve_P256:
|
|
||||||
pk, err := ecdsa.ParseRawPrivateKey(elliptic.P256(), key)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
sp := func(certBytes []byte) ([]byte, error) {
|
|
||||||
// We need to hash first for ECDSA
|
|
||||||
// - https://pkg.go.dev/crypto/ecdsa#SignASN1
|
|
||||||
hashed := sha256.Sum256(certBytes)
|
|
||||||
return ecdsa.SignASN1(rand.Reader, pk, hashed[:])
|
|
||||||
}
|
|
||||||
return t.SignWith(signer, curve, sp)
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("invalid curve: %s", t.Curve)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// SignWith does the same thing as sign, but uses the function in `sp` to calculate the signature.
|
|
||||||
// You should only use SignWith if you do not have direct access to your private key.
|
|
||||||
func (t *TBSCertificate) SignWith(signer Certificate, curve Curve, sp SignerLambda) (Certificate, error) {
|
|
||||||
if curve != t.Curve {
|
|
||||||
return nil, fmt.Errorf("curve in cert and private key supplied don't match")
|
|
||||||
}
|
|
||||||
|
|
||||||
if signer != nil {
|
|
||||||
if t.IsCA {
|
|
||||||
return nil, fmt.Errorf("can not sign a CA certificate with another")
|
|
||||||
}
|
|
||||||
|
|
||||||
err := checkCAConstraints(signer, t.NotBefore, t.NotAfter, t.Groups, t.Networks, t.UnsafeNetworks)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
issuer, err := signer.Fingerprint()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("error computing issuer: %v", err)
|
|
||||||
}
|
|
||||||
t.issuer = issuer
|
|
||||||
} else {
|
|
||||||
if !t.IsCA {
|
|
||||||
return nil, fmt.Errorf("self signed certificates must have IsCA set to true")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var c beingSignedCertificate
|
|
||||||
switch t.Version {
|
|
||||||
case Version1:
|
|
||||||
c = &certificateV1{}
|
|
||||||
err := c.fromTBSCertificate(t)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
case Version2:
|
|
||||||
c = &certificateV2{}
|
|
||||||
err := c.fromTBSCertificate(t)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("unknown cert version %d", t.Version)
|
|
||||||
}
|
|
||||||
|
|
||||||
certBytes, err := c.marshalForSigning()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
sig, err := sp(certBytes)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if curve == Curve_P256 {
|
|
||||||
sig, err = p256.Normalize(sig)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = c.setSignature(sig)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
sc, ok := c.(Certificate)
|
|
||||||
if !ok {
|
|
||||||
return nil, fmt.Errorf("invalid certificate")
|
|
||||||
}
|
|
||||||
|
|
||||||
return sc, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func comparePrefix(a, b netip.Prefix) int {
|
|
||||||
addr := a.Addr().Compare(b.Addr())
|
|
||||||
if addr == 0 {
|
|
||||||
return a.Bits() - b.Bits()
|
|
||||||
}
|
|
||||||
return addr
|
|
||||||
}
|
|
||||||
|
|
||||||
// findDuplicatePrefix returns an error if there is a duplicate prefix in the pre-sorted input slice sortedPrefixes
|
|
||||||
func findDuplicatePrefix(sortedPrefixes []netip.Prefix) error {
|
|
||||||
if len(sortedPrefixes) < 2 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
for i := 1; i < len(sortedPrefixes); i++ {
|
|
||||||
if comparePrefix(sortedPrefixes[i], sortedPrefixes[i-1]) == 0 {
|
|
||||||
return NewErrInvalidCertificateProperties("duplicate network detected: %v", sortedPrefixes[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,137 +0,0 @@
|
|||||||
package cert
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/ed25519"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert/p256"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestCertificateV1_Sign(t *testing.T) {
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
tbs := TBSCertificate{
|
|
||||||
Version: Version1,
|
|
||||||
Name: "testing",
|
|
||||||
Networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
UnsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/24"),
|
|
||||||
},
|
|
||||||
Groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
NotBefore: before,
|
|
||||||
NotAfter: after,
|
|
||||||
PublicKey: pubKey,
|
|
||||||
IsCA: false,
|
|
||||||
}
|
|
||||||
|
|
||||||
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
|
||||||
c, err := tbs.Sign(&certificateV1{details: detailsV1{notBefore: before, notAfter: after}}, Curve_CURVE25519, priv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, c)
|
|
||||||
assert.True(t, c.CheckSignature(pub))
|
|
||||||
|
|
||||||
b, err := c.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
uc, err := unmarshalCertificateV1(b, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, uc)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificateV1_SignP256(t *testing.T) {
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("01234567890abcedfghij1234567890ab1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
tbs := TBSCertificate{
|
|
||||||
Version: Version1,
|
|
||||||
Name: "testing",
|
|
||||||
Networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
UnsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
Groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
NotBefore: before,
|
|
||||||
NotAfter: after,
|
|
||||||
PublicKey: pubKey,
|
|
||||||
IsCA: false,
|
|
||||||
Curve: Curve_P256,
|
|
||||||
}
|
|
||||||
|
|
||||||
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y)
|
|
||||||
rawPriv := priv.D.FillBytes(make([]byte, 32))
|
|
||||||
|
|
||||||
c, err := tbs.Sign(&certificateV1{details: detailsV1{notBefore: before, notAfter: after}}, Curve_P256, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, c)
|
|
||||||
assert.True(t, c.CheckSignature(pub))
|
|
||||||
|
|
||||||
b, err := c.Marshal()
|
|
||||||
require.NoError(t, err)
|
|
||||||
uc, err := unmarshalCertificateV1(b, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, uc)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCertificate_SignP256_AlwaysNormalized(t *testing.T) {
|
|
||||||
before := time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
after := time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
pubKey := []byte("01234567890abcedfghij1234567890ab1234567890abcedfghij1234567890ab")
|
|
||||||
|
|
||||||
tbs := TBSCertificate{
|
|
||||||
Version: Version1,
|
|
||||||
Name: "testing",
|
|
||||||
Networks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("10.1.1.1/24"),
|
|
||||||
mustParsePrefixUnmapped("10.1.1.2/16"),
|
|
||||||
},
|
|
||||||
UnsafeNetworks: []netip.Prefix{
|
|
||||||
mustParsePrefixUnmapped("9.1.1.2/24"),
|
|
||||||
mustParsePrefixUnmapped("9.1.1.3/16"),
|
|
||||||
},
|
|
||||||
Groups: []string{"test-group1", "test-group2", "test-group3"},
|
|
||||||
NotBefore: before,
|
|
||||||
NotAfter: after,
|
|
||||||
PublicKey: pubKey,
|
|
||||||
IsCA: true,
|
|
||||||
Curve: Curve_P256,
|
|
||||||
}
|
|
||||||
|
|
||||||
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y)
|
|
||||||
rawPriv := priv.D.FillBytes(make([]byte, 32))
|
|
||||||
|
|
||||||
for i := 0; i < 1000; i++ {
|
|
||||||
if i&1 == 1 {
|
|
||||||
tbs.Version = Version1
|
|
||||||
} else {
|
|
||||||
tbs.Version = Version2
|
|
||||||
}
|
|
||||||
c, err := tbs.Sign(nil, Curve_P256, rawPriv)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.NotNil(t, c)
|
|
||||||
assert.True(t, c.CheckSignature(pub))
|
|
||||||
normie, err := p256.IsNormalized(c.Signature())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.True(t, normie)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,217 +0,0 @@
|
|||||||
package cert_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/ecdh"
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"io"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"golang.org/x/crypto/curve25519"
|
|
||||||
"golang.org/x/crypto/ed25519"
|
|
||||||
)
|
|
||||||
|
|
||||||
// NewTestCaCert will create a new ca certificate
|
|
||||||
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
|
||||||
var err error
|
|
||||||
var pub, priv []byte
|
|
||||||
|
|
||||||
switch curve {
|
|
||||||
case cert.Curve_CURVE25519:
|
|
||||||
pub, priv, err = ed25519.GenerateKey(rand.Reader)
|
|
||||||
case cert.Curve_P256:
|
|
||||||
privk, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub = elliptic.Marshal(elliptic.P256(), privk.PublicKey.X, privk.PublicKey.Y)
|
|
||||||
priv = privk.D.FillBytes(make([]byte, 32))
|
|
||||||
default:
|
|
||||||
// There is no default to allow the underlying lib to respond with an error
|
|
||||||
}
|
|
||||||
|
|
||||||
if before.IsZero() {
|
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
}
|
|
||||||
if after.IsZero() {
|
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
|
||||||
Curve: curve,
|
|
||||||
Version: version,
|
|
||||||
Name: "test ca",
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pub,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
IsCA: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := t.Sign(nil, curve, priv)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pem, err := c.MarshalPEM()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, pub, priv, pem
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewTestCert will generate a signed certificate with the provided details.
|
|
||||||
// Expiry times are defaulted if you do not pass them in
|
|
||||||
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
|
|
||||||
if before.IsZero() {
|
|
||||||
before = time.Now().Add(time.Second * -60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
if after.IsZero() {
|
|
||||||
after = time.Now().Add(time.Second * 60).Round(time.Second)
|
|
||||||
}
|
|
||||||
|
|
||||||
var pub, priv []byte
|
|
||||||
switch curve {
|
|
||||||
case cert.Curve_CURVE25519:
|
|
||||||
pub, priv = X25519Keypair()
|
|
||||||
case cert.Curve_P256:
|
|
||||||
pub, priv = P256Keypair()
|
|
||||||
default:
|
|
||||||
panic("unknown curve")
|
|
||||||
}
|
|
||||||
|
|
||||||
nc := &cert.TBSCertificate{
|
|
||||||
Version: v,
|
|
||||||
Curve: curve,
|
|
||||||
Name: name,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pub,
|
|
||||||
IsCA: false,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := nc.Sign(ca, ca.Curve(), key)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pem, err := c.MarshalPEM()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, pub, cert.MarshalPrivateKeyToPEM(curve, priv), pem
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewTestCertDifferentVersion(c cert.Certificate, v cert.Version, ca cert.Certificate, key []byte) (cert.Certificate, []byte) {
|
|
||||||
nc := &cert.TBSCertificate{
|
|
||||||
Version: v,
|
|
||||||
Curve: c.Curve(),
|
|
||||||
Name: c.Name(),
|
|
||||||
Networks: c.Networks(),
|
|
||||||
UnsafeNetworks: c.UnsafeNetworks(),
|
|
||||||
Groups: c.Groups(),
|
|
||||||
NotBefore: time.Unix(c.NotBefore().Unix(), 0),
|
|
||||||
NotAfter: time.Unix(c.NotAfter().Unix(), 0),
|
|
||||||
PublicKey: c.PublicKey(),
|
|
||||||
IsCA: false,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := nc.Sign(ca, ca.Curve(), key)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pem, err := c.MarshalPEM()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, pem
|
|
||||||
}
|
|
||||||
|
|
||||||
func X25519Keypair() ([]byte, []byte) {
|
|
||||||
privkey := make([]byte, 32)
|
|
||||||
if _, err := io.ReadFull(rand.Reader, privkey); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
pubkey, err := curve25519.X25519(privkey, curve25519.Basepoint)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return pubkey, privkey
|
|
||||||
}
|
|
||||||
|
|
||||||
func P256Keypair() ([]byte, []byte) {
|
|
||||||
privkey, err := ecdh.P256().GenerateKey(rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
pubkey := privkey.PublicKey()
|
|
||||||
return pubkey.Bytes(), privkey.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
// DummyCert is a minimal cert.Certificate implementation for testing error paths.
|
|
||||||
type DummyCert struct {
|
|
||||||
Version_ cert.Version
|
|
||||||
Curve_ cert.Curve
|
|
||||||
Groups_ []string
|
|
||||||
IsCA_ bool
|
|
||||||
Issuer_ string
|
|
||||||
Name_ string
|
|
||||||
Networks_ []netip.Prefix
|
|
||||||
NotAfter_ time.Time
|
|
||||||
NotBefore_ time.Time
|
|
||||||
PublicKey_ []byte
|
|
||||||
Signature_ []byte
|
|
||||||
UnsafeNetworks_ []netip.Prefix
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *DummyCert) Version() cert.Version { return d.Version_ }
|
|
||||||
func (d *DummyCert) Curve() cert.Curve { return d.Curve_ }
|
|
||||||
func (d *DummyCert) Groups() []string { return d.Groups_ }
|
|
||||||
func (d *DummyCert) IsCA() bool { return d.IsCA_ }
|
|
||||||
func (d *DummyCert) Issuer() string { return d.Issuer_ }
|
|
||||||
func (d *DummyCert) Name() string { return d.Name_ }
|
|
||||||
func (d *DummyCert) Networks() []netip.Prefix { return d.Networks_ }
|
|
||||||
func (d *DummyCert) NotAfter() time.Time { return d.NotAfter_ }
|
|
||||||
func (d *DummyCert) NotBefore() time.Time { return d.NotBefore_ }
|
|
||||||
func (d *DummyCert) PublicKey() []byte { return d.PublicKey_ }
|
|
||||||
func (d *DummyCert) Signature() []byte { return d.Signature_ }
|
|
||||||
func (d *DummyCert) UnsafeNetworks() []netip.Prefix { return d.UnsafeNetworks_ }
|
|
||||||
func (d *DummyCert) Fingerprint() (string, error) { return "", nil }
|
|
||||||
func (d *DummyCert) CheckSignature(key []byte) bool { return false }
|
|
||||||
func (d *DummyCert) MarshalForHandshakes() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) MarshalPEM() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) MarshalJSON() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) Marshal() ([]byte, error) { return nil, nil }
|
|
||||||
func (d *DummyCert) String() string { return "dummy" }
|
|
||||||
func (d *DummyCert) Copy() cert.Certificate { return d }
|
|
||||||
func (d *DummyCert) VerifyPrivateKey(c cert.Curve, k []byte) error { return nil }
|
|
||||||
func (d *DummyCert) Expired(time.Time) bool { return false }
|
|
||||||
func (d *DummyCert) MarshalPublicKeyPEM() []byte { return nil }
|
|
||||||
func (d *DummyCert) PublicKeyPEM() []byte { return nil }
|
|
||||||
|
|
||||||
// NewTestCAPool creates a CAPool from the given CA certificates, panicking on error.
|
|
||||||
func NewTestCAPool(cas ...cert.Certificate) *cert.CAPool {
|
|
||||||
pool := cert.NewCAPool()
|
|
||||||
for _, ca := range cas {
|
|
||||||
if err := pool.AddCA(ca); err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pool
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
package cidr
|
||||||
|
|
||||||
|
import "net"
|
||||||
|
|
||||||
|
// Parse is a convenience function that returns only the IPNet
|
||||||
|
// This function ignores errors since it is primarily a test helper, the result could be nil
|
||||||
|
func Parse(s string) *net.IPNet {
|
||||||
|
_, c, _ := net.ParseCIDR(s)
|
||||||
|
return c
|
||||||
|
}
|
||||||
+167
@@ -0,0 +1,167 @@
|
|||||||
|
package cidr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Node struct {
|
||||||
|
left *Node
|
||||||
|
right *Node
|
||||||
|
parent *Node
|
||||||
|
value interface{}
|
||||||
|
}
|
||||||
|
|
||||||
|
type entry struct {
|
||||||
|
CIDR *net.IPNet
|
||||||
|
Value *interface{}
|
||||||
|
}
|
||||||
|
|
||||||
|
type Tree4 struct {
|
||||||
|
root *Node
|
||||||
|
list []entry
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
startbit = iputil.VpnIp(0x80000000)
|
||||||
|
)
|
||||||
|
|
||||||
|
func NewTree4() *Tree4 {
|
||||||
|
tree := new(Tree4)
|
||||||
|
tree.root = &Node{}
|
||||||
|
tree.list = []entry{}
|
||||||
|
return tree
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tree *Tree4) AddCIDR(cidr *net.IPNet, val interface{}) {
|
||||||
|
bit := startbit
|
||||||
|
node := tree.root
|
||||||
|
next := tree.root
|
||||||
|
|
||||||
|
ip := iputil.Ip2VpnIp(cidr.IP)
|
||||||
|
mask := iputil.Ip2VpnIp(cidr.Mask)
|
||||||
|
|
||||||
|
// Find our last ancestor in the tree
|
||||||
|
for bit&mask != 0 {
|
||||||
|
if ip&bit != 0 {
|
||||||
|
next = node.right
|
||||||
|
} else {
|
||||||
|
next = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
if next == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
bit = bit >> 1
|
||||||
|
node = next
|
||||||
|
}
|
||||||
|
|
||||||
|
// We already have this range so update the value
|
||||||
|
if next != nil {
|
||||||
|
addCIDR := cidr.String()
|
||||||
|
for i, v := range tree.list {
|
||||||
|
if addCIDR == v.CIDR.String() {
|
||||||
|
tree.list = append(tree.list[:i], tree.list[i+1:]...)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
tree.list = append(tree.list, entry{CIDR: cidr, Value: &val})
|
||||||
|
node.value = val
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build up the rest of the tree we don't already have
|
||||||
|
for bit&mask != 0 {
|
||||||
|
next = &Node{}
|
||||||
|
next.parent = node
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node.right = next
|
||||||
|
} else {
|
||||||
|
node.left = next
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
node = next
|
||||||
|
}
|
||||||
|
|
||||||
|
// Final node marks our cidr, set the value
|
||||||
|
node.value = val
|
||||||
|
tree.list = append(tree.list, entry{CIDR: cidr, Value: &val})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Contains finds the first match, which may be the least specific
|
||||||
|
func (tree *Tree4) Contains(ip iputil.VpnIp) (value interface{}) {
|
||||||
|
bit := startbit
|
||||||
|
node := tree.root
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
if node.value != nil {
|
||||||
|
return node.value
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// MostSpecificContains finds the most specific match
|
||||||
|
func (tree *Tree4) MostSpecificContains(ip iputil.VpnIp) (value interface{}) {
|
||||||
|
bit := startbit
|
||||||
|
node := tree.root
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
if node.value != nil {
|
||||||
|
value = node.value
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match finds the most specific match
|
||||||
|
func (tree *Tree4) Match(ip iputil.VpnIp) (value interface{}) {
|
||||||
|
bit := startbit
|
||||||
|
node := tree.root
|
||||||
|
lastNode := node
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
lastNode = node
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
}
|
||||||
|
|
||||||
|
if bit == 0 && lastNode != nil {
|
||||||
|
value = lastNode.value
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// List will return all CIDRs and their current values. Do not modify the contents!
|
||||||
|
func (tree *Tree4) List() []entry {
|
||||||
|
return tree.list
|
||||||
|
}
|
||||||
@@ -0,0 +1,167 @@
|
|||||||
|
package cidr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCIDRTree_List(t *testing.T) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/16"), "1")
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/8"), "2")
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/16"), "3")
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/16"), "4")
|
||||||
|
list := tree.List()
|
||||||
|
assert.Len(t, list, 2)
|
||||||
|
assert.Equal(t, "1.0.0.0/8", list[0].CIDR.String())
|
||||||
|
assert.Equal(t, "2", *list[0].Value)
|
||||||
|
assert.Equal(t, "1.0.0.0/16", list[1].CIDR.String())
|
||||||
|
assert.Equal(t, "4", *list[1].Value)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDRTree_Contains(t *testing.T) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/8"), "1")
|
||||||
|
tree.AddCIDR(Parse("2.1.0.0/16"), "2")
|
||||||
|
tree.AddCIDR(Parse("3.1.1.0/24"), "3")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.0/24"), "4a")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/32"), "4b")
|
||||||
|
tree.AddCIDR(Parse("4.1.2.1/32"), "4c")
|
||||||
|
tree.AddCIDR(Parse("254.0.0.0/4"), "5")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
Result interface{}
|
||||||
|
IP string
|
||||||
|
}{
|
||||||
|
{"1", "1.0.0.0"},
|
||||||
|
{"1", "1.255.255.255"},
|
||||||
|
{"2", "2.1.0.0"},
|
||||||
|
{"2", "2.1.255.255"},
|
||||||
|
{"3", "3.1.1.0"},
|
||||||
|
{"3", "3.1.1.255"},
|
||||||
|
{"4a", "4.1.1.255"},
|
||||||
|
{"4a", "4.1.1.1"},
|
||||||
|
{"5", "240.0.0.0"},
|
||||||
|
{"5", "255.255.255.255"},
|
||||||
|
{nil, "239.0.0.0"},
|
||||||
|
{nil, "4.1.2.2"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
assert.Equal(t, tt.Result, tree.Contains(iputil.Ip2VpnIp(net.ParseIP(tt.IP))))
|
||||||
|
}
|
||||||
|
|
||||||
|
tree = NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.1.1.1/0"), "cool")
|
||||||
|
assert.Equal(t, "cool", tree.Contains(iputil.Ip2VpnIp(net.ParseIP("0.0.0.0"))))
|
||||||
|
assert.Equal(t, "cool", tree.Contains(iputil.Ip2VpnIp(net.ParseIP("255.255.255.255"))))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDRTree_MostSpecificContains(t *testing.T) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/8"), "1")
|
||||||
|
tree.AddCIDR(Parse("2.1.0.0/16"), "2")
|
||||||
|
tree.AddCIDR(Parse("3.1.1.0/24"), "3")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.0/24"), "4a")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.0/30"), "4b")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/32"), "4c")
|
||||||
|
tree.AddCIDR(Parse("254.0.0.0/4"), "5")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
Result interface{}
|
||||||
|
IP string
|
||||||
|
}{
|
||||||
|
{"1", "1.0.0.0"},
|
||||||
|
{"1", "1.255.255.255"},
|
||||||
|
{"2", "2.1.0.0"},
|
||||||
|
{"2", "2.1.255.255"},
|
||||||
|
{"3", "3.1.1.0"},
|
||||||
|
{"3", "3.1.1.255"},
|
||||||
|
{"4a", "4.1.1.255"},
|
||||||
|
{"4b", "4.1.1.2"},
|
||||||
|
{"4c", "4.1.1.1"},
|
||||||
|
{"5", "240.0.0.0"},
|
||||||
|
{"5", "255.255.255.255"},
|
||||||
|
{nil, "239.0.0.0"},
|
||||||
|
{nil, "4.1.2.2"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
assert.Equal(t, tt.Result, tree.MostSpecificContains(iputil.Ip2VpnIp(net.ParseIP(tt.IP))))
|
||||||
|
}
|
||||||
|
|
||||||
|
tree = NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.1.1.1/0"), "cool")
|
||||||
|
assert.Equal(t, "cool", tree.MostSpecificContains(iputil.Ip2VpnIp(net.ParseIP("0.0.0.0"))))
|
||||||
|
assert.Equal(t, "cool", tree.MostSpecificContains(iputil.Ip2VpnIp(net.ParseIP("255.255.255.255"))))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDRTree_Match(t *testing.T) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("4.1.1.0/32"), "1a")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/32"), "1b")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
Result interface{}
|
||||||
|
IP string
|
||||||
|
}{
|
||||||
|
{"1a", "4.1.1.0"},
|
||||||
|
{"1b", "4.1.1.1"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
assert.Equal(t, tt.Result, tree.Match(iputil.Ip2VpnIp(net.ParseIP(tt.IP))))
|
||||||
|
}
|
||||||
|
|
||||||
|
tree = NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.1.1.1/0"), "cool")
|
||||||
|
assert.Equal(t, "cool", tree.Contains(iputil.Ip2VpnIp(net.ParseIP("0.0.0.0"))))
|
||||||
|
assert.Equal(t, "cool", tree.Contains(iputil.Ip2VpnIp(net.ParseIP("255.255.255.255"))))
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkCIDRTree_Contains(b *testing.B) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.1.0.0/16"), "1")
|
||||||
|
tree.AddCIDR(Parse("1.2.1.1/32"), "1")
|
||||||
|
tree.AddCIDR(Parse("192.2.1.1/32"), "1")
|
||||||
|
tree.AddCIDR(Parse("172.2.1.1/32"), "1")
|
||||||
|
|
||||||
|
ip := iputil.Ip2VpnIp(net.ParseIP("1.2.1.1"))
|
||||||
|
b.Run("found", func(b *testing.B) {
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
tree.Contains(ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
ip = iputil.Ip2VpnIp(net.ParseIP("1.2.1.255"))
|
||||||
|
b.Run("not found", func(b *testing.B) {
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
tree.Contains(ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkCIDRTree_Match(b *testing.B) {
|
||||||
|
tree := NewTree4()
|
||||||
|
tree.AddCIDR(Parse("1.1.0.0/16"), "1")
|
||||||
|
tree.AddCIDR(Parse("1.2.1.1/32"), "1")
|
||||||
|
tree.AddCIDR(Parse("192.2.1.1/32"), "1")
|
||||||
|
tree.AddCIDR(Parse("172.2.1.1/32"), "1")
|
||||||
|
|
||||||
|
ip := iputil.Ip2VpnIp(net.ParseIP("1.2.1.1"))
|
||||||
|
b.Run("found", func(b *testing.B) {
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
tree.Match(ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
ip = iputil.Ip2VpnIp(net.ParseIP("1.2.1.255"))
|
||||||
|
b.Run("not found", func(b *testing.B) {
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
tree.Match(ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+185
@@ -0,0 +1,185 @@
|
|||||||
|
package cidr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
)
|
||||||
|
|
||||||
|
const startbit6 = uint64(1 << 63)
|
||||||
|
|
||||||
|
type Tree6 struct {
|
||||||
|
root4 *Node
|
||||||
|
root6 *Node
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTree6() *Tree6 {
|
||||||
|
tree := new(Tree6)
|
||||||
|
tree.root4 = &Node{}
|
||||||
|
tree.root6 = &Node{}
|
||||||
|
return tree
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tree *Tree6) AddCIDR(cidr *net.IPNet, val interface{}) {
|
||||||
|
var node, next *Node
|
||||||
|
|
||||||
|
cidrIP, ipv4 := isIPV4(cidr.IP)
|
||||||
|
if ipv4 {
|
||||||
|
node = tree.root4
|
||||||
|
next = tree.root4
|
||||||
|
|
||||||
|
} else {
|
||||||
|
node = tree.root6
|
||||||
|
next = tree.root6
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := 0; i < len(cidrIP); i += 4 {
|
||||||
|
ip := iputil.Ip2VpnIp(cidrIP[i : i+4])
|
||||||
|
mask := iputil.Ip2VpnIp(cidr.Mask[i : i+4])
|
||||||
|
bit := startbit
|
||||||
|
|
||||||
|
// Find our last ancestor in the tree
|
||||||
|
for bit&mask != 0 {
|
||||||
|
if ip&bit != 0 {
|
||||||
|
next = node.right
|
||||||
|
} else {
|
||||||
|
next = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
if next == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
bit = bit >> 1
|
||||||
|
node = next
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build up the rest of the tree we don't already have
|
||||||
|
for bit&mask != 0 {
|
||||||
|
next = &Node{}
|
||||||
|
next.parent = node
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node.right = next
|
||||||
|
} else {
|
||||||
|
node.left = next
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
node = next
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Final node marks our cidr, set the value
|
||||||
|
node.value = val
|
||||||
|
}
|
||||||
|
|
||||||
|
// Finds the most specific match
|
||||||
|
func (tree *Tree6) MostSpecificContains(ip net.IP) (value interface{}) {
|
||||||
|
var node *Node
|
||||||
|
|
||||||
|
wholeIP, ipv4 := isIPV4(ip)
|
||||||
|
if ipv4 {
|
||||||
|
node = tree.root4
|
||||||
|
} else {
|
||||||
|
node = tree.root6
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := 0; i < len(wholeIP); i += 4 {
|
||||||
|
ip := iputil.Ip2VpnIp(wholeIP[i : i+4])
|
||||||
|
bit := startbit
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
if node.value != nil {
|
||||||
|
value = node.value
|
||||||
|
}
|
||||||
|
|
||||||
|
if bit == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tree *Tree6) MostSpecificContainsIpV4(ip iputil.VpnIp) (value interface{}) {
|
||||||
|
bit := startbit
|
||||||
|
node := tree.root4
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
if node.value != nil {
|
||||||
|
value = node.value
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tree *Tree6) MostSpecificContainsIpV6(hi, lo uint64) (value interface{}) {
|
||||||
|
ip := hi
|
||||||
|
node := tree.root6
|
||||||
|
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
bit := startbit6
|
||||||
|
|
||||||
|
for node != nil {
|
||||||
|
if node.value != nil {
|
||||||
|
value = node.value
|
||||||
|
}
|
||||||
|
|
||||||
|
if bit == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
if ip&bit != 0 {
|
||||||
|
node = node.right
|
||||||
|
} else {
|
||||||
|
node = node.left
|
||||||
|
}
|
||||||
|
|
||||||
|
bit >>= 1
|
||||||
|
}
|
||||||
|
|
||||||
|
ip = lo
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIPV4(ip net.IP) (net.IP, bool) {
|
||||||
|
if len(ip) == net.IPv4len {
|
||||||
|
return ip, true
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(ip) == net.IPv6len && isZeros(ip[0:10]) && ip[10] == 0xff && ip[11] == 0xff {
|
||||||
|
return ip[12:16], true
|
||||||
|
}
|
||||||
|
|
||||||
|
return ip, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func isZeros(p net.IP) bool {
|
||||||
|
for i := 0; i < len(p); i++ {
|
||||||
|
if p[i] != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,81 @@
|
|||||||
|
package cidr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCIDR6Tree_MostSpecificContains(t *testing.T) {
|
||||||
|
tree := NewTree6()
|
||||||
|
tree.AddCIDR(Parse("1.0.0.0/8"), "1")
|
||||||
|
tree.AddCIDR(Parse("2.1.0.0/16"), "2")
|
||||||
|
tree.AddCIDR(Parse("3.1.1.0/24"), "3")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/24"), "4a")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/30"), "4b")
|
||||||
|
tree.AddCIDR(Parse("4.1.1.1/32"), "4c")
|
||||||
|
tree.AddCIDR(Parse("254.0.0.0/4"), "5")
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/64"), "6a")
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/80"), "6b")
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/96"), "6c")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
Result interface{}
|
||||||
|
IP string
|
||||||
|
}{
|
||||||
|
{"1", "1.0.0.0"},
|
||||||
|
{"1", "1.255.255.255"},
|
||||||
|
{"2", "2.1.0.0"},
|
||||||
|
{"2", "2.1.255.255"},
|
||||||
|
{"3", "3.1.1.0"},
|
||||||
|
{"3", "3.1.1.255"},
|
||||||
|
{"4a", "4.1.1.255"},
|
||||||
|
{"4b", "4.1.1.2"},
|
||||||
|
{"4c", "4.1.1.1"},
|
||||||
|
{"5", "240.0.0.0"},
|
||||||
|
{"5", "255.255.255.255"},
|
||||||
|
{"6a", "1:2:0:4:1:1:1:1"},
|
||||||
|
{"6b", "1:2:0:4:5:1:1:1"},
|
||||||
|
{"6c", "1:2:0:4:5:0:0:0"},
|
||||||
|
{nil, "239.0.0.0"},
|
||||||
|
{nil, "4.1.2.2"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
assert.Equal(t, tt.Result, tree.MostSpecificContains(net.ParseIP(tt.IP)))
|
||||||
|
}
|
||||||
|
|
||||||
|
tree = NewTree6()
|
||||||
|
tree.AddCIDR(Parse("1.1.1.1/0"), "cool")
|
||||||
|
tree.AddCIDR(Parse("::/0"), "cool6")
|
||||||
|
assert.Equal(t, "cool", tree.MostSpecificContains(net.ParseIP("0.0.0.0")))
|
||||||
|
assert.Equal(t, "cool", tree.MostSpecificContains(net.ParseIP("255.255.255.255")))
|
||||||
|
assert.Equal(t, "cool6", tree.MostSpecificContains(net.ParseIP("::")))
|
||||||
|
assert.Equal(t, "cool6", tree.MostSpecificContains(net.ParseIP("1:2:3:4:5:6:7:8")))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDR6Tree_MostSpecificContainsIpV6(t *testing.T) {
|
||||||
|
tree := NewTree6()
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/64"), "6a")
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/80"), "6b")
|
||||||
|
tree.AddCIDR(Parse("1:2:0:4:5:0:0:0/96"), "6c")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
Result interface{}
|
||||||
|
IP string
|
||||||
|
}{
|
||||||
|
{"6a", "1:2:0:4:1:1:1:1"},
|
||||||
|
{"6b", "1:2:0:4:5:1:1:1"},
|
||||||
|
{"6c", "1:2:0:4:5:0:0:0"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
ip := net.ParseIP(tt.IP)
|
||||||
|
hi := binary.BigEndian.Uint64(ip[:8])
|
||||||
|
lo := binary.BigEndian.Uint64(ip[8:])
|
||||||
|
|
||||||
|
assert.Equal(t, tt.Result, tree.MostSpecificContainsIpV6(hi, lo))
|
||||||
|
}
|
||||||
|
}
|
||||||
+93
-165
@@ -7,15 +7,15 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"io/ioutil"
|
||||||
"math"
|
"math"
|
||||||
"net/netip"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/skip2/go-qrcode"
|
"github.com/skip2/go-qrcode"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/pkclient"
|
|
||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -27,43 +27,32 @@ type caFlags struct {
|
|||||||
outCertPath *string
|
outCertPath *string
|
||||||
outQRPath *string
|
outQRPath *string
|
||||||
groups *string
|
groups *string
|
||||||
networks *string
|
ips *string
|
||||||
unsafeNetworks *string
|
subnets *string
|
||||||
argonMemory *uint
|
argonMemory *uint
|
||||||
argonIterations *uint
|
argonIterations *uint
|
||||||
argonParallelism *uint
|
argonParallelism *uint
|
||||||
encryption *bool
|
encryption *bool
|
||||||
version *uint
|
|
||||||
|
|
||||||
curve *string
|
curve *string
|
||||||
p11url *string
|
|
||||||
|
|
||||||
// Deprecated options
|
|
||||||
ips *string
|
|
||||||
subnets *string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCaFlags() *caFlags {
|
func newCaFlags() *caFlags {
|
||||||
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
cf := caFlags{set: flag.NewFlagSet("ca", flag.ContinueOnError)}
|
||||||
cf.set.Usage = func() {}
|
cf.set.Usage = func() {}
|
||||||
cf.name = cf.set.String("name", "", "Required: name of the certificate authority")
|
cf.name = cf.set.String("name", "", "Required: name of the certificate authority")
|
||||||
cf.version = cf.set.Uint("version", uint(cert.Version2), "Optional: version of the certificate format to use")
|
|
||||||
cf.duration = cf.set.Duration("duration", time.Duration(time.Hour*8760), "Optional: amount of time the certificate should be valid for. Valid time units are seconds: \"s\", minutes: \"m\", hours: \"h\"")
|
cf.duration = cf.set.Duration("duration", time.Duration(time.Hour*8760), "Optional: amount of time the certificate should be valid for. Valid time units are seconds: \"s\", minutes: \"m\", hours: \"h\"")
|
||||||
cf.outKeyPath = cf.set.String("out-key", "ca.key", "Optional: path to write the private key to")
|
cf.outKeyPath = cf.set.String("out-key", "ca.key", "Optional: path to write the private key to")
|
||||||
cf.outCertPath = cf.set.String("out-crt", "ca.crt", "Optional: path to write the certificate to")
|
cf.outCertPath = cf.set.String("out-crt", "ca.crt", "Optional: path to write the certificate to")
|
||||||
cf.outQRPath = cf.set.String("out-qr", "", "Optional: output a qr code image (png) of the certificate")
|
cf.outQRPath = cf.set.String("out-qr", "", "Optional: output a qr code image (png) of the certificate")
|
||||||
cf.groups = cf.set.String("groups", "", "Optional: comma separated list of groups. This will limit which groups subordinate certs can use")
|
cf.groups = cf.set.String("groups", "", "Optional: comma separated list of groups. This will limit which groups subordinate certs can use")
|
||||||
cf.networks = cf.set.String("networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in networks")
|
cf.ips = cf.set.String("ips", "", "Optional: comma separated list of ipv4 address and network in CIDR notation. This will limit which ipv4 addresses and networks subordinate certs can use for ip addresses")
|
||||||
cf.unsafeNetworks = cf.set.String("unsafe-networks", "", "Optional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in unsafe networks")
|
cf.subnets = cf.set.String("subnets", "", "Optional: comma separated list of ipv4 address and network in CIDR notation. This will limit which ipv4 addresses and networks subordinate certs can use in subnets")
|
||||||
cf.argonMemory = cf.set.Uint("argon-memory", 2*1024*1024, "Optional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase")
|
cf.argonMemory = cf.set.Uint("argon-memory", 2*1024*1024, "Optional: Argon2 memory parameter (in KiB) used for encrypted private key passphrase")
|
||||||
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
cf.argonParallelism = cf.set.Uint("argon-parallelism", 4, "Optional: Argon2 parallelism parameter used for encrypted private key passphrase")
|
||||||
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
cf.argonIterations = cf.set.Uint("argon-iterations", 1, "Optional: Argon2 iterations parameter used for encrypted private key passphrase")
|
||||||
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
cf.encryption = cf.set.Bool("encrypt", false, "Optional: prompt for passphrase and write out-key in an encrypted format")
|
||||||
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "EdDSA/ECDSA Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
|
||||||
|
|
||||||
cf.ips = cf.set.String("ips", "", "Deprecated, see -networks")
|
|
||||||
cf.subnets = cf.set.String("subnets", "", "Deprecated, see -unsafe-networks")
|
|
||||||
return &cf
|
return &cf
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -88,21 +77,17 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
isP11 := len(*cf.p11url) > 0
|
|
||||||
|
|
||||||
if err := mustFlagString("name", cf.name); err != nil {
|
if err := mustFlagString("name", cf.name); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if !isP11 {
|
if err := mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
return err
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
var kdfParams *cert.Argon2Parameters
|
var kdfParams *cert.Argon2Parameters
|
||||||
if !isP11 && *cf.encryption {
|
if *cf.encryption {
|
||||||
if kdfParams, err = parseArgonParameters(*cf.argonMemory, *cf.argonParallelism, *cf.argonIterations); err != nil {
|
if kdfParams, err = parseArgonParameters(*cf.argonMemory, *cf.argonParallelism, *cf.argonIterations); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -122,190 +107,133 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
version := cert.Version(*cf.version)
|
var ips []*net.IPNet
|
||||||
if version != cert.Version1 && version != cert.Version2 {
|
if *cf.ips != "" {
|
||||||
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
for _, rs := range strings.Split(*cf.ips, ",") {
|
||||||
}
|
|
||||||
|
|
||||||
var networks []netip.Prefix
|
|
||||||
if *cf.networks == "" && *cf.ips != "" {
|
|
||||||
// Pull up deprecated -ips flag if needed
|
|
||||||
*cf.networks = *cf.ips
|
|
||||||
}
|
|
||||||
|
|
||||||
if *cf.networks != "" {
|
|
||||||
for _, rs := range strings.Split(*cf.networks, ",") {
|
|
||||||
rs := strings.Trim(rs, " ")
|
rs := strings.Trim(rs, " ")
|
||||||
if rs != "" {
|
if rs != "" {
|
||||||
n, err := netip.ParsePrefix(rs)
|
ip, ipNet, err := net.ParseCIDR(rs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return newHelpErrorf("invalid -networks definition: %s", rs)
|
return newHelpErrorf("invalid ip definition: %s", err)
|
||||||
}
|
}
|
||||||
if version == cert.Version1 && !n.Addr().Is4() {
|
if ip.To4() == nil {
|
||||||
return newHelpErrorf("invalid -networks definition: v1 certificates can only be ipv4, have %s", rs)
|
return newHelpErrorf("invalid ip definition: can only be ipv4, have %s", rs)
|
||||||
}
|
}
|
||||||
networks = append(networks, n)
|
|
||||||
|
ipNet.IP = ip
|
||||||
|
ips = append(ips, ipNet)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var unsafeNetworks []netip.Prefix
|
var subnets []*net.IPNet
|
||||||
if *cf.unsafeNetworks == "" && *cf.subnets != "" {
|
if *cf.subnets != "" {
|
||||||
// Pull up deprecated -subnets flag if needed
|
for _, rs := range strings.Split(*cf.subnets, ",") {
|
||||||
*cf.unsafeNetworks = *cf.subnets
|
|
||||||
}
|
|
||||||
|
|
||||||
if *cf.unsafeNetworks != "" {
|
|
||||||
for _, rs := range strings.Split(*cf.unsafeNetworks, ",") {
|
|
||||||
rs := strings.Trim(rs, " ")
|
rs := strings.Trim(rs, " ")
|
||||||
if rs != "" {
|
if rs != "" {
|
||||||
n, err := netip.ParsePrefix(rs)
|
_, s, err := net.ParseCIDR(rs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return newHelpErrorf("invalid -unsafe-networks definition: %s", rs)
|
return newHelpErrorf("invalid subnet definition: %s", err)
|
||||||
}
|
}
|
||||||
if version == cert.Version1 && !n.Addr().Is4() {
|
if s.IP.To4() == nil {
|
||||||
return newHelpErrorf("invalid -unsafe-networks definition: v1 certificates can only be ipv4, have %s", rs)
|
return newHelpErrorf("invalid subnet definition: can only be ipv4, have %s", rs)
|
||||||
}
|
}
|
||||||
unsafeNetworks = append(unsafeNetworks, n)
|
subnets = append(subnets, s)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var passphrase []byte
|
var passphrase []byte
|
||||||
if !isP11 && *cf.encryption {
|
if *cf.encryption {
|
||||||
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
for i := 0; i < 5; i++ {
|
||||||
|
out.Write([]byte("Enter passphrase: "))
|
||||||
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
|
if err == ErrNoTerminal {
|
||||||
|
return fmt.Errorf("out-key must be encrypted interactively")
|
||||||
|
} else if err != nil {
|
||||||
|
return fmt.Errorf("error reading passphrase: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(passphrase) > 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if len(passphrase) == 0 {
|
if len(passphrase) == 0 {
|
||||||
for i := 0; i < 5; i++ {
|
return fmt.Errorf("no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
||||||
out.Write([]byte("Enter passphrase: "))
|
|
||||||
passphrase, err = pr.ReadPassword()
|
|
||||||
|
|
||||||
if err == ErrNoTerminal {
|
|
||||||
return fmt.Errorf("out-key must be encrypted interactively")
|
|
||||||
} else if err != nil {
|
|
||||||
return fmt.Errorf("error reading passphrase: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(passphrase) > 0 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(passphrase) == 0 {
|
|
||||||
return fmt.Errorf("no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
var pub, rawPriv []byte
|
var pub, rawPriv []byte
|
||||||
var p11Client *pkclient.PKClient
|
switch *cf.curve {
|
||||||
|
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||||
if isP11 {
|
curve = cert.Curve_CURVE25519
|
||||||
switch *cf.curve {
|
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
||||||
case "P256":
|
|
||||||
curve = cert.Curve_P256
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve for PKCS#11: %s", *cf.curve)
|
|
||||||
}
|
|
||||||
|
|
||||||
p11Client, err = pkclient.FromUrl(*cf.p11url)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while creating PKCS#11 client: %w", err)
|
return fmt.Errorf("error while generating ed25519 keys: %s", err)
|
||||||
}
|
}
|
||||||
defer func(client *pkclient.PKClient) {
|
case "P256":
|
||||||
_ = client.Close()
|
var key *ecdsa.PrivateKey
|
||||||
}(p11Client)
|
curve = cert.Curve_P256
|
||||||
pub, err = p11Client.GetPubKey()
|
key, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while getting public key with PKCS#11: %w", err)
|
return fmt.Errorf("error while generating ecdsa keys: %s", err)
|
||||||
}
|
|
||||||
} else {
|
|
||||||
switch *cf.curve {
|
|
||||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
|
||||||
curve = cert.Curve_CURVE25519
|
|
||||||
pub, rawPriv, err = ed25519.GenerateKey(rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while generating ed25519 keys: %s", err)
|
|
||||||
}
|
|
||||||
case "P256":
|
|
||||||
var key *ecdsa.PrivateKey
|
|
||||||
curve = cert.Curve_P256
|
|
||||||
key, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while generating ecdsa keys: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ecdh.PrivateKey lets us get at the encoded bytes, even though
|
|
||||||
// we aren't using ECDH here.
|
|
||||||
eKey, err := key.ECDH()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while converting ecdsa key: %s", err)
|
|
||||||
}
|
|
||||||
rawPriv = eKey.Bytes()
|
|
||||||
pub = eKey.PublicKey().Bytes()
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", *cf.curve)
|
|
||||||
}
|
}
|
||||||
|
// ref: https://github.com/golang/go/blob/go1.19/src/crypto/x509/sec1.go#L60
|
||||||
|
rawPriv = key.D.FillBytes(make([]byte, 32))
|
||||||
|
pub = elliptic.Marshal(elliptic.P256(), key.X, key.Y)
|
||||||
}
|
}
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
nc := cert.NebulaCertificate{
|
||||||
Version: version,
|
Details: cert.NebulaCertificateDetails{
|
||||||
Name: *cf.name,
|
Name: *cf.name,
|
||||||
Groups: groups,
|
Groups: groups,
|
||||||
Networks: networks,
|
Ips: ips,
|
||||||
UnsafeNetworks: unsafeNetworks,
|
Subnets: subnets,
|
||||||
NotBefore: time.Now(),
|
NotBefore: time.Now(),
|
||||||
NotAfter: time.Now().Add(*cf.duration),
|
NotAfter: time.Now().Add(*cf.duration),
|
||||||
PublicKey: pub,
|
PublicKey: pub,
|
||||||
IsCA: true,
|
IsCA: true,
|
||||||
Curve: curve,
|
Curve: curve,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 {
|
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
||||||
if _, err := os.Stat(*cf.outKeyPath); err == nil {
|
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
||||||
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
if _, err := os.Stat(*cf.outCertPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
var c cert.Certificate
|
err = nc.Sign(curve, rawPriv)
|
||||||
var b []byte
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while signing: %s", err)
|
||||||
if isP11 {
|
|
||||||
c, err = t.SignWith(nil, curve, p11Client.SignASN1)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing with PKCS#11: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
c, err = t.Sign(nil, curve, rawPriv)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if *cf.encryption {
|
|
||||||
b, err = cert.EncryptAndMarshalSigningPrivateKey(curve, rawPriv, passphrase, kdfParams)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while encrypting out-key: %s", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outKeyPath, b, 0600)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
b, err = c.MarshalPEM()
|
if *cf.encryption {
|
||||||
|
b, err := cert.EncryptAndMarshalSigningPrivateKey(curve, rawPriv, passphrase, kdfParams)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while encrypting out-key: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = ioutil.WriteFile(*cf.outKeyPath, b, 0600)
|
||||||
|
} else {
|
||||||
|
err = ioutil.WriteFile(*cf.outKeyPath, cert.MarshalSigningPrivateKey(curve, rawPriv), 0600)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := nc.MarshalToPEM()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while marshalling certificate: %s", err)
|
return fmt.Errorf("error while marshalling certificate: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outCertPath, b, 0600)
|
err = ioutil.WriteFile(*cf.outCertPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -316,7 +244,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*cf.outQRPath, b, 0600)
|
err = ioutil.WriteFile(*cf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+74
-91
@@ -7,6 +7,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"encoding/pem"
|
"encoding/pem"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -14,9 +15,10 @@ import (
|
|||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
//TODO: test file permissions
|
||||||
|
|
||||||
func Test_caSummary(t *testing.T) {
|
func Test_caSummary(t *testing.T) {
|
||||||
assert.Equal(t, "ca <flags>: create a self signed certificate authority", caSummary())
|
assert.Equal(t, "ca <flags>: create a self signed certificate authority", caSummary())
|
||||||
}
|
}
|
||||||
@@ -42,24 +44,17 @@ func Test_caHelp(t *testing.T) {
|
|||||||
" -groups string\n"+
|
" -groups string\n"+
|
||||||
" \tOptional: comma separated list of groups. This will limit which groups subordinate certs can use\n"+
|
" \tOptional: comma separated list of groups. This will limit which groups subordinate certs can use\n"+
|
||||||
" -ips string\n"+
|
" -ips string\n"+
|
||||||
" Deprecated, see -networks\n"+
|
" \tOptional: comma separated list of ipv4 address and network in CIDR notation. This will limit which ipv4 addresses and networks subordinate certs can use for ip addresses\n"+
|
||||||
" -name string\n"+
|
" -name string\n"+
|
||||||
" \tRequired: name of the certificate authority\n"+
|
" \tRequired: name of the certificate authority\n"+
|
||||||
" -networks string\n"+
|
|
||||||
" \tOptional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in networks\n"+
|
|
||||||
" -out-crt string\n"+
|
" -out-crt string\n"+
|
||||||
" \tOptional: path to write the certificate to (default \"ca.crt\")\n"+
|
" \tOptional: path to write the certificate to (default \"ca.crt\")\n"+
|
||||||
" -out-key string\n"+
|
" -out-key string\n"+
|
||||||
" \tOptional: path to write the private key to (default \"ca.key\")\n"+
|
" \tOptional: path to write the private key to (default \"ca.key\")\n"+
|
||||||
" -out-qr string\n"+
|
" -out-qr string\n"+
|
||||||
" \tOptional: output a qr code image (png) of the certificate\n"+
|
" \tOptional: output a qr code image (png) of the certificate\n"+
|
||||||
optionalPkcs11String(" -pkcs11 string\n \tOptional: PKCS#11 URI to an existing private key\n")+
|
|
||||||
" -subnets string\n"+
|
" -subnets string\n"+
|
||||||
" \tDeprecated, see -unsafe-networks\n"+
|
" \tOptional: comma separated list of ipv4 address and network in CIDR notation. This will limit which ipv4 addresses and networks subordinate certs can use in subnets\n",
|
||||||
" -unsafe-networks string\n"+
|
|
||||||
" \tOptional: comma separated list of ip address and network in CIDR notation. This will limit which ip addresses and networks subordinate certs can use in unsafe networks\n"+
|
|
||||||
" -version uint\n"+
|
|
||||||
" \tOptional: version of the certificate format to use (default 2)\n",
|
|
||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -88,105 +83,93 @@ func Test_ca(t *testing.T) {
|
|||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, ca(
|
assertHelpError(t, ca(
|
||||||
[]string{"-version", "1", "-out-key", "nope", "-out-crt", "nope", "duration", "100m"}, ob, eb, nopw,
|
[]string{"-out-key", "nope", "-out-crt", "nope", "duration", "100m"}, ob, eb, nopw,
|
||||||
), "-name is required")
|
), "-name is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// ipv4 only ips
|
// ipv4 only ips
|
||||||
assertHelpError(t, ca([]string{"-version", "1", "-name", "ipv6", "-ips", "100::100/100"}, ob, eb, nopw), "invalid -networks definition: v1 certificates can only be ipv4, have 100::100/100")
|
assertHelpError(t, ca([]string{"-name", "ipv6", "-ips", "100::100/100"}, ob, eb, nopw), "invalid ip definition: can only be ipv4, have 100::100/100")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// ipv4 only subnets
|
// ipv4 only subnets
|
||||||
assertHelpError(t, ca([]string{"-version", "1", "-name", "ipv6", "-subnets", "100::100/100"}, ob, eb, nopw), "invalid -unsafe-networks definition: v1 certificates can only be ipv4, have 100::100/100")
|
assertHelpError(t, ca([]string{"-name", "ipv6", "-subnets", "100::100/100"}, ob, eb, nopw), "invalid subnet definition: can only be ipv4, have 100::100/100")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// failed key write
|
// failed key write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args := []string{"-version", "1", "-name", "test", "-duration", "100m", "-out-crt", "/do/not/write/pleasecrt", "-out-key", "/do/not/write/pleasekey"}
|
args := []string{"-name", "test", "-duration", "100m", "-out-crt", "/do/not/write/pleasecrt", "-out-key", "/do/not/write/pleasekey"}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
assert.EqualError(t, ca(args, ob, eb, nopw), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// create temp key file
|
// create temp key file
|
||||||
keyF, err := os.CreateTemp("", "test.key")
|
keyF, err := ioutil.TempFile("", "test.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, os.Remove(keyF.Name()))
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
// failed cert write
|
// failed cert write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-name", "test", "-duration", "100m", "-out-crt", "/do/not/write/pleasecrt", "-out-key", keyF.Name()}
|
args = []string{"-name", "test", "-duration", "100m", "-out-crt", "/do/not/write/pleasecrt", "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "error while writing out-crt: open /do/not/write/pleasecrt: "+NoSuchDirError)
|
assert.EqualError(t, ca(args, ob, eb, nopw), "error while writing out-crt: open /do/not/write/pleasecrt: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// create temp cert file
|
// create temp cert file
|
||||||
crtF, err := os.CreateTemp("", "test.crt")
|
crtF, err := ioutil.TempFile("", "test.crt")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, os.Remove(crtF.Name()))
|
os.Remove(crtF.Name())
|
||||||
require.NoError(t, os.Remove(keyF.Name()))
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
// test proper cert with removed empty groups and subnets
|
// test proper cert with removed empty groups and subnets
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, ca(args, ob, eb, nopw))
|
assert.Nil(t, ca(args, ob, eb, nopw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// read cert and key files
|
// read cert and key files
|
||||||
rb, _ := os.ReadFile(keyF.Name())
|
rb, _ := ioutil.ReadFile(keyF.Name())
|
||||||
lKey, b, c, err := cert.UnmarshalSigningPrivateKeyFromPEM(rb)
|
lKey, b, err := cert.UnmarshalEd25519PrivateKey(rb)
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, c)
|
assert.Len(t, b, 0)
|
||||||
assert.Empty(t, b)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Len(t, lKey, 64)
|
assert.Len(t, lKey, 64)
|
||||||
|
|
||||||
rb, _ = os.ReadFile(crtF.Name())
|
rb, _ = ioutil.ReadFile(crtF.Name())
|
||||||
lCrt, b, err := cert.UnmarshalCertificateFromPEM(rb)
|
lCrt, b, err := cert.UnmarshalNebulaCertificateFromPEM(rb)
|
||||||
assert.Empty(t, b)
|
assert.Len(t, b, 0)
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
|
|
||||||
assert.Equal(t, "test", lCrt.Name())
|
assert.Equal(t, "test", lCrt.Details.Name)
|
||||||
assert.Empty(t, lCrt.Networks())
|
assert.Len(t, lCrt.Details.Ips, 0)
|
||||||
assert.True(t, lCrt.IsCA())
|
assert.True(t, lCrt.Details.IsCA)
|
||||||
assert.Equal(t, []string{"1", "2", "3", "4", "5"}, lCrt.Groups())
|
assert.Equal(t, []string{"1", "2", "3", "4", "5"}, lCrt.Details.Groups)
|
||||||
assert.Empty(t, lCrt.UnsafeNetworks())
|
assert.Len(t, lCrt.Details.Subnets, 0)
|
||||||
assert.Len(t, lCrt.PublicKey(), 32)
|
assert.Len(t, lCrt.Details.PublicKey, 32)
|
||||||
assert.Equal(t, time.Duration(time.Minute*100), lCrt.NotAfter().Sub(lCrt.NotBefore()))
|
assert.Equal(t, time.Duration(time.Minute*100), lCrt.Details.NotAfter.Sub(lCrt.Details.NotBefore))
|
||||||
assert.Empty(t, lCrt.Issuer())
|
assert.Equal(t, "", lCrt.Details.Issuer)
|
||||||
assert.True(t, lCrt.CheckSignature(lCrt.PublicKey()))
|
assert.True(t, lCrt.CheckSignature(lCrt.Details.PublicKey))
|
||||||
|
|
||||||
// test encrypted key
|
// test encrypted key
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, ca(args, ob, eb, testpw))
|
assert.Nil(t, ca(args, ob, eb, testpw))
|
||||||
assert.Equal(t, pwPromptOb, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// test encrypted key with passphrase environment variable
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
|
||||||
require.NoError(t, ca(args, ob, eb, testpw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
|
||||||
|
|
||||||
// read encrypted key file and verify default params
|
// read encrypted key file and verify default params
|
||||||
rb, _ = os.ReadFile(keyF.Name())
|
rb, _ = ioutil.ReadFile(keyF.Name())
|
||||||
k, _ := pem.Decode(rb)
|
k, _ := pem.Decode(rb)
|
||||||
ned, err := cert.UnmarshalNebulaEncryptedData(k.Bytes)
|
ned, err := cert.UnmarshalNebulaEncryptedData(k.Bytes)
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
// we won't know salt in advance, so just check start of string
|
// we won't know salt in advance, so just check start of string
|
||||||
assert.Equal(t, uint32(2*1024*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
assert.Equal(t, uint32(2*1024*1024), ned.EncryptionMetadata.Argon2Parameters.Memory)
|
||||||
assert.Equal(t, uint8(4), ned.EncryptionMetadata.Argon2Parameters.Parallelism)
|
assert.Equal(t, uint8(4), ned.EncryptionMetadata.Argon2Parameters.Parallelism)
|
||||||
@@ -196,54 +179,54 @@ func Test_ca(t *testing.T) {
|
|||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
curve, lKey, b, err = cert.DecryptAndUnmarshalSigningPrivateKey(passphrase, rb)
|
curve, lKey, b, err = cert.DecryptAndUnmarshalSigningPrivateKey(passphrase, rb)
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
assert.Empty(t, b)
|
assert.Len(t, b, 0)
|
||||||
assert.Len(t, lKey, 64)
|
assert.Len(t, lKey, 64)
|
||||||
|
|
||||||
// test when reading password results in an error
|
// test when reading passsword results in an error
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.Error(t, ca(args, ob, eb, errpw))
|
assert.Error(t, ca(args, ob, eb, errpw))
|
||||||
assert.Equal(t, pwPromptOb, ob.String())
|
assert.Equal(t, pwPromptOb, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// test when user fails to enter a password
|
// test when user fails to enter a password
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
assert.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
|
||||||
assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
|
assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, ca(args, ob, eb, nopw))
|
assert.Nil(t, ca(args, ob, eb, nopw))
|
||||||
|
|
||||||
// test that we won't overwrite existing certificate file
|
// test that we won't overwrite existing certificate file
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "refusing to overwrite existing CA key: "+keyF.Name())
|
assert.EqualError(t, ca(args, ob, eb, nopw), "refusing to overwrite existing CA key: "+keyF.Name())
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// test that we won't overwrite existing key file
|
// test that we won't overwrite existing key file
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-name", "test", "-duration", "100m", "-groups", "1,, 2 , ,,,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, ca(args, ob, eb, nopw), "refusing to overwrite existing CA cert: "+crtF.Name())
|
assert.EqualError(t, ca(args, ob, eb, nopw), "refusing to overwrite existing CA cert: "+crtF.Name())
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
+21
-49
@@ -4,10 +4,9 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/pkclient"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -15,8 +14,8 @@ type keygenFlags struct {
|
|||||||
set *flag.FlagSet
|
set *flag.FlagSet
|
||||||
outKeyPath *string
|
outKeyPath *string
|
||||||
outPubPath *string
|
outPubPath *string
|
||||||
curve *string
|
|
||||||
p11url *string
|
curve *string
|
||||||
}
|
}
|
||||||
|
|
||||||
func newKeygenFlags() *keygenFlags {
|
func newKeygenFlags() *keygenFlags {
|
||||||
@@ -25,7 +24,6 @@ func newKeygenFlags() *keygenFlags {
|
|||||||
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
cf.outPubPath = cf.set.String("out-pub", "", "Required: path to write the public key to")
|
||||||
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
cf.outKeyPath = cf.set.String("out-key", "", "Required: path to write the private key to")
|
||||||
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
cf.curve = cf.set.String("curve", "25519", "ECDH Curve (25519, P256)")
|
||||||
cf.p11url = p11Flag(cf.set)
|
|
||||||
return &cf
|
return &cf
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -36,58 +34,32 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
isP11 := len(*cf.p11url) > 0
|
if err := mustFlagString("out-key", cf.outKeyPath); err != nil {
|
||||||
|
return err
|
||||||
if !isP11 {
|
|
||||||
if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
|
if err := mustFlagString("out-pub", cf.outPubPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var pub, rawPriv []byte
|
var pub, rawPriv []byte
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
if isP11 {
|
switch *cf.curve {
|
||||||
switch *cf.curve {
|
case "25519", "X25519", "Curve25519", "CURVE25519":
|
||||||
case "P256":
|
pub, rawPriv = x25519Keypair()
|
||||||
curve = cert.Curve_P256
|
curve = cert.Curve_CURVE25519
|
||||||
default:
|
case "P256":
|
||||||
return fmt.Errorf("invalid curve for PKCS#11: %s", *cf.curve)
|
pub, rawPriv = p256Keypair()
|
||||||
}
|
curve = cert.Curve_P256
|
||||||
} else {
|
default:
|
||||||
switch *cf.curve {
|
return fmt.Errorf("invalid curve: %s", *cf.curve)
|
||||||
case "25519", "X25519", "Curve25519", "CURVE25519":
|
|
||||||
pub, rawPriv = x25519Keypair()
|
|
||||||
curve = cert.Curve_CURVE25519
|
|
||||||
case "P256":
|
|
||||||
pub, rawPriv = p256Keypair()
|
|
||||||
curve = cert.Curve_P256
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("invalid curve: %s", *cf.curve)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if isP11 {
|
err = ioutil.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKey(curve, rawPriv), 0600)
|
||||||
p11Client, err := pkclient.FromUrl(*cf.p11url)
|
if err != nil {
|
||||||
if err != nil {
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
return fmt.Errorf("error while creating PKCS#11 client: %w", err)
|
|
||||||
}
|
|
||||||
defer func(client *pkclient.PKClient) {
|
|
||||||
_ = client.Close()
|
|
||||||
}(p11Client)
|
|
||||||
pub, err = p11Client.GetPubKey()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while getting public key: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
err = os.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
err = os.WriteFile(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600)
|
|
||||||
|
err = ioutil.WriteFile(*cf.outPubPath, cert.MarshalPublicKey(curve, pub), 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-pub: %s", err)
|
return fmt.Errorf("error while writing out-pub: %s", err)
|
||||||
}
|
}
|
||||||
@@ -101,7 +73,7 @@ func keygenSummary() string {
|
|||||||
|
|
||||||
func keygenHelp(out io.Writer) {
|
func keygenHelp(out io.Writer) {
|
||||||
cf := newKeygenFlags()
|
cf := newKeygenFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
|
||||||
cf.set.SetOutput(out)
|
cf.set.SetOutput(out)
|
||||||
cf.set.PrintDefaults()
|
cf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,14 +2,16 @@ package main
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
//TODO: test file permissions
|
||||||
|
|
||||||
func Test_keygenSummary(t *testing.T) {
|
func Test_keygenSummary(t *testing.T) {
|
||||||
assert.Equal(t, "keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`", keygenSummary())
|
assert.Equal(t, "keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`", keygenSummary())
|
||||||
}
|
}
|
||||||
@@ -25,8 +27,7 @@ func Test_keygenHelp(t *testing.T) {
|
|||||||
" -out-key string\n"+
|
" -out-key string\n"+
|
||||||
" \tRequired: path to write the private key to\n"+
|
" \tRequired: path to write the private key to\n"+
|
||||||
" -out-pub string\n"+
|
" -out-pub string\n"+
|
||||||
" \tRequired: path to write the public key to\n"+
|
" \tRequired: path to write the public key to\n",
|
||||||
optionalPkcs11String(" -pkcs11 string\n \tOptional: PKCS#11 URI to an existing private key\n"),
|
|
||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -37,59 +38,57 @@ func Test_keygen(t *testing.T) {
|
|||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, keygen([]string{"-out-pub", "nope"}, ob, eb), "-out-key is required")
|
assertHelpError(t, keygen([]string{"-out-pub", "nope"}, ob, eb), "-out-key is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
assertHelpError(t, keygen([]string{"-out-key", "nope"}, ob, eb), "-out-pub is required")
|
assertHelpError(t, keygen([]string{"-out-key", "nope"}, ob, eb), "-out-pub is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// failed key write
|
// failed key write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args := []string{"-out-pub", "/do/not/write/pleasepub", "-out-key", "/do/not/write/pleasekey"}
|
args := []string{"-out-pub", "/do/not/write/pleasepub", "-out-key", "/do/not/write/pleasekey"}
|
||||||
require.EqualError(t, keygen(args, ob, eb), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
assert.EqualError(t, keygen(args, ob, eb), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// create temp key file
|
// create temp key file
|
||||||
keyF, err := os.CreateTemp("", "test.key")
|
keyF, err := ioutil.TempFile("", "test.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(keyF.Name())
|
defer os.Remove(keyF.Name())
|
||||||
|
|
||||||
// failed pub write
|
// failed pub write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-out-pub", "/do/not/write/pleasepub", "-out-key", keyF.Name()}
|
args = []string{"-out-pub", "/do/not/write/pleasepub", "-out-key", keyF.Name()}
|
||||||
require.EqualError(t, keygen(args, ob, eb), "error while writing out-pub: open /do/not/write/pleasepub: "+NoSuchDirError)
|
assert.EqualError(t, keygen(args, ob, eb), "error while writing out-pub: open /do/not/write/pleasepub: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// create temp pub file
|
// create temp pub file
|
||||||
pubF, err := os.CreateTemp("", "test.pub")
|
pubF, err := ioutil.TempFile("", "test.pub")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(pubF.Name())
|
defer os.Remove(pubF.Name())
|
||||||
|
|
||||||
// test proper keygen
|
// test proper keygen
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-out-pub", pubF.Name(), "-out-key", keyF.Name()}
|
args = []string{"-out-pub", pubF.Name(), "-out-key", keyF.Name()}
|
||||||
require.NoError(t, keygen(args, ob, eb))
|
assert.Nil(t, keygen(args, ob, eb))
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// read cert and key files
|
// read cert and key files
|
||||||
rb, _ := os.ReadFile(keyF.Name())
|
rb, _ := ioutil.ReadFile(keyF.Name())
|
||||||
lKey, b, curve, err := cert.UnmarshalPrivateKeyFromPEM(rb)
|
lKey, b, err := cert.UnmarshalX25519PrivateKey(rb)
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
assert.Len(t, b, 0)
|
||||||
assert.Empty(t, b)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Len(t, lKey, 32)
|
assert.Len(t, lKey, 32)
|
||||||
|
|
||||||
rb, _ = os.ReadFile(pubF.Name())
|
rb, _ = ioutil.ReadFile(pubF.Name())
|
||||||
lPub, b, curve, err := cert.UnmarshalPublicKeyFromPEM(rb)
|
lPub, b, err := cert.UnmarshalX25519PublicKey(rb)
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
assert.Len(t, b, 0)
|
||||||
assert.Empty(t, b)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Len(t, lPub, 32)
|
assert.Len(t, lPub, 32)
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-19
@@ -5,28 +5,10 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"runtime/debug"
|
|
||||||
"strings"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// A version string that can be set with
|
|
||||||
//
|
|
||||||
// -ldflags "-X main.Build=SOMEVERSION"
|
|
||||||
//
|
|
||||||
// at compile-time.
|
|
||||||
var Build string
|
var Build string
|
||||||
|
|
||||||
func init() {
|
|
||||||
if Build == "" {
|
|
||||||
info, ok := debug.ReadBuildInfo()
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
Build = strings.TrimPrefix(info.Main.Version, "v")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type helpError struct {
|
type helpError struct {
|
||||||
s string
|
s string
|
||||||
}
|
}
|
||||||
@@ -35,7 +17,7 @@ func (he *helpError) Error() string {
|
|||||||
return he.s
|
return he.s
|
||||||
}
|
}
|
||||||
|
|
||||||
func newHelpErrorf(s string, v ...any) error {
|
func newHelpErrorf(s string, v ...interface{}) error {
|
||||||
return &helpError{s: fmt.Sprintf(s, v...)}
|
return &helpError{s: fmt.Sprintf(s, v...)}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,15 +3,15 @@ package main
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
//TODO: all flag parsing continueOnError will print to stderr on its own currently
|
||||||
|
|
||||||
func Test_help(t *testing.T) {
|
func Test_help(t *testing.T) {
|
||||||
expected := "Usage of " + os.Args[0] + " <global flags> <mode>:\n" +
|
expected := "Usage of " + os.Args[0] + " <global flags> <mode>:\n" +
|
||||||
" Global flags:\n" +
|
" Global flags:\n" +
|
||||||
@@ -77,16 +77,8 @@ func assertHelpError(t *testing.T, err error, msg string) {
|
|||||||
case *helpError:
|
case *helpError:
|
||||||
// good
|
// good
|
||||||
default:
|
default:
|
||||||
t.Fatal(fmt.Sprintf("err was not a helpError: %q, expected %q", err, msg))
|
t.Fatal("err was not a helpError")
|
||||||
}
|
}
|
||||||
|
|
||||||
require.EqualError(t, err, msg)
|
assert.EqualError(t, err, msg)
|
||||||
}
|
|
||||||
|
|
||||||
func optionalPkcs11String(msg string) string {
|
|
||||||
if p11Supported() {
|
|
||||||
return msg
|
|
||||||
} else {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,15 +0,0 @@
|
|||||||
//go:build cgo && pkcs11
|
|
||||||
|
|
||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"flag"
|
|
||||||
)
|
|
||||||
|
|
||||||
func p11Supported() bool {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func p11Flag(set *flag.FlagSet) *string {
|
|
||||||
return set.String("pkcs11", "", "Optional: PKCS#11 URI to an existing private key")
|
|
||||||
}
|
|
||||||
@@ -1,16 +0,0 @@
|
|||||||
//go:build !cgo || !pkcs11
|
|
||||||
|
|
||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"flag"
|
|
||||||
)
|
|
||||||
|
|
||||||
func p11Supported() bool {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func p11Flag(set *flag.FlagSet) *string {
|
|
||||||
var ret = ""
|
|
||||||
return &ret
|
|
||||||
}
|
|
||||||
+12
-16
@@ -5,6 +5,7 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -40,32 +41,33 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := os.ReadFile(*pf.path)
|
rawCert, err := ioutil.ReadFile(*pf.path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("unable to read cert; %s", err)
|
return fmt.Errorf("unable to read cert; %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var c cert.Certificate
|
var c *cert.NebulaCertificate
|
||||||
var qrBytes []byte
|
var qrBytes []byte
|
||||||
part := 0
|
part := 0
|
||||||
|
|
||||||
var jsonCerts []cert.Certificate
|
|
||||||
|
|
||||||
for {
|
for {
|
||||||
c, rawCert, err = cert.UnmarshalCertificateFromPEM(rawCert)
|
c, rawCert, err = cert.UnmarshalNebulaCertificateFromPEM(rawCert)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
return fmt.Errorf("error while unmarshaling cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json {
|
if *pf.json {
|
||||||
jsonCerts = append(jsonCerts, c)
|
b, _ := json.Marshal(c)
|
||||||
|
out.Write(b)
|
||||||
|
out.Write([]byte("\n"))
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
_, _ = out.Write([]byte(c.String()))
|
out.Write([]byte(c.String()))
|
||||||
_, _ = out.Write([]byte("\n"))
|
out.Write([]byte("\n"))
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
b, err := c.MarshalPEM()
|
b, err := c.MarshalToPEM()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while marshalling cert to PEM: %s", err)
|
return fmt.Errorf("error while marshalling cert to PEM: %s", err)
|
||||||
}
|
}
|
||||||
@@ -79,19 +81,13 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
part++
|
part++
|
||||||
}
|
}
|
||||||
|
|
||||||
if *pf.json {
|
|
||||||
b, _ := json.Marshal(jsonCerts)
|
|
||||||
_, _ = out.Write(b)
|
|
||||||
_, _ = out.Write([]byte("\n"))
|
|
||||||
}
|
|
||||||
|
|
||||||
if *pf.outQRPath != "" {
|
if *pf.outQRPath != "" {
|
||||||
b, err := qrcode.Encode(string(qrBytes), qrcode.Medium, -5)
|
b, err := qrcode.Encode(string(qrBytes), qrcode.Medium, -5)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*pf.outQRPath, b, 0600)
|
err = ioutil.WriteFile(*pf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-159
@@ -2,17 +2,13 @@ package main
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"crypto/ed25519"
|
"io/ioutil"
|
||||||
"crypto/rand"
|
|
||||||
"encoding/hex"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func Test_printSummary(t *testing.T) {
|
func Test_printSummary(t *testing.T) {
|
||||||
@@ -43,203 +39,84 @@ func Test_printCert(t *testing.T) {
|
|||||||
|
|
||||||
// no path
|
// no path
|
||||||
err := printCert([]string{}, ob, eb)
|
err := printCert([]string{}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
assertHelpError(t, err, "-path is required")
|
assertHelpError(t, err, "-path is required")
|
||||||
|
|
||||||
// no cert at path
|
// no cert at path
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
err = printCert([]string{"-path", "does_not_exist"}, ob, eb)
|
err = printCert([]string{"-path", "does_not_exist"}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.EqualError(t, err, "unable to read cert; open does_not_exist: "+NoSuchFileError)
|
assert.EqualError(t, err, "unable to read cert; open does_not_exist: "+NoSuchFileError)
|
||||||
|
|
||||||
// invalid cert at path
|
// invalid cert at path
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
tf, err := os.CreateTemp("", "print-cert")
|
tf, err := ioutil.TempFile("", "print-cert")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(tf.Name())
|
defer os.Remove(tf.Name())
|
||||||
|
|
||||||
tf.WriteString("-----BEGIN NOPE-----")
|
tf.WriteString("-----BEGIN NOPE-----")
|
||||||
err = printCert([]string{"-path", tf.Name()}, ob, eb)
|
err = printCert([]string{"-path", tf.Name()}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.EqualError(t, err, "error while unmarshaling cert: input did not contain a valid PEM encoded block")
|
assert.EqualError(t, err, "error while unmarshaling cert: input did not contain a valid PEM encoded block")
|
||||||
|
|
||||||
// test multiple certs
|
// test multiple certs
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
tf.Truncate(0)
|
tf.Truncate(0)
|
||||||
tf.Seek(0, 0)
|
tf.Seek(0, 0)
|
||||||
ca, caKey := NewTestCaCert("test ca", nil, nil, time.Time{}, time.Time{}, nil, nil, nil)
|
c := cert.NebulaCertificate{
|
||||||
c, _ := NewTestCert(ca, caKey, "test", time.Time{}, time.Time{}, []netip.Prefix{netip.MustParsePrefix("10.0.0.123/8")}, nil, []string{"hi"})
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "test",
|
||||||
|
Groups: []string{"hi"},
|
||||||
|
PublicKey: []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2},
|
||||||
|
},
|
||||||
|
Signature: []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2},
|
||||||
|
}
|
||||||
|
|
||||||
p, _ := c.MarshalPEM()
|
p, _ := c.MarshalToPEM()
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
|
|
||||||
err = printCert([]string{"-path", tf.Name()}, ob, eb)
|
err = printCert([]string{"-path", tf.Name()}, ob, eb)
|
||||||
fp, _ := c.Fingerprint()
|
assert.Nil(t, err)
|
||||||
pk := hex.EncodeToString(c.PublicKey())
|
|
||||||
sig := hex.EncodeToString(c.Signature())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
//"NebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: "+c.Issuer()+"\n\t\tPublic key: "+pk+"\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: "+fp+"\n\tSignature: "+sig+"\n}\nNebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: "+c.Issuer()+"\n\t\tPublic key: "+pk+"\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: "+fp+"\n\tSignature: "+sig+"\n}\nNebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: "+c.Issuer()+"\n\t\tPublic key: "+pk+"\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: "+fp+"\n\tSignature: "+sig+"\n}\n",
|
"NebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: \n\t\tPublic key: 0102030405060708090001020304050607080900010203040506070809000102\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\n\tSignature: 0102030405060708090001020304050607080900010203040506070809000102\n}\nNebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: \n\t\tPublic key: 0102030405060708090001020304050607080900010203040506070809000102\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\n\tSignature: 0102030405060708090001020304050607080900010203040506070809000102\n}\nNebulaCertificate {\n\tDetails {\n\t\tName: test\n\t\tIps: []\n\t\tSubnets: []\n\t\tGroups: [\n\t\t\t\"hi\"\n\t\t]\n\t\tNot before: 0001-01-01 00:00:00 +0000 UTC\n\t\tNot After: 0001-01-01 00:00:00 +0000 UTC\n\t\tIs CA: false\n\t\tIssuer: \n\t\tPublic key: 0102030405060708090001020304050607080900010203040506070809000102\n\t\tCurve: CURVE25519\n\t}\n\tFingerprint: cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\n\tSignature: 0102030405060708090001020304050607080900010203040506070809000102\n}\n",
|
||||||
`{
|
|
||||||
"details": {
|
|
||||||
"curve": "CURVE25519",
|
|
||||||
"groups": [
|
|
||||||
"hi"
|
|
||||||
],
|
|
||||||
"isCa": false,
|
|
||||||
"issuer": "`+c.Issuer()+`",
|
|
||||||
"name": "test",
|
|
||||||
"networks": [
|
|
||||||
"10.0.0.123/8"
|
|
||||||
],
|
|
||||||
"notAfter": "0001-01-01T00:00:00Z",
|
|
||||||
"notBefore": "0001-01-01T00:00:00Z",
|
|
||||||
"publicKey": "`+pk+`",
|
|
||||||
"unsafeNetworks": []
|
|
||||||
},
|
|
||||||
"fingerprint": "`+fp+`",
|
|
||||||
"signature": "`+sig+`",
|
|
||||||
"version": 1
|
|
||||||
}
|
|
||||||
{
|
|
||||||
"details": {
|
|
||||||
"curve": "CURVE25519",
|
|
||||||
"groups": [
|
|
||||||
"hi"
|
|
||||||
],
|
|
||||||
"isCa": false,
|
|
||||||
"issuer": "`+c.Issuer()+`",
|
|
||||||
"name": "test",
|
|
||||||
"networks": [
|
|
||||||
"10.0.0.123/8"
|
|
||||||
],
|
|
||||||
"notAfter": "0001-01-01T00:00:00Z",
|
|
||||||
"notBefore": "0001-01-01T00:00:00Z",
|
|
||||||
"publicKey": "`+pk+`",
|
|
||||||
"unsafeNetworks": []
|
|
||||||
},
|
|
||||||
"fingerprint": "`+fp+`",
|
|
||||||
"signature": "`+sig+`",
|
|
||||||
"version": 1
|
|
||||||
}
|
|
||||||
{
|
|
||||||
"details": {
|
|
||||||
"curve": "CURVE25519",
|
|
||||||
"groups": [
|
|
||||||
"hi"
|
|
||||||
],
|
|
||||||
"isCa": false,
|
|
||||||
"issuer": "`+c.Issuer()+`",
|
|
||||||
"name": "test",
|
|
||||||
"networks": [
|
|
||||||
"10.0.0.123/8"
|
|
||||||
],
|
|
||||||
"notAfter": "0001-01-01T00:00:00Z",
|
|
||||||
"notBefore": "0001-01-01T00:00:00Z",
|
|
||||||
"publicKey": "`+pk+`",
|
|
||||||
"unsafeNetworks": []
|
|
||||||
},
|
|
||||||
"fingerprint": "`+fp+`",
|
|
||||||
"signature": "`+sig+`",
|
|
||||||
"version": 1
|
|
||||||
}
|
|
||||||
`,
|
|
||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// test json
|
// test json
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
tf.Truncate(0)
|
tf.Truncate(0)
|
||||||
tf.Seek(0, 0)
|
tf.Seek(0, 0)
|
||||||
|
c = cert.NebulaCertificate{
|
||||||
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "test",
|
||||||
|
Groups: []string{"hi"},
|
||||||
|
PublicKey: []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2},
|
||||||
|
},
|
||||||
|
Signature: []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2},
|
||||||
|
}
|
||||||
|
|
||||||
|
p, _ = c.MarshalToPEM()
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
tf.Write(p)
|
tf.Write(p)
|
||||||
|
|
||||||
err = printCert([]string{"-json", "-path", tf.Name()}, ob, eb)
|
err = printCert([]string{"-json", "-path", tf.Name()}, ob, eb)
|
||||||
fp, _ = c.Fingerprint()
|
assert.Nil(t, err)
|
||||||
pk = hex.EncodeToString(c.PublicKey())
|
|
||||||
sig = hex.EncodeToString(c.Signature())
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(
|
assert.Equal(
|
||||||
t,
|
t,
|
||||||
`[{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1},{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1},{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1}]
|
"{\"details\":{\"curve\":\"CURVE25519\",\"groups\":[\"hi\"],\"ips\":[],\"isCa\":false,\"issuer\":\"\",\"name\":\"test\",\"notAfter\":\"0001-01-01T00:00:00Z\",\"notBefore\":\"0001-01-01T00:00:00Z\",\"publicKey\":\"0102030405060708090001020304050607080900010203040506070809000102\",\"subnets\":[]},\"fingerprint\":\"cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\",\"signature\":\"0102030405060708090001020304050607080900010203040506070809000102\"}\n{\"details\":{\"curve\":\"CURVE25519\",\"groups\":[\"hi\"],\"ips\":[],\"isCa\":false,\"issuer\":\"\",\"name\":\"test\",\"notAfter\":\"0001-01-01T00:00:00Z\",\"notBefore\":\"0001-01-01T00:00:00Z\",\"publicKey\":\"0102030405060708090001020304050607080900010203040506070809000102\",\"subnets\":[]},\"fingerprint\":\"cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\",\"signature\":\"0102030405060708090001020304050607080900010203040506070809000102\"}\n{\"details\":{\"curve\":\"CURVE25519\",\"groups\":[\"hi\"],\"ips\":[],\"isCa\":false,\"issuer\":\"\",\"name\":\"test\",\"notAfter\":\"0001-01-01T00:00:00Z\",\"notBefore\":\"0001-01-01T00:00:00Z\",\"publicKey\":\"0102030405060708090001020304050607080900010203040506070809000102\",\"subnets\":[]},\"fingerprint\":\"cc3492c0e9c48f17547f5987ea807462ebb3451e622590a10bb3763c344c82bd\",\"signature\":\"0102030405060708090001020304050607080900010203040506070809000102\"}\n",
|
||||||
`,
|
|
||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
}
|
|
||||||
|
|
||||||
// NewTestCaCert will generate a CA cert
|
|
||||||
func NewTestCaCert(name string, pubKey, privKey []byte, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte) {
|
|
||||||
var err error
|
|
||||||
if pubKey == nil || privKey == nil {
|
|
||||||
pubKey, privKey, err = ed25519.GenerateKey(rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
|
||||||
Version: cert.Version1,
|
|
||||||
Name: name,
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pubKey,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
IsCA: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := t.Sign(nil, cert.Curve_CURVE25519, privKey)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, privKey
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewTestCert(ca cert.Certificate, signerKey []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte) {
|
|
||||||
if before.IsZero() {
|
|
||||||
before = ca.NotBefore()
|
|
||||||
}
|
|
||||||
|
|
||||||
if after.IsZero() {
|
|
||||||
after = ca.NotAfter()
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(networks) == 0 {
|
|
||||||
networks = []netip.Prefix{netip.MustParsePrefix("10.0.0.123/8")}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub, rawPriv := x25519Keypair()
|
|
||||||
nc := &cert.TBSCertificate{
|
|
||||||
Version: cert.Version1,
|
|
||||||
Name: name,
|
|
||||||
Networks: networks,
|
|
||||||
UnsafeNetworks: unsafeNetworks,
|
|
||||||
Groups: groups,
|
|
||||||
NotBefore: time.Unix(before.Unix(), 0),
|
|
||||||
NotAfter: time.Unix(after.Unix(), 0),
|
|
||||||
PublicKey: pub,
|
|
||||||
IsCA: false,
|
|
||||||
}
|
|
||||||
|
|
||||||
c, err := nc.Sign(ca, ca.Curve(), signerKey)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return c, rawPriv
|
|
||||||
}
|
}
|
||||||
|
|||||||
+116
-249
@@ -3,63 +3,51 @@ package main
|
|||||||
import (
|
import (
|
||||||
"crypto/ecdh"
|
"crypto/ecdh"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"errors"
|
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"io/ioutil"
|
||||||
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/skip2/go-qrcode"
|
"github.com/skip2/go-qrcode"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/pkclient"
|
|
||||||
"golang.org/x/crypto/curve25519"
|
"golang.org/x/crypto/curve25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
type signFlags struct {
|
type signFlags struct {
|
||||||
set *flag.FlagSet
|
set *flag.FlagSet
|
||||||
version *uint
|
caKeyPath *string
|
||||||
caKeyPath *string
|
caCertPath *string
|
||||||
caCertPath *string
|
name *string
|
||||||
name *string
|
ip *string
|
||||||
networks *string
|
duration *time.Duration
|
||||||
unsafeNetworks *string
|
inPubPath *string
|
||||||
duration *time.Duration
|
outKeyPath *string
|
||||||
inPubPath *string
|
outCertPath *string
|
||||||
outKeyPath *string
|
outQRPath *string
|
||||||
outCertPath *string
|
groups *string
|
||||||
outQRPath *string
|
subnets *string
|
||||||
groups *string
|
|
||||||
|
|
||||||
p11url *string
|
|
||||||
|
|
||||||
// Deprecated options
|
|
||||||
ip *string
|
|
||||||
subnets *string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newSignFlags() *signFlags {
|
func newSignFlags() *signFlags {
|
||||||
sf := signFlags{set: flag.NewFlagSet("sign", flag.ContinueOnError)}
|
sf := signFlags{set: flag.NewFlagSet("sign", flag.ContinueOnError)}
|
||||||
sf.set.Usage = func() {}
|
sf.set.Usage = func() {}
|
||||||
sf.version = sf.set.Uint("version", 0, "Optional: version of the certificate format to use. The default is to match the version of the signing CA")
|
|
||||||
sf.caKeyPath = sf.set.String("ca-key", "ca.key", "Optional: path to the signing CA key")
|
sf.caKeyPath = sf.set.String("ca-key", "ca.key", "Optional: path to the signing CA key")
|
||||||
sf.caCertPath = sf.set.String("ca-crt", "ca.crt", "Optional: path to the signing CA cert")
|
sf.caCertPath = sf.set.String("ca-crt", "ca.crt", "Optional: path to the signing CA cert")
|
||||||
sf.name = sf.set.String("name", "", "Required: name of the cert, usually a hostname")
|
sf.name = sf.set.String("name", "", "Required: name of the cert, usually a hostname")
|
||||||
sf.networks = sf.set.String("networks", "", "Required: comma separated list of ip address and network in CIDR notation to assign to this cert")
|
sf.ip = sf.set.String("ip", "", "Required: ipv4 address and network in CIDR notation to assign the cert")
|
||||||
sf.unsafeNetworks = sf.set.String("unsafe-networks", "", "Optional: comma separated list of ip address and network in CIDR notation. Unsafe networks this cert can route for")
|
|
||||||
sf.duration = sf.set.Duration("duration", 0, "Optional: how long the cert should be valid for. The default is 1 second before the signing cert expires. Valid time units are seconds: \"s\", minutes: \"m\", hours: \"h\"")
|
sf.duration = sf.set.Duration("duration", 0, "Optional: how long the cert should be valid for. The default is 1 second before the signing cert expires. Valid time units are seconds: \"s\", minutes: \"m\", hours: \"h\"")
|
||||||
sf.inPubPath = sf.set.String("in-pub", "", "Optional (if out-key not set): path to read a previously generated public key")
|
sf.inPubPath = sf.set.String("in-pub", "", "Optional (if out-key not set): path to read a previously generated public key")
|
||||||
sf.outKeyPath = sf.set.String("out-key", "", "Optional (if in-pub not set): path to write the private key to")
|
sf.outKeyPath = sf.set.String("out-key", "", "Optional (if in-pub not set): path to write the private key to")
|
||||||
sf.outCertPath = sf.set.String("out-crt", "", "Optional: path to write the certificate to")
|
sf.outCertPath = sf.set.String("out-crt", "", "Optional: path to write the certificate to")
|
||||||
sf.outQRPath = sf.set.String("out-qr", "", "Optional: output a qr code image (png) of the certificate")
|
sf.outQRPath = sf.set.String("out-qr", "", "Optional: output a qr code image (png) of the certificate")
|
||||||
sf.groups = sf.set.String("groups", "", "Optional: comma separated list of groups")
|
sf.groups = sf.set.String("groups", "", "Optional: comma separated list of groups")
|
||||||
sf.p11url = p11Flag(sf.set)
|
sf.subnets = sf.set.String("subnets", "", "Optional: comma separated list of ipv4 address and network in CIDR notation. Subnets this cert can serve for")
|
||||||
|
|
||||||
sf.ip = sf.set.String("ip", "", "Deprecated, see -networks")
|
|
||||||
sf.subnets = sf.set.String("subnets", "", "Deprecated, see -unsafe-networks")
|
|
||||||
return &sf
|
return &sf
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error {
|
func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error {
|
||||||
@@ -69,12 +57,8 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
isP11 := len(*sf.p11url) > 0
|
if err := mustFlagString("ca-key", sf.caKeyPath); err != nil {
|
||||||
|
return err
|
||||||
if !isP11 {
|
|
||||||
if err := mustFlagString("ca-key", sf.caKeyPath); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if err := mustFlagString("ca-crt", sf.caCertPath); err != nil {
|
if err := mustFlagString("ca-crt", sf.caCertPath); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -82,144 +66,90 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
if err := mustFlagString("name", sf.name); err != nil {
|
if err := mustFlagString("name", sf.name); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
if err := mustFlagString("ip", sf.ip); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if *sf.inPubPath != "" && *sf.outKeyPath != "" {
|
||||||
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
return newHelpErrorf("cannot set both -in-pub and -out-key")
|
||||||
}
|
}
|
||||||
|
|
||||||
var v4Networks []netip.Prefix
|
rawCAKey, err := ioutil.ReadFile(*sf.caKeyPath)
|
||||||
var v6Networks []netip.Prefix
|
if err != nil {
|
||||||
if *sf.networks == "" && *sf.ip != "" {
|
return fmt.Errorf("error while reading ca-key: %s", err)
|
||||||
// Pull up deprecated -ip flag if needed
|
|
||||||
*sf.networks = *sf.ip
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(*sf.networks) == 0 {
|
|
||||||
return newHelpErrorf("-networks is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
version := cert.Version(*sf.version)
|
|
||||||
if version != 0 && version != cert.Version1 && version != cert.Version2 {
|
|
||||||
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var curve cert.Curve
|
var curve cert.Curve
|
||||||
var caKey []byte
|
var caKey []byte
|
||||||
|
|
||||||
if !isP11 {
|
// naively attempt to decode the private key as though it is not encrypted
|
||||||
var rawCAKey []byte
|
caKey, _, curve, err = cert.UnmarshalSigningPrivateKey(rawCAKey)
|
||||||
rawCAKey, err := os.ReadFile(*sf.caKeyPath)
|
if err == cert.ErrPrivateKeyEncrypted {
|
||||||
|
// ask for a passphrase until we get one
|
||||||
|
var passphrase []byte
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
out.Write([]byte("Enter passphrase: "))
|
||||||
|
passphrase, err = pr.ReadPassword()
|
||||||
|
|
||||||
|
if err == ErrNoTerminal {
|
||||||
|
return fmt.Errorf("ca-key is encrypted and must be decrypted interactively")
|
||||||
|
} else if err != nil {
|
||||||
|
return fmt.Errorf("error reading password: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(passphrase) > 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(passphrase) == 0 {
|
||||||
|
return fmt.Errorf("cannot open encrypted ca-key without passphrase")
|
||||||
|
}
|
||||||
|
|
||||||
|
curve, caKey, _, err = cert.DecryptAndUnmarshalSigningPrivateKey(passphrase, rawCAKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-key: %s", err)
|
return fmt.Errorf("error while parsing encrypted ca-key: %s", err)
|
||||||
}
|
|
||||||
|
|
||||||
// naively attempt to decode the private key as though it is not encrypted
|
|
||||||
caKey, _, curve, err = cert.UnmarshalSigningPrivateKeyFromPEM(rawCAKey)
|
|
||||||
if errors.Is(err, cert.ErrPrivateKeyEncrypted) {
|
|
||||||
var passphrase []byte
|
|
||||||
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
|
|
||||||
if len(passphrase) == 0 {
|
|
||||||
// ask for a passphrase until we get one
|
|
||||||
for i := 0; i < 5; i++ {
|
|
||||||
out.Write([]byte("Enter passphrase: "))
|
|
||||||
passphrase, err = pr.ReadPassword()
|
|
||||||
|
|
||||||
if errors.Is(err, ErrNoTerminal) {
|
|
||||||
return fmt.Errorf("ca-key is encrypted and must be decrypted interactively")
|
|
||||||
} else if err != nil {
|
|
||||||
return fmt.Errorf("error reading password: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(passphrase) > 0 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(passphrase) == 0 {
|
|
||||||
return fmt.Errorf("cannot open encrypted ca-key without passphrase")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
curve, caKey, _, err = cert.DecryptAndUnmarshalSigningPrivateKey(passphrase, rawCAKey)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while parsing encrypted ca-key: %s", err)
|
|
||||||
}
|
|
||||||
} else if err != nil {
|
|
||||||
return fmt.Errorf("error while parsing ca-key: %s", err)
|
|
||||||
}
|
}
|
||||||
|
} else if err != nil {
|
||||||
|
return fmt.Errorf("error while parsing ca-key: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCACert, err := os.ReadFile(*sf.caCertPath)
|
rawCACert, err := ioutil.ReadFile(*sf.caCertPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading ca-crt: %s", err)
|
return fmt.Errorf("error while reading ca-crt: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
caCert, _, err := cert.UnmarshalCertificateFromPEM(rawCACert)
|
caCert, _, err := cert.UnmarshalNebulaCertificateFromPEM(rawCACert)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while parsing ca-crt: %s", err)
|
return fmt.Errorf("error while parsing ca-crt: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 {
|
if err := caCert.VerifyPrivateKey(curve, caKey); err != nil {
|
||||||
if err := caCert.VerifyPrivateKey(curve, caKey); err != nil {
|
return fmt.Errorf("refusing to sign, root certificate does not match private key")
|
||||||
return fmt.Errorf("refusing to sign, root certificate does not match private key")
|
}
|
||||||
}
|
|
||||||
|
issuer, err := caCert.Sha256Sum()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while getting -ca-crt fingerprint: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if caCert.Expired(time.Now()) {
|
if caCert.Expired(time.Now()) {
|
||||||
return fmt.Errorf("ca certificate is expired")
|
return fmt.Errorf("ca certificate is expired")
|
||||||
}
|
}
|
||||||
|
|
||||||
if version == 0 {
|
|
||||||
version = caCert.Version()
|
|
||||||
}
|
|
||||||
|
|
||||||
// if no duration is given, expire one second before the root expires
|
// if no duration is given, expire one second before the root expires
|
||||||
if *sf.duration <= 0 {
|
if *sf.duration <= 0 {
|
||||||
*sf.duration = time.Until(caCert.NotAfter()) - time.Second*1
|
*sf.duration = time.Until(caCert.Details.NotAfter) - time.Second*1
|
||||||
}
|
}
|
||||||
|
|
||||||
if *sf.networks != "" {
|
ip, ipNet, err := net.ParseCIDR(*sf.ip)
|
||||||
for _, rs := range strings.Split(*sf.networks, ",") {
|
if err != nil {
|
||||||
rs := strings.Trim(rs, " ")
|
return newHelpErrorf("invalid ip definition: %s", err)
|
||||||
if rs != "" {
|
|
||||||
n, err := netip.ParsePrefix(rs)
|
|
||||||
if err != nil {
|
|
||||||
return newHelpErrorf("invalid -networks definition: %s", rs)
|
|
||||||
}
|
|
||||||
|
|
||||||
if n.Addr().Is4() {
|
|
||||||
v4Networks = append(v4Networks, n)
|
|
||||||
} else {
|
|
||||||
v6Networks = append(v6Networks, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
if ip.To4() == nil {
|
||||||
var v4UnsafeNetworks []netip.Prefix
|
return newHelpErrorf("invalid ip definition: can only be ipv4, have %s", *sf.ip)
|
||||||
var v6UnsafeNetworks []netip.Prefix
|
|
||||||
if *sf.unsafeNetworks == "" && *sf.subnets != "" {
|
|
||||||
// Pull up deprecated -subnets flag if needed
|
|
||||||
*sf.unsafeNetworks = *sf.subnets
|
|
||||||
}
|
}
|
||||||
|
ipNet.IP = ip
|
||||||
|
|
||||||
if *sf.unsafeNetworks != "" {
|
groups := []string{}
|
||||||
for _, rs := range strings.Split(*sf.unsafeNetworks, ",") {
|
|
||||||
rs := strings.Trim(rs, " ")
|
|
||||||
if rs != "" {
|
|
||||||
n, err := netip.ParsePrefix(rs)
|
|
||||||
if err != nil {
|
|
||||||
return newHelpErrorf("invalid -unsafe-networks definition: %s", rs)
|
|
||||||
}
|
|
||||||
|
|
||||||
if n.Addr().Is4() {
|
|
||||||
v4UnsafeNetworks = append(v4UnsafeNetworks, n)
|
|
||||||
} else {
|
|
||||||
v6UnsafeNetworks = append(v6UnsafeNetworks, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var groups []string
|
|
||||||
if *sf.groups != "" {
|
if *sf.groups != "" {
|
||||||
for _, rg := range strings.Split(*sf.groups, ",") {
|
for _, rg := range strings.Split(*sf.groups, ",") {
|
||||||
g := strings.TrimSpace(rg)
|
g := strings.TrimSpace(rg)
|
||||||
@@ -229,43 +159,60 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var pub, rawPriv []byte
|
subnets := []*net.IPNet{}
|
||||||
var p11Client *pkclient.PKClient
|
if *sf.subnets != "" {
|
||||||
|
for _, rs := range strings.Split(*sf.subnets, ",") {
|
||||||
if isP11 {
|
rs := strings.Trim(rs, " ")
|
||||||
curve = cert.Curve_P256
|
if rs != "" {
|
||||||
p11Client, err = pkclient.FromUrl(*sf.p11url)
|
_, s, err := net.ParseCIDR(rs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while creating PKCS#11 client: %w", err)
|
return newHelpErrorf("invalid subnet definition: %s", err)
|
||||||
|
}
|
||||||
|
if s.IP.To4() == nil {
|
||||||
|
return newHelpErrorf("invalid subnet definition: can only be ipv4, have %s", rs)
|
||||||
|
}
|
||||||
|
subnets = append(subnets, s)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
defer func(client *pkclient.PKClient) {
|
|
||||||
_ = client.Close()
|
|
||||||
}(p11Client)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var pub, rawPriv []byte
|
||||||
if *sf.inPubPath != "" {
|
if *sf.inPubPath != "" {
|
||||||
var pubCurve cert.Curve
|
rawPub, err := ioutil.ReadFile(*sf.inPubPath)
|
||||||
rawPub, err := os.ReadFile(*sf.inPubPath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while reading in-pub: %s", err)
|
return fmt.Errorf("error while reading in-pub: %s", err)
|
||||||
}
|
}
|
||||||
|
var pubCurve cert.Curve
|
||||||
pub, _, pubCurve, err = cert.UnmarshalPublicKeyFromPEM(rawPub)
|
pub, _, pubCurve, err = cert.UnmarshalPublicKey(rawPub)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while parsing in-pub: %s", err)
|
return fmt.Errorf("error while parsing in-pub: %s", err)
|
||||||
}
|
}
|
||||||
if pubCurve != curve {
|
if pubCurve != curve {
|
||||||
return fmt.Errorf("curve of in-pub does not match ca")
|
return fmt.Errorf("curve of in-pub does not match ca")
|
||||||
}
|
}
|
||||||
} else if isP11 {
|
|
||||||
pub, err = p11Client.GetPubKey()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while getting public key with PKCS#11: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
} else {
|
||||||
pub, rawPriv = newKeypair(curve)
|
pub, rawPriv = newKeypair(curve)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
nc := cert.NebulaCertificate{
|
||||||
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: *sf.name,
|
||||||
|
Ips: []*net.IPNet{ipNet},
|
||||||
|
Groups: groups,
|
||||||
|
Subnets: subnets,
|
||||||
|
NotBefore: time.Now(),
|
||||||
|
NotAfter: time.Now().Add(*sf.duration),
|
||||||
|
PublicKey: pub,
|
||||||
|
IsCA: false,
|
||||||
|
Issuer: issuer,
|
||||||
|
Curve: curve,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := nc.CheckRootConstrains(caCert); err != nil {
|
||||||
|
return fmt.Errorf("refusing to sign, root certificate constraints violated: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
if *sf.outKeyPath == "" {
|
if *sf.outKeyPath == "" {
|
||||||
*sf.outKeyPath = *sf.name + ".key"
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
}
|
}
|
||||||
@@ -278,108 +225,28 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
var crts []cert.Certificate
|
err = nc.Sign(curve, caKey)
|
||||||
|
if err != nil {
|
||||||
notBefore := time.Now()
|
return fmt.Errorf("error while signing: %s", err)
|
||||||
notAfter := notBefore.Add(*sf.duration)
|
|
||||||
|
|
||||||
switch version {
|
|
||||||
case cert.Version1:
|
|
||||||
// Make sure we have only one ipv4 address
|
|
||||||
if len(v4Networks) != 1 {
|
|
||||||
return newHelpErrorf("invalid -networks definition: v1 certificates can only have a single ipv4 address")
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(v6Networks) > 0 {
|
|
||||||
return newHelpErrorf("invalid -networks definition: v1 certificates can only contain ipv4 addresses")
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(v6UnsafeNetworks) > 0 {
|
|
||||||
return newHelpErrorf("invalid -unsafe-networks definition: v1 certificates can only contain ipv4 addresses")
|
|
||||||
}
|
|
||||||
|
|
||||||
t := &cert.TBSCertificate{
|
|
||||||
Version: cert.Version1,
|
|
||||||
Name: *sf.name,
|
|
||||||
Networks: []netip.Prefix{v4Networks[0]},
|
|
||||||
Groups: groups,
|
|
||||||
UnsafeNetworks: v4UnsafeNetworks,
|
|
||||||
NotBefore: notBefore,
|
|
||||||
NotAfter: notAfter,
|
|
||||||
PublicKey: pub,
|
|
||||||
IsCA: false,
|
|
||||||
Curve: curve,
|
|
||||||
}
|
|
||||||
|
|
||||||
var nc cert.Certificate
|
|
||||||
if p11Client == nil {
|
|
||||||
nc, err = t.Sign(caCert, curve, caKey)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
nc, err = t.SignWith(caCert, curve, p11Client.SignASN1)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing with PKCS#11: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
crts = append(crts, nc)
|
|
||||||
|
|
||||||
case cert.Version2:
|
|
||||||
t := &cert.TBSCertificate{
|
|
||||||
Version: cert.Version2,
|
|
||||||
Name: *sf.name,
|
|
||||||
Networks: append(v4Networks, v6Networks...),
|
|
||||||
Groups: groups,
|
|
||||||
UnsafeNetworks: append(v4UnsafeNetworks, v6UnsafeNetworks...),
|
|
||||||
NotBefore: notBefore,
|
|
||||||
NotAfter: notAfter,
|
|
||||||
PublicKey: pub,
|
|
||||||
IsCA: false,
|
|
||||||
Curve: curve,
|
|
||||||
}
|
|
||||||
|
|
||||||
var nc cert.Certificate
|
|
||||||
if p11Client == nil {
|
|
||||||
nc, err = t.Sign(caCert, curve, caKey)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
nc, err = t.SignWith(caCert, curve, p11Client.SignASN1)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while signing with PKCS#11: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
crts = append(crts, nc)
|
|
||||||
default:
|
|
||||||
// this should be unreachable
|
|
||||||
return fmt.Errorf("invalid version: %d", version)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if !isP11 && *sf.inPubPath == "" {
|
if *sf.inPubPath == "" {
|
||||||
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
if _, err := os.Stat(*sf.outKeyPath); err == nil {
|
||||||
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
|
err = ioutil.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKey(curve, rawPriv), 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-key: %s", err)
|
return fmt.Errorf("error while writing out-key: %s", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var b []byte
|
b, err := nc.MarshalToPEM()
|
||||||
for _, c := range crts {
|
if err != nil {
|
||||||
sb, err := c.MarshalPEM()
|
return fmt.Errorf("error while marshalling certificate: %s", err)
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while marshalling certificate: %s", err)
|
|
||||||
}
|
|
||||||
b = append(b, sb...)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outCertPath, b, 0600)
|
err = ioutil.WriteFile(*sf.outCertPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-crt: %s", err)
|
return fmt.Errorf("error while writing out-crt: %s", err)
|
||||||
}
|
}
|
||||||
@@ -390,7 +257,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
|
|||||||
return fmt.Errorf("error while generating qr code: %s", err)
|
return fmt.Errorf("error while generating qr code: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(*sf.outQRPath, b, 0600)
|
err = ioutil.WriteFile(*sf.outQRPath, b, 0600)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error while writing out-qr: %s", err)
|
return fmt.Errorf("error while writing out-qr: %s", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+120
-139
@@ -7,16 +7,18 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
//TODO: test file permissions
|
||||||
|
|
||||||
func Test_signSummary(t *testing.T) {
|
func Test_signSummary(t *testing.T) {
|
||||||
assert.Equal(t, "sign <flags>: create and sign a certificate", signSummary())
|
assert.Equal(t, "sign <flags>: create and sign a certificate", signSummary())
|
||||||
}
|
}
|
||||||
@@ -38,24 +40,17 @@ func Test_signHelp(t *testing.T) {
|
|||||||
" -in-pub string\n"+
|
" -in-pub string\n"+
|
||||||
" \tOptional (if out-key not set): path to read a previously generated public key\n"+
|
" \tOptional (if out-key not set): path to read a previously generated public key\n"+
|
||||||
" -ip string\n"+
|
" -ip string\n"+
|
||||||
" \tDeprecated, see -networks\n"+
|
" \tRequired: ipv4 address and network in CIDR notation to assign the cert\n"+
|
||||||
" -name string\n"+
|
" -name string\n"+
|
||||||
" \tRequired: name of the cert, usually a hostname\n"+
|
" \tRequired: name of the cert, usually a hostname\n"+
|
||||||
" -networks string\n"+
|
|
||||||
" \tRequired: comma separated list of ip address and network in CIDR notation to assign to this cert\n"+
|
|
||||||
" -out-crt string\n"+
|
" -out-crt string\n"+
|
||||||
" \tOptional: path to write the certificate to\n"+
|
" \tOptional: path to write the certificate to\n"+
|
||||||
" -out-key string\n"+
|
" -out-key string\n"+
|
||||||
" \tOptional (if in-pub not set): path to write the private key to\n"+
|
" \tOptional (if in-pub not set): path to write the private key to\n"+
|
||||||
" -out-qr string\n"+
|
" -out-qr string\n"+
|
||||||
" \tOptional: output a qr code image (png) of the certificate\n"+
|
" \tOptional: output a qr code image (png) of the certificate\n"+
|
||||||
optionalPkcs11String(" -pkcs11 string\n \tOptional: PKCS#11 URI to an existing private key\n")+
|
|
||||||
" -subnets string\n"+
|
" -subnets string\n"+
|
||||||
" \tDeprecated, see -unsafe-networks\n"+
|
" \tOptional: comma separated list of ipv4 address and network in CIDR notation. Subnets this cert can serve for\n",
|
||||||
" -unsafe-networks string\n"+
|
|
||||||
" \tOptional: comma separated list of ip address and network in CIDR notation. Unsafe networks this cert can route for\n"+
|
|
||||||
" -version uint\n"+
|
|
||||||
" \tOptional: version of the certificate format to use. The default is to match the version of the signing CA\n",
|
|
||||||
ob.String(),
|
ob.String(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -82,20 +77,20 @@ func Test_signCert(t *testing.T) {
|
|||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, signCert(
|
assertHelpError(t, signCert(
|
||||||
[]string{"-version", "1", "-ca-crt", "./nope", "-ca-key", "./nope", "-ip", "1.1.1.1/24", "-out-key", "nope", "-out-crt", "nope"}, ob, eb, nopw,
|
[]string{"-ca-crt", "./nope", "-ca-key", "./nope", "-ip", "1.1.1.1/24", "-out-key", "nope", "-out-crt", "nope"}, ob, eb, nopw,
|
||||||
), "-name is required")
|
), "-name is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
assertHelpError(t, signCert(
|
assertHelpError(t, signCert(
|
||||||
[]string{"-version", "1", "-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-out-key", "nope", "-out-crt", "nope"}, ob, eb, nopw,
|
[]string{"-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-out-key", "nope", "-out-crt", "nope"}, ob, eb, nopw,
|
||||||
), "-networks is required")
|
), "-ip is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// cannot set -in-pub and -out-key
|
// cannot set -in-pub and -out-key
|
||||||
assertHelpError(t, signCert(
|
assertHelpError(t, signCert(
|
||||||
[]string{"-version", "1", "-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-in-pub", "nope", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope"}, ob, eb, nopw,
|
[]string{"-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-in-pub", "nope", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope"}, ob, eb, nopw,
|
||||||
), "cannot set both -in-pub and -out-key")
|
), "cannot set both -in-pub and -out-key")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
@@ -103,18 +98,18 @@ func Test_signCert(t *testing.T) {
|
|||||||
// failed to read key
|
// failed to read key
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args := []string{"-version", "1", "-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args := []string{"-ca-crt", "./nope", "-ca-key", "./nope", "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while reading ca-key: open ./nope: "+NoSuchFileError)
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while reading ca-key: open ./nope: "+NoSuchFileError)
|
||||||
|
|
||||||
// failed to unmarshal key
|
// failed to unmarshal key
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
caKeyF, err := os.CreateTemp("", "sign-cert.key")
|
caKeyF, err := ioutil.TempFile("", "sign-cert.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caKeyF.Name())
|
defer os.Remove(caKeyF.Name())
|
||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", "./nope", "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args = []string{"-ca-crt", "./nope", "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing ca-key: input did not contain a valid PEM encoded block")
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing ca-key: input did not contain a valid PEM encoded block")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
@@ -122,46 +117,54 @@ func Test_signCert(t *testing.T) {
|
|||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
caKeyF.Write(cert.MarshalSigningPrivateKeyToPEM(cert.Curve_CURVE25519, caPriv))
|
caKeyF.Write(cert.MarshalEd25519PrivateKey(caPriv))
|
||||||
|
|
||||||
// failed to read cert
|
// failed to read cert
|
||||||
args = []string{"-version", "1", "-ca-crt", "./nope", "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args = []string{"-ca-crt", "./nope", "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while reading ca-crt: open ./nope: "+NoSuchFileError)
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while reading ca-crt: open ./nope: "+NoSuchFileError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// failed to unmarshal cert
|
// failed to unmarshal cert
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
caCrtF, err := os.CreateTemp("", "sign-cert.crt")
|
caCrtF, err := ioutil.TempFile("", "sign-cert.crt")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caCrtF.Name())
|
defer os.Remove(caCrtF.Name())
|
||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing ca-crt: input did not contain a valid PEM encoded block")
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing ca-crt: input did not contain a valid PEM encoded block")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// write a proper ca cert for later
|
// write a proper ca cert for later
|
||||||
ca, _ := NewTestCaCert("ca", caPub, caPriv, time.Now(), time.Now().Add(time.Minute*200), nil, nil, nil)
|
ca := cert.NebulaCertificate{
|
||||||
b, _ := ca.MarshalPEM()
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "ca",
|
||||||
|
NotBefore: time.Now(),
|
||||||
|
NotAfter: time.Now().Add(time.Minute * 200),
|
||||||
|
PublicKey: caPub,
|
||||||
|
IsCA: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
b, _ := ca.MarshalToPEM()
|
||||||
caCrtF.Write(b)
|
caCrtF.Write(b)
|
||||||
|
|
||||||
// failed to read pub
|
// failed to read pub
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-in-pub", "./nope", "-duration", "100m"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-in-pub", "./nope", "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while reading in-pub: open ./nope: "+NoSuchFileError)
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while reading in-pub: open ./nope: "+NoSuchFileError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// failed to unmarshal pub
|
// failed to unmarshal pub
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
inPubF, err := os.CreateTemp("", "in.pub")
|
inPubF, err := ioutil.TempFile("", "in.pub")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(inPubF.Name())
|
defer os.Remove(inPubF.Name())
|
||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-in-pub", inPubF.Name(), "-duration", "100m"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-in-pub", inPubF.Name(), "-duration", "100m"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing in-pub: input did not contain a valid PEM encoded block")
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing in-pub: input did not contain a valid PEM encoded block")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
@@ -169,124 +172,116 @@ func Test_signCert(t *testing.T) {
|
|||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
inPub, _ := x25519Keypair()
|
inPub, _ := x25519Keypair()
|
||||||
inPubF.Write(cert.MarshalPublicKeyToPEM(cert.Curve_CURVE25519, inPub))
|
inPubF.Write(cert.MarshalX25519PublicKey(inPub))
|
||||||
|
|
||||||
// bad ip cidr
|
// bad ip cidr
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "a1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "a1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid -networks definition: a1.1.1.1/24")
|
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid ip definition: invalid CIDR address: a1.1.1.1/24")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "100::100/100", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "100::100/100", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
||||||
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid -networks definition: v1 certificates can only have a single ipv4 address")
|
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid ip definition: can only be ipv4, have 100::100/100")
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24,1.1.1.2/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
|
|
||||||
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid -networks definition: v1 certificates can only have a single ipv4 address")
|
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// bad subnet cidr
|
// bad subnet cidr
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "a"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "a"}
|
||||||
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid -unsafe-networks definition: a")
|
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid subnet definition: invalid CIDR address: a")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "100::100/100"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "100::100/100"}
|
||||||
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid -unsafe-networks definition: v1 certificates can only contain ipv4 addresses")
|
assertHelpError(t, signCert(args, ob, eb, nopw), "invalid subnet definition: can only be ipv4, have 100::100/100")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// mismatched ca key
|
// mismatched ca key
|
||||||
_, caPriv2, _ := ed25519.GenerateKey(rand.Reader)
|
_, caPriv2, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
caKeyF2, err := os.CreateTemp("", "sign-cert-2.key")
|
caKeyF2, err := ioutil.TempFile("", "sign-cert-2.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caKeyF2.Name())
|
defer os.Remove(caKeyF2.Name())
|
||||||
caKeyF2.Write(cert.MarshalSigningPrivateKeyToPEM(cert.Curve_CURVE25519, caPriv2))
|
caKeyF2.Write(cert.MarshalEd25519PrivateKey(caPriv2))
|
||||||
|
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF2.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "a"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF2.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m", "-subnets", "a"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "refusing to sign, root certificate does not match private key")
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "refusing to sign, root certificate does not match private key")
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// failed key write
|
// failed key write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "/do/not/write/pleasecrt", "-out-key", "/do/not/write/pleasekey", "-duration", "100m", "-subnets", "10.1.1.1/32"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "/do/not/write/pleasecrt", "-out-key", "/do/not/write/pleasekey", "-duration", "100m", "-subnets", "10.1.1.1/32"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while writing out-key: open /do/not/write/pleasekey: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// create temp key file
|
// create temp key file
|
||||||
keyF, err := os.CreateTemp("", "test.key")
|
keyF, err := ioutil.TempFile("", "test.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
// failed cert write
|
// failed cert write
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "/do/not/write/pleasecrt", "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", "/do/not/write/pleasecrt", "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while writing out-crt: open /do/not/write/pleasecrt: "+NoSuchDirError)
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "error while writing out-crt: open /do/not/write/pleasecrt: "+NoSuchDirError)
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
|
|
||||||
// create temp cert file
|
// create temp cert file
|
||||||
crtF, err := os.CreateTemp("", "test.crt")
|
crtF, err := ioutil.TempFile("", "test.crt")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
|
|
||||||
// test proper cert with removed empty groups and subnets
|
// test proper cert with removed empty groups and subnets
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
assert.Nil(t, signCert(args, ob, eb, nopw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// read cert and key files
|
// read cert and key files
|
||||||
rb, _ := os.ReadFile(keyF.Name())
|
rb, _ := ioutil.ReadFile(keyF.Name())
|
||||||
lKey, b, curve, err := cert.UnmarshalPrivateKeyFromPEM(rb)
|
lKey, b, err := cert.UnmarshalX25519PrivateKey(rb)
|
||||||
assert.Equal(t, cert.Curve_CURVE25519, curve)
|
assert.Len(t, b, 0)
|
||||||
assert.Empty(t, b)
|
assert.Nil(t, err)
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Len(t, lKey, 32)
|
assert.Len(t, lKey, 32)
|
||||||
|
|
||||||
rb, _ = os.ReadFile(crtF.Name())
|
rb, _ = ioutil.ReadFile(crtF.Name())
|
||||||
lCrt, b, err := cert.UnmarshalCertificateFromPEM(rb)
|
lCrt, b, err := cert.UnmarshalNebulaCertificateFromPEM(rb)
|
||||||
assert.Empty(t, b)
|
assert.Len(t, b, 0)
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
|
|
||||||
assert.Equal(t, "test", lCrt.Name())
|
assert.Equal(t, "test", lCrt.Details.Name)
|
||||||
assert.Equal(t, "1.1.1.1/24", lCrt.Networks()[0].String())
|
assert.Equal(t, "1.1.1.1/24", lCrt.Details.Ips[0].String())
|
||||||
assert.Len(t, lCrt.Networks(), 1)
|
assert.Len(t, lCrt.Details.Ips, 1)
|
||||||
assert.False(t, lCrt.IsCA())
|
assert.False(t, lCrt.Details.IsCA)
|
||||||
assert.Equal(t, []string{"1", "2", "3", "4", "5"}, lCrt.Groups())
|
assert.Equal(t, []string{"1", "2", "3", "4", "5"}, lCrt.Details.Groups)
|
||||||
assert.Len(t, lCrt.UnsafeNetworks(), 3)
|
assert.Len(t, lCrt.Details.Subnets, 3)
|
||||||
assert.Len(t, lCrt.PublicKey(), 32)
|
assert.Len(t, lCrt.Details.PublicKey, 32)
|
||||||
assert.Equal(t, time.Duration(time.Minute*100), lCrt.NotAfter().Sub(lCrt.NotBefore()))
|
assert.Equal(t, time.Duration(time.Minute*100), lCrt.Details.NotAfter.Sub(lCrt.Details.NotBefore))
|
||||||
|
|
||||||
sns := []string{}
|
sns := []string{}
|
||||||
for _, sn := range lCrt.UnsafeNetworks() {
|
for _, sn := range lCrt.Details.Subnets {
|
||||||
sns = append(sns, sn.String())
|
sns = append(sns, sn.String())
|
||||||
}
|
}
|
||||||
assert.Equal(t, []string{"10.1.1.1/32", "10.2.2.2/32", "10.5.5.5/32"}, sns)
|
assert.Equal(t, []string{"10.1.1.1/32", "10.2.2.2/32", "10.5.5.5/32"}, sns)
|
||||||
|
|
||||||
issuer, _ := ca.Fingerprint()
|
issuer, _ := ca.Sha256Sum()
|
||||||
assert.Equal(t, issuer, lCrt.Issuer())
|
assert.Equal(t, issuer, lCrt.Details.Issuer)
|
||||||
|
|
||||||
assert.True(t, lCrt.CheckSignature(caPub))
|
assert.True(t, lCrt.CheckSignature(caPub))
|
||||||
|
|
||||||
@@ -295,55 +290,53 @@ func Test_signCert(t *testing.T) {
|
|||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-in-pub", inPubF.Name(), "-duration", "100m", "-groups", "1"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-in-pub", inPubF.Name(), "-duration", "100m", "-groups", "1"}
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
assert.Nil(t, signCert(args, ob, eb, nopw))
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// read cert file and check pub key matches in-pub
|
// read cert file and check pub key matches in-pub
|
||||||
rb, _ = os.ReadFile(crtF.Name())
|
rb, _ = ioutil.ReadFile(crtF.Name())
|
||||||
lCrt, b, err = cert.UnmarshalCertificateFromPEM(rb)
|
lCrt, b, err = cert.UnmarshalNebulaCertificateFromPEM(rb)
|
||||||
assert.Empty(t, b)
|
assert.Len(t, b, 0)
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
assert.Equal(t, lCrt.PublicKey(), inPub)
|
assert.Equal(t, lCrt.Details.PublicKey, inPub)
|
||||||
|
|
||||||
// test refuse to sign cert with duration beyond root
|
// test refuse to sign cert with duration beyond root
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
os.Remove(keyF.Name())
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "1000m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
os.Remove(crtF.Name())
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "refusing to sign, root certificate constraints violated: certificate expires after signing certificate")
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "1000m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while signing: certificate expires after signing certificate")
|
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
assert.Nil(t, signCert(args, ob, eb, nopw))
|
||||||
|
|
||||||
// test that we won't overwrite existing key file
|
// test that we won't overwrite existing key file
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "refusing to overwrite existing key: "+keyF.Name())
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "refusing to overwrite existing key: "+keyF.Name())
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// create valid cert/key for overwrite tests
|
// create valid cert/key for overwrite tests
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
os.Remove(crtF.Name())
|
os.Remove(crtF.Name())
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, nopw))
|
assert.Nil(t, signCert(args, ob, eb, nopw))
|
||||||
|
|
||||||
// test that we won't overwrite existing certificate file
|
// test that we won't overwrite existing certificate file
|
||||||
os.Remove(keyF.Name())
|
os.Remove(keyF.Name())
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "refusing to overwrite existing cert: "+crtF.Name())
|
assert.EqualError(t, signCert(args, ob, eb, nopw), "refusing to overwrite existing cert: "+crtF.Name())
|
||||||
assert.Empty(t, ob.String())
|
assert.Empty(t, ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
@@ -355,12 +348,12 @@ func Test_signCert(t *testing.T) {
|
|||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
|
|
||||||
caKeyF, err = os.CreateTemp("", "sign-cert.key")
|
caKeyF, err = ioutil.TempFile("", "sign-cert.key")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caKeyF.Name())
|
defer os.Remove(caKeyF.Name())
|
||||||
|
|
||||||
caCrtF, err = os.CreateTemp("", "sign-cert.crt")
|
caCrtF, err = ioutil.TempFile("", "sign-cert.crt")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caCrtF.Name())
|
defer os.Remove(caCrtF.Name())
|
||||||
|
|
||||||
// generate the encrypted key
|
// generate the encrypted key
|
||||||
@@ -369,52 +362,40 @@ func Test_signCert(t *testing.T) {
|
|||||||
b, _ = cert.EncryptAndMarshalSigningPrivateKey(cert.Curve_CURVE25519, caPriv, passphrase, kdfParams)
|
b, _ = cert.EncryptAndMarshalSigningPrivateKey(cert.Curve_CURVE25519, caPriv, passphrase, kdfParams)
|
||||||
caKeyF.Write(b)
|
caKeyF.Write(b)
|
||||||
|
|
||||||
ca, _ = NewTestCaCert("ca", caPub, caPriv, time.Now(), time.Now().Add(time.Minute*200), nil, nil, nil)
|
ca = cert.NebulaCertificate{
|
||||||
b, _ = ca.MarshalPEM()
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "ca",
|
||||||
|
NotBefore: time.Now(),
|
||||||
|
NotAfter: time.Now().Add(time.Minute * 200),
|
||||||
|
PublicKey: caPub,
|
||||||
|
IsCA: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
b, _ = ca.MarshalToPEM()
|
||||||
caCrtF.Write(b)
|
caCrtF.Write(b)
|
||||||
|
|
||||||
// test with the proper password
|
// test with the proper password
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
assert.Nil(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test with the proper password in the environment
|
|
||||||
os.Remove(crtF.Name())
|
|
||||||
os.Remove(keyF.Name())
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
|
|
||||||
require.NoError(t, signCert(args, ob, eb, testpw))
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
|
||||||
|
|
||||||
// test with the wrong password
|
// test with the wrong password
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
|
|
||||||
testpw.password = []byte("invalid password")
|
testpw.password = []byte("invalid password")
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, testpw))
|
assert.Error(t, signCert(args, ob, eb, testpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
|
|
||||||
// test with the wrong password in environment
|
|
||||||
ob.Reset()
|
|
||||||
eb.Reset()
|
|
||||||
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "invalid password")
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
|
||||||
require.EqualError(t, signCert(args, ob, eb, nopw), "error while parsing encrypted ca-key: invalid passphrase or corrupt private key")
|
|
||||||
assert.Empty(t, ob.String())
|
|
||||||
assert.Empty(t, eb.String())
|
|
||||||
os.Setenv("NEBULA_CA_PASSPHRASE", "")
|
|
||||||
|
|
||||||
// test with the user not entering a password
|
// test with the user not entering a password
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, nopw))
|
assert.Error(t, signCert(args, ob, eb, nopw))
|
||||||
// normally the user hitting enter on the prompt would add newlines between these
|
// normally the user hitting enter on the prompt would add newlines between these
|
||||||
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
|
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
@@ -423,8 +404,8 @@ func Test_signCert(t *testing.T) {
|
|||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
|
|
||||||
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
args = []string{"-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
|
||||||
require.Error(t, signCert(args, ob, eb, errpw))
|
assert.Error(t, signCert(args, ob, eb, errpw))
|
||||||
assert.Equal(t, "Enter passphrase: ", ob.String())
|
assert.Equal(t, "Enter passphrase: ", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Empty(t, eb.String())
|
||||||
}
|
}
|
||||||
|
|||||||
+28
-31
@@ -1,11 +1,12 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
@@ -39,43 +40,39 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
caFile, err := os.Open(*vf.caPath)
|
rawCACert, err := ioutil.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: %s", err)
|
||||||
}
|
|
||||||
defer caFile.Close()
|
|
||||||
|
|
||||||
caPool, err := cert.NewCAPoolFromPEMReader(caFile)
|
|
||||||
if err != nil && !errors.Is(err, cert.ErrExpired) {
|
|
||||||
return fmt.Errorf("error while adding ca cert to pool: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
rawCert, err := os.ReadFile(*vf.certPath)
|
caPool := cert.NewCAPool()
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("unable to read crt: %w", err)
|
|
||||||
}
|
|
||||||
var errs []error
|
|
||||||
for {
|
for {
|
||||||
if len(rawCert) == 0 {
|
rawCACert, err = caPool.AddCACertificate(rawCACert)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while adding ca cert to pool: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if rawCACert == nil || len(rawCACert) == 0 || strings.TrimSpace(string(rawCACert)) == "" {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
c, extra, err := cert.UnmarshalCertificateFromPEM(rawCert)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("error while parsing crt: %w", err)
|
|
||||||
}
|
|
||||||
rawCert = extra
|
|
||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
if err != nil {
|
|
||||||
switch {
|
|
||||||
case errors.Is(err, cert.ErrCaNotFound):
|
|
||||||
errs = append(errs, fmt.Errorf("error while verifying certificate v%d %s with issuer %s: %w", c.Version(), c.Name(), c.Issuer(), err))
|
|
||||||
default:
|
|
||||||
errs = append(errs, fmt.Errorf("error while verifying certificate %+v: %w", c, err))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.Join(errs...)
|
rawCert, err := ioutil.ReadFile(*vf.certPath)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("unable to read crt; %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
c, _, err := cert.UnmarshalNebulaCertificateFromPEM(rawCert)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("error while parsing crt: %s", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
good, err := c.Verify(time.Now(), caPool)
|
||||||
|
if !good {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func verifySummary() string {
|
func verifySummary() string {
|
||||||
@@ -84,7 +81,7 @@ func verifySummary() string {
|
|||||||
|
|
||||||
func verifyHelp(out io.Writer) {
|
func verifyHelp(out io.Writer) {
|
||||||
vf := newVerifyFlags()
|
vf := newVerifyFlags()
|
||||||
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
|
||||||
vf.set.SetOutput(out)
|
vf.set.SetOutput(out)
|
||||||
vf.set.PrintDefaults()
|
vf.set.PrintDefaults()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,13 +3,13 @@ package main
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"golang.org/x/crypto/ed25519"
|
"golang.org/x/crypto/ed25519"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -38,87 +38,105 @@ func Test_verify(t *testing.T) {
|
|||||||
|
|
||||||
// required args
|
// required args
|
||||||
assertHelpError(t, verify([]string{"-ca", "derp"}, ob, eb), "-crt is required")
|
assertHelpError(t, verify([]string{"-ca", "derp"}, ob, eb), "-crt is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
assertHelpError(t, verify([]string{"-crt", "derp"}, ob, eb), "-ca is required")
|
assertHelpError(t, verify([]string{"-crt", "derp"}, ob, eb), "-ca is required")
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
|
|
||||||
// no ca at path
|
// no ca at path
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
err := verify([]string{"-ca", "does_not_exist", "-crt", "does_not_exist"}, ob, eb)
|
err := verify([]string{"-ca", "does_not_exist", "-crt", "does_not_exist"}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.EqualError(t, err, "error while reading ca: open does_not_exist: "+NoSuchFileError)
|
assert.EqualError(t, err, "error while reading ca: open does_not_exist: "+NoSuchFileError)
|
||||||
|
|
||||||
// invalid ca at path
|
// invalid ca at path
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
caFile, err := os.CreateTemp("", "verify-ca")
|
caFile, err := ioutil.TempFile("", "verify-ca")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(caFile.Name())
|
defer os.Remove(caFile.Name())
|
||||||
|
|
||||||
caFile.WriteString("-----BEGIN NOPE-----")
|
caFile.WriteString("-----BEGIN NOPE-----")
|
||||||
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.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.ErrorIs(t, err, cert.ErrInvalidPEMBlock)
|
assert.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)
|
||||||
ca, _ := NewTestCaCert("test-ca", caPub, caPriv, time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour*2), nil, nil, nil)
|
ca := cert.NebulaCertificate{
|
||||||
b, _ := ca.MarshalPEM()
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "test-ca",
|
||||||
|
NotBefore: time.Now().Add(time.Hour * -1),
|
||||||
|
NotAfter: time.Now().Add(time.Hour * 2),
|
||||||
|
PublicKey: caPub,
|
||||||
|
IsCA: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
ca.Sign(cert.Curve_CURVE25519, caPriv)
|
||||||
|
b, _ := ca.MarshalToPEM()
|
||||||
caFile.Truncate(0)
|
caFile.Truncate(0)
|
||||||
caFile.Seek(0, 0)
|
caFile.Seek(0, 0)
|
||||||
caFile.Write(b)
|
caFile.Write(b)
|
||||||
|
|
||||||
// no crt at path
|
// no crt at path
|
||||||
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.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.EqualError(t, err, "unable to read crt: open does_not_exist: "+NoSuchFileError)
|
assert.EqualError(t, err, "unable to read crt; open does_not_exist: "+NoSuchFileError)
|
||||||
|
|
||||||
// invalid crt at path
|
// invalid crt at path
|
||||||
ob.Reset()
|
ob.Reset()
|
||||||
eb.Reset()
|
eb.Reset()
|
||||||
certFile, err := os.CreateTemp("", "verify-cert")
|
certFile, err := ioutil.TempFile("", "verify-cert")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
defer os.Remove(certFile.Name())
|
defer os.Remove(certFile.Name())
|
||||||
|
|
||||||
certFile.WriteString("-----BEGIN NOPE-----")
|
certFile.WriteString("-----BEGIN NOPE-----")
|
||||||
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.EqualError(t, err, "error while parsing crt: input did not contain a valid PEM encoded block")
|
assert.EqualError(t, err, "error while parsing crt: input did not contain a valid PEM encoded block")
|
||||||
|
|
||||||
// unverifiable cert at path
|
// unverifiable cert at path
|
||||||
crt, _ := NewTestCert(ca, caPriv, "test-cert", time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour), nil, nil, nil)
|
_, badPriv, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
// Slightly evil hack to modify the certificate after it was sealed to generate an invalid signature
|
certPub, _ := x25519Keypair()
|
||||||
pub := crt.PublicKey()
|
signer, _ := ca.Sha256Sum()
|
||||||
for i, _ := range pub {
|
crt := cert.NebulaCertificate{
|
||||||
pub[i] = 0
|
Details: cert.NebulaCertificateDetails{
|
||||||
|
Name: "test-cert",
|
||||||
|
NotBefore: time.Now().Add(time.Hour * -1),
|
||||||
|
NotAfter: time.Now().Add(time.Hour),
|
||||||
|
PublicKey: certPub,
|
||||||
|
IsCA: false,
|
||||||
|
Issuer: signer,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
b, _ = crt.MarshalPEM()
|
|
||||||
|
crt.Sign(cert.Curve_CURVE25519, badPriv)
|
||||||
|
b, _ = crt.MarshalToPEM()
|
||||||
certFile.Truncate(0)
|
certFile.Truncate(0)
|
||||||
certFile.Seek(0, 0)
|
certFile.Seek(0, 0)
|
||||||
certFile.Write(b)
|
certFile.Write(b)
|
||||||
|
|
||||||
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.ErrorIs(t, err, cert.ErrSignatureMismatch)
|
assert.EqualError(t, err, "certificate signature did not match")
|
||||||
|
|
||||||
// verified cert at path
|
// verified cert at path
|
||||||
crt, _ = NewTestCert(ca, caPriv, "test-cert", time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour), nil, nil, nil)
|
crt.Sign(cert.Curve_CURVE25519, caPriv)
|
||||||
b, _ = crt.MarshalPEM()
|
b, _ = crt.MarshalToPEM()
|
||||||
certFile.Truncate(0)
|
certFile.Truncate(0)
|
||||||
certFile.Seek(0, 0)
|
certFile.Seek(0, 0)
|
||||||
certFile.Write(b)
|
certFile.Write(b)
|
||||||
|
|
||||||
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
err = verify([]string{"-ca", caFile.Name(), "-crt", certFile.Name()}, ob, eb)
|
||||||
assert.Empty(t, ob.String())
|
assert.Equal(t, "", ob.String())
|
||||||
assert.Empty(t, eb.String())
|
assert.Equal(t, "", eb.String())
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,15 +3,8 @@
|
|||||||
|
|
||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import "github.com/sirupsen/logrus"
|
||||||
"log/slog"
|
|
||||||
"os"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
func HookLogger(l *logrus.Logger) {
|
||||||
)
|
// Do nothing, let the logs flow to stdout/stderr
|
||||||
|
|
||||||
// newPlatformLogger returns a *slog.Logger that writes to stdout. Non-Windows
|
|
||||||
// platforms have no special sink to integrate with.
|
|
||||||
func newPlatformLogger() *slog.Logger {
|
|
||||||
return logging.NewLogger(os.Stdout)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,86 +1,54 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"fmt"
|
||||||
"log/slog"
|
"io/ioutil"
|
||||||
"strings"
|
"os"
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
"github.com/kardianos/service"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
// newPlatformLogger returns a *slog.Logger that routes every log record
|
// HookLogger routes the logrus logs through the service logger so that they end up in the Windows Event Viewer
|
||||||
// through the Windows service logger so records end up in the Windows
|
// logrus output will be discarded
|
||||||
// Event Log. All the heavy lifting (level management, format swap,
|
func HookLogger(l *logrus.Logger) {
|
||||||
// timestamp toggle, WithAttrs/WithGroup) comes from logging.NewHandler;
|
l.AddHook(newLogHook(logger))
|
||||||
// this file only contributes:
|
l.SetOutput(ioutil.Discard)
|
||||||
//
|
|
||||||
// - an io.Writer that forwards each formatted line to the service
|
|
||||||
// logger at the current record's Event Log severity, and
|
|
||||||
// - a thin severityTag that embeds *logging.Handler and overrides
|
|
||||||
// only Handle / WithAttrs / WithGroup, so Event Viewer's severity
|
|
||||||
// column and severity-based filters keep working the way they did
|
|
||||||
// before the slog migration.
|
|
||||||
//
|
|
||||||
// Format (text vs json) is carried by the embedded *logging.Handler, so
|
|
||||||
// logging.format: json in config still produces JSON lines in Event
|
|
||||||
// Viewer, same as the pre-slog logrus setup.
|
|
||||||
func newPlatformLogger() *slog.Logger {
|
|
||||||
w := &eventLogWriter{}
|
|
||||||
return slog.New(&severityTag{Handler: logging.NewHandler(w), w: w})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// eventLogWriter forwards slog-formatted lines to the Windows service
|
type logHook struct {
|
||||||
// logger at the severity most recently stashed by severityTag.Handle.
|
sl service.Logger
|
||||||
// The mutex serializes the stash + inner.Handle + Write cycle per record
|
|
||||||
// across all concurrent goroutines; slog's builtin text/json handlers
|
|
||||||
// each hold their own mutex around Write, but that only protects the
|
|
||||||
// Write call itself, not our stash-then-handle sequence.
|
|
||||||
type eventLogWriter struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
level slog.Level
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *eventLogWriter) Write(p []byte) (int, error) {
|
func newLogHook(sl service.Logger) *logHook {
|
||||||
line := strings.TrimRight(string(p), "\n")
|
return &logHook{sl: sl}
|
||||||
switch {
|
}
|
||||||
case w.level >= slog.LevelError:
|
|
||||||
return len(p), logger.Error(line)
|
func (h *logHook) Fire(entry *logrus.Entry) error {
|
||||||
case w.level >= slog.LevelWarn:
|
line, err := entry.String()
|
||||||
return len(p), logger.Warning(line)
|
if err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Unable to read entry, %v", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch entry.Level {
|
||||||
|
case logrus.PanicLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.FatalLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.ErrorLevel:
|
||||||
|
return h.sl.Error(line)
|
||||||
|
case logrus.WarnLevel:
|
||||||
|
return h.sl.Warning(line)
|
||||||
|
case logrus.InfoLevel:
|
||||||
|
return h.sl.Info(line)
|
||||||
|
case logrus.DebugLevel:
|
||||||
|
return h.sl.Info(line)
|
||||||
default:
|
default:
|
||||||
return len(p), logger.Info(line)
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// severityTag embeds *logging.Handler to pick up everything it does for
|
func (h *logHook) Levels() []logrus.Level {
|
||||||
// free (Enabled, SetLevel, GetLevel, SetFormat, GetFormat,
|
return logrus.AllLevels
|
||||||
// SetDisableTimestamp) and overrides only Handle / WithAttrs / WithGroup
|
|
||||||
// so each record's slog.Level is stashed on the writer before formatting
|
|
||||||
// and so derived handlers stay wrapped as severityTag rather than
|
|
||||||
// downgrading to bare *logging.Handler.
|
|
||||||
type severityTag struct {
|
|
||||||
*logging.Handler
|
|
||||||
w *eventLogWriter
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) Handle(ctx context.Context, r slog.Record) error {
|
|
||||||
s.w.mu.Lock()
|
|
||||||
defer s.w.mu.Unlock()
|
|
||||||
s.w.level = r.Level
|
|
||||||
return s.Handler.Handle(ctx, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) WithAttrs(attrs []slog.Attr) slog.Handler {
|
|
||||||
if len(attrs) == 0 {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &severityTag{Handler: s.Handler.WithAttrs(attrs).(*logging.Handler), w: s.w}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *severityTag) WithGroup(name string) slog.Handler {
|
|
||||||
if name == "" {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &severityTag{Handler: s.Handler.WithGroup(name).(*logging.Handler), w: s.w}
|
|
||||||
}
|
}
|
||||||
|
|||||||
+15
-47
@@ -4,12 +4,10 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"runtime/debug"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -20,17 +18,6 @@ import (
|
|||||||
// at compile-time.
|
// at compile-time.
|
||||||
var Build string
|
var Build string
|
||||||
|
|
||||||
func init() {
|
|
||||||
if Build == "" {
|
|
||||||
info, ok := debug.ReadBuildInfo()
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
Build = strings.TrimPrefix(info.Main.Version, "v")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
serviceFlag := flag.String("service", "", "Control the system service.")
|
serviceFlag := flag.String("service", "", "Control the system service.")
|
||||||
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
|
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
|
||||||
@@ -50,14 +37,9 @@ func main() {
|
|||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
|
||||||
|
|
||||||
if *serviceFlag != "" {
|
if *serviceFlag != "" {
|
||||||
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
doService(configPath, configTest, Build, serviceFlag)
|
||||||
l.Error("Service command failed", "error", err)
|
os.Exit(1)
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
@@ -66,6 +48,9 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
l := logrus.New()
|
||||||
|
l.Out = os.Stdout
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*configPath)
|
err := c.Load(*configPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -73,37 +58,20 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
fmt.Printf("failed to apply logging config: %s", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
switch v := err.(type) {
|
||||||
|
case util.ContextualError:
|
||||||
|
v.Log(l)
|
||||||
|
os.Exit(1)
|
||||||
|
case error:
|
||||||
|
l.WithError(err).Error("Failed to start")
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
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.Error("Nebula stopped due to fatal error", "error", err)
|
|
||||||
os.Exit(2)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.Info("Goodbye")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
|
|||||||
@@ -7,9 +7,9 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var logger service.Logger
|
var logger service.Logger
|
||||||
@@ -25,7 +25,8 @@ func (p *program) Start(s service.Service) error {
|
|||||||
// Start should not block.
|
// Start should not block.
|
||||||
logger.Info("Nebula service starting.")
|
logger.Info("Nebula service starting.")
|
||||||
|
|
||||||
l := newPlatformLogger()
|
l := logrus.New()
|
||||||
|
HookLogger(l)
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*p.configPath)
|
err := c.Load(*p.configPath)
|
||||||
@@ -33,15 +34,6 @@ func (p *program) Start(s service.Service) error {
|
|||||||
return fmt.Errorf("failed to load config: %s", err)
|
return fmt.Errorf("failed to load config: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
return fmt.Errorf("failed to apply logging config: %s", err)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -65,11 +57,11 @@ func fileExists(filename string) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) {
|
||||||
if *configPath == "" {
|
if *configPath == "" {
|
||||||
ex, err := os.Executable()
|
ex, err := os.Executable()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
panic(err)
|
||||||
}
|
}
|
||||||
*configPath = filepath.Dir(ex) + "/config.yaml"
|
*configPath = filepath.Dir(ex) + "/config.yaml"
|
||||||
if !fileExists(*configPath) {
|
if !fileExists(*configPath) {
|
||||||
@@ -93,16 +85,16 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
// Here are what the different loggers are doing:
|
// Here are what the different loggers are doing:
|
||||||
// - `log` is the standard go log utility, meant to be used while the process is still attached to stdout/stderr
|
// - `log` is the standard go log utility, meant to be used while the process is still attached to stdout/stderr
|
||||||
// - `logger` is the service log utility that may be attached to a special place depending on OS (Windows will have it attached to the event log)
|
// - `logger` is the service log utility that may be attached to a special place depending on OS (Windows will have it attached to the event log)
|
||||||
// - in program.Start we build a *slog.Logger via newPlatformLogger; on non-Windows that is a stdout-backed slog logger, on Windows it routes records through the service logger
|
// - above, in `Run` we create a `logrus.Logger` which is what nebula expects to use
|
||||||
s, err := service.New(prg, svcConfig)
|
s, err := service.New(prg, svcConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
log.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
errs := make(chan error, 5)
|
errs := make(chan error, 5)
|
||||||
logger, err = s.Logger(errs)
|
logger, err = s.Logger(errs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
log.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
@@ -117,16 +109,18 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
|||||||
|
|
||||||
switch *serviceFlag {
|
switch *serviceFlag {
|
||||||
case "run":
|
case "run":
|
||||||
if err := s.Run(); err != nil {
|
err = s.Run()
|
||||||
|
if err != nil {
|
||||||
// Route any errors to the system logger
|
// Route any errors to the system logger
|
||||||
logger.Error(err)
|
logger.Error(err)
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
if err := service.Control(s, *serviceFlag); err != nil {
|
err := service.Control(s, *serviceFlag)
|
||||||
|
if err != nil {
|
||||||
log.Printf("Valid actions: %q\n", service.ControlAction)
|
log.Printf("Valid actions: %q\n", service.ControlAction)
|
||||||
return err
|
log.Fatal(err)
|
||||||
}
|
}
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
+12
-42
@@ -4,12 +4,10 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"runtime/debug"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula"
|
"github.com/slackhq/nebula"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -20,17 +18,6 @@ import (
|
|||||||
// at compile-time.
|
// at compile-time.
|
||||||
var Build string
|
var Build string
|
||||||
|
|
||||||
func init() {
|
|
||||||
if Build == "" {
|
|
||||||
info, ok := debug.ReadBuildInfo()
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
Build = strings.TrimPrefix(info.Main.Version, "v")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
|
configPath := flag.String("config", "", "Path to either a file or directory to load configuration from")
|
||||||
configTest := flag.Bool("test", false, "Test the config and print the end result. Non zero exit indicates a faulty config")
|
configTest := flag.Bool("test", false, "Test the config and print the end result. Non zero exit indicates a faulty config")
|
||||||
@@ -55,7 +42,8 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
l := logging.NewLogger(os.Stdout)
|
l := logrus.New()
|
||||||
|
l.Out = os.Stdout
|
||||||
|
|
||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
err := c.Load(*configPath)
|
err := c.Load(*configPath)
|
||||||
@@ -64,38 +52,20 @@ func main() {
|
|||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
fmt.Printf("failed to apply logging config: %s", err)
|
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := logging.ApplyConfig(l, c); err != nil {
|
|
||||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||||
if err != nil {
|
|
||||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
switch v := err.(type) {
|
||||||
|
case util.ContextualError:
|
||||||
|
v.Log(l)
|
||||||
|
os.Exit(1)
|
||||||
|
case error:
|
||||||
|
l.WithError(err).Error("Failed to start")
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
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()
|
|
||||||
notifyReady(l)
|
|
||||||
|
|
||||||
if err := wait(); err != nil {
|
|
||||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
|
||||||
os.Exit(2)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.Info("Goodbye")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
os.Exit(0)
|
os.Exit(0)
|
||||||
|
|||||||
@@ -1,41 +0,0 @@
|
|||||||
package main
|
|
||||||
|
|
||||||
import (
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
// SdNotifyReady tells systemd the service is ready and dependent services can now be started
|
|
||||||
// https://www.freedesktop.org/software/systemd/man/sd_notify.html
|
|
||||||
// https://www.freedesktop.org/software/systemd/man/systemd.service.html
|
|
||||||
const SdNotifyReady = "READY=1"
|
|
||||||
|
|
||||||
func notifyReady(l *slog.Logger) {
|
|
||||||
sockName := os.Getenv("NOTIFY_SOCKET")
|
|
||||||
if sockName == "" {
|
|
||||||
l.Debug("NOTIFY_SOCKET systemd env var not set, not sending ready signal")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
conn, err := net.DialTimeout("unixgram", sockName, time.Second)
|
|
||||||
if err != nil {
|
|
||||||
l.Error("failed to connect to systemd notification socket", "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
err = conn.SetWriteDeadline(time.Now().Add(time.Second))
|
|
||||||
if err != nil {
|
|
||||||
l.Error("failed to set the write deadline for the systemd notification socket", "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err = conn.Write([]byte(SdNotifyReady)); err != nil {
|
|
||||||
l.Error("failed to signal the systemd notification socket", "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
l.Debug("notified systemd the service is ready")
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
//go:build !linux
|
|
||||||
// +build !linux
|
|
||||||
|
|
||||||
package main
|
|
||||||
|
|
||||||
import "log/slog"
|
|
||||||
|
|
||||||
func notifyReady(_ *slog.Logger) {
|
|
||||||
// No init service to notify
|
|
||||||
}
|
|
||||||
+26
-64
@@ -4,8 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"io/ioutil"
|
||||||
"math"
|
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -16,23 +15,24 @@ import (
|
|||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"github.com/imdario/mergo"
|
||||||
"go.yaml.in/yaml/v3"
|
"github.com/sirupsen/logrus"
|
||||||
|
"gopkg.in/yaml.v2"
|
||||||
)
|
)
|
||||||
|
|
||||||
type C struct {
|
type C struct {
|
||||||
path string
|
path string
|
||||||
files []string
|
files []string
|
||||||
Settings map[string]any
|
Settings map[interface{}]interface{}
|
||||||
oldSettings map[string]any
|
oldSettings map[interface{}]interface{}
|
||||||
callbacks []func(*C)
|
callbacks []func(*C)
|
||||||
l *slog.Logger
|
l *logrus.Logger
|
||||||
reloadLock sync.Mutex
|
reloadLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewC(l *slog.Logger) *C {
|
func NewC(l *logrus.Logger) *C {
|
||||||
return &C{
|
return &C{
|
||||||
Settings: make(map[string]any),
|
Settings: make(map[interface{}]interface{}),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,8 +92,8 @@ func (c *C) HasChanged(k string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
nv any
|
nv interface{}
|
||||||
ov any
|
ov interface{}
|
||||||
)
|
)
|
||||||
|
|
||||||
if k == "" {
|
if k == "" {
|
||||||
@@ -107,18 +107,12 @@ func (c *C) HasChanged(k string) bool {
|
|||||||
|
|
||||||
newVals, err := yaml.Marshal(nv)
|
newVals, err := yaml.Marshal(nv)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error while marshaling new config",
|
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling new config")
|
||||||
"config_path", k,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
oldVals, err := yaml.Marshal(ov)
|
oldVals, err := yaml.Marshal(ov)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error while marshaling old config",
|
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling old config")
|
||||||
"config_path", k,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return string(newVals) != string(oldVals)
|
return string(newVals) != string(oldVals)
|
||||||
@@ -127,10 +121,6 @@ func (c *C) HasChanged(k string) bool {
|
|||||||
// CatchHUP will listen for the HUP signal in a go routine and reload all configs found in the
|
// CatchHUP will listen for the HUP signal in a go routine and reload all configs found in the
|
||||||
// original path provided to Load. The old settings are shallow copied for change detection after the reload.
|
// original path provided to Load. The old settings are shallow copied for change detection after the reload.
|
||||||
func (c *C) CatchHUP(ctx context.Context) {
|
func (c *C) CatchHUP(ctx context.Context) {
|
||||||
if c.path == "" {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
ch := make(chan os.Signal, 1)
|
ch := make(chan os.Signal, 1)
|
||||||
signal.Notify(ch, syscall.SIGHUP)
|
signal.Notify(ch, syscall.SIGHUP)
|
||||||
|
|
||||||
@@ -153,17 +143,14 @@ func (c *C) ReloadConfig() {
|
|||||||
c.reloadLock.Lock()
|
c.reloadLock.Lock()
|
||||||
defer c.reloadLock.Unlock()
|
defer c.reloadLock.Unlock()
|
||||||
|
|
||||||
c.oldSettings = make(map[string]any)
|
c.oldSettings = make(map[interface{}]interface{})
|
||||||
for k, v := range c.Settings {
|
for k, v := range c.Settings {
|
||||||
c.oldSettings[k] = v
|
c.oldSettings[k] = v
|
||||||
}
|
}
|
||||||
|
|
||||||
err := c.Load(c.path)
|
err := c.Load(c.path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.l.Error("Error occurred while reloading config",
|
c.l.WithField("config_path", c.path).WithError(err).Error("Error occurred while reloading config")
|
||||||
"config_path", c.path,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -176,7 +163,7 @@ func (c *C) ReloadConfigString(raw string) error {
|
|||||||
c.reloadLock.Lock()
|
c.reloadLock.Lock()
|
||||||
defer c.reloadLock.Unlock()
|
defer c.reloadLock.Unlock()
|
||||||
|
|
||||||
c.oldSettings = make(map[string]any)
|
c.oldSettings = make(map[interface{}]interface{})
|
||||||
for k, v := range c.Settings {
|
for k, v := range c.Settings {
|
||||||
c.oldSettings[k] = v
|
c.oldSettings[k] = v
|
||||||
}
|
}
|
||||||
@@ -210,7 +197,7 @@ func (c *C) GetStringSlice(k string, d []string) []string {
|
|||||||
return d
|
return d
|
||||||
}
|
}
|
||||||
|
|
||||||
rv, ok := r.([]any)
|
rv, ok := r.([]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return d
|
return d
|
||||||
}
|
}
|
||||||
@@ -224,13 +211,13 @@ func (c *C) GetStringSlice(k string, d []string) []string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetMap will get the map for k or return the default d if not found or invalid
|
// GetMap will get the map for k or return the default d if not found or invalid
|
||||||
func (c *C) GetMap(k string, d map[string]any) map[string]any {
|
func (c *C) GetMap(k string, d map[interface{}]interface{}) map[interface{}]interface{} {
|
||||||
r := c.Get(k)
|
r := c.Get(k)
|
||||||
if r == nil {
|
if r == nil {
|
||||||
return d
|
return d
|
||||||
}
|
}
|
||||||
|
|
||||||
v, ok := r.(map[string]any)
|
v, ok := r.(map[interface{}]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return d
|
return d
|
||||||
}
|
}
|
||||||
@@ -249,15 +236,6 @@ func (c *C) GetInt(k string, d int) int {
|
|||||||
return v
|
return v
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUint32 will get the uint32 for k or return the default d if not found or invalid
|
|
||||||
func (c *C) GetUint32(k string, d uint32) uint32 {
|
|
||||||
r := c.GetInt(k, int(d))
|
|
||||||
if r < 0 || uint64(r) > uint64(math.MaxUint32) {
|
|
||||||
return d
|
|
||||||
}
|
|
||||||
return uint32(r)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetBool will get the bool for k or return the default d if not found or invalid
|
// GetBool will get the bool for k or return the default d if not found or invalid
|
||||||
func (c *C) GetBool(k string, d bool) bool {
|
func (c *C) GetBool(k string, d bool) bool {
|
||||||
r := strings.ToLower(c.GetString(k, fmt.Sprintf("%v", d)))
|
r := strings.ToLower(c.GetString(k, fmt.Sprintf("%v", d)))
|
||||||
@@ -275,22 +253,6 @@ func (c *C) GetBool(k string, d bool) bool {
|
|||||||
return v
|
return v
|
||||||
}
|
}
|
||||||
|
|
||||||
func AsBool(v any) (value bool, ok bool) {
|
|
||||||
switch x := v.(type) {
|
|
||||||
case bool:
|
|
||||||
return x, true
|
|
||||||
case string:
|
|
||||||
switch x {
|
|
||||||
case "y", "yes":
|
|
||||||
return true, true
|
|
||||||
case "n", "no":
|
|
||||||
return false, true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return false, false
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetDuration will get the duration for k or return the default d if not found or invalid
|
// GetDuration will get the duration for k or return the default d if not found or invalid
|
||||||
func (c *C) GetDuration(k string, d time.Duration) time.Duration {
|
func (c *C) GetDuration(k string, d time.Duration) time.Duration {
|
||||||
r := c.GetString(k, "")
|
r := c.GetString(k, "")
|
||||||
@@ -301,7 +263,7 @@ func (c *C) GetDuration(k string, d time.Duration) time.Duration {
|
|||||||
return v
|
return v
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *C) Get(k string) any {
|
func (c *C) Get(k string) interface{} {
|
||||||
return c.get(k, c.Settings)
|
return c.get(k, c.Settings)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -309,10 +271,10 @@ func (c *C) IsSet(k string) bool {
|
|||||||
return c.get(k, c.Settings) != nil
|
return c.get(k, c.Settings) != nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *C) get(k string, v any) any {
|
func (c *C) get(k string, v interface{}) interface{} {
|
||||||
parts := strings.Split(k, ".")
|
parts := strings.Split(k, ".")
|
||||||
for _, p := range parts {
|
for _, p := range parts {
|
||||||
m, ok := v.(map[string]any)
|
m, ok := v.(map[interface{}]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -371,7 +333,7 @@ func (c *C) addFile(path string, direct bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *C) parseRaw(b []byte) error {
|
func (c *C) parseRaw(b []byte) error {
|
||||||
var m map[string]any
|
var m map[interface{}]interface{}
|
||||||
|
|
||||||
err := yaml.Unmarshal(b, &m)
|
err := yaml.Unmarshal(b, &m)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -383,15 +345,15 @@ func (c *C) parseRaw(b []byte) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *C) parse() error {
|
func (c *C) parse() error {
|
||||||
var m map[string]any
|
var m map[interface{}]interface{}
|
||||||
|
|
||||||
for _, path := range c.files {
|
for _, path := range c.files {
|
||||||
b, err := os.ReadFile(path)
|
b, err := ioutil.ReadFile(path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var nm map[string]any
|
var nm map[interface{}]interface{}
|
||||||
err = yaml.Unmarshal(b, &nm)
|
err = yaml.Unmarshal(b, &nm)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
+43
-39
@@ -1,55 +1,59 @@
|
|||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"io/ioutil"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"github.com/imdario/mergo"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"go.yaml.in/yaml/v3"
|
"gopkg.in/yaml.v2"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestConfig_Load(t *testing.T) {
|
func TestConfig_Load(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
dir, err := os.MkdirTemp("", "config-test")
|
dir, err := ioutil.TempDir("", "config-test")
|
||||||
// invalid yaml
|
// invalid yaml
|
||||||
c := NewC(l)
|
c := NewC(l)
|
||||||
os.WriteFile(filepath.Join(dir, "01.yaml"), []byte(" invalid yaml"), 0644)
|
ioutil.WriteFile(filepath.Join(dir, "01.yaml"), []byte(" invalid yaml"), 0644)
|
||||||
require.EqualError(t, c.Load(dir), "yaml: unmarshal errors:\n line 1: cannot unmarshal !!str `invalid...` into map[string]interface {}")
|
assert.EqualError(t, c.Load(dir), "yaml: unmarshal errors:\n line 1: cannot unmarshal !!str `invalid...` into map[interface {}]interface {}")
|
||||||
|
|
||||||
// simple multi config merge
|
// simple multi config merge
|
||||||
c = NewC(l)
|
c = NewC(l)
|
||||||
os.RemoveAll(dir)
|
os.RemoveAll(dir)
|
||||||
os.Mkdir(dir, 0755)
|
os.Mkdir(dir, 0755)
|
||||||
|
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
|
|
||||||
os.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: hi"), 0644)
|
ioutil.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: hi"), 0644)
|
||||||
os.WriteFile(filepath.Join(dir, "02.yml"), []byte("outer:\n inner: override\nnew: hi"), 0644)
|
ioutil.WriteFile(filepath.Join(dir, "02.yml"), []byte("outer:\n inner: override\nnew: hi"), 0644)
|
||||||
require.NoError(t, c.Load(dir))
|
assert.Nil(t, c.Load(dir))
|
||||||
expected := map[string]any{
|
expected := map[interface{}]interface{}{
|
||||||
"outer": map[string]any{
|
"outer": map[interface{}]interface{}{
|
||||||
"inner": "override",
|
"inner": "override",
|
||||||
},
|
},
|
||||||
"new": "hi",
|
"new": "hi",
|
||||||
}
|
}
|
||||||
assert.Equal(t, expected, c.Settings)
|
assert.Equal(t, expected, c.Settings)
|
||||||
|
|
||||||
|
//TODO: test symlinked file
|
||||||
|
//TODO: test symlinked directory
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConfig_Get(t *testing.T) {
|
func TestConfig_Get(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// test simple type
|
// test simple type
|
||||||
c := NewC(l)
|
c := NewC(l)
|
||||||
c.Settings["firewall"] = map[string]any{"outbound": "hi"}
|
c.Settings["firewall"] = map[interface{}]interface{}{"outbound": "hi"}
|
||||||
assert.Equal(t, "hi", c.Get("firewall.outbound"))
|
assert.Equal(t, "hi", c.Get("firewall.outbound"))
|
||||||
|
|
||||||
// test complex type
|
// test complex type
|
||||||
inner := []map[string]any{{"port": "1", "code": "2"}}
|
inner := []map[interface{}]interface{}{{"port": "1", "code": "2"}}
|
||||||
c.Settings["firewall"] = map[string]any{"outbound": inner}
|
c.Settings["firewall"] = map[interface{}]interface{}{"outbound": inner}
|
||||||
assert.EqualValues(t, inner, c.Get("firewall.outbound"))
|
assert.EqualValues(t, inner, c.Get("firewall.outbound"))
|
||||||
|
|
||||||
// test missing
|
// test missing
|
||||||
@@ -59,7 +63,7 @@ func TestConfig_Get(t *testing.T) {
|
|||||||
func TestConfig_GetStringSlice(t *testing.T) {
|
func TestConfig_GetStringSlice(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
c := NewC(l)
|
c := NewC(l)
|
||||||
c.Settings["slice"] = []any{"one", "two"}
|
c.Settings["slice"] = []interface{}{"one", "two"}
|
||||||
assert.Equal(t, []string{"one", "two"}, c.GetStringSlice("slice", []string{}))
|
assert.Equal(t, []string{"one", "two"}, c.GetStringSlice("slice", []string{}))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -67,28 +71,28 @@ func TestConfig_GetBool(t *testing.T) {
|
|||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
c := NewC(l)
|
c := NewC(l)
|
||||||
c.Settings["bool"] = true
|
c.Settings["bool"] = true
|
||||||
assert.True(t, c.GetBool("bool", false))
|
assert.Equal(t, true, c.GetBool("bool", false))
|
||||||
|
|
||||||
c.Settings["bool"] = "true"
|
c.Settings["bool"] = "true"
|
||||||
assert.True(t, c.GetBool("bool", false))
|
assert.Equal(t, true, c.GetBool("bool", false))
|
||||||
|
|
||||||
c.Settings["bool"] = false
|
c.Settings["bool"] = false
|
||||||
assert.False(t, c.GetBool("bool", true))
|
assert.Equal(t, false, c.GetBool("bool", true))
|
||||||
|
|
||||||
c.Settings["bool"] = "false"
|
c.Settings["bool"] = "false"
|
||||||
assert.False(t, c.GetBool("bool", true))
|
assert.Equal(t, false, c.GetBool("bool", true))
|
||||||
|
|
||||||
c.Settings["bool"] = "Y"
|
c.Settings["bool"] = "Y"
|
||||||
assert.True(t, c.GetBool("bool", false))
|
assert.Equal(t, true, c.GetBool("bool", false))
|
||||||
|
|
||||||
c.Settings["bool"] = "yEs"
|
c.Settings["bool"] = "yEs"
|
||||||
assert.True(t, c.GetBool("bool", false))
|
assert.Equal(t, true, c.GetBool("bool", false))
|
||||||
|
|
||||||
c.Settings["bool"] = "N"
|
c.Settings["bool"] = "N"
|
||||||
assert.False(t, c.GetBool("bool", true))
|
assert.Equal(t, false, c.GetBool("bool", true))
|
||||||
|
|
||||||
c.Settings["bool"] = "nO"
|
c.Settings["bool"] = "nO"
|
||||||
assert.False(t, c.GetBool("bool", true))
|
assert.Equal(t, false, c.GetBool("bool", true))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConfig_HasChanged(t *testing.T) {
|
func TestConfig_HasChanged(t *testing.T) {
|
||||||
@@ -101,14 +105,14 @@ func TestConfig_HasChanged(t *testing.T) {
|
|||||||
// Test key change
|
// Test key change
|
||||||
c = NewC(l)
|
c = NewC(l)
|
||||||
c.Settings["test"] = "hi"
|
c.Settings["test"] = "hi"
|
||||||
c.oldSettings = map[string]any{"test": "no"}
|
c.oldSettings = map[interface{}]interface{}{"test": "no"}
|
||||||
assert.True(t, c.HasChanged("test"))
|
assert.True(t, c.HasChanged("test"))
|
||||||
assert.True(t, c.HasChanged(""))
|
assert.True(t, c.HasChanged(""))
|
||||||
|
|
||||||
// No key change
|
// No key change
|
||||||
c = NewC(l)
|
c = NewC(l)
|
||||||
c.Settings["test"] = "hi"
|
c.Settings["test"] = "hi"
|
||||||
c.oldSettings = map[string]any{"test": "hi"}
|
c.oldSettings = map[interface{}]interface{}{"test": "hi"}
|
||||||
assert.False(t, c.HasChanged("test"))
|
assert.False(t, c.HasChanged("test"))
|
||||||
assert.False(t, c.HasChanged(""))
|
assert.False(t, c.HasChanged(""))
|
||||||
}
|
}
|
||||||
@@ -116,18 +120,18 @@ func TestConfig_HasChanged(t *testing.T) {
|
|||||||
func TestConfig_ReloadConfig(t *testing.T) {
|
func TestConfig_ReloadConfig(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
done := make(chan bool, 1)
|
done := make(chan bool, 1)
|
||||||
dir, err := os.MkdirTemp("", "config-test")
|
dir, err := ioutil.TempDir("", "config-test")
|
||||||
require.NoError(t, err)
|
assert.Nil(t, err)
|
||||||
os.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: hi"), 0644)
|
ioutil.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: hi"), 0644)
|
||||||
|
|
||||||
c := NewC(l)
|
c := NewC(l)
|
||||||
require.NoError(t, c.Load(dir))
|
assert.Nil(t, c.Load(dir))
|
||||||
|
|
||||||
assert.False(t, c.HasChanged("outer.inner"))
|
assert.False(t, c.HasChanged("outer.inner"))
|
||||||
assert.False(t, c.HasChanged("outer"))
|
assert.False(t, c.HasChanged("outer"))
|
||||||
assert.False(t, c.HasChanged(""))
|
assert.False(t, c.HasChanged(""))
|
||||||
|
|
||||||
os.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: ho"), 0644)
|
ioutil.WriteFile(filepath.Join(dir, "01.yaml"), []byte("outer:\n inner: ho"), 0644)
|
||||||
|
|
||||||
c.RegisterReloadCallback(func(c *C) {
|
c.RegisterReloadCallback(func(c *C) {
|
||||||
done <- true
|
done <- true
|
||||||
@@ -184,11 +188,11 @@ firewall:
|
|||||||
`),
|
`),
|
||||||
}
|
}
|
||||||
|
|
||||||
var m map[string]any
|
var m map[any]any
|
||||||
|
|
||||||
// merge the same way config.parse() merges
|
// merge the same way config.parse() merges
|
||||||
for _, b := range configs {
|
for _, b := range configs {
|
||||||
var nm map[string]any
|
var nm map[any]any
|
||||||
err := yaml.Unmarshal(b, &nm)
|
err := yaml.Unmarshal(b, &nm)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
@@ -205,15 +209,15 @@ firewall:
|
|||||||
t.Logf("Merged Config as YAML:\n%s", mYaml)
|
t.Logf("Merged Config as YAML:\n%s", mYaml)
|
||||||
|
|
||||||
// If a bug is present, some items might be replaced instead of merged like we expect
|
// If a bug is present, some items might be replaced instead of merged like we expect
|
||||||
expected := map[string]any{
|
expected := map[any]any{
|
||||||
"firewall": map[string]any{
|
"firewall": map[any]any{
|
||||||
"inbound": []any{
|
"inbound": []any{
|
||||||
map[string]any{"host": "any", "port": "any", "proto": "icmp"},
|
map[any]any{"host": "any", "port": "any", "proto": "icmp"},
|
||||||
map[string]any{"groups": []any{"server"}, "port": 443, "proto": "tcp"},
|
map[any]any{"groups": []any{"server"}, "port": 443, "proto": "tcp"},
|
||||||
map[string]any{"groups": []any{"webapp"}, "port": 443, "proto": "tcp"}},
|
map[any]any{"groups": []any{"webapp"}, "port": 443, "proto": "tcp"}},
|
||||||
"outbound": []any{
|
"outbound": []any{
|
||||||
map[string]any{"host": "any", "port": "any", "proto": "any"}}},
|
map[any]any{"host": "any", "port": "any", "proto": "any"}}},
|
||||||
"listen": map[string]any{
|
"listen": map[any]any{
|
||||||
"host": "0.0.0.0",
|
"host": "0.0.0.0",
|
||||||
"port": 4242,
|
"port": 4242,
|
||||||
},
|
},
|
||||||
|
|||||||
+257
-364
@@ -3,18 +3,15 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
type trafficDecision int
|
type trafficDecision int
|
||||||
@@ -26,131 +23,133 @@ const (
|
|||||||
swapPrimary trafficDecision = 3
|
swapPrimary trafficDecision = 3
|
||||||
migrateRelays trafficDecision = 4
|
migrateRelays trafficDecision = 4
|
||||||
tryRehandshake trafficDecision = 5
|
tryRehandshake trafficDecision = 5
|
||||||
sendTestPacket trafficDecision = 6
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type connectionManager struct {
|
type connectionManager struct {
|
||||||
|
in map[uint32]struct{}
|
||||||
|
inLock *sync.RWMutex
|
||||||
|
|
||||||
|
out map[uint32]struct{}
|
||||||
|
outLock *sync.RWMutex
|
||||||
|
|
||||||
// relayUsed holds which relay localIndexs are in use
|
// relayUsed holds which relay localIndexs are in use
|
||||||
relayUsed map[uint32]struct{}
|
relayUsed map[uint32]struct{}
|
||||||
relayUsedLock *sync.RWMutex
|
relayUsedLock *sync.RWMutex
|
||||||
|
|
||||||
hostMap *HostMap
|
hostMap *HostMap
|
||||||
trafficTimer *LockingTimerWheel[uint32]
|
trafficTimer *LockingTimerWheel[uint32]
|
||||||
intf *Interface
|
intf *Interface
|
||||||
punchy *Punchy
|
pendingDeletion map[uint32]struct{}
|
||||||
|
punchy *Punchy
|
||||||
// Configuration settings
|
|
||||||
checkInterval time.Duration
|
checkInterval time.Duration
|
||||||
pendingDeletionInterval time.Duration
|
pendingDeletionInterval time.Duration
|
||||||
inactivityTimeout atomic.Int64
|
metricsTxPunchy metrics.Counter
|
||||||
dropInactive atomic.Bool
|
|
||||||
|
|
||||||
metricsTxPunchy metrics.Counter
|
l *logrus.Logger
|
||||||
|
|
||||||
l *slog.Logger
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
func newConnectionManager(ctx context.Context, l *logrus.Logger, intf *Interface, checkInterval, pendingDeletionInterval time.Duration, punchy *Punchy) *connectionManager {
|
||||||
cm := &connectionManager{
|
var max time.Duration
|
||||||
hostMap: hm,
|
if checkInterval < pendingDeletionInterval {
|
||||||
l: l,
|
max = pendingDeletionInterval
|
||||||
punchy: p,
|
} else {
|
||||||
relayUsed: make(map[uint32]struct{}),
|
max = checkInterval
|
||||||
relayUsedLock: &sync.RWMutex{},
|
|
||||||
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.reload(c, true)
|
nc := &connectionManager{
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
hostMap: intf.hostMap,
|
||||||
cm.reload(c, false)
|
in: make(map[uint32]struct{}),
|
||||||
})
|
inLock: &sync.RWMutex{},
|
||||||
|
out: make(map[uint32]struct{}),
|
||||||
return cm
|
outLock: &sync.RWMutex{},
|
||||||
}
|
relayUsed: make(map[uint32]struct{}),
|
||||||
|
relayUsedLock: &sync.RWMutex{},
|
||||||
func (cm *connectionManager) reload(c *config.C, initial bool) {
|
trafficTimer: NewLockingTimerWheel[uint32](time.Millisecond*500, max),
|
||||||
if initial {
|
intf: intf,
|
||||||
cm.checkInterval = time.Duration(c.GetInt("timers.connection_alive_interval", 5)) * time.Second
|
pendingDeletion: make(map[uint32]struct{}),
|
||||||
cm.pendingDeletionInterval = time.Duration(c.GetInt("timers.pending_deletion_interval", 10)) * time.Second
|
checkInterval: checkInterval,
|
||||||
|
pendingDeletionInterval: pendingDeletionInterval,
|
||||||
// We want at least a minimum resolution of 500ms per tick so that we can hit these intervals
|
punchy: punchy,
|
||||||
// pretty close to their configured duration.
|
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||||
// The inactivity duration is checked each time a hostinfo ticks through so we don't need the wheel to contain it.
|
l: l,
|
||||||
minDuration := min(time.Millisecond*500, cm.checkInterval, cm.pendingDeletionInterval)
|
|
||||||
maxDuration := max(cm.checkInterval, cm.pendingDeletionInterval)
|
|
||||||
cm.trafficTimer = NewLockingTimerWheel[uint32](minDuration, maxDuration)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if initial || c.HasChanged("tunnels.inactivity_timeout") {
|
nc.Start(ctx)
|
||||||
old := cm.getInactivityTimeout()
|
return nc
|
||||||
cm.inactivityTimeout.Store((int64)(c.GetDuration("tunnels.inactivity_timeout", 10*time.Minute)))
|
|
||||||
if !initial {
|
|
||||||
cm.l.Info("Inactivity timeout has changed",
|
|
||||||
"oldDuration", old,
|
|
||||||
"newDuration", cm.getInactivityTimeout(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if initial || c.HasChanged("tunnels.drop_inactive") {
|
|
||||||
old := cm.dropInactive.Load()
|
|
||||||
cm.dropInactive.Store(c.GetBool("tunnels.drop_inactive", false))
|
|
||||||
if !initial {
|
|
||||||
cm.l.Info("Drop inactive setting has changed",
|
|
||||||
"oldBool", old,
|
|
||||||
"newBool", cm.dropInactive.Load(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) getInactivityTimeout() time.Duration {
|
func (n *connectionManager) In(localIndex uint32) {
|
||||||
return (time.Duration)(cm.inactivityTimeout.Load())
|
n.inLock.RLock()
|
||||||
}
|
|
||||||
|
|
||||||
func (cm *connectionManager) In(h *HostInfo) {
|
|
||||||
h.in.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (cm *connectionManager) Out(h *HostInfo) {
|
|
||||||
h.out.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (cm *connectionManager) RelayUsed(localIndex uint32) {
|
|
||||||
cm.relayUsedLock.RLock()
|
|
||||||
// If this already exists, return
|
// If this already exists, return
|
||||||
if _, ok := cm.relayUsed[localIndex]; ok {
|
if _, ok := n.in[localIndex]; ok {
|
||||||
cm.relayUsedLock.RUnlock()
|
n.inLock.RUnlock()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
cm.relayUsedLock.RUnlock()
|
n.inLock.RUnlock()
|
||||||
cm.relayUsedLock.Lock()
|
n.inLock.Lock()
|
||||||
cm.relayUsed[localIndex] = struct{}{}
|
n.in[localIndex] = struct{}{}
|
||||||
cm.relayUsedLock.Unlock()
|
n.inLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *connectionManager) Out(localIndex uint32) {
|
||||||
|
n.outLock.RLock()
|
||||||
|
// If this already exists, return
|
||||||
|
if _, ok := n.out[localIndex]; ok {
|
||||||
|
n.outLock.RUnlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
n.outLock.RUnlock()
|
||||||
|
n.outLock.Lock()
|
||||||
|
n.out[localIndex] = struct{}{}
|
||||||
|
n.outLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *connectionManager) RelayUsed(localIndex uint32) {
|
||||||
|
n.relayUsedLock.RLock()
|
||||||
|
// If this already exists, return
|
||||||
|
if _, ok := n.relayUsed[localIndex]; ok {
|
||||||
|
n.relayUsedLock.RUnlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
n.relayUsedLock.RUnlock()
|
||||||
|
n.relayUsedLock.Lock()
|
||||||
|
n.relayUsed[localIndex] = struct{}{}
|
||||||
|
n.relayUsedLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
|
||||||
// resets the state for this local index
|
// resets the state for this local index
|
||||||
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
|
func (n *connectionManager) getAndResetTrafficCheck(localIndex uint32) (bool, bool) {
|
||||||
in := h.in.Swap(false)
|
n.inLock.Lock()
|
||||||
out := h.out.Swap(false)
|
n.outLock.Lock()
|
||||||
if in || out {
|
_, in := n.in[localIndex]
|
||||||
h.lastUsed = now
|
_, out := n.out[localIndex]
|
||||||
}
|
delete(n.in, localIndex)
|
||||||
|
delete(n.out, localIndex)
|
||||||
|
n.inLock.Unlock()
|
||||||
|
n.outLock.Unlock()
|
||||||
return in, out
|
return in, out
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddTrafficWatch must be called for every new HostInfo.
|
func (n *connectionManager) AddTrafficWatch(localIndex uint32) {
|
||||||
// We will continue to monitor the HostInfo until the tunnel is dropped.
|
// Use a write lock directly because it should be incredibly rare that we are ever already tracking this index
|
||||||
func (cm *connectionManager) AddTrafficWatch(h *HostInfo) {
|
n.outLock.Lock()
|
||||||
if h.out.Swap(true) == false {
|
if _, ok := n.out[localIndex]; ok {
|
||||||
cm.trafficTimer.Add(h.localIndexId, cm.checkInterval)
|
n.outLock.Unlock()
|
||||||
cm.intf.pmtudManager.OnTunnelUp(h)
|
return
|
||||||
}
|
}
|
||||||
|
n.out[localIndex] = struct{}{}
|
||||||
|
n.trafficTimer.Add(localIndex, n.checkInterval)
|
||||||
|
n.outLock.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) Start(ctx context.Context) {
|
func (n *connectionManager) Start(ctx context.Context) {
|
||||||
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
|
go n.Run(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n *connectionManager) Run(ctx context.Context) {
|
||||||
|
//TODO: this tick should be based on the min wheel tick? Check firewall
|
||||||
|
clockSource := time.NewTicker(500 * time.Millisecond)
|
||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -163,123 +162,107 @@ func (cm *connectionManager) Start(ctx context.Context) {
|
|||||||
return
|
return
|
||||||
|
|
||||||
case now := <-clockSource.C:
|
case now := <-clockSource.C:
|
||||||
cm.trafficTimer.Advance(now)
|
n.trafficTimer.Advance(now)
|
||||||
for {
|
for {
|
||||||
localIndex, has := cm.trafficTimer.Purge()
|
localIndex, has := n.trafficTimer.Purge()
|
||||||
if !has {
|
if !has {
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.doTrafficCheck(localIndex, p, nb, out, now)
|
n.doTrafficCheck(localIndex, p, nb, out, now)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) doTrafficCheck(localIndex uint32, p, nb, out []byte, now time.Time) {
|
func (n *connectionManager) doTrafficCheck(localIndex uint32, p, nb, out []byte, now time.Time) {
|
||||||
decision, hostinfo, primary := cm.makeTrafficDecision(localIndex, now)
|
decision, hostinfo, primary := n.makeTrafficDecision(localIndex, p, nb, out, now)
|
||||||
|
|
||||||
switch decision {
|
switch decision {
|
||||||
case deleteTunnel:
|
case deleteTunnel:
|
||||||
cm.intf.pmtudManager.OnTunnelDown(hostinfo)
|
if n.hostMap.DeleteHostInfo(hostinfo) {
|
||||||
if cm.hostMap.DeleteHostInfo(hostinfo) {
|
|
||||||
// Only clearing the lighthouse cache if this is the last hostinfo for this vpn ip in the hostmap
|
// Only clearing the lighthouse cache if this is the last hostinfo for this vpn ip in the hostmap
|
||||||
cm.intf.lightHouse.DeleteVpnAddrs(hostinfo.vpnAddrs)
|
n.intf.lightHouse.DeleteVpnIp(hostinfo.vpnIp)
|
||||||
}
|
}
|
||||||
|
|
||||||
case closeTunnel:
|
case closeTunnel:
|
||||||
cm.intf.sendCloseTunnel(hostinfo)
|
n.intf.sendCloseTunnel(hostinfo)
|
||||||
cm.intf.closeTunnel(hostinfo)
|
n.intf.closeTunnel(hostinfo)
|
||||||
|
|
||||||
case swapPrimary:
|
case swapPrimary:
|
||||||
cm.swapPrimary(hostinfo, primary)
|
n.swapPrimary(hostinfo, primary)
|
||||||
|
|
||||||
case migrateRelays:
|
case migrateRelays:
|
||||||
cm.migrateRelayUsed(hostinfo, primary)
|
n.migrateRelayUsed(hostinfo, primary)
|
||||||
|
|
||||||
case tryRehandshake:
|
case tryRehandshake:
|
||||||
cm.tryRehandshake(hostinfo)
|
n.tryRehandshake(hostinfo)
|
||||||
|
|
||||||
case sendTestPacket:
|
|
||||||
// Defer to pmtud if it has a confirmed PMTU > floor for this peer:
|
|
||||||
// the probe at the confirmed size verifies both liveness AND that
|
|
||||||
// the discovered PMTU still fits, so we don't burn a separate test
|
|
||||||
// packet on top of it. If pmtud declines (disabled, peer unsupported,
|
|
||||||
// or no confirmed size yet) we fall back to the regular test.
|
|
||||||
if !cm.intf.pmtudManager.MaybeProbeAsTest(hostinfo) {
|
|
||||||
cm.intf.SendMessageToHostInfo(header.Test, header.TestRequest, hostinfo, p, nb, out)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.resetRelayTrafficCheck(hostinfo)
|
n.resetRelayTrafficCheck(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) resetRelayTrafficCheck(hostinfo *HostInfo) {
|
func (n *connectionManager) resetRelayTrafficCheck(hostinfo *HostInfo) {
|
||||||
if hostinfo != nil {
|
if hostinfo != nil {
|
||||||
cm.relayUsedLock.Lock()
|
n.relayUsedLock.Lock()
|
||||||
defer cm.relayUsedLock.Unlock()
|
defer n.relayUsedLock.Unlock()
|
||||||
// No need to migrate any relays, delete usage info now.
|
// No need to migrate any relays, delete usage info now.
|
||||||
for _, idx := range hostinfo.relayState.CopyRelayForIdxs() {
|
for _, idx := range hostinfo.relayState.CopyRelayForIdxs() {
|
||||||
delete(cm.relayUsed, idx)
|
delete(n.relayUsed, idx)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo) {
|
func (n *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo) {
|
||||||
relayFor := oldhostinfo.relayState.CopyAllRelayFor()
|
relayFor := oldhostinfo.relayState.CopyAllRelayFor()
|
||||||
|
|
||||||
for _, r := range relayFor {
|
for _, r := range relayFor {
|
||||||
existing, ok := newhostinfo.relayState.QueryRelayForByIp(r.PeerAddr)
|
existing, ok := newhostinfo.relayState.QueryRelayForByIp(r.PeerIp)
|
||||||
|
|
||||||
var index uint32
|
var index uint32
|
||||||
var relayFrom netip.Addr
|
var relayFrom iputil.VpnIp
|
||||||
var relayTo netip.Addr
|
var relayTo iputil.VpnIp
|
||||||
switch {
|
switch {
|
||||||
case ok:
|
case ok && existing.State == Established:
|
||||||
switch existing.State {
|
// This relay already exists in newhostinfo, then do nothing.
|
||||||
case Established, PeerRequested, Disestablished:
|
continue
|
||||||
// This relay already exists in newhostinfo, then do nothing.
|
case ok && existing.State == Requested:
|
||||||
continue
|
// The relay exists in a Requested state; re-send the request
|
||||||
case Requested:
|
index = existing.LocalIndex
|
||||||
// The relay exists in a Requested state; re-send the request
|
switch r.Type {
|
||||||
index = existing.LocalIndex
|
case TerminalType:
|
||||||
switch r.Type {
|
relayFrom = newhostinfo.vpnIp
|
||||||
case TerminalType:
|
relayTo = existing.PeerIp
|
||||||
relayFrom = cm.intf.myVpnAddrs[0]
|
case ForwardingType:
|
||||||
relayTo = existing.PeerAddr
|
relayFrom = existing.PeerIp
|
||||||
case ForwardingType:
|
relayTo = newhostinfo.vpnIp
|
||||||
relayFrom = existing.PeerAddr
|
default:
|
||||||
relayTo = newhostinfo.vpnAddrs[0]
|
// should never happen
|
||||||
default:
|
|
||||||
// should never happen
|
|
||||||
panic(fmt.Sprintf("Migrating unknown relay type: %v", r.Type))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
case !ok:
|
case !ok:
|
||||||
cm.relayUsedLock.RLock()
|
n.relayUsedLock.RLock()
|
||||||
if _, relayUsed := cm.relayUsed[r.LocalIndex]; !relayUsed {
|
if _, relayUsed := n.relayUsed[r.LocalIndex]; !relayUsed {
|
||||||
// The relay hasn't been used; don't migrate it.
|
// The relay hasn't been used; don't migrate it.
|
||||||
cm.relayUsedLock.RUnlock()
|
n.relayUsedLock.RUnlock()
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
cm.relayUsedLock.RUnlock()
|
n.relayUsedLock.RUnlock()
|
||||||
// The relay doesn't exist at all; create some relay state and send the request.
|
// The relay doesn't exist at all; create some relay state and send the request.
|
||||||
var err error
|
var err error
|
||||||
index, err = AddRelay(cm.l, newhostinfo, cm.hostMap, r.PeerAddr, nil, r.Type, Requested)
|
index, err = AddRelay(n.l, newhostinfo, n.hostMap, r.PeerIp, nil, r.Type, Requested)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cm.l.Error("failed to migrate relay to new hostinfo", "error", err)
|
n.l.WithError(err).Error("failed to migrate relay to new hostinfo")
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
switch r.Type {
|
switch r.Type {
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
relayFrom = cm.intf.myVpnAddrs[0]
|
relayFrom = newhostinfo.vpnIp
|
||||||
relayTo = r.PeerAddr
|
relayTo = r.PeerIp
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
relayFrom = r.PeerAddr
|
relayFrom = r.PeerIp
|
||||||
relayTo = newhostinfo.vpnAddrs[0]
|
relayTo = newhostinfo.vpnIp
|
||||||
default:
|
default:
|
||||||
// should never happen
|
// should never happen
|
||||||
panic(fmt.Sprintf("Migrating unknown relay type: %v", r.Type))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -287,86 +270,65 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
req := NebulaControl{
|
req := NebulaControl{
|
||||||
Type: NebulaControl_CreateRelayRequest,
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
InitiatorRelayIndex: index,
|
InitiatorRelayIndex: index,
|
||||||
|
RelayFromIp: uint32(relayFrom),
|
||||||
|
RelayToIp: uint32(relayTo),
|
||||||
}
|
}
|
||||||
|
|
||||||
switch newhostinfo.GetCert().Certificate.Version() {
|
|
||||||
case cert.Version1:
|
|
||||||
if !relayFrom.Is4() {
|
|
||||||
cm.l.Error("can not migrate v1 relay with a v6 network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if !relayTo.Is4() {
|
|
||||||
cm.l.Error("can not migrate v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
b := relayFrom.As4()
|
|
||||||
req.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
b = relayTo.As4()
|
|
||||||
req.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
case cert.Version2:
|
|
||||||
req.RelayFromAddr = netAddrToProtoAddr(relayFrom)
|
|
||||||
req.RelayToAddr = netAddrToProtoAddr(relayTo)
|
|
||||||
default:
|
|
||||||
newhostinfo.logger(cm.l).Error("Unknown certificate version found while attempting to migrate relay")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
msg, err := req.Marshal()
|
msg, err := req.Marshal()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
cm.l.Error("failed to marshal Control message to migrate relay", "error", err)
|
n.l.WithError(err).Error("failed to marshal Control message to migrate relay")
|
||||||
} else {
|
} else {
|
||||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
n.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
cm.l.Info("send CreateRelayRequest",
|
n.l.WithFields(logrus.Fields{
|
||||||
"relayFrom", req.RelayFromAddr,
|
"relayFrom": iputil.VpnIp(req.RelayFromIp),
|
||||||
"relayTo", req.RelayToAddr,
|
"relayTo": iputil.VpnIp(req.RelayToIp),
|
||||||
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
"initiatorRelayIndex": req.InitiatorRelayIndex,
|
||||||
"responderRelayIndex", req.ResponderRelayIndex,
|
"responderRelayIndex": req.ResponderRelayIndex,
|
||||||
"vpnAddrs", newhostinfo.vpnAddrs,
|
"vpnIp": newhostinfo.vpnIp}).
|
||||||
)
|
Info("send CreateRelayRequest")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Time) (trafficDecision, *HostInfo, *HostInfo) {
|
func (n *connectionManager) makeTrafficDecision(localIndex uint32, p, nb, out []byte, now time.Time) (trafficDecision, *HostInfo, *HostInfo) {
|
||||||
// Read lock the main hostmap to order decisions based on tunnels being the primary tunnel
|
n.hostMap.RLock()
|
||||||
cm.hostMap.RLock()
|
defer n.hostMap.RUnlock()
|
||||||
defer cm.hostMap.RUnlock()
|
|
||||||
|
|
||||||
hostinfo := cm.hostMap.Indexes[localIndex]
|
hostinfo := n.hostMap.Indexes[localIndex]
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
cm.l.Debug("Not found in hostmap", "localIndex", localIndex)
|
n.l.WithField("localIndex", localIndex).Debugf("Not found in hostmap")
|
||||||
|
delete(n.pendingDeletion, localIndex)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if cm.isInvalidCertificate(now, hostinfo) {
|
if n.isInvalidCertificate(now, hostinfo) {
|
||||||
|
delete(n.pendingDeletion, hostinfo.localIndexId)
|
||||||
return closeTunnel, hostinfo, nil
|
return closeTunnel, hostinfo, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
primary := cm.hostMap.Hosts[hostinfo.vpnAddrs[0]]
|
primary := n.hostMap.Hosts[hostinfo.vpnIp]
|
||||||
mainHostInfo := true
|
mainHostInfo := true
|
||||||
if primary != nil && primary != hostinfo {
|
if primary != nil && primary != hostinfo {
|
||||||
mainHostInfo = false
|
mainHostInfo = false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check for traffic on this hostinfo
|
// Check for traffic on this hostinfo
|
||||||
inTraffic, outTraffic := cm.getAndResetTrafficCheck(hostinfo, now)
|
inTraffic, outTraffic := n.getAndResetTrafficCheck(localIndex)
|
||||||
|
|
||||||
// A hostinfo is determined alive if there is incoming traffic
|
// A hostinfo is determined alive if there is incoming traffic
|
||||||
if inTraffic {
|
if inTraffic {
|
||||||
decision := doNothing
|
decision := doNothing
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if n.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(n.l).
|
||||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
WithField("tunnelCheck", m{"state": "alive", "method": "passive"}).
|
||||||
)
|
Debug("Tunnel status")
|
||||||
}
|
}
|
||||||
hostinfo.pendingDeletion.Store(false)
|
delete(n.pendingDeletion, hostinfo.localIndexId)
|
||||||
|
|
||||||
if mainHostInfo {
|
if mainHostInfo {
|
||||||
decision = tryRehandshake
|
decision = tryRehandshake
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
if cm.shouldSwapPrimary(hostinfo) {
|
if n.shouldSwapPrimary(hostinfo, primary) {
|
||||||
decision = swapPrimary
|
decision = swapPrimary
|
||||||
} else {
|
} else {
|
||||||
// migrate the relays to the primary, if in use.
|
// migrate the relays to the primary, if in use.
|
||||||
@@ -374,224 +336,155 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
n.trafficTimer.Add(hostinfo.localIndexId, n.checkInterval)
|
||||||
|
|
||||||
if !outTraffic {
|
if !outTraffic {
|
||||||
// Send a punch packet to keep the NAT state alive
|
// Send a punch packet to keep the NAT state alive
|
||||||
cm.sendPunch(hostinfo)
|
n.sendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
}
|
}
|
||||||
|
|
||||||
if hostinfo.pendingDeletion.Load() {
|
if _, ok := n.pendingDeletion[hostinfo.localIndexId]; ok {
|
||||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
hostinfo.logger(n.l).
|
||||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
WithField("tunnelCheck", m{"state": "dead", "method": "active"}).
|
||||||
)
|
Info("Tunnel status")
|
||||||
|
|
||||||
|
delete(n.pendingDeletion, hostinfo.localIndexId)
|
||||||
return deleteTunnel, hostinfo, nil
|
return deleteTunnel, hostinfo, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
decision := doNothing
|
|
||||||
if hostinfo != nil && hostinfo.ConnectionState != nil && mainHostInfo {
|
if hostinfo != nil && hostinfo.ConnectionState != nil && mainHostInfo {
|
||||||
if !outTraffic {
|
if !outTraffic {
|
||||||
inactiveFor, isInactive := cm.isInactive(hostinfo, now)
|
|
||||||
if isInactive {
|
|
||||||
// Tunnel is inactive, tear it down
|
|
||||||
hostinfo.logger(cm.l).Info("Dropping tunnel due to inactivity",
|
|
||||||
"inactiveDuration", inactiveFor,
|
|
||||||
"primary", mainHostInfo,
|
|
||||||
)
|
|
||||||
|
|
||||||
return closeTunnel, hostinfo, primary
|
|
||||||
}
|
|
||||||
|
|
||||||
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
||||||
// Just maintain NAT state if configured to do so.
|
// Just maintain NAT state if configured to do so.
|
||||||
cm.sendPunch(hostinfo)
|
n.sendPunch(hostinfo)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
n.trafficTimer.Add(hostinfo.localIndexId, n.checkInterval)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if cm.punchy.GetTargetEverything() {
|
if n.punchy.GetTargetEverything() {
|
||||||
// This is similar to the old punchy behavior with a slight optimization.
|
// This is similar to the old punchy behavior with a slight optimization.
|
||||||
// We aren't receiving traffic but we are sending it, punch on all known
|
// We aren't receiving traffic but we are sending it, punch on all known
|
||||||
// ips in case we need to re-prime NAT state
|
// ips in case we need to re-prime NAT state
|
||||||
cm.sendPunch(hostinfo)
|
n.sendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if n.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(n.l).
|
||||||
"tunnelCheck", m{"state": "testing", "method": "active"},
|
WithField("tunnelCheck", m{"state": "testing", "method": "active"}).
|
||||||
)
|
Debug("Tunnel status")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||||
decision = sendTestPacket
|
n.intf.SendMessageToHostInfo(header.Test, header.TestRequest, hostinfo, p, nb, out)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if n.l.Level >= logrus.DebugLevel {
|
||||||
hostinfo.logger(cm.l).Debug("Hostinfo sadness")
|
hostinfo.logger(n.l).Debugf("Hostinfo sadness")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo.pendingDeletion.Store(true)
|
n.pendingDeletion[hostinfo.localIndexId] = struct{}{}
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
|
n.trafficTimer.Add(hostinfo.localIndexId, n.pendingDeletionInterval)
|
||||||
return decision, hostinfo, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) isInactive(hostinfo *HostInfo, now time.Time) (time.Duration, bool) {
|
func (n *connectionManager) shouldSwapPrimary(current, primary *HostInfo) bool {
|
||||||
if cm.dropInactive.Load() == false {
|
|
||||||
// We aren't configured to drop inactive tunnels
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
|
|
||||||
inactiveDuration := now.Sub(hostinfo.lastUsed)
|
|
||||||
if inactiveDuration < cm.getInactivityTimeout() {
|
|
||||||
// It's not considered inactive
|
|
||||||
return inactiveDuration, false
|
|
||||||
}
|
|
||||||
|
|
||||||
// The tunnel is inactive
|
|
||||||
return inactiveDuration, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (cm *connectionManager) shouldSwapPrimary(current *HostInfo) bool {
|
|
||||||
// The primary tunnel is the most recent handshake to complete locally and should work entirely fine.
|
// The primary tunnel is the most recent handshake to complete locally and should work entirely fine.
|
||||||
// If we are here then we have multiple tunnels for a host pair and neither side believes the same tunnel is primary.
|
// If we are here then we have multiple tunnels for a host pair and neither side believes the same tunnel is primary.
|
||||||
// Let's sort this out.
|
// Let's sort this out.
|
||||||
|
|
||||||
// Only one side should swap because if both swap then we may never resolve to a single tunnel.
|
if current.vpnIp < n.intf.myVpnIp {
|
||||||
// vpn addr is static across all tunnels for this host pair so lets
|
// Only one side should flip primary because if both flip then we may never resolve to a single tunnel.
|
||||||
// use that to determine if we should consider swapping.
|
// vpn ip is static across all tunnels for this host pair so lets use that to determine who is flipping.
|
||||||
if current.vpnAddrs[0].Compare(cm.intf.myVpnAddrs[0]) < 0 {
|
// The remotes vpn ip is lower than mine. I will not flip.
|
||||||
// Their primary vpn addr is less than mine. Do not swap.
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
crt := cm.intf.pki.getCertState().getCertificate(current.ConnectionState.myCert.Version())
|
certState := n.intf.certState.Load()
|
||||||
if crt == nil {
|
return bytes.Equal(current.ConnectionState.certState.certificate.Signature, certState.certificate.Signature)
|
||||||
//my cert was reloaded away. We should definitely swap from this tunnel
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
// If this tunnel is using the latest certificate then we should swap it to primary for a bit and see if things
|
|
||||||
// settle down.
|
|
||||||
return bytes.Equal(current.ConnectionState.myCert.Signature(), crt.Signature())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) swapPrimary(current, primary *HostInfo) {
|
func (n *connectionManager) swapPrimary(current, primary *HostInfo) {
|
||||||
cm.hostMap.Lock()
|
n.hostMap.Lock()
|
||||||
// Make sure the primary is still the same after the write lock. This avoids a race with a rehandshake.
|
// Make sure the primary is still the same after the write lock. This avoids a race with a rehandshake.
|
||||||
if cm.hostMap.Hosts[current.vpnAddrs[0]] == primary {
|
if n.hostMap.Hosts[current.vpnIp] == primary {
|
||||||
cm.hostMap.unlockedMakePrimary(current)
|
n.hostMap.unlockedMakePrimary(current)
|
||||||
}
|
}
|
||||||
cm.hostMap.Unlock()
|
n.hostMap.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// isInvalidCertificate decides if we should destroy a tunnel.
|
// isInvalidCertificate will check if we should destroy a tunnel if pki.disconnect_invalid is true and
|
||||||
// returns true if pki.disconnect_invalid is true and the certificate is no longer valid.
|
// the certificate is no longer valid. Block listed certificates will skip the pki.disconnect_invalid
|
||||||
// Blocklisted certificates will skip the pki.disconnect_invalid check and return true.
|
// check and return true.
|
||||||
func (cm *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostInfo) bool {
|
func (n *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostInfo) bool {
|
||||||
remoteCert := hostinfo.GetCert()
|
remoteCert := hostinfo.GetCert()
|
||||||
if remoteCert == nil {
|
if remoteCert == nil {
|
||||||
return false //don't tear down tunnels for handshakes in progress
|
|
||||||
}
|
|
||||||
|
|
||||||
caPool := cm.intf.pki.GetCAPool()
|
|
||||||
err := caPool.VerifyCachedCertificate(now, remoteCert)
|
|
||||||
if err == nil {
|
|
||||||
return false //cert is still valid! yay!
|
|
||||||
} else if err == cert.ErrBlockListed { //avoiding errors.Is for speed
|
|
||||||
// Block listed certificates should always be disconnected
|
|
||||||
hostinfo.logger(cm.l).Info("Remote certificate is blocked, tearing down the tunnel",
|
|
||||||
"error", err,
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
)
|
|
||||||
return true
|
|
||||||
} else if cm.intf.disconnectInvalid.Load() {
|
|
||||||
hostinfo.logger(cm.l).Info("Remote certificate is no longer valid, tearing down the tunnel",
|
|
||||||
"error", err,
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
)
|
|
||||||
return true
|
|
||||||
} else {
|
|
||||||
//if we reach here, the cert is no longer valid, but we're configured to keep tunnels from now-invalid certs open
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
valid, err := remoteCert.VerifyWithCache(now, n.intf.caPool)
|
||||||
|
if valid {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if !n.intf.disconnectInvalid && err != cert.ErrBlockListed {
|
||||||
|
// Block listed certificates should always be disconnected
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
fingerprint, _ := remoteCert.Sha256Sum()
|
||||||
|
hostinfo.logger(n.l).WithError(err).
|
||||||
|
WithField("fingerprint", fingerprint).
|
||||||
|
Info("Remote certificate is no longer valid, tearing down the tunnel")
|
||||||
|
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) sendPunch(hostinfo *HostInfo) {
|
func (n *connectionManager) sendPunch(hostinfo *HostInfo) {
|
||||||
if !cm.punchy.GetPunch() {
|
if !n.punchy.GetPunch() {
|
||||||
// Punching is disabled
|
// Punching is disabled
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if cm.intf.lightHouse.IsAnyLighthouseAddr(hostinfo.vpnAddrs) {
|
if n.punchy.GetTargetEverything() {
|
||||||
// Do not punch to lighthouses, we assume our lighthouse update interval is good enough.
|
hostinfo.remotes.ForEach(n.hostMap.preferredRanges, func(addr *udp.Addr, preferred bool) {
|
||||||
// In the event the update interval is not sufficient to maintain NAT state then a publicly available lighthouse
|
n.metricsTxPunchy.Inc(1)
|
||||||
// would lose the ability to notify us and punchy.respond would become unreliable.
|
n.intf.outside.WriteTo([]byte{1}, addr)
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if cm.punchy.GetTargetEverything() {
|
|
||||||
hostinfo.remotes.ForEach(cm.hostMap.GetPreferredRanges(), func(addr netip.AddrPort, preferred bool) {
|
|
||||||
cm.metricsTxPunchy.Inc(1)
|
|
||||||
cm.intf.outside.WriteTo([]byte{1}, addr)
|
|
||||||
})
|
})
|
||||||
|
|
||||||
} else if hostinfo.remote.IsValid() {
|
} else if hostinfo.remote != nil {
|
||||||
cm.metricsTxPunchy.Inc(1)
|
n.metricsTxPunchy.Inc(1)
|
||||||
cm.intf.outside.WriteTo([]byte{1}, hostinfo.remote)
|
n.intf.outside.WriteTo([]byte{1}, hostinfo.remote)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
func (n *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||||
cs := cm.intf.pki.getCertState()
|
certState := n.intf.certState.Load()
|
||||||
curCrt := hostinfo.ConnectionState.myCert
|
if bytes.Equal(hostinfo.ConnectionState.certState.certificate.Signature, certState.certificate.Signature) {
|
||||||
curCrtVersion := curCrt.Version()
|
|
||||||
myCrt := cs.getCertificate(curCrtVersion)
|
|
||||||
if myCrt == nil {
|
|
||||||
cm.l.Info("Re-handshaking with remote",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"version", curCrtVersion,
|
|
||||||
"reason", "local certificate removed",
|
|
||||||
)
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
peerCrt := hostinfo.ConnectionState.peerCert
|
|
||||||
if peerCrt != nil && curCrtVersion < peerCrt.Certificate.Version() {
|
n.l.WithField("vpnIp", hostinfo.vpnIp).
|
||||||
// if our certificate version is less than theirs, and we have a matching version available, rehandshake?
|
WithField("reason", "local certificate is not current").
|
||||||
if cs.getCertificate(peerCrt.Certificate.Version()) != nil {
|
Info("Re-handshaking with remote")
|
||||||
cm.l.Info("Re-handshaking with remote",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
//TODO: this is copied from getOrHandshake to keep the extra checks out of the hot path, figure it out
|
||||||
"version", curCrtVersion,
|
newHostinfo := n.intf.handshakeManager.AddVpnIp(hostinfo.vpnIp, n.intf.initHostInfo)
|
||||||
"peerVersion", peerCrt.Certificate.Version(),
|
if !newHostinfo.HandshakeReady {
|
||||||
"reason", "local certificate version lower than peer, attempting to correct",
|
ixHandshakeStage0(n.intf, newHostinfo.vpnIp, newHostinfo)
|
||||||
)
|
}
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(hh *HandshakeHostInfo) {
|
|
||||||
hh.initiatingVersionOverride = peerCrt.Certificate.Version()
|
//If this is a static host, we don't need to wait for the HostQueryReply
|
||||||
})
|
//We can trigger the handshake right now
|
||||||
return
|
if _, ok := n.intf.lightHouse.GetStaticHostList()[hostinfo.vpnIp]; ok {
|
||||||
|
select {
|
||||||
|
case n.intf.handshakeManager.trigger <- hostinfo.vpnIp:
|
||||||
|
default:
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !bytes.Equal(curCrt.Signature(), myCrt.Signature()) {
|
|
||||||
cm.l.Info("Re-handshaking with remote",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"reason", "local certificate is not current",
|
|
||||||
)
|
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if curCrtVersion < cs.initiatingVersion {
|
|
||||||
cm.l.Info("Re-handshaking with remote",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"reason", "current cert version < pki.initiatingVersion",
|
|
||||||
)
|
|
||||||
|
|
||||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
+135
-353
@@ -1,29 +1,31 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"crypto/ed25519"
|
"crypto/ed25519"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"net/netip"
|
"net"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"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/overlaytest"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"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"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var vpnIp iputil.VpnIp
|
||||||
|
|
||||||
func newTestLighthouse() *LightHouse {
|
func newTestLighthouse() *LightHouse {
|
||||||
lh := &LightHouse{
|
lh := &LightHouse{
|
||||||
l: test.NewLogger(),
|
l: test.NewLogger(),
|
||||||
addrMap: map[netip.Addr]*RemoteList{},
|
addrMap: map[iputil.VpnIp]*RemoteList{},
|
||||||
queryChan: make(chan netip.Addr, 10),
|
|
||||||
}
|
}
|
||||||
lighthouses := []netip.Addr{}
|
lighthouses := map[iputil.VpnIp]struct{}{}
|
||||||
staticList := map[netip.Addr]struct{}{}
|
staticList := map[iputil.VpnIp]struct{}{}
|
||||||
|
|
||||||
lh.lighthouses.Store(&lighthouses)
|
lh.lighthouses.Store(&lighthouses)
|
||||||
lh.staticList.Store(&staticList)
|
lh.staticList.Store(&staticList)
|
||||||
@@ -34,265 +36,162 @@ func newTestLighthouse() *LightHouse {
|
|||||||
func Test_NewConnectionManagerTest(t *testing.T) {
|
func Test_NewConnectionManagerTest(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
//_, tuncidr, _ := net.ParseCIDR("1.1.1.1/24")
|
//_, tuncidr, _ := net.ParseCIDR("1.1.1.1/24")
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
_, vpncidr, _ := net.ParseCIDR("172.1.1.1/24")
|
||||||
vpnIp := netip.MustParseAddr("172.1.1.2")
|
_, localrange, _ := net.ParseCIDR("10.1.1.1/24")
|
||||||
preferredRanges := []netip.Prefix{localrange}
|
vpnIp = iputil.Ip2VpnIp(net.ParseIP("172.1.1.2"))
|
||||||
|
preferredRanges := []*net.IPNet{localrange}
|
||||||
|
|
||||||
// Very incomplete mock objects
|
// Very incomplete mock objects
|
||||||
hostMap := newHostMap(l)
|
hostMap := NewHostMap(l, "test", vpncidr, preferredRanges)
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
initiatingVersion: cert.Version1,
|
rawCertificate: []byte{},
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
certificate: &cert.NebulaCertificate{},
|
||||||
v1Credential: nil,
|
rawCertificateNoKey: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
pki: &PKI{},
|
handshakeManager: NewHandshakeManager(l, vpncidr, preferredRanges, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
ifce.pki.cs.Store(cs)
|
ifce.certState.Store(cs)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
defer cancel()
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
punchy := NewPunchyFromConfig(l, config.NewC(l))
|
||||||
nc.intf = ifce
|
nc := newConnectionManager(ctx, l, ifce, 5, 10, punchy)
|
||||||
ifce.pmtudManager = newPMTUDManagerFromConfig(test.NewLogger(), conf, ifce.inside)
|
|
||||||
ifce.pmtudManager.intf = ifce
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
|
|
||||||
// Add an ip we have established a connection w/ to hostmap
|
// Add an ip we have established a connection w/ to hostmap
|
||||||
hostinfo := &HostInfo{
|
hostinfo := &HostInfo{
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
vpnIp: vpnIp,
|
||||||
localIndexId: 1099,
|
localIndexId: 1099,
|
||||||
remoteIndexId: 9901,
|
remoteIndexId: 9901,
|
||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
certState: cs,
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
// We saw traffic out to vpnIp
|
// We saw traffic out to vpnIp
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo.localIndexId)
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.True(t, hostinfo.out.Load())
|
assert.Contains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.True(t, hostinfo.in.Load())
|
|
||||||
|
|
||||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.NotContains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.in, hostinfo.localIndexId)
|
||||||
|
|
||||||
// Do another traffic check tick, this host should be pending deletion now
|
// Do another traffic check tick, this host should be pending deletion now
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo.localIndexId)
|
||||||
assert.True(t, hostinfo.out.Load())
|
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
assert.Contains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.NotContains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.in, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
|
|
||||||
// Do a final traffic check tick, the host should now be removed
|
// Do a final traffic check tick, the host should now be removed
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.NotContains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs)
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
|
assert.NotContains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
assert.NotContains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.NotContains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_NewConnectionManagerTest2(t *testing.T) {
|
func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
//_, tuncidr, _ := net.ParseCIDR("1.1.1.1/24")
|
//_, tuncidr, _ := net.ParseCIDR("1.1.1.1/24")
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
_, vpncidr, _ := net.ParseCIDR("172.1.1.1/24")
|
||||||
vpnIp := netip.MustParseAddr("172.1.1.2")
|
_, localrange, _ := net.ParseCIDR("10.1.1.1/24")
|
||||||
preferredRanges := []netip.Prefix{localrange}
|
preferredRanges := []*net.IPNet{localrange}
|
||||||
|
|
||||||
// Very incomplete mock objects
|
// Very incomplete mock objects
|
||||||
hostMap := newHostMap(l)
|
hostMap := NewHostMap(l, "test", vpncidr, preferredRanges)
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
initiatingVersion: cert.Version1,
|
rawCertificate: []byte{},
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
certificate: &cert.NebulaCertificate{},
|
||||||
v1Credential: nil,
|
rawCertificateNoKey: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
pki: &PKI{},
|
handshakeManager: NewHandshakeManager(l, vpncidr, preferredRanges, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
ifce.pki.cs.Store(cs)
|
ifce.certState.Store(cs)
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
defer cancel()
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
punchy := NewPunchyFromConfig(l, config.NewC(l))
|
||||||
nc.intf = ifce
|
nc := newConnectionManager(ctx, l, ifce, 5, 10, punchy)
|
||||||
ifce.pmtudManager = newPMTUDManagerFromConfig(test.NewLogger(), conf, ifce.inside)
|
|
||||||
ifce.pmtudManager.intf = ifce
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
|
|
||||||
// Add an ip we have established a connection w/ to hostmap
|
// Add an ip we have established a connection w/ to hostmap
|
||||||
hostinfo := &HostInfo{
|
hostinfo := &HostInfo{
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
vpnIp: vpnIp,
|
||||||
localIndexId: 1099,
|
localIndexId: 1099,
|
||||||
remoteIndexId: 9901,
|
remoteIndexId: 9901,
|
||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
certState: cs,
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
// We saw traffic out to vpnIp
|
// We saw traffic out to vpnIp
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo.localIndexId)
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo.localIndexId)
|
||||||
assert.True(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.vpnIp)
|
||||||
assert.True(t, hostinfo.out.Load())
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
|
|
||||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.NotContains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.in, hostinfo.localIndexId)
|
||||||
|
|
||||||
// Do another traffic check tick, this host should be pending deletion now
|
// Do another traffic check tick, this host should be pending deletion now
|
||||||
nc.Out(hostinfo)
|
nc.Out(hostinfo.localIndexId)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
assert.Contains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.NotContains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.in, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
|
|
||||||
// We saw traffic, should no longer be pending deletion
|
// We saw traffic, should no longer be pending deletion
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo.localIndexId)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.NotContains(t, nc.pendingDeletion, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.NotContains(t, nc.out, hostinfo.localIndexId)
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.NotContains(t, nc.in, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnIp)
|
||||||
}
|
|
||||||
|
|
||||||
func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|
||||||
l := test.NewLogger()
|
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
|
||||||
vpnAddrs := []netip.Addr{netip.MustParseAddr("172.1.1.2")}
|
|
||||||
preferredRanges := []netip.Prefix{localrange}
|
|
||||||
|
|
||||||
// Very incomplete mock objects
|
|
||||||
hostMap := newHostMap(l)
|
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
|
||||||
|
|
||||||
cs := &CertState{
|
|
||||||
initiatingVersion: cert.Version1,
|
|
||||||
privateKey: []byte{},
|
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
|
||||||
v1Credential: nil,
|
|
||||||
}
|
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
|
||||||
ifce := &Interface{
|
|
||||||
hostMap: hostMap,
|
|
||||||
inside: &overlaytest.NoopTun{},
|
|
||||||
outside: &udp.NoopConn{},
|
|
||||||
firewall: &Firewall{},
|
|
||||||
lightHouse: lh,
|
|
||||||
pki: &PKI{},
|
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
ifce.pki.cs.Store(cs)
|
|
||||||
|
|
||||||
// Create manager
|
|
||||||
conf := config.NewC(test.NewLogger())
|
|
||||||
conf.Settings["tunnels"] = map[string]any{
|
|
||||||
"drop_inactive": true,
|
|
||||||
}
|
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
|
||||||
assert.True(t, nc.dropInactive.Load())
|
|
||||||
nc.intf = ifce
|
|
||||||
|
|
||||||
// Add an ip we have established a connection w/ to hostmap
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
vpnAddrs: vpnAddrs,
|
|
||||||
localIndexId: 1099,
|
|
||||||
remoteIndexId: 9901,
|
|
||||||
}
|
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
|
||||||
}
|
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
|
||||||
|
|
||||||
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
|
|
||||||
nc.Out(hostinfo)
|
|
||||||
nc.In(hostinfo)
|
|
||||||
assert.True(t, hostinfo.out.Load())
|
|
||||||
assert.True(t, hostinfo.in.Load())
|
|
||||||
|
|
||||||
now := time.Now()
|
|
||||||
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
|
|
||||||
assert.Equal(t, tryRehandshake, decision)
|
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
|
||||||
assert.False(t, hostinfo.out.Load())
|
|
||||||
assert.False(t, hostinfo.in.Load())
|
|
||||||
|
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
|
|
||||||
assert.Equal(t, doNothing, decision)
|
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
|
||||||
assert.False(t, hostinfo.out.Load())
|
|
||||||
assert.False(t, hostinfo.in.Load())
|
|
||||||
|
|
||||||
// Do another traffic check tick, should still not be pending deletion
|
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
|
|
||||||
assert.Equal(t, doNothing, decision)
|
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
|
||||||
assert.False(t, hostinfo.out.Load())
|
|
||||||
assert.False(t, hostinfo.in.Load())
|
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
|
||||||
|
|
||||||
// Finally advance beyond the inactivity timeout
|
|
||||||
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
|
|
||||||
assert.Equal(t, closeTunnel, decision)
|
|
||||||
assert.Equal(t, now, hostinfo.lastUsed)
|
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
|
||||||
assert.False(t, hostinfo.out.Load())
|
|
||||||
assert.False(t, hostinfo.in.Load())
|
|
||||||
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
|
||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if we can disconnect the peer.
|
// Check if we can disconnect the peer.
|
||||||
@@ -301,82 +200,80 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
|
ipNet := net.IPNet{
|
||||||
vpncidr := netip.MustParsePrefix("172.1.1.1/24")
|
IP: net.IPv4(172, 1, 1, 2),
|
||||||
localrange := netip.MustParsePrefix("10.1.1.1/24")
|
Mask: net.IPMask{255, 255, 255, 0},
|
||||||
vpnIp := netip.MustParseAddr("172.1.1.2")
|
}
|
||||||
preferredRanges := []netip.Prefix{localrange}
|
_, vpncidr, _ := net.ParseCIDR("172.1.1.1/24")
|
||||||
hostMap := newHostMap(l)
|
_, localrange, _ := net.ParseCIDR("10.1.1.1/24")
|
||||||
hostMap.preferredRanges.Store(&preferredRanges)
|
preferredRanges := []*net.IPNet{localrange}
|
||||||
|
hostMap := NewHostMap(l, "test", vpncidr, preferredRanges)
|
||||||
|
|
||||||
// Generate keys for CA and peer's cert.
|
// Generate keys for CA and peer's cert.
|
||||||
pubCA, privCA, _ := ed25519.GenerateKey(rand.Reader)
|
pubCA, privCA, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
tbs := &cert.TBSCertificate{
|
caCert := cert.NebulaCertificate{
|
||||||
Version: 1,
|
Details: cert.NebulaCertificateDetails{
|
||||||
Name: "ca",
|
Name: "ca",
|
||||||
IsCA: true,
|
NotBefore: now,
|
||||||
NotBefore: now,
|
NotAfter: now.Add(1 * time.Hour),
|
||||||
NotAfter: now.Add(1 * time.Hour),
|
IsCA: true,
|
||||||
PublicKey: pubCA,
|
PublicKey: pubCA,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
caCert.Sign(cert.Curve_CURVE25519, privCA)
|
||||||
caCert, err := tbs.Sign(nil, cert.Curve_CURVE25519, privCA)
|
ncp := &cert.NebulaCAPool{
|
||||||
require.NoError(t, err)
|
CAs: cert.NewCAPool().CAs,
|
||||||
ncp := cert.NewCAPool()
|
}
|
||||||
require.NoError(t, ncp.AddCA(caCert))
|
ncp.CAs["ca"] = &caCert
|
||||||
|
|
||||||
pubCrt, _, _ := ed25519.GenerateKey(rand.Reader)
|
pubCrt, _, _ := ed25519.GenerateKey(rand.Reader)
|
||||||
tbs = &cert.TBSCertificate{
|
peerCert := cert.NebulaCertificate{
|
||||||
Version: 1,
|
Details: cert.NebulaCertificateDetails{
|
||||||
Name: "host",
|
Name: "host",
|
||||||
Networks: []netip.Prefix{vpncidr},
|
Ips: []*net.IPNet{&ipNet},
|
||||||
NotBefore: now,
|
Subnets: []*net.IPNet{},
|
||||||
NotAfter: now.Add(60 * time.Second),
|
NotBefore: now,
|
||||||
PublicKey: pubCrt,
|
NotAfter: now.Add(60 * time.Second),
|
||||||
|
PublicKey: pubCrt,
|
||||||
|
IsCA: false,
|
||||||
|
Issuer: "ca",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
peerCert, err := tbs.Sign(caCert, cert.Curve_CURVE25519, privCA)
|
peerCert.Sign(cert.Curve_CURVE25519, privCA)
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
|
||||||
|
|
||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
privateKey: []byte{},
|
rawCertificate: []byte{},
|
||||||
v1Cert: &dummyCert{},
|
privateKey: []byte{},
|
||||||
v1Credential: nil,
|
certificate: &cert.NebulaCertificate{},
|
||||||
|
rawCertificateNoKey: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
ifce := &Interface{
|
ifce := &Interface{
|
||||||
hostMap: hostMap,
|
hostMap: hostMap,
|
||||||
inside: &overlaytest.NoopTun{},
|
inside: &test.NoopTun{},
|
||||||
outside: &udp.NoopConn{},
|
outside: &udp.NoopConn{},
|
||||||
firewall: &Firewall{},
|
firewall: &Firewall{},
|
||||||
lightHouse: lh,
|
lightHouse: lh,
|
||||||
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
handshakeManager: NewHandshakeManager(l, vpncidr, preferredRanges, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
|
||||||
l: l,
|
l: l,
|
||||||
pki: &PKI{},
|
disconnectInvalid: true,
|
||||||
|
caPool: ncp,
|
||||||
}
|
}
|
||||||
ifce.pki.cs.Store(cs)
|
ifce.certState.Store(cs)
|
||||||
ifce.pki.caPool.Store(ncp)
|
|
||||||
ifce.disconnectInvalid.Store(true)
|
|
||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
defer cancel()
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
punchy := NewPunchyFromConfig(l, config.NewC(l))
|
||||||
nc.intf = ifce
|
nc := newConnectionManager(ctx, l, ifce, 5, 10, punchy)
|
||||||
ifce.pmtudManager = newPMTUDManagerFromConfig(test.NewLogger(), conf, ifce.inside)
|
|
||||||
ifce.pmtudManager.intf = ifce
|
|
||||||
ifce.connectionManager = nc
|
ifce.connectionManager = nc
|
||||||
|
hostinfo, _ := nc.hostMap.AddVpnIp(vpnIp, nil)
|
||||||
hostinfo := &HostInfo{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
certState: cs,
|
||||||
ConnectionState: &ConnectionState{
|
peerCert: &peerCert,
|
||||||
myCert: &dummyCert{},
|
H: &noise.HandshakeState{},
|
||||||
peerCert: cachedPeerCert,
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
|
||||||
|
|
||||||
// Move ahead 45s.
|
// Move ahead 45s.
|
||||||
// Check if to disconnect with invalid certificate.
|
// Check if to disconnect with invalid certificate.
|
||||||
@@ -392,118 +289,3 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
invalid = nc.isInvalidCertificate(nextTick, hostinfo)
|
invalid = nc.isInvalidCertificate(nextTick, hostinfo)
|
||||||
assert.True(t, invalid)
|
assert.True(t, invalid)
|
||||||
}
|
}
|
||||||
|
|
||||||
type dummyCert struct {
|
|
||||||
version cert.Version
|
|
||||||
curve cert.Curve
|
|
||||||
groups []string
|
|
||||||
isCa bool
|
|
||||||
issuer string
|
|
||||||
name string
|
|
||||||
networks []netip.Prefix
|
|
||||||
notAfter time.Time
|
|
||||||
notBefore time.Time
|
|
||||||
publicKey []byte
|
|
||||||
signature []byte
|
|
||||||
unsafeNetworks []netip.Prefix
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Version() cert.Version {
|
|
||||||
return d.version
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Curve() cert.Curve {
|
|
||||||
return d.curve
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Groups() []string {
|
|
||||||
return d.groups
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) IsCA() bool {
|
|
||||||
return d.isCa
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Issuer() string {
|
|
||||||
return d.issuer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Name() string {
|
|
||||||
return d.name
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Networks() []netip.Prefix {
|
|
||||||
return d.networks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) NotAfter() time.Time {
|
|
||||||
return d.notAfter
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) NotBefore() time.Time {
|
|
||||||
return d.notBefore
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) PublicKey() []byte {
|
|
||||||
return d.publicKey
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) MarshalPublicKeyPEM() []byte {
|
|
||||||
return cert.MarshalPublicKeyToPEM(d.curve, d.publicKey)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Signature() []byte {
|
|
||||||
return d.signature
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) UnsafeNetworks() []netip.Prefix {
|
|
||||||
return d.unsafeNetworks
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) MarshalForHandshakes() ([]byte, error) {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Sign(curve cert.Curve, key []byte) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) CheckSignature(key []byte) bool {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Expired(t time.Time) bool {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) CheckRootConstraints(signer cert.Certificate) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) VerifyPrivateKey(curve cert.Curve, key []byte) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) String() string {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Marshal() ([]byte, error) {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) MarshalPEM() ([]byte, error) {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Fingerprint() (string, error) {
|
|
||||||
return "", nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) MarshalJSON() ([]byte, error) {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dummyCert) Copy() cert.Certificate {
|
|
||||||
return d
|
|
||||||
}
|
|
||||||
|
|||||||
+55
-22
@@ -1,12 +1,15 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 1024
|
||||||
@@ -14,31 +17,64 @@ const ReplayWindow = 1024
|
|||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey *NebulaCipherState
|
eKey *NebulaCipherState
|
||||||
dKey *NebulaCipherState
|
dKey *NebulaCipherState
|
||||||
myCert cert.Certificate
|
H *noise.HandshakeState
|
||||||
peerCert *cert.CachedCertificate
|
certState *CertState
|
||||||
|
peerCert *cert.NebulaCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
messageCounter atomic.Uint64
|
messageCounter atomic.Uint64
|
||||||
window *Bits
|
window *Bits
|
||||||
|
queueLock sync.Mutex
|
||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
|
ready bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
func (f *Interface) newConnectionState(l *logrus.Logger, initiator bool, pattern noise.HandshakePattern, psk []byte, pskStage int) *ConnectionState {
|
||||||
// completed handshake.Result. It seeds messageCounter and the replay window so
|
var dhFunc noise.DHFunc
|
||||||
// that the post-handshake message indices already used on the wire don't count
|
curCertState := f.certState.Load()
|
||||||
// as missed traffic in the data plane.
|
|
||||||
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
switch curCertState.certificate.Details.Curve {
|
||||||
|
case cert.Curve_CURVE25519:
|
||||||
|
dhFunc = noise.DH25519
|
||||||
|
case cert.Curve_P256:
|
||||||
|
dhFunc = noiseutil.DHP256
|
||||||
|
default:
|
||||||
|
l.Errorf("invalid curve: %s", curCertState.certificate.Details.Curve)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cs := noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256)
|
||||||
|
if f.cipher == "chachapoly" {
|
||||||
|
cs = noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
}
|
||||||
|
|
||||||
|
static := noise.DHKey{Private: curCertState.privateKey, Public: curCertState.publicKey}
|
||||||
|
|
||||||
|
b := NewBits(ReplayWindow)
|
||||||
|
// Clear out bit 0, we never transmit it and we don't want it showing as packet loss
|
||||||
|
b.Update(l, 0)
|
||||||
|
|
||||||
|
hs, err := noise.NewHandshakeState(noise.Config{
|
||||||
|
CipherSuite: cs,
|
||||||
|
Random: rand.Reader,
|
||||||
|
Pattern: pattern,
|
||||||
|
Initiator: initiator,
|
||||||
|
StaticKeypair: static,
|
||||||
|
PresharedKey: psk,
|
||||||
|
PresharedKeyPlacement: pskStage,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// The queue and ready params prevent a counter race that would happen when
|
||||||
|
// sending stored packets and simultaneously accepting new traffic.
|
||||||
ci := &ConnectionState{
|
ci := &ConnectionState{
|
||||||
myCert: r.MyCert,
|
H: hs,
|
||||||
initiator: r.Initiator,
|
initiator: initiator,
|
||||||
peerCert: r.RemoteCert,
|
window: b,
|
||||||
eKey: NewNebulaCipherState(r.EKey),
|
ready: false,
|
||||||
dKey: NewNebulaCipherState(r.DKey),
|
certState: curCertState,
|
||||||
window: NewBits(ReplayWindow),
|
|
||||||
}
|
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
|
||||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
|
||||||
ci.window.Update(nil, i)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return ci
|
return ci
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -47,9 +83,6 @@ func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
|||||||
"certificate": cs.peerCert,
|
"certificate": cs.peerCert,
|
||||||
"initiator": cs.initiator,
|
"initiator": cs.initiator,
|
||||||
"message_counter": cs.messageCounter.Load(),
|
"message_counter": cs.messageCounter.Load(),
|
||||||
|
"ready": cs.ready,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) Curve() cert.Curve {
|
|
||||||
return cs.myCert.Curve()
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,114 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
ct "github.com/slackhq/nebula/cert_test"
|
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// runTestHandshake runs a complete IX handshake between two freshly-built
|
|
||||||
// peers and returns the initiator and responder Results. Used to produce
|
|
||||||
// real cipher states for tests that need to exercise post-handshake glue.
|
|
||||||
func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
|
||||||
)
|
|
||||||
caPool := ct.NewTestCAPool(ca)
|
|
||||||
|
|
||||||
makeCreds := func(name string, networks []netip.Prefix) handshake.GetCredentialFunc {
|
|
||||||
c, _, rawKey, _ := ct.NewTestCert(
|
|
||||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
|
||||||
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
|
||||||
)
|
|
||||||
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawKey)
|
|
||||||
require.NoError(t, err)
|
|
||||||
hsBytes, err := c.MarshalForHandshakes()
|
|
||||||
require.NoError(t, err)
|
|
||||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
cred := handshake.NewCredential(c, hsBytes, priv, ncs)
|
|
||||||
return func(v cert.Version) *handshake.Credential {
|
|
||||||
if v == cert.Version2 {
|
|
||||||
return cred
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
verifier := func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
|
||||||
return caPool.VerifyCertificate(time.Now(), c)
|
|
||||||
}
|
|
||||||
|
|
||||||
initCreds := makeCreds("initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
|
||||||
respCreds := makeCreds("responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
|
||||||
|
|
||||||
initM, err := handshake.NewMachine(
|
|
||||||
cert.Version2, initCreds, verifier,
|
|
||||||
func() (uint32, error) { return 1000, nil },
|
|
||||||
true, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
respM, err := handshake.NewMachine(
|
|
||||||
cert.Version2, respCreds, verifier,
|
|
||||||
func() (uint32, error) { return 2000, nil },
|
|
||||||
false, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
msg1, err := initM.Initiate(nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
resp, respR, err := respM.ProcessPacket(nil, msg1)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, respR)
|
|
||||||
|
|
||||||
_, initR, err = initM.ProcessPacket(nil, resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotNil(t, initR)
|
|
||||||
|
|
||||||
return initR, respR
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewConnectionStateFromResult(t *testing.T) {
|
|
||||||
initR, respR := runTestHandshake(t)
|
|
||||||
|
|
||||||
t.Run("initiator", func(t *testing.T) {
|
|
||||||
ci := newConnectionStateFromResult(initR)
|
|
||||||
assert.True(t, ci.initiator)
|
|
||||||
assert.Equal(t, initR.MyCert, ci.myCert)
|
|
||||||
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
|
||||||
assert.NotNil(t, ci.eKey)
|
|
||||||
assert.NotNil(t, ci.dKey)
|
|
||||||
|
|
||||||
// IX has 2 handshake messages; the next data-plane send is counter=3.
|
|
||||||
assert.Equal(t, uint64(2), ci.messageCounter.Load(),
|
|
||||||
"messageCounter must equal Result.MessageIndex so the next send is N+1")
|
|
||||||
|
|
||||||
// Both handshake counters must be marked seen so they don't appear lost.
|
|
||||||
// Check returns false if an index has already been recorded.
|
|
||||||
assert.False(t, ci.window.Check(nil, 1), "counter 1 must already be seen")
|
|
||||||
assert.False(t, ci.window.Check(nil, 2), "counter 2 must already be seen")
|
|
||||||
// Counter 3 is the next data-plane message and must NOT be pre-marked.
|
|
||||||
assert.True(t, ci.window.Check(nil, 3), "counter 3 must not be pre-seeded")
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("responder", func(t *testing.T) {
|
|
||||||
ci := newConnectionStateFromResult(respR)
|
|
||||||
assert.False(t, ci.initiator)
|
|
||||||
assert.Equal(t, respR.MyCert, ci.myCert)
|
|
||||||
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
|
||||||
assert.NotNil(t, ci.eKey)
|
|
||||||
assert.NotNil(t, ci.dKey)
|
|
||||||
assert.Equal(t, uint64(2), ci.messageCounter.Load())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
+89
-212
@@ -2,164 +2,74 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"net"
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"sync"
|
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
type controlEach func(h *HostInfo)
|
|
||||||
|
|
||||||
type controlHostLister interface {
|
|
||||||
QueryVpnAddr(vpnAddr netip.Addr) *HostInfo
|
|
||||||
ForEachIndex(each controlEach)
|
|
||||||
ForEachVpnAddr(each controlEach)
|
|
||||||
GetPreferredRanges() []netip.Prefix
|
|
||||||
}
|
|
||||||
|
|
||||||
type Control struct {
|
type Control struct {
|
||||||
stateLock sync.Mutex
|
f *Interface
|
||||||
state RunState
|
l *logrus.Logger
|
||||||
|
cancel context.CancelFunc
|
||||||
f *Interface
|
sshStart func()
|
||||||
l *slog.Logger
|
httpStart func()
|
||||||
ctx context.Context
|
dnsStart func()
|
||||||
cancel context.CancelFunc
|
|
||||||
sshStart func()
|
|
||||||
statsStart func()
|
|
||||||
dnsStart func()
|
|
||||||
lighthouseStart func()
|
|
||||||
connectionManagerStart func(context.Context)
|
|
||||||
pmtudManagerStart func(context.Context)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ControlHostInfo struct {
|
type ControlHostInfo struct {
|
||||||
VpnAddrs []netip.Addr `json:"vpnAddrs"`
|
VpnIp net.IP `json:"vpnIp"`
|
||||||
LocalIndex uint32 `json:"localIndex"`
|
LocalIndex uint32 `json:"localIndex"`
|
||||||
RemoteIndex uint32 `json:"remoteIndex"`
|
RemoteIndex uint32 `json:"remoteIndex"`
|
||||||
RemoteAddrs []netip.AddrPort `json:"remoteAddrs"`
|
RemoteAddrs []*udp.Addr `json:"remoteAddrs"`
|
||||||
Cert cert.Certificate `json:"cert"`
|
CachedPackets int `json:"cachedPackets"`
|
||||||
MessageCounter uint64 `json:"messageCounter"`
|
Cert *cert.NebulaCertificate `json:"cert"`
|
||||||
CurrentRemote netip.AddrPort `json:"currentRemote"`
|
MessageCounter uint64 `json:"messageCounter"`
|
||||||
CurrentRelaysToMe []netip.Addr `json:"currentRelaysToMe"`
|
CurrentRemote *udp.Addr `json:"currentRemote"`
|
||||||
CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"`
|
CurrentRelaysToMe []iputil.VpnIp `json:"currentRelaysToMe"`
|
||||||
|
CurrentRelaysThroughMe []iputil.VpnIp `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 {
|
||||||
go c.sshStart()
|
go c.sshStart()
|
||||||
}
|
}
|
||||||
if c.statsStart != nil {
|
if c.httpStart != nil {
|
||||||
go c.statsStart()
|
go c.httpStart()
|
||||||
}
|
}
|
||||||
if c.dnsStart != nil {
|
if c.dnsStart != nil {
|
||||||
go c.dnsStart()
|
go c.dnsStart()
|
||||||
}
|
}
|
||||||
if c.connectionManagerStart != nil {
|
|
||||||
go c.connectionManagerStart(c.ctx)
|
|
||||||
}
|
|
||||||
if c.pmtudManagerStart != nil {
|
|
||||||
go c.pmtudManagerStart(c.ctx)
|
|
||||||
}
|
|
||||||
if c.lighthouseStart != nil {
|
|
||||||
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 {
|
// Stop signals nebula to shutdown, returns after the shutdown is complete
|
||||||
c.stateLock.Lock()
|
|
||||||
defer c.stateLock.Unlock()
|
|
||||||
return c.state
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Control) Context() context.Context {
|
|
||||||
return c.ctx
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
|
||||||
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()
|
||||||
|
|
||||||
c.CloseAllTunnels(false)
|
c.CloseAllTunnels(false)
|
||||||
if err := c.f.Close(); err != nil {
|
if err := c.f.Close(); err != nil {
|
||||||
c.l.Error("Close interface failed", "error", err)
|
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
|
||||||
@@ -170,7 +80,7 @@ func (c *Control) ShutdownBlock() {
|
|||||||
|
|
||||||
rawSig := <-sigChan
|
rawSig := <-sigChan
|
||||||
sig := rawSig.String()
|
sig := rawSig.String()
|
||||||
c.l.Info("Caught signal, shutting down", "signal", sig)
|
c.l.WithField("signal", sig).Info("Caught signal, shutting down")
|
||||||
c.Stop()
|
c.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -179,7 +89,7 @@ func (c *Control) RebindUDPServer() {
|
|||||||
_ = c.f.outside.Rebind()
|
_ = c.f.outside.Rebind()
|
||||||
|
|
||||||
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
|
||||||
c.f.lightHouse.SendUpdate()
|
c.f.lightHouse.SendUpdate(c.f)
|
||||||
|
|
||||||
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
|
||||||
c.f.rebindCount++
|
c.f.rebindCount++
|
||||||
@@ -188,7 +98,7 @@ func (c *Control) RebindUDPServer() {
|
|||||||
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
|
||||||
func (c *Control) ListHostmapHosts(pendingMap bool) []ControlHostInfo {
|
func (c *Control) ListHostmapHosts(pendingMap bool) []ControlHostInfo {
|
||||||
if pendingMap {
|
if pendingMap {
|
||||||
return listHostMapHosts(c.f.handshakeManager)
|
return listHostMapHosts(c.f.handshakeManager.pendingHostMap)
|
||||||
} else {
|
} else {
|
||||||
return listHostMapHosts(c.f.hostMap)
|
return listHostMapHosts(c.f.hostMap)
|
||||||
}
|
}
|
||||||
@@ -197,87 +107,46 @@ func (c *Control) ListHostmapHosts(pendingMap bool) []ControlHostInfo {
|
|||||||
// ListHostmapIndexes returns details about the actual or pending (handshaking) hostmap by local index id
|
// ListHostmapIndexes returns details about the actual or pending (handshaking) hostmap by local index id
|
||||||
func (c *Control) ListHostmapIndexes(pendingMap bool) []ControlHostInfo {
|
func (c *Control) ListHostmapIndexes(pendingMap bool) []ControlHostInfo {
|
||||||
if pendingMap {
|
if pendingMap {
|
||||||
return listHostMapIndexes(c.f.handshakeManager)
|
return listHostMapIndexes(c.f.handshakeManager.pendingHostMap)
|
||||||
} else {
|
} else {
|
||||||
return listHostMapIndexes(c.f.hostMap)
|
return listHostMapIndexes(c.f.hostMap)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetCertByVpnIp returns the authenticated certificate of the given vpn IP, or nil if not found
|
// GetHostInfoByVpnIp returns a single tunnels hostInfo, or nil if not found
|
||||||
func (c *Control) GetCertByVpnIp(vpnIp netip.Addr) cert.Certificate {
|
func (c *Control) GetHostInfoByVpnIp(vpnIp iputil.VpnIp, pending bool) *ControlHostInfo {
|
||||||
if c.f.myVpnAddrsTable.Contains(vpnIp) {
|
var hm *HostMap
|
||||||
// Only returning the default certificate since its impossible
|
|
||||||
// for any other host but ourselves to have more than 1
|
|
||||||
return c.f.pki.getCertState().GetDefaultCertificate().Copy()
|
|
||||||
}
|
|
||||||
hi := c.f.hostMap.QueryVpnAddr(vpnIp)
|
|
||||||
if hi == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return hi.GetCert().Certificate.Copy()
|
|
||||||
}
|
|
||||||
|
|
||||||
// CreateTunnel creates a new tunnel to the given vpn ip.
|
|
||||||
func (c *Control) CreateTunnel(vpnIp netip.Addr) {
|
|
||||||
c.f.handshakeManager.StartHandshake(vpnIp, nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
// PrintTunnel creates a new tunnel to the given vpn ip.
|
|
||||||
func (c *Control) PrintTunnel(vpnIp netip.Addr) *ControlHostInfo {
|
|
||||||
hi := c.f.hostMap.QueryVpnAddr(vpnIp)
|
|
||||||
if hi == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
chi := copyHostInfo(hi, c.f.hostMap.GetPreferredRanges())
|
|
||||||
return &chi
|
|
||||||
}
|
|
||||||
|
|
||||||
// QueryLighthouse queries the lighthouse.
|
|
||||||
func (c *Control) QueryLighthouse(vpnIp netip.Addr) *CacheMap {
|
|
||||||
hi := c.f.lightHouse.Query(vpnIp)
|
|
||||||
if hi == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return hi.CopyCache()
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetHostInfoByVpnAddr returns a single tunnels hostInfo, or nil if not found
|
|
||||||
// Caller should take care to Unmap() any 4in6 addresses prior to calling.
|
|
||||||
func (c *Control) GetHostInfoByVpnAddr(vpnAddr netip.Addr, pending bool) *ControlHostInfo {
|
|
||||||
var hl controlHostLister
|
|
||||||
if pending {
|
if pending {
|
||||||
hl = c.f.handshakeManager
|
hm = c.f.handshakeManager.pendingHostMap
|
||||||
} else {
|
} else {
|
||||||
hl = c.f.hostMap
|
hm = c.f.hostMap
|
||||||
}
|
}
|
||||||
|
|
||||||
h := hl.QueryVpnAddr(vpnAddr)
|
h, err := hm.QueryVpnIp(vpnIp)
|
||||||
if h == nil {
|
if err != nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
ch := copyHostInfo(h, c.f.hostMap.GetPreferredRanges())
|
ch := copyHostInfo(h, c.f.hostMap.preferredRanges)
|
||||||
return &ch
|
return &ch
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetRemoteForTunnel forces a tunnel to use a specific remote
|
// SetRemoteForTunnel forces a tunnel to use a specific remote
|
||||||
// Caller should take care to Unmap() any 4in6 addresses prior to calling.
|
func (c *Control) SetRemoteForTunnel(vpnIp iputil.VpnIp, addr udp.Addr) *ControlHostInfo {
|
||||||
func (c *Control) SetRemoteForTunnel(vpnIp netip.Addr, addr netip.AddrPort) *ControlHostInfo {
|
hostInfo, err := c.f.hostMap.QueryVpnIp(vpnIp)
|
||||||
hostInfo := c.f.hostMap.QueryVpnAddr(vpnIp)
|
if err != nil {
|
||||||
if hostInfo == nil {
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
hostInfo.SetRemote(addr)
|
hostInfo.SetRemote(addr.Copy())
|
||||||
ch := copyHostInfo(hostInfo, c.f.hostMap.GetPreferredRanges())
|
ch := copyHostInfo(hostInfo, c.f.hostMap.preferredRanges)
|
||||||
return &ch
|
return &ch
|
||||||
}
|
}
|
||||||
|
|
||||||
// CloseTunnel closes a fully established tunnel. If localOnly is false it will notify the remote end as well.
|
// CloseTunnel closes a fully established tunnel. If localOnly is false it will notify the remote end as well.
|
||||||
// Caller should take care to Unmap() any 4in6 addresses prior to calling.
|
func (c *Control) CloseTunnel(vpnIp iputil.VpnIp, localOnly bool) bool {
|
||||||
func (c *Control) CloseTunnel(vpnIp netip.Addr, localOnly bool) bool {
|
hostInfo, err := c.f.hostMap.QueryVpnIp(vpnIp)
|
||||||
hostInfo := c.f.hostMap.QueryVpnAddr(vpnIp)
|
if err != nil {
|
||||||
if hostInfo == nil {
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -300,26 +169,29 @@ func (c *Control) CloseTunnel(vpnIp netip.Addr, localOnly bool) bool {
|
|||||||
// CloseAllTunnels is just like CloseTunnel except it goes through and shuts them all down, optionally you can avoid shutting down lighthouse tunnels
|
// CloseAllTunnels is just like CloseTunnel except it goes through and shuts them all down, optionally you can avoid shutting down lighthouse tunnels
|
||||||
// the int returned is a count of tunnels closed
|
// the int returned is a count of tunnels closed
|
||||||
func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
||||||
|
//TODO: this is probably better as a function in ConnectionManager or HostMap directly
|
||||||
|
lighthouses := c.f.lightHouse.GetLighthouses()
|
||||||
|
|
||||||
shutdown := func(h *HostInfo) {
|
shutdown := func(h *HostInfo) {
|
||||||
if excludeLighthouses && c.f.lightHouse.IsAnyLighthouseAddr(h.vpnAddrs) {
|
if excludeLighthouses {
|
||||||
return
|
if _, ok := lighthouses[h.vpnIp]; ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
c.f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
c.f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
||||||
c.f.closeTunnel(h)
|
c.f.closeTunnel(h)
|
||||||
|
|
||||||
c.l.Debug("Sending close tunnel message",
|
c.l.WithField("vpnIp", h.vpnIp).WithField("udpAddr", h.remote).
|
||||||
"vpnAddrs", h.vpnAddrs,
|
Debug("Sending close tunnel message")
|
||||||
"udpAddr", h.remote,
|
|
||||||
)
|
|
||||||
closed++
|
closed++
|
||||||
}
|
}
|
||||||
|
|
||||||
// Learn which hosts are being used as relays, so we can shut them down last.
|
// Learn which hosts are being used as relays, so we can shut them down last.
|
||||||
relayingHosts := map[netip.Addr]*HostInfo{}
|
relayingHosts := map[iputil.VpnIp]*HostInfo{}
|
||||||
// Grab the hostMap lock to access the Relays map
|
// Grab the hostMap lock to access the Relays map
|
||||||
c.f.hostMap.Lock()
|
c.f.hostMap.Lock()
|
||||||
for _, relayingHost := range c.f.hostMap.Relays {
|
for _, relayingHost := range c.f.hostMap.Relays {
|
||||||
relayingHosts[relayingHost.vpnAddrs[0]] = relayingHost
|
relayingHosts[relayingHost.vpnIp] = relayingHost
|
||||||
}
|
}
|
||||||
c.f.hostMap.Unlock()
|
c.f.hostMap.Unlock()
|
||||||
|
|
||||||
@@ -327,7 +199,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
|||||||
// Grab the hostMap lock to access the Hosts map
|
// Grab the hostMap lock to access the Hosts map
|
||||||
c.f.hostMap.Lock()
|
c.f.hostMap.Lock()
|
||||||
for _, relayHost := range c.f.hostMap.Indexes {
|
for _, relayHost := range c.f.hostMap.Indexes {
|
||||||
if _, ok := relayingHosts[relayHost.vpnAddrs[0]]; !ok {
|
if _, ok := relayingHosts[relayHost.vpnIp]; !ok {
|
||||||
hostInfos = append(hostInfos, relayHost)
|
hostInfos = append(hostInfos, relayHost)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -342,23 +214,16 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) Device() overlay.Device {
|
func copyHostInfo(h *HostInfo, preferredRanges []*net.IPNet) ControlHostInfo {
|
||||||
return c.f.inside
|
|
||||||
}
|
|
||||||
|
|
||||||
func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
|
||||||
chi := ControlHostInfo{
|
chi := ControlHostInfo{
|
||||||
VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)),
|
VpnIp: h.vpnIp.ToIP(),
|
||||||
LocalIndex: h.localIndexId,
|
LocalIndex: h.localIndexId,
|
||||||
RemoteIndex: h.remoteIndexId,
|
RemoteIndex: h.remoteIndexId,
|
||||||
RemoteAddrs: h.remotes.CopyAddrs(preferredRanges),
|
RemoteAddrs: h.remotes.CopyAddrs(preferredRanges),
|
||||||
|
CachedPackets: len(h.packetStore),
|
||||||
CurrentRelaysToMe: h.relayState.CopyRelayIps(),
|
CurrentRelaysToMe: h.relayState.CopyRelayIps(),
|
||||||
CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(),
|
CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(),
|
||||||
CurrentRemote: h.remote,
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, a := range h.vpnAddrs {
|
|
||||||
chi.VpnAddrs[i] = a
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if h.ConnectionState != nil {
|
if h.ConnectionState != nil {
|
||||||
@@ -366,26 +231,38 @@ func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if c := h.GetCert(); c != nil {
|
if c := h.GetCert(); c != nil {
|
||||||
chi.Cert = c.Certificate.Copy()
|
chi.Cert = c.Copy()
|
||||||
|
}
|
||||||
|
|
||||||
|
if h.remote != nil {
|
||||||
|
chi.CurrentRemote = h.remote.Copy()
|
||||||
}
|
}
|
||||||
|
|
||||||
return chi
|
return chi
|
||||||
}
|
}
|
||||||
|
|
||||||
func listHostMapHosts(hl controlHostLister) []ControlHostInfo {
|
func listHostMapHosts(hm *HostMap) []ControlHostInfo {
|
||||||
hosts := make([]ControlHostInfo, 0)
|
hm.RLock()
|
||||||
pr := hl.GetPreferredRanges()
|
hosts := make([]ControlHostInfo, len(hm.Hosts))
|
||||||
hl.ForEachVpnAddr(func(hostinfo *HostInfo) {
|
i := 0
|
||||||
hosts = append(hosts, copyHostInfo(hostinfo, pr))
|
for _, v := range hm.Hosts {
|
||||||
})
|
hosts[i] = copyHostInfo(v, hm.preferredRanges)
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
hm.RUnlock()
|
||||||
|
|
||||||
return hosts
|
return hosts
|
||||||
}
|
}
|
||||||
|
|
||||||
func listHostMapIndexes(hl controlHostLister) []ControlHostInfo {
|
func listHostMapIndexes(hm *HostMap) []ControlHostInfo {
|
||||||
hosts := make([]ControlHostInfo, 0)
|
hm.RLock()
|
||||||
pr := hl.GetPreferredRanges()
|
hosts := make([]ControlHostInfo, len(hm.Indexes))
|
||||||
hl.ForEachIndex(func(hostinfo *HostInfo) {
|
i := 0
|
||||||
hosts = append(hosts, copyHostInfo(hostinfo, pr))
|
for _, v := range hm.Indexes {
|
||||||
})
|
hosts[i] = copyHostInfo(v, hm.preferredRanges)
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
hm.RUnlock()
|
||||||
|
|
||||||
return hosts
|
return hosts
|
||||||
}
|
}
|
||||||
|
|||||||
+51
-47
@@ -2,66 +2,71 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
|
||||||
"reflect"
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
||||||
//TODO: CERT-V2 with multiple certificate versions we have a problem with this test
|
|
||||||
// Some certs versions have different characteristics and each version implements their own Copy() func
|
|
||||||
// which means this is not a good place to test for exposing memory
|
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Special care must be taken to re-use all objects provided to the hostmap and certificate in the expectedInfo object
|
// Special care must be taken to re-use all objects provided to the hostmap and certificate in the expectedInfo object
|
||||||
// To properly ensure we are not exposing core memory to the caller
|
// To properly ensure we are not exposing core memory to the caller
|
||||||
hm := newHostMap(l)
|
hm := NewHostMap(l, "test", &net.IPNet{}, make([]*net.IPNet, 0))
|
||||||
hm.preferredRanges.Store(&[]netip.Prefix{})
|
remote1 := udp.NewAddr(net.ParseIP("0.0.0.100"), 4444)
|
||||||
|
remote2 := udp.NewAddr(net.ParseIP("1:2:3:4:5:6:7:8"), 4444)
|
||||||
remote1 := netip.MustParseAddrPort("0.0.0.100:4444")
|
|
||||||
remote2 := netip.MustParseAddrPort("[1:2:3:4:5:6:7:8]:4444")
|
|
||||||
|
|
||||||
ipNet := net.IPNet{
|
ipNet := net.IPNet{
|
||||||
IP: remote1.Addr().AsSlice(),
|
IP: net.IPv4(1, 2, 3, 4),
|
||||||
Mask: net.IPMask{255, 255, 255, 0},
|
Mask: net.IPMask{255, 255, 255, 0},
|
||||||
}
|
}
|
||||||
|
|
||||||
ipNet2 := net.IPNet{
|
ipNet2 := net.IPNet{
|
||||||
IP: remote2.Addr().AsSlice(),
|
IP: net.ParseIP("1:2:3:4:5:6:7:8"),
|
||||||
Mask: net.IPMask{255, 255, 255, 0},
|
Mask: net.IPMask{255, 255, 255, 0},
|
||||||
}
|
}
|
||||||
|
|
||||||
remotes := NewRemoteList([]netip.Addr{netip.IPv4Unspecified()}, nil)
|
crt := &cert.NebulaCertificate{
|
||||||
remotes.unlockedPrependV4(netip.IPv4Unspecified(), netAddrToProtoV4AddrPort(remote1.Addr(), remote1.Port()))
|
Details: cert.NebulaCertificateDetails{
|
||||||
remotes.unlockedPrependV6(netip.IPv4Unspecified(), netAddrToProtoV6AddrPort(remote2.Addr(), remote2.Port()))
|
Name: "test",
|
||||||
|
Ips: []*net.IPNet{&ipNet},
|
||||||
|
Subnets: []*net.IPNet{},
|
||||||
|
Groups: []string{"default-group"},
|
||||||
|
NotBefore: time.Unix(1, 0),
|
||||||
|
NotAfter: time.Unix(2, 0),
|
||||||
|
PublicKey: []byte{5, 6, 7, 8},
|
||||||
|
IsCA: false,
|
||||||
|
Issuer: "the-issuer",
|
||||||
|
InvertedGroups: map[string]struct{}{"default-group": {}},
|
||||||
|
},
|
||||||
|
Signature: []byte{1, 2, 1, 2, 1, 3},
|
||||||
|
}
|
||||||
|
|
||||||
vpnIp, ok := netip.AddrFromSlice(ipNet.IP)
|
remotes := NewRemoteList(nil)
|
||||||
assert.True(t, ok)
|
remotes.unlockedPrependV4(0, NewIp4AndPort(remote1.IP, uint32(remote1.Port)))
|
||||||
|
remotes.unlockedPrependV6(0, NewIp6AndPort(remote2.IP, uint32(remote2.Port)))
|
||||||
crt := &dummyCert{}
|
hm.Add(iputil.Ip2VpnIp(ipNet.IP), &HostInfo{
|
||||||
hm.unlockedAddHostInfo(&HostInfo{
|
|
||||||
remote: remote1,
|
remote: remote1,
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
peerCert: &cert.CachedCertificate{Certificate: crt},
|
peerCert: crt,
|
||||||
},
|
},
|
||||||
remoteIndexId: 200,
|
remoteIndexId: 200,
|
||||||
localIndexId: 201,
|
localIndexId: 201,
|
||||||
vpnAddrs: []netip.Addr{vpnIp},
|
vpnIp: iputil.Ip2VpnIp(ipNet.IP),
|
||||||
relayState: RelayState{
|
relayState: RelayState{
|
||||||
relays: nil,
|
relays: map[iputil.VpnIp]struct{}{},
|
||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
relayForByIp: map[iputil.VpnIp]*Relay{},
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
},
|
},
|
||||||
}, &Interface{})
|
})
|
||||||
|
|
||||||
vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP)
|
hm.Add(iputil.Ip2VpnIp(ipNet2.IP), &HostInfo{
|
||||||
assert.True(t, ok)
|
|
||||||
|
|
||||||
hm.unlockedAddHostInfo(&HostInfo{
|
|
||||||
remote: remote1,
|
remote: remote1,
|
||||||
remotes: remotes,
|
remotes: remotes,
|
||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
@@ -69,48 +74,47 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
|||||||
},
|
},
|
||||||
remoteIndexId: 200,
|
remoteIndexId: 200,
|
||||||
localIndexId: 201,
|
localIndexId: 201,
|
||||||
vpnAddrs: []netip.Addr{vpnIp2},
|
vpnIp: iputil.Ip2VpnIp(ipNet2.IP),
|
||||||
relayState: RelayState{
|
relayState: RelayState{
|
||||||
relays: nil,
|
relays: map[iputil.VpnIp]struct{}{},
|
||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
relayForByIp: map[iputil.VpnIp]*Relay{},
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
},
|
},
|
||||||
}, &Interface{})
|
})
|
||||||
|
|
||||||
c := Control{
|
c := Control{
|
||||||
state: StateReady,
|
|
||||||
f: &Interface{
|
f: &Interface{
|
||||||
hostMap: hm,
|
hostMap: hm,
|
||||||
},
|
},
|
||||||
l: test.NewLogger(),
|
l: logrus.New(),
|
||||||
}
|
}
|
||||||
|
|
||||||
thi := c.GetHostInfoByVpnAddr(vpnIp, false)
|
thi := c.GetHostInfoByVpnIp(iputil.Ip2VpnIp(ipNet.IP), false)
|
||||||
|
|
||||||
expectedInfo := ControlHostInfo{
|
expectedInfo := ControlHostInfo{
|
||||||
VpnAddrs: []netip.Addr{vpnIp},
|
VpnIp: net.IPv4(1, 2, 3, 4).To4(),
|
||||||
LocalIndex: 201,
|
LocalIndex: 201,
|
||||||
RemoteIndex: 200,
|
RemoteIndex: 200,
|
||||||
RemoteAddrs: []netip.AddrPort{remote2, remote1},
|
RemoteAddrs: []*udp.Addr{remote2, remote1},
|
||||||
|
CachedPackets: 0,
|
||||||
Cert: crt.Copy(),
|
Cert: crt.Copy(),
|
||||||
MessageCounter: 0,
|
MessageCounter: 0,
|
||||||
CurrentRemote: remote1,
|
CurrentRemote: udp.NewAddr(net.ParseIP("0.0.0.100"), 4444),
|
||||||
CurrentRelaysToMe: []netip.Addr{},
|
CurrentRelaysToMe: []iputil.VpnIp{},
|
||||||
CurrentRelaysThroughMe: []netip.Addr{},
|
CurrentRelaysThroughMe: []iputil.VpnIp{},
|
||||||
}
|
}
|
||||||
|
|
||||||
// Make sure we don't have any unexpected fields
|
// Make sure we don't have any unexpected fields
|
||||||
assertFields(t, []string{"VpnAddrs", "LocalIndex", "RemoteIndex", "RemoteAddrs", "Cert", "MessageCounter", "CurrentRemote", "CurrentRelaysToMe", "CurrentRelaysThroughMe"}, thi)
|
assertFields(t, []string{"VpnIp", "LocalIndex", "RemoteIndex", "RemoteAddrs", "CachedPackets", "Cert", "MessageCounter", "CurrentRemote", "CurrentRelaysToMe", "CurrentRelaysThroughMe"}, thi)
|
||||||
assert.Equal(t, &expectedInfo, thi)
|
|
||||||
test.AssertDeepCopyEqual(t, &expectedInfo, thi)
|
test.AssertDeepCopyEqual(t, &expectedInfo, thi)
|
||||||
|
|
||||||
// Make sure we don't panic if the host info doesn't have a cert yet
|
// Make sure we don't panic if the host info doesn't have a cert yet
|
||||||
assert.NotPanics(t, func() {
|
assert.NotPanics(t, func() {
|
||||||
thi = c.GetHostInfoByVpnAddr(vpnIp2, false)
|
thi = c.GetHostInfoByVpnIp(iputil.Ip2VpnIp(ipNet2.IP), false)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func assertFields(t *testing.T, expected []string, actualStruct any) {
|
func assertFields(t *testing.T, expected []string, actualStruct interface{}) {
|
||||||
val := reflect.ValueOf(actualStruct).Elem()
|
val := reflect.ValueOf(actualStruct).Elem()
|
||||||
fields := make([]string, val.NumField())
|
fields := make([]string, val.NumField())
|
||||||
for i := 0; i < val.NumField(); i++ {
|
for i := 0; i < val.NumField(); i++ {
|
||||||
|
|||||||
+50
-57
@@ -1,13 +1,17 @@
|
|||||||
//go:build e2e_testing
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/netip"
|
"net"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
@@ -46,30 +50,37 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
|||||||
|
|
||||||
// InjectLightHouseAddr will push toAddr into the local lighthouse cache for the vpnIp
|
// InjectLightHouseAddr will push toAddr into the local lighthouse cache for the vpnIp
|
||||||
// This is necessary if you did not configure static hosts or are not running a lighthouse
|
// This is necessary if you did not configure static hosts or are not running a lighthouse
|
||||||
func (c *Control) InjectLightHouseAddr(vpnIp netip.Addr, toAddr netip.AddrPort) {
|
func (c *Control) InjectLightHouseAddr(vpnIp net.IP, toAddr *net.UDPAddr) {
|
||||||
c.f.lightHouse.Lock()
|
c.f.lightHouse.Lock()
|
||||||
remoteList := c.f.lightHouse.unlockedGetRemoteList([]netip.Addr{vpnIp})
|
remoteList := c.f.lightHouse.unlockedGetRemoteList(iputil.Ip2VpnIp(vpnIp))
|
||||||
remoteList.Lock()
|
remoteList.Lock()
|
||||||
defer remoteList.Unlock()
|
defer remoteList.Unlock()
|
||||||
c.f.lightHouse.Unlock()
|
c.f.lightHouse.Unlock()
|
||||||
|
|
||||||
if toAddr.Addr().Is4() {
|
iVpnIp := iputil.Ip2VpnIp(vpnIp)
|
||||||
remoteList.unlockedPrependV4(vpnIp, netAddrToProtoV4AddrPort(toAddr.Addr(), toAddr.Port()))
|
if v4 := toAddr.IP.To4(); v4 != nil {
|
||||||
|
remoteList.unlockedPrependV4(iVpnIp, NewIp4AndPort(v4, uint32(toAddr.Port)))
|
||||||
} else {
|
} else {
|
||||||
remoteList.unlockedPrependV6(vpnIp, netAddrToProtoV6AddrPort(toAddr.Addr(), toAddr.Port()))
|
remoteList.unlockedPrependV6(iVpnIp, NewIp6AndPort(toAddr.IP, uint32(toAddr.Port)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectRelays will push relayVpnIps into the local lighthouse cache for the vpnIp
|
// InjectRelays will push relayVpnIps into the local lighthouse cache for the vpnIp
|
||||||
// This is necessary to inform an initiator of possible relays for communicating with a responder
|
// This is necessary to inform an initiator of possible relays for communicating with a responder
|
||||||
func (c *Control) InjectRelays(vpnIp netip.Addr, relayVpnIps []netip.Addr) {
|
func (c *Control) InjectRelays(vpnIp net.IP, relayVpnIps []net.IP) {
|
||||||
c.f.lightHouse.Lock()
|
c.f.lightHouse.Lock()
|
||||||
remoteList := c.f.lightHouse.unlockedGetRemoteList([]netip.Addr{vpnIp})
|
remoteList := c.f.lightHouse.unlockedGetRemoteList(iputil.Ip2VpnIp(vpnIp))
|
||||||
remoteList.Lock()
|
remoteList.Lock()
|
||||||
defer remoteList.Unlock()
|
defer remoteList.Unlock()
|
||||||
c.f.lightHouse.Unlock()
|
c.f.lightHouse.Unlock()
|
||||||
|
|
||||||
remoteList.unlockedSetRelay(vpnIp, relayVpnIps)
|
iVpnIp := iputil.Ip2VpnIp(vpnIp)
|
||||||
|
uVpnIp := []uint32{}
|
||||||
|
for _, rVPnIp := range relayVpnIps {
|
||||||
|
uVpnIp = append(uVpnIp, uint32(iputil.Ip2VpnIp(rVPnIp)))
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteList.unlockedSetRelay(iVpnIp, iVpnIp, uVpnIp)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetFromTun will pull a packet off the tun side of nebula
|
// GetFromTun will pull a packet off the tun side of nebula
|
||||||
@@ -96,42 +107,20 @@ func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
||||||
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
func (c *Control) InjectTunUDPPacket(toIp net.IP, toPort uint16, fromPort uint16, data []byte) {
|
||||||
serialize := make([]gopacket.SerializableLayer, 0)
|
ip := layers.IPv4{
|
||||||
var netLayer gopacket.NetworkLayer
|
Version: 4,
|
||||||
if toAddr.Is6() {
|
TTL: 64,
|
||||||
if !fromAddr.Is6() {
|
Protocol: layers.IPProtocolUDP,
|
||||||
panic("Cant send ipv6 to ipv4")
|
SrcIP: c.f.inside.Cidr().IP,
|
||||||
}
|
DstIP: toIp,
|
||||||
ip := &layers.IPv6{
|
|
||||||
Version: 6,
|
|
||||||
NextHeader: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
} else {
|
|
||||||
if !fromAddr.Is4() {
|
|
||||||
panic("Cant send ipv4 to ipv6")
|
|
||||||
}
|
|
||||||
|
|
||||||
ip := &layers.IPv4{
|
|
||||||
Version: 4,
|
|
||||||
TTL: 64,
|
|
||||||
Protocol: layers.IPProtocolUDP,
|
|
||||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
|
||||||
DstIP: toAddr.Unmap().AsSlice(),
|
|
||||||
}
|
|
||||||
serialize = append(serialize, ip)
|
|
||||||
netLayer = ip
|
|
||||||
}
|
}
|
||||||
|
|
||||||
udp := layers.UDP{
|
udp := layers.UDP{
|
||||||
SrcPort: layers.UDPPort(fromPort),
|
SrcPort: layers.UDPPort(fromPort),
|
||||||
DstPort: layers.UDPPort(toPort),
|
DstPort: layers.UDPPort(toPort),
|
||||||
}
|
}
|
||||||
err := udp.SetNetworkLayerForChecksum(netLayer)
|
err := udp.SetNetworkLayerForChecksum(&ip)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
@@ -141,9 +130,7 @@ func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr
|
|||||||
ComputeChecksums: true,
|
ComputeChecksums: true,
|
||||||
FixLengths: true,
|
FixLengths: true,
|
||||||
}
|
}
|
||||||
|
err = gopacket.SerializeLayers(buffer, opt, &ip, &udp, gopacket.Payload(data))
|
||||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
|
||||||
err = gopacket.SerializeLayers(buffer, opt, serialize...)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
@@ -151,21 +138,21 @@ func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr
|
|||||||
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
func (c *Control) GetVpnIp() iputil.VpnIp {
|
||||||
return c.f.myVpnAddrs
|
return c.f.myVpnIp
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetUDPAddr() netip.AddrPort {
|
func (c *Control) GetUDPAddr() string {
|
||||||
return c.f.outside.(*udp.TesterConn).Addr
|
return c.f.outside.(*udp.TesterConn).Addr.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
|
func (c *Control) KillPendingTunnel(vpnIp net.IP) bool {
|
||||||
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
|
hostinfo, ok := c.f.handshakeManager.pendingHostMap.Hosts[iputil.Ip2VpnIp(vpnIp)]
|
||||||
if hostinfo == nil {
|
if !ok {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
c.f.handshakeManager.DeleteHostInfo(hostinfo)
|
c.f.handshakeManager.pendingHostMap.DeleteHostInfo(hostinfo)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -173,14 +160,20 @@ func (c *Control) GetHostmap() *HostMap {
|
|||||||
return c.f.hostMap
|
return c.f.hostMap
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetF() *Interface {
|
func (c *Control) GetCert() *cert.NebulaCertificate {
|
||||||
return c.f
|
return c.f.certState.Load().certificate
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetCertState() *CertState {
|
func (c *Control) ReHandshake(vpnIp iputil.VpnIp) {
|
||||||
return c.f.pki.getCertState()
|
hostinfo := c.f.handshakeManager.AddVpnIp(vpnIp, c.f.initHostInfo)
|
||||||
}
|
ixHandshakeStage0(c.f, vpnIp, hostinfo)
|
||||||
|
|
||||||
func (c *Control) ReHandshake(vpnIp netip.Addr) {
|
// If this is a static host, we don't need to wait for the HostQueryReply
|
||||||
c.f.handshakeManager.StartHandshake(vpnIp, nil)
|
// We can trigger the handshake right now
|
||||||
|
if _, ok := c.f.lightHouse.GetStaticHostList()[hostinfo.vpnIp]; ok {
|
||||||
|
select {
|
||||||
|
case c.f.handshakeManager.trigger <- hostinfo.vpnIp:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+13
@@ -0,0 +1,13 @@
|
|||||||
|
[Unit]
|
||||||
|
Description=Nebula overlay networking tool
|
||||||
|
Wants=basic.target network-online.target nss-lookup.target time-sync.target
|
||||||
|
After=basic.target network.target network-online.target
|
||||||
|
|
||||||
|
[Service]
|
||||||
|
SyslogIdentifier=nebula
|
||||||
|
ExecReload=/bin/kill -HUP $MAINPID
|
||||||
|
ExecStart=/usr/bin/nebula -config /etc/nebula/config.yml
|
||||||
|
Restart=always
|
||||||
|
|
||||||
|
[Install]
|
||||||
|
WantedBy=multi-user.target
|
||||||
Vendored
+14
@@ -0,0 +1,14 @@
|
|||||||
|
[Unit]
|
||||||
|
Description=Nebula overlay networking tool
|
||||||
|
Wants=basic.target network-online.target nss-lookup.target time-sync.target
|
||||||
|
After=basic.target network.target network-online.target
|
||||||
|
Before=sshd.service
|
||||||
|
|
||||||
|
[Service]
|
||||||
|
SyslogIdentifier=nebula
|
||||||
|
ExecReload=/bin/kill -HUP $MAINPID
|
||||||
|
ExecStart=/usr/bin/nebula -config /etc/nebula/config.yml
|
||||||
|
Restart=always
|
||||||
|
|
||||||
|
[Install]
|
||||||
|
WantedBy=multi-user.target
|
||||||
Vendored
+14
-8
@@ -84,24 +84,30 @@ end
|
|||||||
|
|
||||||
function nebula.prefs_changed()
|
function nebula.prefs_changed()
|
||||||
if default_settings.all_ports == nebula.prefs.all_ports and default_settings.port == nebula.prefs.port then
|
if default_settings.all_ports == nebula.prefs.all_ports and default_settings.port == nebula.prefs.port then
|
||||||
|
-- Nothing changed, bail
|
||||||
return
|
return
|
||||||
end
|
end
|
||||||
|
|
||||||
-- Remove all existing registrations
|
-- Remove our old dissector
|
||||||
DissectorTable.get("udp.port"):remove_all(nebula)
|
DissectorTable.get("udp.port"):remove_all(nebula)
|
||||||
|
|
||||||
if nebula.prefs.all_ports then
|
if nebula.prefs.all_ports and default_settings.all_ports ~= nebula.prefs.all_ports then
|
||||||
-- Register on every port for hole punch capture
|
default_settings.all_port = nebula.prefs.all_ports
|
||||||
|
|
||||||
for i=0, 65535 do
|
for i=0, 65535 do
|
||||||
DissectorTable.get("udp.port"):add(i, nebula)
|
DissectorTable.get("udp.port"):add(i, nebula)
|
||||||
end
|
end
|
||||||
else
|
|
||||||
-- Register on the configured port only
|
-- no need to establish again on specific ports
|
||||||
DissectorTable.get("udp.port"):add(nebula.prefs.port, nebula)
|
return
|
||||||
end
|
end
|
||||||
|
|
||||||
default_settings.all_ports = nebula.prefs.all_ports
|
|
||||||
default_settings.port = nebula.prefs.port
|
if default_settings.all_ports ~= nebula.prefs.all_ports then
|
||||||
|
-- Add our new port dissector
|
||||||
|
default_settings.port = nebula.prefs.port
|
||||||
|
DissectorTable.get("udp.port"):add(default_settings.port, nebula)
|
||||||
|
end
|
||||||
end
|
end
|
||||||
|
|
||||||
DissectorTable.get("udp.port"):add(default_settings.port, nebula)
|
DissectorTable.get("udp.port"):add(default_settings.port, nebula)
|
||||||
|
|||||||
+84
-312
@@ -1,351 +1,93 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/iputil"
|
||||||
)
|
)
|
||||||
|
|
||||||
type dnsServer struct {
|
// This whole thing should be rewritten to use context
|
||||||
|
|
||||||
|
var dnsR *dnsRecords
|
||||||
|
var dnsServer *dns.Server
|
||||||
|
var dnsAddr string
|
||||||
|
|
||||||
|
type dnsRecords struct {
|
||||||
sync.RWMutex
|
sync.RWMutex
|
||||||
l *slog.Logger
|
dnsMap map[string]string
|
||||||
ctx context.Context
|
hostMap *HostMap
|
||||||
dnsMap4 map[string]netip.Addr
|
|
||||||
dnsMap6 map[string]netip.Addr
|
|
||||||
hostMap *HostMap
|
|
||||||
myVpnAddrsTable *bart.Lite
|
|
||||||
|
|
||||||
mux *dns.ServeMux
|
|
||||||
|
|
||||||
// enabled mirrors `lighthouse.serve_dns && lighthouse.am_lighthouse`.
|
|
||||||
// Start, Add, and reload consult it so callers don't need to know the
|
|
||||||
// gating rules. When it toggles off via reload, accumulated records are
|
|
||||||
// cleared so a later re-enable starts with a fresh map populated from
|
|
||||||
// new handshakes.
|
|
||||||
enabled atomic.Bool
|
|
||||||
|
|
||||||
serverMu sync.Mutex
|
|
||||||
server *dns.Server
|
|
||||||
// started is closed once `server` has finished binding (or after
|
|
||||||
// ListenAndServe returns on a bind failure). Stop waits on it before
|
|
||||||
// calling Shutdown to avoid the miekg/dns "server not started" race
|
|
||||||
// where a Shutdown that arrives before bind completes is silently
|
|
||||||
// ignored, leaving the listener running forever.
|
|
||||||
started chan struct{}
|
|
||||||
addr string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// newDnsServerFromConfig builds a dnsServer, applies the initial config, and
|
func newDnsRecords(hostMap *HostMap) *dnsRecords {
|
||||||
// registers a reload callback. The reload callback is registered before the
|
return &dnsRecords{
|
||||||
// initial config is applied, so a SIGHUP can later enable, fix, or disable
|
dnsMap: make(map[string]string),
|
||||||
// DNS even if the initial application failed.
|
hostMap: hostMap,
|
||||||
//
|
|
||||||
// The dnsServer internally gates on `lighthouse.serve_dns &&
|
|
||||||
// lighthouse.am_lighthouse`. Start and Add are safe to call unconditionally,
|
|
||||||
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
|
|
||||||
// watcher that tears the listener down on nebula shutdown. The returned
|
|
||||||
// pointer is always non-nil, even on error.
|
|
||||||
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
|
||||||
ds := &dnsServer{
|
|
||||||
l: l,
|
|
||||||
ctx: ctx,
|
|
||||||
dnsMap4: make(map[string]netip.Addr),
|
|
||||||
dnsMap6: make(map[string]netip.Addr),
|
|
||||||
hostMap: hostMap,
|
|
||||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
|
||||||
}
|
|
||||||
ds.mux = dns.NewServeMux()
|
|
||||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
|
||||||
|
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
|
||||||
if err := ds.reload(c, false); err != nil {
|
|
||||||
ds.l.Error("Failed to reload DNS responder from config", "error", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
if err := ds.reload(c, true); err != nil {
|
|
||||||
return ds, err
|
|
||||||
}
|
|
||||||
return ds, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// reload applies the latest config and reconciles the running state with it:
|
|
||||||
// - enabled toggled on -> spawn a runner
|
|
||||||
// - enabled toggled off -> stop the runner
|
|
||||||
// - listen address changed (while running) -> restart on the new address
|
|
||||||
// - everything else -> no-op
|
|
||||||
//
|
|
||||||
// On the initial call it only records configuration; Control.Start is what
|
|
||||||
// launches the first runner via dnsStart.
|
|
||||||
func (d *dnsServer) reload(c *config.C, initial bool) error {
|
|
||||||
wantsDns := c.GetBool("lighthouse.serve_dns", false)
|
|
||||||
amLighthouse := c.GetBool("lighthouse.am_lighthouse", false)
|
|
||||||
enabled := wantsDns && amLighthouse
|
|
||||||
newAddr := getDnsServerAddr(c)
|
|
||||||
|
|
||||||
d.serverMu.Lock()
|
|
||||||
running := d.server
|
|
||||||
runningStarted := d.started
|
|
||||||
sameAddr := d.addr == newAddr
|
|
||||||
d.addr = newAddr
|
|
||||||
d.enabled.Store(enabled)
|
|
||||||
d.serverMu.Unlock()
|
|
||||||
|
|
||||||
if initial {
|
|
||||||
if wantsDns && !amLighthouse {
|
|
||||||
d.l.Warn("DNS server refusing to run because this host is not a lighthouse.")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if !enabled {
|
|
||||||
if running != nil {
|
|
||||||
d.Stop()
|
|
||||||
}
|
|
||||||
// Drop any records that accumulated while enabled; a later re-enable
|
|
||||||
// will repopulate from fresh handshakes.
|
|
||||||
d.clearRecords()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if running == nil {
|
|
||||||
// Was disabled (or never started); bring it up now.
|
|
||||||
go d.Start()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if sameAddr {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
d.shutdownServer(running, runningStarted, "reload")
|
|
||||||
// Old Start goroutine has now exited; bring up a fresh listener on the
|
|
||||||
// new address.
|
|
||||||
go d.Start()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// shutdownServer waits for the server to finish binding (so Shutdown actually
|
|
||||||
// stops it rather than no-oping) and then shuts it down.
|
|
||||||
func (d *dnsServer) shutdownServer(srv *dns.Server, started chan struct{}, reason string) {
|
|
||||||
if srv == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if started != nil {
|
|
||||||
<-started
|
|
||||||
}
|
|
||||||
if err := srv.Shutdown(); err != nil {
|
|
||||||
d.l.Warn("Failed to shut down the DNS responder", "reason", reason, "error", err)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Start binds and serves the DNS responder. Blocks until Stop is called or
|
func (d *dnsRecords) Query(data string) string {
|
||||||
// the listener errors. Safe to call when DNS is disabled (returns
|
|
||||||
// immediately). This is what Control.dnsStart points at.
|
|
||||||
//
|
|
||||||
// Must be invoked after the tun device is active so that lighthouse.dns.host
|
|
||||||
// may bind to a nebula IP.
|
|
||||||
func (d *dnsServer) Start() {
|
|
||||||
if !d.enabled.Load() {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
started := make(chan struct{})
|
|
||||||
d.serverMu.Lock()
|
|
||||||
if d.ctx.Err() != nil {
|
|
||||||
d.serverMu.Unlock()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
addr := d.addr
|
|
||||||
server := &dns.Server{
|
|
||||||
Addr: addr,
|
|
||||||
Net: "udp",
|
|
||||||
Handler: d.mux,
|
|
||||||
NotifyStartedFunc: func() { close(started) },
|
|
||||||
}
|
|
||||||
d.server = server
|
|
||||||
d.started = started
|
|
||||||
d.serverMu.Unlock()
|
|
||||||
|
|
||||||
// Per-invocation ctx watcher. Exits when Start does, so we don't leak a
|
|
||||||
// watcher per reload-driven restart.
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
select {
|
|
||||||
case <-d.ctx.Done():
|
|
||||||
d.shutdownServer(server, started, "shutdown")
|
|
||||||
case <-done:
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
d.l.Info("Starting DNS responder", "dnsListener", addr)
|
|
||||||
err := server.ListenAndServe()
|
|
||||||
close(done)
|
|
||||||
|
|
||||||
// If the listener never bound (bind error) NotifyStartedFunc never fires,
|
|
||||||
// so close started here to release any Stop caller waiting on it.
|
|
||||||
select {
|
|
||||||
case <-started:
|
|
||||||
default:
|
|
||||||
close(started)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stop shuts down the active server, if any. Idempotent.
|
|
||||||
func (d *dnsServer) Stop() {
|
|
||||||
d.serverMu.Lock()
|
|
||||||
srv := d.server
|
|
||||||
started := d.started
|
|
||||||
d.server = nil
|
|
||||||
d.started = nil
|
|
||||||
d.serverMu.Unlock()
|
|
||||||
d.shutdownServer(srv, started, "stop")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Query returns the address for the given name and query type. The second
|
|
||||||
// return value reports whether the name is known at all (in either A or AAAA),
|
|
||||||
// which lets callers distinguish NODATA from NXDOMAIN.
|
|
||||||
func (d *dnsServer) Query(q uint16, data string) (netip.Addr, bool) {
|
|
||||||
data = strings.ToLower(data)
|
|
||||||
d.RLock()
|
d.RLock()
|
||||||
defer d.RUnlock()
|
defer d.RUnlock()
|
||||||
addr4, haveV4 := d.dnsMap4[data]
|
if r, ok := d.dnsMap[strings.ToLower(data)]; ok {
|
||||||
addr6, haveV6 := d.dnsMap6[data]
|
return r
|
||||||
nameExists := haveV4 || haveV6
|
|
||||||
switch q {
|
|
||||||
case dns.TypeA:
|
|
||||||
if haveV4 {
|
|
||||||
return addr4, nameExists
|
|
||||||
}
|
|
||||||
case dns.TypeAAAA:
|
|
||||||
if haveV6 {
|
|
||||||
return addr6, nameExists
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
return ""
|
||||||
return netip.Addr{}, nameExists
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) QueryCert(data string) string {
|
func (d *dnsRecords) QueryCert(data string) string {
|
||||||
if len(data) < 2 {
|
ip := net.ParseIP(data[:len(data)-1])
|
||||||
|
if ip == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
ip, err := netip.ParseAddr(data[:len(data)-1])
|
iip := iputil.Ip2VpnIp(ip)
|
||||||
|
hostinfo, err := d.hostMap.QueryVpnIp(iip)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo := d.hostMap.QueryVpnAddr(ip)
|
|
||||||
if hostinfo == nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
q := hostinfo.GetCert()
|
q := hostinfo.GetCert()
|
||||||
if q == nil {
|
if q == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
cert := q.Details
|
||||||
b, err := q.Certificate.MarshalJSON()
|
c := fmt.Sprintf("\"Name: %s\" \"Ips: %s\" \"Subnets %s\" \"Groups %s\" \"NotBefore %s\" \"NotAFter %s\" \"PublicKey %x\" \"IsCA %t\" \"Issuer %s\"", cert.Name, cert.Ips, cert.Subnets, cert.Groups, cert.NotBefore, cert.NotAfter, cert.PublicKey, cert.IsCA, cert.Issuer)
|
||||||
if err != nil {
|
return c
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return string(b)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearRecords drops all DNS records.
|
func (d *dnsRecords) Add(host, data string) {
|
||||||
func (d *dnsServer) clearRecords() {
|
|
||||||
d.Lock()
|
d.Lock()
|
||||||
defer d.Unlock()
|
defer d.Unlock()
|
||||||
clear(d.dnsMap4)
|
d.dnsMap[strings.ToLower(host)] = data
|
||||||
clear(d.dnsMap6)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
func parseQuery(l *logrus.Logger, m *dns.Msg, w dns.ResponseWriter) {
|
||||||
func (d *dnsServer) Add(host string, addresses []netip.Addr) {
|
|
||||||
if !d.enabled.Load() {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
host = strings.ToLower(host)
|
|
||||||
d.Lock()
|
|
||||||
defer d.Unlock()
|
|
||||||
haveV4 := false
|
|
||||||
haveV6 := false
|
|
||||||
for _, addr := range addresses {
|
|
||||||
if addr.Is4() && !haveV4 {
|
|
||||||
d.dnsMap4[host] = addr
|
|
||||||
haveV4 = true
|
|
||||||
} else if addr.Is6() && !haveV6 {
|
|
||||||
d.dnsMap6[host] = addr
|
|
||||||
haveV6 = true
|
|
||||||
}
|
|
||||||
if haveV4 && haveV6 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dnsServer) isSelfNebulaOrLocalhost(addr string) bool {
|
|
||||||
a, _, _ := net.SplitHostPort(addr)
|
|
||||||
b, err := netip.ParseAddr(a)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if b.IsLoopback() {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
//if we found it in this table, it's good
|
|
||||||
return d.myVpnAddrsTable.Contains(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
|
||||||
debugEnabled := d.l.Enabled(context.Background(), slog.LevelDebug)
|
|
||||||
// Per RFC 2308 §2.2, a name that exists but has no record of the requested
|
|
||||||
// type must be answered with NOERROR and an empty answer section (NODATA),
|
|
||||||
// not NXDOMAIN (RFC 2308 §2.1), which is reserved for names that do not
|
|
||||||
// exist at all.
|
|
||||||
anyNameExists := false
|
|
||||||
for _, q := range m.Question {
|
for _, q := range m.Question {
|
||||||
switch q.Qtype {
|
switch q.Qtype {
|
||||||
case dns.TypeA, dns.TypeAAAA:
|
case dns.TypeA:
|
||||||
qType := dns.TypeToString[q.Qtype]
|
l.Debugf("Query for A %s", q.Name)
|
||||||
if debugEnabled {
|
ip := dnsR.Query(q.Name)
|
||||||
d.l.Debug("DNS query", "type", qType, "name", q.Name)
|
if ip != "" {
|
||||||
}
|
rr, err := dns.NewRR(fmt.Sprintf("%s A %s", q.Name, ip))
|
||||||
ip, nameExists := d.Query(q.Qtype, q.Name)
|
|
||||||
if nameExists {
|
|
||||||
anyNameExists = true
|
|
||||||
}
|
|
||||||
if ip.IsValid() {
|
|
||||||
rr, err := dns.NewRR(fmt.Sprintf("%s %s %s", q.Name, qType, ip))
|
|
||||||
if err == nil {
|
if err == nil {
|
||||||
m.Answer = append(m.Answer, rr)
|
m.Answer = append(m.Answer, rr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
case dns.TypeTXT:
|
case dns.TypeTXT:
|
||||||
// We only answer these queries from nebula nodes or localhost
|
a, _, _ := net.SplitHostPort(w.RemoteAddr().String())
|
||||||
if !d.isSelfNebulaOrLocalhost(w.RemoteAddr().String()) {
|
b := net.ParseIP(a)
|
||||||
|
// We don't answer these queries from non nebula nodes or localhost
|
||||||
|
//l.Debugf("Does %s contain %s", b, dnsR.hostMap.vpnCIDR)
|
||||||
|
if !dnsR.hostMap.vpnCIDR.Contains(b) && a != "127.0.0.1" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if debugEnabled {
|
l.Debugf("Query for TXT %s", q.Name)
|
||||||
d.l.Debug("DNS query", "type", "TXT", "name", q.Name)
|
ip := dnsR.QueryCert(q.Name)
|
||||||
}
|
|
||||||
ip := d.QueryCert(q.Name)
|
|
||||||
if ip != "" {
|
if ip != "" {
|
||||||
rr, err := dns.NewRR(fmt.Sprintf("%s TXT %s", q.Name, ip))
|
rr, err := dns.NewRR(fmt.Sprintf("%s TXT %s", q.Name, ip))
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -354,30 +96,60 @@ func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(m.Answer) == 0 && !anyNameExists {
|
|
||||||
m.Rcode = dns.RcodeNameError
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *dnsServer) handleDnsRequest(w dns.ResponseWriter, r *dns.Msg) {
|
func handleDnsRequest(l *logrus.Logger, w dns.ResponseWriter, r *dns.Msg) {
|
||||||
m := new(dns.Msg)
|
m := new(dns.Msg)
|
||||||
m.SetReply(r)
|
m.SetReply(r)
|
||||||
m.Compress = false
|
m.Compress = false
|
||||||
|
|
||||||
switch r.Opcode {
|
switch r.Opcode {
|
||||||
case dns.OpcodeQuery:
|
case dns.OpcodeQuery:
|
||||||
d.parseQuery(m, w)
|
parseQuery(l, m, w)
|
||||||
}
|
}
|
||||||
|
|
||||||
w.WriteMsg(m)
|
w.WriteMsg(m)
|
||||||
}
|
}
|
||||||
|
|
||||||
func getDnsServerAddr(c *config.C) string {
|
func dnsMain(l *logrus.Logger, hostMap *HostMap, c *config.C) func() {
|
||||||
dnsHost := strings.TrimSpace(c.GetString("lighthouse.dns.host", ""))
|
dnsR = newDnsRecords(hostMap)
|
||||||
// Old guidance was to provide the literal `[::]` in `lighthouse.dns.host` but that won't resolve.
|
|
||||||
if dnsHost == "[::]" {
|
// attach request handler func
|
||||||
dnsHost = "::"
|
dns.HandleFunc(".", func(w dns.ResponseWriter, r *dns.Msg) {
|
||||||
|
handleDnsRequest(l, w, r)
|
||||||
|
})
|
||||||
|
|
||||||
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
|
reloadDns(l, c)
|
||||||
|
})
|
||||||
|
|
||||||
|
return func() {
|
||||||
|
startDns(l, c)
|
||||||
}
|
}
|
||||||
return net.JoinHostPort(dnsHost, strconv.Itoa(c.GetInt("lighthouse.dns.port", 53)))
|
}
|
||||||
|
|
||||||
|
func getDnsServerAddr(c *config.C) string {
|
||||||
|
return c.GetString("lighthouse.dns.host", "") + ":" + strconv.Itoa(c.GetInt("lighthouse.dns.port", 53))
|
||||||
|
}
|
||||||
|
|
||||||
|
func startDns(l *logrus.Logger, c *config.C) {
|
||||||
|
dnsAddr = getDnsServerAddr(c)
|
||||||
|
dnsServer = &dns.Server{Addr: dnsAddr, Net: "udp"}
|
||||||
|
l.WithField("dnsListener", dnsAddr).Info("Starting DNS responder")
|
||||||
|
err := dnsServer.ListenAndServe()
|
||||||
|
defer dnsServer.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
l.Errorf("Failed to start server: %s\n ", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func reloadDns(l *logrus.Logger, c *config.C) {
|
||||||
|
if dnsAddr == getDnsServerAddr(c) {
|
||||||
|
l.Debug("No DNS server config change detected")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
l.Debug("Restarting DNS server")
|
||||||
|
dnsServer.Shutdown()
|
||||||
|
go startDns(l, c)
|
||||||
}
|
}
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user