Files
baalvpn/server/main.go

286 lines
6.4 KiB
Go

// TODO: slice bounds checking
// TODO: authentication
// TODO: somehow persist IPs
// TODO: key rotation
// TODO: dont trust claimed IP at all
// TODO: replay attacks are still a thing
package main
import (
"crypto/hmac"
"crypto/hpke"
"crypto/rand"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"log"
"net"
"os"
"baalvpn/shared"
)
type Peer struct {
RealAddr *net.UDPAddr
InternalIP string
HexPublicKey string
AuthSecret []byte
}
var (
peers = shared.NewTMap[string, *Peer]()
peersByInternal = shared.NewTMap[string, *Peer]()
peersByPubKey = shared.NewTMap[string, *Peer]()
conn *net.UDPConn
multicastKey = make([]byte, 32)
authSecretsByPubKey = shared.NewTMap[string, []byte]()
)
func main() {
if os.Geteuid() != 0 {
panic("root permissions needed")
}
if _, err := rand.Read(multicastKey); err != nil {
panic(err)
}
receivePackets()
}
func receivePackets() {
log.Println("Listening for packets...")
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()
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) {
defer func() {
if r := recover(); r != nil {
log.Printf("recovered from panic while handling request from %s: %v\n", addr.String(), r)
}
}()
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)
}
authSecretsByPubKey.Set(hexPubKey, authSecret)
ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-challenge"), authSecret)
if err != nil {
panic(err)
}
sendTo(addr, shared.BuildPkt(shared.RESP_GET_CHALLENGE, ciphertext))
case shared.REQ_REGISTER:
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 {
panic(err)
}
sendTo(addr, shared.BuildPkt(shared.RESP_REGISTER, net.ParseIP(internalIP), ciphertext))
case shared.ENC_PKT:
peer := peers.Get(addr.String())
if peer == nil {
log.Println("ENC_PKT 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 ENC_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 ENC_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:
peer := peers.Get(addr.String())
if peer == nil {
log.Println("BROADCAST_PKT 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 BROADCAST_PKT")
return
}
log.Println(peer.InternalIP + " -> *")
broadcast(req, peer.InternalIP)
case shared.REQ_GET_PUBKEY:
peerIP := net.IP(req[4:20]).String()
peer := peersByInternal.Get(peerIP)
if peer != nil {
pubKey, err := hex.DecodeString(peer.HexPublicKey)
if err != nil {
panic(err)
}
sendTo(addr, shared.BuildPkt(shared.RESP_GET_PUBKEY, net.ParseIP(peerIP), pubKey))
}
case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH:
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 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)
}
return ip
}