package main import ( "crypto/hmac" "crypto/hpke" "crypto/sha256" "encoding/binary" "encoding/hex" "fmt" "log" "net" "os" "sync" "time" "baalvpn/shared" "github.com/cloudflare/circl/sign/mldsa/mldsa44" ) type Config struct { ListenAddr *net.UDPAddr ServerAddr *net.UDPAddr ServerPubKey mldsa44.PublicKey } var ( config Config conn *net.UDPConn internalIP net.IP privKey hpke.PrivateKey multicastKey []byte authSecret []byte challengeCh = make(chan []byte, 1) establishMutex sync.Mutex peerPubKeys = shared.NewTMap[string, hpke.PublicKey]() peerSessionKeys = shared.NewTMap[string, []byte]() peerEstablishAcks = shared.NewTMap[string, bool]() ) func main() { if len(os.Args) != 2 { fmt.Fprintln(os.Stderr, "Usage: baalvpn-client ") os.Exit(1) } parseConfig() if !isAdmin() { panic("this program must be ran with administrative privileges") } var err error conn, err = net.ListenUDP("udp", config.ListenAddr) if err != nil { panic(err) } defer conn.Close() go receivePackets() register() <-make(chan int) // block forever } func parseConfig() { conf, err := shared.ParseConfFile(os.Args[1]) if err != nil { panic(err) } pubKeyBytes, err := hex.DecodeString(conf["SERVER_PUBKEY"]) if err != nil { panic(err) } var serverPubKey mldsa44.PublicKey if err := serverPubKey.UnmarshalBinary(pubKeyBytes); err != nil { panic(err) } listenAddr, err := net.ResolveUDPAddr("udp", conf["LISTEN_ADDR"]) if err != nil { panic(err) } serverAddr, err := net.ResolveUDPAddr("udp", conf["SERVER_ADDR"]) if err != nil { panic(err) } config = Config{ ListenAddr: listenAddr, ServerAddr: serverAddr, ServerPubKey: serverPubKey, } } func register() { var err error privKey, err = hpke.MLKEM768X25519().GenerateKey() if err != nil { panic(err) } pubKey := privKey.PublicKey() // TODO: all this probably should be retried until we get RESP_REGISTER log.Println("requesting authSecret") send(shared.BuildPkt(shared.REQ_GET_CHALLENGE, pubKey.Bytes())) challenge := <-challengeCh authSecret, err = hpke.Open(privKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-challenge"), challenge) if err != nil { panic(err) } log.Println("received authSecret, registering...") mac := hmac.New(sha256.New, authSecret) mac.Write(pubKey.Bytes()) send(shared.BuildPkt(shared.REQ_REGISTER, mac.Sum(nil), pubKey.Bytes())) } func receivePackets() { buffer := make([]byte, 50000) for { n, addr, err := conn.ReadFromUDP(buffer[:]) if err != nil { log.Println(err) continue } if n == 0 { continue } if addr.String() != config.ServerAddr.String() { log.Println("non-server connection rejected") continue } req := make([]byte, n) copy(req, buffer[:n]) go handleIncomingPkt(req) } } func handleIncomingPkt(pkt []byte) { if len(pkt) <= 4 { log.Println("packet too short") return } version := binary.LittleEndian.Uint16(pkt[0:]) if version != shared.PROTO_VERSION { log.Println("mismatched packet version") return } pktType := binary.LittleEndian.Uint16(pkt[2:]) switch pktType { case shared.RESP_GET_CHALLENGE: if !verifyPkt(pkt) { log.Println("failed to verify signature") return } challengeCh <- pkt[2424:] case shared.RESP_REGISTER: if !verifyPkt(pkt) { log.Println("failed to verify signature") return } internalIP = net.IP(pkt[2424:2440]) var err error multicastKey, err = hpke.Open(privKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-register"), pkt[2440:]) if err != nil { panic(err) } log.Println("got assigned IP by the server:", internalIP.String()) setupInterface() go relayPackets() case shared.UNICAST_PKT: srcIP := net.IP(pkt[52:68]).String() sessionKey := getOrEstablishSessionKey(srcIP) if sessionKey == nil { return } decryptedPkt, err := DecryptSym(sessionKey, pkt[68:]) if err != nil { panic(err) } if len(decryptedPkt) >= 24 { innerSrcIP := net.IP(decryptedPkt[8:24]).String() if innerSrcIP != srcIP { return } } log.Println("received pkt:", hex.EncodeToString(decryptedPkt)) if err := writePkt(decryptedPkt); err != nil { log.Println(err) } case shared.BROADCAST_PKT: decryptedPkt, err := DecryptSym(multicastKey, pkt[52:]) if err != nil { panic(err) } if len(decryptedPkt) >= 24 { outerSrcIP := net.IP(pkt[36:52]).String() innerSrcIP := net.IP(decryptedPkt[8:24]).String() if innerSrcIP != outerSrcIP { return } } log.Println("received broadcast:", hex.EncodeToString(decryptedPkt)) if err := writePkt(decryptedPkt); err != nil { log.Println(err) } case shared.RESP_GET_PUBKEY: if !verifyPkt(pkt) { log.Println("failed to verify signature") return } respIP := net.IP(pkt[2424:2440]).String() pubKey, err := hpke.MLKEM768X25519().NewPublicKey(pkt[2440:]) if err != nil { panic(err) } peerPubKeys.Set(respIP, pubKey) case shared.REQ_ESTABLISH: srcIP := net.IP(pkt[52:68]).String() ciphertext := pkt[68:] r, err := hpke.NewRecipient(ciphertext, privKey, hpke.HKDFSHA256(), hpke.ExportOnly(), nil) if err != nil { panic(err) } sessionKey, err := r.Export("baalvpn-establish", 32) if err != nil { panic(err) } log.Println("received session key from", srcIP) peerSessionKeys.Set(srcIP, sessionKey) send(buildAuthPkt(shared.RESP_ESTABLISH, net.ParseIP(srcIP), internalIP)) case shared.RESP_ESTABLISH: srcIP := net.IP(pkt[52:68]).String() log.Println("established session key with", srcIP) peerEstablishAcks.Set(srcIP, true) default: log.Println("unknown packet type") } } 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 } log.Println("lets try to get " + ip + "'s pubkey") // TODO: this definitely shouldnt block the main thread for range 30 { // request pubkey every 50ms time.Sleep(50 * time.Millisecond) peerPubKey, ok := peerPubKeys.GetOK(ip) if !ok { send(buildAuthPkt(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-establish", 32) if err != nil { panic(err) } send(buildAuthPkt(shared.REQ_ESTABLISH, net.ParseIP(ip), internalIP, ciphertext)) success := false for range 30 { // TODO: eww time.Sleep(50 * time.Millisecond) if peerEstablishAcks.Get(ip) { peerEstablishAcks.Delete(ip) success = true break } } if !success { log.Println("failed to establish") return nil } peerSessionKeys.Set(ip, sessionKey) return sessionKey } log.Println("failed to get the pubkey") return nil } func relayPackets() { log.Println("Listening for packets...") for { pkt, err := readPkt() if err != nil { panic(err) } version := pkt[0] >> 4 if version != 6 { releasePkt(pkt) continue } destIP := net.IP(pkt[24:40]).String() if pkt[24] == 0xff { // multicast encryptedPkt, err := EncryptSym(multicastKey, pkt) if err != nil { panic(err) } send(buildAuthPkt(shared.BROADCAST_PKT, internalIP, encryptedPkt)) } else { destSessionKey := getOrEstablishSessionKey(destIP) if destSessionKey == nil { releasePkt(pkt) continue } encryptedPkt, err := EncryptSym(destSessionKey, pkt) if err != nil { panic(err) } send(buildAuthPkt(shared.UNICAST_PKT, net.ParseIP(destIP), internalIP, encryptedPkt)) } releasePkt(pkt) } } func buildAuthPkt(pktType uint16, parts ...[]byte) []byte { pkt := shared.BuildPkt(pktType, parts...) mac := hmac.New(sha256.New, authSecret) // this ignores version and packet type, fine for now mac.Write(pkt[4:]) return append(pkt[:4], append(mac.Sum(nil), pkt[4:]...)...) } func send(req []byte) { if _, err := conn.WriteToUDP(req, config.ServerAddr); err != nil { panic(err) } } func verifyPkt(pkt []byte) bool { // this ignores version and packet type, fine for now return mldsa44.Verify(&config.ServerPubKey, pkt[2424:], nil, pkt[4:2424]) }