encrypt multicast with shared key

This commit is contained in:
2026-07-22 17:25:18 +02:00
parent f9f0df478d
commit 412426fd70
3 changed files with 117 additions and 98 deletions

View File

@@ -22,11 +22,12 @@ const (
) )
var ( var (
internalIP net.IP = nil
serverAddr *net.UDPAddr
iface *water.Interface iface *water.Interface
conn *net.UDPConn conn *net.UDPConn
serverAddr *net.UDPAddr
internalIP net.IP = nil
privKey hpke.PrivateKey privKey hpke.PrivateKey
multicastKey []byte
establishMutex sync.Mutex establishMutex sync.Mutex
peerPubKeys = shared.NewTMap[string, hpke.PublicKey]() peerPubKeys = shared.NewTMap[string, hpke.PublicKey]()
@@ -62,74 +63,13 @@ func main() {
func register() { func register() {
var err error var err error
privKey, err = hpke.MLKEM768X25519().GenerateKey() privKey, err = hpke.MLKEM768X25519().GenerateKey()
if err != nil { if err != nil {
panic(err) panic(err)
} }
pubKey := privKey.PublicKey() pubKey := privKey.PublicKey()
req := shared.BuildPkt(shared.REQ_REGISTER, pubKey.Bytes()) send(shared.BuildPkt(shared.REQ_REGISTER, pubKey.Bytes()))
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
panic(err)
}
}
func getOrEstablishSessionKey(ip string) []byte {
if key, ok := peerSessionKeys.GetOK(ip); ok {
return key
}
establishMutex.Lock()
defer establishMutex.Unlock()
// check again after the other goroutine finished
if key, ok := peerSessionKeys.GetOK(ip); ok {
return key
}
// TODO: this definitely shouldnt block the main thread
for {
// request pubkey every 50ms
time.Sleep(50 * time.Millisecond)
peerPubKey, ok := peerPubKeys.GetOK(ip)
if !ok {
req := shared.BuildPkt(shared.REQ_GET_PUBKEY, net.ParseIP(ip))
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
panic(err)
}
continue
}
log.Println("got pubkey of " + ip)
ciphertext, sender, err := hpke.NewSender(peerPubKey, hpke.HKDFSHA256(), hpke.ExportOnly(), nil)
if err != nil {
panic(err)
}
sessionKey, err := sender.Export("baalvpn", 32)
if err != nil {
panic(err)
}
reqData := append(net.ParseIP(ip), append(internalIP, ciphertext...)...)
req := shared.BuildPkt(shared.REQ_ESTABLISH, reqData)
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
panic(err)
}
for {
// TODO: eww
time.Sleep(50 * time.Millisecond)
if peerEstablishAcks.Get(ip) {
peerEstablishAcks.Delete(ip)
break
}
}
peerSessionKeys.Set(ip, sessionKey)
return sessionKey
}
} }
func receivePackets() { func receivePackets() {
@@ -148,11 +88,11 @@ func receivePackets() {
continue continue
} }
go handlePkt(buffer[:n]) go handleIncomingPkt(buffer[:n])
} }
} }
func handlePkt(pkt []byte) { func handleIncomingPkt(pkt []byte) {
defer func() { defer func() {
if r := recover(); r != nil { if r := recover(); r != nil {
log.Printf("recovered from panic while handling a packet: %v\n", r) log.Printf("recovered from panic while handling a packet: %v\n", r)
@@ -173,10 +113,18 @@ func handlePkt(pkt []byte) {
switch pktType { switch pktType {
case shared.RESP_REGISTER: case shared.RESP_REGISTER:
internalIP = net.IP(pkt[4:]) internalIP = net.IP(pkt[4:20])
encryptedMulticastKey := net.IP(pkt[20:])
var err error
multicastKey, err = hpke.Open(privKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn"), encryptedMulticastKey)
if err != nil {
panic(err)
}
log.Println("got assigned IP by the server:", internalIP.String()) log.Println("got assigned IP by the server:", internalIP.String())
setupInterface() setupInterface()
go sendPackets() go relayPackets()
case shared.ENC_PKT: case shared.ENC_PKT:
srcIP := net.IP(pkt[20:36]).String() srcIP := net.IP(pkt[20:36]).String()
@@ -192,8 +140,12 @@ func handlePkt(pkt []byte) {
log.Println(err) log.Println(err)
} }
case shared.BROADCAST_PKT: case shared.BROADCAST_PKT:
log.Println("received pkt:", hex.EncodeToString(pkt[4:])) decryptedPkt, err := DecryptSym(multicastKey, pkt[4:])
if _, err := iface.Write(pkt[4:]); err != nil { if err != nil {
panic(err)
}
log.Println("received broadcast:", hex.EncodeToString(decryptedPkt))
if _, err := iface.Write(decryptedPkt); err != nil {
log.Println(err) log.Println(err)
} }
case shared.RESP_GET_PUBKEY: case shared.RESP_GET_PUBKEY:
@@ -220,10 +172,7 @@ func handlePkt(pkt []byte) {
peerSessionKeys.Set(srcIP, sessionKey) peerSessionKeys.Set(srcIP, sessionKey)
reqData := append(net.ParseIP(srcIP), internalIP...) reqData := append(net.ParseIP(srcIP), internalIP...)
req := shared.BuildPkt(shared.RESP_ESTABLISH, reqData) send(shared.BuildPkt(shared.RESP_ESTABLISH, reqData))
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
panic(err)
}
case shared.RESP_ESTABLISH: case shared.RESP_ESTABLISH:
srcIP := net.IP(pkt[20:36]).String() srcIP := net.IP(pkt[20:36]).String()
log.Println("established session key with", srcIP) log.Println("established session key with", srcIP)
@@ -233,6 +182,57 @@ func handlePkt(pkt []byte) {
} }
} }
func getOrEstablishSessionKey(ip string) []byte {
if key, ok := peerSessionKeys.GetOK(ip); ok {
return key
}
establishMutex.Lock()
defer establishMutex.Unlock()
// check again after the other goroutine finished
if key, ok := peerSessionKeys.GetOK(ip); ok {
return key
}
// TODO: this definitely shouldnt block the main thread
for {
// request pubkey every 50ms
time.Sleep(50 * time.Millisecond)
peerPubKey, ok := peerPubKeys.GetOK(ip)
if !ok {
send(shared.BuildPkt(shared.REQ_GET_PUBKEY, net.ParseIP(ip)))
continue
}
log.Println("got pubkey of " + ip)
ciphertext, sender, err := hpke.NewSender(peerPubKey, hpke.HKDFSHA256(), hpke.ExportOnly(), nil)
if err != nil {
panic(err)
}
sessionKey, err := sender.Export("baalvpn", 32)
if err != nil {
panic(err)
}
reqData := append(net.ParseIP(ip), append(internalIP, ciphertext...)...)
send(shared.BuildPkt(shared.REQ_ESTABLISH, reqData))
for {
// TODO: eww
time.Sleep(50 * time.Millisecond)
if peerEstablishAcks.Get(ip) {
peerEstablishAcks.Delete(ip)
break
}
}
peerSessionKeys.Set(ip, sessionKey)
return sessionKey
}
}
func setupInterface() { func setupInterface() {
config := water.Config{DeviceType: water.TAP} config := water.Config{DeviceType: water.TAP}
config.Name = IFACE_NAME config.Name = IFACE_NAME
@@ -251,7 +251,7 @@ func setupInterface() {
} }
} }
func sendPackets() { func relayPackets() {
log.Println("Listening for packets...") log.Println("Listening for packets...")
for { for {
@@ -272,10 +272,12 @@ func sendPackets() {
destIP := net.IP(payload[24:40]).String() destIP := net.IP(payload[24:40]).String()
if payload[24] == 0xff { // multicast if payload[24] == 0xff { // multicast
req := shared.BuildPkt(shared.BROADCAST_PKT, pkt) encryptedPkt, err := EncryptSym(multicastKey, pkt)
if _, err := conn.WriteToUDP(req, serverAddr); err != nil { if err != nil {
log.Println(err) panic(err)
} }
send(shared.BuildPkt(shared.BROADCAST_PKT, encryptedPkt))
} else { } else {
destSessionKey := getOrEstablishSessionKey(destIP) destSessionKey := getOrEstablishSessionKey(destIP)
@@ -285,10 +287,13 @@ func sendPackets() {
} }
reqBody := append(net.ParseIP(destIP), append(internalIP, encryptedPkt...)...) reqBody := append(net.ParseIP(destIP), append(internalIP, encryptedPkt...)...)
req := shared.BuildPkt(shared.ENC_PKT, reqBody) send(shared.BuildPkt(shared.ENC_PKT, reqBody))
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
log.Println(err)
}
} }
} }
} }
func send(req []byte) {
if _, err := conn.WriteToUDP(req, serverAddr); err != nil {
panic(err)
}
}

View File

@@ -6,6 +6,7 @@
package main package main
import ( import (
"crypto/hpke"
"crypto/rand" "crypto/rand"
"encoding/binary" "encoding/binary"
"log" "log"
@@ -18,12 +19,13 @@ import (
type Peer struct { type Peer struct {
RealAddr *net.UDPAddr RealAddr *net.UDPAddr
InternalIP string InternalIP string
PublicKey []byte PublicKey hpke.PublicKey
} }
var ( var (
peers = shared.NewTMap[string, *Peer]() peers = shared.NewTMap[string, *Peer]()
conn *net.UDPConn conn *net.UDPConn
multicastKey = make([]byte, 32)
) )
func main() { func main() {
@@ -31,6 +33,10 @@ func main() {
panic("root permissions needed") panic("root permissions needed")
} }
if _, err := rand.Read(multicastKey); err != nil {
panic(err)
}
receivePackets() receivePackets()
} }
@@ -83,7 +89,11 @@ func handleReq(req []byte, addr *net.UDPAddr) {
switch pktType { switch pktType {
case shared.REQ_REGISTER: case shared.REQ_REGISTER:
pubKey := req[4:] pubKey, err := hpke.MLKEM768X25519().NewPublicKey(req[4:])
if err != nil {
log.Println("invalid pubkey")
return
}
internalIP := randomIP().String() internalIP := randomIP().String()
@@ -97,7 +107,12 @@ func handleReq(req []byte, addr *net.UDPAddr) {
}) })
} }
resp := shared.BuildPkt(shared.RESP_REGISTER, net.ParseIP(internalIP)) ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn"), multicastKey)
if err != nil {
panic(err)
}
resp := shared.BuildPkt(shared.RESP_REGISTER, append(net.ParseIP(internalIP), ciphertext...))
sendTo(addr, resp) sendTo(addr, resp)
case shared.ENC_PKT: case shared.ENC_PKT:
peer := peers.Get(addr.IP.String()) peer := peers.Get(addr.IP.String())
@@ -126,8 +141,7 @@ func handleReq(req []byte, addr *net.UDPAddr) {
return return
} }
destIP := net.IP(req[42:58]).String() log.Println(peer.InternalIP + " -> *")
log.Println(peer.InternalIP + " -> " + destIP)
broadcast(req, peer.InternalIP) broadcast(req, peer.InternalIP)
case shared.REQ_GET_PUBKEY: case shared.REQ_GET_PUBKEY:
peerIP := net.IP(req[4:20]).String() peerIP := net.IP(req[4:20]).String()
@@ -137,7 +151,7 @@ func handleReq(req []byte, addr *net.UDPAddr) {
return return
} }
resp := shared.BuildPkt(shared.RESP_GET_PUBKEY, append(net.ParseIP(peer.InternalIP), peer.PublicKey...)) resp := shared.BuildPkt(shared.RESP_GET_PUBKEY, append(net.ParseIP(peer.InternalIP), peer.PublicKey.Bytes()...))
sendTo(addr, resp) sendTo(addr, resp)
case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH: case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH:
peer := peers.Get(addr.IP.String()) peer := peers.Get(addr.IP.String())

View File

@@ -5,18 +5,18 @@ import "encoding/binary"
const PROTO_VERSION uint16 = 1 const PROTO_VERSION uint16 = 1
const ( const (
unused uint16 = iota _ uint16 = iota
// (client -> server) requests an IP // (client -> server) requests an IP
// [ pubkey - 1216 bytes ] // [ pubkey - 1216 bytes ]
REQ_REGISTER REQ_REGISTER
// (server -> client) assigns an IP // (server -> client) returns the assigned IP and encapsulated multicast key
// [ ip - 16 bytes ] // [ ip - 16 bytes ] [ ciphertext - 1168 bytes ]
RESP_REGISTER RESP_REGISTER
// (client -> server -> client2) relays an encrypted packet to a specified peer // (client -> server -> client2) relays an encrypted packet to a specified peer
// [ destIP - 16 bytes ] [ srcIP - 16 bytes ] [ encrypted pkt ] // [ destIP - 16 bytes ] [ srcIP - 16 bytes ] [ encrypted pkt ]
ENC_PKT ENC_PKT
// (client -> server -> *) broadcasts an unsecured packet // (client -> server -> *) broadcasts an encrypted packet
// [ plaintext pkt ] // [ encrypted pkt ]
BROADCAST_PKT BROADCAST_PKT
// (client -> server) requests peer's pubkey from the server for encapsulation // (client -> server) requests peer's pubkey from the server for encapsulation
// [ ip - 16 bytes ] // [ ip - 16 bytes ]