add hmac to REQ_GET_PUBKEY

This commit is contained in:
2026-07-24 15:15:58 +02:00
parent fa9348592e
commit 48811df6ed
4 changed files with 57 additions and 40 deletions

View File

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