381 lines
8.2 KiB
Go
381 lines
8.2 KiB
Go
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 <configPath>")
|
|
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])
|
|
}
|