// TODO: somehow persist internal IPs // TODO: key rotation // TODO: replay attacks are still a thing package main import ( "crypto/hmac" "crypto/hpke" "crypto/rand" "crypto/sha256" "encoding/binary" "encoding/hex" "fmt" "log" "net" "os" "baalvpn/shared" "github.com/cloudflare/circl/sign/mldsa/mldsa44" ) type Peer struct { RealAddr *net.UDPAddr InternalIP string HexPublicKey string AuthSecret []byte } var ( conn *net.UDPConn multicastKey = make([]byte, 32) authSecretsByPubKey = shared.NewTMap[string, []byte]() serverPrivKey mldsa44.PrivateKey peers = shared.NewTMap[string, *Peer]() peersByInternal = shared.NewTMap[string, *Peer]() peersByPubKey = shared.NewTMap[string, *Peer]() ) func main() { if len(os.Args) != 2 { fmt.Fprintln(os.Stderr, "Usage: baalvpn-server ") fmt.Fprintln(os.Stderr, " baalvpn-server keygen") os.Exit(1) } if os.Args[1] == "keygen" { pubKey, privKey, err := mldsa44.GenerateKey(rand.Reader) if err != nil { panic(err) } if err := os.WriteFile("client.conf", fmt.Appendf(nil, "LISTEN_ADDR=:38000\nSERVER_ADDR=CHANGEME:38000\nSERVER_PUBKEY=%s", hex.EncodeToString(pubKey.Bytes())), 0600); err != nil { panic(err) } if err := os.WriteFile("server.conf", fmt.Appendf(nil, "SERVER_PRIVKEY=%s", hex.EncodeToString(privKey.Bytes())), 0600); err != nil { panic(err) } fmt.Println("client.conf and server.conf generated") return } parseConfig() if os.Geteuid() != 0 { panic("root permissions needed") } if _, err := rand.Read(multicastKey); err != nil { panic(err) // should never happen } receivePackets() } func parseConfig() { conf, err := shared.ParseConfFile(os.Args[1]) if err != nil { panic(err) } privKeyBytes, err := hex.DecodeString(conf["SERVER_PRIVKEY"]) if err != nil { panic(err) } if err := serverPrivKey.UnmarshalBinary(privKeyBytes); err != nil { panic(err) } } func receivePackets() { log.Println("Listening for packets...") listenAddr, err := net.ResolveUDPAddr("udp", ":38000") if err != nil { panic(err) // cant continue } conn, err = net.ListenUDP("udp", listenAddr) if err != nil { panic(err) // cant continue } defer conn.Close() buffer := make([]byte, 50000) for { n, addr, err := conn.ReadFromUDP(buffer) if err != nil { log.Println(err) continue } if n == 0 { continue } req := make([]byte, n) copy(req, buffer[:n]) go handleReq(req, addr) } } func handleReq(req []byte, addr *net.UDPAddr) { 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_GET_CHALLENGE: pubKey, err := hpke.MLKEM768X25519().NewPublicKey(req[4:]) if err != nil { log.Println("invalid pubkey") return } hexPubKey := hex.EncodeToString(pubKey.Bytes()) authSecret := make([]byte, 32) if _, err = rand.Read(authSecret); err != nil { panic(err) // should never happen } authSecretsByPubKey.Set(hexPubKey, authSecret) ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-challenge"), authSecret) if err != nil { log.Println("failed to Seal the authSecret:", err) return } sendTo(addr, buildSignedPkt(shared.RESP_GET_CHALLENGE, ciphertext)) case shared.REQ_REGISTER: if len(req) <= 36 { log.Println("REGISTER packet too short") return } auth := req[4:36] pubKey, err := hpke.MLKEM768X25519().NewPublicKey(req[36:]) if err != nil { log.Println("invalid pubkey") return } hexPubKey := hex.EncodeToString(pubKey.Bytes()) authSecret, ok := authSecretsByPubKey.GetOK(hexPubKey) if !ok { log.Println("unknown pubkey tried to REGISTER") return } defer authSecretsByPubKey.Delete(hexPubKey) mac := hmac.New(sha256.New, authSecret) mac.Write(pubKey.Bytes()) if !hmac.Equal(mac.Sum(nil), auth) { log.Println("failed to authenticate REGISTER") return } var internalIP string if peer, ok := peersByPubKey.GetOK(hexPubKey); ok { internalIP = peer.InternalIP peer.AuthSecret = authSecret peers.Delete(peer.RealAddr.String()) peer.RealAddr = addr peers.Set(addr.String(), peer) } else { internalIP = randomIP().String() peer := &Peer{ RealAddr: addr, InternalIP: internalIP, HexPublicKey: hexPubKey, AuthSecret: authSecret, } peers.Set(addr.String(), peer) peersByInternal.Set(internalIP, peer) peersByPubKey.Set(hexPubKey, peer) } ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-register"), multicastKey) if err != nil { log.Println("failed to Seal the multicastKey:", err) return } sendTo(addr, buildSignedPkt(shared.RESP_REGISTER, net.ParseIP(internalIP), ciphertext)) case shared.UNICAST_PKT: if len(req) < 68 { log.Println("UNICAST_PKT packet too short") return } peer := peers.Get(addr.String()) if peer == nil { log.Println("UNICAST_PKT from an unregistered peer") return } mac := hmac.New(sha256.New, peer.AuthSecret) mac.Write(req[36:]) if !hmac.Equal(mac.Sum(nil), req[4:36]) { log.Println("failed to authenticate UNICAST_PKT") return } destIP := net.IP(req[36:52]).String() srcIP := net.IP(req[52:68]).String() if srcIP != peer.InternalIP { log.Println("rejected spoofed srcIP in UNICAST_PKT") return } if destPeer := peersByInternal.Get(destIP); destPeer != nil { log.Println(srcIP + " -> " + destIP) sendTo(destPeer.RealAddr, req) } else { log.Println(srcIP + " -/> " + destIP + " (unrecognized IP)") } case shared.BROADCAST_PKT: if len(req) <= 52 { log.Println("BROADCAST_PKT packet too short") return } peer := peers.Get(addr.String()) if peer == nil { log.Println("BROADCAST_PKT from unregistered peer:", addr.String()) return } srcIP := net.IP(req[36:52]).String() if srcIP != peer.InternalIP { log.Println("rejected spoofed srcIP in BROADCAST_PKT") return } mac := hmac.New(sha256.New, peer.AuthSecret) mac.Write(req[36:]) if !hmac.Equal(mac.Sum(nil), req[4:36]) { log.Println("failed to authenticate BROADCAST_PKT") return } log.Println(peer.InternalIP + " -> *") broadcast(req, peer.InternalIP) case shared.REQ_GET_PUBKEY: if len(req) < 52 { log.Println("REQ_GET_PUBKEY packet too short") return } peer := peers.Get(addr.String()) if peer == nil { log.Println("REQ_GET_PUBKEY from unregistered peer:", addr.String()) return } mac := hmac.New(sha256.New, peer.AuthSecret) mac.Write(req[36:]) if !hmac.Equal(mac.Sum(nil), req[4:36]) { log.Println("failed to authenticate REQ_GET_PUBKEY") return } peerIP := net.IP(req[36:52]).String() targetPeer := peersByInternal.Get(peerIP) if targetPeer != nil { pubKey, err := hex.DecodeString(targetPeer.HexPublicKey) if err != nil { panic(err) // should never happen } sendTo(addr, buildSignedPkt(shared.RESP_GET_PUBKEY, net.ParseIP(peerIP), pubKey)) } case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH: if len(req) < 68 { log.Println("ESTABLISH packet too short") return } peer := peers.Get(addr.String()) if peer == nil { log.Println("ESTABLISH from unregistered peer:", addr.String()) return } mac := hmac.New(sha256.New, peer.AuthSecret) mac.Write(req[36:]) if !hmac.Equal(mac.Sum(nil), req[4:36]) { log.Println("failed to authenticate ESTABLISH") return } srcIP := net.IP(req[52:68]).String() if srcIP != peer.InternalIP { log.Println("rejected spoofed srcIP in ESTABLISH") return } dstIP := net.IP(req[36:52]).String() dstPeer := peersByInternal.Get(dstIP) if dstPeer == nil { log.Println("tried to ESTABLISH with an unknown peer") return } sendTo(dstPeer.RealAddr, req) default: log.Println("unknown packet type") } } func buildSignedPkt(pktType uint16, parts ...[]byte) []byte { pkt := shared.BuildPkt(pktType, parts...) sig, err := serverPrivKey.Sign(rand.Reader, pkt[4:], nil) if err != nil { panic(err) // should never happen } return append(pkt[:4], append(sig, pkt[4:]...)...) } 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) } } // 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) // should never happen } return ip }