From a2d647f3afd3a45c380b1a0424e6793b99815e2a Mon Sep 17 00:00:00 2001 From: Toni Date: Wed, 22 Jul 2026 14:28:11 +0200 Subject: [PATCH] thread safe map --- client/main.go | 99 ++++++++++++++++++++++++++++---------------------- server/main.go | 63 +++++++++++++++----------------- shared/tmap.go | 47 ++++++++++++++++++++++++ 3 files changed, 132 insertions(+), 77 deletions(-) create mode 100644 shared/tmap.go diff --git a/client/main.go b/client/main.go index e6bf140..b5a4f3c 100644 --- a/client/main.go +++ b/client/main.go @@ -8,6 +8,7 @@ import ( "net" "os" "os/exec" + "sync" "time" "baalvpn/shared" @@ -28,12 +29,10 @@ var ( conn *net.UDPConn privKey *xwing.DecapsulationKey - // TODO: this should probably have a mutex - peerPubKeys = map[string][]byte{} - // TODO: this should probably have a mutex - peerSessionKeys = map[string][]byte{} - // TODO: this should probably have a mutex - peerEstablishAcks = map[string]bool{} + establishMutex sync.Mutex + peerPubKeys = shared.NewTMap[string, []byte]() + peerSessionKeys = shared.NewTMap[string, []byte]() + peerEstablishAcks = shared.NewTMap[string, bool]() ) func main() { @@ -77,45 +76,53 @@ func register() { } func getOrEstablishSessionKey(ip string) []byte { - if key, ok := peerSessionKeys[ip]; ok { + if key, ok := peerSessionKeys.GetOK(ip); ok { return key - } else { - 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[ip]; ok { - log.Println("got pubkey of " + ip) + establishMutex.Lock() + defer establishMutex.Unlock() - ciphertext, sharedSecret, err := xwing.Encapsulate(peerPubKey) - if err != nil { - panic(err) - } + // check again after the other goroutine finished + if key, ok := peerSessionKeys.GetOK(ip); ok { + return key + } - 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) - } + 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 peerEstablishAcks[ip] { - delete(peerEstablishAcks, ip) - break - } - } + for { + // TODO: eww + time.Sleep(50 * time.Millisecond) + if peerPubKey, ok := peerPubKeys.GetOK(ip); ok { + log.Println("got pubkey of " + ip) - sessionKey := sha256.Sum256(sharedSecret) - peerSessionKeys[ip] = sessionKey[:] - return sessionKey[:] + 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[:] } } } @@ -136,11 +143,17 @@ func receivePackets() { continue } - go receivePacket(buffer[:n]) + go handlePkt(buffer[:n]) } } -func receivePacket(pkt []byte) { +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 @@ -181,7 +194,7 @@ func receivePacket(pkt []byte) { case shared.RESP_GET_PUBKEY: respIP := net.IP(pkt[4:20]).String() pubKey := pkt[20:] - peerPubKeys[respIP] = pubKey + peerPubKeys.Set(respIP, pubKey) case shared.REQ_ESTABLISH: srcIP := net.IP(pkt[20:36]).String() ciphertext := pkt[36:] @@ -192,7 +205,7 @@ func receivePacket(pkt []byte) { } sessionKey := sha256.Sum256(sharedSecret) - peerSessionKeys[srcIP] = sessionKey[:] + peerSessionKeys.Set(srcIP, sessionKey[:]) log.Println("received session key from", srcIP) reqData := append(net.ParseIP(srcIP), internalIP...) @@ -203,7 +216,7 @@ func receivePacket(pkt []byte) { case shared.RESP_ESTABLISH: srcIP := net.IP(pkt[20:36]).String() log.Println("established session key with", srcIP) - peerEstablishAcks[srcIP] = true + peerEstablishAcks.Set(srcIP, true) default: log.Println("unknown packet type") } diff --git a/server/main.go b/server/main.go index aa7e54a..937a373 100644 --- a/server/main.go +++ b/server/main.go @@ -4,7 +4,6 @@ // TODO: handle multiple peers behind one NAT // TODO: key rotation // TODO: sha256 -> HKDF -// TODO: recovery package main import ( @@ -13,7 +12,6 @@ import ( "log" "net" "os" - "sync" "baalvpn/shared" ) @@ -25,9 +23,8 @@ type Peer struct { } var ( - peers = map[string]*Peer{} - peersMutex sync.Mutex - conn *net.UDPConn + peers = shared.NewTMap[string, *Peer]() + conn *net.UDPConn ) func main() { @@ -67,6 +64,12 @@ func receivePackets() { } 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.IP.String(), r) + } + }() + if len(req) < 4 { log.Println("packet too short") return @@ -81,29 +84,24 @@ func handleReq(req []byte, addr *net.UDPAddr) { switch pktType { case shared.REQ_REGISTER: - peersMutex.Lock() - defer peersMutex.Unlock() - pubKey := req[4:] internalIP := randomIP().String() - if destPeer, ok := peers[addr.IP.String()]; ok { + if destPeer, ok := peers.GetOK(addr.IP.String()); ok { internalIP = destPeer.InternalIP } else { - peers[addr.IP.String()] = &Peer{ + peers.Set(addr.IP.String(), &Peer{ RealAddr: addr, InternalIP: internalIP, PublicKey: pubKey, - } + }) } resp := shared.BuildPkt(shared.RESP_REGISTER, net.ParseIP(internalIP)) sendTo(addr, resp) case shared.ENC_PKT: - peersMutex.Lock() - peer := peers[addr.IP.String()] - peersMutex.Unlock() + peer := peers.Get(addr.IP.String()) if peer == nil { log.Println("data from unregistered peer:", addr.String()) return @@ -123,9 +121,7 @@ func handleReq(req []byte, addr *net.UDPAddr) { log.Println(srcIP + " -/> " + destIP + " (unrecognized IP)") } case shared.BROADCAST_PKT: - peersMutex.Lock() - peer := peers[addr.IP.String()] - peersMutex.Unlock() + peer := peers.Get(addr.IP.String()) if peer == nil { log.Println("data from unregistered peer:", addr.String()) return @@ -145,9 +141,11 @@ func handleReq(req []byte, addr *net.UDPAddr) { resp := shared.BuildPkt(shared.RESP_GET_PUBKEY, append(net.ParseIP(peer.InternalIP), peer.PublicKey...)) sendTo(addr, resp) case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH: - peersMutex.Lock() - peer := peers[addr.IP.String()] - peersMutex.Unlock() + peer := peers.Get(addr.IP.String()) + if peer == nil { + log.Println("unregistered peer", addr.String(), "tried to ESTABLISH") + return + } srcIP := net.IP(req[20:36]).String() if srcIP != peer.InternalIP { @@ -164,14 +162,12 @@ func handleReq(req []byte, addr *net.UDPAddr) { } func broadcast(pkt []byte, except string) { - peersMutex.Lock() - defer peersMutex.Unlock() - for _, p := range peers { - if p.InternalIP == except { - continue + peers.ForEach(func(k string, v *Peer) { + if v.InternalIP == except { + return } - sendTo(p.RealAddr, pkt) - } + go sendTo(v.RealAddr, pkt) + }) } func sendTo(addr *net.UDPAddr, pkt []byte) { @@ -181,14 +177,13 @@ func sendTo(addr *net.UDPAddr, pkt []byte) { } func getPeerByInternalIP(internalIP string) *Peer { - peersMutex.Lock() - defer peersMutex.Unlock() - for _, p := range peers { - if p.InternalIP == internalIP { - return p + var out *Peer = nil + peers.ForEach(func(k string, v *Peer) { + if v.InternalIP == internalIP { + out = v } - } - return nil + }) + return out } // fd00:baa1::/32 diff --git a/shared/tmap.go b/shared/tmap.go new file mode 100644 index 0000000..da2a963 --- /dev/null +++ b/shared/tmap.go @@ -0,0 +1,47 @@ +package shared + +import "sync" + +type TMap[K comparable, V any] struct { + data map[K]V + mu sync.RWMutex +} + +func NewTMap[K comparable, V any]() TMap[K, V] { + return TMap[K, V]{ + data: map[K]V{}, + } +} + +func (m *TMap[K, V]) Get(k K) V { + m.mu.RLock() + defer m.mu.RUnlock() + return m.data[k] +} + +func (m *TMap[K, V]) GetOK(k K) (V, bool) { + m.mu.RLock() + defer m.mu.RUnlock() + v, ok := m.data[k] + return v, ok +} + +func (m *TMap[K, V]) Set(k K, v V) { + m.mu.Lock() + defer m.mu.Unlock() + m.data[k] = v +} + +func (m *TMap[K, V]) Delete(k K) { + m.mu.Lock() + defer m.mu.Unlock() + delete(m.data, k) +} + +func (m *TMap[K, V]) ForEach(f func(k K, v V)) { + m.mu.RLock() + defer m.mu.RUnlock() + for k, v := range m.data { + f(k, v) + } +}