package main import ( "crypto/hmac" "crypto/hpke" "crypto/sha256" "encoding/binary" "encoding/hex" "log" "net" "os" "os/exec" "sync" "time" "baalvpn/shared" "github.com/songgao/water" ) // TODO: parse some sort of config+key file const ( SERVER_IP = "172.20.12.47" IFACE_NAME = "baalvpn" ) var ( iface *water.Interface conn *net.UDPConn serverAddr *net.UDPAddr internalIP net.IP privKey hpke.PrivateKey multicastKey []byte challengeCh = make(chan []byte, 1) authSecret []byte establishMutex sync.Mutex peerPubKeys = shared.NewTMap[string, hpke.PublicKey]() peerSessionKeys = shared.NewTMap[string, []byte]() peerEstablishAcks = shared.NewTMap[string, bool]() ) func main() { if os.Geteuid() != 0 { panic("root permissions needed") } 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() serverAddr, err = net.ResolveUDPAddr("udp", SERVER_IP+":38000") if err != nil { panic(err) } go receivePackets() register() <-make(chan int) // block forever } 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 register challenge") send(shared.BuildPkt(shared.REQ_GET_CHALLENGE, pubKey.Bytes())) challenge := <-challengeCh log.Println("got register challenge") authSecret, err = hpke.Open(privKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-challenge"), challenge) if err != nil { panic(err) } log.Println("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() != serverAddr.String() { log.Println("non-server connection rejected") continue } req := make([]byte, n) copy(req, buffer[:n]) go handleIncomingPkt(req) } } func handleIncomingPkt(pkt []byte) { defer func() { if r := recover(); r != nil { log.Printf("recovered from panic while handling a packet: %v\n", r) } }() 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: challengeCh <- pkt[4:] case shared.RESP_REGISTER: internalIP = net.IP(pkt[4:20]) var err error multicastKey, err = hpke.Open(privKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-register"), pkt[20:]) if err != nil { panic(err) } log.Println("got assigned IP by the server:", internalIP.String()) setupInterface() go relayPackets() case shared.ENC_PKT: srcIP := net.IP(pkt[52:68]).String() sessionKey := getOrEstablishSessionKey(srcIP) decryptedPkt, err := DecryptSym(sessionKey, pkt[68:]) if err != nil { panic(err) } log.Println("received pkt:", hex.EncodeToString(decryptedPkt)) if _, err := iface.Write(decryptedPkt); err != nil { log.Println(err) } case shared.BROADCAST_PKT: decryptedPkt, err := DecryptSym(multicastKey, pkt[36:]) 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: respIP := net.IP(pkt[4:20]).String() pubKey, err := hpke.MLKEM768X25519().NewPublicKey(pkt[20:]) 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 } // TODO: this definitely shouldnt block the main thread // TODO: this should try like 10 times tops since server doesnt respond at all if it doesnt have the pubkey 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-establish", 32) if err != nil { panic(err) } send(buildAuthPkt(shared.REQ_ESTABLISH, net.ParseIP(ip), internalIP, ciphertext)) 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.TUN} config.Name = IFACE_NAME var err error iface, err = water.New(config) if err != nil { panic(err) } if err := exec.Command("ip", "addr", "add", internalIP.String()+"/32", "dev", IFACE_NAME).Run(); err != nil { panic(err) } if err := exec.Command("ip", "link", "set", "dev", IFACE_NAME, "up").Run(); err != nil { panic(err) } } func relayPackets() { log.Println("Listening for packets...") for { pkt := make([]byte, 2000) n, err := iface.Read([]byte(pkt)) if err != nil { panic(err) } pkt = pkt[:n] version := pkt[0] >> 4 if version != 6 { 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, encryptedPkt)) } else { destSessionKey := getOrEstablishSessionKey(destIP) encryptedPkt, err := EncryptSym(destSessionKey, pkt) if err != nil { panic(err) } send(buildAuthPkt(shared.ENC_PKT, net.ParseIP(destIP), internalIP, encryptedPkt)) } } } func buildAuthPkt(pktType uint16, parts ...[]byte) []byte { pkt := shared.BuildPkt(pktType, parts...) mac := hmac.New(sha256.New, authSecret) mac.Write(pkt[4:]) return append(pkt[:4], append(mac.Sum(nil), pkt[4:]...)...) } func send(req []byte) { if _, err := conn.WriteToUDP(req, serverAddr); err != nil { panic(err) } }