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