diff --git a/.gitignore b/.gitignore index 7a187b8..8c3970a 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1 @@ Justfile -/test diff --git a/client/main.go b/client/main.go index 2f6526e..5fa8d88 100644 --- a/client/main.go +++ b/client/main.go @@ -118,13 +118,7 @@ func receivePackets() { } 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 { + if len(pkt) <= 4 { log.Println("packet too short") return } @@ -138,13 +132,13 @@ func handleIncomingPkt(pkt []byte) { switch pktType { case shared.RESP_GET_CHALLENGE: - if !verifyPacket(pkt) { + if !verifyPkt(pkt) { log.Println("failed to verify signature") return } challengeCh <- pkt[2424:] case shared.RESP_REGISTER: - if !verifyPacket(pkt) { + if !verifyPkt(pkt) { log.Println("failed to verify signature") return } @@ -183,7 +177,7 @@ func handleIncomingPkt(pkt []byte) { log.Println(err) } case shared.RESP_GET_PUBKEY: - if !verifyPacket(pkt) { + if !verifyPkt(pkt) { log.Println("failed to verify signature") return } @@ -239,7 +233,7 @@ func getOrEstablishSessionKey(ip string) []byte { time.Sleep(50 * time.Millisecond) peerPubKey, ok := peerPubKeys.GetOK(ip) if !ok { - send(shared.BuildPkt(shared.REQ_GET_PUBKEY, net.ParseIP(ip))) + send(buildAuthPkt(shared.REQ_GET_PUBKEY, net.ParseIP(ip))) continue } @@ -329,6 +323,7 @@ func relayPackets() { 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:]...)...) } @@ -339,15 +334,16 @@ func send(req []byte) { } } -func verifyPacket(pkt []byte) bool { +func verifyPkt(pkt []byte) bool { pubKeyBytes, err := hex.DecodeString(SERVER_PUBKEY) if err != nil { - panic(err) + panic(err) // fatal misconfiguration } var pubKey mldsa44.PublicKey if err := pubKey.UnmarshalBinary(pubKeyBytes); err != nil { - panic(err) + panic(err) // fatal misconfiguration } + // this ignores version and packet type, fine for now return mldsa44.Verify(&pubKey, pkt[2424:], nil, pkt[4:2424]) } diff --git a/server/main.go b/server/main.go index 32006bd..cf6a90a 100644 --- a/server/main.go +++ b/server/main.go @@ -1,4 +1,3 @@ -// TODO: slice bounds checking // TODO: somehow persist internal IPs // TODO: key rotation // TODO: dont trust claimed IP at all @@ -45,7 +44,7 @@ func main() { if len(os.Args) > 1 && os.Args[1] == "keygen" { pubKey, privKey, err := mldsa44.GenerateKey(rand.Reader) if err != nil { - panic(err) + panic(err) // cant continue } fmt.Println("pubKey: " + hex.EncodeToString(pubKey.Bytes())) fmt.Println("privKey: " + hex.EncodeToString(privKey.Bytes())) @@ -57,7 +56,7 @@ func main() { } if _, err := rand.Read(multicastKey); err != nil { - panic(err) + panic(err) // should never happen } receivePackets() @@ -68,11 +67,11 @@ func receivePackets() { listenAddr, err := net.ResolveUDPAddr("udp", ":38000") if err != nil { - panic(err) + panic(err) // cant continue } conn, err = net.ListenUDP("udp", listenAddr) if err != nil { - panic(err) + panic(err) // cant continue } defer conn.Close() @@ -94,13 +93,7 @@ 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.String(), r) - } - }() - - if len(req) < 4 { + if len(req) <= 4 { log.Println("packet too short") return } @@ -124,18 +117,23 @@ func handleReq(req []byte, addr *net.UDPAddr) { authSecret := make([]byte, 32) if _, err = rand.Read(authSecret); err != nil { - panic(err) + panic(err) // should never happen } authSecretsByPubKey.Set(hexPubKey, authSecret) ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-challenge"), authSecret) if err != nil { - panic(err) + log.Println("failed to Seal the authSecret:", err) + return } sendTo(addr, buildSignedPkt(shared.RESP_GET_CHALLENGE, ciphertext)) case shared.REQ_REGISTER: + if len(req) <= 36 { + log.Println("REGISTER packet too short") + return + } auth := req[4:36] pubKey, err := hpke.MLKEM768X25519().NewPublicKey(req[36:]) @@ -182,28 +180,33 @@ func handleReq(req []byte, addr *net.UDPAddr) { ciphertext, err := hpke.Seal(pubKey, hpke.HKDFSHA256(), hpke.ChaCha20Poly1305(), []byte("baalvpn-register"), multicastKey) if err != nil { - panic(err) + log.Println("failed to Seal the multicastKey:", err) + return } sendTo(addr, buildSignedPkt(shared.RESP_REGISTER, net.ParseIP(internalIP), ciphertext)) case shared.UNICAST_PKT: + if len(req) < 68 { + log.Println("UNICAST_PKT packet too short") + return + } peer := peers.Get(addr.String()) if peer == nil { - log.Println("ENC_PKT from unregistered peer:", addr.String()) + log.Println("UNICAST_PKT from an unregistered peer") 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") + log.Println("failed to authenticate UNICAST_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") + log.Println("rejected spoofed srcIP in UNICAST_PKT") return } @@ -214,6 +217,10 @@ func handleReq(req []byte, addr *net.UDPAddr) { log.Println(srcIP + " -/> " + destIP + " (unrecognized IP)") } case shared.BROADCAST_PKT: + if len(req) <= 36 { + log.Println("BROADCAST_PKT packet too short") + return + } peer := peers.Get(addr.String()) if peer == nil { log.Println("BROADCAST_PKT from unregistered peer:", addr.String()) @@ -230,16 +237,31 @@ func handleReq(req []byte, addr *net.UDPAddr) { log.Println(peer.InternalIP + " -> *") broadcast(req, peer.InternalIP) case shared.REQ_GET_PUBKEY: - peerIP := net.IP(req[4:20]).String() + if len(req) < 52 { + log.Println("REQ_GET_PUBKEY packet too short") + return + } + peerIP := net.IP(req[36:52]).String() peer := peersByInternal.Get(peerIP) if peer != nil { + 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 REQ_GET_PUBKEY") + return + } + pubKey, err := hex.DecodeString(peer.HexPublicKey) if err != nil { - panic(err) + panic(err) // should never happen } sendTo(addr, buildSignedPkt(shared.RESP_GET_PUBKEY, net.ParseIP(peerIP), pubKey)) } case shared.REQ_ESTABLISH, shared.RESP_ESTABLISH: + if len(req) < 68 { + log.Println("ESTABLISH packet too short") + return + } peer := peers.Get(addr.String()) if peer == nil { log.Println("ESTABLISH from unregistered peer:", addr.String()) @@ -276,16 +298,16 @@ func buildSignedPkt(pktType uint16, parts ...[]byte) []byte { privKeyBytes, err := hex.DecodeString(SERVER_PRIVKEY) if err != nil { - panic(err) + panic(err) // fatal misconfiguration } var privKey mldsa44.PrivateKey if err := privKey.UnmarshalBinary(privKeyBytes); err != nil { - panic(err) + panic(err) // fatal misconfiguration } sig, err := privKey.Sign(rand.Reader, pkt[4:], nil) if err != nil { - panic(err) + panic(err) // should never happen } return append(pkt[:4], append(sig, pkt[4:]...)...) @@ -314,7 +336,7 @@ func randomIP() net.IP { ip[2] = 0xba ip[3] = 0xa1 if _, err := rand.Read(ip[4:]); err != nil { - panic(err) + panic(err) // should never happen } return ip } diff --git a/shared/proto.go b/shared/proto.go index f6165ef..80d5ccd 100644 --- a/shared/proto.go +++ b/shared/proto.go @@ -34,7 +34,7 @@ const ( // [ ds.Sign(serverPrivKey, rest) - 2420 bytes ] [ ip - 16 bytes ] [ hpke.Seal(pubKey, multicastKey) - 1168 bytes ] RESP_REGISTER // (client -> server) requests peer's pubKey from the server for encapsulation - // [ ip - 16 bytes ] + // [ HMAC(authSecret, rest) - 32 bytes ] [ ip - 16 bytes ] REQ_GET_PUBKEY // (server -> client) provides requested pubKey // [ ds.Sign(serverPrivKey, rest) - 2420 bytes ] [ ip - 16 bytes ] [ pubKey - 1216 bytes ]