diff --git a/client/main.go b/client/main.go index e1b4422..fd9e809 100644 --- a/client/main.go +++ b/client/main.go @@ -22,11 +22,12 @@ const ( ) var ( - internalIP net.IP = nil - serverAddr *net.UDPAddr - iface *water.Interface - conn *net.UDPConn - privKey hpke.PrivateKey + iface *water.Interface + conn *net.UDPConn + serverAddr *net.UDPAddr + internalIP net.IP = nil + privKey hpke.PrivateKey + multicastKey []byte establishMutex sync.Mutex peerPubKeys = shared.NewTMap[string, hpke.PublicKey]() @@ -62,74 +63,13 @@ func main() { func register() { var err error - privKey, err = hpke.MLKEM768X25519().GenerateKey() if err != nil { panic(err) } pubKey := privKey.PublicKey() - req := 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 - } + send(shared.BuildPkt(shared.REQ_REGISTER, pubKey.Bytes())) } func receivePackets() { @@ -148,11 +88,11 @@ func receivePackets() { continue } - go handlePkt(buffer[:n]) + go handleIncomingPkt(buffer[:n]) } } -func handlePkt(pkt []byte) { +func handleIncomingPkt(pkt []byte) { defer func() { if r := recover(); r != nil { log.Printf("recovered from panic while handling a packet: %v\n", r) @@ -173,10 +113,18 @@ func handlePkt(pkt []byte) { switch pktType { 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()) setupInterface() - go sendPackets() + go relayPackets() case shared.ENC_PKT: srcIP := net.IP(pkt[20:36]).String() @@ -192,8 +140,12 @@ func handlePkt(pkt []byte) { log.Println(err) } case shared.BROADCAST_PKT: - log.Println("received pkt:", hex.EncodeToString(pkt[4:])) - if _, err := iface.Write(pkt[4:]); err != nil { + decryptedPkt, err := DecryptSym(multicastKey, pkt[4:]) + if err != nil { + panic(err) + } + log.Println("received broadcast:", hex.EncodeToString(decryptedPkt)) + if _, err := iface.Write(decryptedPkt); err != nil { log.Println(err) } case shared.RESP_GET_PUBKEY: @@ -220,10 +172,7 @@ func handlePkt(pkt []byte) { peerSessionKeys.Set(srcIP, sessionKey) reqData := append(net.ParseIP(srcIP), internalIP...) - req := shared.BuildPkt(shared.RESP_ESTABLISH, reqData) - if _, err := conn.WriteToUDP(req, serverAddr); err != nil { - panic(err) - } + send(shared.BuildPkt(shared.RESP_ESTABLISH, reqData)) case shared.RESP_ESTABLISH: srcIP := net.IP(pkt[20:36]).String() 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() { config := water.Config{DeviceType: water.TAP} config.Name = IFACE_NAME @@ -251,7 +251,7 @@ func setupInterface() { } } -func sendPackets() { +func relayPackets() { log.Println("Listening for packets...") for { @@ -272,10 +272,12 @@ func sendPackets() { destIP := net.IP(payload[24:40]).String() if payload[24] == 0xff { // multicast - req := shared.BuildPkt(shared.BROADCAST_PKT, pkt) - if _, err := conn.WriteToUDP(req, serverAddr); err != nil { - log.Println(err) + encryptedPkt, err := EncryptSym(multicastKey, pkt) + if err != nil { + panic(err) } + + send(shared.BuildPkt(shared.BROADCAST_PKT, encryptedPkt)) } else { destSessionKey := getOrEstablishSessionKey(destIP) @@ -285,10 +287,13 @@ func sendPackets() { } reqBody := append(net.ParseIP(destIP), append(internalIP, encryptedPkt...)...) - req := shared.BuildPkt(shared.ENC_PKT, reqBody) - if _, err := conn.WriteToUDP(req, serverAddr); err != nil { - log.Println(err) - } + send(shared.BuildPkt(shared.ENC_PKT, reqBody)) } } } + +func send(req []byte) { + if _, err := conn.WriteToUDP(req, serverAddr); err != nil { + panic(err) + } +} diff --git a/server/main.go b/server/main.go index e30f153..1f52278 100644 --- a/server/main.go +++ b/server/main.go @@ -6,6 +6,7 @@ package main import ( + "crypto/hpke" "crypto/rand" "encoding/binary" "log" @@ -18,12 +19,13 @@ import ( type Peer struct { RealAddr *net.UDPAddr InternalIP string - PublicKey []byte + PublicKey hpke.PublicKey } var ( - peers = shared.NewTMap[string, *Peer]() - conn *net.UDPConn + peers = shared.NewTMap[string, *Peer]() + conn *net.UDPConn + multicastKey = make([]byte, 32) ) func main() { @@ -31,6 +33,10 @@ func main() { panic("root permissions needed") } + if _, err := rand.Read(multicastKey); err != nil { + panic(err) + } + receivePackets() } @@ -83,7 +89,11 @@ func handleReq(req []byte, addr *net.UDPAddr) { switch pktType { 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() @@ -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) case shared.ENC_PKT: peer := peers.Get(addr.IP.String()) @@ -126,8 +141,7 @@ func handleReq(req []byte, addr *net.UDPAddr) { return } - destIP := net.IP(req[42:58]).String() - log.Println(peer.InternalIP + " -> " + destIP) + log.Println(peer.InternalIP + " -> *") broadcast(req, peer.InternalIP) case shared.REQ_GET_PUBKEY: peerIP := net.IP(req[4:20]).String() @@ -137,7 +151,7 @@ func handleReq(req []byte, addr *net.UDPAddr) { 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) case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH: peer := peers.Get(addr.IP.String()) diff --git a/shared/proto.go b/shared/proto.go index fc00e65..226856f 100644 --- a/shared/proto.go +++ b/shared/proto.go @@ -5,18 +5,18 @@ import "encoding/binary" const PROTO_VERSION uint16 = 1 const ( - unused uint16 = iota + _ uint16 = iota // (client -> server) requests an IP // [ pubkey - 1216 bytes ] REQ_REGISTER - // (server -> client) assigns an IP - // [ ip - 16 bytes ] + // (server -> client) returns the assigned IP and encapsulated multicast key + // [ ip - 16 bytes ] [ ciphertext - 1168 bytes ] RESP_REGISTER // (client -> server -> client2) relays an encrypted packet to a specified peer // [ destIP - 16 bytes ] [ srcIP - 16 bytes ] [ encrypted pkt ] ENC_PKT - // (client -> server -> *) broadcasts an unsecured packet - // [ plaintext pkt ] + // (client -> server -> *) broadcasts an encrypted packet + // [ encrypted pkt ] BROADCAST_PKT // (client -> server) requests peer's pubkey from the server for encapsulation // [ ip - 16 bytes ]