diff --git a/client/internal/pqkem/kem.go b/client/internal/pqkem/kem.go new file mode 100644 index 000000000..c2961eced --- /dev/null +++ b/client/internal/pqkem/kem.go @@ -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 +} diff --git a/client/internal/pqkem/kem_test.go b/client/internal/pqkem/kem_test.go new file mode 100644 index 000000000..3bdc99635 --- /dev/null +++ b/client/internal/pqkem/kem_test.go @@ -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) +}