mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-16 17:57:00 +02:00
Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 86cef88744 | |||
| b7e9939e92 | |||
| 33c2d7277c | |||
| f141cebe8d | |||
| 9ec8cf10f3 |
@@ -163,3 +163,55 @@ func P256Keypair() ([]byte, []byte) {
|
|||||||
pubkey := privkey.PublicKey()
|
pubkey := privkey.PublicKey()
|
||||||
return pubkey.Bytes(), privkey.Bytes()
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -153,8 +153,8 @@ func (cm *connectionManager) Start(ctx context.Context) {
|
|||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
// Long-lived buf for the traffic-check goroutine; never released.
|
||||||
out := make([]byte, mtu)
|
buf := cm.intf.bufAlloc.Acquire()
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
@@ -169,13 +169,13 @@ func (cm *connectionManager) Start(ctx context.Context) {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.doTrafficCheck(localIndex, p, nb, out, now)
|
cm.doTrafficCheck(localIndex, p, buf, now)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cm *connectionManager) doTrafficCheck(localIndex uint32, p, nb, out []byte, now time.Time) {
|
func (cm *connectionManager) doTrafficCheck(localIndex uint32, p []byte, buf *WireBuffer, now time.Time) {
|
||||||
decision, hostinfo, primary := cm.makeTrafficDecision(localIndex, now)
|
decision, hostinfo, primary := cm.makeTrafficDecision(localIndex, now)
|
||||||
|
|
||||||
switch decision {
|
switch decision {
|
||||||
@@ -199,7 +199,7 @@ func (cm *connectionManager) doTrafficCheck(localIndex uint32, p, nb, out []byte
|
|||||||
cm.tryRehandshake(hostinfo)
|
cm.tryRehandshake(hostinfo)
|
||||||
|
|
||||||
case sendTestPacket:
|
case sendTestPacket:
|
||||||
cm.intf.SendMessageToHostInfo(header.Test, header.TestRequest, hostinfo, p, nb, out)
|
cm.intf.SendMessageToHostInfo(header.Test, header.TestRequest, hostinfo, p, buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.resetRelayTrafficCheck(hostinfo)
|
cm.resetRelayTrafficCheck(hostinfo)
|
||||||
@@ -308,7 +308,9 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
cm.l.Error("failed to marshal Control message to migrate relay", "error", err)
|
cm.l.Error("failed to marshal Control message to migrate relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
migBuf := cm.intf.bufAlloc.Acquire()
|
||||||
|
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, migBuf)
|
||||||
|
cm.intf.bufAlloc.Release(migBuf)
|
||||||
cm.l.Info("send CreateRelayRequest",
|
cm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", req.RelayFromAddr,
|
"relayFrom", req.RelayFromAddr,
|
||||||
"relayTo", req.RelayToAddr,
|
"relayTo", req.RelayToAddr,
|
||||||
|
|||||||
+14
-19
@@ -7,7 +7,6 @@ import (
|
|||||||
"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/overlay/overlaytest"
|
||||||
@@ -47,7 +46,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1HandshakeBytes: []byte{},
|
v1Credential: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -68,9 +67,9 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
buf := NewWireBuffer(mtu, 0)
|
||||||
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{
|
||||||
@@ -80,7 +79,6 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
H: &noise.HandshakeState{},
|
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -94,7 +92,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
assert.True(t, hostinfo.in.Load())
|
assert.True(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
|
// 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, buf, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.False(t, hostinfo.in.Load())
|
||||||
@@ -102,7 +100,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
// 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)
|
||||||
assert.True(t, hostinfo.out.Load())
|
assert.True(t, hostinfo.out.Load())
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, buf, time.Now())
|
||||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.False(t, hostinfo.in.Load())
|
||||||
@@ -110,7 +108,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
|
||||||
|
|
||||||
// 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, buf, time.Now())
|
||||||
assert.NotContains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs)
|
assert.NotContains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs)
|
||||||
assert.NotContains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
assert.NotContains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
|
||||||
}
|
}
|
||||||
@@ -130,7 +128,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1HandshakeBytes: []byte{},
|
v1Credential: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -151,9 +149,9 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
|
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
nb := make([]byte, 12, 12)
|
buf := NewWireBuffer(mtu, 0)
|
||||||
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{
|
||||||
@@ -163,7 +161,6 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
H: &noise.HandshakeState{},
|
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -177,14 +174,14 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
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, buf, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.False(t, hostinfo.in.Load())
|
||||||
|
|
||||||
// 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)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, buf, time.Now())
|
||||||
assert.True(t, hostinfo.pendingDeletion.Load())
|
assert.True(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.False(t, hostinfo.in.Load())
|
||||||
@@ -193,7 +190,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
|
|
||||||
// We saw traffic, should no longer be pending deletion
|
// We saw traffic, should no longer be pending deletion
|
||||||
nc.In(hostinfo)
|
nc.In(hostinfo)
|
||||||
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
|
nc.doTrafficCheck(hostinfo.localIndexId, p, buf, time.Now())
|
||||||
assert.False(t, hostinfo.pendingDeletion.Load())
|
assert.False(t, hostinfo.pendingDeletion.Load())
|
||||||
assert.False(t, hostinfo.out.Load())
|
assert.False(t, hostinfo.out.Load())
|
||||||
assert.False(t, hostinfo.in.Load())
|
assert.False(t, hostinfo.in.Load())
|
||||||
@@ -215,7 +212,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1HandshakeBytes: []byte{},
|
v1Credential: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -249,7 +246,6 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
H: &noise.HandshakeState{},
|
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -342,7 +338,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{},
|
v1Cert: &dummyCert{},
|
||||||
v1HandshakeBytes: []byte{},
|
v1Credential: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -372,7 +368,6 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
myCert: &dummyCert{},
|
myCert: &dummyCert{},
|
||||||
peerCert: cachedPeerCert,
|
peerCert: cachedPeerCert,
|
||||||
H: &noise.HandshakeState{},
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|||||||
+16
-51
@@ -1,15 +1,12 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/rand"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/handshake"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 1024
|
||||||
@@ -17,7 +14,6 @@ const ReplayWindow = 1024
|
|||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey *NebulaCipherState
|
eKey *NebulaCipherState
|
||||||
dKey *NebulaCipherState
|
dKey *NebulaCipherState
|
||||||
H *noise.HandshakeState
|
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
@@ -26,55 +22,24 @@ type ConnectionState struct {
|
|||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewConnectionState(cs *CertState, crt cert.Certificate, initiator bool, pattern noise.HandshakePattern) (*ConnectionState, error) {
|
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
||||||
var dhFunc noise.DHFunc
|
// completed handshake.Result. It seeds messageCounter and the replay window so
|
||||||
switch crt.Curve() {
|
// that the post-handshake message indices already used on the wire don't count
|
||||||
case cert.Curve_CURVE25519:
|
// as missed traffic in the data plane.
|
||||||
dhFunc = noise.DH25519
|
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||||
case cert.Curve_P256:
|
|
||||||
if cs.pkcs11Backed {
|
|
||||||
dhFunc = noiseutil.DHP256PKCS11
|
|
||||||
} else {
|
|
||||||
dhFunc = noiseutil.DHP256
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return nil, fmt.Errorf("invalid curve: %s", crt.Curve())
|
|
||||||
}
|
|
||||||
|
|
||||||
var ncs noise.CipherSuite
|
|
||||||
if cs.cipher == "chachapoly" {
|
|
||||||
ncs = noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256)
|
|
||||||
} else {
|
|
||||||
ncs = noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256)
|
|
||||||
}
|
|
||||||
|
|
||||||
static := noise.DHKey{Private: cs.privateKey, Public: crt.PublicKey()}
|
|
||||||
hs, err := noise.NewHandshakeState(noise.Config{
|
|
||||||
CipherSuite: ncs,
|
|
||||||
Random: rand.Reader,
|
|
||||||
Pattern: pattern,
|
|
||||||
Initiator: initiator,
|
|
||||||
StaticKeypair: static,
|
|
||||||
//NOTE: These should come from CertState (pki.go) when we finally implement it
|
|
||||||
PresharedKey: []byte{},
|
|
||||||
PresharedKeyPlacement: 0,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("NewConnectionState: %s", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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{
|
||||||
H: hs,
|
myCert: r.MyCert,
|
||||||
initiator: initiator,
|
initiator: r.Initiator,
|
||||||
|
peerCert: r.RemoteCert,
|
||||||
|
eKey: NewNebulaCipherState(r.EKey),
|
||||||
|
dKey: NewNebulaCipherState(r.DKey),
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
myCert: crt,
|
|
||||||
}
|
}
|
||||||
// always start the counter from 2, as packet 1 and packet 2 are handshake packets.
|
ci.messageCounter.Add(r.MessageIndex)
|
||||||
ci.messageCounter.Add(2)
|
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||||
|
ci.window.Update(nil, i)
|
||||||
return ci, nil
|
}
|
||||||
|
return ci
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||||
|
|||||||
@@ -0,0 +1,114 @@
|
|||||||
|
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())
|
||||||
|
})
|
||||||
|
}
|
||||||
+7
-10
@@ -278,15 +278,9 @@ func (c *Control) CloseTunnel(vpnIp netip.Addr, localOnly bool) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !localOnly {
|
if !localOnly {
|
||||||
c.f.send(
|
buf := c.f.bufAlloc.Acquire()
|
||||||
header.CloseTunnel,
|
c.f.send(header.CloseTunnel, 0, hostInfo.ConnectionState, hostInfo, []byte{}, buf)
|
||||||
0,
|
c.f.bufAlloc.Release(buf)
|
||||||
hostInfo.ConnectionState,
|
|
||||||
hostInfo,
|
|
||||||
[]byte{},
|
|
||||||
make([]byte, 12, 12),
|
|
||||||
make([]byte, mtu),
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
c.f.closeTunnel(hostInfo)
|
c.f.closeTunnel(hostInfo)
|
||||||
@@ -296,11 +290,14 @@ 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) {
|
||||||
|
// One WireBuffer for the whole shutdown loop.
|
||||||
|
buf := c.f.bufAlloc.Acquire()
|
||||||
|
defer c.f.bufAlloc.Release(buf)
|
||||||
shutdown := func(h *HostInfo) {
|
shutdown := func(h *HostInfo) {
|
||||||
if excludeLighthouses && c.f.lightHouse.IsAnyLighthouseAddr(h.vpnAddrs) {
|
if excludeLighthouses && c.f.lightHouse.IsAnyLighthouseAddr(h.vpnAddrs) {
|
||||||
return
|
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{}, buf)
|
||||||
c.f.closeTunnel(h)
|
c.f.closeTunnel(h)
|
||||||
|
|
||||||
c.l.Debug("Sending close tunnel message",
|
c.l.Debug("Sending close tunnel message",
|
||||||
|
|||||||
+12
-60
@@ -5,8 +5,6 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/google/gopacket"
|
|
||||||
"github.com/google/gopacket/layers"
|
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -22,7 +20,9 @@ func (c *Control) WaitForType(msgType header.MessageType, subType header.Message
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
if h.Type == msgType && h.Subtype == subType {
|
match := h.Type == msgType && h.Subtype == subType
|
||||||
|
p.Release()
|
||||||
|
if match {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -38,7 +38,9 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
if h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType {
|
match := h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType
|
||||||
|
p.Release()
|
||||||
|
if match {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -90,65 +92,15 @@ func (c *Control) GetTunTxChan() <-chan []byte {
|
|||||||
return c.f.inside.(*overlay.TestTun).TxPackets
|
return c.f.inside.(*overlay.TestTun).TxPackets
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectUDPPacket will inject a packet into the udp side of nebula
|
// InjectUDPPacket injects a packet into the udp side. We copy internally so the caller keeps ownership of p.
|
||||||
|
// The copy comes from the freelist so steady-state alloc is zero.
|
||||||
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
||||||
c.f.outside.(*udp.TesterConn).Send(p)
|
c.f.outside.(*udp.TesterConn).Send(p.Copy())
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
// InjectTunPacket pushes an IP packet onto the tun interface.
|
||||||
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
func (c *Control) InjectTunPacket(packet []byte) {
|
||||||
serialize := make([]gopacket.SerializableLayer, 0)
|
c.f.inside.(*overlay.TestTun).Send(packet)
|
||||||
var netLayer gopacket.NetworkLayer
|
|
||||||
if toAddr.Is6() {
|
|
||||||
if !fromAddr.Is6() {
|
|
||||||
panic("Cant send ipv6 to ipv4")
|
|
||||||
}
|
|
||||||
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{
|
|
||||||
SrcPort: layers.UDPPort(fromPort),
|
|
||||||
DstPort: layers.UDPPort(toPort),
|
|
||||||
}
|
|
||||||
err := udp.SetNetworkLayerForChecksum(netLayer)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
buffer := gopacket.NewSerializeBuffer()
|
|
||||||
opt := gopacket.SerializeOptions{
|
|
||||||
ComputeChecksums: true,
|
|
||||||
FixLengths: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
|
||||||
err = gopacket.SerializeLayers(buffer, opt, serialize...)
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||||
|
|||||||
@@ -0,0 +1,68 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
|
)
|
||||||
|
|
||||||
|
// BenchmarkHandshake measures end-to-end tunnel establishment time. The two
|
||||||
|
// nodes and the router are constructed once before the loop so the timed window
|
||||||
|
// is just the handshake itself: trigger packet -> handshake1 -> handshake2 ->
|
||||||
|
// cached packet replay -> arrival on the remote TUN. Between iterations we
|
||||||
|
// tear down both sides locally (no CloseTunnel notification on the wire) and
|
||||||
|
// re-inject the lighthouse address that closeTunnel cleared, so the next
|
||||||
|
// iteration runs through a fresh handshake against the same harness.
|
||||||
|
func BenchmarkHandshake(b *testing.B) {
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
|
// Default try_interval is 100ms. The handshake manager schedules handshake1
|
||||||
|
// on its OutboundHandshakeTimer rather than firing immediately on trigger
|
||||||
|
// (the trigger channel only fast-paths static hosts), so a 100ms default
|
||||||
|
// drowns the actual handshake cost. Drop it to 1ms so the bench reflects
|
||||||
|
// the computation, not the wheel cadence.
|
||||||
|
bovr := m{"handshakes": m{"try_interval": "1ms"}}
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", bovr)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24", bovr)
|
||||||
|
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
defer myControl.Stop()
|
||||||
|
defer theirControl.Stop()
|
||||||
|
|
||||||
|
r := router.NewR(b, myControl, theirControl)
|
||||||
|
r.CancelFlowLogs()
|
||||||
|
r.EnableFanIn()
|
||||||
|
|
||||||
|
trigger := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
|
||||||
|
for n := 0; n < b.N; n++ {
|
||||||
|
myControl.InjectTunPacket(trigger)
|
||||||
|
// RouteForAllUntilTxTun returns the moment the cached packet arrives at
|
||||||
|
// the remote TUN, which is also when both sides are fully established.
|
||||||
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
|
b.StopTimer()
|
||||||
|
// Local-only close removes hostmap state on both sides without putting a
|
||||||
|
// CloseTunnel packet on the wire that we'd then have to drain. The
|
||||||
|
// closeTunnel path also clears learned lighthouse state for the peer
|
||||||
|
// when the last hostinfo for that addr goes away, so we re-inject.
|
||||||
|
myControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
|
||||||
|
theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true)
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
b.StartTimer()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,6 +28,7 @@ func makeHandshakePacket(from, to netip.AddrPort, subtype header.MessageSubType,
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify the responder correctly handles receiving the same msg1 multiple times
|
// Verify the responder correctly handles receiving the same msg1 multiple times
|
||||||
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
||||||
// and the cached response is resent.
|
// and the cached response is resent.
|
||||||
@@ -46,7 +47,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me to them")
|
t.Log("Trigger handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Grab my msg1")
|
t.Log("Grab my msg1")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -78,6 +79,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify that a truncated handshake packet is ignored and the real
|
// Verify that a truncated handshake packet is ignored and the real
|
||||||
// packet can still complete the handshake.
|
// packet can still complete the handshake.
|
||||||
|
|
||||||
@@ -95,7 +97,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake")
|
t.Log("Trigger handshake")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Get msg1 and deliver to responder")
|
t.Log("Get msg1 and deliver to responder")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -126,6 +128,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// A msg2 arriving with no matching pending index should be silently dropped
|
// A msg2 arriving with no matching pending index should be silently dropped
|
||||||
// with no response sent and no state changes.
|
// with no response sent and no state changes.
|
||||||
|
|
||||||
@@ -143,7 +146,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake")
|
t.Log("Complete a normal handshake")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -168,6 +171,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// A handshake packet with an unexpected message counter should be silently
|
// A handshake packet with an unexpected message counter should be silently
|
||||||
// dropped with no side effects and no UDP response.
|
// dropped with no side effects and no UDP response.
|
||||||
|
|
||||||
@@ -199,6 +203,7 @@ func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownSubtype(t *testing.T) {
|
func TestHandshakeUnknownSubtype(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// A handshake packet with an unknown subtype should be silently dropped.
|
// A handshake packet with an unknown subtype should be silently dropped.
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -224,6 +229,7 @@ func TestHandshakeUnknownSubtype(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeLateResponse(t *testing.T) {
|
func TestHandshakeLateResponse(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// After a handshake times out, a late response should be silently ignored
|
// After a handshake times out, a late response should be silently ignored
|
||||||
// with no new tunnels created.
|
// with no new tunnels created.
|
||||||
|
|
||||||
@@ -242,7 +248,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
|
|
||||||
t.Log("Grab msg1 but don't deliver")
|
t.Log("Grab msg1 but don't deliver")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -273,6 +279,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify that a node rejects a handshake containing its own VPN IP in the
|
// Verify that a node rejects a handshake containing its own VPN IP in the
|
||||||
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
||||||
|
|
||||||
@@ -285,7 +292,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
myControl.Start()
|
myControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Drain any handshake retransmits before injecting")
|
t.Log("Drain any handshake retransmits before injecting")
|
||||||
@@ -321,6 +328,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
||||||
|
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -341,6 +349,7 @@ func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRemoteAllowList(t *testing.T) {
|
func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify that a handshake from a blocked underlay IP is dropped with no
|
// Verify that a handshake from a blocked underlay IP is dropped with no
|
||||||
// response and no state changes. Then verify the same packet from an
|
// response and no state changes. Then verify the same packet from an
|
||||||
// allowed IP succeeds.
|
// allowed IP succeeds.
|
||||||
@@ -366,7 +375,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from them")
|
t.Log("Trigger handshake from them")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
msg1 := theirControl.GetFromUDP(true)
|
msg1 := theirControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Rewrite the source to a blocked IP and inject")
|
t.Log("Rewrite the source to a blocked IP and inject")
|
||||||
@@ -399,6 +408,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
||||||
// remains functional and hostmap index count is stable.
|
// remains functional and hostmap index count is stable.
|
||||||
|
|
||||||
@@ -416,7 +426,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake via the router")
|
t.Log("Complete a normal handshake via the router")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -427,7 +437,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
originalRemote := hi.CurrentRemote
|
originalRemote := hi.CurrentRemote
|
||||||
|
|
||||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam")))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
t.Log("Verify tunnel still works")
|
t.Log("Verify tunnel still works")
|
||||||
@@ -445,6 +455,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify that when the wrong host responds, the cached packets are
|
// Verify that when the wrong host responds, the cached packets are
|
||||||
// transferred to the new handshake, the evil tunnel is closed, evil's
|
// transferred to the new handshake, the evil tunnel is closed, evil's
|
||||||
// address is blocked, and the correct tunnel is eventually established.
|
// address is blocked, and the correct tunnel is eventually established.
|
||||||
@@ -464,8 +475,8 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Send multiple packets to them (cached during handshake)")
|
t.Log("Send multiple packets to them (cached during handshake)")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1")))
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2")))
|
||||||
|
|
||||||
t.Log("Route until evil tunnel is closed")
|
t.Log("Route until evil tunnel is closed")
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
@@ -508,6 +519,7 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRelayComplete(t *testing.T) {
|
func TestHandshakeRelayComplete(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// Verify that a relay handshake completes correctly and relay state is
|
// Verify that a relay handshake completes correctly and relay state is
|
||||||
// properly maintained on all three nodes.
|
// properly maintained on all three nodes.
|
||||||
|
|
||||||
@@ -528,7 +540,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake via relay")
|
t.Log("Trigger handshake via relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -556,7 +568,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
||||||
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
// BuildTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
// framework. The check is in handshake_manager.go handleOutbound relay
|
||||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
||||||
// address is IPv6, the relay is skipped.
|
// address is IPv6, the relay is skipped.
|
||||||
|
|||||||
+68
-31
@@ -16,6 +16,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/overlay"
|
||||||
"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"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -39,11 +40,22 @@ func BenchmarkHotPath(b *testing.B) {
|
|||||||
r.CancelFlowLogs()
|
r.CancelFlowLogs()
|
||||||
|
|
||||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
|
// Pre-build the IP packet bytes once so the bench measures the data plane,
|
||||||
|
// not gopacket SerializeLayers overhead.
|
||||||
|
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
|
// EnableFanIn switches the router to a 0-alloc routing path. Required
|
||||||
|
// for hot-path benchmarks; would conflict with GetFromUDP-using tests.
|
||||||
|
r.EnableFanIn()
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(prebuilt)
|
||||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
// Release the TUN-side bytes back to the harness freelist; the bench
|
||||||
|
// just confirms a packet arrived, the contents aren't inspected.
|
||||||
|
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -71,11 +83,15 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||||
|
|
||||||
|
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
r.EnableFanIn()
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(prebuilt)
|
||||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -84,6 +100,7 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshake(t *testing.T) {
|
func TestGoodHandshake(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -96,7 +113,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -134,6 +151,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
func TestGoodHandshakeNoOverlap(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
||||||
@@ -147,7 +165,7 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
|||||||
|
|
||||||
empty := []byte{}
|
empty := []byte{}
|
||||||
t.Log("do something to cause a handshake")
|
t.Log("do something to cause a handshake")
|
||||||
myControl.GetF().SendMessageToVpnAddr(header.Test, header.MessageNone, theirVpnIpNet[0].Addr(), empty, empty, empty)
|
myControl.GetF().SendMessageToVpnAddr(header.Test, header.MessageNone, theirVpnIpNet[0].Addr(), empty, nebula.NewWireBuffer(9001, 0))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -169,6 +187,7 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshake(t *testing.T) {
|
func TestWrongResponderHandshake(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
||||||
@@ -188,7 +207,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -245,6 +264,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
||||||
@@ -269,7 +289,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -327,6 +347,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1Race(t *testing.T) {
|
func TestStage1Race(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
||||||
// But will eventually collapse down to a single tunnel
|
// But will eventually collapse down to a single tunnel
|
||||||
|
|
||||||
@@ -347,8 +368,8 @@ func TestStage1Race(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake to start on both me and them")
|
t.Log("Trigger a handshake to start on both me and them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
t.Log("Get both stage 1 handshake packets")
|
t.Log("Get both stage 1 handshake packets")
|
||||||
myHsForThem := myControl.GetFromUDP(true)
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
@@ -407,6 +428,7 @@ func TestStage1Race(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -424,7 +446,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -435,7 +457,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -456,6 +478,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -473,7 +496,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -485,7 +508,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
||||||
@@ -507,6 +530,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelays(t *testing.T) {
|
func TestRelays(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -527,7 +551,7 @@ func TestRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -536,6 +560,7 @@ func TestRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelaysDontCareAboutIps(t *testing.T) {
|
func TestRelaysDontCareAboutIps(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -556,7 +581,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -565,6 +590,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestReestablishRelays(t *testing.T) {
|
func TestReestablishRelays(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -585,14 +611,14 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("Ensure packet traversal from them to me via the relay")
|
t.Log("Ensure packet traversal from them to me via the relay")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -607,7 +633,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
for curIndexes >= start {
|
for curIndexes >= start {
|
||||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
curIndexes = len(myControl.GetHostmap().Indexes)
|
||||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
||||||
|
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
return router.RouteAndExit
|
return router.RouteAndExit
|
||||||
@@ -624,7 +650,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -659,7 +685,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
t.Log("Assert the tunnel works the other way, too")
|
t.Log("Assert the tunnel works the other way, too")
|
||||||
for {
|
for {
|
||||||
t.Log("RouteForAllUntilTxTun")
|
t.Log("RouteForAllUntilTxTun")
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -696,6 +722,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays(t *testing.T) {
|
func TestStage1RaceRelays(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -728,8 +755,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me")
|
r.Log("Wait for a packet from them to me")
|
||||||
p := r.RouteForAllUntilTxTun(myControl)
|
p := r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -743,6 +770,7 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays2(t *testing.T) {
|
func TestStage1RaceRelays2(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -775,8 +803,8 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||||
|
|
||||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||||
@@ -819,6 +847,7 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelays(t *testing.T) {
|
func TestRehandshakingRelays(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -839,7 +868,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -922,6 +951,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -943,7 +973,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -1026,6 +1056,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshaking(t *testing.T) {
|
func TestRehandshaking(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
||||||
@@ -1121,6 +1152,7 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingLoser(t *testing.T) {
|
func TestRehandshakingLoser(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
||||||
// Should be the one with the new certificate
|
// Should be the one with the new certificate
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -1219,6 +1251,7 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRaceRegression(t *testing.T) {
|
func TestRaceRegression(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
||||||
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
||||||
// caused a cross-linked hostinfo
|
// caused a cross-linked hostinfo
|
||||||
@@ -1242,8 +1275,8 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||||
|
|
||||||
t.Log("Start both handshakes")
|
t.Log("Start both handshakes")
|
||||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||||
|
|
||||||
t.Log("Get both stage 1")
|
t.Log("Get both stage 1")
|
||||||
myStage1ForThem := myControl.GetFromUDP(true)
|
myStage1ForThem := myControl.GetFromUDP(true)
|
||||||
@@ -1279,6 +1312,7 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1319,6 +1353,7 @@ func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1359,6 +1394,7 @@ func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestLighthouseUpdateOnReload(t *testing.T) {
|
func TestLighthouseUpdateOnReload(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
// Create the lighthouse
|
// Create the lighthouse
|
||||||
@@ -1434,6 +1470,7 @@ func TestLighthouseUpdateOnReload(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
unsafePrefix := "192.168.6.0/24"
|
unsafePrefix := "192.168.6.0/24"
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
||||||
@@ -1455,7 +1492,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -1483,7 +1520,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||||
|
|
||||||
//reply
|
//reply
|
||||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
||||||
//wait for reply
|
//wait for reply
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
theirControl.WaitForType(1, 0, myControl)
|
||||||
theirCachedPacket := myControl.GetFromTun(true)
|
theirCachedPacket := myControl.GetFromTun(true)
|
||||||
|
|||||||
+57
-2
@@ -294,12 +294,12 @@ func deadline(t *testing.T, seconds time.Duration) doneCb {
|
|||||||
|
|
||||||
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||||
// Send a packet from them to me
|
// Send a packet from them to me
|
||||||
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
||||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||||
|
|
||||||
// And once more from me to them
|
// And once more from me to them
|
||||||
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
||||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||||
}
|
}
|
||||||
@@ -408,3 +408,58 @@ func testLogLevelName() string {
|
|||||||
}
|
}
|
||||||
return "info"
|
return "info"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// BuildTunUDPPacket assembles an IP+UDP packet suitable for Control.InjectTunPacket.
|
||||||
|
// Using UDP here because it's a simpler protocol.
|
||||||
|
func BuildTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) []byte {
|
||||||
|
serialize := make([]gopacket.SerializableLayer, 0)
|
||||||
|
var netLayer gopacket.NetworkLayer
|
||||||
|
if toAddr.Is6() {
|
||||||
|
if !fromAddr.Is6() {
|
||||||
|
panic("Cant send ipv6 to ipv4")
|
||||||
|
}
|
||||||
|
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{
|
||||||
|
SrcPort: layers.UDPPort(fromPort),
|
||||||
|
DstPort: layers.UDPPort(toPort),
|
||||||
|
}
|
||||||
|
if err := udp.SetNetworkLayerForChecksum(netLayer); err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
buffer := gopacket.NewSerializeBuffer()
|
||||||
|
opt := gopacket.SerializeOptions{
|
||||||
|
ComputeChecksums: true,
|
||||||
|
FixLengths: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
serialize = append(serialize, &udp, gopacket.Payload(data))
|
||||||
|
if err := gopacket.SerializeLayers(buffer, opt, serialize...); err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return buffer.Bytes()
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
|
"go.uber.org/goleak"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestNoGoroutineLeaks brings up two nebula instances, completes a tunnel,
|
||||||
|
// stops both, and asserts no goroutines leak past the shutdown. goleak's
|
||||||
|
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
||||||
|
// before failing the assertion.
|
||||||
|
//
|
||||||
|
// IgnoreCurrent is necessary in the parallelized suite: other tests can
|
||||||
|
// leave goroutines mid-shutdown when this one runs (Stop is async, the
|
||||||
|
// wg.Wait() drain is not blocking on test return). We're checking that
|
||||||
|
// *this* test's setup tears down cleanly, not that the whole suite is
|
||||||
|
// idle at this moment. Intentionally NOT t.Parallel()'d for the same
|
||||||
|
// reason — concurrent test goroutines would always show up.
|
||||||
|
func TestNoGoroutineLeaks(t *testing.T) {
|
||||||
|
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
||||||
|
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||||
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
|
|
||||||
|
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
|
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||||
|
|
||||||
|
myControl.Start()
|
||||||
|
theirControl.Start()
|
||||||
|
|
||||||
|
r := router.NewR(t, myControl, theirControl)
|
||||||
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
|
myControl.Stop()
|
||||||
|
theirControl.Stop()
|
||||||
|
r.RenderFlow()
|
||||||
|
|
||||||
|
// Settle period: Stop() is non-blocking; the wg-driven goroutines need
|
||||||
|
// a moment to drain. goleak retries internally too, but a short explicit
|
||||||
|
// settle reduces flakes when the suite is busy.
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
}
|
||||||
+188
-54
@@ -13,6 +13,7 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -24,6 +25,19 @@ import (
|
|||||||
"golang.org/x/exp/maps"
|
"golang.org/x/exp/maps"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// outNatKey is the (from, to) pair used by outNat. Comparable struct, so it works as a map key without the
|
||||||
|
// allocation cost of a string-concat key.
|
||||||
|
type outNatKey struct {
|
||||||
|
from, to netip.AddrPort
|
||||||
|
}
|
||||||
|
|
||||||
|
// fannedPacket pairs a UDP TX packet with its source control so the router can route it after popping from
|
||||||
|
// the fan-in channel.
|
||||||
|
type fannedPacket struct {
|
||||||
|
from *nebula.Control
|
||||||
|
pkt *udp.Packet
|
||||||
|
}
|
||||||
|
|
||||||
type R struct {
|
type R struct {
|
||||||
// Simple map of the ip:port registered on a control to the control
|
// Simple map of the ip:port registered on a control to the control
|
||||||
// Basically a router, right?
|
// Basically a router, right?
|
||||||
@@ -34,12 +48,28 @@ type R struct {
|
|||||||
|
|
||||||
// A last used map, if an inbound packet hit the inNat map then
|
// A last used map, if an inbound packet hit the inNat map then
|
||||||
// all return packets should use the same last used inbound address for the outbound sender
|
// all return packets should use the same last used inbound address for the outbound sender
|
||||||
// map[from address + ":" + to address] => ip:port to rewrite in the udp packet to receiver
|
outNat map[outNatKey]netip.AddrPort
|
||||||
outNat map[string]netip.AddrPort
|
|
||||||
|
|
||||||
// A map of vpn ip to the nebula control it belongs to
|
// A map of vpn ip to the nebula control it belongs to
|
||||||
vpnControls map[netip.Addr]*nebula.Control
|
vpnControls map[netip.Addr]*nebula.Control
|
||||||
|
|
||||||
|
// Cached select infrastructure for RouteForAllUntilTxTun.
|
||||||
|
// The controls map is immutable after NewR so the cases are good for the test lifetime.
|
||||||
|
// We only rebuild if a different receiver is asked.
|
||||||
|
selRecvCtl *nebula.Control
|
||||||
|
selCases []reflect.SelectCase
|
||||||
|
selCtls []*nebula.Control
|
||||||
|
|
||||||
|
// Optional fan-in mode for hot-path benchmarks: one forwarder goroutine per control drains UDP TX into udpFanIn,
|
||||||
|
// so RouteForAllUntilTxTun can do a fixed 2-way native select instead of paying reflect.Select per call.
|
||||||
|
// Off by default (would otherwise interleave with tests that use GetFromUDP directly on the same control).
|
||||||
|
// Enabled by EnableFanIn.
|
||||||
|
udpFanIn chan fannedPacket
|
||||||
|
stopFanIn chan struct{}
|
||||||
|
fanInWG sync.WaitGroup
|
||||||
|
fanInMu sync.Mutex
|
||||||
|
fanInOn atomic.Bool
|
||||||
|
|
||||||
ignoreFlows []ignoreFlow
|
ignoreFlows []ignoreFlow
|
||||||
flow []flowEntry
|
flow []flowEntry
|
||||||
|
|
||||||
@@ -119,7 +149,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
controls: make(map[netip.AddrPort]*nebula.Control),
|
controls: make(map[netip.AddrPort]*nebula.Control),
|
||||||
vpnControls: make(map[netip.Addr]*nebula.Control),
|
vpnControls: make(map[netip.Addr]*nebula.Control),
|
||||||
inNat: make(map[netip.AddrPort]*nebula.Control),
|
inNat: make(map[netip.AddrPort]*nebula.Control),
|
||||||
outNat: make(map[string]netip.AddrPort),
|
outNat: make(map[outNatKey]netip.AddrPort),
|
||||||
flow: []flowEntry{},
|
flow: []flowEntry{},
|
||||||
ignoreFlows: []ignoreFlow{},
|
ignoreFlows: []ignoreFlow{},
|
||||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||||
@@ -153,8 +183,10 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-clockSource.C:
|
case <-clockSource.C:
|
||||||
|
r.Lock()
|
||||||
r.renderHostmaps("clock tick")
|
r.renderHostmaps("clock tick")
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
|
r.Unlock()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -180,15 +212,21 @@ func (r *R) AddRoute(ip netip.Addr, port uint16, c *nebula.Control) {
|
|||||||
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
||||||
func (r *R) RenderFlow() {
|
func (r *R) RenderFlow() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
}
|
}
|
||||||
|
|
||||||
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
||||||
func (r *R) CancelFlowLogs() {
|
func (r *R) CancelFlowLogs() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
|
r.Lock()
|
||||||
r.flow = nil
|
r.flow = nil
|
||||||
|
r.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// renderFlow writes the flow log to disk. Caller must hold r.Lock. renderFlow reads r.flow / r.additionalGraphs and
|
||||||
|
// the *packet pointers stashed inside, all of which are mutated under the same lock by routing paths.
|
||||||
func (r *R) renderFlow() {
|
func (r *R) renderFlow() {
|
||||||
if r.flow == nil {
|
if r.flow == nil {
|
||||||
return
|
return
|
||||||
@@ -434,68 +472,157 @@ func (r *R) RouteUntilTxTun(sender *nebula.Control, receiver *nebula.Control) []
|
|||||||
panic("No control for udp tx " + a.String())
|
panic("No control for udp tx " + a.String())
|
||||||
}
|
}
|
||||||
fp := r.unlockedInjectFlow(sender, c, p, false)
|
fp := r.unlockedInjectFlow(sender, c, p, false)
|
||||||
c.InjectUDPPacket(p)
|
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on receivers tun
|
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on the receiver's tun.
|
||||||
// If the router doesn't have the nebula controller for that address, we panic
|
// If a control's UDP TX address can't be matched to a registered control, we panic.
|
||||||
|
//
|
||||||
|
// For allocation-sensitive callers (hot-path benchmarks, in particular relay
|
||||||
|
// benches with 3+ controls), call EnableFanIn() first.
|
||||||
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
||||||
|
if r.fanInOn.Load() {
|
||||||
|
return r.routeFanIn(receiver)
|
||||||
|
}
|
||||||
|
return r.routeReflect(receiver)
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeFanIn is the alloc-free path used when EnableFanIn is in effect.
|
||||||
|
func (r *R) routeFanIn(receiver *nebula.Control) []byte {
|
||||||
|
tunTx := receiver.GetTunTxChan()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case p := <-tunTx:
|
||||||
|
r.Lock()
|
||||||
|
if r.flow != nil {
|
||||||
|
np := udp.Packet{Data: make([]byte, len(p))}
|
||||||
|
copy(np.Data, p)
|
||||||
|
r.unlockedInjectFlow(receiver, receiver, &np, true)
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
return p
|
||||||
|
case fp := <-r.udpFanIn:
|
||||||
|
r.routeUDP(fp.from, fp.pkt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeReflect is the default reflect.Select-based path. Pays the boxing allocation per call but doesn't interfere
|
||||||
|
// with tests that pull packets directly from controls' UDP TX channels via GetFromUDP.
|
||||||
|
func (r *R) routeReflect(receiver *nebula.Control) []byte {
|
||||||
|
sc, cm := r.selectCasesFor(receiver)
|
||||||
|
for {
|
||||||
|
x, rx, _ := reflect.Select(sc)
|
||||||
|
if x == 0 {
|
||||||
|
p := rx.Interface().([]byte)
|
||||||
|
r.Lock()
|
||||||
|
if r.flow != nil {
|
||||||
|
np := udp.Packet{Data: make([]byte, len(p))}
|
||||||
|
copy(np.Data, p)
|
||||||
|
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
||||||
|
}
|
||||||
|
r.Unlock()
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
r.routeUDP(cm[x], rx.Interface().(*udp.Packet))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EnableFanIn switches RouteForAllUntilTxTun to the alloc-free fan-in path.
|
||||||
|
// One forwarder goroutine per registered control drains UDP TX into a shared channel that RouteForAllUntilTxTun selects
|
||||||
|
// on alongside the receiver's TUN TX channel.
|
||||||
|
func (r *R) EnableFanIn() {
|
||||||
|
r.fanInMu.Lock()
|
||||||
|
defer r.fanInMu.Unlock()
|
||||||
|
if r.fanInOn.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.udpFanIn = make(chan fannedPacket, 32)
|
||||||
|
r.stopFanIn = make(chan struct{})
|
||||||
|
for _, c := range r.controls {
|
||||||
|
r.startFanInWorker(c)
|
||||||
|
}
|
||||||
|
r.fanInOn.Store(true)
|
||||||
|
r.t.Cleanup(r.stopFanInWorkers)
|
||||||
|
}
|
||||||
|
|
||||||
|
// startFanInWorker spawns a goroutine that drains c's UDP TX into r.udpFanIn.
|
||||||
|
func (r *R) startFanInWorker(c *nebula.Control) {
|
||||||
|
r.fanInWG.Add(1)
|
||||||
|
udpTx := c.GetUDPTxChan()
|
||||||
|
go func() {
|
||||||
|
defer r.fanInWG.Done()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-r.stopFanIn:
|
||||||
|
return
|
||||||
|
case p := <-udpTx:
|
||||||
|
select {
|
||||||
|
case <-r.stopFanIn:
|
||||||
|
p.Release()
|
||||||
|
return
|
||||||
|
case r.udpFanIn <- fannedPacket{from: c, pkt: p}:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopFanInWorkers signals the fan-in goroutines to exit and waits for them.
|
||||||
|
func (r *R) stopFanInWorkers() {
|
||||||
|
r.fanInMu.Lock()
|
||||||
|
wasOn := r.fanInOn.Swap(false)
|
||||||
|
r.fanInMu.Unlock()
|
||||||
|
if !wasOn {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
close(r.stopFanIn)
|
||||||
|
r.fanInWG.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
// routeUDP forwards a UDP TX packet from the named source control to the destination control derived from p.To,
|
||||||
|
// releasing the source packet after InjectUDPPacket has copied its bytes into a fresh pool slot.
|
||||||
|
func (r *R) routeUDP(from *nebula.Control, p *udp.Packet) {
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
|
a := from.GetUDPAddr()
|
||||||
|
c := r.getControl(a, p.To, p)
|
||||||
|
if c == nil {
|
||||||
|
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
||||||
|
}
|
||||||
|
fp := r.unlockedInjectFlow(from, c, p, false)
|
||||||
|
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||||
|
fp.WasReceived()
|
||||||
|
p.Release()
|
||||||
|
}
|
||||||
|
|
||||||
|
// selectCasesFor returns the SelectCase array used by routeReflect: one slot for the receiver's TUN TX channel followed
|
||||||
|
// by one per control's UDP TX channel. Cached for the test lifetime, only rebuilt if the receiver changes.
|
||||||
|
func (r *R) selectCasesFor(receiver *nebula.Control) ([]reflect.SelectCase, []*nebula.Control) {
|
||||||
|
r.Lock()
|
||||||
|
defer r.Unlock()
|
||||||
|
if r.selRecvCtl == receiver && r.selCases != nil {
|
||||||
|
return r.selCases, r.selCtls
|
||||||
|
}
|
||||||
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
||||||
cm := make([]*nebula.Control, len(r.controls)+1)
|
cm := make([]*nebula.Control, len(r.controls)+1)
|
||||||
|
sc[0] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(receiver.GetTunTxChan())}
|
||||||
i := 0
|
cm[0] = receiver
|
||||||
sc[i] = reflect.SelectCase{
|
i := 1
|
||||||
Dir: reflect.SelectRecv,
|
|
||||||
Chan: reflect.ValueOf(receiver.GetTunTxChan()),
|
|
||||||
Send: reflect.Value{},
|
|
||||||
}
|
|
||||||
cm[i] = receiver
|
|
||||||
|
|
||||||
i++
|
|
||||||
for _, c := range r.controls {
|
for _, c := range r.controls {
|
||||||
sc[i] = reflect.SelectCase{
|
sc[i] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(c.GetUDPTxChan())}
|
||||||
Dir: reflect.SelectRecv,
|
|
||||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
|
||||||
Send: reflect.Value{},
|
|
||||||
}
|
|
||||||
|
|
||||||
cm[i] = c
|
cm[i] = c
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
|
r.selRecvCtl = receiver
|
||||||
for {
|
r.selCases = sc
|
||||||
x, rx, _ := reflect.Select(sc)
|
r.selCtls = cm
|
||||||
r.Lock()
|
return sc, cm
|
||||||
|
|
||||||
if x == 0 {
|
|
||||||
// we are the tun tx, we can exit
|
|
||||||
p := rx.Interface().([]byte)
|
|
||||||
np := udp.Packet{Data: make([]byte, len(p))}
|
|
||||||
copy(np.Data, p)
|
|
||||||
|
|
||||||
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
|
||||||
r.Unlock()
|
|
||||||
return p
|
|
||||||
|
|
||||||
} else {
|
|
||||||
// we are a udp tx, route and continue
|
|
||||||
p := rx.Interface().(*udp.Packet)
|
|
||||||
a := cm[x].GetUDPAddr()
|
|
||||||
c := r.getControl(a, p.To, p)
|
|
||||||
if c == nil {
|
|
||||||
r.Unlock()
|
|
||||||
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
|
||||||
}
|
|
||||||
fp := r.unlockedInjectFlow(cm[x], c, p, false)
|
|
||||||
c.InjectUDPPacket(p)
|
|
||||||
fp.WasReceived()
|
|
||||||
}
|
|
||||||
r.Unlock()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
||||||
@@ -522,6 +649,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -529,6 +657,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -541,6 +670,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -641,6 +771,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -648,6 +779,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -659,6 +791,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||||
}
|
}
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -702,19 +835,20 @@ func (r *R) FlushAll() {
|
|||||||
}
|
}
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
||||||
// This is an internal router function, the caller must hold the lock
|
// This is an internal router function, the caller must hold the lock
|
||||||
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
||||||
if newAddr, ok := r.outNat[fromAddr.String()+":"+toAddr.String()]; ok {
|
if newAddr, ok := r.outNat[outNatKey{from: fromAddr, to: toAddr}]; ok {
|
||||||
p.From = newAddr
|
p.From = newAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
c, ok := r.inNat[toAddr]
|
c, ok := r.inNat[toAddr]
|
||||||
if ok {
|
if ok {
|
||||||
r.outNat[c.GetUDPAddr().String()+":"+fromAddr.String()] = toAddr
|
r.outNat[outNatKey{from: c.GetUDPAddr(), to: fromAddr}] = toAddr
|
||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+8
-2
@@ -19,6 +19,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestDropInactiveTunnels(t *testing.T) {
|
func TestDropInactiveTunnels(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -63,6 +64,7 @@ func TestDropInactiveTunnels(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertUpgrade(t *testing.T) {
|
func TestCertUpgrade(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -157,6 +159,7 @@ func TestCertUpgrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertDowngrade(t *testing.T) {
|
func TestCertDowngrade(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -255,6 +258,7 @@ func TestCertDowngrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertMismatchCorrection(t *testing.T) {
|
func TestCertMismatchCorrection(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -322,6 +326,7 @@ func TestCertMismatchCorrection(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCrossStackRelaysWork(t *testing.T) {
|
func TestCrossStackRelaysWork(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}})
|
||||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
||||||
@@ -350,14 +355,14 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))
|
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("reply?")
|
t.Log("reply?")
|
||||||
theirControl.InjectTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))
|
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -369,6 +374,7 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCloseTunnelAuthenticated(t *testing.T) {
|
func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "5s"}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "5s"}})
|
||||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
||||||
|
|||||||
+1
-1
@@ -1033,7 +1033,7 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
|||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Test a bad rule definition
|
// Test a bad rule definition
|
||||||
c := &dummyCert{}
|
c := &dummyCert{}
|
||||||
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil)
|
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil, "aes")
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ require (
|
|||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
|
go.uber.org/goleak v1.3.0
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.50.0
|
golang.org/x/crypto v0.50.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Credential holds everything needed to participate in a handshake
|
||||||
|
// at a given cert version. Version and Curve are read from Cert; the public
|
||||||
|
// half of the static keypair likewise comes from Cert.PublicKey().
|
||||||
|
type Credential struct {
|
||||||
|
Cert cert.Certificate // the certificate
|
||||||
|
Bytes []byte // pre-marshaled certificate bytes
|
||||||
|
privateKey []byte // static private key (public half lives in Cert)
|
||||||
|
cipherSuite noise.CipherSuite // pre-built cipher suite (DH + cipher + hash)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCredential creates a Credential with all material needed for handshake
|
||||||
|
// participation. The cipherSuite should be pre-built by the caller with the
|
||||||
|
// appropriate DH function, cipher, and hash.
|
||||||
|
func NewCredential(
|
||||||
|
c cert.Certificate,
|
||||||
|
hsBytes []byte,
|
||||||
|
privateKey []byte,
|
||||||
|
cipherSuite noise.CipherSuite,
|
||||||
|
) *Credential {
|
||||||
|
return &Credential{
|
||||||
|
Cert: c,
|
||||||
|
Bytes: hsBytes,
|
||||||
|
privateKey: privateKey,
|
||||||
|
cipherSuite: cipherSuite,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildHandshakeState creates a noise.HandshakeState from this credential.
|
||||||
|
func (hc *Credential) buildHandshakeState(initiator bool, pattern noise.HandshakePattern) (*noise.HandshakeState, error) {
|
||||||
|
return noise.NewHandshakeState(noise.Config{
|
||||||
|
CipherSuite: hc.cipherSuite,
|
||||||
|
Random: rand.Reader,
|
||||||
|
Pattern: pattern,
|
||||||
|
Initiator: initiator,
|
||||||
|
StaticKeypair: noise.DHKey{Private: hc.privateKey, Public: hc.Cert.PublicKey()},
|
||||||
|
PresharedKey: []byte{},
|
||||||
|
PresharedKeyPlacement: 0,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCredentialFunc returns the handshake credential for the given version,
|
||||||
|
// or nil if that version is not available.
|
||||||
|
//
|
||||||
|
// Implementations must return credentials drawn from a snapshot stable for
|
||||||
|
// the lifetime of any single Machine. The Machine may call this multiple
|
||||||
|
// times during a handshake (e.g. when negotiating to the peer's version)
|
||||||
|
// and assumes the underlying static keypair is consistent across calls.
|
||||||
|
type GetCredentialFunc func(v cert.Version) *Credential
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import "errors"
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrInitiateOnResponder = errors.New("initiate called on responder")
|
||||||
|
ErrInitiateAlreadyCalled = errors.New("initiate already called")
|
||||||
|
ErrInitiateNotCalled = errors.New("initiate must be called before ProcessPacket for initiators")
|
||||||
|
ErrPacketTooShort = errors.New("packet too short")
|
||||||
|
ErrPublicKeyMismatch = errors.New("public key mismatch between certificate and handshake")
|
||||||
|
ErrIncompleteHandshake = errors.New("handshake completed without receiving required content")
|
||||||
|
ErrMachineFailed = errors.New("handshake machine has failed")
|
||||||
|
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
||||||
|
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
||||||
|
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
||||||
|
ErrIndexAllocation = errors.New("failed to allocate local index")
|
||||||
|
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
||||||
|
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
||||||
|
ErrMultiMessageUnsupported = errors.New("multi-message handshake patterns are not yet supported by the manager")
|
||||||
|
ErrSubtypeMismatch = errors.New("packet subtype does not match handshake machine subtype")
|
||||||
|
)
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
// This file documents the wire format the nebula handshake speaks. It is
|
||||||
|
// not run through protoc; the encoder/decoder in payload.go is hand-written
|
||||||
|
// against this shape directly to keep the parser narrow and panic-free.
|
||||||
|
//
|
||||||
|
// Any change to the wire format must be reflected here, and adding a new
|
||||||
|
// field requires updating MarshalPayload / unmarshalPayloadDetails together
|
||||||
|
// with the field-uniqueness and wire-type checks in those functions.
|
||||||
|
|
||||||
|
syntax = "proto3";
|
||||||
|
package nebula.handshake;
|
||||||
|
|
||||||
|
message NebulaHandshake {
|
||||||
|
NebulaHandshakeDetails Details = 1;
|
||||||
|
bytes Hmac = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message NebulaHandshakeDetails {
|
||||||
|
bytes Cert = 1;
|
||||||
|
uint32 InitiatorIndex = 2;
|
||||||
|
uint32 ResponderIndex = 3;
|
||||||
|
// Cookie was reserved for an anti-DoS mechanism that was never
|
||||||
|
// implemented. No released version of nebula has ever populated it; the
|
||||||
|
// hand-written parser silently skips it on read.
|
||||||
|
uint64 Cookie = 4 [deprecated = true];
|
||||||
|
uint64 Time = 5;
|
||||||
|
uint32 CertVersion = 8;
|
||||||
|
// reserved for WIP multiport
|
||||||
|
reserved 6, 7;
|
||||||
|
}
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
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/header"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testCertState holds cert material for a test peer.
|
||||||
|
type testCertState struct {
|
||||||
|
version cert.Version
|
||||||
|
creds map[cert.Version]*Credential
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *testCertState) getCredential(v cert.Version) *Credential {
|
||||||
|
return s.creds[v]
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestCertState(
|
||||||
|
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
||||||
|
) *testCertState {
|
||||||
|
return newTestCertStateWithCipher(t, ca, caKey, name, networks, noise.CipherChaChaPoly)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestCertStateWithCipher(
|
||||||
|
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
||||||
|
cipher noise.CipherFunc,
|
||||||
|
) *testCertState {
|
||||||
|
t.Helper()
|
||||||
|
c, _, rawPrivKey, _ := ct.NewTestCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
||||||
|
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawPrivKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
hsBytes, err := c.MarshalForHandshakes()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ncs := noise.NewCipherSuite(noise.DH25519, cipher, noise.HashSHA256)
|
||||||
|
return &testCertState{
|
||||||
|
version: cert.Version2,
|
||||||
|
creds: map[cert.Version]*Credential{
|
||||||
|
cert.Version2: NewCredential(c, hsBytes, priv, ncs),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func testVerifier(pool *cert.CAPool) CertVerifier {
|
||||||
|
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
||||||
|
return pool.VerifyCertificate(time.Now(), c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestMachine(
|
||||||
|
t *testing.T,
|
||||||
|
cs *testCertState,
|
||||||
|
verifier CertVerifier,
|
||||||
|
initiator bool,
|
||||||
|
localIndex uint32,
|
||||||
|
) *Machine {
|
||||||
|
t.Helper()
|
||||||
|
m, err := NewMachine(
|
||||||
|
cs.version, cs.getCredential,
|
||||||
|
verifier, func() (uint32, error) { return localIndex, nil },
|
||||||
|
initiator, header.HandshakeIXPSK0,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func initiateHandshake(
|
||||||
|
t *testing.T,
|
||||||
|
initCS *testCertState, initVerifier CertVerifier,
|
||||||
|
respCS *testCertState, respVerifier CertVerifier,
|
||||||
|
) (initM, respM *Machine, respResult *Result, resp []byte, err error) {
|
||||||
|
t.Helper()
|
||||||
|
initM = newTestMachine(t, initCS, initVerifier, true, 100)
|
||||||
|
msg1, merr := initM.Initiate(nil)
|
||||||
|
require.NoError(t, merr)
|
||||||
|
|
||||||
|
respM = newTestMachine(t, respCS, respVerifier, false, 200)
|
||||||
|
resp, respResult, err = respM.ProcessPacket(nil, msg1)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func doFullHandshake(
|
||||||
|
t *testing.T, initCS, respCS *testCertState, caPool *cert.CAPool,
|
||||||
|
) (initResult, respResult *Result) {
|
||||||
|
t.Helper()
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
initM := newTestMachine(t, initCS, v, true, 1000)
|
||||||
|
respM := newTestMachine(t, respCS, v, false, 2000)
|
||||||
|
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
resp, respResult, err := respM.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, respResult)
|
||||||
|
require.NotEmpty(t, resp)
|
||||||
|
|
||||||
|
_, initResult, err = initM.ProcessPacket(nil, resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, initResult)
|
||||||
|
|
||||||
|
return initResult, respResult
|
||||||
|
}
|
||||||
@@ -0,0 +1,444 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
)
|
||||||
|
|
||||||
|
// IndexAllocator is called by the Machine to allocate a local index for the
|
||||||
|
// handshake. It is called at most once, when the first outgoing message that
|
||||||
|
// carries a payload is built.
|
||||||
|
//
|
||||||
|
// Implementations MUST NOT return 0. Zero is reserved as a sentinel meaning
|
||||||
|
// "no index assigned" on the wire and in the payload-presence checks. If an
|
||||||
|
// allocator ever returned 0, a legitimate handshake's payload could be
|
||||||
|
// indistinguishable from an empty one and would be rejected.
|
||||||
|
type IndexAllocator func() (uint32, error)
|
||||||
|
|
||||||
|
// CertVerifier is called by the Machine after reconstructing the peer's
|
||||||
|
// certificate from the handshake. The verifier performs all validation
|
||||||
|
// (CA trust, expiry, policy checks, allow lists).
|
||||||
|
type CertVerifier func(cert.Certificate) (*cert.CachedCertificate, error)
|
||||||
|
|
||||||
|
// Result contains the results of a successful handshake.
|
||||||
|
// Returned by ProcessPacket when the handshake is complete.
|
||||||
|
type Result struct {
|
||||||
|
EKey *noise.CipherState
|
||||||
|
DKey *noise.CipherState
|
||||||
|
MyCert cert.Certificate
|
||||||
|
RemoteCert *cert.CachedCertificate
|
||||||
|
RemoteIndex uint32
|
||||||
|
LocalIndex uint32
|
||||||
|
HandshakeTime uint64
|
||||||
|
MessageIndex uint64 // number of messages exchanged during the handshake
|
||||||
|
Initiator bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// Machine drives a Noise handshake through N messages. It handles Noise
|
||||||
|
// protocol operations, certificate reconstruction, and payload encoding.
|
||||||
|
// Certificate validation is delegated to the caller via CertVerifier.
|
||||||
|
//
|
||||||
|
// A Machine is not safe for concurrent use. The caller must ensure that
|
||||||
|
// Initiate and ProcessPacket are not called concurrently.
|
||||||
|
//
|
||||||
|
// Error contract: when ProcessPacket or Initiate returns an error, callers
|
||||||
|
// must check Failed() to decide what to do next. If Failed() is false the
|
||||||
|
// underlying noise state was not advanced (the packet was rejected before
|
||||||
|
// ReadMessage took effect, or the rejection is non-fatal like a stale
|
||||||
|
// retransmit) and the Machine can accept another packet. If Failed() is
|
||||||
|
// true the Machine is unrecoverable and the caller must abandon it.
|
||||||
|
type Machine struct {
|
||||||
|
hs *noise.HandshakeState
|
||||||
|
getCred GetCredentialFunc
|
||||||
|
allocIndex IndexAllocator
|
||||||
|
verifier CertVerifier
|
||||||
|
result *Result
|
||||||
|
msgs []msgFlags
|
||||||
|
myVersion cert.Version
|
||||||
|
subtype header.MessageSubType
|
||||||
|
indexAllocated bool
|
||||||
|
remoteCertSet bool
|
||||||
|
payloadSet bool
|
||||||
|
failed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMachine creates a handshake state machine. The subtype determines both
|
||||||
|
// the noise pattern and the per-message content layout. The credential for
|
||||||
|
// `version` is fetched via getCred and used to seed the noise.HandshakeState.
|
||||||
|
// IndexAllocator is called lazily when the first outgoing payload is built.
|
||||||
|
func NewMachine(
|
||||||
|
version cert.Version,
|
||||||
|
getCred GetCredentialFunc,
|
||||||
|
verifier CertVerifier,
|
||||||
|
allocIndex IndexAllocator,
|
||||||
|
initiator bool,
|
||||||
|
subtype header.MessageSubType,
|
||||||
|
) (*Machine, error) {
|
||||||
|
info, err := subtypeInfoFor(subtype)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
cred := getCred(version)
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("%w: %v", ErrNoCredential, version)
|
||||||
|
}
|
||||||
|
|
||||||
|
hs, err := cred.buildHandshakeState(initiator, info.pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("build noise state: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Machine{
|
||||||
|
hs: hs,
|
||||||
|
subtype: subtype,
|
||||||
|
msgs: info.msgs,
|
||||||
|
getCred: getCred,
|
||||||
|
allocIndex: allocIndex,
|
||||||
|
verifier: verifier,
|
||||||
|
myVersion: version,
|
||||||
|
result: &Result{
|
||||||
|
Initiator: initiator,
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Failed returns true if the Machine is in an unrecoverable state.
|
||||||
|
func (m *Machine) Failed() bool {
|
||||||
|
return m.failed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Subtype returns the handshake subtype this Machine was built for.
|
||||||
|
func (m *Machine) Subtype() header.MessageSubType {
|
||||||
|
return m.subtype
|
||||||
|
}
|
||||||
|
|
||||||
|
// MessageIndex returns the noise handshake message index, which equals the
|
||||||
|
// wire counter of the most recently sent or received message.
|
||||||
|
func (m *Machine) MessageIndex() int {
|
||||||
|
return m.hs.MessageIndex()
|
||||||
|
}
|
||||||
|
|
||||||
|
// requireComplete checks that both a peer cert and payload have been received.
|
||||||
|
// Marks the machine as failed if not.
|
||||||
|
func (m *Machine) requireComplete() error {
|
||||||
|
if !m.payloadSet || !m.remoteCertSet {
|
||||||
|
m.failed = true
|
||||||
|
return ErrIncompleteHandshake
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// myMsgFlags returns the flags for the current outgoing message.
|
||||||
|
func (m *Machine) myMsgFlags() msgFlags {
|
||||||
|
idx := m.hs.MessageIndex()
|
||||||
|
if idx < len(m.msgs) {
|
||||||
|
return m.msgs[idx]
|
||||||
|
}
|
||||||
|
return msgFlags{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// peerMsgFlags returns the flags for the message we just read.
|
||||||
|
func (m *Machine) peerMsgFlags() msgFlags {
|
||||||
|
idx := m.hs.MessageIndex() - 1
|
||||||
|
if idx >= 0 && idx < len(m.msgs) {
|
||||||
|
return m.msgs[idx]
|
||||||
|
}
|
||||||
|
return msgFlags{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initiate produces the first handshake message. Only valid for initiators,
|
||||||
|
// and must be called exactly once before ProcessPacket.
|
||||||
|
//
|
||||||
|
// out is a destination buffer the message is appended to and returned. Pass
|
||||||
|
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
||||||
|
// buf[:0]) with sufficient capacity to avoid allocation.
|
||||||
|
//
|
||||||
|
// An error return may not indicate a fatal condition, check Failed() to
|
||||||
|
// determine if the Machine can still be used.
|
||||||
|
func (m *Machine) Initiate(out []byte) ([]byte, error) {
|
||||||
|
if m.failed {
|
||||||
|
return nil, ErrMachineFailed
|
||||||
|
}
|
||||||
|
if !m.result.Initiator {
|
||||||
|
m.failed = true
|
||||||
|
return nil, ErrInitiateOnResponder
|
||||||
|
}
|
||||||
|
if m.hs.MessageIndex() != 0 {
|
||||||
|
m.failed = true
|
||||||
|
return nil, ErrInitiateAlreadyCalled
|
||||||
|
}
|
||||||
|
|
||||||
|
// At MessageIndex=0 with RemoteIndex still zero, buildResponse produces
|
||||||
|
// header counter 1 and remote index 0, which is what the initial message needs.
|
||||||
|
out, _, _, err := m.buildResponse(out)
|
||||||
|
if err != nil {
|
||||||
|
m.failed = true
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessPacket handles an incoming handshake message. It advances the Noise
|
||||||
|
// state, validates the peer certificate via the verifier, and optionally
|
||||||
|
// produces a response.
|
||||||
|
//
|
||||||
|
// out is a destination buffer the response is appended to and returned. Pass
|
||||||
|
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
||||||
|
// buf[:0]) with sufficient capacity to avoid allocation. The returned slice
|
||||||
|
// is nil when no outgoing message is produced (handshake complete on this
|
||||||
|
// side, or final message of a multi-message pattern).
|
||||||
|
//
|
||||||
|
// Returns a non-nil Result when the handshake is complete.
|
||||||
|
// An error return may not indicate a fatal condition, check Failed() to
|
||||||
|
// determine if the Machine can still be used.
|
||||||
|
func (m *Machine) ProcessPacket(out, packet []byte) ([]byte, *Result, error) {
|
||||||
|
if m.failed {
|
||||||
|
return nil, nil, ErrMachineFailed
|
||||||
|
}
|
||||||
|
if len(packet) < header.Len {
|
||||||
|
return nil, nil, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
// Reject packets whose subtype doesn't match the one this Machine was
|
||||||
|
// built for. A pending handshake that suddenly receives a different
|
||||||
|
// subtype on its index is either a stray packet that matched by chance
|
||||||
|
// or a peer protocol violation; drop it without failing the Machine so
|
||||||
|
// the legitimate retransmit can still complete.
|
||||||
|
if header.MessageSubType(packet[1]) != m.subtype {
|
||||||
|
return nil, nil, ErrSubtypeMismatch
|
||||||
|
}
|
||||||
|
if m.result.Initiator && m.hs.MessageIndex() == 0 {
|
||||||
|
m.failed = true
|
||||||
|
return nil, nil, ErrInitiateNotCalled
|
||||||
|
}
|
||||||
|
|
||||||
|
// The (eKey, dKey) ordering here is correct for IX, where the initiator
|
||||||
|
// completes the handshake by reading the responder's stage-2 message.
|
||||||
|
// noise returns (cs1, cs2) where cs1 is the initiator->responder cipher.
|
||||||
|
// For 3-message patterns where a responder finishes by reading the final
|
||||||
|
// message, this ordering would be wrong; revisit when XX/pqIX lands.
|
||||||
|
msg, eKey, dKey, err := m.hs.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
// Noise ReadMessage failed. The noise library checkpoints and rolls back
|
||||||
|
// on failure, so the Machine is still alive. The caller can retry with
|
||||||
|
// a different packet.
|
||||||
|
return nil, nil, fmt.Errorf("noise ReadMessage: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// From here on, noise state has advanced. Any error is fatal.
|
||||||
|
flags := m.peerMsgFlags()
|
||||||
|
|
||||||
|
if err := m.processPayload(msg, flags); err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// If ReadMessage derived keys, the handshake is complete. Noise should
|
||||||
|
// always produce both keys together; asymmetry is a protocol invariant
|
||||||
|
// violation.
|
||||||
|
if eKey != nil || dKey != nil {
|
||||||
|
if eKey == nil || dKey == nil {
|
||||||
|
m.failed = true
|
||||||
|
return nil, nil, ErrAsymmetricCipherKeys
|
||||||
|
}
|
||||||
|
if err := m.requireComplete(); err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return nil, m.completed(eKey, dKey), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadMessage didn't complete, produce the next outgoing message
|
||||||
|
out, dk, ek, err := m.buildResponse(out)
|
||||||
|
if err != nil {
|
||||||
|
m.failed = true
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if ek != nil || dk != nil {
|
||||||
|
if ek == nil || dk == nil {
|
||||||
|
m.failed = true
|
||||||
|
return nil, nil, ErrAsymmetricCipherKeys
|
||||||
|
}
|
||||||
|
if err := m.requireComplete(); err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return out, m.completed(ek, dk), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Machine) completed(eKey, dKey *noise.CipherState) *Result {
|
||||||
|
m.result.EKey = eKey
|
||||||
|
m.result.DKey = dKey
|
||||||
|
m.result.MessageIndex = uint64(m.hs.MessageIndex())
|
||||||
|
return m.result
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
||||||
|
if len(msg) == 0 {
|
||||||
|
if flags.expectsPayload || flags.expectsCert {
|
||||||
|
m.failed = true
|
||||||
|
return ErrMissingContent
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
payload, err := UnmarshalPayload(msg)
|
||||||
|
if err != nil {
|
||||||
|
m.failed = true
|
||||||
|
return fmt.Errorf("unmarshal handshake: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Assert the payload contains exactly what we expect
|
||||||
|
hasPayloadData := payload.InitiatorIndex != 0 || payload.ResponderIndex != 0 || payload.Time != 0
|
||||||
|
if hasPayloadData != flags.expectsPayload {
|
||||||
|
m.failed = true
|
||||||
|
return ErrUnexpectedContent
|
||||||
|
}
|
||||||
|
|
||||||
|
hasCertData := len(payload.Cert) > 0
|
||||||
|
if hasCertData != flags.expectsCert {
|
||||||
|
m.failed = true
|
||||||
|
return ErrUnexpectedContent
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process payload
|
||||||
|
if flags.expectsPayload {
|
||||||
|
if m.result.Initiator {
|
||||||
|
m.result.RemoteIndex = payload.ResponderIndex
|
||||||
|
} else {
|
||||||
|
m.result.RemoteIndex = payload.InitiatorIndex
|
||||||
|
}
|
||||||
|
m.result.HandshakeTime = payload.Time
|
||||||
|
m.payloadSet = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process certificate
|
||||||
|
if flags.expectsCert {
|
||||||
|
if err := m.validateCert(payload); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Machine) validateCert(payload Payload) error {
|
||||||
|
cred := m.getCred(m.myVersion)
|
||||||
|
if cred == nil {
|
||||||
|
m.failed = true
|
||||||
|
return fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
||||||
|
}
|
||||||
|
rc, err := cert.Recombine(
|
||||||
|
cert.Version(payload.CertVersion),
|
||||||
|
payload.Cert,
|
||||||
|
m.hs.PeerStatic(),
|
||||||
|
cred.Cert.Curve(),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
m.failed = true
|
||||||
|
return fmt.Errorf("recombine cert: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(rc.PublicKey(), m.hs.PeerStatic()) {
|
||||||
|
m.failed = true
|
||||||
|
return ErrPublicKeyMismatch
|
||||||
|
}
|
||||||
|
|
||||||
|
// Version negotiation, if the peer sent a different version and we have it, switch
|
||||||
|
if rc.Version() != m.myVersion {
|
||||||
|
if m.getCred(rc.Version()) != nil {
|
||||||
|
m.myVersion = rc.Version()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
verified, err := m.verifier(rc)
|
||||||
|
if err != nil {
|
||||||
|
m.failed = true
|
||||||
|
return fmt.Errorf("verify cert: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m.result.RemoteCert = verified
|
||||||
|
m.remoteCertSet = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Machine) marshalOutgoing(flags msgFlags) ([]byte, error) {
|
||||||
|
if !flags.expectsPayload && !flags.expectsCert {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var p Payload
|
||||||
|
if flags.expectsPayload {
|
||||||
|
if !m.indexAllocated {
|
||||||
|
index, err := m.allocIndex()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%w: %w", ErrIndexAllocation, err)
|
||||||
|
}
|
||||||
|
m.result.LocalIndex = index
|
||||||
|
m.indexAllocated = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if m.result.Initiator {
|
||||||
|
p.InitiatorIndex = m.result.LocalIndex
|
||||||
|
} else {
|
||||||
|
p.ResponderIndex = m.result.LocalIndex
|
||||||
|
p.InitiatorIndex = m.result.RemoteIndex
|
||||||
|
}
|
||||||
|
p.Time = uint64(time.Now().UnixNano())
|
||||||
|
}
|
||||||
|
if flags.expectsCert {
|
||||||
|
cred := m.getCred(m.myVersion)
|
||||||
|
if cred == nil {
|
||||||
|
return nil, fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
||||||
|
}
|
||||||
|
p.Cert = cred.Bytes
|
||||||
|
p.CertVersion = uint32(cred.Cert.Version())
|
||||||
|
m.result.MyCert = cred.Cert
|
||||||
|
}
|
||||||
|
|
||||||
|
return MarshalPayload(nil, p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Machine) buildResponse(out []byte) ([]byte, *noise.CipherState, *noise.CipherState, error) {
|
||||||
|
flags := m.myMsgFlags()
|
||||||
|
hsBytes, err := m.marshalOutgoing(flags)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extend out by header.Len to make room for the header. slices.Grow is a
|
||||||
|
// no-op when the cap is already sufficient (the zero-copy case where the
|
||||||
|
// caller passed a pre-sized buffer). header.Encode overwrites the new
|
||||||
|
// bytes, so they don't need to be zeroed.
|
||||||
|
start := len(out)
|
||||||
|
out = slices.Grow(out, header.Len)[:start+header.Len]
|
||||||
|
header.Encode(
|
||||||
|
out[start:],
|
||||||
|
header.Version, header.Handshake, m.subtype,
|
||||||
|
m.result.RemoteIndex,
|
||||||
|
uint64(m.hs.MessageIndex()+1),
|
||||||
|
)
|
||||||
|
|
||||||
|
// noise.WriteMessage appends the encrypted handshake message to out,
|
||||||
|
// reusing capacity when present.
|
||||||
|
//
|
||||||
|
// The (dKey, eKey) ordering here is correct for IX, where the responder
|
||||||
|
// completes the handshake by writing the stage-2 message. noise returns
|
||||||
|
// (cs1, cs2) where cs1 is the initiator->responder cipher (which is the
|
||||||
|
// responder's decrypt key). For 3-message patterns where an initiator
|
||||||
|
// finishes by writing the final message, this ordering would be wrong;
|
||||||
|
// revisit when XX/pqIX lands.
|
||||||
|
out, dKey, eKey, err := m.hs.WriteMessage(out, hsBytes)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, nil, fmt.Errorf("noise WriteMessage: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, dKey, eKey, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,662 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
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/header"
|
||||||
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMachineIXHappyPath(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
respCS := newTestCertState(t, ca, caKey, "responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
|
||||||
|
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||||
|
|
||||||
|
assert.Equal(t, "responder", initR.RemoteCert.Certificate.Name())
|
||||||
|
assert.Equal(t, "initiator", respR.RemoteCert.Certificate.Name())
|
||||||
|
|
||||||
|
assert.Equal(t, uint32(1000), initR.LocalIndex)
|
||||||
|
assert.Equal(t, uint32(2000), initR.RemoteIndex)
|
||||||
|
assert.Equal(t, uint32(2000), respR.LocalIndex)
|
||||||
|
assert.Equal(t, uint32(1000), respR.RemoteIndex)
|
||||||
|
|
||||||
|
assert.Equal(t, uint64(2), initR.MessageIndex, "IX has 2 messages")
|
||||||
|
assert.Equal(t, uint64(2), respR.MessageIndex, "IX has 2 messages")
|
||||||
|
|
||||||
|
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("hello"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("hello"), pt1)
|
||||||
|
|
||||||
|
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("world"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("world"), pt2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineInitiateErrors(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
t.Run("initiate on responder", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
_, err := m.Initiate(nil)
|
||||||
|
require.ErrorIs(t, err, ErrInitiateOnResponder)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("initiate called twice", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, true, 100)
|
||||||
|
_, err := m.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, err = m.Initiate(nil)
|
||||||
|
require.ErrorIs(t, err, ErrInitiateAlreadyCalled)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("process packet before initiate on initiator", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, true, 100)
|
||||||
|
_, _, err := m.ProcessPacket(nil, make([]byte, 100))
|
||||||
|
require.ErrorIs(t, err, ErrInitiateNotCalled)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("calling failed machine", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
_, err := m.Initiate(nil) // fails: responder
|
||||||
|
require.Error(t, err)
|
||||||
|
_, err = m.Initiate(nil) // fails: already failed
|
||||||
|
require.ErrorIs(t, err, ErrMachineFailed)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineProcessPacketErrors(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
t.Run("packet too short", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
_, _, err := m.ProcessPacket(nil, []byte{1, 2, 3})
|
||||||
|
require.ErrorIs(t, err, ErrPacketTooShort)
|
||||||
|
assert.False(t, m.Failed(), "short packet should not kill machine")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("noise decryption failure is recoverable", func(t *testing.T) {
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
initM := newTestMachine(t, initCS, v, true, 100)
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
respM := newTestMachine(t, cs, v, false, 200)
|
||||||
|
resp, _, err := respM.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
corrupted := make([]byte, len(resp))
|
||||||
|
copy(corrupted, resp)
|
||||||
|
for i := header.Len; i < len(corrupted); i++ {
|
||||||
|
corrupted[i] ^= 0xff
|
||||||
|
}
|
||||||
|
_, _, err = initM.ProcessPacket(nil, corrupted)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.False(t, initM.Failed(), "noise failure should be recoverable")
|
||||||
|
|
||||||
|
// And the machine should still complete a real handshake afterward.
|
||||||
|
_, result, err := initM.ProcessPacket(nil, resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result, "initiator should complete on the legitimate response")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("invalid cert is fatal", func(t *testing.T) {
|
||||||
|
otherCA, _, otherCAKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
otherCS := newTestCertState(t, otherCA, otherCAKey, "other", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
|
||||||
|
initM := newTestMachine(t, otherCS, testVerifier(ct.NewTestCAPool(otherCA)), true, 100)
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
respM := newTestMachine(t, cs, v, false, 200)
|
||||||
|
_, _, err = respM.ProcessPacket(nil, msg1)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.True(t, respM.Failed(), "cert validation failure should kill machine")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("subtype mismatch is recoverable", func(t *testing.T) {
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
initM := newTestMachine(t, initCS, v, true, 100)
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Mutate the subtype byte (offset 1 in the header) to a value the
|
||||||
|
// responder Machine wasn't built for.
|
||||||
|
bad := make([]byte, len(msg1))
|
||||||
|
copy(bad, msg1)
|
||||||
|
bad[1] = 0xff
|
||||||
|
|
||||||
|
respM := newTestMachine(t, cs, v, false, 200)
|
||||||
|
_, _, err = respM.ProcessPacket(nil, bad)
|
||||||
|
require.ErrorIs(t, err, ErrSubtypeMismatch)
|
||||||
|
assert.False(t, respM.Failed(), "subtype mismatch should not kill the machine")
|
||||||
|
|
||||||
|
// And the machine should still complete a real handshake afterward.
|
||||||
|
resp, result, err := respM.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result, "responder should complete on the legitimate stage-1 packet")
|
||||||
|
assert.NotEmpty(t, resp, "responder should produce a stage-2 reply")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMachineProcessPayload exercises processPayload's internal validation
|
||||||
|
// directly. Most of these failure modes can't be reached black-box once the
|
||||||
|
// subtype check at the top of ProcessPacket gates external callers, so we
|
||||||
|
// drive them by hand here for coverage.
|
||||||
|
func TestMachineProcessPayload(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
t.Run("empty message with expects fails", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
err := m.processPayload(nil, msgFlags{expectsPayload: true, expectsCert: true})
|
||||||
|
require.ErrorIs(t, err, ErrMissingContent)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty message with no expects passes", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
err := m.processPayload(nil, msgFlags{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.False(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("malformed protobuf is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
err := m.processPayload([]byte{0xff, 0xff, 0xff}, msgFlags{expectsPayload: true, expectsCert: true})
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("unexpected payload data is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
// A payload with index data when none was expected.
|
||||||
|
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1})
|
||||||
|
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
||||||
|
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("unexpected cert data is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
// A payload with cert when none was expected.
|
||||||
|
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
||||||
|
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
||||||
|
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("missing payload data when expected is fatal", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
// Cert present, but no index/time fields.
|
||||||
|
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
||||||
|
err := m.processPayload(bytes, msgFlags{expectsPayload: true, expectsCert: true})
|
||||||
|
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
||||||
|
// directly. Like processPayload above this isn't reachable from a normal IX
|
||||||
|
// flow, so we drive it by hand.
|
||||||
|
func TestMachineRequireComplete(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
t.Run("missing both fails", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
err := m.requireComplete()
|
||||||
|
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("payload only fails", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
m.payloadSet = true
|
||||||
|
err := m.requireComplete()
|
||||||
|
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("cert only fails", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
m.remoteCertSet = true
|
||||||
|
err := m.requireComplete()
|
||||||
|
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||||
|
assert.True(t, m.Failed())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("both set passes", func(t *testing.T) {
|
||||||
|
m := newTestMachine(t, cs, v, false, 100)
|
||||||
|
m.payloadSet = true
|
||||||
|
m.remoteCertSet = true
|
||||||
|
err := m.requireComplete()
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.False(t, m.Failed())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineAESCipher(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
|
||||||
|
initCS := newTestCertStateWithCipher(
|
||||||
|
t, ca, caKey, "init",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||||
|
noiseutil.CipherAESGCM,
|
||||||
|
)
|
||||||
|
respCS := newTestCertStateWithCipher(
|
||||||
|
t, ca, caKey, "resp",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||||
|
noiseutil.CipherAESGCM,
|
||||||
|
)
|
||||||
|
|
||||||
|
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||||
|
|
||||||
|
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("works"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("works"), pt1)
|
||||||
|
|
||||||
|
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("back"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("back"), pt2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResultFields(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
|
||||||
|
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||||
|
|
||||||
|
assert.True(t, initR.Initiator)
|
||||||
|
assert.False(t, respR.Initiator)
|
||||||
|
assert.NotZero(t, initR.HandshakeTime)
|
||||||
|
assert.NotZero(t, respR.HandshakeTime)
|
||||||
|
assert.NotNil(t, initR.RemoteCert)
|
||||||
|
assert.NotNil(t, respR.RemoteCert)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineBufferReuse(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
initM := newTestMachine(t, initCS, v, true, 1000)
|
||||||
|
respM := newTestMachine(t, respCS, v, false, 2000)
|
||||||
|
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Run("response writes into provided buffer", func(t *testing.T) {
|
||||||
|
buf := make([]byte, 0, 4096)
|
||||||
|
resp, result, err := respM.ProcessPacket(buf, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
assert.NotEmpty(t, resp, "response should have content")
|
||||||
|
assert.Equal(t, &buf[:1][0], &resp[:1][0],
|
||||||
|
"response should reuse the provided buffer's backing array")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("initiate writes into provided buffer", func(t *testing.T) {
|
||||||
|
initM2 := newTestMachine(t, initCS, v, true, 3000)
|
||||||
|
buf := make([]byte, 0, 4096)
|
||||||
|
msg, err := initM2.Initiate(buf)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.NotEmpty(t, msg, "initiate should have content")
|
||||||
|
assert.Equal(t, &buf[:1][0], &msg[:1][0],
|
||||||
|
"initiate should reuse the provided buffer's backing array")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil out still works", func(t *testing.T) {
|
||||||
|
initM2 := newTestMachine(t, initCS, v, true, 4000)
|
||||||
|
respM2 := newTestMachine(t, respCS, v, false, 5000)
|
||||||
|
|
||||||
|
msg1, err := initM2.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
resp, _, err := respM2.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
out, result, err := initM2.ProcessPacket(nil, resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Nil(t, out, "initiator should have no response for IX msg2")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineMsgIndexTracking(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
initM := newTestMachine(t, initCS, v, true, 100)
|
||||||
|
respM := newTestMachine(t, respCS, v, false, 200)
|
||||||
|
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
resp1, result1, err := respM.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotNil(t, result1)
|
||||||
|
|
||||||
|
_, result2, err := initM.ProcessPacket(nil, resp1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotNil(t, result2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineThreeMessagePattern(t *testing.T) {
|
||||||
|
registerTestXXInfo(t)
|
||||||
|
|
||||||
|
// Use HandshakeXX (3 messages) to verify the Machine handles multi-message
|
||||||
|
// patterns correctly. XX flow:
|
||||||
|
// msg1 (I->R): [E] - payload only, no cert
|
||||||
|
// msg2 (R->I): [E, ee, S, es] - payload + cert
|
||||||
|
// msg3 (I->R): [S, se] - cert only (no payload, not first two)
|
||||||
|
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||||
|
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||||
|
|
||||||
|
initM, err := NewMachine(
|
||||||
|
cert.Version2,
|
||||||
|
initCS.getCredential, v,
|
||||||
|
func() (uint32, error) { return 1000, nil },
|
||||||
|
true, header.HandshakeXXPSK0,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
respM, err := NewMachine(
|
||||||
|
cert.Version2,
|
||||||
|
respCS.getCredential, v,
|
||||||
|
func() (uint32, error) { return 2000, nil },
|
||||||
|
false, header.HandshakeXXPSK0,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// msg1: initiator -> responder (E only, no cert)
|
||||||
|
msg1, err := initM.Initiate(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, msg1)
|
||||||
|
|
||||||
|
// Responder processes msg1, should not complete yet, should produce msg2
|
||||||
|
msg2, result, err := respM.ProcessPacket(nil, msg1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Nil(t, result, "XX should not complete on msg1")
|
||||||
|
assert.NotEmpty(t, msg2, "responder should produce msg2")
|
||||||
|
|
||||||
|
// Initiator processes msg2: gets responder's cert, produces msg3, and
|
||||||
|
// completes (WriteMessage for msg3 derives keys)
|
||||||
|
msg3, initResult, err := initM.ProcessPacket(nil, msg2)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, initResult, "XX initiator should complete after reading msg2 and writing msg3")
|
||||||
|
assert.NotEmpty(t, msg3, "initiator should produce msg3")
|
||||||
|
assert.Equal(t, "resp", initResult.RemoteCert.Certificate.Name())
|
||||||
|
|
||||||
|
// Responder processes msg3: gets initiator's cert and completes
|
||||||
|
_, respResult, err := respM.ProcessPacket(nil, msg3)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, respResult, "XX responder should complete on msg3")
|
||||||
|
assert.Equal(t, "init", respResult.RemoteCert.Certificate.Name())
|
||||||
|
|
||||||
|
assert.Equal(t, uint64(3), initResult.MessageIndex, "XX has 3 messages")
|
||||||
|
assert.Equal(t, uint64(3), respResult.MessageIndex, "XX has 3 messages")
|
||||||
|
|
||||||
|
// Verify keys work
|
||||||
|
ct1, err := initResult.EKey.Encrypt(nil, nil, []byte("three messages"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pt1, err := respResult.DKey.Decrypt(nil, nil, ct1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("three messages"), pt1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NOTE: ErrIncompleteHandshake is tested implicitly. It can't be triggered with
|
||||||
|
// IX since the cert is always in the payload. A 3-message pattern test (HybridIX)
|
||||||
|
// should exercise the case where cert arrives in msg3 and verify that completing
|
||||||
|
// without it fails.
|
||||||
|
|
||||||
|
func TestMachineExpiredCert(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519,
|
||||||
|
time.Now().Add(-24*time.Hour), time.Now().Add(24*time.Hour),
|
||||||
|
nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
|
||||||
|
expCert, _, expKeyPEM, _ := ct.NewTestCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
||||||
|
"expired", time.Now().Add(-2*time.Hour), time.Now().Add(-1*time.Hour),
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}, nil, nil,
|
||||||
|
)
|
||||||
|
expKey, _, _, err := cert.UnmarshalPrivateKeyFromPEM(expKeyPEM)
|
||||||
|
require.NoError(t, err)
|
||||||
|
expHsBytes, err := expCert.MarshalForHandshakes()
|
||||||
|
require.NoError(t, err)
|
||||||
|
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
|
||||||
|
expiredCS := &testCertState{
|
||||||
|
version: cert.Version2,
|
||||||
|
creds: map[cert.Version]*Credential{
|
||||||
|
cert.Version2: NewCredential(expCert, expHsBytes, expKey, ncs),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
respCS := newTestCertState(
|
||||||
|
t, ca, caKey, "responder",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||||
|
)
|
||||||
|
|
||||||
|
_, respM, _, _, err := initiateHandshake(
|
||||||
|
t, expiredCS, testVerifier(caPool),
|
||||||
|
respCS, testVerifier(caPool),
|
||||||
|
)
|
||||||
|
require.ErrorContains(t, err, "verify cert")
|
||||||
|
assert.True(t, respM.Failed())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineNoCertNetworks(t *testing.T) {
|
||||||
|
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca)
|
||||||
|
|
||||||
|
caHsBytes, err := ca.MarshalForHandshakes()
|
||||||
|
require.NoError(t, err)
|
||||||
|
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
|
||||||
|
noNetCS := &testCertState{
|
||||||
|
version: cert.Version2,
|
||||||
|
creds: map[cert.Version]*Credential{
|
||||||
|
cert.Version2: NewCredential(ca, caHsBytes, caKey, ncs),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
respCS := newTestCertState(
|
||||||
|
t, ca, caKey, "responder",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||||
|
)
|
||||||
|
|
||||||
|
_, respM, _, _, err := initiateHandshake(
|
||||||
|
t, noNetCS, testVerifier(caPool),
|
||||||
|
respCS, testVerifier(caPool),
|
||||||
|
)
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.True(t, respM.Failed())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineDifferentCAs(t *testing.T) {
|
||||||
|
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
initCS := newTestCertState(
|
||||||
|
t, ca1, caKey1, "init",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||||
|
)
|
||||||
|
respCS := newTestCertState(
|
||||||
|
t, ca2, caKey2, "resp",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||||
|
)
|
||||||
|
|
||||||
|
_, respM, _, _, err := initiateHandshake(
|
||||||
|
t, initCS, testVerifier(ct.NewTestCAPool(ca1)),
|
||||||
|
respCS, testVerifier(ct.NewTestCAPool(ca2)),
|
||||||
|
)
|
||||||
|
require.ErrorContains(t, err, "verify cert")
|
||||||
|
assert.True(t, respM.Failed())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMachineVersionNegotiation(t *testing.T) {
|
||||||
|
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version1, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
||||||
|
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||||
|
)
|
||||||
|
caPool := ct.NewTestCAPool(ca1, ca2)
|
||||||
|
|
||||||
|
makeMultiVersionResp := func(t *testing.T) *testCertState {
|
||||||
|
t.Helper()
|
||||||
|
respCertV1, _, respKeyPEM, _ := ct.NewTestCert(
|
||||||
|
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
||||||
|
ca1.NotBefore(), ca1.NotAfter(),
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
||||||
|
)
|
||||||
|
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
||||||
|
respCertV2, _ := ct.NewTestCertDifferentVersion(respCertV1, cert.Version2, ca2, caKey2)
|
||||||
|
respHsV1, _ := respCertV1.MarshalForHandshakes()
|
||||||
|
respHsV2, _ := respCertV2.MarshalForHandshakes()
|
||||||
|
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
return &testCertState{
|
||||||
|
version: cert.Version1,
|
||||||
|
creds: map[cert.Version]*Credential{
|
||||||
|
cert.Version1: NewCredential(respCertV1, respHsV1, respKey, ncs),
|
||||||
|
cert.Version2: NewCredential(respCertV2, respHsV2, respKey, ncs),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("responder matches initiator version", func(t *testing.T) {
|
||||||
|
initCS := newTestCertState(
|
||||||
|
t, ca2, caKey2, "init",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||||
|
)
|
||||||
|
respCS := makeMultiVersionResp(t)
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
|
||||||
|
initM, _, respResult, resp, err := initiateHandshake(
|
||||||
|
t, initCS, v,
|
||||||
|
respCS, v,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, respResult)
|
||||||
|
|
||||||
|
assert.Equal(t, cert.Version2, respResult.MyCert.Version(),
|
||||||
|
"responder should negotiate to initiator's version")
|
||||||
|
|
||||||
|
_, initResult, err := initM.ProcessPacket(nil, resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, initResult)
|
||||||
|
assert.Equal(t, cert.Version2, initResult.RemoteCert.Certificate.Version(),
|
||||||
|
"initiator should see V2 cert from responder")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("responder keeps version when no match available", func(t *testing.T) {
|
||||||
|
initCS := newTestCertState(
|
||||||
|
t, ca2, caKey2, "init",
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||||
|
)
|
||||||
|
|
||||||
|
respCert, _, respKeyPEM, _ := ct.NewTestCert(
|
||||||
|
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
||||||
|
ca1.NotBefore(), ca1.NotAfter(),
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
||||||
|
)
|
||||||
|
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
||||||
|
respHs, _ := respCert.MarshalForHandshakes()
|
||||||
|
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||||
|
respCS := &testCertState{
|
||||||
|
version: cert.Version1,
|
||||||
|
creds: map[cert.Version]*Credential{
|
||||||
|
cert.Version1: NewCredential(respCert, respHs, respKey, ncs),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
v := testVerifier(caPool)
|
||||||
|
_, _, respResult, _, err := initiateHandshake(
|
||||||
|
t, initCS, v,
|
||||||
|
respCS, v,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, respResult)
|
||||||
|
|
||||||
|
assert.Equal(t, cert.Version1, respResult.MyCert.Version(),
|
||||||
|
"responder should keep V1 when V2 not available")
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
)
|
||||||
|
|
||||||
|
// msgFlags tracks what application data a handshake message carries.
|
||||||
|
type msgFlags struct {
|
||||||
|
expectsPayload bool // message carries indexes and time
|
||||||
|
expectsCert bool // message carries the certificate
|
||||||
|
}
|
||||||
|
|
||||||
|
// subtypeInfo bundles the noise pattern with the per-message flags for a
|
||||||
|
// given handshake subtype.
|
||||||
|
type subtypeInfo struct {
|
||||||
|
pattern noise.HandshakePattern
|
||||||
|
msgs []msgFlags
|
||||||
|
}
|
||||||
|
|
||||||
|
// subtypeInfos defines the noise pattern and message content layout for each
|
||||||
|
// handshake subtype.
|
||||||
|
var subtypeInfos = map[header.MessageSubType]subtypeInfo{
|
||||||
|
// IX: 2 messages, both carry payload and cert
|
||||||
|
header.HandshakeIXPSK0: {
|
||||||
|
pattern: noise.HandshakeIX,
|
||||||
|
msgs: []msgFlags{
|
||||||
|
{expectsPayload: true, expectsCert: true},
|
||||||
|
{expectsPayload: true, expectsCert: true},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
|
||||||
|
// XX: 3 messages
|
||||||
|
// msg1 (I->R): payload only
|
||||||
|
// msg2 (R->I): payload + cert
|
||||||
|
// msg3 (I->R): cert only
|
||||||
|
//header.HandshakeXXPSK0: {
|
||||||
|
// pattern: noise.HandshakeXX,
|
||||||
|
// msgs: []msgFlags{
|
||||||
|
// {expectsPayload: true, expectsCert: false},
|
||||||
|
// {expectsPayload: true, expectsCert: true},
|
||||||
|
// {expectsPayload: false, expectsCert: true},
|
||||||
|
// },
|
||||||
|
//},
|
||||||
|
}
|
||||||
|
|
||||||
|
func subtypeInfoFor(subtype header.MessageSubType) (subtypeInfo, error) {
|
||||||
|
if info, ok := subtypeInfos[subtype]; ok {
|
||||||
|
return info, nil
|
||||||
|
}
|
||||||
|
return subtypeInfo{}, fmt.Errorf("%w: %d", ErrUnknownSubtype, subtype)
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSubtypeInfo(t *testing.T) {
|
||||||
|
t.Run("IX", func(t *testing.T) {
|
||||||
|
info, err := subtypeInfoFor(header.HandshakeIXPSK0)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, noise.HandshakeIX.Name, info.pattern.Name)
|
||||||
|
require.Len(t, info.msgs, 2)
|
||||||
|
// msg1: payload + cert
|
||||||
|
assert.True(t, info.msgs[0].expectsPayload)
|
||||||
|
assert.True(t, info.msgs[0].expectsCert)
|
||||||
|
// msg2: payload + cert
|
||||||
|
assert.True(t, info.msgs[1].expectsPayload)
|
||||||
|
assert.True(t, info.msgs[1].expectsCert)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("XX", func(t *testing.T) {
|
||||||
|
registerTestXXInfo(t)
|
||||||
|
info, err := subtypeInfoFor(header.HandshakeXXPSK0)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, noise.HandshakeXX.Name, info.pattern.Name)
|
||||||
|
require.Len(t, info.msgs, 3)
|
||||||
|
// msg1: payload only
|
||||||
|
assert.True(t, info.msgs[0].expectsPayload)
|
||||||
|
assert.False(t, info.msgs[0].expectsCert)
|
||||||
|
// msg2: payload + cert
|
||||||
|
assert.True(t, info.msgs[1].expectsPayload)
|
||||||
|
assert.True(t, info.msgs[1].expectsCert)
|
||||||
|
// msg3: cert only
|
||||||
|
assert.False(t, info.msgs[2].expectsPayload)
|
||||||
|
assert.True(t, info.msgs[2].expectsCert)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("unknown subtype returns error", func(t *testing.T) {
|
||||||
|
_, err := subtypeInfoFor(99)
|
||||||
|
require.ErrorIs(t, err, ErrUnknownSubtype)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerTestXXInfo temporarily registers XX subtype info for testing.
|
||||||
|
func registerTestXXInfo(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
subtypeInfos[header.HandshakeXXPSK0] = subtypeInfo{
|
||||||
|
pattern: noise.HandshakeXX,
|
||||||
|
msgs: []msgFlags{
|
||||||
|
{expectsPayload: true, expectsCert: false},
|
||||||
|
{expectsPayload: true, expectsCert: true},
|
||||||
|
{expectsPayload: false, expectsCert: true},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
delete(subtypeInfos, header.HandshakeXXPSK0)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,173 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
|
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errInvalidHandshakeMessage = errors.New("invalid handshake message")
|
||||||
|
errInvalidHandshakeDetails = errors.New("invalid handshake details")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Payload represents the decoded fields of a handshake message.
|
||||||
|
// Wire format is protobuf-compatible with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
||||||
|
type Payload struct {
|
||||||
|
Cert []byte
|
||||||
|
InitiatorIndex uint32
|
||||||
|
ResponderIndex uint32
|
||||||
|
Time uint64
|
||||||
|
CertVersion uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Proto field numbers for NebulaHandshakeDetails
|
||||||
|
const (
|
||||||
|
fieldCert = 1 // bytes
|
||||||
|
fieldInitiatorIndex = 2 // uint32
|
||||||
|
fieldResponderIndex = 3 // uint32
|
||||||
|
fieldTime = 5 // uint64
|
||||||
|
fieldCertVersion = 8 // uint32
|
||||||
|
)
|
||||||
|
|
||||||
|
// MarshalPayload encodes a handshake payload in protobuf wire format compatible
|
||||||
|
// with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
||||||
|
// Returns out (which may be nil), with the marshalled Payload appended to it.
|
||||||
|
func MarshalPayload(out []byte, p Payload) []byte {
|
||||||
|
var details []byte
|
||||||
|
|
||||||
|
if len(p.Cert) > 0 {
|
||||||
|
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
||||||
|
details = protowire.AppendBytes(details, p.Cert)
|
||||||
|
}
|
||||||
|
if p.InitiatorIndex != 0 {
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, uint64(p.InitiatorIndex))
|
||||||
|
}
|
||||||
|
if p.ResponderIndex != 0 {
|
||||||
|
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, uint64(p.ResponderIndex))
|
||||||
|
}
|
||||||
|
if p.Time != 0 {
|
||||||
|
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, p.Time)
|
||||||
|
}
|
||||||
|
if p.CertVersion != 0 {
|
||||||
|
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, uint64(p.CertVersion))
|
||||||
|
}
|
||||||
|
|
||||||
|
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
||||||
|
out = protowire.AppendBytes(out, details)
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalPayload decodes a protobuf-encoded NebulaHandshake message.
|
||||||
|
func UnmarshalPayload(b []byte) (Payload, error) {
|
||||||
|
var p Payload
|
||||||
|
|
||||||
|
for len(b) > 0 {
|
||||||
|
num, typ, n := protowire.ConsumeTag(b)
|
||||||
|
if n < 0 {
|
||||||
|
return p, errInvalidHandshakeMessage
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case num == 1 && typ == protowire.BytesType:
|
||||||
|
details, n := protowire.ConsumeBytes(b)
|
||||||
|
if n < 0 {
|
||||||
|
return p, errInvalidHandshakeMessage
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
if err := unmarshalPayloadDetails(&p, details); err != nil {
|
||||||
|
return p, err
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||||
|
if n < 0 {
|
||||||
|
return p, errInvalidHandshakeMessage
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return p, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func unmarshalPayloadDetails(p *Payload, b []byte) error {
|
||||||
|
for len(b) > 0 {
|
||||||
|
num, typ, n := protowire.ConsumeTag(b)
|
||||||
|
if n < 0 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
|
||||||
|
// For known field numbers, reject any non-matching wire type as a
|
||||||
|
// hard error rather than silently skipping. The caller will catch
|
||||||
|
// missing-field cases downstream, but a wire-type mismatch on a tag
|
||||||
|
// we know is a peer protocol violation worth flagging here.
|
||||||
|
// Repeated occurrences of a singular field follow proto3 last-wins.
|
||||||
|
switch num {
|
||||||
|
case fieldCert:
|
||||||
|
if typ != protowire.BytesType {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
v, n := protowire.ConsumeBytes(b)
|
||||||
|
if n < 0 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
p.Cert = append([]byte(nil), v...)
|
||||||
|
b = b[n:]
|
||||||
|
case fieldInitiatorIndex:
|
||||||
|
if typ != protowire.VarintType {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
v, n := protowire.ConsumeVarint(b)
|
||||||
|
if n < 0 || v > math.MaxUint32 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
p.InitiatorIndex = uint32(v)
|
||||||
|
b = b[n:]
|
||||||
|
case fieldResponderIndex:
|
||||||
|
if typ != protowire.VarintType {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
v, n := protowire.ConsumeVarint(b)
|
||||||
|
if n < 0 || v > math.MaxUint32 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
p.ResponderIndex = uint32(v)
|
||||||
|
b = b[n:]
|
||||||
|
case fieldTime:
|
||||||
|
if typ != protowire.VarintType {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
v, n := protowire.ConsumeVarint(b)
|
||||||
|
if n < 0 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
p.Time = v
|
||||||
|
b = b[n:]
|
||||||
|
case fieldCertVersion:
|
||||||
|
if typ != protowire.VarintType {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
v, n := protowire.ConsumeVarint(b)
|
||||||
|
if n < 0 || v > math.MaxUint32 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
p.CertVersion = uint32(v)
|
||||||
|
b = b[n:]
|
||||||
|
default:
|
||||||
|
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||||
|
if n < 0 {
|
||||||
|
return errInvalidHandshakeDetails
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,361 @@
|
|||||||
|
package handshake
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPayloadRoundTrip(t *testing.T) {
|
||||||
|
t.Run("all fields set", func(t *testing.T) {
|
||||||
|
data := MarshalPayload(nil, Payload{
|
||||||
|
Cert: []byte("test-cert-bytes"),
|
||||||
|
CertVersion: 2,
|
||||||
|
InitiatorIndex: 12345,
|
||||||
|
ResponderIndex: 67890,
|
||||||
|
Time: 1234567890,
|
||||||
|
})
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, []byte("test-cert-bytes"), got.Cert)
|
||||||
|
assert.Equal(t, uint32(12345), got.InitiatorIndex)
|
||||||
|
assert.Equal(t, uint32(67890), got.ResponderIndex)
|
||||||
|
assert.Equal(t, uint64(1234567890), got.Time)
|
||||||
|
assert.Equal(t, uint32(2), got.CertVersion)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("minimal fields", func(t *testing.T) {
|
||||||
|
data := MarshalPayload(nil, Payload{InitiatorIndex: 1})
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, uint32(1), got.InitiatorIndex)
|
||||||
|
assert.Equal(t, uint32(0), got.ResponderIndex)
|
||||||
|
assert.Equal(t, uint64(0), got.Time)
|
||||||
|
assert.Nil(t, got.Cert)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty payload", func(t *testing.T) {
|
||||||
|
data := MarshalPayload(nil, Payload{})
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("large cert bytes", func(t *testing.T) {
|
||||||
|
bigCert := make([]byte, 4096)
|
||||||
|
for i := range bigCert {
|
||||||
|
bigCert[i] = byte(i % 256)
|
||||||
|
}
|
||||||
|
|
||||||
|
data := MarshalPayload(nil, Payload{
|
||||||
|
Cert: bigCert,
|
||||||
|
CertVersion: 2,
|
||||||
|
InitiatorIndex: 999,
|
||||||
|
})
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
assert.Equal(t, bigCert, got.Cert)
|
||||||
|
assert.Equal(t, uint32(999), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("append to existing buffer", func(t *testing.T) {
|
||||||
|
prefix := []byte("prefix")
|
||||||
|
data := MarshalPayload(prefix, Payload{InitiatorIndex: 42})
|
||||||
|
|
||||||
|
assert.Equal(t, []byte("prefix"), data[:6])
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data[6:])
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPayloadUnknownFields(t *testing.T) {
|
||||||
|
t.Run("unknown field in outer message is skipped", func(t *testing.T) {
|
||||||
|
// Marshal a normal payload then append an unknown field (field 99, varint)
|
||||||
|
data := MarshalPayload(nil, Payload{InitiatorIndex: 42})
|
||||||
|
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
||||||
|
data = protowire.AppendVarint(data, 12345)
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("unknown field in details is skipped", func(t *testing.T) {
|
||||||
|
// Build details with a known field + unknown field
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 77)
|
||||||
|
// Unknown field 50, varint
|
||||||
|
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 9999)
|
||||||
|
// Another known field after the unknown one
|
||||||
|
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 88)
|
||||||
|
|
||||||
|
// Wrap in outer message
|
||||||
|
var data []byte
|
||||||
|
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||||
|
data = protowire.AppendBytes(data, details)
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(77), got.InitiatorIndex)
|
||||||
|
assert.Equal(t, uint32(88), got.ResponderIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("reserved fields 6 and 7 are skipped", func(t *testing.T) {
|
||||||
|
// Fields 6 and 7 are reserved in the proto definition
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 100)
|
||||||
|
details = protowire.AppendTag(details, 6, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 1)
|
||||||
|
details = protowire.AppendTag(details, 7, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 2)
|
||||||
|
|
||||||
|
var data []byte
|
||||||
|
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||||
|
data = protowire.AppendBytes(data, details)
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(100), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPayloadBytesConsumed(t *testing.T) {
|
||||||
|
t.Run("all bytes consumed on valid input", func(t *testing.T) {
|
||||||
|
original := Payload{
|
||||||
|
Cert: []byte("cert"),
|
||||||
|
CertVersion: 2,
|
||||||
|
InitiatorIndex: 100,
|
||||||
|
ResponderIndex: 200,
|
||||||
|
Time: 999,
|
||||||
|
}
|
||||||
|
data := MarshalPayload(nil, original)
|
||||||
|
|
||||||
|
got, err := UnmarshalPayload(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Re-marshal and compare — proves we consumed and reproduced all fields
|
||||||
|
remarshaled := MarshalPayload(nil, got)
|
||||||
|
assert.Equal(t, data, remarshaled)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// wrapDetails wraps raw detail bytes in the outer NebulaHandshake envelope
|
||||||
|
// so UnmarshalPayload can reach unmarshalPayloadDetails.
|
||||||
|
func wrapDetails(details []byte) []byte {
|
||||||
|
var out []byte
|
||||||
|
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
||||||
|
out = protowire.AppendBytes(out, details)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPayloadUnmarshalErrors(t *testing.T) {
|
||||||
|
t.Run("nil input", func(t *testing.T) {
|
||||||
|
got, err := UnmarshalPayload(nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated outer tag", func(t *testing.T) {
|
||||||
|
_, err := UnmarshalPayload([]byte{0x80})
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated outer details field", func(t *testing.T) {
|
||||||
|
_, err := UnmarshalPayload([]byte{0x0a, 0x64, 0x01, 0x02, 0x03, 0x04, 0x05})
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated outer unknown field", func(t *testing.T) {
|
||||||
|
// Valid tag for unknown field 99 varint, but no value follows
|
||||||
|
var data []byte
|
||||||
|
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
||||||
|
_, err := UnmarshalPayload(data)
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated details tag", func(t *testing.T) {
|
||||||
|
_, err := UnmarshalPayload(wrapDetails([]byte{0x80}))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated cert bytes", func(t *testing.T) {
|
||||||
|
// Field 1 (cert), bytes type, length 10 but only 2 bytes
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
||||||
|
details = append(details, 0x0a, 0x01, 0x02) // length 10, only 2 bytes
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated initiator index varint", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = append(details, 0x80) // incomplete varint
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated responder index varint", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||||
|
details = append(details, 0x80)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated time varint", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
||||||
|
details = append(details, 0x80)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated cert version varint", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||||
|
details = append(details, 0x80)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("truncated unknown field in details", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
||||||
|
details = append(details, 0x80) // incomplete varint
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("cert with wrong wire type rejected", func(t *testing.T) {
|
||||||
|
// fieldCert as Varint instead of Bytes.
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldCert, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 42)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("initiator index with wrong wire type rejected", func(t *testing.T) {
|
||||||
|
// fieldInitiatorIndex as Bytes instead of Varint.
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.BytesType)
|
||||||
|
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("time with wrong wire type rejected", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldTime, protowire.BytesType)
|
||||||
|
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("cert version with wrong wire type rejected", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldCertVersion, protowire.BytesType)
|
||||||
|
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("repeated singular field follows proto3 last-wins", func(t *testing.T) {
|
||||||
|
// Per proto3, multiple instances of a singular field are accepted and
|
||||||
|
// the last value wins. We keep this behavior so that peers using
|
||||||
|
// alternative encoders aren't rejected.
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 1)
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, 42)
|
||||||
|
got, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("initiator index varint overflow rejected", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("cert version varint overflow rejected", func(t *testing.T) {
|
||||||
|
var details []byte
|
||||||
|
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||||
|
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
||||||
|
_, err := UnmarshalPayload(wrapDetails(details))
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
// FuzzPayload feeds arbitrary bytes through UnmarshalPayload to confirm it
|
||||||
|
// never panics, and for any input that parses cleanly, that re-marshal +
|
||||||
|
// re-parse is a fix-point. Inputs come from an authenticated peer (post-
|
||||||
|
// noise-decrypt), so the threat model is "valid peer behaving arbitrarily,"
|
||||||
|
// not "unauthenticated injection."
|
||||||
|
func FuzzPayload(f *testing.F) {
|
||||||
|
// Seed corpus with a handful of known-good shapes.
|
||||||
|
f.Add(MarshalPayload(nil, Payload{}))
|
||||||
|
f.Add(MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2}))
|
||||||
|
f.Add(MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1}))
|
||||||
|
f.Add(MarshalPayload(nil, Payload{
|
||||||
|
Cert: []byte("seed-cert"),
|
||||||
|
InitiatorIndex: 1,
|
||||||
|
ResponderIndex: 2,
|
||||||
|
Time: 3,
|
||||||
|
CertVersion: 2,
|
||||||
|
}))
|
||||||
|
f.Add([]byte{})
|
||||||
|
f.Add([]byte{0xff})
|
||||||
|
|
||||||
|
f.Fuzz(func(t *testing.T, data []byte) {
|
||||||
|
p1, err := UnmarshalPayload(data)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// For any input that parses, re-marshaling and re-parsing must
|
||||||
|
// yield an equivalent Payload. This catches dispatch bugs (e.g.
|
||||||
|
// emitting a field on marshal that we don't accept on parse) and
|
||||||
|
// any non-idempotent parsing behavior.
|
||||||
|
b2 := MarshalPayload(nil, p1)
|
||||||
|
p2, err := UnmarshalPayload(b2)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("re-parse of self-marshaled payload failed: %v\nintermediate: %x\n", err, b2)
|
||||||
|
}
|
||||||
|
if !payloadsEqual(p1, p2) {
|
||||||
|
t.Fatalf("re-marshal not idempotent\nfirst: %+v\nsecond: %+v", p1, p2)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func payloadsEqual(a, b Payload) bool {
|
||||||
|
return bytes.Equal(a.Cert, b.Cert) &&
|
||||||
|
a.InitiatorIndex == b.InitiatorIndex &&
|
||||||
|
a.ResponderIndex == b.ResponderIndex &&
|
||||||
|
a.Time == b.Time &&
|
||||||
|
a.CertVersion == b.CertVersion
|
||||||
|
}
|
||||||
-813
@@ -1,813 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
|
||||||
"github.com/slackhq/nebula/header"
|
|
||||||
)
|
|
||||||
|
|
||||||
// NOISE IX Handshakes
|
|
||||||
|
|
||||||
// This function constructs a handshake packet, but does not actually send it
|
|
||||||
// Sending is done by the handshake manager
|
|
||||||
func ixHandshakeStage0(f *Interface, hh *HandshakeHostInfo) bool {
|
|
||||||
err := f.handshakeManager.allocateIndex(hh)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to generate index",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
cs := f.pki.getCertState()
|
|
||||||
v := cs.initiatingVersion
|
|
||||||
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
|
||||||
v = hh.initiatingVersionOverride
|
|
||||||
} else if v < cert.Version2 {
|
|
||||||
// If we're connecting to a v6 address we should encourage use of a V2 cert
|
|
||||||
for _, a := range hh.hostinfo.vpnAddrs {
|
|
||||||
if a.Is6() {
|
|
||||||
v = cert.Version2
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
crt := cs.getCertificate(v)
|
|
||||||
if crt == nil {
|
|
||||||
f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
"certVersion", v,
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
crtHs := cs.getHandshakeBytes(v)
|
|
||||||
if crtHs == nil {
|
|
||||||
f.l.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
"certVersion", v,
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
ci, err := NewConnectionState(cs, crt, true, noise.HandshakeIX)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to create connection state",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
"certVersion", v,
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
hh.hostinfo.ConnectionState = ci
|
|
||||||
|
|
||||||
hs := &NebulaHandshake{
|
|
||||||
Details: &NebulaHandshakeDetails{
|
|
||||||
InitiatorIndex: hh.hostinfo.localIndexId,
|
|
||||||
Time: uint64(time.Now().UnixNano()),
|
|
||||||
Cert: crtHs,
|
|
||||||
CertVersion: uint32(v),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
hsBytes, err := hs.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to marshal handshake message",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"certVersion", v,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
h := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, 0, 1)
|
|
||||||
|
|
||||||
msg, _, _, err := ci.H.WriteMessage(h, hsBytes)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to call noise.WriteMessage",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// We are sending handshake packet 1, so we don't expect to receive
|
|
||||||
// handshake packet 1 from the responder
|
|
||||||
ci.window.Update(f.l, 1)
|
|
||||||
|
|
||||||
hh.hostinfo.HandshakePacket[0] = msg
|
|
||||||
hh.ready = true
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func ixHandshakeStage1(f *Interface, via ViaSender, packet []byte, h *header.H) {
|
|
||||||
cs := f.pki.getCertState()
|
|
||||||
crt := cs.GetDefaultCertificate()
|
|
||||||
if crt == nil {
|
|
||||||
f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
|
||||||
"certVersion", cs.initiatingVersion,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
ci, err := NewConnectionState(cs, crt, false, noise.HandshakeIX)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to create connection state",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Mark packet 1 as seen so it doesn't show up as missed
|
|
||||||
ci.window.Update(f.l, 1)
|
|
||||||
|
|
||||||
msg, _, _, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to call noise.ReadMessage",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hs := &NebulaHandshake{}
|
|
||||||
err = hs.Unmarshal(msg)
|
|
||||||
if err != nil || hs.Details == nil {
|
|
||||||
f.l.Error("Failed unmarshal handshake message",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
|
||||||
if err != nil {
|
|
||||||
f.l.Info("Handshake did not contain a certificate",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
|
||||||
if err != nil {
|
|
||||||
fp, fperr := rc.Fingerprint()
|
|
||||||
if fperr != nil {
|
|
||||||
fp = "<error generating certificate fingerprint>"
|
|
||||||
}
|
|
||||||
|
|
||||||
attrs := []slog.Attr{
|
|
||||||
slog.Any("error", err),
|
|
||||||
slog.Any("from", via),
|
|
||||||
slog.Any("handshake", m{"stage": 1, "style": "ix_psk0"}),
|
|
||||||
slog.Any("certVpnNetworks", rc.Networks()),
|
|
||||||
slog.String("certFingerprint", fp),
|
|
||||||
}
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
attrs = append(attrs, slog.Any("cert", rc))
|
|
||||||
}
|
|
||||||
|
|
||||||
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
|
||||||
// callers grow conditionally, which has no pair-form equivalent.
|
|
||||||
//nolint:sloglint
|
|
||||||
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
|
||||||
f.l.Info("public key mismatch between certificate and handshake",
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
"cert", remoteCert,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if remoteCert.Certificate.Version() != ci.myCert.Version() {
|
|
||||||
// We started off using the wrong certificate version, lets see if we can match the version that was sent to us
|
|
||||||
myCertOtherVersion := cs.getCertificate(remoteCert.Certificate.Version())
|
|
||||||
if myCertOtherVersion == nil {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
f.l.Debug("Might be unable to handshake with host due to missing certificate version",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
"cert", remoteCert,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// Record the certificate we are actually using
|
|
||||||
ci.myCert = myCertOtherVersion
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(remoteCert.Certificate.Networks()) == 0 {
|
|
||||||
f.l.Info("No networks in certificate",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"cert", remoteCert,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
certName := remoteCert.Certificate.Name()
|
|
||||||
certVersion := remoteCert.Certificate.Version()
|
|
||||||
fingerprint := remoteCert.Fingerprint
|
|
||||||
issuer := remoteCert.Certificate.Issuer()
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
|
||||||
f.l.Error("Refusing to handshake with myself",
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
// We only want to apply the remote allow list for direct tunnels here
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
myIndex, err := generateIndex(f.l)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to generate index",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
ConnectionState: ci,
|
|
||||||
localIndexId: myIndex,
|
|
||||||
remoteIndexId: hs.Details.InitiatorIndex,
|
|
||||||
vpnAddrs: vpnAddrs,
|
|
||||||
HandshakePacket: make(map[uint8][]byte, 0),
|
|
||||||
lastHandshakeTime: hs.Details.Time,
|
|
||||||
relayState: RelayState{
|
|
||||||
relays: nil,
|
|
||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
msgRxL := f.l.With(
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
|
|
||||||
if anyVpnAddrsInCommon {
|
|
||||||
msgRxL.Info("Handshake message received")
|
|
||||||
} else {
|
|
||||||
//todo warn if not lighthouse or relay?
|
|
||||||
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
|
||||||
}
|
|
||||||
|
|
||||||
hs.Details.ResponderIndex = myIndex
|
|
||||||
hs.Details.Cert = cs.getHandshakeBytes(ci.myCert.Version())
|
|
||||||
if hs.Details.Cert == nil {
|
|
||||||
msgRxL.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
|
||||||
"myCertVersion", ci.myCert.Version(),
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hs.Details.CertVersion = uint32(ci.myCert.Version())
|
|
||||||
// Update the time in case their clock is way off from ours
|
|
||||||
hs.Details.Time = uint64(time.Now().UnixNano())
|
|
||||||
|
|
||||||
hsBytes, err := hs.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to marshal handshake message",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
nh := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, hs.Details.InitiatorIndex, 2)
|
|
||||||
msg, dKey, eKey, err := ci.H.WriteMessage(nh, hsBytes)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to call noise.WriteMessage",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
} else if dKey == nil || eKey == nil {
|
|
||||||
f.l.Error("Noise did not arrive at a key",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.HandshakePacket[0] = make([]byte, len(packet[header.Len:]))
|
|
||||||
copy(hostinfo.HandshakePacket[0], packet[header.Len:])
|
|
||||||
|
|
||||||
// Regardless of whether you are the sender or receiver, you should arrive here
|
|
||||||
// and complete standing up the connection.
|
|
||||||
hostinfo.HandshakePacket[2] = make([]byte, len(msg))
|
|
||||||
copy(hostinfo.HandshakePacket[2], msg)
|
|
||||||
|
|
||||||
// We are sending handshake packet 2, so we don't expect to receive
|
|
||||||
// handshake packet 2 from the initiator.
|
|
||||||
ci.window.Update(f.l, 2)
|
|
||||||
|
|
||||||
ci.peerCert = remoteCert
|
|
||||||
ci.dKey = NewNebulaCipherState(dKey)
|
|
||||||
ci.eKey = NewNebulaCipherState(eKey)
|
|
||||||
|
|
||||||
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
}
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
existing, err := f.handshakeManager.CheckAndComplete(hostinfo, 0, f)
|
|
||||||
if err != nil {
|
|
||||||
switch err {
|
|
||||||
case ErrAlreadySeen:
|
|
||||||
// Update remote if preferred
|
|
||||||
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
|
||||||
// Send a test packet to ensure the other side has also switched to
|
|
||||||
// the preferred remote
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
}
|
|
||||||
|
|
||||||
msg = existing.HandshakePacket[2]
|
|
||||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
|
||||||
if !via.IsRelayed {
|
|
||||||
err := f.outside.WriteTo(msg, via.UdpAddr)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to send handshake message",
|
|
||||||
"vpnAddrs", existing.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"cached", true,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
f.l.Info("Handshake message sent",
|
|
||||||
"vpnAddrs", existing.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"cached", true,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
} else {
|
|
||||||
if via.relay == nil {
|
|
||||||
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
|
||||||
f.l.Info("Handshake message sent",
|
|
||||||
"vpnAddrs", existing.vpnAddrs,
|
|
||||||
"relay", via.relayHI.vpnAddrs[0],
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"cached", true,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
case ErrExistingHostInfo:
|
|
||||||
// This means there was an existing tunnel and this handshake was older than the one we are currently based on
|
|
||||||
f.l.Info("Handshake too old",
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"oldHandshakeTime", existing.lastHandshakeTime,
|
|
||||||
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
|
|
||||||
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
return
|
|
||||||
case ErrLocalIndexCollision:
|
|
||||||
// This means we failed to insert because of collision on localIndexId. Just let the next handshake packet retry
|
|
||||||
f.l.Error("Failed to add HostInfo due to localIndex collision",
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
"localIndex", hostinfo.localIndexId,
|
|
||||||
"collision", existing.vpnAddrs,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
default:
|
|
||||||
// Shouldn't happen, but just in case someone adds a new error type to CheckAndComplete
|
|
||||||
// And we forget to update it here
|
|
||||||
f.l.Error("Failed to add HostInfo to HostMap",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Do the send
|
|
||||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
|
||||||
if !via.IsRelayed {
|
|
||||||
err = f.outside.WriteTo(msg, via.UdpAddr)
|
|
||||||
log := f.l.With(
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
log.Error("Failed to send handshake", "error", err)
|
|
||||||
} else {
|
|
||||||
log.Info("Handshake message sent")
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
if via.relay == nil {
|
|
||||||
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
// I successfully received a handshake. Just in case I marked this tunnel as 'Disestablished', ensure
|
|
||||||
// it's correctly marked as working.
|
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
|
||||||
f.l.Info("Handshake message sent",
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"relay", via.relayHI.vpnAddrs[0],
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func ixHandshakeStage2(f *Interface, via ViaSender, hh *HandshakeHostInfo, packet []byte, h *header.H) bool {
|
|
||||||
if hh == nil {
|
|
||||||
// Nothing here to tear down, got a bogus stage 2 packet
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
hh.Lock()
|
|
||||||
defer hh.Unlock()
|
|
||||||
|
|
||||||
hostinfo := hh.hostinfo
|
|
||||||
if !via.IsRelayed {
|
|
||||||
// The vpnAddr we know about is the one we tried to handshake with, use it to apply the remote allow list.
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
ci := hostinfo.ConnectionState
|
|
||||||
msg, eKey, dKey, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to call noise.ReadMessage",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"header", h,
|
|
||||||
)
|
|
||||||
|
|
||||||
// We don't want to tear down the connection on a bad ReadMessage because it could be an attacker trying
|
|
||||||
// to DOS us. Every other error condition after should to allow a possible good handshake to complete in the
|
|
||||||
// near future
|
|
||||||
return false
|
|
||||||
} else if dKey == nil || eKey == nil {
|
|
||||||
f.l.Error("Noise did not arrive at a key",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
|
|
||||||
// This should be impossible in IX but just in case, if we get here then there is no chance to recover
|
|
||||||
// the handshake state machine. Tear it down
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
hs := &NebulaHandshake{}
|
|
||||||
err = hs.Unmarshal(msg)
|
|
||||||
if err != nil || hs.Details == nil {
|
|
||||||
f.l.Error("Failed unmarshal handshake message",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
|
|
||||||
// The handshake state machine is complete, if things break now there is no chance to recover. Tear down and start again
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
|
||||||
if err != nil {
|
|
||||||
f.l.Info("Handshake did not contain a certificate",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
|
||||||
if err != nil {
|
|
||||||
fp, err := rc.Fingerprint()
|
|
||||||
if err != nil {
|
|
||||||
fp = "<error generating certificate fingerprint>"
|
|
||||||
}
|
|
||||||
|
|
||||||
attrs := []slog.Attr{
|
|
||||||
slog.Any("error", err),
|
|
||||||
slog.Any("from", via),
|
|
||||||
slog.Any("vpnAddrs", hostinfo.vpnAddrs),
|
|
||||||
slog.Any("handshake", m{"stage": 2, "style": "ix_psk0"}),
|
|
||||||
slog.String("certFingerprint", fp),
|
|
||||||
slog.Any("certVpnNetworks", rc.Networks()),
|
|
||||||
}
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
attrs = append(attrs, slog.Any("cert", rc))
|
|
||||||
}
|
|
||||||
|
|
||||||
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
|
||||||
// callers grow conditionally, which has no pair-form equivalent.
|
|
||||||
//nolint:sloglint
|
|
||||||
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
|
||||||
f.l.Info("public key mismatch between certificate and handshake",
|
|
||||||
"from", via,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"cert", remoteCert,
|
|
||||||
)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(remoteCert.Certificate.Networks()) == 0 {
|
|
||||||
f.l.Info("No networks in certificate",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"cert", remoteCert,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
certName := remoteCert.Certificate.Name()
|
|
||||||
certVersion := remoteCert.Certificate.Version()
|
|
||||||
fingerprint := remoteCert.Fingerprint
|
|
||||||
issuer := remoteCert.Certificate.Issuer()
|
|
||||||
|
|
||||||
hostinfo.remoteIndexId = hs.Details.ResponderIndex
|
|
||||||
hostinfo.lastHandshakeTime = hs.Details.Time
|
|
||||||
|
|
||||||
// Store their cert and our symmetric keys
|
|
||||||
ci.peerCert = remoteCert
|
|
||||||
ci.dKey = NewNebulaCipherState(dKey)
|
|
||||||
ci.eKey = NewNebulaCipherState(eKey)
|
|
||||||
|
|
||||||
// Make sure the current udpAddr being used is set for responding
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
} else {
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
correctHostResponded := false
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
if hostinfo.vpnAddrs[0] == network.Addr() {
|
|
||||||
// todo is it more correct to see if any of hostinfo.vpnAddrs are in the cert? it should have len==1, but one day it might not?
|
|
||||||
correctHostResponded = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ensure the right host responded
|
|
||||||
if !correctHostResponded {
|
|
||||||
f.l.Info("Incorrect host responded to handshake",
|
|
||||||
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"haveVpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
)
|
|
||||||
|
|
||||||
// Release our old handshake from pending, it should not continue
|
|
||||||
f.handshakeManager.DeleteHostInfo(hostinfo)
|
|
||||||
|
|
||||||
// Create a new hostinfo/handshake for the intended vpn ip
|
|
||||||
//TODO is hostinfo.vpnAddrs[0] always the address to use?
|
|
||||||
f.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
|
||||||
// Block the current used address
|
|
||||||
newHH.hostinfo.remotes = hostinfo.remotes
|
|
||||||
newHH.hostinfo.remotes.BlockRemote(via)
|
|
||||||
|
|
||||||
f.l.Info("Blocked addresses for handshakes",
|
|
||||||
"blockedUdpAddrs", newHH.hostinfo.remotes.CopyBlockedRemotes(),
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"remotes", newHH.hostinfo.remotes.CopyAddrs(f.hostMap.GetPreferredRanges()),
|
|
||||||
)
|
|
||||||
|
|
||||||
// Swap the packet store to benefit the original intended recipient
|
|
||||||
newHH.packetStore = hh.packetStore
|
|
||||||
hh.packetStore = []*cachedPacket{}
|
|
||||||
|
|
||||||
// Finally, put the correct vpn addrs in the host info, tell them to close the tunnel, and return true to tear down
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
f.sendCloseTunnel(hostinfo)
|
|
||||||
})
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// Mark packet 2 as seen so it doesn't show up as missed
|
|
||||||
ci.window.Update(f.l, 2)
|
|
||||||
|
|
||||||
duration := time.Since(hh.startTime).Nanoseconds()
|
|
||||||
msgRxL := f.l.With(
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", certName,
|
|
||||||
"certVersion", certVersion,
|
|
||||||
"fingerprint", fingerprint,
|
|
||||||
"issuer", issuer,
|
|
||||||
"initiatorIndex", hs.Details.InitiatorIndex,
|
|
||||||
"responderIndex", hs.Details.ResponderIndex,
|
|
||||||
"remoteIndex", h.RemoteIndex,
|
|
||||||
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
|
||||||
"durationNs", duration,
|
|
||||||
"sentCachedPackets", len(hh.packetStore),
|
|
||||||
)
|
|
||||||
if anyVpnAddrsInCommon {
|
|
||||||
msgRxL.Info("Handshake message received")
|
|
||||||
} else {
|
|
||||||
//todo warn if not lighthouse or relay?
|
|
||||||
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build up the radix for the firewall if we have subnets in the cert
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
// Complete our handshake and update metrics, this will replace any existing tunnels for the vpnAddrs here
|
|
||||||
f.handshakeManager.Complete(hostinfo, f)
|
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("Sending stored packets",
|
|
||||||
"count", len(hh.packetStore),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(hh.packetStore) > 0 {
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
out := make([]byte, mtu)
|
|
||||||
for _, cp := range hh.packetStore {
|
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
|
||||||
}
|
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
|
||||||
f.metricHandshakes.Update(duration)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
+610
-169
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/handshake"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
@@ -22,7 +23,18 @@ const (
|
|||||||
DefaultHandshakeTryInterval = time.Millisecond * 100
|
DefaultHandshakeTryInterval = time.Millisecond * 100
|
||||||
DefaultHandshakeRetries = 10
|
DefaultHandshakeRetries = 10
|
||||||
DefaultHandshakeTriggerBuffer = 64
|
DefaultHandshakeTriggerBuffer = 64
|
||||||
DefaultUseRelays = true
|
|
||||||
|
// maxCachedPackets is how many unsent packets we'll buffer per pending
|
||||||
|
// handshake before dropping further ones.
|
||||||
|
maxCachedPackets = 100
|
||||||
|
|
||||||
|
// HandshakePacket map keys mirror the IX protocol stage convention:
|
||||||
|
// stage 0 = the initiator's first message (and what the responder
|
||||||
|
// receives, stripped of header)
|
||||||
|
// stage 2 = the responder's reply
|
||||||
|
// Other handshake patterns will need new keys when added.
|
||||||
|
handshakePacketStage0 uint8 = 0
|
||||||
|
handshakePacketStage2 uint8 = 2
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -30,7 +42,6 @@ var (
|
|||||||
tryInterval: DefaultHandshakeTryInterval,
|
tryInterval: DefaultHandshakeTryInterval,
|
||||||
retries: DefaultHandshakeRetries,
|
retries: DefaultHandshakeRetries,
|
||||||
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
||||||
useRelays: DefaultUseRelays,
|
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -38,7 +49,6 @@ type HandshakeConfig struct {
|
|||||||
tryInterval time.Duration
|
tryInterval time.Duration
|
||||||
retries int64
|
retries int64
|
||||||
triggerBuffer int
|
triggerBuffer int
|
||||||
useRelays bool
|
|
||||||
|
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
}
|
}
|
||||||
@@ -76,10 +86,11 @@ type HandshakeHostInfo struct {
|
|||||||
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
||||||
|
|
||||||
hostinfo *HostInfo
|
hostinfo *HostInfo
|
||||||
|
machine *handshake.Machine // The handshake state machine, set during stage 0 (initiator) or beginHandshake (responder multi-message)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
||||||
if len(hh.packetStore) < 100 {
|
if len(hh.packetStore) < maxCachedPackets {
|
||||||
tempPacket := make([]byte, len(packet))
|
tempPacket := make([]byte, len(packet))
|
||||||
copy(tempPacket, packet)
|
copy(tempPacket, packet)
|
||||||
|
|
||||||
@@ -137,6 +148,18 @@ func (hm *HandshakeManager) Run(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
||||||
|
// Gate on known handshake subtypes. Unknown subtypes (or future ones we
|
||||||
|
// don't yet support) are dropped here rather than silently routed through
|
||||||
|
// the IX path. Add a case when introducing a new pattern.
|
||||||
|
switch h.Subtype {
|
||||||
|
case header.HandshakeIXPSK0:
|
||||||
|
// supported
|
||||||
|
default:
|
||||||
|
hm.l.Debug("dropping handshake with unsupported subtype",
|
||||||
|
"from", via, "subtype", h.Subtype)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// First remote allow list check before we know the vpnIp
|
// First remote allow list check before we know the vpnIp
|
||||||
if !via.IsRelayed {
|
if !via.IsRelayed {
|
||||||
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
||||||
@@ -145,19 +168,27 @@ func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *head
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
switch h.Subtype {
|
// First message of a new handshake. The wire format requires RemoteIndex
|
||||||
case header.HandshakeIXPSK0:
|
// to be zero here (the initiator has no responder index to fill in yet),
|
||||||
switch h.MessageCounter {
|
// and generateIndex never allocates 0, so any non-zero RemoteIndex on a
|
||||||
case 1:
|
// stage-1 packet is malformed or someone probing for an index collision.
|
||||||
ixHandshakeStage1(hm.f, via, packet, h)
|
// Drop without paying the cost of running noise on a pending Machine.
|
||||||
|
if h.MessageCounter == 1 {
|
||||||
|
if h.RemoteIndex != 0 {
|
||||||
|
hm.l.Debug("dropping stage-1 handshake with non-zero RemoteIndex",
|
||||||
|
"from", via, "remoteIndex", h.RemoteIndex)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hm.beginHandshake(via, packet, h)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
case 2:
|
// Continuation message must match a pending handshake by index.
|
||||||
newHostinfo := hm.queryIndex(h.RemoteIndex)
|
// Anything else is an orphaned packet (e.g., late retransmit after
|
||||||
tearDown := ixHandshakeStage2(hm.f, via, newHostinfo, packet, h)
|
// timeout) and is dropped.
|
||||||
if tearDown && newHostinfo != nil {
|
if hh := hm.queryIndex(h.RemoteIndex); hh != nil {
|
||||||
hm.DeleteHostInfo(newHostinfo.hostinfo)
|
hm.continueHandshake(via, hh, packet)
|
||||||
}
|
return
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -183,13 +214,22 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo := hh.hostinfo
|
hostinfo := hh.hostinfo
|
||||||
// If we are out of time, clean up
|
// If we are out of time, clean up
|
||||||
if hh.counter >= hm.config.retries {
|
if hh.counter >= hm.config.retries {
|
||||||
hh.hostinfo.logger(hm.l).Info("Handshake timed out",
|
fields := []any{
|
||||||
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
||||||
"initiatorIndex", hh.hostinfo.localIndexId,
|
"initiatorIndex", hh.hostinfo.localIndexId,
|
||||||
"remoteIndex", hh.hostinfo.remoteIndexId,
|
"remoteIndex", hh.hostinfo.remoteIndexId,
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
|
||||||
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
||||||
)
|
}
|
||||||
|
// hh.machine can be nil here if buildStage0Packet never succeeded
|
||||||
|
// (e.g., no certificate available). In that case there's no useful
|
||||||
|
// handshake metadata to log.
|
||||||
|
if hh.machine != nil {
|
||||||
|
fields = append(fields, "handshake", m{
|
||||||
|
"stage": uint64(hh.machine.MessageIndex()),
|
||||||
|
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
hh.hostinfo.logger(hm.l).Info("Handshake timed out", fields...)
|
||||||
hm.metricTimedOut.Inc(1)
|
hm.metricTimedOut.Inc(1)
|
||||||
hm.DeleteHostInfo(hostinfo)
|
hm.DeleteHostInfo(hostinfo)
|
||||||
return
|
return
|
||||||
@@ -200,12 +240,25 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
|
|
||||||
// Check if we have a handshake packet to transmit yet
|
// Check if we have a handshake packet to transmit yet
|
||||||
if !hh.ready {
|
if !hh.ready {
|
||||||
if !ixHandshakeStage0(hm.f, hh) {
|
if !hm.buildStage0Packet(hh) {
|
||||||
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TODO: this hardcodes "always retransmit stage 0", which is correct for
|
||||||
|
// IX (the initiator only ever sends one packet, msg1) but wrong the
|
||||||
|
// moment a 3+ message pattern lands. The retry loop should resend the
|
||||||
|
// most recent outgoing message, not always stage 0. That implies
|
||||||
|
// HandshakeHostInfo tracking a single "currentOutbound" packet (bytes +
|
||||||
|
// header metadata) that gets replaced as the handshake progresses,
|
||||||
|
// instead of indexing into HandshakePacket.
|
||||||
|
stage0 := hostinfo.HandshakePacket[handshakePacketStage0]
|
||||||
|
hsFields := m{
|
||||||
|
"stage": uint64(hh.machine.MessageIndex()),
|
||||||
|
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
||||||
|
}
|
||||||
|
|
||||||
// Get a remotes object if we don't already have one.
|
// Get a remotes object if we don't already have one.
|
||||||
// This is mainly to protect us as this should never be the case
|
// This is mainly to protect us as this should never be the case
|
||||||
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
||||||
@@ -239,13 +292,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
||||||
var sentTo []netip.AddrPort
|
var sentTo []netip.AddrPort
|
||||||
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
||||||
hm.messageMetrics.Tx(header.Handshake, header.MessageSubType(hostinfo.HandshakePacket[0][1]), 1)
|
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
||||||
err := hm.outside.WriteTo(hostinfo.HandshakePacket[0], addr)
|
err := hm.outside.WriteTo(stage0, addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||||
"udpAddr", addr,
|
"udpAddr", addr,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
"handshake", hsFields,
|
||||||
"error", err,
|
"error", err,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -260,156 +313,17 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo.logger(hm.l).Info("Handshake message sent",
|
hostinfo.logger(hm.l).Info("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
"handshake", hsFields,
|
||||||
)
|
)
|
||||||
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
"handshake", hsFields,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
if hm.config.useRelays && len(hostinfo.remotes.relays) > 0 {
|
hm.f.relayManager.StartRelays(hm.f, vpnIp, hostinfo, stage0)
|
||||||
hostinfo.logger(hm.l).Info("Attempt to relay through hosts", "relays", hostinfo.remotes.relays)
|
|
||||||
// Send a RelayRequest to all known Relay IP's
|
|
||||||
for _, relay := range hostinfo.remotes.relays {
|
|
||||||
// Don't relay through the host I'm trying to connect to
|
|
||||||
if relay == vpnIp {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Don't relay to myself
|
|
||||||
if hm.f.myVpnAddrsTable.Contains(relay) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
relayHostInfo := hm.mainHostMap.QueryVpnAddr(relay)
|
|
||||||
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
|
||||||
hostinfo.logger(hm.l).Info("Establish tunnel to relay target", "relay", relay.String())
|
|
||||||
hm.f.Handshake(relay)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// Check the relay HostInfo to see if we already established a relay through
|
|
||||||
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
|
||||||
if !ok {
|
|
||||||
// No relays exist or requested yet.
|
|
||||||
if relayHostInfo.remote.IsValid() {
|
|
||||||
idx, err := AddRelay(hm.l, relayHostInfo, hm.mainHostMap, vpnIp, nil, TerminalType, Requested)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(hm.l).Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
m := NebulaControl{
|
|
||||||
Type: NebulaControl_CreateRelayRequest,
|
|
||||||
InitiatorRelayIndex: idx,
|
|
||||||
}
|
|
||||||
|
|
||||||
switch relayHostInfo.GetCert().Certificate.Version() {
|
|
||||||
case cert.Version1:
|
|
||||||
if !hm.f.myVpnAddrs[0].Is4() {
|
|
||||||
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if !vpnIp.Is4() {
|
|
||||||
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
b := hm.f.myVpnAddrs[0].As4()
|
|
||||||
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
b = vpnIp.As4()
|
|
||||||
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
case cert.Version2:
|
|
||||||
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
|
||||||
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
|
||||||
default:
|
|
||||||
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
msg, err := m.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
|
||||||
} else {
|
|
||||||
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
|
||||||
hm.l.Info("send CreateRelayRequest",
|
|
||||||
"relayFrom", hm.f.myVpnAddrs[0],
|
|
||||||
"relayTo", vpnIp,
|
|
||||||
"initiatorRelayIndex", idx,
|
|
||||||
"relay", relay,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
switch existingRelay.State {
|
|
||||||
case Established:
|
|
||||||
hostinfo.logger(hm.l).Info("Send handshake via relay", "relay", relay.String())
|
|
||||||
hm.f.SendVia(relayHostInfo, existingRelay, hostinfo.HandshakePacket[0], make([]byte, 12), make([]byte, mtu), false)
|
|
||||||
case Disestablished:
|
|
||||||
// Mark this relay as 'requested'
|
|
||||||
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
|
||||||
fallthrough
|
|
||||||
case Requested:
|
|
||||||
hostinfo.logger(hm.l).Info("Re-send CreateRelay request", "relay", relay.String())
|
|
||||||
// Re-send the CreateRelay request, in case the previous one was lost.
|
|
||||||
m := NebulaControl{
|
|
||||||
Type: NebulaControl_CreateRelayRequest,
|
|
||||||
InitiatorRelayIndex: existingRelay.LocalIndex,
|
|
||||||
}
|
|
||||||
|
|
||||||
switch relayHostInfo.GetCert().Certificate.Version() {
|
|
||||||
case cert.Version1:
|
|
||||||
if !hm.f.myVpnAddrs[0].Is4() {
|
|
||||||
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if !vpnIp.Is4() {
|
|
||||||
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
b := hm.f.myVpnAddrs[0].As4()
|
|
||||||
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
b = vpnIp.As4()
|
|
||||||
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
|
||||||
case cert.Version2:
|
|
||||||
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
|
||||||
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
|
||||||
default:
|
|
||||||
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
msg, err := m.Marshal()
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
|
||||||
} else {
|
|
||||||
// This must send over the hostinfo, not over hm.Hosts[ip]
|
|
||||||
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
|
||||||
hm.l.Info("send CreateRelayRequest",
|
|
||||||
"relayFrom", hm.f.myVpnAddrs[0],
|
|
||||||
"relayTo", vpnIp,
|
|
||||||
"initiatorRelayIndex", existingRelay.LocalIndex,
|
|
||||||
"relay", relay,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
case PeerRequested:
|
|
||||||
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
|
||||||
fallthrough
|
|
||||||
default:
|
|
||||||
hostinfo.logger(hm.l).Error("Relay unexpected state",
|
|
||||||
"vpnIp", vpnIp,
|
|
||||||
"state", existingRelay.State,
|
|
||||||
"relay", relay,
|
|
||||||
)
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
||||||
if !lighthouseTriggered {
|
if !lighthouseTriggered {
|
||||||
@@ -587,7 +501,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// allocateIndex generates a unique localIndexId for this HostInfo
|
// allocateIndex generates a unique localIndexId for this HostInfo
|
||||||
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
||||||
// a unique localIndexId
|
// a unique localIndexId
|
||||||
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) error {
|
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error) {
|
||||||
hm.mainHostMap.RLock()
|
hm.mainHostMap.RLock()
|
||||||
defer hm.mainHostMap.RUnlock()
|
defer hm.mainHostMap.RUnlock()
|
||||||
hm.Lock()
|
hm.Lock()
|
||||||
@@ -596,7 +510,7 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) error {
|
|||||||
for range 32 {
|
for range 32 {
|
||||||
index, err := generateIndex(hm.l)
|
index, err := generateIndex(hm.l)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
_, inPending := hm.indexes[index]
|
_, inPending := hm.indexes[index]
|
||||||
@@ -605,11 +519,11 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) error {
|
|||||||
if !inMain && !inPending {
|
if !inMain && !inPending {
|
||||||
hh.hostinfo.localIndexId = index
|
hh.hostinfo.localIndexId = index
|
||||||
hm.indexes[index] = hh
|
hm.indexes[index] = hh
|
||||||
return nil
|
return index, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.New("failed to generate unique localIndexId")
|
return 0, errors.New("failed to generate unique localIndexId")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||||
@@ -728,3 +642,530 @@ func generateIndex(l *slog.Logger) (uint32, error) {
|
|||||||
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
||||||
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildStage0Packet creates the initial handshake packet for the initiator.
|
||||||
|
func (hm *HandshakeManager) buildStage0Packet(hh *HandshakeHostInfo) bool {
|
||||||
|
cs := hm.f.pki.getCertState()
|
||||||
|
v := cs.DefaultVersion()
|
||||||
|
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
||||||
|
v = hh.initiatingVersionOverride
|
||||||
|
} else if v < cert.Version2 {
|
||||||
|
for _, a := range hh.hostinfo.vpnAddrs {
|
||||||
|
if a.Is6() {
|
||||||
|
v = cert.Version2
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cred := cs.GetCredential(v)
|
||||||
|
if cred == nil {
|
||||||
|
hm.f.l.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs, "certVersion", v)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
machine, err := handshake.NewMachine(
|
||||||
|
v, cs.GetCredential,
|
||||||
|
hm.certVerifier(), func() (uint32, error) { return hm.allocateIndex(hh) },
|
||||||
|
true, header.HandshakeIXPSK0,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
hm.f.l.Error("Failed to create handshake machine",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
msg, err := machine.Initiate(nil)
|
||||||
|
if err != nil {
|
||||||
|
hm.f.l.Error("Failed to initiate handshake",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// hostinfo.ConnectionState stays nil until the handshake completes in
|
||||||
|
// continueHandshake. Pre-completion control surfaces guard with nil
|
||||||
|
// checks; the data plane never observes a pending hostinfo.
|
||||||
|
hh.hostinfo.HandshakePacket[handshakePacketStage0] = msg
|
||||||
|
hh.machine = machine
|
||||||
|
hh.ready = true
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// beginHandshake handles an incoming handshake packet that doesn't match any
|
||||||
|
// existing pending handshake. It creates a new responder Machine and processes
|
||||||
|
// the first message.
|
||||||
|
func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *header.H) {
|
||||||
|
f := hm.f
|
||||||
|
cs := f.pki.getCertState()
|
||||||
|
|
||||||
|
v := cs.DefaultVersion()
|
||||||
|
if cs.GetCredential(v) == nil {
|
||||||
|
f.l.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"from", via, "certVersion", v)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
machine, err := handshake.NewMachine(
|
||||||
|
v, cs.GetCredential,
|
||||||
|
hm.certVerifier(), func() (uint32, error) { return generateIndex(f.l) },
|
||||||
|
false, header.HandshakeIXPSK0,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to create handshake machine", "from", via, "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
response, result, err := machine.ProcessPacket(nil, packet)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to process handshake packet", "from", via, "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if result == nil {
|
||||||
|
// Multi-message pattern: the responder Machine would need to be
|
||||||
|
// registered in hm.indexes so a future inbound packet finds it via
|
||||||
|
// continueHandshake. The current manager doesn't do that yet, so
|
||||||
|
// fail loudly rather than silently dropping the in-flight handshake.
|
||||||
|
// TODO: support multi-message responder flows (XX, pqIX, etc.).
|
||||||
|
// See also the IX-shaped cipher key assignment in handshake.Machine.
|
||||||
|
f.l.Error("multi-message handshake responder is not supported",
|
||||||
|
"from", via, "error", handshake.ErrMultiMessageUnsupported)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
remoteCert := result.RemoteCert
|
||||||
|
if remoteCert == nil {
|
||||||
|
f.l.Error("Handshake did not produce a peer certificate", "from", via)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate peer identity
|
||||||
|
vpnAddrs, anyVpnAddrsInCommon, ok := hm.validatePeerCert(via, remoteCert)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo := &HostInfo{
|
||||||
|
ConnectionState: newConnectionStateFromResult(result),
|
||||||
|
localIndexId: result.LocalIndex,
|
||||||
|
remoteIndexId: result.RemoteIndex,
|
||||||
|
vpnAddrs: vpnAddrs,
|
||||||
|
HandshakePacket: make(map[uint8][]byte, 0),
|
||||||
|
lastHandshakeTime: result.HandshakeTime,
|
||||||
|
relayState: RelayState{
|
||||||
|
relays: nil,
|
||||||
|
relayForByAddr: map[netip.Addr]*Relay{},
|
||||||
|
relayForByIdx: map[uint32]*Relay{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := "Handshake message received"
|
||||||
|
if !anyVpnAddrsInCommon {
|
||||||
|
msg = "Handshake message received, but no vpnNetworks in common."
|
||||||
|
}
|
||||||
|
f.l.Info(msg,
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", remoteCert.Certificate.Name(),
|
||||||
|
"certVersion", remoteCert.Certificate.Version(),
|
||||||
|
"fingerprint", remoteCert.Fingerprint,
|
||||||
|
"issuer", remoteCert.Certificate.Issuer(),
|
||||||
|
"initiatorIndex", result.RemoteIndex,
|
||||||
|
"responderIndex", result.LocalIndex,
|
||||||
|
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||||
|
)
|
||||||
|
|
||||||
|
// packet aliases the listener's incoming buffer, so this copy must stay.
|
||||||
|
hostinfo.HandshakePacket[handshakePacketStage0] = make([]byte, len(packet[header.Len:]))
|
||||||
|
copy(hostinfo.HandshakePacket[handshakePacketStage0], packet[header.Len:])
|
||||||
|
|
||||||
|
// response was freshly allocated by ProcessPacket; safe to retain directly.
|
||||||
|
if response != nil {
|
||||||
|
hostinfo.HandshakePacket[handshakePacketStage2] = response
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
}
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
existing, err := hm.CheckAndComplete(hostinfo, handshakePacketStage0, f)
|
||||||
|
if err != nil {
|
||||||
|
hm.handleCheckAndCompleteError(err, existing, hostinfo, via)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// continueHandshake feeds an incoming packet to an existing pending handshake Machine.
|
||||||
|
func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostInfo, packet []byte) {
|
||||||
|
f := hm.f
|
||||||
|
|
||||||
|
hh.Lock()
|
||||||
|
defer hh.Unlock()
|
||||||
|
|
||||||
|
// Re-verify hh is still tracked. Between queryIndex returning and us taking
|
||||||
|
// hh.Lock, handleOutbound may have timed out and deleted it. Once we hold
|
||||||
|
// hh.Lock no other deleter can race our index: handleOutbound also takes
|
||||||
|
// hh.Lock first, and handleRecvError targets a main-hostmap entry with a
|
||||||
|
// different localIndexId.
|
||||||
|
hm.RLock()
|
||||||
|
cur, ok := hm.indexes[hh.hostinfo.localIndexId]
|
||||||
|
hm.RUnlock()
|
||||||
|
if !ok || cur != hh {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo := hh.hostinfo
|
||||||
|
if !via.IsRelayed {
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
machine := hh.machine
|
||||||
|
if machine == nil {
|
||||||
|
f.l.Error("No handshake machine available for continuation",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
||||||
|
hm.DeleteHostInfo(hostinfo)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
response, result, err := machine.ProcessPacket(nil, packet)
|
||||||
|
if err != nil {
|
||||||
|
// Recoverable errors are routine noise, log at Debug. Fatal errors get a Warn.
|
||||||
|
if machine.Failed() {
|
||||||
|
f.l.Warn("Failed to process handshake packet, abandoning",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
||||||
|
hm.DeleteHostInfo(hostinfo)
|
||||||
|
} else {
|
||||||
|
f.l.Debug("Failed to process handshake packet",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if response != nil {
|
||||||
|
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
||||||
|
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
||||||
|
|
||||||
|
remoteCert := result.RemoteCert
|
||||||
|
if remoteCert == nil {
|
||||||
|
f.l.Error("Handshake completed without peer certificate",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
||||||
|
hm.DeleteHostInfo(hostinfo)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
hostinfo.remoteIndexId = result.RemoteIndex
|
||||||
|
hostinfo.lastHandshakeTime = result.HandshakeTime
|
||||||
|
|
||||||
|
if !via.IsRelayed {
|
||||||
|
hostinfo.SetRemote(via.UdpAddr)
|
||||||
|
} else {
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify correct host responded (initiator check)
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
correctHostResponded := false
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
// inside.go drops self-routed packets at the firewall stage, but we'd
|
||||||
|
// rather not let a self-handshake complete in the first place: it
|
||||||
|
// wastes a hostmap slot, suppresses no log, and obscures routing
|
||||||
|
// misconfig. Explicit refusal here mirrors the responder-side check
|
||||||
|
// in validatePeerCert.
|
||||||
|
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
||||||
|
f.l.Error("Refusing to handshake with myself",
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", remoteCert.Certificate.Name(),
|
||||||
|
"certVersion", remoteCert.Certificate.Version(),
|
||||||
|
"fingerprint", remoteCert.Fingerprint,
|
||||||
|
"issuer", remoteCert.Certificate.Issuer(),
|
||||||
|
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||||
|
)
|
||||||
|
hm.DeleteHostInfo(hostinfo)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if hostinfo.vpnAddrs[0] == network.Addr() {
|
||||||
|
correctHostResponded = true
|
||||||
|
}
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !correctHostResponded {
|
||||||
|
f.l.Info("Incorrect host responded to handshake",
|
||||||
|
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"haveVpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", remoteCert.Certificate.Name(),
|
||||||
|
"certVersion", remoteCert.Certificate.Version(),
|
||||||
|
"fingerprint", remoteCert.Fingerprint,
|
||||||
|
"issuer", remoteCert.Certificate.Issuer(),
|
||||||
|
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||||
|
)
|
||||||
|
|
||||||
|
hm.DeleteHostInfo(hostinfo)
|
||||||
|
hm.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
||||||
|
newHH.hostinfo.remotes = hostinfo.remotes
|
||||||
|
newHH.hostinfo.remotes.BlockRemote(via)
|
||||||
|
newHH.packetStore = hh.packetStore
|
||||||
|
hh.packetStore = []*cachedPacket{}
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
f.sendCloseTunnel(hostinfo)
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
duration := time.Since(hh.startTime).Nanoseconds()
|
||||||
|
msg := "Handshake message received"
|
||||||
|
if !anyVpnAddrsInCommon {
|
||||||
|
msg = "Handshake message received, but no vpnNetworks in common."
|
||||||
|
}
|
||||||
|
f.l.Info(msg,
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", remoteCert.Certificate.Name(),
|
||||||
|
"certVersion", remoteCert.Certificate.Version(),
|
||||||
|
"fingerprint", remoteCert.Fingerprint,
|
||||||
|
"issuer", remoteCert.Certificate.Issuer(),
|
||||||
|
"initiatorIndex", result.LocalIndex,
|
||||||
|
"responderIndex", result.RemoteIndex,
|
||||||
|
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
||||||
|
"durationNs", duration,
|
||||||
|
"sentCachedPackets", len(hh.packetStore),
|
||||||
|
)
|
||||||
|
|
||||||
|
hostinfo.vpnAddrs = vpnAddrs
|
||||||
|
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||||
|
|
||||||
|
hm.Complete(hostinfo, f)
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
if len(hh.packetStore) > 0 {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Sending stored packets", "count", len(hh.packetStore))
|
||||||
|
}
|
||||||
|
buf := f.bufAlloc.Acquire()
|
||||||
|
for _, cp := range hh.packetStore {
|
||||||
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, buf)
|
||||||
|
}
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
f.metricHandshakes.Update(duration)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// validatePeerCert checks the peer certificate for self-connection and remote allow list.
|
||||||
|
// Returns the VPN addrs, whether any of them fall within one of our own VPN
|
||||||
|
// networks, and true if valid; false if rejected.
|
||||||
|
func (hm *HandshakeManager) validatePeerCert(via ViaSender, remoteCert *cert.CachedCertificate) ([]netip.Addr, bool, bool) {
|
||||||
|
f := hm.f
|
||||||
|
vpnNetworks := remoteCert.Certificate.Networks()
|
||||||
|
|
||||||
|
// The cert package rejects host certs with no networks at parse time, so
|
||||||
|
// reaching this state would mean an invariant was bypassed elsewhere.
|
||||||
|
// Refuse explicitly so downstream code (which indexes vpnAddrs[0]) can't
|
||||||
|
// panic if that invariant ever changes.
|
||||||
|
if len(vpnNetworks) == 0 {
|
||||||
|
f.l.Info("No networks in certificate",
|
||||||
|
"from", via, "cert", remoteCert)
|
||||||
|
return nil, false, false
|
||||||
|
}
|
||||||
|
|
||||||
|
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||||
|
anyVpnAddrsInCommon := false
|
||||||
|
|
||||||
|
for i, network := range vpnNetworks {
|
||||||
|
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
||||||
|
f.l.Error("Refusing to handshake with myself",
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", remoteCert.Certificate.Name(),
|
||||||
|
"certVersion", remoteCert.Certificate.Version(),
|
||||||
|
"fingerprint", remoteCert.Fingerprint,
|
||||||
|
"issuer", remoteCert.Certificate.Issuer(),
|
||||||
|
)
|
||||||
|
return nil, false, false
|
||||||
|
}
|
||||||
|
vpnAddrs[i] = network.Addr()
|
||||||
|
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||||
|
anyVpnAddrsInCommon = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !via.IsRelayed {
|
||||||
|
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", vpnAddrs, "from", via)
|
||||||
|
return nil, false, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return vpnAddrs, anyVpnAddrsInCommon, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendHandshakeResponse sends a handshake response via the appropriate transport.
|
||||||
|
// cached is true when msg is a stored response being retransmitted because
|
||||||
|
// the peer's stage-1 retransmit landed (the ErrAlreadySeen path); false on a
|
||||||
|
// fresh response.
|
||||||
|
func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hostinfo *HostInfo, cached bool) {
|
||||||
|
if msg == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f := hm.f
|
||||||
|
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||||
|
|
||||||
|
// Common log fields. peerCert may be nil during intermediate
|
||||||
|
// multi-message flows (handshake hasn't completed yet); skip the cert
|
||||||
|
// block if so.
|
||||||
|
logFields := []any{
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": uint64(2), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)},
|
||||||
|
"cached", cached,
|
||||||
|
"initiatorIndex", hostinfo.remoteIndexId,
|
||||||
|
"responderIndex", hostinfo.localIndexId,
|
||||||
|
}
|
||||||
|
if peerCert := hostinfo.ConnectionState.peerCert; peerCert != nil {
|
||||||
|
logFields = append(logFields,
|
||||||
|
"certName", peerCert.Certificate.Name(),
|
||||||
|
"certVersion", peerCert.Certificate.Version(),
|
||||||
|
"fingerprint", peerCert.Fingerprint,
|
||||||
|
"issuer", peerCert.Certificate.Issuer(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !via.IsRelayed {
|
||||||
|
fields := append(logFields, "from", via)
|
||||||
|
err := f.outside.WriteTo(msg, via.UdpAddr)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to send handshake message", append(fields, "error", err)...)
|
||||||
|
} else {
|
||||||
|
f.l.Info("Handshake message sent", fields...)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if via.relay == nil {
|
||||||
|
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||||
|
// We received a valid handshake on this relay, so make sure the relay
|
||||||
|
// state reflects that, in case it had been marked Disestablished.
|
||||||
|
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||||
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendVia(via.relayHI, via.relay, msg, buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
|
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleCheckAndCompleteError handles errors from CheckAndComplete.
|
||||||
|
// This only fires from the responder-side beginHandshake path, after the
|
||||||
|
// peer cert has been validated and ConnectionState populated, so peerCert
|
||||||
|
// is always non-nil for the cases that log it.
|
||||||
|
func (hm *HandshakeManager) handleCheckAndCompleteError(err error, existing, hostinfo *HostInfo, via ViaSender) {
|
||||||
|
f := hm.f
|
||||||
|
peerCert := hostinfo.ConnectionState.peerCert
|
||||||
|
hsFields := m{"stage": uint64(1), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)}
|
||||||
|
|
||||||
|
switch err {
|
||||||
|
case ErrAlreadySeen:
|
||||||
|
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
||||||
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
|
}
|
||||||
|
// Resend the original response. The peer is committed to that response's
|
||||||
|
// ephemeral keys; a freshly-built one would have different keys and break
|
||||||
|
// the tunnel even though both sides "completed" the handshake.
|
||||||
|
if msg := existing.HandshakePacket[handshakePacketStage2]; msg != nil {
|
||||||
|
hm.sendHandshakeResponse(via, msg, existing, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
case ErrExistingHostInfo:
|
||||||
|
f.l.Info("Handshake too old",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", peerCert.Certificate.Name(),
|
||||||
|
"certVersion", peerCert.Certificate.Version(),
|
||||||
|
"fingerprint", peerCert.Fingerprint,
|
||||||
|
"issuer", peerCert.Certificate.Issuer(),
|
||||||
|
"oldHandshakeTime", existing.lastHandshakeTime,
|
||||||
|
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
||||||
|
"initiatorIndex", hostinfo.remoteIndexId,
|
||||||
|
"responderIndex", hostinfo.localIndexId,
|
||||||
|
"handshake", hsFields,
|
||||||
|
)
|
||||||
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
|
|
||||||
|
case ErrLocalIndexCollision:
|
||||||
|
f.l.Error("Failed to add HostInfo due to localIndex collision",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", peerCert.Certificate.Name(),
|
||||||
|
"certVersion", peerCert.Certificate.Version(),
|
||||||
|
"fingerprint", peerCert.Fingerprint,
|
||||||
|
"issuer", peerCert.Certificate.Issuer(),
|
||||||
|
"localIndex", hostinfo.localIndexId,
|
||||||
|
"initiatorIndex", hostinfo.remoteIndexId,
|
||||||
|
"responderIndex", hostinfo.localIndexId,
|
||||||
|
"handshake", hsFields,
|
||||||
|
)
|
||||||
|
|
||||||
|
default:
|
||||||
|
f.l.Error("Failed to add HostInfo to HostMap",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"error", err,
|
||||||
|
"certName", peerCert.Certificate.Name(),
|
||||||
|
"certVersion", peerCert.Certificate.Version(),
|
||||||
|
"fingerprint", peerCert.Fingerprint,
|
||||||
|
"issuer", peerCert.Certificate.Issuer(),
|
||||||
|
"initiatorIndex", hostinfo.remoteIndexId,
|
||||||
|
"responderIndex", hostinfo.localIndexId,
|
||||||
|
"handshake", hsFields,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// certVerifier returns a CertVerifier that validates certs against the current CA pool.
|
||||||
|
func (hm *HandshakeManager) certVerifier() handshake.CertVerifier {
|
||||||
|
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
||||||
|
return hm.f.pki.GetCAPool().VerifyCertificate(time.Now(), c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+139
-4
@@ -5,6 +5,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
@@ -27,7 +28,7 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1HandshakeBytes: []byte{},
|
v1Credential: nil,
|
||||||
}
|
}
|
||||||
|
|
||||||
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
@@ -79,15 +80,15 @@ func testCountTimerWheelEntries(tw *LockingTimerWheel[netip.Addr]) (c int) {
|
|||||||
type mockEncWriter struct {
|
type mockEncWriter struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.MessageSubType, _ netip.Addr, _, _, _ []byte) {
|
func (mw *mockEncWriter) SendMessageToVpnAddr(_ header.MessageType, _ header.MessageSubType, _ netip.Addr, _ []byte, _ *WireBuffer) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _, _, _ []byte, _ bool) {
|
func (mw *mockEncWriter) SendVia(_ *HostInfo, _ *Relay, _ []byte, _ *WireBuffer) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
func (mw *mockEncWriter) SendMessageToHostInfo(_ header.MessageType, _ header.MessageSubType, _ *HostInfo, _, _, _ []byte) {
|
func (mw *mockEncWriter) SendMessageToHostInfo(_ header.MessageType, _ header.MessageSubType, _ *HostInfo, _ []byte, _ *WireBuffer) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -100,3 +101,137 @@ func (mw *mockEncWriter) GetHostInfo(_ netip.Addr) *HostInfo {
|
|||||||
func (mw *mockEncWriter) GetCertState() *CertState {
|
func (mw *mockEncWriter) GetCertState() *CertState {
|
||||||
return &CertState{initiatingVersion: cert.Version2}
|
return &CertState{initiatingVersion: cert.Version2}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidatePeerCert(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
myNetwork := netip.MustParsePrefix("10.0.0.1/24")
|
||||||
|
myAddrTable := new(bart.Lite)
|
||||||
|
myAddrTable.Insert(netip.PrefixFrom(myNetwork.Addr(), myNetwork.Addr().BitLen()))
|
||||||
|
myNetTable := new(bart.Lite)
|
||||||
|
myNetTable.Insert(myNetwork.Masked())
|
||||||
|
|
||||||
|
newHM := func() *HandshakeManager {
|
||||||
|
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
|
hm.f = &Interface{
|
||||||
|
handshakeManager: hm,
|
||||||
|
pki: &PKI{},
|
||||||
|
l: l,
|
||||||
|
myVpnAddrsTable: myAddrTable,
|
||||||
|
myVpnNetworksTable: myNetTable,
|
||||||
|
lightHouse: hm.lightHouse,
|
||||||
|
}
|
||||||
|
return hm
|
||||||
|
}
|
||||||
|
|
||||||
|
cached := func(networks ...netip.Prefix) *cert.CachedCertificate {
|
||||||
|
return &cert.CachedCertificate{
|
||||||
|
Certificate: &dummyCert{name: "peer", networks: networks},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
via := ViaSender{
|
||||||
|
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
||||||
|
IsRelayed: true, // skip the remote allow list (covered separately)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("addr inside our networks sets anyVpnAddrsInCommon", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
// 10.0.0.2 falls inside our 10.0.0.0/24
|
||||||
|
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.2/24")))
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.True(t, common)
|
||||||
|
assert.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.2")}, addrs)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("addr outside our networks leaves anyVpnAddrsInCommon false", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("192.168.1.5/24")))
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.False(t, common)
|
||||||
|
assert.Equal(t, []netip.Addr{netip.MustParseAddr("192.168.1.5")}, addrs)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("any matching network is enough", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
addrs, common, ok := hm.validatePeerCert(via, cached(
|
||||||
|
netip.MustParsePrefix("192.168.1.5/24"),
|
||||||
|
netip.MustParsePrefix("10.0.0.42/24"),
|
||||||
|
))
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.True(t, common)
|
||||||
|
assert.Len(t, addrs, 2)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("self-handshake is rejected", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
// 10.0.0.1 is in myVpnAddrsTable
|
||||||
|
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.1/24")))
|
||||||
|
assert.False(t, ok)
|
||||||
|
assert.False(t, common)
|
||||||
|
assert.Nil(t, addrs)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("cert with no networks is rejected", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
addrs, common, ok := hm.validatePeerCert(via, cached())
|
||||||
|
assert.False(t, ok)
|
||||||
|
assert.False(t, common)
|
||||||
|
assert.Nil(t, addrs)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleIncomingDispatch(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
|
||||||
|
newHM := func() *HandshakeManager {
|
||||||
|
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
|
hm.f = &Interface{
|
||||||
|
handshakeManager: hm,
|
||||||
|
pki: &PKI{},
|
||||||
|
l: l,
|
||||||
|
}
|
||||||
|
return hm
|
||||||
|
}
|
||||||
|
|
||||||
|
via := ViaSender{
|
||||||
|
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
||||||
|
IsRelayed: true, // bypass remote allow list
|
||||||
|
}
|
||||||
|
|
||||||
|
// A packet body of zero length is fine for these tests: dispatch is
|
||||||
|
// gated on header fields, and we assert that we never reach noise/cert
|
||||||
|
// processing for any of the malformed shapes here.
|
||||||
|
pkt := make([]byte, header.Len)
|
||||||
|
|
||||||
|
t.Run("unsupported subtype dropped", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
h := &header.H{Type: header.Handshake, Subtype: header.MessageSubType(99), MessageCounter: 1}
|
||||||
|
hm.HandleIncoming(via, pkt, h)
|
||||||
|
assert.Empty(t, hm.indexes, "no pending handshake should be created")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("stage-1 with non-zero RemoteIndex dropped", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
h := &header.H{
|
||||||
|
Type: header.Handshake,
|
||||||
|
Subtype: header.HandshakeIXPSK0,
|
||||||
|
RemoteIndex: 0xdeadbeef,
|
||||||
|
MessageCounter: 1,
|
||||||
|
}
|
||||||
|
hm.HandleIncoming(via, pkt, h)
|
||||||
|
assert.Empty(t, hm.indexes, "spoofed stage-1 must not create a pending machine")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("continuation with no matching pending index dropped", func(t *testing.T) {
|
||||||
|
hm := newHM()
|
||||||
|
h := &header.H{
|
||||||
|
Type: header.Handshake,
|
||||||
|
Subtype: header.HandshakeIXPSK0,
|
||||||
|
RemoteIndex: 0xcafef00d,
|
||||||
|
MessageCounter: 2,
|
||||||
|
}
|
||||||
|
hm.HandleIncoming(via, pkt, h)
|
||||||
|
assert.Empty(t, hm.indexes, "orphan stage-2 must not create state")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
+4
-2
@@ -308,7 +308,7 @@ type cachedPacket struct {
|
|||||||
packet []byte
|
packet []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
type packetCallback func(t header.MessageType, st header.MessageSubType, h *HostInfo, p, nb, out []byte)
|
type packetCallback func(t header.MessageType, st header.MessageSubType, h *HostInfo, p []byte, buf *WireBuffer)
|
||||||
|
|
||||||
type cachedPacketMetrics struct {
|
type cachedPacketMetrics struct {
|
||||||
sent metrics.Counter
|
sent metrics.Counter
|
||||||
@@ -691,6 +691,7 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
buf := ifce.bufAlloc.Acquire()
|
||||||
i.remotes.ForEach(preferredRanges, func(addr netip.AddrPort, preferred bool) {
|
i.remotes.ForEach(preferredRanges, func(addr netip.AddrPort, preferred bool) {
|
||||||
if remote.IsValid() && (!addr.IsValid() || !preferred) {
|
if remote.IsValid() && (!addr.IsValid() || !preferred) {
|
||||||
return
|
return
|
||||||
@@ -698,8 +699,9 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
|
|||||||
|
|
||||||
// Try to send a test packet to that host, this should
|
// Try to send a test packet to that host, this should
|
||||||
// cause it to detect a roaming event and switch remotes
|
// cause it to detect a roaming event and switch remotes
|
||||||
ifce.sendTo(header.Test, header.TestRequest, i.ConnectionState, i, addr, []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
ifce.sendTo(header.Test, header.TestRequest, i.ConnectionState, i, addr, []byte(""), buf)
|
||||||
})
|
})
|
||||||
|
ifce.bufAlloc.Release(buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Re query our lighthouses for new remotes occasionally
|
// Re query our lighthouses for new remotes occasionally
|
||||||
|
|||||||
@@ -8,12 +8,13 @@ import (
|
|||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(buf *WireBuffer, q int, localCache firewall.ConntrackCache) {
|
||||||
err := newPacket(packet, false, fwPacket)
|
packet := buf.IPPacket()
|
||||||
|
|
||||||
|
err := newPacket(packet, false, buf.FwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("Error while validating outbound packet",
|
f.l.Debug("Error while validating outbound packet",
|
||||||
@@ -26,12 +27,12 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
|
|
||||||
// Ignore local broadcast packets
|
// Ignore local broadcast packets
|
||||||
if f.dropLocalBroadcast {
|
if f.dropLocalBroadcast {
|
||||||
if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) {
|
if f.myBroadcastAddrsTable.Contains(buf.FwPacket.RemoteAddr) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if f.myVpnAddrsTable.Contains(fwPacket.RemoteAddr) {
|
if f.myVpnAddrsTable.Contains(buf.FwPacket.RemoteAddr) {
|
||||||
// Immediately forward packets from self to self.
|
// Immediately forward packets from self to self.
|
||||||
// This should only happen on Darwin-based and FreeBSD hosts, which
|
// This should only happen on Darwin-based and FreeBSD hosts, which
|
||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
@@ -48,20 +49,20 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Ignore multicast packets
|
// Ignore multicast packets
|
||||||
if f.dropMulticast && fwPacket.RemoteAddr.IsMulticast() {
|
if f.dropMulticast && buf.FwPacket.RemoteAddr.IsMulticast() {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(buf.FwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, buf.Out, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", buf.FwPacket.RemoteAddr,
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", buf.FwPacket,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
@@ -71,15 +72,15 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*buf.FwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, buf, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, buf.Out, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", buf.FwPacket,
|
||||||
"reason", dropReason,
|
"reason", dropReason,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -102,27 +103,27 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
|
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, scratch []byte, buf *WireBuffer, q int) {
|
||||||
if !f.firewall.OutSendReject {
|
if !f.firewall.OutSendReject {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
out = iputil.CreateRejectPacket(packet, out)
|
rejectIP := iputil.CreateRejectPacket(packet, scratch)
|
||||||
if len(out) == 0 {
|
if len(rejectIP) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(out) > iputil.MaxRejectPacketSize {
|
if len(rejectIP) > iputil.MaxRejectPacketSize {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||||
f.l.Info("rejectOutside: packet too big, not sending",
|
f.l.Info("rejectOutside: packet too big, not sending",
|
||||||
"packet", packet,
|
"packet", packet,
|
||||||
"outPacket", out,
|
"outPacket", rejectIP,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, out, nb, packet, q)
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, rejectIP, buf, q)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
// Handshake will attempt to initiate a tunnel with the provided vpn address. This is a no-op if the tunnel is already established or being established
|
||||||
@@ -215,7 +216,7 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p []byte, buf *WireBuffer) {
|
||||||
fp := &firewall.Packet{}
|
fp := &firewall.Packet{}
|
||||||
err := newPacket(p, false, fp)
|
err := newPacket(p, false, fp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -235,12 +236,12 @@ func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubTyp
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.sendNoMetrics(header.Message, st, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, p, nb, out, 0)
|
f.sendNoMetrics(header.Message, st, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, p, buf, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendMessageToVpnAddr handles real addr:port lookup and sends to the current best known address for vpnAddr.
|
// SendMessageToVpnAddr handles real addr:port lookup and sends to the current best known address for vpnAddr.
|
||||||
// This function ignores myVpnNetworksTable, and will always attempt to treat the address as a vpnAddr
|
// This function ignores myVpnNetworksTable, and will always attempt to treat the address as a vpnAddr
|
||||||
func (f *Interface) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte) {
|
func (f *Interface) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p []byte, buf *WireBuffer) {
|
||||||
hostInfo, ready := f.handshakeManager.GetOrHandshake(vpnAddr, func(hh *HandshakeHostInfo) {
|
hostInfo, ready := f.handshakeManager.GetOrHandshake(vpnAddr, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, t, st, p, f.SendMessageToHostInfo, f.cachedPacketMetrics)
|
hh.cachePacket(f.l, t, st, p, f.SendMessageToHostInfo, f.cachedPacketMetrics)
|
||||||
})
|
})
|
||||||
@@ -258,113 +259,73 @@ func (f *Interface) SendMessageToVpnAddr(t header.MessageType, st header.Message
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.SendMessageToHostInfo(t, st, hostInfo, p, nb, out)
|
f.SendMessageToHostInfo(t, st, hostInfo, p, buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hi *HostInfo, p, nb, out []byte) {
|
func (f *Interface) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hi *HostInfo, p []byte, buf *WireBuffer) {
|
||||||
f.send(t, st, hi.ConnectionState, hi, p, nb, out)
|
f.send(t, st, hi.ConnectionState, hi, p, buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) send(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, p, nb, out []byte) {
|
func (f *Interface) send(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, p []byte, buf *WireBuffer) {
|
||||||
f.messageMetrics.Tx(t, st, 1)
|
f.messageMetrics.Tx(t, st, 1)
|
||||||
f.sendNoMetrics(t, st, ci, hostinfo, netip.AddrPort{}, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, netip.AddrPort{}, p, buf, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte) {
|
func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p []byte, buf *WireBuffer) {
|
||||||
f.messageMetrics.Tx(t, st, 1)
|
f.messageMetrics.Tx(t, st, 1)
|
||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, buf, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
// handshake messages to peers through relay tunnels.
|
// handshake messages to peers through relay tunnels.
|
||||||
// via is the HostInfo through which the message is relayed.
|
//
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
// via is the HostInfo through which the message is relayed. ad is staged into
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
// the inner-payload slot of buf and then AAD-only sealed under via's key by
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
// SealRelayInPlace. The sendNoMetrics relay-forward path skips this entry
|
||||||
// q indicates which writer to use to send the packet.
|
// point and calls sendViaInPlace directly because its inner ciphertext is
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
// already in place from the encrypt step.
|
||||||
relay *Relay,
|
func (f *Interface) SendVia(via *HostInfo, relay *Relay, ad []byte, buf *WireBuffer) {
|
||||||
ad,
|
if header.Len+len(ad)+via.ConnectionState.eKey.Overhead() > cap(buf.Out) {
|
||||||
nb,
|
|
||||||
out []byte,
|
|
||||||
nocopy bool,
|
|
||||||
) {
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
|
||||||
via.ConnectionState.writeLock.Lock()
|
|
||||||
}
|
|
||||||
c := via.ConnectionState.messageCounter.Add(1)
|
|
||||||
|
|
||||||
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
|
|
||||||
f.connectionManager.Out(via)
|
|
||||||
|
|
||||||
// Authenticate the header and payload, but do not encrypt for this message type.
|
|
||||||
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
|
|
||||||
if len(out)+len(ad)+via.ConnectionState.eKey.Overhead() > cap(out) {
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
via.ConnectionState.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
via.logger(f.l).Error("SendVia out buffer not large enough for relay",
|
via.logger(f.l).Error("SendVia out buffer not large enough for relay",
|
||||||
"outCap", cap(out),
|
"outCap", cap(buf.Out),
|
||||||
"payloadLen", len(ad),
|
"payloadLen", len(ad),
|
||||||
"headerLen", len(out),
|
"headerLen", header.Len,
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
buf.StageRelayInner(ad)
|
||||||
|
f.sendViaInPlace(via, relay, len(ad), buf)
|
||||||
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// sendViaInPlace stamps the outer relay header, AAD-seals over the [outer
|
||||||
offset := len(out)
|
// header | inner-already-staged] region, and writes the result to via.remote.
|
||||||
out = out[:offset+len(ad)]
|
// Called from SendVia (after staging ad) and from sendNoMetrics' relay-forward
|
||||||
|
// path (where the inner ciphertext is already in place from SealForRelay).
|
||||||
// In one call path, the associated data _is_ already stored in out. In other call paths, the associated data must
|
func (f *Interface) sendViaInPlace(via *HostInfo, relay *Relay, innerLen int, buf *WireBuffer) {
|
||||||
// be copied into 'out'.
|
f.connectionManager.Out(via)
|
||||||
if !nocopy {
|
out, err := buf.SealRelayInPlace(via.ConnectionState, relay.RemoteIndex, innerLen)
|
||||||
copy(out[offset:], ad)
|
|
||||||
}
|
|
||||||
|
|
||||||
var err error
|
|
||||||
out, err = via.ConnectionState.eKey.EncryptDanger(out, out, nil, c, nb)
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
via.ConnectionState.writeLock.Unlock()
|
|
||||||
}
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.remote)
|
if err := f.writers[0].WriteTo(out, via.remote); err != nil {
|
||||||
if err != nil {
|
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
// sendNoMetrics encrypts and writes one outbound nebula packet (data, control,
|
||||||
|
// lighthouse, etc) using buf as the per-call wire scratch. When the hostinfo
|
||||||
|
// has no direct remote we encrypt into the relay-reserved slot via
|
||||||
|
// SealForRelay so sendViaInPlace can wrap it without an extra copy.
|
||||||
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p []byte, buf *WireBuffer, q int) {
|
||||||
if ci.eKey == nil {
|
if ci.eKey == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
useRelay := !remote.IsValid() && !hostinfo.remote.IsValid()
|
useRelay := !remote.IsValid() && !hostinfo.remote.IsValid()
|
||||||
fullOut := out
|
|
||||||
|
|
||||||
if useRelay {
|
|
||||||
if len(out) < header.Len {
|
|
||||||
// out always has a capacity of mtu, but not always a length greater than the header.Len.
|
|
||||||
// Grow it to make sure the next operation works.
|
|
||||||
out = out[:header.Len]
|
|
||||||
}
|
|
||||||
// Save a header's worth of data at the front of the 'out' buffer.
|
|
||||||
out = out[header.Len:]
|
|
||||||
}
|
|
||||||
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
|
||||||
ci.writeLock.Lock()
|
|
||||||
}
|
|
||||||
c := ci.messageCounter.Add(1)
|
|
||||||
|
|
||||||
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
|
|
||||||
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
|
|
||||||
f.connectionManager.Out(hostinfo)
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
|
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
|
||||||
@@ -381,50 +342,42 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var out []byte
|
||||||
var err error
|
var err error
|
||||||
out, err = ci.eKey.EncryptDanger(out, out, p, c, nb)
|
if useRelay {
|
||||||
if noiseutil.EncryptLockNeeded {
|
out, err = buf.SealForRelay(ci, t, st, hostinfo.remoteIndexId, p)
|
||||||
ci.writeLock.Unlock()
|
} else {
|
||||||
|
out, err = buf.Seal(ci, t, st, hostinfo.remoteIndexId, p)
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"udpAddr", remote,
|
"udpAddr", remote,
|
||||||
"counter", c,
|
|
||||||
"attemptedCounter", c,
|
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if remote.IsValid() {
|
switch {
|
||||||
err = f.writers[q].WriteTo(out, remote)
|
case remote.IsValid():
|
||||||
if err != nil {
|
if err := f.writers[q].WriteTo(out, remote); err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
hostinfo.logger(f.l).Error("Failed to write outgoing packet", "error", err, "udpAddr", remote)
|
||||||
"error", err,
|
|
||||||
"udpAddr", remote,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
} else if hostinfo.remote.IsValid() {
|
case hostinfo.remote.IsValid():
|
||||||
err = f.writers[q].WriteTo(out, hostinfo.remote)
|
if err := f.writers[q].WriteTo(out, hostinfo.remote); err != nil {
|
||||||
if err != nil {
|
hostinfo.logger(f.l).Error("Failed to write outgoing packet", "error", err, "udpAddr", hostinfo.remote)
|
||||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
|
||||||
"error", err,
|
|
||||||
"udpAddr", remote,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
} else {
|
default:
|
||||||
// Try to send via a relay
|
// SealForRelay placed the inner ciphertext at buf.Out[header.Len:],
|
||||||
|
// so sendViaInPlace can wrap it with the outer relay header without
|
||||||
|
// an extra copy.
|
||||||
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||||
relayHostInfo, relay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
relayHostInfo, relay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.relayState.DeleteRelay(relayIP)
|
hostinfo.relayState.DeleteRelay(relayIP)
|
||||||
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo", "relay", relayIP, "error", err)
|
||||||
"relay", relayIP,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
f.sendViaInPlace(relayHostInfo, relay, len(out), buf)
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+18
-21
@@ -101,19 +101,19 @@ type Interface struct {
|
|||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
cachedPacketMetrics *cachedPacketMetrics
|
cachedPacketMetrics *cachedPacketMetrics
|
||||||
|
|
||||||
|
// bufAlloc hands out reusable WireBuffers sized for this interface's
|
||||||
|
// inside Device. All buf consumers (hot-path data-plane goroutines,
|
||||||
|
// long-lived workers, and cold callers) acquire from here so sizing
|
||||||
|
// is centralized and consistent. Long-lived owners just don't release.
|
||||||
|
bufAlloc WireBufferAllocator
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
type EncWriter interface {
|
type EncWriter interface {
|
||||||
SendVia(via *HostInfo,
|
SendVia(via *HostInfo, relay *Relay, ad []byte, buf *WireBuffer)
|
||||||
relay *Relay,
|
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p []byte, buf *WireBuffer)
|
||||||
ad,
|
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p []byte, buf *WireBuffer)
|
||||||
nb,
|
|
||||||
out []byte,
|
|
||||||
nocopy bool,
|
|
||||||
)
|
|
||||||
SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnAddr netip.Addr, p, nb, out []byte)
|
|
||||||
SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte)
|
|
||||||
Handshake(vpnAddr netip.Addr)
|
Handshake(vpnAddr netip.Addr)
|
||||||
GetHostInfo(vpnAddr netip.Addr) *HostInfo
|
GetHostInfo(vpnAddr netip.Addr) *HostInfo
|
||||||
GetCertState() *CertState
|
GetCertState() *CertState
|
||||||
@@ -204,6 +204,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
dropped: metrics.GetOrRegisterCounter("hostinfo.cached_packets.dropped", nil),
|
dropped: metrics.GetOrRegisterCounter("hostinfo.cached_packets.dropped", nil),
|
||||||
},
|
},
|
||||||
|
|
||||||
|
bufAlloc: NewWireBufferPool(mtu, c.Inside.TunPrefixLen()),
|
||||||
|
|
||||||
l: c.l,
|
l: c.l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -311,13 +313,11 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
// Long-lived per-receive-goroutine buf; never released back to the pool.
|
||||||
h := &header.H{}
|
buf := f.bufAlloc.Acquire()
|
||||||
fwPacket := &firewall.Packet{}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, buf, payload, lhh, i, ctCache.Get())
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
if err != nil && !f.closed.Load() {
|
||||||
@@ -329,15 +329,12 @@ func (f *Interface) listenOut(i int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||||
packet := make([]byte, mtu)
|
// Long-lived per-tun-reader buf; never released back to the pool.
|
||||||
out := make([]byte, mtu)
|
buf := f.bufAlloc.Acquire()
|
||||||
fwPacket := &firewall.Packet{}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
_, err := buf.ReadIPFromTUN(reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !f.closed.Load() {
|
if !f.closed.Load() {
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||||
@@ -346,7 +343,7 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
f.consumeInsidePacket(buf, i, conntrackCache.Get())
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", i)
|
||||||
|
|||||||
+46
-23
@@ -63,6 +63,10 @@ type LightHouse struct {
|
|||||||
interval atomic.Int64
|
interval atomic.Int64
|
||||||
updateCancel context.CancelFunc
|
updateCancel context.CancelFunc
|
||||||
ifce EncWriter
|
ifce EncWriter
|
||||||
|
// bufAlloc lets the lighthouse query/update workers, request handlers
|
||||||
|
// and punchback goroutines acquire correctly sized WireBuffers from
|
||||||
|
// the same pool as the data plane. Set by main.go alongside ifce.
|
||||||
|
bufAlloc WireBufferAllocator
|
||||||
nebulaPort uint32 // 32 bits because protobuf does not have a uint16
|
nebulaPort uint32 // 32 bits because protobuf does not have a uint16
|
||||||
|
|
||||||
advertiseAddrs atomic.Pointer[[]netip.AddrPort]
|
advertiseAddrs atomic.Pointer[[]netip.AddrPort]
|
||||||
@@ -109,6 +113,10 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
punchy: p,
|
punchy: p,
|
||||||
updateTrigger: make(chan struct{}, 1),
|
updateTrigger: make(chan struct{}, 1),
|
||||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||||
|
// Default to a no-prefix pool so the query/update workers and
|
||||||
|
// request handlers have a working WireBufferAllocator before
|
||||||
|
// main.go wires up the real one from the Interface.
|
||||||
|
bufAlloc: NewWireBufferPool(mtu, 0),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
lighthouses := make([]netip.Addr, 0)
|
lighthouses := make([]netip.Addr, 0)
|
||||||
@@ -758,21 +766,22 @@ func (lh *LightHouse) startQueryWorker() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
nb := make([]byte, 12, 12)
|
// Long-lived per-worker WireBuffer; reused for every lighthouse query
|
||||||
out := make([]byte, mtu)
|
// this worker issues for the life of the goroutine.
|
||||||
|
buf := lh.bufAlloc.Acquire()
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-lh.ctx.Done():
|
case <-lh.ctx.Done():
|
||||||
return
|
return
|
||||||
case addr := <-lh.queryChan:
|
case addr := <-lh.queryChan:
|
||||||
lh.innerQueryServer(addr, nb, out)
|
lh.innerQueryServer(addr, buf)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
func (lh *LightHouse) innerQueryServer(addr netip.Addr, buf *WireBuffer) {
|
||||||
if lh.IsLighthouseAddr(addr) {
|
if lh.IsLighthouseAddr(addr) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -821,7 +830,7 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v1Query, nb, out)
|
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v1Query, buf)
|
||||||
queried++
|
queried++
|
||||||
|
|
||||||
} else if v == cert.Version2 {
|
} else if v == cert.Version2 {
|
||||||
@@ -840,7 +849,7 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v2Query, nb, out)
|
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v2Query, buf)
|
||||||
queried++
|
queried++
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
@@ -869,8 +878,12 @@ func (lh *LightHouse) StartUpdateWorker() {
|
|||||||
go func() {
|
go func() {
|
||||||
defer clockSource.Stop()
|
defer clockSource.Stop()
|
||||||
|
|
||||||
|
// Long-lived per-worker WireBuffer; reused across every periodic
|
||||||
|
// update for the life of this goroutine.
|
||||||
|
buf := lh.bufAlloc.Acquire()
|
||||||
|
|
||||||
for {
|
for {
|
||||||
lh.SendUpdate()
|
lh.sendUpdate(buf)
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case <-updateCtx.Done():
|
case <-updateCtx.Done():
|
||||||
@@ -884,6 +897,15 @@ func (lh *LightHouse) StartUpdateWorker() {
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SendUpdate is the public entry point that triggers a one-shot lighthouse
|
||||||
|
// update outside the worker loop (e.g. tests or reload paths). It allocates
|
||||||
|
// its own WireBuffer since callers don't already own one.
|
||||||
|
func (lh *LightHouse) SendUpdate() {
|
||||||
|
buf := lh.bufAlloc.Acquire()
|
||||||
|
defer lh.bufAlloc.Release(buf)
|
||||||
|
lh.sendUpdate(buf)
|
||||||
|
}
|
||||||
|
|
||||||
// TriggerUpdate requests an immediate lighthouse update. This is a non-blocking
|
// TriggerUpdate requests an immediate lighthouse update. This is a non-blocking
|
||||||
// operation intended to be called after a handshake completes with a lighthouse,
|
// operation intended to be called after a handshake completes with a lighthouse,
|
||||||
// so the lighthouse has our current addresses without waiting for the next
|
// so the lighthouse has our current addresses without waiting for the next
|
||||||
@@ -895,7 +917,7 @@ func (lh *LightHouse) TriggerUpdate() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (lh *LightHouse) SendUpdate() {
|
func (lh *LightHouse) sendUpdate(buf *WireBuffer) {
|
||||||
var v4 []*V4AddrPort
|
var v4 []*V4AddrPort
|
||||||
var v6 []*V6AddrPort
|
var v6 []*V6AddrPort
|
||||||
|
|
||||||
@@ -921,9 +943,6 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
out := make([]byte, mtu)
|
|
||||||
|
|
||||||
var v1Update, v2Update []byte
|
var v1Update, v2Update []byte
|
||||||
var err error
|
var err error
|
||||||
updated := 0
|
updated := 0
|
||||||
@@ -974,7 +993,7 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v1Update, nb, out)
|
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v1Update, buf)
|
||||||
updated++
|
updated++
|
||||||
|
|
||||||
} else if v == cert.Version2 {
|
} else if v == cert.Version2 {
|
||||||
@@ -1003,7 +1022,7 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v2Update, nb, out)
|
lh.ifce.SendMessageToVpnAddr(header.LightHouse, 0, lhVpnAddr, v2Update, buf)
|
||||||
updated++
|
updated++
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
@@ -1020,8 +1039,10 @@ func (lh *LightHouse) SendUpdate() {
|
|||||||
|
|
||||||
type LightHouseHandler struct {
|
type LightHouseHandler struct {
|
||||||
lh *LightHouse
|
lh *LightHouse
|
||||||
nb []byte
|
// buf is the long-lived per-handler wire scratch. NewRequestHandler is
|
||||||
out []byte
|
// called once per data-plane receive goroutine, so buf is owned by that
|
||||||
|
// goroutine and reused for every lighthouse send the handler issues.
|
||||||
|
buf *WireBuffer
|
||||||
pb []byte
|
pb []byte
|
||||||
meta *NebulaMeta
|
meta *NebulaMeta
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
@@ -1030,8 +1051,7 @@ type LightHouseHandler struct {
|
|||||||
func (lh *LightHouse) NewRequestHandler() *LightHouseHandler {
|
func (lh *LightHouse) NewRequestHandler() *LightHouseHandler {
|
||||||
lhh := &LightHouseHandler{
|
lhh := &LightHouseHandler{
|
||||||
lh: lh,
|
lh: lh,
|
||||||
nb: make([]byte, 12, 12),
|
buf: lh.bufAlloc.Acquire(),
|
||||||
out: make([]byte, mtu),
|
|
||||||
l: lh.l,
|
l: lh.l,
|
||||||
pb: make([]byte, mtu),
|
pb: make([]byte, mtu),
|
||||||
|
|
||||||
@@ -1168,7 +1188,7 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
|
|||||||
}
|
}
|
||||||
|
|
||||||
lhh.lh.metricTx(NebulaMeta_HostQueryReply, 1)
|
lhh.lh.metricTx(NebulaMeta_HostQueryReply, 1)
|
||||||
w.SendMessageToVpnAddr(header.LightHouse, 0, fromVpnAddrs[0], lhh.pb[:ln], lhh.nb, lhh.out[:0])
|
w.SendMessageToVpnAddr(header.LightHouse, 0, fromVpnAddrs[0], lhh.pb[:ln], lhh.buf)
|
||||||
|
|
||||||
lhh.sendHostPunchNotification(n, fromVpnAddrs, queryVpnAddr, w)
|
lhh.sendHostPunchNotification(n, fromVpnAddrs, queryVpnAddr, w)
|
||||||
}
|
}
|
||||||
@@ -1228,7 +1248,7 @@ func (lhh *LightHouseHandler) sendHostPunchNotification(n *NebulaMeta, fromVpnAd
|
|||||||
}
|
}
|
||||||
|
|
||||||
lhh.lh.metricTx(NebulaMeta_HostPunchNotification, 1)
|
lhh.lh.metricTx(NebulaMeta_HostPunchNotification, 1)
|
||||||
w.SendMessageToVpnAddr(header.LightHouse, 0, punchNotifDest, lhh.pb[:ln], lhh.nb, lhh.out[:0])
|
w.SendMessageToVpnAddr(header.LightHouse, 0, punchNotifDest, lhh.pb[:ln], lhh.buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (lhh *LightHouseHandler) coalesceAnswers(v cert.Version, c *cache, n *NebulaMeta) {
|
func (lhh *LightHouseHandler) coalesceAnswers(v cert.Version, c *cache, n *NebulaMeta) {
|
||||||
@@ -1385,7 +1405,7 @@ func (lhh *LightHouseHandler) handleHostUpdateNotification(n *NebulaMeta, fromVp
|
|||||||
}
|
}
|
||||||
|
|
||||||
lhh.lh.metricTx(NebulaMeta_HostUpdateNotificationAck, 1)
|
lhh.lh.metricTx(NebulaMeta_HostUpdateNotificationAck, 1)
|
||||||
w.SendMessageToVpnAddr(header.LightHouse, 0, fromVpnAddrs[0], lhh.pb[:ln], lhh.nb, lhh.out[:0])
|
w.SendMessageToVpnAddr(header.LightHouse, 0, fromVpnAddrs[0], lhh.pb[:ln], lhh.buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpnAddrs []netip.Addr, w EncWriter) {
|
func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpnAddrs []netip.Addr, w EncWriter) {
|
||||||
@@ -1452,10 +1472,13 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
|
|||||||
"vpnAddr", detailsVpnAddr,
|
"vpnAddr", detailsVpnAddr,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
//NOTE: we have to allocate a new output buffer here since we are spawning a new goroutine
|
// We acquire and release a fresh buf within this goroutine so it
|
||||||
// for each punchBack packet. We should move this into a timerwheel or a single goroutine
|
// returns to the pool once the punchback send completes. We
|
||||||
|
// should move this into a timerwheel or a single goroutine
|
||||||
// managed by a channel.
|
// managed by a channel.
|
||||||
w.SendMessageToVpnAddr(header.Test, header.TestRequest, detailsVpnAddr, []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
pbuf := lhh.lh.bufAlloc.Acquire()
|
||||||
|
defer lhh.lh.bufAlloc.Release(pbuf)
|
||||||
|
w.SendMessageToVpnAddr(header.Test, header.TestRequest, detailsVpnAddr, []byte(""), pbuf)
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+3
-3
@@ -372,12 +372,12 @@ type testEncWriter struct {
|
|||||||
protocolVersion cert.Version
|
protocolVersion cert.Version
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad, nb, out []byte, nocopy bool) {
|
func (tw *testEncWriter) SendVia(via *HostInfo, relay *Relay, ad []byte, buf *WireBuffer) {
|
||||||
}
|
}
|
||||||
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
func (tw *testEncWriter) Handshake(vpnIp netip.Addr) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tw *testEncWriter) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, _, _ []byte) {
|
func (tw *testEncWriter) SendMessageToHostInfo(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p []byte, _ *WireBuffer) {
|
||||||
msg := &NebulaMeta{}
|
msg := &NebulaMeta{}
|
||||||
err := msg.Unmarshal(p)
|
err := msg.Unmarshal(p)
|
||||||
if tw.metaFilter == nil || msg.Type == *tw.metaFilter {
|
if tw.metaFilter == nil || msg.Type == *tw.metaFilter {
|
||||||
@@ -394,7 +394,7 @@ func (tw *testEncWriter) SendMessageToHostInfo(t header.MessageType, st header.M
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tw *testEncWriter) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnIp netip.Addr, p, _, _ []byte) {
|
func (tw *testEncWriter) SendMessageToVpnAddr(t header.MessageType, st header.MessageSubType, vpnIp netip.Addr, p []byte, _ *WireBuffer) {
|
||||||
msg := &NebulaMeta{}
|
msg := &NebulaMeta{}
|
||||||
err := msg.Unmarshal(p)
|
err := msg.Unmarshal(p)
|
||||||
if tw.metaFilter == nil || msg.Type == *tw.metaFilter {
|
if tw.metaFilter == nil || msg.Type == *tw.metaFilter {
|
||||||
|
|||||||
@@ -184,14 +184,10 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
messageMetrics = newMessageMetricsOnlyRecvError()
|
messageMetrics = newMessageMetricsOnlyRecvError()
|
||||||
}
|
}
|
||||||
|
|
||||||
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
|
||||||
|
|
||||||
handshakeConfig := HandshakeConfig{
|
handshakeConfig := HandshakeConfig{
|
||||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||||
useRelays: useRelays,
|
|
||||||
|
|
||||||
messageMetrics: messageMetrics,
|
messageMetrics: messageMetrics,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -236,6 +232,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
ifce.writers = udpConns
|
ifce.writers = udpConns
|
||||||
lightHouse.ifce = ifce
|
lightHouse.ifce = ifce
|
||||||
|
lightHouse.bufAlloc = ifce.bufAlloc
|
||||||
|
|
||||||
ifce.RegisterConfigChangeCallbacks(c)
|
ifce.RegisterConfigChangeCallbacks(c)
|
||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
|
|||||||
+45
-632
@@ -124,7 +124,7 @@ func (x NebulaControl_MessageType) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{8, 0}
|
return fileDescriptor_2d65afa7693df5ef, []int{6, 0}
|
||||||
}
|
}
|
||||||
|
|
||||||
type NebulaMeta struct {
|
type NebulaMeta struct {
|
||||||
@@ -489,142 +489,6 @@ func (m *NebulaPing) GetTime() uint64 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
type NebulaHandshake struct {
|
|
||||||
Details *NebulaHandshakeDetails `protobuf:"bytes,1,opt,name=Details,proto3" json:"Details,omitempty"`
|
|
||||||
Hmac []byte `protobuf:"bytes,2,opt,name=Hmac,proto3" json:"Hmac,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshake) Reset() { *m = NebulaHandshake{} }
|
|
||||||
func (m *NebulaHandshake) String() string { return proto.CompactTextString(m) }
|
|
||||||
func (*NebulaHandshake) ProtoMessage() {}
|
|
||||||
func (*NebulaHandshake) Descriptor() ([]byte, []int) {
|
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshake) XXX_Unmarshal(b []byte) error {
|
|
||||||
return m.Unmarshal(b)
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshake) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
|
||||||
if deterministic {
|
|
||||||
return xxx_messageInfo_NebulaHandshake.Marshal(b, m, deterministic)
|
|
||||||
} else {
|
|
||||||
b = b[:cap(b)]
|
|
||||||
n, err := m.MarshalToSizedBuffer(b)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return b[:n], nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshake) XXX_Merge(src proto.Message) {
|
|
||||||
xxx_messageInfo_NebulaHandshake.Merge(m, src)
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshake) XXX_Size() int {
|
|
||||||
return m.Size()
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshake) XXX_DiscardUnknown() {
|
|
||||||
xxx_messageInfo_NebulaHandshake.DiscardUnknown(m)
|
|
||||||
}
|
|
||||||
|
|
||||||
var xxx_messageInfo_NebulaHandshake proto.InternalMessageInfo
|
|
||||||
|
|
||||||
func (m *NebulaHandshake) GetDetails() *NebulaHandshakeDetails {
|
|
||||||
if m != nil {
|
|
||||||
return m.Details
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshake) GetHmac() []byte {
|
|
||||||
if m != nil {
|
|
||||||
return m.Hmac
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type NebulaHandshakeDetails struct {
|
|
||||||
Cert []byte `protobuf:"bytes,1,opt,name=Cert,proto3" json:"Cert,omitempty"`
|
|
||||||
InitiatorIndex uint32 `protobuf:"varint,2,opt,name=InitiatorIndex,proto3" json:"InitiatorIndex,omitempty"`
|
|
||||||
ResponderIndex uint32 `protobuf:"varint,3,opt,name=ResponderIndex,proto3" json:"ResponderIndex,omitempty"`
|
|
||||||
Cookie uint64 `protobuf:"varint,4,opt,name=Cookie,proto3" json:"Cookie,omitempty"`
|
|
||||||
Time uint64 `protobuf:"varint,5,opt,name=Time,proto3" json:"Time,omitempty"`
|
|
||||||
CertVersion uint32 `protobuf:"varint,8,opt,name=CertVersion,proto3" json:"CertVersion,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) Reset() { *m = NebulaHandshakeDetails{} }
|
|
||||||
func (m *NebulaHandshakeDetails) String() string { return proto.CompactTextString(m) }
|
|
||||||
func (*NebulaHandshakeDetails) ProtoMessage() {}
|
|
||||||
func (*NebulaHandshakeDetails) Descriptor() ([]byte, []int) {
|
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{7}
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) XXX_Unmarshal(b []byte) error {
|
|
||||||
return m.Unmarshal(b)
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
|
||||||
if deterministic {
|
|
||||||
return xxx_messageInfo_NebulaHandshakeDetails.Marshal(b, m, deterministic)
|
|
||||||
} else {
|
|
||||||
b = b[:cap(b)]
|
|
||||||
n, err := m.MarshalToSizedBuffer(b)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return b[:n], nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) XXX_Merge(src proto.Message) {
|
|
||||||
xxx_messageInfo_NebulaHandshakeDetails.Merge(m, src)
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) XXX_Size() int {
|
|
||||||
return m.Size()
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) XXX_DiscardUnknown() {
|
|
||||||
xxx_messageInfo_NebulaHandshakeDetails.DiscardUnknown(m)
|
|
||||||
}
|
|
||||||
|
|
||||||
var xxx_messageInfo_NebulaHandshakeDetails proto.InternalMessageInfo
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetCert() []byte {
|
|
||||||
if m != nil {
|
|
||||||
return m.Cert
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetInitiatorIndex() uint32 {
|
|
||||||
if m != nil {
|
|
||||||
return m.InitiatorIndex
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetResponderIndex() uint32 {
|
|
||||||
if m != nil {
|
|
||||||
return m.ResponderIndex
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetCookie() uint64 {
|
|
||||||
if m != nil {
|
|
||||||
return m.Cookie
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetTime() uint64 {
|
|
||||||
if m != nil {
|
|
||||||
return m.Time
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) GetCertVersion() uint32 {
|
|
||||||
if m != nil {
|
|
||||||
return m.CertVersion
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
type NebulaControl struct {
|
type NebulaControl struct {
|
||||||
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
||||||
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
||||||
@@ -639,7 +503,7 @@ func (m *NebulaControl) Reset() { *m = NebulaControl{} }
|
|||||||
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
||||||
func (*NebulaControl) ProtoMessage() {}
|
func (*NebulaControl) ProtoMessage() {}
|
||||||
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{8}
|
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
||||||
}
|
}
|
||||||
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
||||||
return m.Unmarshal(b)
|
return m.Unmarshal(b)
|
||||||
@@ -729,65 +593,55 @@ func init() {
|
|||||||
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
||||||
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
||||||
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
||||||
proto.RegisterType((*NebulaHandshake)(nil), "nebula.NebulaHandshake")
|
|
||||||
proto.RegisterType((*NebulaHandshakeDetails)(nil), "nebula.NebulaHandshakeDetails")
|
|
||||||
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
||||||
|
|
||||||
var fileDescriptor_2d65afa7693df5ef = []byte{
|
var fileDescriptor_2d65afa7693df5ef = []byte{
|
||||||
// 785 bytes of a gzipped FileDescriptorProto
|
// 665 bytes of a gzipped FileDescriptorProto
|
||||||
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x55, 0xcd, 0x6e, 0xeb, 0x44,
|
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x54, 0xcd, 0x6e, 0xd3, 0x5c,
|
||||||
0x14, 0x8e, 0x1d, 0x27, 0x4e, 0x4f, 0x7e, 0xae, 0x39, 0x15, 0xc1, 0x41, 0x22, 0x0a, 0x5e, 0x54,
|
0x10, 0x8d, 0x1d, 0x27, 0x69, 0x27, 0x4d, 0x3e, 0x7f, 0x53, 0x51, 0x12, 0x24, 0xac, 0xe0, 0x45,
|
||||||
0x57, 0x2c, 0x72, 0x51, 0x5a, 0xae, 0x58, 0x72, 0x1b, 0x84, 0xd2, 0xaa, 0x3f, 0x61, 0x54, 0x8a,
|
0x55, 0xb1, 0x48, 0x51, 0x5a, 0xba, 0xa6, 0x2d, 0x42, 0xa9, 0xd4, 0x9f, 0x70, 0x55, 0x8a, 0xc4,
|
||||||
0xc4, 0x06, 0xb9, 0xf6, 0xd0, 0x58, 0x71, 0x3c, 0xa9, 0x3d, 0x41, 0xcd, 0x5b, 0xf0, 0x30, 0x3c,
|
0xce, 0xb5, 0x2f, 0x8d, 0x55, 0xc7, 0x37, 0xb5, 0x6f, 0x50, 0xf3, 0x16, 0x3c, 0x0c, 0x0f, 0x01,
|
||||||
0x04, 0xec, 0xba, 0x42, 0x2c, 0x51, 0xbb, 0x64, 0xc9, 0x0b, 0xa0, 0x19, 0xff, 0x27, 0x86, 0xbb,
|
0xbb, 0x2e, 0x59, 0xa2, 0x66, 0xc9, 0x92, 0x17, 0x40, 0xf7, 0xfa, 0xbf, 0x31, 0xb0, 0xbb, 0x33,
|
||||||
0x9b, 0x73, 0xbe, 0xef, 0x3b, 0x73, 0xe6, 0xf3, 0x9c, 0x31, 0x74, 0x02, 0x7a, 0xb7, 0xf1, 0xed,
|
0xe7, 0x9c, 0x99, 0xc9, 0xc9, 0x8c, 0x61, 0xcd, 0xa7, 0x97, 0x33, 0xcf, 0xea, 0x4f, 0x03, 0xc6,
|
||||||
0xf1, 0x3a, 0x64, 0x9c, 0x61, 0x33, 0x8e, 0xac, 0xbf, 0x55, 0x80, 0x2b, 0xb9, 0xbc, 0xa4, 0xdc,
|
0x19, 0xd6, 0xa3, 0xc8, 0xfc, 0xa9, 0x02, 0x9c, 0xca, 0xe7, 0x09, 0xe5, 0x16, 0x0e, 0x40, 0x3b,
|
||||||
0xc6, 0x09, 0x68, 0x37, 0xdb, 0x35, 0x35, 0x95, 0x91, 0xf2, 0xba, 0x37, 0x19, 0x8e, 0x13, 0x4d,
|
0x9f, 0x4f, 0x69, 0x47, 0xe9, 0x29, 0x5b, 0xed, 0x81, 0xd1, 0x8f, 0x35, 0x19, 0xa3, 0x7f, 0x42,
|
||||||
0xce, 0x18, 0x5f, 0xd2, 0x28, 0xb2, 0xef, 0xa9, 0x60, 0x11, 0xc9, 0xc5, 0x63, 0xd0, 0xbf, 0xa6,
|
0xc3, 0xd0, 0xba, 0xa2, 0x82, 0x45, 0x24, 0x17, 0x77, 0xa0, 0xf1, 0x9a, 0x72, 0xcb, 0xf5, 0xc2,
|
||||||
0xdc, 0xf6, 0xfc, 0xc8, 0x54, 0x47, 0xca, 0xeb, 0xf6, 0x64, 0xb0, 0x2f, 0x4b, 0x08, 0x24, 0x65,
|
0x8e, 0xda, 0x53, 0xb6, 0x9a, 0x83, 0xee, 0xb2, 0x2c, 0x26, 0x90, 0x84, 0x69, 0xfe, 0x52, 0xa0,
|
||||||
0x5a, 0xff, 0x28, 0xd0, 0x2e, 0x94, 0xc2, 0x16, 0x68, 0x57, 0x2c, 0xa0, 0x46, 0x0d, 0xbb, 0x70,
|
0x99, 0x2b, 0x85, 0x2b, 0xa0, 0x9d, 0x32, 0x9f, 0xea, 0x15, 0x6c, 0xc1, 0xea, 0x90, 0x85, 0xfc,
|
||||||
0x30, 0x63, 0x11, 0xff, 0x76, 0x43, 0xc3, 0xad, 0xa1, 0x20, 0x42, 0x2f, 0x0b, 0x09, 0x5d, 0xfb,
|
0xed, 0x8c, 0x06, 0x73, 0x5d, 0x41, 0x84, 0x76, 0x1a, 0x12, 0x3a, 0xf5, 0xe6, 0xba, 0x8a, 0x4f,
|
||||||
0x5b, 0x43, 0xc5, 0x8f, 0xa1, 0x2f, 0x72, 0xdf, 0xad, 0x5d, 0x9b, 0xd3, 0x2b, 0xc6, 0xbd, 0x9f,
|
0x60, 0x43, 0xe4, 0xde, 0x4d, 0x1d, 0x8b, 0xd3, 0x53, 0xc6, 0xdd, 0x8f, 0xae, 0x6d, 0x71, 0x97,
|
||||||
0x3c, 0xc7, 0xe6, 0x1e, 0x0b, 0x8c, 0x3a, 0x0e, 0xe0, 0x43, 0x81, 0x5d, 0xb2, 0x9f, 0xa9, 0x5b,
|
0xf9, 0x7a, 0x15, 0xbb, 0xf0, 0x48, 0x60, 0x27, 0xec, 0x13, 0x75, 0x0a, 0x90, 0x96, 0x40, 0xa3,
|
||||||
0x82, 0xb4, 0x14, 0x9a, 0x6f, 0x02, 0x67, 0x51, 0x82, 0x1a, 0xd8, 0x03, 0x10, 0xd0, 0xf7, 0x0b,
|
0x99, 0x6f, 0x8f, 0x0b, 0x50, 0x0d, 0xdb, 0x00, 0x02, 0x7a, 0x3f, 0x66, 0xd6, 0xc4, 0xd5, 0xeb,
|
||||||
0x66, 0xaf, 0x3c, 0xa3, 0x89, 0x87, 0xf0, 0x2a, 0x8f, 0xe3, 0x6d, 0x75, 0xd1, 0xd9, 0xdc, 0xe6,
|
0xb8, 0x0e, 0xff, 0x65, 0x71, 0xd4, 0xb6, 0x21, 0x26, 0x1b, 0x59, 0x7c, 0x7c, 0x38, 0xa6, 0xf6,
|
||||||
0x8b, 0xe9, 0x82, 0x3a, 0x4b, 0xa3, 0x25, 0x3a, 0xcb, 0xc2, 0x98, 0x72, 0x80, 0x9f, 0xc0, 0xa0,
|
0xb5, 0xbe, 0x22, 0x26, 0x4b, 0xc3, 0x88, 0xb2, 0x8a, 0x4f, 0xa1, 0x5b, 0x3e, 0xd9, 0xbe, 0x7d,
|
||||||
0xba, 0xb3, 0x77, 0xce, 0xd2, 0x00, 0xeb, 0x77, 0x15, 0x3e, 0xd8, 0x33, 0x05, 0x2d, 0x80, 0x6b,
|
0xad, 0x83, 0xf9, 0x4d, 0x85, 0xff, 0x97, 0x4c, 0x41, 0x13, 0xe0, 0xcc, 0x73, 0x2e, 0xa6, 0xfe,
|
||||||
0xdf, 0xbd, 0x5d, 0x07, 0xef, 0x5c, 0x37, 0x94, 0xd6, 0x77, 0x4f, 0x55, 0x53, 0x21, 0x85, 0x2c,
|
0xbe, 0xe3, 0x04, 0xd2, 0xfa, 0xd6, 0x81, 0xda, 0x51, 0x48, 0x2e, 0x8b, 0x9b, 0xd0, 0x48, 0x08,
|
||||||
0x1e, 0x81, 0x9e, 0x12, 0x9a, 0xd2, 0xe4, 0x4e, 0x6a, 0xb2, 0xc8, 0x91, 0x14, 0xc4, 0x31, 0x18,
|
0x75, 0x69, 0xf2, 0x5a, 0x62, 0xb2, 0xc8, 0x91, 0x04, 0xc4, 0x3e, 0xe8, 0x67, 0x9e, 0x43, 0xa8,
|
||||||
0xd7, 0xbe, 0x4b, 0xa8, 0x6f, 0x6f, 0x93, 0x54, 0x64, 0x36, 0x46, 0xf5, 0xa4, 0xe2, 0x1e, 0x86,
|
0x67, 0xcd, 0xe3, 0x54, 0xd8, 0xa9, 0xf5, 0xaa, 0x71, 0xc5, 0x25, 0x0c, 0x07, 0xd0, 0x2a, 0x92,
|
||||||
0x13, 0xe8, 0x96, 0xc9, 0xfa, 0xa8, 0xbe, 0x57, 0xbd, 0x4c, 0xc1, 0x13, 0x68, 0xdf, 0x9e, 0x88,
|
0x1b, 0xbd, 0xea, 0x52, 0xf5, 0x22, 0x05, 0x77, 0xa1, 0x79, 0xb1, 0x2b, 0x9e, 0x23, 0x16, 0x70,
|
||||||
0xe5, 0x9c, 0x85, 0x5c, 0x7c, 0x74, 0xa1, 0xc0, 0x54, 0x91, 0x43, 0xa4, 0x48, 0x93, 0xaa, 0xb7,
|
0xf1, 0xa7, 0x0b, 0x05, 0x26, 0x8a, 0x0c, 0x22, 0x79, 0x9a, 0x54, 0xed, 0x65, 0x2a, 0xed, 0x81,
|
||||||
0xb9, 0x4a, 0xdb, 0x51, 0xbd, 0x2d, 0xa8, 0x72, 0x1a, 0x9a, 0xa0, 0x3b, 0x6c, 0x13, 0x70, 0x1a,
|
0x6a, 0x2f, 0xa7, 0xca, 0x68, 0xd8, 0x81, 0x86, 0xcd, 0x66, 0x3e, 0xa7, 0x41, 0xa7, 0x2a, 0x8c,
|
||||||
0x9a, 0x75, 0x61, 0x0c, 0x49, 0x43, 0xeb, 0x08, 0x34, 0x79, 0xe2, 0x1e, 0xa8, 0x33, 0x4f, 0xba,
|
0x21, 0x49, 0x68, 0x6e, 0x82, 0x26, 0x7f, 0x71, 0x1b, 0xd4, 0xa1, 0x2b, 0x5d, 0xd3, 0x88, 0x3a,
|
||||||
0xa6, 0x11, 0x75, 0xe6, 0x89, 0xf8, 0x82, 0xc9, 0x9b, 0xa8, 0x11, 0xf5, 0x82, 0x59, 0x27, 0x00,
|
0x74, 0x45, 0x7c, 0xcc, 0xe4, 0x26, 0x6a, 0x44, 0x3d, 0x66, 0xe6, 0x2e, 0x40, 0x36, 0x06, 0x62,
|
||||||
0x79, 0x1b, 0x88, 0xb1, 0x2a, 0x76, 0x99, 0xc4, 0x15, 0x10, 0x34, 0x81, 0x49, 0x4d, 0x97, 0xc8,
|
0xa4, 0x8a, 0x5c, 0x26, 0x51, 0x05, 0x04, 0x4d, 0x60, 0x52, 0xd3, 0x22, 0xf2, 0x6d, 0xbe, 0x02,
|
||||||
0xb5, 0xf5, 0x15, 0x40, 0xde, 0xc6, 0xfb, 0xf6, 0xc8, 0x2a, 0xd4, 0x0b, 0x15, 0x1e, 0xd3, 0xc1,
|
0xc8, 0xc6, 0xf8, 0x57, 0x8f, 0xb4, 0x42, 0x35, 0x57, 0xe1, 0x36, 0x39, 0xac, 0x91, 0xeb, 0x5f,
|
||||||
0x9a, 0x7b, 0xc1, 0xfd, 0xff, 0x0f, 0x96, 0x60, 0x54, 0x0c, 0x16, 0x82, 0x76, 0xe3, 0xad, 0x68,
|
0xfd, 0xfd, 0xb0, 0x04, 0xa3, 0xe4, 0xb0, 0x10, 0xb4, 0x73, 0x77, 0x42, 0xe3, 0x3e, 0xf2, 0x6d,
|
||||||
0xb2, 0x8f, 0x5c, 0x5b, 0xd6, 0xde, 0xd8, 0x08, 0xb1, 0x51, 0xc3, 0x03, 0x68, 0xc4, 0x97, 0x50,
|
0x9a, 0x4b, 0x67, 0x23, 0xc4, 0x7a, 0x05, 0x57, 0xa1, 0x16, 0x2d, 0xa1, 0x62, 0x7e, 0xa9, 0x42,
|
||||||
0xb1, 0x7e, 0x84, 0x57, 0x71, 0xdd, 0x99, 0x1d, 0xb8, 0xd1, 0xc2, 0x5e, 0x52, 0xfc, 0x32, 0x9f,
|
0x2b, 0x2a, 0x7c, 0xc8, 0x7c, 0x1e, 0x30, 0x0f, 0x5f, 0x16, 0xba, 0x3f, 0x2b, 0x76, 0x8f, 0x49,
|
||||||
0x51, 0x45, 0x5e, 0x9f, 0x9d, 0x0e, 0x32, 0xe6, 0xee, 0xa0, 0x8a, 0x26, 0x66, 0x2b, 0xdb, 0x91,
|
0x25, 0x03, 0xbc, 0x80, 0xf5, 0x23, 0xdf, 0xe5, 0xae, 0xc5, 0x59, 0x20, 0x57, 0xe0, 0xc8, 0x77,
|
||||||
0x4d, 0x74, 0x88, 0x5c, 0x5b, 0x7f, 0x28, 0xd0, 0xaf, 0xd6, 0x09, 0xfa, 0x94, 0x86, 0x5c, 0xee,
|
0xe8, 0x6d, 0xec, 0x53, 0x19, 0x24, 0x14, 0x84, 0x86, 0x53, 0xe6, 0x3b, 0x34, 0xaf, 0x88, 0x7c,
|
||||||
0xd2, 0x21, 0x72, 0x8d, 0x47, 0xd0, 0x3b, 0x0b, 0x3c, 0xee, 0xd9, 0x9c, 0x85, 0x67, 0x81, 0x4b,
|
0x29, 0x83, 0xf0, 0x39, 0xb4, 0x93, 0xa5, 0x3c, 0x67, 0xf2, 0xaf, 0xd1, 0xd2, 0x03, 0x78, 0x80,
|
||||||
0x1f, 0x13, 0xa7, 0x77, 0xb2, 0x82, 0x47, 0x68, 0xb4, 0x66, 0x81, 0x4b, 0x13, 0x5e, 0xec, 0xe7,
|
0xe4, 0x97, 0xfb, 0x4d, 0xc0, 0x26, 0x92, 0x5d, 0x4b, 0xd9, 0x4b, 0x18, 0xf6, 0xa1, 0x99, 0x2f,
|
||||||
0x4e, 0x16, 0xfb, 0xd0, 0x9c, 0x32, 0xb6, 0xf4, 0xa8, 0xa9, 0x49, 0x67, 0x92, 0x28, 0xf3, 0xab,
|
0x5c, 0x76, 0x38, 0x79, 0x42, 0x7a, 0x0c, 0x69, 0xf1, 0x46, 0x89, 0xa2, 0x48, 0x31, 0x87, 0x7f,
|
||||||
0x91, 0xfb, 0x85, 0x23, 0x68, 0x8b, 0x1e, 0x6e, 0x69, 0x18, 0x79, 0x2c, 0x30, 0x5b, 0xb2, 0x60,
|
0xfa, 0x8e, 0x6d, 0x00, 0x1e, 0x06, 0xd4, 0xe2, 0x54, 0xf2, 0x09, 0xbd, 0x99, 0xd1, 0x90, 0xeb,
|
||||||
0x31, 0x75, 0xae, 0xb5, 0x9a, 0x86, 0x7e, 0xae, 0xb5, 0x74, 0xa3, 0x65, 0xfd, 0x5a, 0x87, 0x6e,
|
0x0a, 0x3e, 0x86, 0xf5, 0x42, 0x5e, 0x58, 0x12, 0x52, 0x5d, 0x3d, 0xd8, 0xf9, 0x7a, 0x6f, 0x28,
|
||||||
0x7c, 0xb0, 0x29, 0x0b, 0x78, 0xc8, 0x7c, 0xfc, 0xa2, 0xf4, 0xdd, 0x3e, 0x2d, 0xbb, 0x96, 0x90,
|
0x77, 0xf7, 0x86, 0xf2, 0xe3, 0xde, 0x50, 0x3e, 0x2f, 0x8c, 0xca, 0xdd, 0xc2, 0xa8, 0x7c, 0x5f,
|
||||||
0x2a, 0x3e, 0xdd, 0xe7, 0x70, 0x98, 0x1d, 0x4e, 0x0e, 0x4f, 0xf1, 0xdc, 0x55, 0x90, 0x50, 0x64,
|
0x18, 0x95, 0x0f, 0xdd, 0x2b, 0x97, 0x8f, 0x67, 0x97, 0x7d, 0x9b, 0x4d, 0xb6, 0x43, 0xcf, 0xb2,
|
||||||
0xc7, 0x2c, 0x28, 0x62, 0x07, 0xaa, 0x20, 0xfc, 0x0c, 0x7a, 0xe9, 0x38, 0xdf, 0x30, 0x79, 0xa9,
|
0xaf, 0xc7, 0x37, 0xdb, 0xd1, 0x48, 0x97, 0x75, 0xf9, 0x39, 0xdf, 0xf9, 0x1d, 0x00, 0x00, 0xff,
|
||||||
0xb5, 0xec, 0xe9, 0xd8, 0x41, 0x8a, 0xcf, 0xc2, 0x37, 0x21, 0x5b, 0x49, 0x76, 0x23, 0x63, 0xef,
|
0xff, 0x51, 0x0a, 0xe3, 0xd7, 0xde, 0x05, 0x00, 0x00,
|
||||||
0x61, 0x38, 0x86, 0x76, 0xb1, 0x70, 0xd5, 0x93, 0x53, 0x24, 0x64, 0xcf, 0x48, 0x56, 0x5c, 0xaf,
|
|
||||||
0x50, 0x94, 0x29, 0xd6, 0xec, 0xbf, 0xfe, 0x00, 0x7d, 0xc0, 0x69, 0x48, 0x6d, 0x4e, 0x25, 0x9f,
|
|
||||||
0xd0, 0x87, 0x0d, 0x8d, 0xb8, 0xa1, 0xe0, 0x47, 0x70, 0x58, 0xca, 0x0b, 0x4b, 0x22, 0x6a, 0xa8,
|
|
||||||
0xa7, 0xc7, 0xbf, 0x3d, 0x0f, 0x95, 0xa7, 0xe7, 0xa1, 0xf2, 0xd7, 0xf3, 0x50, 0xf9, 0xe5, 0x65,
|
|
||||||
0x58, 0x7b, 0x7a, 0x19, 0xd6, 0xfe, 0x7c, 0x19, 0xd6, 0x7e, 0x18, 0xdc, 0x7b, 0x7c, 0xb1, 0xb9,
|
|
||||||
0x1b, 0x3b, 0x6c, 0xf5, 0x26, 0xf2, 0x6d, 0x67, 0xb9, 0x78, 0x78, 0x13, 0xb7, 0x74, 0xd7, 0x94,
|
|
||||||
0x3f, 0xc2, 0xe3, 0x7f, 0x03, 0x00, 0x00, 0xff, 0xff, 0xea, 0x6f, 0xbc, 0x50, 0x18, 0x07, 0x00,
|
|
||||||
0x00,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
||||||
@@ -1072,103 +926,6 @@ func (m *NebulaPing) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
|||||||
return len(dAtA) - i, nil
|
return len(dAtA) - i, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *NebulaHandshake) Marshal() (dAtA []byte, err error) {
|
|
||||||
size := m.Size()
|
|
||||||
dAtA = make([]byte, size)
|
|
||||||
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return dAtA[:n], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshake) MarshalTo(dAtA []byte) (int, error) {
|
|
||||||
size := m.Size()
|
|
||||||
return m.MarshalToSizedBuffer(dAtA[:size])
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshake) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
|
||||||
i := len(dAtA)
|
|
||||||
_ = i
|
|
||||||
var l int
|
|
||||||
_ = l
|
|
||||||
if len(m.Hmac) > 0 {
|
|
||||||
i -= len(m.Hmac)
|
|
||||||
copy(dAtA[i:], m.Hmac)
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(len(m.Hmac)))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x12
|
|
||||||
}
|
|
||||||
if m.Details != nil {
|
|
||||||
{
|
|
||||||
size, err := m.Details.MarshalToSizedBuffer(dAtA[:i])
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
i -= size
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(size))
|
|
||||||
}
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0xa
|
|
||||||
}
|
|
||||||
return len(dAtA) - i, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) Marshal() (dAtA []byte, err error) {
|
|
||||||
size := m.Size()
|
|
||||||
dAtA = make([]byte, size)
|
|
||||||
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return dAtA[:n], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) MarshalTo(dAtA []byte) (int, error) {
|
|
||||||
size := m.Size()
|
|
||||||
return m.MarshalToSizedBuffer(dAtA[:size])
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
|
||||||
i := len(dAtA)
|
|
||||||
_ = i
|
|
||||||
var l int
|
|
||||||
_ = l
|
|
||||||
if m.CertVersion != 0 {
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(m.CertVersion))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x40
|
|
||||||
}
|
|
||||||
if m.Time != 0 {
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(m.Time))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x28
|
|
||||||
}
|
|
||||||
if m.Cookie != 0 {
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(m.Cookie))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x20
|
|
||||||
}
|
|
||||||
if m.ResponderIndex != 0 {
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(m.ResponderIndex))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x18
|
|
||||||
}
|
|
||||||
if m.InitiatorIndex != 0 {
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(m.InitiatorIndex))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0x10
|
|
||||||
}
|
|
||||||
if len(m.Cert) > 0 {
|
|
||||||
i -= len(m.Cert)
|
|
||||||
copy(dAtA[i:], m.Cert)
|
|
||||||
i = encodeVarintNebula(dAtA, i, uint64(len(m.Cert)))
|
|
||||||
i--
|
|
||||||
dAtA[i] = 0xa
|
|
||||||
}
|
|
||||||
return len(dAtA) - i, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
||||||
size := m.Size()
|
size := m.Size()
|
||||||
dAtA = make([]byte, size)
|
dAtA = make([]byte, size)
|
||||||
@@ -1375,51 +1132,6 @@ func (m *NebulaPing) Size() (n int) {
|
|||||||
return n
|
return n
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *NebulaHandshake) Size() (n int) {
|
|
||||||
if m == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
var l int
|
|
||||||
_ = l
|
|
||||||
if m.Details != nil {
|
|
||||||
l = m.Details.Size()
|
|
||||||
n += 1 + l + sovNebula(uint64(l))
|
|
||||||
}
|
|
||||||
l = len(m.Hmac)
|
|
||||||
if l > 0 {
|
|
||||||
n += 1 + l + sovNebula(uint64(l))
|
|
||||||
}
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaHandshakeDetails) Size() (n int) {
|
|
||||||
if m == nil {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
var l int
|
|
||||||
_ = l
|
|
||||||
l = len(m.Cert)
|
|
||||||
if l > 0 {
|
|
||||||
n += 1 + l + sovNebula(uint64(l))
|
|
||||||
}
|
|
||||||
if m.InitiatorIndex != 0 {
|
|
||||||
n += 1 + sovNebula(uint64(m.InitiatorIndex))
|
|
||||||
}
|
|
||||||
if m.ResponderIndex != 0 {
|
|
||||||
n += 1 + sovNebula(uint64(m.ResponderIndex))
|
|
||||||
}
|
|
||||||
if m.Cookie != 0 {
|
|
||||||
n += 1 + sovNebula(uint64(m.Cookie))
|
|
||||||
}
|
|
||||||
if m.Time != 0 {
|
|
||||||
n += 1 + sovNebula(uint64(m.Time))
|
|
||||||
}
|
|
||||||
if m.CertVersion != 0 {
|
|
||||||
n += 1 + sovNebula(uint64(m.CertVersion))
|
|
||||||
}
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *NebulaControl) Size() (n int) {
|
func (m *NebulaControl) Size() (n int) {
|
||||||
if m == nil {
|
if m == nil {
|
||||||
return 0
|
return 0
|
||||||
@@ -2236,305 +1948,6 @@ func (m *NebulaPing) Unmarshal(dAtA []byte) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (m *NebulaHandshake) Unmarshal(dAtA []byte) error {
|
|
||||||
l := len(dAtA)
|
|
||||||
iNdEx := 0
|
|
||||||
for iNdEx < l {
|
|
||||||
preIndex := iNdEx
|
|
||||||
var wire uint64
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
wire |= uint64(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
fieldNum := int32(wire >> 3)
|
|
||||||
wireType := int(wire & 0x7)
|
|
||||||
if wireType == 4 {
|
|
||||||
return fmt.Errorf("proto: NebulaHandshake: wiretype end group for non-group")
|
|
||||||
}
|
|
||||||
if fieldNum <= 0 {
|
|
||||||
return fmt.Errorf("proto: NebulaHandshake: illegal tag %d (wire type %d)", fieldNum, wire)
|
|
||||||
}
|
|
||||||
switch fieldNum {
|
|
||||||
case 1:
|
|
||||||
if wireType != 2 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field Details", wireType)
|
|
||||||
}
|
|
||||||
var msglen int
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
msglen |= int(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if msglen < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
postIndex := iNdEx + msglen
|
|
||||||
if postIndex < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
if postIndex > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
if m.Details == nil {
|
|
||||||
m.Details = &NebulaHandshakeDetails{}
|
|
||||||
}
|
|
||||||
if err := m.Details.Unmarshal(dAtA[iNdEx:postIndex]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
iNdEx = postIndex
|
|
||||||
case 2:
|
|
||||||
if wireType != 2 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field Hmac", wireType)
|
|
||||||
}
|
|
||||||
var byteLen int
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
byteLen |= int(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if byteLen < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
postIndex := iNdEx + byteLen
|
|
||||||
if postIndex < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
if postIndex > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
m.Hmac = append(m.Hmac[:0], dAtA[iNdEx:postIndex]...)
|
|
||||||
if m.Hmac == nil {
|
|
||||||
m.Hmac = []byte{}
|
|
||||||
}
|
|
||||||
iNdEx = postIndex
|
|
||||||
default:
|
|
||||||
iNdEx = preIndex
|
|
||||||
skippy, err := skipNebula(dAtA[iNdEx:])
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
if (iNdEx + skippy) > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
iNdEx += skippy
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if iNdEx > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
func (m *NebulaHandshakeDetails) Unmarshal(dAtA []byte) error {
|
|
||||||
l := len(dAtA)
|
|
||||||
iNdEx := 0
|
|
||||||
for iNdEx < l {
|
|
||||||
preIndex := iNdEx
|
|
||||||
var wire uint64
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
wire |= uint64(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
fieldNum := int32(wire >> 3)
|
|
||||||
wireType := int(wire & 0x7)
|
|
||||||
if wireType == 4 {
|
|
||||||
return fmt.Errorf("proto: NebulaHandshakeDetails: wiretype end group for non-group")
|
|
||||||
}
|
|
||||||
if fieldNum <= 0 {
|
|
||||||
return fmt.Errorf("proto: NebulaHandshakeDetails: illegal tag %d (wire type %d)", fieldNum, wire)
|
|
||||||
}
|
|
||||||
switch fieldNum {
|
|
||||||
case 1:
|
|
||||||
if wireType != 2 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field Cert", wireType)
|
|
||||||
}
|
|
||||||
var byteLen int
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
byteLen |= int(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if byteLen < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
postIndex := iNdEx + byteLen
|
|
||||||
if postIndex < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
if postIndex > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
m.Cert = append(m.Cert[:0], dAtA[iNdEx:postIndex]...)
|
|
||||||
if m.Cert == nil {
|
|
||||||
m.Cert = []byte{}
|
|
||||||
}
|
|
||||||
iNdEx = postIndex
|
|
||||||
case 2:
|
|
||||||
if wireType != 0 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field InitiatorIndex", wireType)
|
|
||||||
}
|
|
||||||
m.InitiatorIndex = 0
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
m.InitiatorIndex |= uint32(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case 3:
|
|
||||||
if wireType != 0 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field ResponderIndex", wireType)
|
|
||||||
}
|
|
||||||
m.ResponderIndex = 0
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
m.ResponderIndex |= uint32(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case 4:
|
|
||||||
if wireType != 0 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field Cookie", wireType)
|
|
||||||
}
|
|
||||||
m.Cookie = 0
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
m.Cookie |= uint64(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case 5:
|
|
||||||
if wireType != 0 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field Time", wireType)
|
|
||||||
}
|
|
||||||
m.Time = 0
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
m.Time |= uint64(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case 8:
|
|
||||||
if wireType != 0 {
|
|
||||||
return fmt.Errorf("proto: wrong wireType = %d for field CertVersion", wireType)
|
|
||||||
}
|
|
||||||
m.CertVersion = 0
|
|
||||||
for shift := uint(0); ; shift += 7 {
|
|
||||||
if shift >= 64 {
|
|
||||||
return ErrIntOverflowNebula
|
|
||||||
}
|
|
||||||
if iNdEx >= l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
b := dAtA[iNdEx]
|
|
||||||
iNdEx++
|
|
||||||
m.CertVersion |= uint32(b&0x7F) << shift
|
|
||||||
if b < 0x80 {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
iNdEx = preIndex
|
|
||||||
skippy, err := skipNebula(dAtA[iNdEx:])
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
|
||||||
return ErrInvalidLengthNebula
|
|
||||||
}
|
|
||||||
if (iNdEx + skippy) > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
iNdEx += skippy
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if iNdEx > l {
|
|
||||||
return io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
||||||
l := len(dAtA)
|
l := len(dAtA)
|
||||||
iNdEx := 0
|
iNdEx := 0
|
||||||
|
|||||||
+3
-15
@@ -60,21 +60,9 @@ message NebulaPing {
|
|||||||
uint64 Time = 2;
|
uint64 Time = 2;
|
||||||
}
|
}
|
||||||
|
|
||||||
message NebulaHandshake {
|
// NebulaHandshake / NebulaHandshakeDetails moved to
|
||||||
NebulaHandshakeDetails Details = 1;
|
// handshake/handshake.proto. The handshake package speaks that wire format
|
||||||
bytes Hmac = 2;
|
// directly via a hand-written encoder/decoder.
|
||||||
}
|
|
||||||
|
|
||||||
message NebulaHandshakeDetails {
|
|
||||||
bytes Cert = 1;
|
|
||||||
uint32 InitiatorIndex = 2;
|
|
||||||
uint32 ResponderIndex = 3;
|
|
||||||
uint64 Cookie = 4;
|
|
||||||
uint64 Time = 5;
|
|
||||||
uint32 CertVersion = 8;
|
|
||||||
// reserved for WIP multiport
|
|
||||||
reserved 6, 7;
|
|
||||||
}
|
|
||||||
|
|
||||||
message NebulaControl {
|
message NebulaControl {
|
||||||
enum MessageType {
|
enum MessageType {
|
||||||
|
|||||||
@@ -14,6 +14,19 @@ type endianness interface {
|
|||||||
|
|
||||||
var noiseEndianness endianness = binary.BigEndian
|
var noiseEndianness endianness = binary.BigEndian
|
||||||
|
|
||||||
|
// NonceSize is the AEAD nonce length used by all ciphers nebula supports
|
||||||
|
// today (AES-GCM and ChaCha20-Poly1305 both use 96-bit nonces). Encrypt-
|
||||||
|
// and DecryptDanger lay out the nonce as 4 zero bytes followed by an 8-byte
|
||||||
|
// big-endian counter; if a future cipher with a different nonce size is
|
||||||
|
// added, this constant and those layouts must change together.
|
||||||
|
const NonceSize = 12
|
||||||
|
|
||||||
|
// AEADOverhead is the AEAD authentication tag length the ciphers nebula
|
||||||
|
// supports append to ciphertext. Both AES-GCM and ChaCha20-Poly1305 use
|
||||||
|
// 128-bit tags. NebulaCipherState.Overhead() returns this dynamically from
|
||||||
|
// the cipher; the constant is for sizing buffers at construction time.
|
||||||
|
const AEADOverhead = 16
|
||||||
|
|
||||||
type NebulaCipherState struct {
|
type NebulaCipherState struct {
|
||||||
c cipher.AEAD
|
c cipher.AEAD
|
||||||
}
|
}
|
||||||
|
|||||||
+35
-33
@@ -20,7 +20,8 @@ const (
|
|||||||
minFwPacketLen = 4
|
minFwPacketLen = 4
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) readOutsidePackets(via ViaSender, buf *WireBuffer, packet []byte, lhf *LightHouseHandler, q int, localCache firewall.ConntrackCache) {
|
||||||
|
h := buf.H
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
@@ -65,7 +66,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
|
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
if !f.decryptToTun(hostinfo, h.MessageCounter, buf, packet, q, localCache) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
case header.MessageRelay:
|
case header.MessageRelay:
|
||||||
@@ -76,8 +77,9 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
// which will gracefully fail in the DecryptDanger call.
|
// which will gracefully fail in the DecryptDanger call.
|
||||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
// AAD-only validation: passing dst=nil since there's no plaintext
|
||||||
if err != nil {
|
// to recover (ciphertext is just the trailing AEAD tag).
|
||||||
|
if _, err = hostinfo.ConnectionState.dKey.DecryptDanger(nil, signedPayload, signatureValue, h.MessageCounter, buf.NB); err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
// Successfully validated the thing. Get rid of the Relay header.
|
// Successfully validated the thing. Get rid of the Relay header.
|
||||||
@@ -110,7 +112,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
IsRelayed: true,
|
||||||
}
|
}
|
||||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
buf.Reset()
|
||||||
|
f.readOutsidePackets(via, buf, signedPayload, lhf, q, localCache)
|
||||||
return
|
return
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// Find the target HostInfo relay object
|
||||||
@@ -130,7 +133,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel
|
||||||
// Find the target HostInfo
|
// Find the target HostInfo
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
f.SendVia(targetHI, targetRelay, signedPayload, buf)
|
||||||
return
|
return
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
@@ -152,7 +155,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
d, err := f.decrypt(hostinfo, h.MessageCounter, buf, packet, h)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt lighthouse packet",
|
hostinfo.logger(f.l).Error("Failed to decrypt lighthouse packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
@@ -173,7 +176,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
d, err := f.decrypt(hostinfo, h.MessageCounter, buf, packet, h)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt test packet",
|
hostinfo.logger(f.l).Error("Failed to decrypt test packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
@@ -185,9 +188,9 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
|
|
||||||
if h.Subtype == header.TestRequest {
|
if h.Subtype == header.TestRequest {
|
||||||
// This testRequest might be from TryPromoteBest, so we should roam
|
// This testRequest might be from TryPromoteBest, so we should roam
|
||||||
// to the new IP address before responding
|
// to the new IP address before responding.
|
||||||
f.handleHostRoaming(hostinfo, via)
|
f.handleHostRoaming(hostinfo, via)
|
||||||
f.send(header.Test, header.TestReply, ci, hostinfo, d, nb, out)
|
f.send(header.Test, header.TestReply, ci, hostinfo, d, buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fallthrough to the bottom to record incoming traffic
|
// Fallthrough to the bottom to record incoming traffic
|
||||||
@@ -210,7 +213,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
if !f.handleEncrypted(ci, via, h) {
|
if !f.handleEncrypted(ci, via, h) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
_, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
_, err = f.decrypt(hostinfo, h.MessageCounter, buf, packet, h)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt CloseTunnel packet",
|
hostinfo.logger(f.l).Error("Failed to decrypt CloseTunnel packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
@@ -230,7 +233,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
d, err := f.decrypt(hostinfo, h.MessageCounter, buf, packet, h)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt Control packet",
|
hostinfo.logger(f.l).Error("Failed to decrypt Control packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
@@ -266,7 +269,9 @@ func (f *Interface) closeTunnel(hostInfo *HostInfo) {
|
|||||||
|
|
||||||
// sendCloseTunnel is a helper function to send a proper close tunnel packet to a remote
|
// sendCloseTunnel is a helper function to send a proper close tunnel packet to a remote
|
||||||
func (f *Interface) sendCloseTunnel(h *HostInfo) {
|
func (f *Interface) sendCloseTunnel(h *HostInfo) {
|
||||||
f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
buf := f.bufAlloc.Acquire()
|
||||||
|
defer f.bufAlloc.Release(buf)
|
||||||
|
f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, buf)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
||||||
@@ -515,9 +520,8 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []byte, h *header.H, nb []byte) ([]byte, error) {
|
func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, buf *WireBuffer, packet []byte, h *header.H) ([]byte, error) {
|
||||||
var err error
|
plaintext, err := buf.DecryptForHandler(hostinfo.ConnectionState, packet, mc)
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], mc, nb)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -529,42 +533,41 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
|||||||
return nil, errors.New("out of window packet")
|
return nil, errors.New("out of window packet")
|
||||||
}
|
}
|
||||||
|
|
||||||
return out, nil
|
return plaintext, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool {
|
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, buf *WireBuffer, packet []byte, q int, localCache firewall.ConntrackCache) bool {
|
||||||
var err error
|
if err := buf.DecryptDatagram(hostinfo.ConnectionState, packet, messageCounter); err != nil {
|
||||||
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
err = newPacket(out, true, fwPacket)
|
ipPacket := buf.IPPacket()
|
||||||
if err != nil {
|
if err := newPacket(ipPacket, true, buf.FwPacket); err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"packet", out,
|
"packet", ipPacket,
|
||||||
)
|
)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, messageCounter) {
|
if !hostinfo.ConnectionState.window.Update(f.l, messageCounter) {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping out of window packet", "fwPacket", fwPacket)
|
hostinfo.logger(f.l).Debug("dropping out of window packet", "fwPacket", buf.FwPacket)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*buf.FwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
// NOTE: We hand `packet` (the original UDP ciphertext we already
|
||||||
// This gives us a buffer to build the reject packet in
|
// decrypted from) as the reject-IP scratch since we no longer
|
||||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
// need its ciphertext, and it's disjoint from buf.Out where
|
||||||
|
// sendNoMetrics will encrypt the wire packet.
|
||||||
|
f.rejectOutside(ipPacket, hostinfo.ConnectionState, hostinfo, packet, buf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", buf.FwPacket,
|
||||||
"reason", dropReason,
|
"reason", dropReason,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -572,8 +575,7 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out
|
|||||||
}
|
}
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
f.connectionManager.In(hostinfo)
|
||||||
_, err = f.readers[q].Write(out)
|
if _, err := buf.WriteIPToTUN(f.readers[q]); err != nil {
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
|
|||||||
@@ -15,4 +15,7 @@ type Device interface {
|
|||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
SupportsMultiqueue() bool
|
SupportsMultiqueue() bool
|
||||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
||||||
|
// TunPrefixLen reports the number of bytes the device prepends to every IP packet on the wire.
|
||||||
|
// Currently only non zero for the BSD tun devices.
|
||||||
|
TunPrefixLen() int
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -50,3 +50,5 @@ func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
|||||||
func (NoopTun) Close() error {
|
func (NoopTun) Close() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (NoopTun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
@@ -102,3 +102,5 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
//go:build (darwin || ios || freebsd || openbsd || netbsd) && !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"syscall"
|
||||||
|
)
|
||||||
|
|
||||||
|
// StampTunPrefix writes the 4-byte AF_INET / AF_INET6 protocol-family marker into buf[0:4] in place,
|
||||||
|
// picking the family from the first byte of the IP packet at buf[4].
|
||||||
|
func StampTunPrefix(buf []byte) error {
|
||||||
|
if len(buf) < 5 {
|
||||||
|
return fmt.Errorf("tun write buffer too small for prefix")
|
||||||
|
}
|
||||||
|
ipVer := buf[4] >> 4
|
||||||
|
buf[0] = 0
|
||||||
|
buf[1] = 0
|
||||||
|
buf[2] = 0
|
||||||
|
switch ipVer {
|
||||||
|
case 4:
|
||||||
|
buf[3] = syscall.AF_INET
|
||||||
|
case 6:
|
||||||
|
buf[3] = syscall.AF_INET6
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unable to determine IP version from packet")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+4
-42
@@ -11,7 +11,6 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"syscall"
|
|
||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
@@ -31,9 +30,6 @@ type tun struct {
|
|||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
linkAddr *netroute.LinkAddr
|
linkAddr *netroute.LinkAddr
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
|
|
||||||
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
|
||||||
out []byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ifReq struct {
|
type ifReq struct {
|
||||||
@@ -502,44 +498,6 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
|
||||||
buf := make([]byte, len(to)+4)
|
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Read(buf)
|
|
||||||
|
|
||||||
copy(to, buf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
|
||||||
buf := t.out
|
|
||||||
if cap(buf) < len(from)+4 {
|
|
||||||
buf = make([]byte, len(from)+4)
|
|
||||||
t.out = buf
|
|
||||||
}
|
|
||||||
buf = buf[:len(from)+4]
|
|
||||||
|
|
||||||
if len(from) == 0 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
// Determine the IP Family for the NULL L2 Header
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
if ipVer == 4 {
|
|
||||||
buf[3] = syscall.AF_INET
|
|
||||||
} else if ipVer == 6 {
|
|
||||||
buf[3] = syscall.AF_INET6
|
|
||||||
} else {
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(buf[4:], from)
|
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Write(buf)
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Networks() []netip.Prefix {
|
func (t *tun) Networks() []netip.Prefix {
|
||||||
return t.vpnNetworks
|
return t.vpnNetworks
|
||||||
}
|
}
|
||||||
@@ -555,3 +513,7 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunPrefixLen reports the 4-byte BSD AF_INET / AF_INET6 protocol-family
|
||||||
|
// marker the kernel prepends on read and expects on write.
|
||||||
|
func (t *tun) TunPrefixLen() int { return 4 }
|
||||||
|
|||||||
@@ -136,3 +136,5 @@ func (p prettyPacket) String() string {
|
|||||||
|
|
||||||
return s.String()
|
return s.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
+18
-45
@@ -158,74 +158,43 @@ func (t *tun) blockOnWrite() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(to []byte) (int, error) {
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
var head [4]byte
|
|
||||||
iovecs := [2]syscall.Iovec{
|
|
||||||
{&head[0], 4},
|
|
||||||
{&to[0], uint64(len(to))},
|
|
||||||
}
|
|
||||||
for {
|
for {
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_READV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
n, err := unix.Read(t.fd, to)
|
||||||
if errno == 0 {
|
if err == nil {
|
||||||
bytesRead := int(n)
|
return n, nil
|
||||||
if bytesRead < 4 {
|
|
||||||
return 0, nil
|
|
||||||
}
|
}
|
||||||
return bytesRead - 4, nil
|
switch err {
|
||||||
}
|
|
||||||
switch errno {
|
|
||||||
case unix.EAGAIN:
|
case unix.EAGAIN:
|
||||||
if err := t.blockOnRead(); err != nil {
|
if berr := t.blockOnRead(); berr != nil {
|
||||||
return 0, err
|
return 0, berr
|
||||||
}
|
}
|
||||||
case unix.EINTR:
|
case unix.EINTR:
|
||||||
// retry
|
// retry
|
||||||
case unix.EBADF:
|
case unix.EBADF:
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
default:
|
default:
|
||||||
return 0, errno
|
return 0, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
if len(from) <= 1 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
var head [4]byte
|
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
switch ipVer {
|
|
||||||
case 4:
|
|
||||||
head[3] = syscall.AF_INET
|
|
||||||
case 6:
|
|
||||||
head[3] = syscall.AF_INET6
|
|
||||||
default:
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
iovecs := [2]syscall.Iovec{
|
|
||||||
{&head[0], 4},
|
|
||||||
{&from[0], uint64(len(from))},
|
|
||||||
}
|
|
||||||
for {
|
for {
|
||||||
n, _, errno := syscall.Syscall(syscall.SYS_WRITEV, uintptr(t.fd), uintptr(unsafe.Pointer(&iovecs[0])), 2)
|
n, err := unix.Write(t.fd, from)
|
||||||
if errno == 0 {
|
if err == nil {
|
||||||
return int(n) - 4, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
switch errno {
|
switch err {
|
||||||
case unix.EAGAIN:
|
case unix.EAGAIN:
|
||||||
if err := t.blockOnWrite(); err != nil {
|
if berr := t.blockOnWrite(); berr != nil {
|
||||||
return 0, err
|
return 0, berr
|
||||||
}
|
}
|
||||||
case unix.EINTR:
|
case unix.EINTR:
|
||||||
// retry
|
// retry
|
||||||
case unix.EBADF:
|
case unix.EBADF:
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
default:
|
default:
|
||||||
return 0, errno
|
return 0, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -732,3 +701,7 @@ func getLinkAddr(name string) (*netroute.LinkAddr, error) {
|
|||||||
|
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunPrefixLen reports the 4-byte BSD AF_INET / AF_INET6 protocol-family
|
||||||
|
// marker the kernel prepends on read and expects on write.
|
||||||
|
func (t *tun) TunPrefixLen() int { return 4 }
|
||||||
|
|||||||
+5
-62
@@ -4,15 +4,12 @@
|
|||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"syscall"
|
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
@@ -36,7 +33,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||||
t := &tun{
|
t := &tun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
ReadWriteCloser: &tunReadCloser{f: file},
|
ReadWriteCloser: file,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -85,64 +82,6 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
// The following is hoisted up from water, we do this so we can inject our own fd on iOS
|
|
||||||
type tunReadCloser struct {
|
|
||||||
f io.ReadWriteCloser
|
|
||||||
|
|
||||||
rMu sync.Mutex
|
|
||||||
rBuf []byte
|
|
||||||
|
|
||||||
wMu sync.Mutex
|
|
||||||
wBuf []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (tr *tunReadCloser) Read(to []byte) (int, error) {
|
|
||||||
tr.rMu.Lock()
|
|
||||||
defer tr.rMu.Unlock()
|
|
||||||
|
|
||||||
if cap(tr.rBuf) < len(to)+4 {
|
|
||||||
tr.rBuf = make([]byte, len(to)+4)
|
|
||||||
}
|
|
||||||
tr.rBuf = tr.rBuf[:len(to)+4]
|
|
||||||
|
|
||||||
n, err := tr.f.Read(tr.rBuf)
|
|
||||||
copy(to, tr.rBuf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (tr *tunReadCloser) Write(from []byte) (int, error) {
|
|
||||||
if len(from) == 0 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
tr.wMu.Lock()
|
|
||||||
defer tr.wMu.Unlock()
|
|
||||||
|
|
||||||
if cap(tr.wBuf) < len(from)+4 {
|
|
||||||
tr.wBuf = make([]byte, len(from)+4)
|
|
||||||
}
|
|
||||||
tr.wBuf = tr.wBuf[:len(from)+4]
|
|
||||||
|
|
||||||
// Determine the IP Family for the NULL L2 Header
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
if ipVer == 4 {
|
|
||||||
tr.wBuf[3] = syscall.AF_INET
|
|
||||||
} else if ipVer == 6 {
|
|
||||||
tr.wBuf[3] = syscall.AF_INET6
|
|
||||||
} else {
|
|
||||||
return 0, errors.New("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(tr.wBuf[4:], from)
|
|
||||||
|
|
||||||
n, err := tr.f.Write(tr.wBuf)
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (tr *tunReadCloser) Close() error {
|
|
||||||
return tr.f.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) Networks() []netip.Prefix {
|
func (t *tun) Networks() []netip.Prefix {
|
||||||
return t.vpnNetworks
|
return t.vpnNetworks
|
||||||
}
|
}
|
||||||
@@ -158,3 +97,7 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunPrefixLen reports the 4-byte BSD AF_INET / AF_INET6 protocol-family
|
||||||
|
// marker the kernel prepends on read and expects on write.
|
||||||
|
func (t *tun) TunPrefixLen() int { return 4 }
|
||||||
|
|||||||
@@ -907,3 +907,5 @@ func (t *tun) Close() error {
|
|||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
+9
-98
@@ -58,13 +58,13 @@ type addrLifetime struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
|
io.ReadWriteCloser
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
MTU int
|
MTU int
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
f *os.File
|
|
||||||
fd int
|
fd int
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,7 +96,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
f: os.NewFile(uintptr(fd), ""),
|
ReadWriteCloser: os.NewFile(uintptr(fd), ""),
|
||||||
fd: fd,
|
fd: fd,
|
||||||
Device: deviceName,
|
Device: deviceName,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
@@ -120,12 +120,12 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.f != nil {
|
if t.ReadWriteCloser != nil {
|
||||||
if err := t.f.Close(); err != nil {
|
if err := t.ReadWriteCloser.Close(); err != nil {
|
||||||
return fmt.Errorf("error closing tun file: %w", err)
|
return fmt.Errorf("error closing tun file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// t.f.Close should have handled it for us but let's be extra sure
|
// Close on the os.File should have handled the fd for us but let's be extra sure
|
||||||
_ = unix.Close(t.fd)
|
_ = unix.Close(t.fd)
|
||||||
|
|
||||||
s, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_DGRAM, syscall.IPPROTO_IP)
|
s, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_DGRAM, syscall.IPPROTO_IP)
|
||||||
@@ -141,99 +141,6 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
|
||||||
rc, err := t.f.SyscallConn()
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("failed to get syscall conn for tun: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var errno syscall.Errno
|
|
||||||
var n uintptr
|
|
||||||
err = rc.Read(func(fd uintptr) bool {
|
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
head := [4]byte{}
|
|
||||||
iovecs := []syscall.Iovec{
|
|
||||||
{&head[0], 4},
|
|
||||||
{&to[0], uint64(len(to))},
|
|
||||||
}
|
|
||||||
|
|
||||||
n, _, errno = syscall.Syscall(syscall.SYS_READV, fd, uintptr(unsafe.Pointer(&iovecs[0])), uintptr(2))
|
|
||||||
if errno.Temporary() {
|
|
||||||
// We got an EAGAIN, EINTR, or EWOULDBLOCK, go again
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
if err == syscall.EBADF || err.Error() == "use of closed file" {
|
|
||||||
// Go doesn't export poll.ErrFileClosing but happily reports it to us so here we are
|
|
||||||
// https://github.com/golang/go/blob/master/src/internal/poll/fd_poll_runtime.go#L121
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
}
|
|
||||||
return 0, fmt.Errorf("failed to make read call for tun: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if errno != 0 {
|
|
||||||
return 0, fmt.Errorf("failed to make inner read call for tun: %w", errno)
|
|
||||||
}
|
|
||||||
|
|
||||||
// fix bytes read number to exclude header
|
|
||||||
bytesRead := int(n)
|
|
||||||
if bytesRead < 0 {
|
|
||||||
return bytesRead, nil
|
|
||||||
} else if bytesRead < 4 {
|
|
||||||
return 0, nil
|
|
||||||
} else {
|
|
||||||
return bytesRead - 4, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
|
||||||
if len(from) <= 1 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
var head [4]byte
|
|
||||||
// first 4 bytes is protocol family, in network byte order
|
|
||||||
if ipVer == 4 {
|
|
||||||
head[3] = syscall.AF_INET
|
|
||||||
} else if ipVer == 6 {
|
|
||||||
head[3] = syscall.AF_INET6
|
|
||||||
} else {
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
rc, err := t.f.SyscallConn()
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
var errno syscall.Errno
|
|
||||||
var n uintptr
|
|
||||||
err = rc.Write(func(fd uintptr) bool {
|
|
||||||
iovecs := []syscall.Iovec{
|
|
||||||
{&head[0], 4},
|
|
||||||
{&from[0], uint64(len(from))},
|
|
||||||
}
|
|
||||||
|
|
||||||
n, _, errno = syscall.Syscall(syscall.SYS_WRITEV, fd, uintptr(unsafe.Pointer(&iovecs[0])), uintptr(2))
|
|
||||||
// According to NetBSD documentation for TUN, writes will only return errors in which
|
|
||||||
// this packet will never be delivered so just go on living life.
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if errno != 0 {
|
|
||||||
return 0, errno
|
|
||||||
}
|
|
||||||
|
|
||||||
return int(n) - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) addIp(cidr netip.Prefix) error {
|
func (t *tun) addIp(cidr netip.Prefix) error {
|
||||||
if cidr.Addr().Is4() {
|
if cidr.Addr().Is4() {
|
||||||
var req ifreqAlias4
|
var req ifreqAlias4
|
||||||
@@ -551,3 +458,7 @@ func delRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
|||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunPrefixLen reports the 4-byte BSD AF_INET / AF_INET6 protocol-family
|
||||||
|
// marker the kernel prepends on read and expects on write.
|
||||||
|
func (t *tun) TunPrefixLen() int { return 4 }
|
||||||
|
|||||||
@@ -0,0 +1,10 @@
|
|||||||
|
//go:build (!darwin && !ios && !freebsd && !openbsd && !netbsd) || e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
// StampTunPrefix is a no-op on platforms whose tun devices have no
|
||||||
|
// protocol-family marker. WireBuffer only invokes it when its prefixLen
|
||||||
|
// is non-zero, so this should never be reached on these platforms.
|
||||||
|
func StampTunPrefix(buf []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+9
-45
@@ -49,16 +49,14 @@ type ifreq struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
|
io.ReadWriteCloser
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
MTU int
|
MTU int
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
f *os.File
|
|
||||||
fd int
|
fd int
|
||||||
// cache out buffer since we need to prepend 4 bytes for tun metadata
|
|
||||||
out []byte
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
@@ -89,7 +87,7 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
f: os.NewFile(uintptr(fd), ""),
|
ReadWriteCloser: os.NewFile(uintptr(fd), ""),
|
||||||
fd: fd,
|
fd: fd,
|
||||||
Device: deviceName,
|
Device: deviceName,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
@@ -113,55 +111,17 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.f != nil {
|
if t.ReadWriteCloser != nil {
|
||||||
if err := t.f.Close(); err != nil {
|
if err := t.ReadWriteCloser.Close(); err != nil {
|
||||||
return fmt.Errorf("error closing tun file: %w", err)
|
return fmt.Errorf("error closing tun file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// t.f.Close should have handled it for us but let's be extra sure
|
// Close on the os.File should have handled the fd for us but let's be extra sure
|
||||||
_ = unix.Close(t.fd)
|
_ = unix.Close(t.fd)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
|
||||||
buf := make([]byte, len(to)+4)
|
|
||||||
|
|
||||||
n, err := t.f.Read(buf)
|
|
||||||
|
|
||||||
copy(to, buf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
|
||||||
buf := t.out
|
|
||||||
if cap(buf) < len(from)+4 {
|
|
||||||
buf = make([]byte, len(from)+4)
|
|
||||||
t.out = buf
|
|
||||||
}
|
|
||||||
buf = buf[:len(from)+4]
|
|
||||||
|
|
||||||
if len(from) == 0 {
|
|
||||||
return 0, syscall.EIO
|
|
||||||
}
|
|
||||||
|
|
||||||
// Determine the IP Family for the NULL L2 Header
|
|
||||||
ipVer := from[0] >> 4
|
|
||||||
if ipVer == 4 {
|
|
||||||
buf[3] = syscall.AF_INET
|
|
||||||
} else if ipVer == 6 {
|
|
||||||
buf[3] = syscall.AF_INET6
|
|
||||||
} else {
|
|
||||||
return 0, fmt.Errorf("unable to determine IP version from packet")
|
|
||||||
}
|
|
||||||
|
|
||||||
copy(buf[4:], from)
|
|
||||||
|
|
||||||
n, err := t.f.Write(buf)
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *tun) addIp(cidr netip.Prefix) error {
|
func (t *tun) addIp(cidr netip.Prefix) error {
|
||||||
if cidr.Addr().Is4() {
|
if cidr.Addr().Is4() {
|
||||||
var req ifreqAlias4
|
var req ifreqAlias4
|
||||||
@@ -471,3 +431,7 @@ func delRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
|||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunPrefixLen reports the 4-byte BSD AF_INET / AF_INET6 protocol-family
|
||||||
|
// marker the kernel prepends on read and expects on write.
|
||||||
|
func (t *tun) TunPrefixLen() int { return 4 }
|
||||||
|
|||||||
+51
-5
@@ -15,6 +15,7 @@ import (
|
|||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
type TestTun struct {
|
type TestTun struct {
|
||||||
@@ -54,9 +55,12 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*TestTu
|
|||||||
return nil, fmt.Errorf("newTunFromFd not supported")
|
return nil, fmt.Errorf("newTunFromFd not supported")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send will place a byte array onto the receive queue for nebula to consume
|
// Send will place a byte array onto the receive queue for nebula to consume.
|
||||||
// These are unencrypted ip layer frames destined for another nebula node.
|
// These are unencrypted ip layer frames destined for another nebula node.
|
||||||
// packets should exit the udp side, capture them with udpConn.Get
|
// packets should exit the udp side, capture them with udpConn.Get.
|
||||||
|
//
|
||||||
|
// Send copies the input via the freelist, so the caller is free to mutate
|
||||||
|
// or reuse it after the call returns.
|
||||||
func (t *TestTun) Send(packet []byte) {
|
func (t *TestTun) Send(packet []byte) {
|
||||||
if t.closed.Load() {
|
if t.closed.Load() {
|
||||||
return
|
return
|
||||||
@@ -65,7 +69,9 @@ func (t *TestTun) Send(packet []byte) {
|
|||||||
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
if t.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
t.l.Debug("Tun receiving injected packet", "dataLen", len(packet))
|
t.l.Debug("Tun receiving injected packet", "dataLen", len(packet))
|
||||||
}
|
}
|
||||||
t.rxPackets <- packet
|
buf := acquireTunBuf(len(packet))
|
||||||
|
copy(buf, packet)
|
||||||
|
t.rxPackets <- buf
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get will pull an unencrypted ip layer frame from the transmit queue
|
// Get will pull an unencrypted ip layer frame from the transmit queue
|
||||||
@@ -110,12 +116,44 @@ func (t *TestTun) Write(b []byte) (n int, err error) {
|
|||||||
return 0, io.ErrClosedPipe
|
return 0, io.ErrClosedPipe
|
||||||
}
|
}
|
||||||
|
|
||||||
packet := make([]byte, len(b), len(b))
|
packet := acquireTunBuf(len(b))
|
||||||
copy(packet, b)
|
copy(packet, b)
|
||||||
t.TxPackets <- packet
|
t.TxPackets <- packet
|
||||||
return len(b), nil
|
return len(b), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ReleaseTunBuf returns a slice from TxPackets to the harness freelist, don't use the bytes after the call.
|
||||||
|
// Channel-backed instead of sync.Pool because putting a []byte in a sync.Pool escapes the slice header to heap.
|
||||||
|
func ReleaseTunBuf(b []byte) {
|
||||||
|
if b == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case tunBufFreelist <- b:
|
||||||
|
default:
|
||||||
|
// Freelist full; drop the buffer for the GC.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// tunBufFreelist retains the backing arrays for TestTun.Write so steady-state allocation drops to zero once the
|
||||||
|
// freelist has saturated for the current MTU.
|
||||||
|
var tunBufFreelist = make(chan []byte, 64)
|
||||||
|
|
||||||
|
func acquireTunBuf(n int) []byte {
|
||||||
|
var b []byte
|
||||||
|
select {
|
||||||
|
case b = <-tunBufFreelist:
|
||||||
|
default:
|
||||||
|
b = make([]byte, 0, udp.MTU)
|
||||||
|
}
|
||||||
|
if cap(b) < n {
|
||||||
|
b = make([]byte, n)
|
||||||
|
} else {
|
||||||
|
b = b[:n]
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
func (t *TestTun) Close() error {
|
func (t *TestTun) Close() error {
|
||||||
if t.closed.CompareAndSwap(false, true) {
|
if t.closed.CompareAndSwap(false, true) {
|
||||||
close(t.rxPackets)
|
close(t.rxPackets)
|
||||||
@@ -129,8 +167,14 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
|||||||
if !ok {
|
if !ok {
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
}
|
}
|
||||||
|
n := len(p)
|
||||||
copy(b, p)
|
copy(b, p)
|
||||||
return len(p), nil
|
// Send always pushes a freelist-acquired slice, return it once we've copied the bytes into the caller's buffer.
|
||||||
|
select {
|
||||||
|
case tunBufFreelist <- p:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) SupportsMultiqueue() bool {
|
func (t *TestTun) SupportsMultiqueue() bool {
|
||||||
@@ -140,3 +184,5 @@ func (t *TestTun) SupportsMultiqueue() bool {
|
|||||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
@@ -296,3 +296,5 @@ func checkWinTunExists() error {
|
|||||||
_, err = syscall.LoadDLL(filepath.Join(filepath.Dir(myPath), "dist", "windows", "wintun", "bin", arch, "wintun.dll"))
|
_, err = syscall.LoadDLL(filepath.Join(filepath.Dir(myPath), "dist", "windows", "wintun", "bin", arch, "wintun.dll"))
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *winTun) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
@@ -69,3 +69,5 @@ func (d *UserDevice) Close() error {
|
|||||||
d.outboundWriter.Close()
|
d.outboundWriter.Close()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) TunPrefixLen() int { return 0 }
|
||||||
|
|||||||
@@ -15,9 +15,12 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/handshake"
|
||||||
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -29,10 +32,10 @@ type PKI struct {
|
|||||||
|
|
||||||
type CertState struct {
|
type CertState struct {
|
||||||
v1Cert cert.Certificate
|
v1Cert cert.Certificate
|
||||||
v1HandshakeBytes []byte
|
v1Credential *handshake.Credential
|
||||||
|
|
||||||
v2Cert cert.Certificate
|
v2Cert cert.Certificate
|
||||||
v2HandshakeBytes []byte
|
v2Credential *handshake.Credential
|
||||||
|
|
||||||
initiatingVersion cert.Version
|
initiatingVersion cert.Version
|
||||||
privateKey []byte
|
privateKey []byte
|
||||||
@@ -92,13 +95,35 @@ func (p *PKI) reload(c *config.C, initial bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *PKI) reloadCerts(c *config.C, initial bool) *util.ContextualError {
|
func (p *PKI) reloadCerts(c *config.C, initial bool) *util.ContextualError {
|
||||||
newState, err := newCertStateFromConfig(c)
|
var cipher string
|
||||||
|
var currentState *CertState
|
||||||
|
if initial {
|
||||||
|
cipher = c.GetString("cipher", "aes")
|
||||||
|
//TODO: this sucks and we should make it not a global
|
||||||
|
switch cipher {
|
||||||
|
case "aes":
|
||||||
|
noiseEndianness = binary.BigEndian
|
||||||
|
case "chachapoly":
|
||||||
|
noiseEndianness = binary.LittleEndian
|
||||||
|
default:
|
||||||
|
return util.NewContextualError(
|
||||||
|
"unknown cipher",
|
||||||
|
m{"cipher": cipher},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Cipher cant be hot swapped so just leave it at what it was before
|
||||||
|
currentState = p.cs.Load()
|
||||||
|
cipher = currentState.cipher
|
||||||
|
}
|
||||||
|
|
||||||
|
newState, err := newCertStateFromConfig(c, cipher)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return util.NewContextualError("Could not load client cert", nil, err)
|
return util.NewContextualError("Could not load client cert", nil, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !initial {
|
if currentState != nil {
|
||||||
currentState := p.cs.Load()
|
|
||||||
if newState.v1Cert != nil {
|
if newState.v1Cert != nil {
|
||||||
if currentState.v1Cert == nil {
|
if currentState.v1Cert == nil {
|
||||||
//adding certs is fine, actually. Networks-in-common confirmed in newCertState().
|
//adding certs is fine, actually. Networks-in-common confirmed in newCertState().
|
||||||
@@ -158,25 +183,6 @@ func (p *PKI) reloadCerts(c *config.C, initial bool) *util.ContextualError {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Cipher cant be hot swapped so just leave it at what it was before
|
|
||||||
newState.cipher = currentState.cipher
|
|
||||||
|
|
||||||
} else {
|
|
||||||
newState.cipher = c.GetString("cipher", "aes")
|
|
||||||
//TODO: this sucks and we should make it not a global
|
|
||||||
switch newState.cipher {
|
|
||||||
case "aes":
|
|
||||||
noiseEndianness = binary.BigEndian
|
|
||||||
case "chachapoly":
|
|
||||||
noiseEndianness = binary.LittleEndian
|
|
||||||
default:
|
|
||||||
return util.NewContextualError(
|
|
||||||
"unknown cipher",
|
|
||||||
m{"cipher": newState.cipher},
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
p.cs.Store(newState)
|
p.cs.Store(newState)
|
||||||
@@ -208,6 +214,20 @@ func (cs *CertState) GetDefaultCertificate() cert.Certificate {
|
|||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DefaultVersion returns the preferred cert version for initiating handshakes.
|
||||||
|
func (cs *CertState) DefaultVersion() cert.Version { return cs.initiatingVersion }
|
||||||
|
|
||||||
|
// GetCredential returns the pre-computed handshake credential for the given version, or nil.
|
||||||
|
func (cs *CertState) GetCredential(v cert.Version) *handshake.Credential {
|
||||||
|
switch v {
|
||||||
|
case cert.Version1:
|
||||||
|
return cs.v1Credential
|
||||||
|
case cert.Version2:
|
||||||
|
return cs.v2Credential
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (cs *CertState) getCertificate(v cert.Version) cert.Certificate {
|
func (cs *CertState) getCertificate(v cert.Version) cert.Certificate {
|
||||||
switch v {
|
switch v {
|
||||||
case cert.Version1:
|
case cert.Version1:
|
||||||
@@ -219,17 +239,25 @@ func (cs *CertState) getCertificate(v cert.Version) cert.Certificate {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// getHandshakeBytes returns the cached bytes to be used in a handshake message for the requested version.
|
func newCipherSuite(curve cert.Curve, pkcs11backed bool, cipher string) (noise.CipherSuite, error) {
|
||||||
// Callers must check if the return []byte is nil.
|
var dhFunc noise.DHFunc
|
||||||
func (cs *CertState) getHandshakeBytes(v cert.Version) []byte {
|
switch curve {
|
||||||
switch v {
|
case cert.Curve_CURVE25519:
|
||||||
case cert.Version1:
|
dhFunc = noise.DH25519
|
||||||
return cs.v1HandshakeBytes
|
case cert.Curve_P256:
|
||||||
case cert.Version2:
|
if pkcs11backed {
|
||||||
return cs.v2HandshakeBytes
|
dhFunc = noiseutil.DHP256PKCS11
|
||||||
default:
|
} else {
|
||||||
return nil
|
dhFunc = noiseutil.DHP256
|
||||||
}
|
}
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("unsupported curve: %s", curve)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cipher == "chachapoly" {
|
||||||
|
return noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256), nil
|
||||||
|
}
|
||||||
|
return noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *CertState) String() string {
|
func (cs *CertState) String() string {
|
||||||
@@ -261,7 +289,7 @@ func (cs *CertState) MarshalJSON() ([]byte, error) {
|
|||||||
return json.Marshal(msg)
|
return json.Marshal(msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCertStateFromConfig(c *config.C) (*CertState, error) {
|
func newCertStateFromConfig(c *config.C, cipher string) (*CertState, error) {
|
||||||
var err error
|
var err error
|
||||||
|
|
||||||
privPathOrPEM := c.GetString("pki.key", "")
|
privPathOrPEM := c.GetString("pki.key", "")
|
||||||
@@ -345,13 +373,14 @@ func newCertStateFromConfig(c *config.C) (*CertState, error) {
|
|||||||
return nil, fmt.Errorf("unknown pki.initiating_version: %v", rawInitiatingVersion)
|
return nil, fmt.Errorf("unknown pki.initiating_version: %v", rawInitiatingVersion)
|
||||||
}
|
}
|
||||||
|
|
||||||
return newCertState(initiatingVersion, v1, v2, isPkcs11, curve, rawKey)
|
return newCertState(initiatingVersion, v1, v2, isPkcs11, curve, rawKey, cipher)
|
||||||
}
|
}
|
||||||
|
|
||||||
func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, privateKeyCurve cert.Curve, privateKey []byte) (*CertState, error) {
|
func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, privateKeyCurve cert.Curve, privateKey []byte, cipher string) (*CertState, error) {
|
||||||
cs := CertState{
|
cs := CertState{
|
||||||
privateKey: privateKey,
|
privateKey: privateKey,
|
||||||
pkcs11Backed: pkcs11backed,
|
pkcs11Backed: pkcs11backed,
|
||||||
|
cipher: cipher,
|
||||||
myVpnNetworksTable: new(bart.Lite),
|
myVpnNetworksTable: new(bart.Lite),
|
||||||
myVpnAddrsTable: new(bart.Lite),
|
myVpnAddrsTable: new(bart.Lite),
|
||||||
myVpnBroadcastAddrsTable: new(bart.Lite),
|
myVpnBroadcastAddrsTable: new(bart.Lite),
|
||||||
@@ -384,10 +413,14 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
|||||||
|
|
||||||
v1hs, err := v1.MarshalForHandshakes()
|
v1hs, err := v1.MarshalForHandshakes()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("error marshalling certificate for handshake: %w", err)
|
return nil, fmt.Errorf("error marshalling v1 certificate for handshake: %w", err)
|
||||||
|
}
|
||||||
|
ncs, err := newCipherSuite(v1.Curve(), pkcs11backed, cipher)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
cs.v1Cert = v1
|
cs.v1Cert = v1
|
||||||
cs.v1HandshakeBytes = v1hs
|
cs.v1Credential = handshake.NewCredential(v1, v1hs, privateKey, ncs)
|
||||||
|
|
||||||
if cs.initiatingVersion == 0 {
|
if cs.initiatingVersion == 0 {
|
||||||
cs.initiatingVersion = cert.Version1
|
cs.initiatingVersion = cert.Version1
|
||||||
@@ -405,10 +438,14 @@ func newCertState(dv cert.Version, v1, v2 cert.Certificate, pkcs11backed bool, p
|
|||||||
|
|
||||||
v2hs, err := v2.MarshalForHandshakes()
|
v2hs, err := v2.MarshalForHandshakes()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("error marshalling certificate for handshake: %w", err)
|
return nil, fmt.Errorf("error marshalling v2 certificate for handshake: %w", err)
|
||||||
|
}
|
||||||
|
ncs, err := newCipherSuite(v2.Curve(), pkcs11backed, cipher)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
cs.v2Cert = v2
|
cs.v2Cert = v2
|
||||||
cs.v2HandshakeBytes = v2hs
|
cs.v2Credential = handshake.NewCredential(v2, v2hs, privateKey, ncs)
|
||||||
|
|
||||||
if cs.initiatingVersion == 0 {
|
if cs.initiatingVersion == 0 {
|
||||||
cs.initiatingVersion = cert.Version2
|
cs.initiatingVersion = cert.Version2
|
||||||
|
|||||||
+168
-7
@@ -18,6 +18,7 @@ type relayManager struct {
|
|||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
hostmap *HostMap
|
hostmap *HostMap
|
||||||
amRelay atomic.Bool
|
amRelay atomic.Bool
|
||||||
|
useRelays atomic.Bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewRelayManager(ctx context.Context, l *slog.Logger, hostmap *HostMap, c *config.C) *relayManager {
|
func NewRelayManager(ctx context.Context, l *slog.Logger, hostmap *HostMap, c *config.C) *relayManager {
|
||||||
@@ -36,8 +37,10 @@ func NewRelayManager(ctx context.Context, l *slog.Logger, hostmap *HostMap, c *c
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (rm *relayManager) reload(c *config.C, initial bool) error {
|
func (rm *relayManager) reload(c *config.C, initial bool) error {
|
||||||
if initial || c.HasChanged("relay.am_relay") {
|
if initial || c.HasChanged("relay.am_relay") || c.HasChanged("relay.use_relays") {
|
||||||
rm.setAmRelay(c.GetBool("relay.am_relay", false))
|
amRelay := c.GetBool("relay.am_relay", false)
|
||||||
|
rm.amRelay.Store(amRelay)
|
||||||
|
rm.useRelays.Store(c.GetBool("relay.use_relays", true) && !amRelay)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -46,8 +49,160 @@ func (rm *relayManager) GetAmRelay() bool {
|
|||||||
return rm.amRelay.Load()
|
return rm.amRelay.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (rm *relayManager) setAmRelay(v bool) {
|
func (rm *relayManager) GetUseRelays() bool {
|
||||||
rm.amRelay.Store(v)
|
return rm.useRelays.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// StartRelays drives the relay-establishment side of an outbound handshake attempt.
|
||||||
|
// For each candidate relay it either kicks off a handshake to the relay, sends a CreateRelayRequest, retransmits
|
||||||
|
// one that may have been lost, or, once the relay is Established, forwards the in-progress
|
||||||
|
// stage 0 handshake packet for vpnIp through it.
|
||||||
|
func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hostinfo *HostInfo, stage0 []byte) {
|
||||||
|
if !rm.GetUseRelays() || len(hostinfo.remotes.relays) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hostinfo.logger(rm.l).Info("Attempt to relay through hosts", "relays", hostinfo.remotes.relays)
|
||||||
|
// One WireBuffer for the whole relay-fanout loop.
|
||||||
|
buf := f.bufAlloc.Acquire()
|
||||||
|
defer f.bufAlloc.Release(buf)
|
||||||
|
// Send a RelayRequest to all known Relay IP's
|
||||||
|
for _, relay := range hostinfo.remotes.relays {
|
||||||
|
// Don't relay through the host I'm trying to connect to
|
||||||
|
if relay == vpnIp {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Don't relay to myself
|
||||||
|
if f.myVpnAddrsTable.Contains(relay) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
relayHostInfo := rm.hostmap.QueryVpnAddr(relay)
|
||||||
|
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
||||||
|
hostinfo.logger(rm.l).Info("Establish tunnel to relay target", "relay", relay.String())
|
||||||
|
f.Handshake(relay)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Check the relay HostInfo to see if we already established a relay through
|
||||||
|
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
||||||
|
if !ok {
|
||||||
|
// No relays exist or requested yet.
|
||||||
|
if relayHostInfo.remote.IsValid() {
|
||||||
|
idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(rm.l).Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: idx,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(rm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(rm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, buf)
|
||||||
|
rm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", idx,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
switch existingRelay.State {
|
||||||
|
case Established:
|
||||||
|
hostinfo.logger(rm.l).Info("Send handshake via relay", "relay", relay.String())
|
||||||
|
f.SendVia(relayHostInfo, existingRelay, stage0, buf)
|
||||||
|
case Disestablished:
|
||||||
|
// Mark this relay as 'requested'
|
||||||
|
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||||
|
fallthrough
|
||||||
|
case Requested:
|
||||||
|
hostinfo.logger(rm.l).Info("Re-send CreateRelay request", "relay", relay.String())
|
||||||
|
// Re-send the CreateRelay request, in case the previous one was lost.
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: existingRelay.LocalIndex,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(rm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(rm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(rm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
// This must send over the hostinfo, not over hm.Hosts[ip]
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, buf)
|
||||||
|
rm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", existingRelay.LocalIndex,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
case PeerRequested:
|
||||||
|
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
||||||
|
fallthrough
|
||||||
|
default:
|
||||||
|
hostinfo.logger(rm.l).Error("Relay unexpected state",
|
||||||
|
"vpnIp", vpnIp,
|
||||||
|
"state", existingRelay.State,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddRelay finds an available relay index on the hostmap, and associates the relay info with it.
|
// AddRelay finds an available relay index on the hostmap, and associates the relay info with it.
|
||||||
@@ -216,7 +371,9 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
rm.l.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
rm.l.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
rm.l.Info("send CreateRelayResponse",
|
rm.l.Info("send CreateRelayResponse",
|
||||||
"relayFrom", resp.RelayFromAddr,
|
"relayFrom", resp.RelayFromAddr,
|
||||||
"relayTo", resp.RelayToAddr,
|
"relayTo", resp.RelayToAddr,
|
||||||
@@ -316,7 +473,9 @@ func (rm *relayManager) handleCreateRelayRequest(v cert.Version, h *HostInfo, f
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
logMsg.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
logMsg.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
f.SendMessageToHostInfo(header.Control, 0, h, msg, make([]byte, 12), make([]byte, mtu))
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, h, msg, buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
rm.l.Info("send CreateRelayResponse",
|
rm.l.Info("send CreateRelayResponse",
|
||||||
"relayFrom", from,
|
"relayFrom", from,
|
||||||
"relayTo", target,
|
"relayTo", target,
|
||||||
@@ -386,7 +545,9 @@ func (rm *relayManager) handleCreateRelayRequest(v cert.Version, h *HostInfo, f
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
logMsg.Error("relayManager Failed to marshal Control message to create relay", "error", err)
|
logMsg.Error("relayManager Failed to marshal Control message to create relay", "error", err)
|
||||||
} else {
|
} else {
|
||||||
f.SendMessageToHostInfo(header.Control, 0, peer, msg, make([]byte, 12), make([]byte, mtu))
|
buf := f.bufAlloc.Acquire()
|
||||||
|
f.SendMessageToHostInfo(header.Control, 0, peer, msg, buf)
|
||||||
|
f.bufAlloc.Release(buf)
|
||||||
rm.l.Info("send CreateRelayRequest",
|
rm.l.Info("send CreateRelayRequest",
|
||||||
"relayFrom", h.vpnAddrs[0],
|
"relayFrom", h.vpnAddrs[0],
|
||||||
"relayTo", target,
|
"relayTo", target,
|
||||||
|
|||||||
@@ -632,15 +632,9 @@ func sshCloseTunnel(ifce *Interface, fs any, a []string, w sshd.StringWriter) er
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !flags.LocalOnly {
|
if !flags.LocalOnly {
|
||||||
ifce.send(
|
buf := ifce.bufAlloc.Acquire()
|
||||||
header.CloseTunnel,
|
ifce.send(header.CloseTunnel, 0, hostInfo.ConnectionState, hostInfo, []byte{}, buf)
|
||||||
0,
|
ifce.bufAlloc.Release(buf)
|
||||||
hostInfo.ConnectionState,
|
|
||||||
hostInfo,
|
|
||||||
[]byte{},
|
|
||||||
make([]byte, 12, 12),
|
|
||||||
make([]byte, mtu),
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
ifce.closeTunnel(hostInfo)
|
ifce.closeTunnel(hostInfo)
|
||||||
|
|||||||
+50
-13
@@ -21,17 +21,48 @@ type Packet struct {
|
|||||||
Data []byte
|
Data []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Copy returns a fresh *Packet (from the freelist) with a duplicate Data buffer.
|
||||||
func (u *Packet) Copy() *Packet {
|
func (u *Packet) Copy() *Packet {
|
||||||
n := &Packet{
|
n := acquirePacket()
|
||||||
To: u.To,
|
n.To = u.To
|
||||||
From: u.From,
|
n.From = u.From
|
||||||
Data: make([]byte, len(u.Data)),
|
if cap(n.Data) < len(u.Data) {
|
||||||
|
n.Data = make([]byte, len(u.Data))
|
||||||
|
} else {
|
||||||
|
n.Data = n.Data[:len(u.Data)]
|
||||||
}
|
}
|
||||||
|
|
||||||
copy(n.Data, u.Data)
|
copy(n.Data, u.Data)
|
||||||
return n
|
return n
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Release returns p to the harness packet freelist.
|
||||||
|
// Callers that pull a *Packet from Get / TxPackets must Release when done.
|
||||||
|
// Channel-backed instead of sync.Pool because sync.Pool's per-P caches drain badly under cross-goroutine Get/Put,
|
||||||
|
// and putting a []byte in a Pool escapes the slice header to heap.
|
||||||
|
func (p *Packet) Release() {
|
||||||
|
if p == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p.Data = p.Data[:0]
|
||||||
|
select {
|
||||||
|
case packetFreelist <- p:
|
||||||
|
default:
|
||||||
|
// Freelist full; drop the *Packet for the GC.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// packetFreelist retains *Packet structs (and their backing Data arrays) so steady-state allocation drops to zero.
|
||||||
|
var packetFreelist = make(chan *Packet, 64)
|
||||||
|
|
||||||
|
func acquirePacket() *Packet {
|
||||||
|
select {
|
||||||
|
case p := <-packetFreelist:
|
||||||
|
return p
|
||||||
|
default:
|
||||||
|
return &Packet{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type TesterConn struct {
|
type TesterConn struct {
|
||||||
Addr netip.AddrPort
|
Addr netip.AddrPort
|
||||||
|
|
||||||
@@ -64,13 +95,15 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, _ bool, _ int) (Conn,
|
|||||||
// this is an encrypted packet or a handshake message in most cases
|
// this is an encrypted packet or a handshake message in most cases
|
||||||
// packets were transmitted from another nebula node, you can send them with Tun.Send
|
// packets were transmitted from another nebula node, you can send them with Tun.Send
|
||||||
func (u *TesterConn) Send(packet *Packet) {
|
func (u *TesterConn) Send(packet *Packet) {
|
||||||
h := &header.H{}
|
if u.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
// Parse the header only under debug logging, otherwise the
|
||||||
|
// allocation would show up in every Send call.
|
||||||
|
var h header.H
|
||||||
if err := h.Parse(packet.Data); err != nil {
|
if err := h.Parse(packet.Data); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if u.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
u.l.Debug("UDP receiving injected packet",
|
u.l.Debug("UDP receiving injected packet",
|
||||||
"header", h,
|
"header", &h,
|
||||||
"udpAddr", packet.From,
|
"udpAddr", packet.From,
|
||||||
"dataLen", len(packet.Data),
|
"dataLen", len(packet.Data),
|
||||||
)
|
)
|
||||||
@@ -107,15 +140,18 @@ func (u *TesterConn) Get(block bool) *Packet {
|
|||||||
//********************************************************************************************************************//
|
//********************************************************************************************************************//
|
||||||
|
|
||||||
func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
||||||
p := &Packet{
|
p := acquirePacket()
|
||||||
Data: make([]byte, len(b), len(b)),
|
if cap(p.Data) < len(b) {
|
||||||
From: u.Addr,
|
p.Data = make([]byte, len(b))
|
||||||
To: addr,
|
} else {
|
||||||
|
p.Data = p.Data[:len(b)]
|
||||||
}
|
}
|
||||||
|
|
||||||
copy(p.Data, b)
|
copy(p.Data, b)
|
||||||
|
p.From = u.Addr
|
||||||
|
p.To = addr
|
||||||
select {
|
select {
|
||||||
case <-u.done:
|
case <-u.done:
|
||||||
|
p.Release()
|
||||||
return io.ErrClosedPipe
|
return io.ErrClosedPipe
|
||||||
case u.TxPackets <- p:
|
case u.TxPackets <- p:
|
||||||
return nil
|
return nil
|
||||||
@@ -129,6 +165,7 @@ func (u *TesterConn) ListenOut(r EncReader) error {
|
|||||||
return os.ErrClosed
|
return os.ErrClosed
|
||||||
case p := <-u.RxPackets:
|
case p := <-u.RxPackets:
|
||||||
r(p.From, p.Data)
|
r(p.From, p.Data)
|
||||||
|
p.Release()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+255
@@ -0,0 +1,255 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/firewall"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
)
|
||||||
|
|
||||||
|
// WireBuffer is the per-goroutine working set for processing one IP packet
|
||||||
|
// through the data plane. It owns:
|
||||||
|
//
|
||||||
|
// - The IP-payload byte buffer used to hold the current inbound or
|
||||||
|
// outbound packet, with prefixLen bytes of slack at the front for
|
||||||
|
// the BSD AF_INET protocol-family marker.
|
||||||
|
// - The fwPacket scratch parsed by newPacket().
|
||||||
|
// - The 12-byte AEAD nonce scratch.
|
||||||
|
// - The header.H parse target used by the receive path.
|
||||||
|
// - An mtu-sized wire-output scratch for sendNoMetrics and for building
|
||||||
|
// reject packets.
|
||||||
|
//
|
||||||
|
// One WireBuffer is allocated per data-plane goroutine (listenIn for the
|
||||||
|
// TUN-side, listenOut for the UDP-side) and reused for every packet. No
|
||||||
|
// per-packet allocation. Future GRO/GSO/TSO and reliable-transport work
|
||||||
|
// will likely extend this to carry batch state and fragment metadata.
|
||||||
|
//
|
||||||
|
// The TUN protocol-family prefix is handled here, not in the overlay
|
||||||
|
// package. On BSDs the kernel writes the 4-byte marker into the slack on
|
||||||
|
// read, and we stamp it into the slack before write. On linux/windows
|
||||||
|
// /userspace devices prefixLen is 0 and the slack is empty.
|
||||||
|
type WireBuffer struct {
|
||||||
|
// FwPacket is the parsed IP packet metadata (5-tuple, fragment flags,
|
||||||
|
// etc.) populated by newPacket().
|
||||||
|
FwPacket *firewall.Packet
|
||||||
|
// NB is a 12-byte scratch the AEAD uses for the nonce; reused so we
|
||||||
|
// don't allocate one per encrypt/decrypt.
|
||||||
|
NB []byte
|
||||||
|
// H is the parse target for inbound nebula headers. Receive path only.
|
||||||
|
H *header.H
|
||||||
|
// Out is an mtu-sized wire-output scratch passed to sendNoMetrics and
|
||||||
|
// rejectInside / rejectOutside. Sized to fit any single wire packet.
|
||||||
|
Out []byte
|
||||||
|
|
||||||
|
// ip is the IP-payload region: a slice of len 0, cap linkMTU sliced
|
||||||
|
// from raw at offset prefixLen. The current packet (if any) is
|
||||||
|
// ip[:bodyN]. The TUN prefix slack lives at raw[0:prefixLen] just
|
||||||
|
// before ip.
|
||||||
|
ip []byte
|
||||||
|
// raw is the backing slab. Layout:
|
||||||
|
// [prefixLen bytes prefix slack | linkMTU bytes IP region | outSize bytes Out scratch]
|
||||||
|
// Holding it lets ReadIPFromTUN / WriteIPToTUN address the slack
|
||||||
|
// region directly.
|
||||||
|
raw []byte
|
||||||
|
prefixLen int
|
||||||
|
bodyN int
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewWireBuffer returns a buffer sized to hold any single IP packet up to
|
||||||
|
// linkMTU, plus a disjoint wire-output scratch sliced from the same backing
|
||||||
|
// slab (the AEAD's Seal contract requires plaintext and dst not to partially
|
||||||
|
// overlap, and keeping them in one slab gives a single allocation per
|
||||||
|
// goroutine). Out is sized for the relay worst case
|
||||||
|
// (linkMTU + 2*header.Len + 2*AEADOverhead).
|
||||||
|
//
|
||||||
|
// prefixLen is the number of bytes the destination tun device prepends/
|
||||||
|
// expects on each IP packet (overlay.Device.TunPrefixLen). On BSDs this
|
||||||
|
// is 4 (AF_INET marker); on linux/windows/userspace devices it is 0.
|
||||||
|
func NewWireBuffer(linkMTU, prefixLen int) *WireBuffer {
|
||||||
|
outSize := linkMTU + 2*header.Len + 2*AEADOverhead
|
||||||
|
raw := make([]byte, prefixLen+linkMTU+outSize)
|
||||||
|
outStart := prefixLen + linkMTU
|
||||||
|
return &WireBuffer{
|
||||||
|
FwPacket: &firewall.Packet{},
|
||||||
|
NB: make([]byte, NonceSize),
|
||||||
|
H: &header.H{},
|
||||||
|
Out: raw[outStart : outStart : outStart+outSize],
|
||||||
|
ip: raw[prefixLen:prefixLen:outStart],
|
||||||
|
raw: raw,
|
||||||
|
prefixLen: prefixLen,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset clears the body-length record so the buffer is ready for another
|
||||||
|
// recv (e.g. relay-receive recursion before a nested decrypt).
|
||||||
|
func (b *WireBuffer) Reset() { b.bodyN = 0 }
|
||||||
|
|
||||||
|
// IPPacket returns the IP packet currently held in the payload region (after
|
||||||
|
// a successful ReadIPFromTUN or DecryptDatagram). The slice aliases the
|
||||||
|
// buffer; do not retain past the next operation.
|
||||||
|
func (b *WireBuffer) IPPacket() []byte {
|
||||||
|
return b.ip[:b.bodyN]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Seal stamps a nebula header at the front of buf.Out and AEAD-seals p as the
|
||||||
|
// payload, treating the header as additional authenticated data. The lock
|
||||||
|
// scope around counter increment + encrypt matches what goboring AESGCMTLS
|
||||||
|
// requires; non-boring builds skip the lock.
|
||||||
|
//
|
||||||
|
// Returns the wire bytes (header || ciphertext || tag), aliased to buf.Out.
|
||||||
|
// The slice is invalidated by the next Seal* call on this buffer.
|
||||||
|
func (b *WireBuffer) Seal(ci *ConnectionState, t header.MessageType, st header.MessageSubType, remoteIndex uint32, p []byte) ([]byte, error) {
|
||||||
|
return b.sealInto(b.Out[:cap(b.Out)], ci, t, st, remoteIndex, p)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SealForRelay is like Seal but reserves header.Len bytes of slack at the front
|
||||||
|
// of buf.Out for an outer relay header. The inner header + ciphertext lands at
|
||||||
|
// offset header.Len so a follow-up SealRelayInPlace can stamp the outer header
|
||||||
|
// without copying. Use this when the caller may need to wrap the result in a
|
||||||
|
// relay envelope after the fact.
|
||||||
|
func (b *WireBuffer) SealForRelay(ci *ConnectionState, t header.MessageType, st header.MessageSubType, remoteIndex uint32, p []byte) ([]byte, error) {
|
||||||
|
return b.sealInto(b.Out[header.Len:cap(b.Out)], ci, t, st, remoteIndex, p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *WireBuffer) sealInto(out []byte, ci *ConnectionState, t header.MessageType, st header.MessageSubType, remoteIndex uint32, p []byte) ([]byte, error) {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
out = header.Encode(out, header.Version, t, st, remoteIndex, c)
|
||||||
|
out, err := ci.eKey.EncryptDanger(out, out, p, c, b.NB)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
return out, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// SealRelayInPlace wraps an inner message that is already staged at
|
||||||
|
// buf.Out[header.Len:header.Len+innerLen] (either from a SealForRelay encrypt
|
||||||
|
// or from a copy via the SendVia entry point). It stamps the outer relay
|
||||||
|
// header into buf.Out[:header.Len] and AAD-only seals over the entire region,
|
||||||
|
// producing the wire bytes for the relay tunnel.
|
||||||
|
//
|
||||||
|
// Returns the wire bytes aliased to buf.Out; invalidated by the next Seal*
|
||||||
|
// call on this buffer.
|
||||||
|
func (b *WireBuffer) SealRelayInPlace(ci *ConnectionState, remoteIndex uint32, innerLen int) ([]byte, error) {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
out := b.Out[:cap(b.Out)]
|
||||||
|
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, remoteIndex, c)
|
||||||
|
out = out[:header.Len+innerLen]
|
||||||
|
out, err := ci.eKey.EncryptDanger(out, out, nil, c, b.NB)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
return out, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// StageRelayInner copies ad into the inner-payload slot at buf.Out[header.Len:]
|
||||||
|
// so SealRelayInPlace can wrap it on the next call. Used by SendVia when ad
|
||||||
|
// did not come from a prior SealForRelay (e.g. a handshake message being
|
||||||
|
// forwarded through a relay tunnel without our own encryption).
|
||||||
|
func (b *WireBuffer) StageRelayInner(ad []byte) int {
|
||||||
|
return copy(b.Out[header.Len:cap(b.Out)], ad)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadIPFromTUN reads one IP packet from r into the payload region and
|
||||||
|
// updates bodyN. On BSDs the kernel writes its 4-byte protocol-family
|
||||||
|
// marker into the slack at raw[0:prefixLen] and the IP packet at
|
||||||
|
// raw[prefixLen:prefixLen+n]; we hand it the slack-prefixed slice so
|
||||||
|
// the kernel can do this in one syscall with no copy. On linux/windows/
|
||||||
|
// userspace devices prefixLen is 0 and the slack is empty.
|
||||||
|
func (b *WireBuffer) ReadIPFromTUN(r io.Reader) (int, error) {
|
||||||
|
n, err := r.Read(b.raw[:b.prefixLen+cap(b.ip)])
|
||||||
|
if err != nil {
|
||||||
|
b.bodyN = 0
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
if n < b.prefixLen {
|
||||||
|
b.bodyN = 0
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
b.bodyN = n - b.prefixLen
|
||||||
|
return b.bodyN, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteIPToTUN writes the IP packet currently in the payload region to w.
|
||||||
|
// On BSDs we stamp the protocol-family marker into the slack at
|
||||||
|
// raw[0:prefixLen] in place and write the entire slack+IP region in a
|
||||||
|
// single syscall, so the kernel sees [marker][ip] back to back without a
|
||||||
|
// userspace copy. On linux/windows/userspace devices the slack is empty
|
||||||
|
// and we just write the IP region.
|
||||||
|
func (b *WireBuffer) WriteIPToTUN(w io.Writer) (int, error) {
|
||||||
|
out := b.raw[:b.prefixLen+b.bodyN]
|
||||||
|
if b.prefixLen > 0 {
|
||||||
|
if err := overlay.StampTunPrefix(out); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return w.Write(out)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DecryptDatagram decrypts an inbound UDP packet into the payload region.
|
||||||
|
func (b *WireBuffer) DecryptDatagram(ci *ConnectionState, packet []byte, mc uint64) error {
|
||||||
|
dst, err := ci.dKey.DecryptDanger(b.ip[:0], packet[:header.Len], packet[header.Len:], mc, b.NB)
|
||||||
|
if err != nil {
|
||||||
|
b.bodyN = 0
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
b.bodyN = len(dst)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DecryptForHandler decrypts an inbound UDP packet (lighthouse, test,
|
||||||
|
// control, close-tunnel) into the payload region and returns the plaintext
|
||||||
|
// slice for the in-process handler. Returned slice aliases the buffer.
|
||||||
|
func (b *WireBuffer) DecryptForHandler(ci *ConnectionState, packet []byte, mc uint64) ([]byte, error) {
|
||||||
|
dst, err := ci.dKey.DecryptDanger(b.ip[:0], packet[:header.Len], packet[header.Len:], mc, b.NB)
|
||||||
|
if err != nil {
|
||||||
|
b.bodyN = 0
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
b.bodyN = len(dst)
|
||||||
|
return dst, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WireBufferAllocator hands out reusable WireBuffers for cold callers that
|
||||||
|
// don't own a long-lived per-goroutine buffer (control plane, relay manager,
|
||||||
|
// connection manager teardown, etc.). Hot-path goroutines hold their own
|
||||||
|
// buffer for the life of the goroutine and don't need to acquire one.
|
||||||
|
type WireBufferAllocator interface {
|
||||||
|
Acquire() *WireBuffer
|
||||||
|
Release(*WireBuffer)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wireBufferPool is a sync.Pool-backed WireBufferAllocator. The pool is
|
||||||
|
// keyed off a single linkMTU and prefixLen; cold callers send across the
|
||||||
|
// data-plane mtu and target the same Device, so we size the pool's
|
||||||
|
// buffers the same way.
|
||||||
|
type wireBufferPool struct {
|
||||||
|
pool sync.Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewWireBufferPool(linkMTU, prefixLen int) *wireBufferPool {
|
||||||
|
return &wireBufferPool{
|
||||||
|
pool: sync.Pool{
|
||||||
|
New: func() any {
|
||||||
|
return NewWireBuffer(linkMTU, prefixLen)
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *wireBufferPool) Acquire() *WireBuffer {
|
||||||
|
return p.pool.Get().(*WireBuffer)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *wireBufferPool) Release(b *WireBuffer) {
|
||||||
|
b.Reset()
|
||||||
|
p.pool.Put(b)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user