package main import ( "crypto/sha256" "encoding/binary" "encoding/hex" "log" "net" "os" "os/exec" "sync" "time" "baalvpn/shared" "filippo.io/mlkem768/xwing" "github.com/songgao/water" ) const ( SERVER_IP = "172.20.12.47" IFACE_NAME = "baalvpn" ) var ( internalIP net.IP = nil serverAddr *net.UDPAddr iface *water.Interface conn *net.UDPConn privKey *xwing.DecapsulationKey establishMutex sync.Mutex peerPubKeys = shared.NewTMap[string, []byte]() 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 = xwing.GenerateKey() if err != nil { panic(err) } pubKey := privKey.EncapsulationKey() req := shared.BuildPkt(shared.REQ_REGISTER, pubKey) 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 } log.Println("requesting pubkey of " + ip) req := shared.BuildPkt(shared.REQ_GET_PUBKEY, net.ParseIP(ip)) if _, err := conn.WriteToUDP(req, serverAddr); err != nil { panic(err) } for { // TODO: eww time.Sleep(50 * time.Millisecond) if peerPubKey, ok := peerPubKeys.GetOK(ip); ok { log.Println("got pubkey of " + ip) ciphertext, sharedSecret, err := xwing.Encapsulate(peerPubKey) 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 } } sessionKey := sha256.Sum256(sharedSecret) peerSessionKeys.Set(ip, sessionKey[:]) return sessionKey[:] } } } func receivePackets() { for { buffer := make([]byte, 50000) 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 } go handlePkt(buffer[:n]) } } func handlePkt(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_REGISTER: internalIP = net.IP(pkt[4:]) log.Println("got assigned IP by the server:", internalIP.String()) setupInterface() go sendPackets() case shared.ENC_PKT: srcIP := net.IP(pkt[20:36]).String() sessionKey := getOrEstablishSessionKey(srcIP) decryptedPkt, err := DecryptSym(sessionKey, pkt[36:]) 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: log.Println("received pkt:", hex.EncodeToString(pkt[4:])) if _, err := iface.Write(pkt[4:]); err != nil { log.Println(err) } case shared.RESP_GET_PUBKEY: respIP := net.IP(pkt[4:20]).String() pubKey := pkt[20:] peerPubKeys.Set(respIP, pubKey) case shared.REQ_ESTABLISH: srcIP := net.IP(pkt[20:36]).String() ciphertext := pkt[36:] sharedSecret, err := xwing.Decapsulate(privKey, ciphertext) if err != nil { panic(err) } sessionKey := sha256.Sum256(sharedSecret) peerSessionKeys.Set(srcIP, sessionKey[:]) log.Println("received session key from", srcIP) reqData := append(net.ParseIP(srcIP), internalIP...) req := shared.BuildPkt(shared.RESP_ESTABLISH, reqData) if _, err := conn.WriteToUDP(req, serverAddr); err != nil { panic(err) } case shared.RESP_ESTABLISH: srcIP := net.IP(pkt[20:36]).String() log.Println("established session key with", srcIP) peerEstablishAcks.Set(srcIP, true) default: log.Println("unknown packet type") } } func setupInterface() { config := water.Config{DeviceType: water.TAP} 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 sendPackets() { log.Println("Listening for packets...") for { pkt := make([]byte, 1500) n, err := iface.Read([]byte(pkt)) if err != nil { panic(err) } pkt = pkt[:n] etherType := binary.BigEndian.Uint16(pkt[12:14]) payload := pkt[14:] if etherType != 0x86dd { // IPv6 continue } 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) } } else { destSessionKey := getOrEstablishSessionKey(destIP) encryptedPkt, err := EncryptSym(destSessionKey, pkt) if err != nil { panic(err) } 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) } } } }