diff --git a/encryption/sharedkey.go b/encryption/sharedkey.go new file mode 100644 index 000000000..1509632c4 --- /dev/null +++ b/encryption/sharedkey.go @@ -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 +} diff --git a/encryption/sharedkey_limit.go b/encryption/sharedkey_limit.go new file mode 100644 index 000000000..b141492f6 --- /dev/null +++ b/encryption/sharedkey_limit.go @@ -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 diff --git a/encryption/sharedkey_limit_mobile.go b/encryption/sharedkey_limit_mobile.go new file mode 100644 index 000000000..f36181f94 --- /dev/null +++ b/encryption/sharedkey_limit_mobile.go @@ -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 diff --git a/encryption/sharedkey_test.go b/encryption/sharedkey_test.go new file mode 100644 index 000000000..003ccb627 --- /dev/null +++ b/encryption/sharedkey_test.go @@ -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") +} diff --git a/shared/signal/client/grpc.go b/shared/signal/client/grpc.go index a0bb2f080..92be57b30 100644 --- a/shared/signal/client/grpc.go +++ b/shared/signal/client/grpc.go @@ -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 }