// TODO: slice bounds checking // TODO: authentication // TODO: somehow persist IPs // TODO: handle multiple peers behind one NAT // TODO: key rotation package main import ( "crypto/hpke" "crypto/rand" "encoding/binary" "log" "net" "os" "baalvpn/shared" ) type Peer struct { RealAddr *net.UDPAddr InternalIP string PublicKey hpke.PublicKey } var ( peers = shared.NewTMap[string, *Peer]() conn *net.UDPConn multicastKey = make([]byte, 32) ) func main() { if os.Geteuid() != 0 { panic("root permissions needed") } if _, err := rand.Read(multicastKey); err != nil { panic(err) } receivePackets() } func receivePackets() { log.Println("Listening for packets...") listenAddr, err := net.ResolveUDPAddr("udp", ":38000") if err != nil { panic(err) } conn, err = net.ListenUDP("udp", listenAddr) if err != nil { panic(err) } defer conn.Close() for { buffer := make([]byte, 50000) n, addr, err := conn.ReadFromUDP(buffer) if err != nil { log.Println(err) continue } if n == 0 { continue } go handleReq(buffer[:n], addr) } } func handleReq(req []byte, addr *net.UDPAddr) { defer func() { if r := recover(); r != nil { log.Printf("recovered from panic while handling request from %s: %v\n", addr.IP.String(), r) } }() if len(req) < 4 { log.Println("packet too short") return } version := binary.LittleEndian.Uint16(req[0:]) if version != shared.PROTO_VERSION { log.Println("mismatched packet version") return } pktType := binary.LittleEndian.Uint16(req[2:]) switch pktType { case shared.REQ_REGISTER: pubKey, err := hpke.MLKEM768X25519().NewPublicKey(req[4:]) if err != nil { log.Println("invalid pubkey") return } internalIP := randomIP().String() if destPeer, ok := peers.GetOK(addr.IP.String()); ok { internalIP = destPeer.InternalIP } else { peers.Set(addr.IP.String(), &Peer{ RealAddr: addr, InternalIP: internalIP, PublicKey: pubKey, }) } 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()) if peer == nil { log.Println("data from unregistered peer:", addr.String()) return } destIP := net.IP(req[4:20]).String() srcIP := net.IP(req[20:36]).String() if srcIP != peer.InternalIP { log.Println("rejected spoofed srcIP in ENC_PKT") return } if destPeer := getPeerByInternalIP(destIP); destPeer != nil { log.Println(srcIP + " -> " + destIP) sendTo(destPeer.RealAddr, req) } else { log.Println(srcIP + " -/> " + destIP + " (unrecognized IP)") } case shared.BROADCAST_PKT: peer := peers.Get(addr.IP.String()) if peer == nil { log.Println("data from unregistered peer:", addr.String()) return } log.Println(peer.InternalIP + " -> *") broadcast(req, peer.InternalIP) case shared.REQ_GET_PUBKEY: peerIP := net.IP(req[4:20]).String() peer := getPeerByInternalIP(peerIP) if peer == nil { log.Println(addr.String(), "tried to get pubkey of unknown peer", peerIP) return } 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()) if peer == nil { log.Println("unregistered peer", addr.String(), "tried to ESTABLISH") return } srcIP := net.IP(req[20:36]).String() if srcIP != peer.InternalIP { log.Println("rejected spoofed srcIP in ESTABLISH") return } dstIP := net.IP(req[4:20]).String() dstPeer := getPeerByInternalIP(dstIP) sendTo(dstPeer.RealAddr, req) default: log.Println("unknown packet type") } } func broadcast(pkt []byte, except string) { peers.ForEach(func(k string, v *Peer) { if v.InternalIP == except { return } go sendTo(v.RealAddr, pkt) }) } func sendTo(addr *net.UDPAddr, pkt []byte) { if _, err := conn.WriteToUDP(pkt, addr); err != nil { log.Println(err) } } func getPeerByInternalIP(internalIP string) *Peer { var out *Peer = nil peers.ForEach(func(k string, v *Peer) { if v.InternalIP == internalIP { out = v } }) return out } // fd00:baa1::/32 func randomIP() net.IP { ip := make(net.IP, 16) ip[0] = 0xfd ip[1] = 0x00 ip[2] = 0xba ip[3] = 0xa1 if _, err := rand.Read(ip[4:]); err != nil { panic(err) } return ip }