[client] Cache the box shared key per remote peer in the Signal client (#7807)

This commit is contained in:
Viktor Liu
2026-09-30 17:03:04 +02:00
committed by GitHub
parent e72be6698f
commit 0dc729c4ea
5 changed files with 352 additions and 2 deletions
+147
View File
@@ -0,0 +1,147 @@
package encryption
import (
"fmt"
"sync"
pb "github.com/golang/protobuf/proto" //nolint
"golang.org/x/crypto/nacl/box"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// SharedKeyCache encrypts and decrypts messages for one local private key, deriving
// the box shared key once per remote public key instead of once per message.
//
// The shared key is a pure function of the two keys, so a cached entry never goes
// stale: a different remote key is a different entry, and a different local key
// needs a different cache. Entries are only dropped to stay under maxSharedKeys.
// Every message still uses its own random nonce.
//
// The cached values are secret key material, as sensitive as the private key.
type SharedKeyCache struct {
privateKey wgtypes.Key
limit int
mu sync.RWMutex
keys map[wgtypes.Key]*[32]byte
closed bool
}
// NewSharedKeyCache returns a cache for messages sent and received with privateKey.
func NewSharedKeyCache(privateKey wgtypes.Key) *SharedKeyCache {
return &SharedKeyCache{
privateKey: privateKey,
limit: maxSharedKeys,
keys: make(map[wgtypes.Key]*[32]byte),
}
}
// Encrypt encrypts msg for peerPublicKey. It is safe for concurrent use.
func (c *SharedKeyCache) Encrypt(msg []byte, peerPublicKey wgtypes.Key) ([]byte, error) {
nonce, err := genNonce()
if err != nil {
return nil, err
}
return box.SealAfterPrecomputation(nonce[:], msg, nonce, c.sharedKey(peerPublicKey)), nil
}
// Decrypt decrypts a message that peerPublicKey encrypted for this cache's private
// key. It is safe for concurrent use.
func (c *SharedKeyCache) Decrypt(encryptedMsg []byte, peerPublicKey wgtypes.Key) ([]byte, error) {
if len(encryptedMsg) < nonceSize {
return nil, fmt.Errorf("invalid encrypted message length")
}
var nonce [nonceSize]byte
copy(nonce[:], encryptedMsg[:nonceSize])
shared, cached := c.cached(peerPublicKey)
if !cached {
shared = c.derive(peerPublicKey)
}
opened, ok := box.OpenAfterPrecomputation(nil, encryptedMsg[nonceSize:], &nonce, shared)
if !ok {
return nil, fmt.Errorf("failed to decrypt message from peer %s", peerPublicKey.String())
}
// The sender key of an incoming message is not authenticated until it opens, so
// only a key that produced a valid message is cached. Forged senders cannot fill
// the cache or evict real peers.
if !cached {
c.store(peerPublicKey, shared)
}
return opened, nil
}
// EncryptMessage marshals message and encrypts it for peerPublicKey.
func (c *SharedKeyCache) EncryptMessage(peerPublicKey wgtypes.Key, message pb.Message) ([]byte, error) {
body, err := pb.Marshal(message)
if err != nil {
return nil, fmt.Errorf("marshal message: %w", err)
}
return c.Encrypt(body, peerPublicKey)
}
// DecryptMessage decrypts a message from peerPublicKey and unmarshals it into message.
func (c *SharedKeyCache) DecryptMessage(peerPublicKey wgtypes.Key, encryptedMessage []byte, message pb.Message) error {
body, err := c.Decrypt(encryptedMessage, peerPublicKey)
if err != nil {
return err
}
if err := pb.Unmarshal(body, message); err != nil {
return fmt.Errorf("unmarshal message from peer %s: %w", peerPublicKey.String(), err)
}
return nil
}
// Close drops every cached shared key and stops caching new ones. Encrypt and
// Decrypt keep working afterwards by deriving the key for each message.
func (c *SharedKeyCache) Close() {
c.mu.Lock()
defer c.mu.Unlock()
c.closed = true
clear(c.keys)
}
func (c *SharedKeyCache) sharedKey(peerPublicKey wgtypes.Key) *[32]byte {
if shared, ok := c.cached(peerPublicKey); ok {
return shared
}
shared := c.derive(peerPublicKey)
c.store(peerPublicKey, shared)
return shared
}
func (c *SharedKeyCache) cached(peerPublicKey wgtypes.Key) (*[32]byte, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
shared, ok := c.keys[peerPublicKey]
return shared, ok
}
// derive computes the shared key outside the lock: two goroutines racing on a new
// peer compute the same value, and holding the lock would serialise the x25519 work
// this cache avoids.
func (c *SharedKeyCache) derive(peerPublicKey wgtypes.Key) *[32]byte {
shared := new([32]byte)
box.Precompute(shared, toByte32(peerPublicKey), toByte32(c.privateKey))
return shared
}
func (c *SharedKeyCache) store(peerPublicKey wgtypes.Key, shared *[32]byte) {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return
}
if len(c.keys) >= c.limit {
// Map iteration order is random, so this evicts an arbitrary entry.
for k := range c.keys {
delete(c.keys, k)
break
}
}
c.keys[peerPublicKey] = shared
}
+8
View File
@@ -0,0 +1,8 @@
//go:build !ios && !android
package encryption
// maxSharedKeys bounds the cache so peers that come and go (ephemeral peers get a
// new key on every registration) cannot grow it for the lifetime of the process.
// An entry costs about 130 bytes, so a full cache is around 8 MB.
const maxSharedKeys = 1 << 16
+8
View File
@@ -0,0 +1,8 @@
//go:build ios || android
package encryption
// maxSharedKeys is small on mobile, where the process runs under a tight memory
// limit. A miss only costs a fresh key derivation. An entry costs about 130 bytes,
// so a full cache is around 130 KB.
const maxSharedKeys = 1 << 10
+184
View File
@@ -0,0 +1,184 @@
package encryption
import (
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
func newKeyPair(t testing.TB) (wgtypes.Key, wgtypes.Key) {
t.Helper()
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err)
return priv, priv.PublicKey()
}
// The cache must stay wire compatible with peers that use the uncached functions,
// in both directions.
func TestSharedKeyCache_InteropWithUncached(t *testing.T) {
alicePriv, alicePub := newKeyPair(t)
bobPriv, bobPub := newKeyPair(t)
alice := NewSharedKeyCache(alicePriv)
msg := []byte("offer")
enc, err := alice.Encrypt(msg, bobPub)
require.NoError(t, err)
dec, err := Decrypt(enc, alicePub, bobPriv)
require.NoError(t, err)
assert.Equal(t, msg, dec, "uncached peer must read a cached sender's message")
enc, err = Encrypt(msg, alicePub, bobPriv)
require.NoError(t, err)
dec, err = alice.Decrypt(enc, bobPub)
require.NoError(t, err)
assert.Equal(t, msg, dec, "cached peer must read an uncached sender's message")
}
// Two messages to the same peer share the derived key but never the nonce, so the
// ciphertexts differ.
func TestSharedKeyCache_FreshNoncePerMessage(t *testing.T) {
priv, _ := newKeyPair(t)
_, peerPub := newKeyPair(t)
c := NewSharedKeyCache(priv)
a, err := c.Encrypt([]byte("same"), peerPub)
require.NoError(t, err)
b, err := c.Encrypt([]byte("same"), peerPub)
require.NoError(t, err)
assert.NotEqual(t, a, b, "ciphertexts of identical plaintext must differ")
assert.Len(t, c.keys, 1, "the shared key must be derived once per peer")
}
// A message from one peer must not decrypt under another peer's cached key.
func TestSharedKeyCache_DoesNotMixPeers(t *testing.T) {
alicePriv, alicePub := newKeyPair(t)
bobPriv, _ := newKeyPair(t)
_, carolPub := newKeyPair(t)
alice := NewSharedKeyCache(alicePriv)
enc, err := Encrypt([]byte("hi"), alicePub, bobPriv)
require.NoError(t, err)
_, err = alice.Decrypt(enc, carolPub)
assert.Error(t, err, "a message from Bob must not open with Carol's key")
}
func TestSharedKeyCache_RejectsShortMessage(t *testing.T) {
priv, _ := newKeyPair(t)
_, peerPub := newKeyPair(t)
_, err := NewSharedKeyCache(priv).Decrypt(make([]byte, nonceSize-1), peerPub)
assert.Error(t, err)
}
func TestSharedKeyCache_StaysBounded(t *testing.T) {
priv, _ := newKeyPair(t)
c := NewSharedKeyCache(priv)
c.limit = 4
for i := 0; i < 20; i++ {
_, peerPub := newKeyPair(t)
_, err := c.Encrypt([]byte("x"), peerPub)
require.NoError(t, err)
assert.LessOrEqual(t, len(c.keys), c.limit, "cache must not grow past its cap")
}
assert.Len(t, c.keys, c.limit, "a full cache keeps evicting one entry per new peer")
c.Close()
assert.Empty(t, c.keys, "Close must drop every entry")
}
func TestSharedKeyCache_Concurrent(t *testing.T) {
alicePriv, alicePub := newKeyPair(t)
bobPriv, bobPub := newKeyPair(t)
alice := NewSharedKeyCache(alicePriv)
bob := NewSharedKeyCache(bobPriv)
var wg sync.WaitGroup
for i := 0; i < 16; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for j := 0; j < 50; j++ {
enc, err := alice.Encrypt([]byte("m"), bobPub)
if !assert.NoError(t, err) {
return
}
dec, err := bob.Decrypt(enc, alicePub)
if !assert.NoError(t, err) || !assert.Equal(t, []byte("m"), dec) {
return
}
}
}()
}
wg.Wait()
}
func BenchmarkEncryptDecryptUncached(b *testing.B) {
alicePriv, alicePub := newKeyPair(b)
bobPriv, bobPub := newKeyPair(b)
msg := make([]byte, 512)
b.ReportAllocs()
for i := 0; i < b.N; i++ {
enc, err := Encrypt(msg, bobPub, alicePriv)
require.NoError(b, err)
_, err = Decrypt(enc, alicePub, bobPriv)
require.NoError(b, err)
}
}
func BenchmarkEncryptDecryptCached(b *testing.B) {
alicePriv, alicePub := newKeyPair(b)
bobPriv, bobPub := newKeyPair(b)
alice := NewSharedKeyCache(alicePriv)
bob := NewSharedKeyCache(bobPriv)
msg := make([]byte, 512)
b.ReportAllocs()
for i := 0; i < b.N; i++ {
enc, err := alice.Encrypt(msg, bobPub)
require.NoError(b, err)
_, err = bob.Decrypt(enc, alicePub)
require.NoError(b, err)
}
}
// A forged sender key must not populate the cache: the key of an incoming message
// is only trusted once the message opens.
func TestSharedKeyCache_FailedDecryptDoesNotCache(t *testing.T) {
alicePriv, alicePub := newKeyPair(t)
bobPriv, bobPub := newKeyPair(t)
_, forgedPub := newKeyPair(t)
alice := NewSharedKeyCache(alicePriv)
enc, err := Encrypt([]byte("hi"), alicePub, bobPriv)
require.NoError(t, err)
_, err = alice.Decrypt(enc, forgedPub)
require.Error(t, err)
assert.Empty(t, alice.keys, "a message that fails to open must not add a cache entry")
_, err = alice.Decrypt(enc, bobPub)
require.NoError(t, err)
assert.Len(t, alice.keys, 1, "a message that opens caches its sender's key")
}
// After Close the cache still works but no longer keeps key material.
func TestSharedKeyCache_ClosedDoesNotRepopulate(t *testing.T) {
alicePriv, alicePub := newKeyPair(t)
bobPriv, bobPub := newKeyPair(t)
alice := NewSharedKeyCache(alicePriv)
_, err := alice.Encrypt([]byte("x"), bobPub)
require.NoError(t, err)
alice.Close()
assert.Empty(t, alice.keys)
enc, err := alice.Encrypt([]byte("y"), bobPub)
require.NoError(t, err)
dec, err := Decrypt(enc, alicePub, bobPriv)
require.NoError(t, err)
assert.Equal(t, []byte("y"), dec, "a closed cache must still encrypt correctly")
assert.Empty(t, alice.keys, "a closed cache must not cache new keys")
}
+5 -2
View File
@@ -53,6 +53,7 @@ type ConnStateNotifier interface {
// GrpcClient Wraps the Signal Exchange Service gRpc client
type GrpcClient struct {
key wgtypes.Key
sharedKeys *encryption.SharedKeyCache
realClient proto.SignalExchangeClient
signalConn *grpc.ClientConn
ctx context.Context
@@ -107,6 +108,7 @@ func NewClient(ctx context.Context, addr string, key wgtypes.Key, tlsEnabled boo
c := &GrpcClient{
ctx: ctx,
key: key,
sharedKeys: encryption.NewSharedKeyCache(key),
mux: sync.Mutex{},
status: StreamDisconnected,
connStateCallbackLock: sync.RWMutex{},
@@ -158,6 +160,7 @@ func (c *GrpcClient) Close() error {
}
c.decryptionWg.Wait()
c.decryptionWorker = nil
c.sharedKeys.Close()
return c.signalConn.Close()
}
@@ -418,7 +421,7 @@ func (c *GrpcClient) decryptMessage(msg *proto.EncryptedMessage) (*proto.Message
}
body := &proto.Body{}
err = encryption.DecryptMessage(remoteKey, c.key, msg.GetBody(), body)
err = c.sharedKeys.DecryptMessage(remoteKey, msg.GetBody(), body)
if err != nil {
return nil, err
}
@@ -438,7 +441,7 @@ func (c *GrpcClient) encryptMessage(msg *proto.Message) (*proto.EncryptedMessage
return nil, err
}
encryptedBody, err := encryption.EncryptMessage(remoteKey, c.key, msg.Body)
encryptedBody, err := c.sharedKeys.EncryptMessage(remoteKey, msg.Body)
if err != nil {
return nil, err
}