diff --git a/peer/app.go b/peer/app.go index 5102beb..cd021f7 100644 --- a/peer/app.go +++ b/peer/app.go @@ -68,7 +68,7 @@ type App struct { selfV6 netip.AddrPort // Monotonically increasing ID for outbound pings (client role only) - nextPingID int64 + nextPingID int64 // TODO: Remove // Event channels fed by background goroutines hubAddCh <-chan HubPeer @@ -104,7 +104,6 @@ func (a *App) Run() error { } } - func (a *App) onShutdown() error { return wginterface.Delete(a.dev.Name()) } diff --git a/peer/app_test.go b/peer/app_test.go index 8d909d7..e49a35d 100644 --- a/peer/app_test.go +++ b/peer/app_test.go @@ -3,10 +3,28 @@ package peer import ( "net/netip" "testing" + "time" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" ) +// addRelayPeer adds a public relay peer and marks it Up so it satisfies +// CanRelay. It does not set a.relay — callers do that explicitly. +func addRelayPeer(t *testing.T, a *App, vpnIP string, ep netip.AddrPort) *Peer { + t.Helper() + key := mustKey(t) + a.onAddPeer(HubPeer{ + PubKey: key, + VPNIP: netip.MustParseAddr(vpnIP), + IsPublic: true, + IsRelay: true, + EndpointV4: ep, + }) + p := a.peersByKey[key] + p.wgPeer.LastHandshakeTime = time.Now() + return p +} + // newTestApp returns a minimal App wired to a fakeWGDevice. // vpnIP is the local VPN address (e.g. "10.0.0.1"). // isPublic / isRelay describe the local node's role. diff --git a/peer/control/ping.go b/peer/control/ping.go new file mode 100644 index 0000000..1809e31 --- /dev/null +++ b/peer/control/ping.go @@ -0,0 +1,76 @@ +// Package control implements the VPN-internal peer control protocol. +// Peers exchange Ping packets over UDP on the VPN control port to maintain +// liveness and discover external endpoints for direct connection attempts. +package control + +import ( + "encoding/binary" + "fmt" + "net/netip" +) + +const ( + version = 1 + Size = 59 // 1 version + 8 ID + 8 PingTS + 6 SrcV4 + 18 SrcV6 + 18 Dst +) + +// Ping is the single control packet type exchanged between VPN peers. +// +// In each peer pair, the peer with the lower VPN IP is the client: it sets ID +// and PingTS and sends pings on a timer. The server echoes ID and PingTS back +// in its response, allowing the client to compute RTT = now - PingTS. +// +// Both client and server populate SrcV4, SrcV6, and Dst on every packet so +// endpoint information flows in both directions. +// +// Dst is the recipient's external endpoint as observed by the sender from the +// WireGuard handshake source. Zero if the sender has not observed a handshake +// from the recipient. +type Ping struct { + ID int64 // Client ping ID. + PingTS int64 // Client ping send time in nanoseconds. + SrcV4 netip.AddrPort // Sender's discovered IPv4 address and port. + SrcV6 netip.AddrPort // Sender's discovered IPv6 address and port. + Dst netip.AddrPort +} + +// Marshal encodes p into a fixed-size 59-byte array. +func (p Ping) Marshal() [Size]byte { + var buf [Size]byte + buf[0] = version + binary.BigEndian.PutUint64(buf[1:9], uint64(p.ID)) + binary.BigEndian.PutUint64(buf[9:17], uint64(p.PingTS)) + if p.SrcV4.IsValid() { + a4 := p.SrcV4.Addr().As4() + copy(buf[17:21], a4[:]) + binary.BigEndian.PutUint16(buf[21:23], p.SrcV4.Port()) + } + a16 := p.SrcV6.Addr().As16() + copy(buf[23:39], a16[:]) + binary.BigEndian.PutUint16(buf[39:41], p.SrcV6.Port()) + a16 = p.Dst.Addr().As16() + copy(buf[41:57], a16[:]) + binary.BigEndian.PutUint16(buf[57:59], p.Dst.Port()) + return buf +} + +// Unmarshal decodes a Ping from a fixed-size 59-byte array. +func Unmarshal(buf [Size]byte) (Ping, error) { + if buf[0] != version { + return Ping{}, fmt.Errorf("unknown ping version %d", buf[0]) + } + p := Ping{ + ID: int64(binary.BigEndian.Uint64(buf[1:9])), + PingTS: int64(binary.BigEndian.Uint64(buf[9:17])), + } + if addr := netip.AddrFrom4([4]byte(buf[17:21])); !addr.IsUnspecified() { + p.SrcV4 = netip.AddrPortFrom(addr, binary.BigEndian.Uint16(buf[21:23])) + } + if addr := netip.AddrFrom16([16]byte(buf[23:39])); !addr.IsUnspecified() { + p.SrcV6 = netip.AddrPortFrom(addr, binary.BigEndian.Uint16(buf[39:41])) + } + if addr := netip.AddrFrom16([16]byte(buf[41:57])).Unmap(); !addr.IsUnspecified() { + p.Dst = netip.AddrPortFrom(addr, binary.BigEndian.Uint16(buf[57:59])) + } + return p, nil +} diff --git a/peer/control/ping_test.go b/peer/control/ping_test.go new file mode 100644 index 0000000..5c9696a --- /dev/null +++ b/peer/control/ping_test.go @@ -0,0 +1,109 @@ +package control_test + +import ( + "net/netip" + "testing" + + "vppn/peer/control" +) + +func TestRoundTrip(t *testing.T) { + cases := []struct { + name string + ping control.Ping + }{ + { + name: "zero", + ping: control.Ping{}, + }, + { + name: "client ping", + ping: control.Ping{ + ID: 42, + PingTS: 1234567890, + SrcV4: netip.MustParseAddrPort("1.2.3.4:51820"), + Dst: netip.MustParseAddrPort("5.6.7.8:51820"), + }, + }, + { + name: "server response", + ping: control.Ping{ + ID: 42, + PingTS: 1234567890, + SrcV4: netip.MustParseAddrPort("5.6.7.8:51820"), + Dst: netip.MustParseAddrPort("1.2.3.4:9999"), + }, + }, + { + name: "IPv6 only", + ping: control.Ping{ + ID: 1, + PingTS: 999, + SrcV6: netip.MustParseAddrPort("[2001:db8::1]:51820"), + Dst: netip.MustParseAddrPort("[2001:db8::2]:51820"), + }, + }, + { + name: "dual stack", + ping: control.Ping{ + ID: 7, + PingTS: 555, + SrcV4: netip.MustParseAddrPort("1.2.3.4:51820"), + SrcV6: netip.MustParseAddrPort("[2001:db8::1]:51820"), + Dst: netip.MustParseAddrPort("5.6.7.8:9999"), + }, + }, + { + name: "no src known", + ping: control.Ping{ + ID: 3, + Dst: netip.MustParseAddrPort("5.6.7.8:51820"), + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + buf := tc.ping.Marshal() + got, err := control.Unmarshal(buf) + if err != nil { + t.Fatalf("Unmarshal: %v", err) + } + if got != tc.ping { + t.Fatalf("round-trip mismatch:\n got %+v\n want %+v", got, tc.ping) + } + }) + } +} + +func TestUnmarshalBadVersion(t *testing.T) { + var buf [control.Size]byte + buf[0] = 99 + if _, err := control.Unmarshal(buf); err == nil { + t.Fatal("expected error for unknown version, got nil") + } +} + +func TestZeroEncoding(t *testing.T) { + buf := (control.Ping{}).Marshal() + for i, b := range buf { + if i == 0 { + continue // version byte + } + if b != 0 { + t.Fatalf("expected zero encoding at byte %d, got %d", i, b) + } + } +} + +func TestRoleFor(t *testing.T) { + lo := netip.MustParseAddr("10.0.0.1") + hi := netip.MustParseAddr("10.0.0.2") + + if control.RoleFor(lo, hi) != control.Client { + t.Error("lower IP should be client") + } + if control.RoleFor(hi, lo) != control.Server { + t.Error("higher IP should be server") + } +} diff --git a/peer/control/role.go b/peer/control/role.go new file mode 100644 index 0000000..8dbd2cd --- /dev/null +++ b/peer/control/role.go @@ -0,0 +1,22 @@ +package control + +import "net/netip" + +// Role identifies a peer's role in a ping exchange with a specific remote peer. +type Role string + +const ( + // Client initiates pings and measures RTT. + Client Role = "CLIENT" + // Server responds to pings. + Server Role = "SERVER" +) + +// RoleFor returns the Role of local relative to remote. +// The peer with the lower VPN IP is the client. +func RoleFor(local, remote netip.Addr) Role { + if local.Compare(remote) < 0 { + return Client + } + return Server +} diff --git a/peer/device.go b/peer/device.go index 5627d69..cc2d8c0 100644 --- a/peer/device.go +++ b/peer/device.go @@ -15,34 +15,38 @@ func (a *App) devPeers() []wgtypes.Peer { return peers } +func (a *App) devAddPeer(p *Peer) { + if err := a.dev.AddPeer(p.PubKey()); err != nil { + log.Fatalf("Failed to add peer %v: %v", p.VPNIP, err) + } +} + func (a *App) devAddDirect(p *Peer, endpoint netip.AddrPort) { - if err := a.dev.AddDirect(p.PubKey, endpoint, p.VPNIP); err != nil { + if err := a.dev.AddDirect(p.PubKey(), endpoint, p.VPNIP); err != nil { log.Fatalf("Failed to add peer %v: %v", p.VPNIP, err) } } func (a *App) devSetRelay(p *Peer, endpoint netip.AddrPort) { - if err := a.dev.SetRelay(p.PubKey, endpoint, a.vpnNet); err != nil { + if err := a.dev.SetRelay(p.PubKey(), endpoint, a.vpnNet); err != nil { log.Fatalf("Failed to add relay %v: %v", p.VPNIP, err) } } func (a *App) devPromote(p *Peer) { - if err := a.dev.Promote(p.PubKey, p.VPNIP); err != nil { - log.Fatalf("Failed to promot peer %v: %v", p.VPNIP, err) + if err := a.dev.Promote(p.PubKey(), p.VPNIP); err != nil { + log.Fatalf("Failed to promote peer %v: %v", p.VPNIP, err) } } func (a *App) devAddProbe(p *Peer, endpoint netip.AddrPort) { - if err := a.dev.AddProbe(p.PubKey, endpoint); err != nil { + if err := a.dev.AddProbe(p.PubKey(), endpoint); err != nil { log.Fatalf("Failed to add probe %v: %v", p.VPNIP, err) } } func (a *App) devRemove(p *Peer) { - if p.State != StateRelayed { - if err := a.dev.RemovePeer(p.PubKey); err != nil { - log.Fatalf("Failed to remove peer %v: %v", p.VPNIP, err) - } + if err := a.dev.RemovePeer(p.PubKey()); err != nil { + log.Fatalf("Failed to remove peer %v: %v", p.VPNIP, err) } } diff --git a/peer/fake_wgdevice_test.go b/peer/fake_wgdevice_test.go index b2651fa..a21cd54 100644 --- a/peer/fake_wgdevice_test.go +++ b/peer/fake_wgdevice_test.go @@ -3,6 +3,7 @@ package peer import ( "net/netip" "sync" + "testing" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" ) @@ -41,6 +42,11 @@ func (f *fakeWGDevice) Peers() ([]wgtypes.Peer, error) { return out, nil } +func (f *fakeWGDevice) AddPeer(pubKey wgtypes.Key) error { + f.record(fakeCall{Method: "AddPeer", PubKey: pubKey}) + return nil +} + func (f *fakeWGDevice) AddDirect(pubKey wgtypes.Key, endpoint netip.AddrPort, vpnIP netip.Addr) error { f.record(fakeCall{Method: "AddDirect", PubKey: pubKey, Endpoint: endpoint, VPNiP: vpnIP}) return nil @@ -65,3 +71,53 @@ func (f *fakeWGDevice) RemovePeer(pubKey wgtypes.Key) error { f.record(fakeCall{Method: "RemovePeer", PubKey: pubKey}) return nil } + +// AssertNoCalls fails the test if any dev calls were recorded. +func (f *fakeWGDevice) AssertNoCalls(t *testing.T) { + t.Helper() + f.mu.Lock() + defer f.mu.Unlock() + if len(f.Calls) != 0 { + t.Fatalf("unexpected dev calls: %v", f.Calls) + } +} + +func (f *fakeWGDevice) AssertAddPeer(t *testing.T, i int, pubKey wgtypes.Key) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "AddPeer", PubKey: pubKey}) +} + +func (f *fakeWGDevice) AssertAddDirect(t *testing.T, i int, pubKey wgtypes.Key, endpoint netip.AddrPort, vpnIP netip.Addr) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "AddDirect", PubKey: pubKey, Endpoint: endpoint, VPNiP: vpnIP}) +} + +func (f *fakeWGDevice) AssertSetRelay(t *testing.T, i int, pubKey wgtypes.Key, endpoint netip.AddrPort, network netip.Prefix) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "SetRelay", PubKey: pubKey, Endpoint: endpoint, Network: network}) +} + +func (f *fakeWGDevice) AssertAddProbe(t *testing.T, i int, pubKey wgtypes.Key, endpoint netip.AddrPort) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "AddProbe", PubKey: pubKey, Endpoint: endpoint}) +} + +func (f *fakeWGDevice) AssertPromote(t *testing.T, i int, pubKey wgtypes.Key, vpnIP netip.Addr) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "Promote", PubKey: pubKey, VPNiP: vpnIP}) +} + +func (f *fakeWGDevice) AssertRemovePeer(t *testing.T, i int, pubKey wgtypes.Key) { + t.Helper() + f.assertCall(t, i, fakeCall{Method: "RemovePeer", PubKey: pubKey}) +} + +func (f *fakeWGDevice) assertCall(t *testing.T, i int, c fakeCall) { + t.Helper() + if len(f.Calls) <= i { + t.Fatalf("no call at index %d: %v", i, c) + } + if c != f.Calls[i] { + t.Fatalf("call[%d]: got %v, want %v", i, f.Calls[i], c) + } +} diff --git a/peer/interfaces.go b/peer/interfaces.go index be13ac5..7aceaa7 100644 --- a/peer/interfaces.go +++ b/peer/interfaces.go @@ -10,6 +10,7 @@ import ( type WGDevice interface { Name() string Peers() ([]wgtypes.Peer, error) + AddPeer(pubKey wgtypes.Key) error AddDirect(pubKey wgtypes.Key, endpoint netip.AddrPort, vpnIP netip.Addr) error SetRelay(pubKey wgtypes.Key, endpoint netip.AddrPort, network netip.Prefix) error AddProbe(pubKey wgtypes.Key, endpoint netip.AddrPort) error diff --git a/peer/on_hub.go b/peer/on_hub.go index 3c83d86..76b1b6e 100644 --- a/peer/on_hub.go +++ b/peer/on_hub.go @@ -10,19 +10,16 @@ import ( ) func (a *App) onAddPeer(p HubPeer) { - if _, exists := a.peersByKey[p.PubKey]; exists { - a.onRemovePeer(p.PubKey) - } + a.onRemovePeer(p.PubKey) peer := &Peer{ - PubKey: p.PubKey, + wgPeer: wgtypes.Peer{PublicKey: p.PubKey}, VPNIP: p.VPNIP, IsRelay: p.IsRelay, IsPublic: p.IsPublic, Endpoint4: p.EndpointV4, Endpoint6: p.EndpointV6, Role: roleFor(a.isPublic, a.vpnIP, p), - State: StateRelayed, } endpoint := peer.PreferredEndpoint() @@ -36,11 +33,10 @@ func (a *App) onAddPeer(p HubPeer) { a.peersByIP[peer.VPNIP] = peer if !peer.IsPublic { + a.devAddPeer(peer) return } - peer.WGEndpoint = endpoint - peer.State = StateDirect a.devAddDirect(peer, endpoint) } @@ -62,7 +58,7 @@ func (a *App) onRemovePeer(key wgtypes.Key) { // switchActiveRelay promotes the lowest-latency relay peer to active. func (a *App) switchActiveRelay() { if a.relay != nil { - a.devAddDirect(a.relay, a.relay.WGEndpoint) + a.devAddDirect(a.relay, a.relay.WGEndpoint()) a.relay = nil } @@ -81,10 +77,11 @@ func (a *App) switchActiveRelay() { return } - a.devSetRelay(best, best.WGEndpoint) + a.devSetRelay(best, best.WGEndpoint()) a.relay = best } +// TODO: Why not < ?? // betterRelay reports whether a is a better relay candidate than b. // Prefers lower RTT; treats zero RTT (no measurement yet) as worst case. func betterRelay(a, b *Peer) bool { diff --git a/peer/on_hub_test.go b/peer/on_hub_test.go new file mode 100644 index 0000000..d69ab5a --- /dev/null +++ b/peer/on_hub_test.go @@ -0,0 +1,315 @@ +package peer + +import ( + "net/netip" + "testing" + "time" + + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +func mustKey(t *testing.T) wgtypes.Key { + t.Helper() + k, err := wgtypes.GeneratePrivateKey() + if err != nil { + t.Fatalf("generate key: %v", err) + } + return k.PublicKey() +} + +func TestOnAddPeer(t *testing.T) { + ep1 := netip.MustParseAddrPort("1.2.3.4:51820") + ep2 := netip.MustParseAddrPort("5.6.7.8:51820") + peerVPNIP := netip.MustParseAddr("10.0.0.2") + + testCases := []struct { + name string + setup func(a *App, key wgtypes.Key) + peer func(key wgtypes.Key) HubPeer + check func(t *testing.T, a *App, dev *fakeWGDevice, key wgtypes.Key) + }{ + { + name: "non-public peer added in StateRelayed with no dev calls", + peer: func(k wgtypes.Key) HubPeer { + return HubPeer{PubKey: k, VPNIP: peerVPNIP} + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice, key wgtypes.Key) { + p := a.peersByKey[key] + if p == nil { + t.Fatal("not in peersByKey") + } + if a.peersByIP[peerVPNIP] == nil { + t.Fatal("not in peersByIP") + } + if p.State != StateRelayed { + t.Fatalf("state = %v, want StateRelayed", p.State) + } + dev.AssertNoCalls(t) + }, + }, + { + name: "public peer with endpoint goes to StateDirect via AddDirect", + peer: func(k wgtypes.Key) HubPeer { + return HubPeer{PubKey: k, VPNIP: peerVPNIP, IsPublic: true, EndpointV4: ep1} + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice, key wgtypes.Key) { + p := a.peersByKey[key] + if p == nil { + t.Fatal("not in peersByKey") + } + if p.State != StateDirect { + t.Fatalf("state = %v, want StateDirect", p.State) + } + if p.WGEndpoint != ep1 { + t.Fatalf("WGEndpoint = %v, want %v", p.WGEndpoint, ep1) + } + dev.AssertAddDirect(t, 0, p.PubKey, p.WGEndpoint, p.VPNIP) + }, + }, + { + name: "public peer with no endpoint is dropped", + peer: func(k wgtypes.Key) HubPeer { + return HubPeer{PubKey: k, VPNIP: peerVPNIP, IsPublic: true} + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice, key wgtypes.Key) { + if a.peersByKey[key] != nil { + t.Fatal("peer should not be in peersByKey") + } + if a.peersByIP[peerVPNIP] != nil { + t.Fatal("peer should not be in peersByIP") + } + dev.AssertNoCalls(t) + }, + }, + { + name: "re-add removes old WG entry before adding new one", + setup: func(a *App, key wgtypes.Key) { + a.onAddPeer(HubPeer{PubKey: key, VPNIP: peerVPNIP, IsPublic: true, EndpointV4: ep1}) + }, + peer: func(k wgtypes.Key) HubPeer { + return HubPeer{PubKey: k, VPNIP: peerVPNIP, IsPublic: true, EndpointV4: ep2} + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice, key wgtypes.Key) { + if len(dev.Calls) != 2 { + t.Fatalf("dev calls = %v, want [RemovePeer, AddDirect]", dev.Calls) + } + dev.AssertRemovePeer(t, 0, key) + dev.AssertAddDirect(t, 1, key, ep2, peerVPNIP) + if len(a.peersByKey) != 1 || len(a.peersByIP) != 1 { + t.Errorf("maps: peersByKey=%d peersByIP=%d, want 1 each", len(a.peersByKey), len(a.peersByIP)) + } + }, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + a, dev := newTestApp(t, "10.0.0.1", false, false) + key := mustKey(t) + if tc.setup != nil { + tc.setup(a, key) + dev.Calls = nil + } + a.onAddPeer(tc.peer(key)) + tc.check(t, a, dev, key) + }) + } +} + +func TestOnRemovePeer(t *testing.T) { + ep1 := netip.MustParseAddrPort("1.2.3.4:51820") + ep2 := netip.MustParseAddrPort("5.6.7.8:51820") + + testCases := []struct { + name string + setup func(t *testing.T, a *App) wgtypes.Key // returns the key to remove + check func(t *testing.T, a *App, dev *fakeWGDevice) + }{ + { + name: "unknown key is a no-op", + setup: func(t *testing.T, a *App) wgtypes.Key { + return mustKey(t) + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + dev.AssertNoCalls(t) + if len(a.peersByKey) != 0 { + t.Errorf("peersByKey should be empty") + } + }, + }, + { + name: "StateRelayed peer removed from maps without RemovePeer", + setup: func(t *testing.T, a *App) wgtypes.Key { + key := mustKey(t) + a.onAddPeer(HubPeer{PubKey: key, VPNIP: netip.MustParseAddr("10.0.0.2")}) + return key + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + dev.AssertNoCalls(t) + if len(a.peersByKey) != 0 || len(a.peersByIP) != 0 { + t.Errorf("maps should be empty after remove") + } + }, + }, + { + name: "StateDirect peer removed from maps with RemovePeer", + setup: func(t *testing.T, a *App) wgtypes.Key { + key := mustKey(t) + a.onAddPeer(HubPeer{PubKey: key, VPNIP: netip.MustParseAddr("10.0.0.2"), IsPublic: true, EndpointV4: ep1}) + return key + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 1 { + t.Fatalf("dev calls = %v, want [RemovePeer]", dev.Calls) + } + dev.AssertRemovePeer(t, 0, dev.Calls[0].PubKey) + if len(a.peersByKey) != 0 || len(a.peersByIP) != 0 { + t.Errorf("maps should be empty after remove") + } + }, + }, + { + name: "removing active relay with no backup clears relay field", + setup: func(t *testing.T, a *App) wgtypes.Key { + relay := addRelayPeer(t, a, "10.0.0.10", ep1) + a.relay = relay + return relay.PubKey + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 1 { + t.Fatalf("dev calls = %v, want [RemovePeer]", dev.Calls) + } + dev.AssertRemovePeer(t, 0, dev.Calls[0].PubKey) + if a.relay != nil { + t.Errorf("relay should be nil after removing only relay") + } + }, + }, + { + name: "removing active relay elects backup via SetRelay", + setup: func(t *testing.T, a *App) wgtypes.Key { + relay1 := addRelayPeer(t, a, "10.0.0.10", ep1) + addRelayPeer(t, a, "10.0.0.11", ep2) + a.relay = relay1 + return relay1.PubKey + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 2 { + t.Fatalf("dev calls = %v, want [RemovePeer, SetRelay]", dev.Calls) + } + dev.AssertRemovePeer(t, 0, dev.Calls[0].PubKey) + dev.AssertSetRelay(t, 1, dev.Calls[1].PubKey, ep2, a.vpnNet) + if a.relay == nil { + t.Errorf("relay should be set to backup after failover") + } + }, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + a, dev := newTestApp(t, "10.0.0.1", false, false) + key := tc.setup(t, a) + dev.Calls = nil + a.onRemovePeer(key) + tc.check(t, a, dev) + }) + } +} + +func TestSwitchActiveRelay(t *testing.T) { + ep1 := netip.MustParseAddrPort("1.2.3.4:51820") + ep2 := netip.MustParseAddrPort("5.6.7.8:51820") + + testCases := []struct { + name string + setup func(t *testing.T, a *App) + check func(t *testing.T, a *App, dev *fakeWGDevice) + }{ + { + name: "no candidates leaves relay nil", + setup: func(t *testing.T, a *App) {}, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + dev.AssertNoCalls(t) + if a.relay != nil { + t.Error("relay should be nil") + } + }, + }, + { + name: "single candidate elected via SetRelay", + setup: func(t *testing.T, a *App) { + addRelayPeer(t, a, "10.0.0.10", ep1) + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 1 { + t.Fatalf("dev calls = %v, want [SetRelay]", dev.Calls) + } + dev.AssertSetRelay(t, 0, dev.Calls[0].PubKey, ep1, a.vpnNet) + if a.relay == nil { + t.Error("relay should be set") + } + }, + }, + { + name: "measured RTT beats zero RTT", + setup: func(t *testing.T, a *App) { + r1 := addRelayPeer(t, a, "10.0.0.10", ep1) + r1.RTT = 10 * time.Millisecond + addRelayPeer(t, a, "10.0.0.11", ep2) // RTT stays 0 + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 1 { + t.Fatalf("dev calls = %v, want [SetRelay]", dev.Calls) + } + dev.AssertSetRelay(t, 0, dev.Calls[0].PubKey, ep1, a.vpnNet) + }, + }, + { + name: "lower RTT wins", + setup: func(t *testing.T, a *App) { + r1 := addRelayPeer(t, a, "10.0.0.10", ep1) + r1.RTT = 5 * time.Millisecond + r2 := addRelayPeer(t, a, "10.0.0.11", ep2) + r2.RTT = 20 * time.Millisecond + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 1 { + t.Fatalf("dev calls = %v, want [SetRelay]", dev.Calls) + } + dev.AssertSetRelay(t, 0, dev.Calls[0].PubKey, ep1, a.vpnNet) + }, + }, + { + name: "stale relay demoted to direct before backup elected", + setup: func(t *testing.T, a *App) { + old := addRelayPeer(t, a, "10.0.0.10", ep1) + old.Up = false // stale — this is what triggers the switch from onTick + a.relay = old + addRelayPeer(t, a, "10.0.0.11", ep2) + }, + check: func(t *testing.T, a *App, dev *fakeWGDevice) { + if len(dev.Calls) != 2 { + t.Fatalf("dev calls = %v, want [AddDirect, SetRelay]", dev.Calls) + } + if dev.Calls[0].Method != "AddDirect" || dev.Calls[0].Endpoint != ep1 { + t.Errorf("call[0]: got %v, want AddDirect with ep1", dev.Calls[0]) + } + dev.AssertSetRelay(t, 1, dev.Calls[1].PubKey, ep2, a.vpnNet) + if a.relay == nil || a.relay.WGEndpoint != ep2 { + t.Error("relay should be the backup peer") + } + }, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + a, dev := newTestApp(t, "10.0.0.1", false, false) + tc.setup(t, a) + dev.Calls = nil + a.switchActiveRelay() + tc.check(t, a, dev) + }) + } +} diff --git a/peer/on_multicast.go b/peer/on_multicast.go index bb6c99d..646ce31 100644 --- a/peer/on_multicast.go +++ b/peer/on_multicast.go @@ -12,7 +12,7 @@ func (a *App) onMulticastDiscovery(e MulticastEvent) { return } - if peer.IsPublic || peer.State == StateDirect { + if peer.IsPublic || peer.State() == StateDirect { return } diff --git a/peer/on_ping.go b/peer/on_ping.go index 3bd41c6..7d861ac 100644 --- a/peer/on_ping.go +++ b/peer/on_ping.go @@ -15,10 +15,9 @@ func (a *App) onPing(e PingEvent) { } now := time.Now() - peer.LastPing = now - peer.Up = true - // If we're the server, respond. + // If we're the server, respond - this is always necessary as it's used to + // know if peers are up or down. if peer.Role == control.Server { a.sendPing(peer, e.ping.ID, e.ping.PingTS) } @@ -35,7 +34,7 @@ func (a *App) onPing(e PingEvent) { // We can only learn our own endpoint from directly-connected peers — Dst // is the sender's observation of our WG handshake source. - if peer.State == StateDirect { + if peer.State() == StateDirect { if dst := e.ping.Dst; dst.IsValid() { if dst.Addr().Is4() { a.selfV4 = dst @@ -51,16 +50,11 @@ func (a *App) onPing(e PingEvent) { func (a *App) addProbe(peer *Peer, v4, v6 netip.AddrPort) { endpoint := preferredEndpoint(v4, v6) - if !endpoint.IsValid() || endpoint == peer.WGEndpoint { + if !endpoint.IsValid() || endpoint == peer.WGEndpoint() { return } peer.Endpoint4 = v4 peer.Endpoint6 = v6 - peer.WGEndpoint = endpoint - - if peer.State == StateRelayed { - peer.State = StateProbing - } a.devAddProbe(peer, endpoint) } diff --git a/peer/on_tick.go b/peer/on_tick.go index 5673162..30c3d64 100644 --- a/peer/on_tick.go +++ b/peer/on_tick.go @@ -5,6 +5,7 @@ import ( "time" "vppn/peer/control" + "vppn/peer/wginterface" ) func (a *App) onTick() { @@ -13,38 +14,36 @@ func (a *App) onTick() { a.nextPingID++ now := time.Now().UnixNano() - // Update Up values. - for _, p := range a.peersByKey { - p.Up = p.Alive() - } - for _, wgPeer := range wgPeers { p, ok := a.peersByKey[wgPeer.PublicKey] if !ok { log.Fatalf("Wireguard peer not in index: %v", wgPeer) } + p.wgPeer = wgPeer // Send pings to peers where we're the client. if p.Role == control.Client { a.sendPing(p, a.nextPingID, now) } - // Promote probing peers to direct once alive (direct path confirmed - // working). - if p.State == StateProbing && time.Since(wgPeer.LastHandshakeTime) < 2*PingInterval { - p.State = StateDirect - a.devAddDirect(p, p.WGEndpoint) - } + switch p.State() { + case StateProbing: + // Promote probing peers to direct once alive (direct path confirmed + // working). + if time.Since(p.LastHandshakeTime()) < wginterface.SessionTimeout { + a.devAddDirect(p, p.WGEndpoint()) + } - // Demote stale non-public direct peers back to probing. - if p.State == StateDirect && !p.IsPublic && !p.Up { - p.State = StateProbing - a.devAddProbe(p, p.WGEndpoint) + case StateDirect: + // Demote stale non-public direct peers back to probing. + if !p.IsPublic && !p.Up() { + a.devAddProbe(p, p.WGEndpoint()) + } } } // Ensure we have a live relay. - if a.relay == nil || !a.relay.Up { + if a.relay == nil || !a.relay.Up() { a.switchActiveRelay() } } diff --git a/peer/remote.go b/peer/remote.go index f8553fd..69b8f53 100644 --- a/peer/remote.go +++ b/peer/remote.go @@ -7,6 +7,7 @@ import ( "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "vppn/peer/control" + "vppn/peer/wginterface" ) type PeerState string @@ -18,27 +19,53 @@ const ( ) type Peer struct { - PubKey wgtypes.Key // WireGuard public key. - VPNIP netip.Addr // VPN IP address. - IsRelay bool // Peer is a relay. - IsPublic bool // Peer has a public IP. - Endpoint4 netip.AddrPort // Reported IPv4 endpoint. - Endpoint6 netip.AddrPort // Reported IPv6 endpoint. - ObservedEndpoint netip.AddrPort // If we're public: WG handshake source endpoint. - LastPing time.Time // Time of last ping received. - RTT time.Duration // Round-trip time. - Role control.Role // Client initiates pings; server responds. - State PeerState // Current connection state. - Up bool // Is the peer alive. - WGEndpoint netip.AddrPort // Endpoint currently in WG (direct/probe). + wgPeer wgtypes.Peer + VPNIP netip.Addr // VPN IP address. + IsRelay bool // Peer is a relay. + IsPublic bool // Peer has a public IP. + Endpoint4 netip.AddrPort // Reported IPv4 endpoint. + Endpoint6 netip.AddrPort // Reported IPv6 endpoint. + RTT time.Duration // Round-trip time. + Role control.Role // Client initiates pings; server responds. } -func (p *Peer) Alive() bool { - return time.Since(p.LastPing) < TimeoutInterval +// PubKey is the wireguard public key. +func (p *Peer) PubKey() wgtypes.Key { + return p.wgPeer.PublicKey +} + +func (p *Peer) State() PeerState { + if len(p.wgPeer.AllowedIPs) > 0 { + return StateDirect + } + if p.wgPeer.Endpoint == nil { + return StateRelayed + } + return StateProbing +} + +func (p *Peer) WGEndpoint() netip.AddrPort { + ep := p.wgPeer.Endpoint + if ep == nil { + return netip.AddrPort{} + } + addr, ok := netip.AddrFromSlice(ep.IP) + if !ok { + return netip.AddrPort{} + } + return netip.AddrPortFrom(addr.Unmap(), uint16(ep.Port)) +} + +func (p *Peer) LastHandshakeTime() time.Time { + return p.wgPeer.LastHandshakeTime +} + +func (p *Peer) Up() bool { + return time.Since(p.wgPeer.LastHandshakeTime) < wginterface.SessionTimeout } func (p *Peer) CanRelay() bool { - return p.IsRelay && p.Up && p.WGEndpoint.IsValid() + return p.IsRelay && p.Up() } func (p *Peer) PreferredEndpoint() netip.AddrPort { diff --git a/peer/wginterface/interface.go b/peer/wginterface/interface.go index f73d609..3354eed 100644 --- a/peer/wginterface/interface.go +++ b/peer/wginterface/interface.go @@ -21,6 +21,7 @@ import ( // Create creates a WireGuard interface named name, assigns vpnIP/prefixLen to // it, and brings it up. func Create(name string, vpnIP net.IP, prefixLen int) error { + _ = Delete(name) // remove any stale interface left by a previous run if err := nlNewLink(name); err != nil { return fmt.Errorf("failed to create wireguard link: %w", err) } diff --git a/peer/wginterface/manage.go b/peer/wginterface/manage.go index 07f4e57..383be2a 100644 --- a/peer/wginterface/manage.go +++ b/peer/wginterface/manage.go @@ -84,6 +84,17 @@ func (d *Device) Peer(pubKey wgtypes.Key) (wgtypes.Peer, error) { return wgtypes.Peer{}, fmt.Errorf("peer %v not found in %q", pubKey, d.name) } +// AddPeer registers a peer with no AllowedIPs and no endpoint. WireGuard will +// accept handshakes from this peer but route no traffic to it yet. +func (d *Device) AddPeer(pubKey wgtypes.Key) error { + return d.client.ConfigureDevice(d.name, wgtypes.Config{ + Peers: []wgtypes.PeerConfig{{ + PublicKey: pubKey, + ReplaceAllowedIPs: true, + }}, + }) +} + // SetRelay configures the relay peer with AllowedIPs covering the entire VPN // network prefix. This is the fallback route for all VPN traffic. func (d *Device) SetRelay(pubKey wgtypes.Key, endpoint netip.AddrPort, network netip.Prefix) error {