mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-01 19:19:07 +02:00
ML-KEM encapsulate/decapsulate module
This commit is contained in:
@@ -0,0 +1,165 @@
|
|||||||
|
// Package pqkem is a spike (NET-1406) for a post-quantum pre-shared-key exchange
|
||||||
|
// that could replace Rosenpass. It performs an X25519MLKEM768 hybrid key
|
||||||
|
// encapsulation and derives a 32-byte WireGuard PSK.
|
||||||
|
//
|
||||||
|
// The exchange is a single round trip designed to ride the (already
|
||||||
|
// authenticated) Signal offer/answer channel:
|
||||||
|
//
|
||||||
|
// initiator --Offer(1216B)--> responder
|
||||||
|
// initiator <--Answer(1120B)-- responder
|
||||||
|
//
|
||||||
|
// Both sides then hold the same PSK, which is bound to the two peers' identities
|
||||||
|
// (their WireGuard static public keys) so the derived key cannot be transplanted
|
||||||
|
// to a different peer pair even if the transport authentication were bypassed.
|
||||||
|
//
|
||||||
|
// Combiner note: this follows the IETF hybrid layout (X25519 ‖ ML-KEM on the
|
||||||
|
// wire; ML-KEM_ss ‖ X25519_ss into the KDF) from
|
||||||
|
// draft-kwiatkowski-tls-ecdhe-mlkem. The spike uses SHA-256 as the KDF; a
|
||||||
|
// production version should use HKDF with the RFC labels — see TODO below.
|
||||||
|
package pqkem
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/ecdh"
|
||||||
|
"crypto/mlkem"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// OfferSize is the initiator message: X25519 public key ‖ ML-KEM-768 encapsulation key.
|
||||||
|
OfferSize = 32 + mlkem.EncapsulationKeySize768 // 1216
|
||||||
|
// AnswerSize is the responder message: ML-KEM-768 ciphertext ‖ X25519 public key.
|
||||||
|
AnswerSize = mlkem.CiphertextSize768 + 32 // 1120
|
||||||
|
|
||||||
|
pskLabel = "netbird-pq-psk-v1"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PSK is the 32-byte pre-shared key handed to WireGuard.
|
||||||
|
type PSK [32]byte
|
||||||
|
|
||||||
|
// Binding identifies the peer pair the PSK is derived for. Callers set both
|
||||||
|
// WireGuard static public keys; the order does not matter (it is canonicalised).
|
||||||
|
type Binding struct {
|
||||||
|
LocalWgPub []byte
|
||||||
|
RemoteWgPub []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initiator holds the ephemeral secrets between Offer and Finish.
|
||||||
|
type Initiator struct {
|
||||||
|
x25519 *ecdh.PrivateKey
|
||||||
|
mlkemDK *mlkem.DecapsulationKey768
|
||||||
|
offer []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewInitiator generates the ephemeral X25519 + ML-KEM-768 keypairs.
|
||||||
|
func NewInitiator() (*Initiator, error) {
|
||||||
|
x, err := ecdh.X25519().GenerateKey(rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("x25519 keygen: %w", err)
|
||||||
|
}
|
||||||
|
dk, err := mlkem.GenerateKey768()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("ml-kem keygen: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
offer := make([]byte, 0, OfferSize)
|
||||||
|
offer = append(offer, x.PublicKey().Bytes()...)
|
||||||
|
offer = append(offer, dk.EncapsulationKey().Bytes()...)
|
||||||
|
|
||||||
|
return &Initiator{x25519: x, mlkemDK: dk, offer: offer}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Offer returns the initiator message to send over Signal.
|
||||||
|
func (i *Initiator) Offer() []byte {
|
||||||
|
return i.offer
|
||||||
|
}
|
||||||
|
|
||||||
|
// Finish consumes the responder's answer and derives the PSK.
|
||||||
|
func (i *Initiator) Finish(answer []byte, b Binding) (PSK, error) {
|
||||||
|
if len(answer) != AnswerSize {
|
||||||
|
return PSK{}, fmt.Errorf("answer: got %d bytes, want %d", len(answer), AnswerSize)
|
||||||
|
}
|
||||||
|
ct := answer[:mlkem.CiphertextSize768]
|
||||||
|
peerX := answer[mlkem.CiphertextSize768:]
|
||||||
|
|
||||||
|
ssMLKEM, err := i.mlkemDK.Decapsulate(ct)
|
||||||
|
if err != nil {
|
||||||
|
return PSK{}, fmt.Errorf("ml-kem decapsulate: %w", err)
|
||||||
|
}
|
||||||
|
pub, err := ecdh.X25519().NewPublicKey(peerX)
|
||||||
|
if err != nil {
|
||||||
|
return PSK{}, fmt.Errorf("parse peer x25519: %w", err)
|
||||||
|
}
|
||||||
|
ssX, err := i.x25519.ECDH(pub)
|
||||||
|
if err != nil {
|
||||||
|
return PSK{}, fmt.Errorf("x25519 ecdh: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return derivePSK(ssMLKEM, ssX, i.offer, answer, b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Respond consumes an initiator offer, produces the answer, and derives the PSK.
|
||||||
|
func Respond(offer []byte, b Binding) (answer []byte, psk PSK, err error) {
|
||||||
|
if len(offer) != OfferSize {
|
||||||
|
return nil, PSK{}, fmt.Errorf("offer: got %d bytes, want %d", len(offer), OfferSize)
|
||||||
|
}
|
||||||
|
peerX := offer[:32]
|
||||||
|
peerEK := offer[32:]
|
||||||
|
|
||||||
|
ek, err := mlkem.NewEncapsulationKey768(peerEK)
|
||||||
|
if err != nil {
|
||||||
|
return nil, PSK{}, fmt.Errorf("parse peer ml-kem key: %w", err)
|
||||||
|
}
|
||||||
|
ssMLKEM, ct := ek.Encapsulate()
|
||||||
|
|
||||||
|
x, err := ecdh.X25519().GenerateKey(rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
return nil, PSK{}, fmt.Errorf("x25519 keygen: %w", err)
|
||||||
|
}
|
||||||
|
pub, err := ecdh.X25519().NewPublicKey(peerX)
|
||||||
|
if err != nil {
|
||||||
|
return nil, PSK{}, fmt.Errorf("parse peer x25519: %w", err)
|
||||||
|
}
|
||||||
|
ssX, err := x.ECDH(pub)
|
||||||
|
if err != nil {
|
||||||
|
return nil, PSK{}, fmt.Errorf("x25519 ecdh: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
answer = make([]byte, 0, AnswerSize)
|
||||||
|
answer = append(answer, ct...)
|
||||||
|
answer = append(answer, x.PublicKey().Bytes()...)
|
||||||
|
|
||||||
|
// derivePSK uses the same argument order on both sides; the responder's local
|
||||||
|
// binding is the mirror of the initiator's, canonicalised inside derivePSK.
|
||||||
|
return answer, derivePSK(ssMLKEM, ssX, offer, answer, b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// derivePSK combines the two shared secrets and binds the result to the full
|
||||||
|
// transcript (offer ‖ answer) and the canonicalised peer identities.
|
||||||
|
//
|
||||||
|
// TODO(NET-1406): replace the SHA-256 concat with the RFC HKDF combiner
|
||||||
|
// (crypto/hkdf, Go 1.24+) and proper labels before this leaves spike status.
|
||||||
|
func derivePSK(ssMLKEM, ssX, offer, answer []byte, b Binding) PSK {
|
||||||
|
lo, hi := canonicalPair(b.LocalWgPub, b.RemoteWgPub)
|
||||||
|
|
||||||
|
h := sha256.New()
|
||||||
|
h.Write([]byte(pskLabel))
|
||||||
|
h.Write(ssMLKEM)
|
||||||
|
h.Write(ssX)
|
||||||
|
h.Write(offer)
|
||||||
|
h.Write(answer)
|
||||||
|
h.Write(lo)
|
||||||
|
h.Write(hi)
|
||||||
|
|
||||||
|
var psk PSK
|
||||||
|
copy(psk[:], h.Sum(nil))
|
||||||
|
return psk
|
||||||
|
}
|
||||||
|
|
||||||
|
func canonicalPair(a, b []byte) (lo, hi []byte) {
|
||||||
|
if string(a) <= string(b) {
|
||||||
|
return a, b
|
||||||
|
}
|
||||||
|
return b, a
|
||||||
|
}
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
package pqkem
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
wgA = []byte("peer-A-wireguard-pubkey-32bytes!")
|
||||||
|
wgB = []byte("peer-B-wireguard-pubkey-32bytes!")
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExchange_DerivesMatchingPSK(t *testing.T) {
|
||||||
|
init, err := NewInitiator()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.Len(t, init.Offer(), OfferSize)
|
||||||
|
|
||||||
|
answer, pskB, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, answer, AnswerSize)
|
||||||
|
|
||||||
|
pskA, err := init.Finish(answer, Binding{LocalWgPub: wgA, RemoteWgPub: wgB})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.Equal(t, pskB, pskA, "both sides must derive the same PSK")
|
||||||
|
require.NotEqual(t, PSK{}, pskA, "PSK must not be zero")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExchange_PSKBoundToPeerIdentities(t *testing.T) {
|
||||||
|
init, err := NewInitiator()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// responder computes with the honest pair...
|
||||||
|
_, pskHonest, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// ...a second responder run with a different peer identity yields a different PSK,
|
||||||
|
// even though the KEM material would otherwise combine identically.
|
||||||
|
wgC := []byte("peer-C-wireguard-pubkey-32bytes!")
|
||||||
|
_, pskWrong, err := Respond(init.Offer(), Binding{LocalWgPub: wgC, RemoteWgPub: wgA})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NotEqual(t, pskHonest, pskWrong, "PSK must be bound to the peer pair")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExchange_RejectsMalformedMessages(t *testing.T) {
|
||||||
|
init, err := NewInitiator()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, _, err = Respond(init.Offer()[:10], Binding{})
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
_, err = init.Finish([]byte("too short"), Binding{})
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestExchange_ReportSizesAndTiming is a spike measurement, not a pass/fail gate.
|
||||||
|
// Run with: go test -run TestExchange_ReportSizesAndTiming -v ./client/internal/pqkem/
|
||||||
|
func TestExchange_ReportSizesAndTiming(t *testing.T) {
|
||||||
|
const iters = 200
|
||||||
|
|
||||||
|
var tInit, tResp, tFinish time.Duration
|
||||||
|
for i := 0; i < iters; i++ {
|
||||||
|
s0 := time.Now()
|
||||||
|
init, err := NewInitiator()
|
||||||
|
require.NoError(t, err)
|
||||||
|
tInit += time.Since(s0)
|
||||||
|
|
||||||
|
s1 := time.Now()
|
||||||
|
answer, _, err := Respond(init.Offer(), Binding{LocalWgPub: wgB, RemoteWgPub: wgA})
|
||||||
|
require.NoError(t, err)
|
||||||
|
tResp += time.Since(s1)
|
||||||
|
|
||||||
|
s2 := time.Now()
|
||||||
|
_, err = init.Finish(answer, Binding{LocalWgPub: wgA, RemoteWgPub: wgB})
|
||||||
|
require.NoError(t, err)
|
||||||
|
tFinish += time.Since(s2)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("wire sizes: offer=%d B answer=%d B (Rosenpass static pubkey ~524160 B)", OfferSize, AnswerSize)
|
||||||
|
t.Logf("total on-wire per handshake: %d B (~%.0fx smaller than RP static key)", OfferSize+AnswerSize, 524160.0/float64(OfferSize+AnswerSize))
|
||||||
|
t.Logf("avg NewInitiator (keygen): %s", tInit/iters)
|
||||||
|
t.Logf("avg Respond (encaps+dh): %s", tResp/iters)
|
||||||
|
t.Logf("avg Finish (decaps+dh): %s", tFinish/iters)
|
||||||
|
t.Logf("avg full handshake CPU: %s", (tInit+tResp+tFinish)/iters)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user