Files
baalvpn/client/main.go

327 lines
7.2 KiB
Go

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)
}
}