diff --git a/hub/api/api.go b/hub/api/api.go index 2c678ef..7b656a6 100644 --- a/hub/api/api.go +++ b/hub/api/api.go @@ -147,6 +147,7 @@ func (a *API) Network_List() ([]*Network, error) { func (a *API) Peer_CreateNew(p *Peer) error { p.Version = idgen.NextID(0) p.WGPubKey = []byte{} + p.SignPubKey = []byte{} p.APIKey = idgen.NewToken() return db.Peer_Insert(a.db, p) @@ -158,6 +159,7 @@ func (a *API) Peer_Init(peer *Peer, args m.PeerInitArgs) error { peer.Version = idgen.NextID(0) peer.WGPubKey = args.WGPubKey + peer.SignPubKey = args.SignPubKey return db.Peer_UpdateFull(a.db, peer) } diff --git a/hub/api/db/generated.go b/hub/api/db/generated.go index 76d5e96..41f331d 100644 --- a/hub/api/db/generated.go +++ b/hub/api/db/generated.go @@ -339,19 +339,20 @@ func Network_List( // ---------------------------------------------------------------------------- type Peer struct { - NetworkID int64 - PeerIP byte - Version int64 - APIKey string - Name string - Addr4 []byte - Addr6 []byte - Port uint16 - Relay bool - WGPubKey []byte + NetworkID int64 + PeerIP byte + Version int64 + APIKey string + Name string + Addr4 []byte + Addr6 []byte + Port uint16 + Relay bool + WGPubKey []byte + SignPubKey []byte } -const Peer_SelectQuery = "SELECT NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey FROM peers" +const Peer_SelectQuery = "SELECT NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey,SignPubKey FROM peers" func Peer_Insert( tx TX, @@ -362,7 +363,7 @@ func Peer_Insert( return err } - _, err = tx.Exec("INSERT INTO peers(NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey) VALUES(?,?,?,?,?,?,?,?,?,?)", row.NetworkID, row.PeerIP, row.Version, row.APIKey, row.Name, row.Addr4, row.Addr6, row.Port, row.Relay, row.WGPubKey) + _, err = tx.Exec("INSERT INTO peers(NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey,SignPubKey) VALUES(?,?,?,?,?,?,?,?,?,?,?)", row.NetworkID, row.PeerIP, row.Version, row.APIKey, row.Name, row.Addr4, row.Addr6, row.Port, row.Relay, row.WGPubKey, row.SignPubKey) return err } @@ -403,7 +404,7 @@ func Peer_UpdateFull( return err } - result, err := tx.Exec("UPDATE peers SET Version=?,APIKey=?,Name=?,Addr4=?,Addr6=?,Port=?,Relay=?,WGPubKey=? WHERE NetworkID=? AND PeerIP=?", row.Version, row.APIKey, row.Name, row.Addr4, row.Addr6, row.Port, row.Relay, row.WGPubKey, row.NetworkID, row.PeerIP) + result, err := tx.Exec("UPDATE peers SET Version=?,APIKey=?,Name=?,Addr4=?,Addr6=?,Port=?,Relay=?,WGPubKey=?,SignPubKey=? WHERE NetworkID=? AND PeerIP=?", row.Version, row.APIKey, row.Name, row.Addr4, row.Addr6, row.Port, row.Relay, row.WGPubKey, row.SignPubKey, row.NetworkID, row.PeerIP) if err != nil { return err } @@ -455,8 +456,8 @@ func Peer_Get( err error, ) { row = &Peer{} - r := tx.QueryRow("SELECT NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey FROM peers WHERE NetworkID=? AND PeerIP=?", NetworkID, PeerIP) - if err = r.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey); err != nil { + r := tx.QueryRow("SELECT NetworkID,PeerIP,Version,APIKey,Name,Addr4,Addr6,Port,Relay,WGPubKey,SignPubKey FROM peers WHERE NetworkID=? AND PeerIP=?", NetworkID, PeerIP) + if err = r.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey, &row.SignPubKey); err != nil { row = nil } return @@ -472,7 +473,7 @@ func Peer_GetWhere( ) { row = &Peer{} r := tx.QueryRow(query, args...) - if err = r.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey); err != nil { + if err = r.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey, &row.SignPubKey); err != nil { row = nil } return @@ -494,7 +495,7 @@ func Peer_Iterate( defer rows.Close() for rows.Next() { row := &Peer{} - err := rows.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey) + err := rows.Scan(&row.NetworkID, &row.PeerIP, &row.Version, &row.APIKey, &row.Name, &row.Addr4, &row.Addr6, &row.Port, &row.Relay, &row.WGPubKey, &row.SignPubKey) if !yield(row, err) { return } diff --git a/hub/api/db/tables.defs b/hub/api/db/tables.defs index 15e134a..3b8c230 100644 --- a/hub/api/db/tables.defs +++ b/hub/api/db/tables.defs @@ -19,5 +19,6 @@ TABLE peers OF Peer ( Addr6 []byte, Port uint16, Relay bool, - WGPubKey []byte NoUpdate + WGPubKey []byte NoUpdate, + SignPubKey []byte NoUpdate ); diff --git a/hub/api/migrations/2024-11-30-init.sql b/hub/api/migrations/2024-11-30-init.sql index 9a0705b..f2f2a79 100644 --- a/hub/api/migrations/2024-11-30-init.sql +++ b/hub/api/migrations/2024-11-30-init.sql @@ -20,5 +20,6 @@ CREATE TABLE peers ( Port INTEGER NOT NULL, Relay INTEGER NOT NULL DEFAULT 0, -- Boolean if peer will forward packets. WGPubKey BLOB NOT NULL, + SignPubKey BLOB NOT NULL PRIMARY KEY(NetworkID, PeerIP) ) WITHOUT ROWID; diff --git a/hub/handlers.go b/hub/handlers.go index f189412..e90eafe 100644 --- a/hub/handlers.go +++ b/hub/handlers.go @@ -313,6 +313,15 @@ func (a *App) _peerInit(peer *api.Peer, w http.ResponseWriter, r *http.Request) return err } + if len(args.WGPubKey) != 32 { + http.Error(w, "invalid WGPubKey", http.StatusBadRequest) + return nil + } + if len(args.SignPubKey) != 32 { + http.Error(w, "invalid SignPubKey", http.StatusBadRequest) + return nil + } + net, err := a.api.Network_Get(peer.NetworkID) if err != nil { return err @@ -353,14 +362,15 @@ func (a *App) peersArray(networkID int64) (peers [256]*m.Peer, err error) { for _, p := range l { if len(p.WGPubKey) != 0 { peers[p.PeerIP] = &m.Peer{ - PeerIP: p.PeerIP, - Version: p.Version, - Name: p.Name, - Addr4: p.Addr4, - Addr6: p.Addr6, - Port: p.Port, - Relay: p.Relay, - WGPubKey: p.WGPubKey, + PeerIP: p.PeerIP, + Version: p.Version, + Name: p.Name, + Addr4: p.Addr4, + Addr6: p.Addr6, + Port: p.Port, + Relay: p.Relay, + WGPubKey: p.WGPubKey, + SignPubKey: p.SignPubKey, } } } diff --git a/m/models.go b/m/models.go index 36c6ff5..d16986c 100644 --- a/m/models.go +++ b/m/models.go @@ -2,7 +2,8 @@ package m type PeerInitArgs struct { - WGPubKey []byte + WGPubKey []byte + SignPubKey []byte } type PeerInitResp struct { @@ -18,8 +19,9 @@ type Peer struct { Addr4 []byte Addr6 []byte Port uint16 - Relay bool - WGPubKey []byte + Relay bool + WGPubKey []byte + SignPubKey []byte } type NetworkState struct { diff --git a/peer/app.go b/peer/app.go index 90d123d..3c69a14 100644 --- a/peer/app.go +++ b/peer/app.go @@ -29,6 +29,7 @@ type HubPeer struct { IsPublic bool EndpointV4 netip.AddrPort // zero if none EndpointV6 netip.AddrPort // zero if none + SignPubKey [32]byte } type PingEvent struct { @@ -36,10 +37,10 @@ type PingEvent struct { ping control.Ping } +// MulticastEvent carries a raw signed beacon for verification in the event loop. type MulticastEvent struct { - pubKey wgtypes.Key - vpnIP netip.Addr - endpoint netip.AddrPort + signed []byte // nacl/sign signed beacon (64-byte sig || 35-byte payload) + src netip.Addr // physical LAN source address } // App is the peer application. All mutable state lives here and is diff --git a/peer/hub_poller.go b/peer/hub_poller.go index c978889..edfc6b9 100644 --- a/peer/hub_poller.go +++ b/peer/hub_poller.go @@ -71,8 +71,14 @@ func (hp *HubPoller) poll() { log.Printf("[HubPoller] fetch: %v", err) return } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + log.Printf("[HubPoller] unexpected status %d", resp.StatusCode) + return + } + body, err := io.ReadAll(resp.Body) - _ = resp.Body.Close() if err != nil { log.Printf("[HubPoller] read body: %v", err) return @@ -93,7 +99,7 @@ func (hp *HubPoller) apply(state m.NetworkState) { netAddr := hp.vpnNet.Addr().As4() for _, p := range state.Peers { - if p == nil || len(p.WGPubKey) != wgtypes.KeyLen { + if p == nil || len(p.WGPubKey) != wgtypes.KeyLen || len(p.SignPubKey) != 32 { continue } @@ -138,6 +144,8 @@ func hubPeerFrom(pubKey wgtypes.Key, vpnIP netip.Addr, p *m.Peer) HubPeer { ep6 = netip.AddrPortFrom(addr, p.Port) } } + var signPubKey [32]byte + copy(signPubKey[:], p.SignPubKey) return HubPeer{ PubKey: pubKey, VPNIP: vpnIP, @@ -145,5 +153,6 @@ func hubPeerFrom(pubKey wgtypes.Key, vpnIP netip.Addr, p *m.Peer) HubPeer { IsPublic: ep4.IsValid() || ep6.IsValid(), EndpointV4: ep4, EndpointV6: ep6, + SignPubKey: signPubKey, } } diff --git a/peer/init.go b/peer/init.go index 8e8010c..9d5fb9f 100644 --- a/peer/init.go +++ b/peer/init.go @@ -2,6 +2,7 @@ package peer import ( "bytes" + "crypto/rand" "encoding/base64" "encoding/json" "fmt" @@ -10,6 +11,7 @@ import ( "os" "path/filepath" + "golang.org/x/crypto/nacl/sign" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "vppn/m" @@ -19,6 +21,7 @@ import ( // loaded on every subsequent run. type LocalState struct { PrivKey wgtypes.Key + SignKey [64]byte // nacl/sign Ed25519 private key VPNIP netip.Addr VPNNet netip.Prefix WGPort uint16 @@ -28,7 +31,8 @@ type LocalState struct { // localStateJSON is the on-disk representation. type localStateJSON struct { - PrivKey string `json:"priv_key"` // standard base64 + PrivKey string `json:"priv_key"` // standard base64 + SignKey string `json:"sign_key"` // standard base64 VPNIP netip.Addr `json:"vpn_ip"` VPNNet netip.Prefix `json:"vpn_net"` WGPort uint16 `json:"wg_port"` @@ -60,8 +64,17 @@ func LoadOrInit(statePath, hubURL, apiKey string) (LocalState, error) { } func initFromHub(hubURL, apiKey string, privKey wgtypes.Key) (LocalState, error) { - pubKey := privKey.PublicKey() - body, _ := json.Marshal(m.PeerInitArgs{WGPubKey: pubKey[:]}) + wgPubKey := privKey.PublicKey() + + signPubKey, signPrivKey, err := sign.GenerateKey(rand.Reader) + if err != nil { + return LocalState{}, fmt.Errorf("generate sign key: %w", err) + } + + body, _ := json.Marshal(m.PeerInitArgs{ + WGPubKey: wgPubKey[:], + SignPubKey: signPubKey[:], + }) req, err := http.NewRequest(http.MethodPost, hubURL+"/peer/init/", bytes.NewReader(body)) if err != nil { @@ -105,6 +118,7 @@ func initFromHub(hubURL, apiKey string, privKey wgtypes.Key) (LocalState, error) return LocalState{ PrivKey: privKey, + SignKey: *signPrivKey, VPNIP: vpnIP, VPNNet: vpnNet, WGPort: wgPort, @@ -126,8 +140,18 @@ func parseLocalState(data []byte) (LocalState, error) { if err != nil { return LocalState{}, fmt.Errorf("invalid key: %w", err) } + signKeyBytes, err := base64.StdEncoding.DecodeString(j.SignKey) + if err != nil { + return LocalState{}, fmt.Errorf("decode sign key: %w", err) + } + if len(signKeyBytes) != 64 { + return LocalState{}, fmt.Errorf("invalid sign key length: %d", len(signKeyBytes)) + } + var signKey [64]byte + copy(signKey[:], signKeyBytes) return LocalState{ PrivKey: key, + SignKey: signKey, VPNIP: j.VPNIP, VPNNet: j.VPNNet, WGPort: j.WGPort, @@ -139,6 +163,7 @@ func parseLocalState(data []byte) (LocalState, error) { func saveLocalState(path string, s LocalState) error { j := localStateJSON{ PrivKey: base64.StdEncoding.EncodeToString(s.PrivKey[:]), + SignKey: base64.StdEncoding.EncodeToString(s.SignKey[:]), VPNIP: s.VPNIP, VPNNet: s.VPNNet, WGPort: s.WGPort, diff --git a/peer/multicast.go b/peer/multicast.go index d8e2561..01d225b 100644 --- a/peer/multicast.go +++ b/peer/multicast.go @@ -8,11 +8,13 @@ import ( "net/netip" "time" + "golang.org/x/crypto/nacl/sign" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" ) const ( - mcBeaconLen = 35 // 1 VPN IP byte + 32 WG pubkey + 2 WG listen port + mcBeaconLen = 35 // 1 VPN IP byte + 32 WG pubkey + 2 WG listen port + mcSignedBeaconLen = sign.Overhead + mcBeaconLen // 64-byte nacl/sign prefix + payload mcBroadcastInterval = 32 * time.Second mcErrorRetryInterval = 16 * time.Second ) @@ -21,18 +23,23 @@ var mcAddr = net.UDPAddrFromAddrPort(netip.AddrPortFrom( netip.AddrFrom4([4]byte{224, 0, 0, 157}), 4560)) -// RunMCWriter broadcasts a beacon on the local multicast group every +// RunMCWriter broadcasts a signed beacon on the local multicast group every // mcBroadcastInterval so that LAN peers can discover our WireGuard endpoint. -func RunMCWriter(selfVPNIP netip.Addr, pubKey wgtypes.Key, wgPort uint16) { +func RunMCWriter(selfVPNIP netip.Addr, pubKey wgtypes.Key, wgPort uint16, signKey *[64]byte) { conn, err := net.ListenMulticastUDP("udp", nil, mcAddr) if err != nil { - log.Fatalf("[MCWriter] bind: %v", err) + log.Printf("[MCWriter] bind: %v", err) + return } - beacon := buildBeacon(selfVPNIP, pubKey, wgPort) + payload := buildBeacon(selfVPNIP, pubKey, wgPort) + signed := sign.Sign(nil, payload, signKey) + if _, err := conn.WriteToUDP(signed, mcAddr); err != nil { + log.Printf("[MCWriter] write: %v", err) + } for range time.Tick(mcBroadcastInterval) { - if _, err := conn.WriteToUDP(beacon, mcAddr); err != nil { + if _, err := conn.WriteToUDP(signed, mcAddr); err != nil { log.Printf("[MCWriter] write: %v", err) } } @@ -64,7 +71,7 @@ func runMCReaderInner(vpnNet netip.Prefix, selfVPNIP netip.Addr, ch chan<- Multi } defer conn.Close() - buf := make([]byte, 64) + buf := make([]byte, mcSignedBeaconLen+1) // +1 to detect oversized packets netAddr := vpnNet.Addr().As4() for { @@ -73,29 +80,25 @@ func runMCReaderInner(vpnNet netip.Prefix, selfVPNIP netip.Addr, ch chan<- Multi if err != nil { return fmt.Errorf("read: %w", err) } - if n != mcBeaconLen { + if n != mcSignedBeaconLen { continue } + // Peek at VPN IP byte (first byte of payload after 64-byte sig prefix) + // to skip our own beacon before verifying. octets := netAddr - octets[3] = buf[0] + octets[3] = buf[sign.Overhead] vpnIP := netip.AddrFrom4(octets) if vpnIP == selfVPNIP { continue } - pubKey, err := wgtypes.NewKey(buf[1:33]) - if err != nil { - continue - } - - wgPort := binary.BigEndian.Uint16(buf[33:35]) - endpoint := netip.AddrPortFrom(src.Addr().Unmap(), wgPort) + signed := make([]byte, mcSignedBeaconLen) + copy(signed, buf[:mcSignedBeaconLen]) ch <- MulticastEvent{ - pubKey: pubKey, - vpnIP: vpnIP, - endpoint: endpoint, + signed: signed, + src: src.Addr().Unmap(), } } } diff --git a/peer/new.go b/peer/new.go index 144ec94..ffe0154 100644 --- a/peer/new.go +++ b/peer/new.go @@ -58,7 +58,7 @@ func New( go cc.run(pingCh) go poller.Run() - go RunMCWriter(state.VPNIP, state.PrivKey.PublicKey(), state.WGPort) + go RunMCWriter(state.VPNIP, state.PrivKey.PublicKey(), state.WGPort, &state.SignKey) if !state.IsPublic { go RunMCReader(state.VPNNet, state.VPNIP, multicastCh) } diff --git a/peer/on_hub.go b/peer/on_hub.go index ec0de5a..be66b89 100644 --- a/peer/on_hub.go +++ b/peer/on_hub.go @@ -15,14 +15,15 @@ func (a *App) onAddPeer(p HubPeer) { a.onRemovePeer(p.PubKey) peer := &Peer{ - wgPeer: wgtypes.Peer{PublicKey: p.PubKey}, - VPNIP: p.VPNIP, - IsRelay: p.IsRelay, - IsPublic: p.IsPublic, - Endpoint4: p.EndpointV4, - Endpoint6: p.EndpointV6, - RTT: time.Duration(math.MaxInt64) * time.Nanosecond, - Role: roleFor(a.isPublic, a.vpnIP, p), + wgPeer: wgtypes.Peer{PublicKey: p.PubKey}, + VPNIP: p.VPNIP, + IsRelay: p.IsRelay, + IsPublic: p.IsPublic, + Endpoint4: p.EndpointV4, + Endpoint6: p.EndpointV6, + RTT: time.Duration(math.MaxInt64) * time.Nanosecond, + Role: roleFor(a.isPublic, a.vpnIP, p), + SignPubKey: p.SignPubKey, } endpoint := peer.PreferredEndpoint() diff --git a/peer/on_multicast.go b/peer/on_multicast.go index 646ce31..000a375 100644 --- a/peer/on_multicast.go +++ b/peer/on_multicast.go @@ -1,13 +1,29 @@ package peer -import "net/netip" +import ( + "encoding/binary" + "net/netip" + + "golang.org/x/crypto/nacl/sign" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) func (a *App) onMulticastDiscovery(e MulticastEvent) { if a.isPublic { return } - peer, ok := a.peersByKey[e.pubKey] + // Peek at the VPN IP byte to find the sender peer before verifying. + // nacl/sign prepends a 64-byte signature, so payload starts at offset sign.Overhead. + if len(e.signed) != mcSignedBeaconLen { + return + } + netAddr := a.vpnNet.Addr().As4() + octets := netAddr + octets[3] = e.signed[sign.Overhead] + vpnIP := netip.AddrFrom4(octets) + + peer, ok := a.peersByIP[vpnIP] if !ok { return } @@ -16,11 +32,28 @@ func (a *App) onMulticastDiscovery(e MulticastEvent) { return } + payload, ok := sign.Open(nil, e.signed, &peer.SignPubKey) + if !ok { + return + } + + // payload: [1 VPN IP byte][32 WG pubkey][2 WG port] + wgPubKey, err := wgtypes.NewKey(payload[1:33]) + if err != nil || wgPubKey != peer.PubKey() { + return + } + + wgPort := binary.BigEndian.Uint16(payload[33:35]) + endpoint := netip.AddrPortFrom(e.src, wgPort) + if !endpoint.IsValid() { + return + } + var v4, v6 netip.AddrPort - if e.endpoint.Addr().Is4() { - v4 = e.endpoint + if e.src.Is4() { + v4 = endpoint } else { - v6 = e.endpoint + v6 = endpoint } a.addProbe(peer, v4, v6) diff --git a/peer/remote.go b/peer/remote.go index 5a6f0c5..d7fd56d 100644 --- a/peer/remote.go +++ b/peer/remote.go @@ -19,14 +19,15 @@ const ( ) type Peer struct { - 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. + 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. + SignPubKey [32]byte // nacl/sign public key for verifying multicast beacons. } // PubKey is the wireguard public key. diff --git a/peer/wginterface/manage.go b/peer/wginterface/manage.go index 383be2a..90789cd 100644 --- a/peer/wginterface/manage.go +++ b/peer/wginterface/manage.go @@ -173,9 +173,12 @@ func (d *Device) RemovePeer(pubKey wgtypes.Key) error { }) } -// EnableForwarding enables IPv4 forwarding on the interface, required for -// relay peers that forward traffic between VPN peers. +// EnableForwarding enables IPv4 forwarding globally and on the interface, +// required for relay peers that forward traffic between VPN peers. func (d *Device) EnableForwarding() error { + if err := os.WriteFile("/proc/sys/net/ipv4/ip_forward", []byte("1\n"), 0644); err != nil { + return err + } path := fmt.Sprintf("/proc/sys/net/ipv4/conf/%s/forwarding", d.name) return os.WriteFile(path, []byte("1\n"), 0644) }