diff --git a/e2ee/datacryptor.go b/e2ee/datacryptor.go index d11bfaa1..0f662094 100644 --- a/e2ee/datacryptor.go +++ b/e2ee/datacryptor.go @@ -37,8 +37,8 @@ type dataCipherState struct { keyBytes []byte } -// DataCryptor handles encryption and decryption of data channel messages. -// It mirrors the JS SDK's DataCryptor class, using AES-128-GCM with no AAD. +// DataCryptor handles encryption and decryption of data channel messages and of raw payloads such +// as data track frames. It mirrors the JS SDK's DataCryptor class, using AES-128-GCM with no AAD. type DataCryptor struct { keyProvider types.KeyProvider cipherCache map[uint32]*dataCipherState @@ -53,9 +53,8 @@ func NewDataCryptor(keyProvider types.KeyProvider) *DataCryptor { } } -// Encrypt wraps a DataPacket's value in an EncryptedPacket. -// The original value is serialized as EncryptedPacketPayload, then -// encrypted with AES-128-GCM using a random IV and no AAD. +// Encrypt wraps a DataPacket's value in an EncryptedPacket: the value is serialized as an +// EncryptedPacketPayload and sealed with EncryptPayload. func (dc *DataCryptor) Encrypt(pck *livekit.DataPacket) (*livekit.DataPacket, error) { payload := DataPacketValueToPayload(pck) if payload == nil { @@ -67,24 +66,11 @@ func (dc *DataCryptor) Encrypt(pck *livekit.DataPacket) (*livekit.DataPacket, er return nil, fmt.Errorf("marshal payload: %w", err) } - keyIndex := dc.keyProvider.CurrentKeyIndex() - block, err := dc.getCipherBlock(keyIndex) - if err != nil { - return nil, fmt.Errorf("get cipher: %w", err) - } - - aesGCM, err := cipher.NewGCMWithNonceSize(block, types.IVLength) + encrypted, err := dc.EncryptPayload(plaintext) if err != nil { return nil, err } - iv := make([]byte, types.IVLength) - if _, err := io.ReadFull(rand.Reader, iv); err != nil { - return nil, fmt.Errorf("generate IV: %w", err) - } - - ciphertext := aesGCM.Seal(nil, iv, plaintext, nil) - return &livekit.DataPacket{ Kind: pck.Kind, //nolint:staticcheck ParticipantIdentity: pck.ParticipantIdentity, @@ -92,9 +78,9 @@ func (dc *DataCryptor) Encrypt(pck *livekit.DataPacket) (*livekit.DataPacket, er Value: &livekit.DataPacket_EncryptedPacket{ EncryptedPacket: &livekit.EncryptedPacket{ EncryptionType: livekit.Encryption_GCM, - Iv: iv, - KeyIndex: keyIndex, - EncryptedValue: ciphertext, + Iv: encrypted.IV, + KeyIndex: encrypted.KeyIndex, + EncryptedValue: encrypted.Ciphertext, }, }, }, nil @@ -102,30 +88,70 @@ func (dc *DataCryptor) Encrypt(pck *livekit.DataPacket) (*livekit.DataPacket, er // Decrypt extracts and decrypts an EncryptedPacket, returning the inner payload. func (dc *DataCryptor) Decrypt(ep *livekit.EncryptedPacket) (*livekit.EncryptedPacketPayload, error) { - if len(ep.Iv) == 0 || len(ep.EncryptedValue) == 0 { - return nil, fmt.Errorf("empty IV or ciphertext") + plaintext, err := dc.DecryptPayload(EncryptedPayload{Ciphertext: ep.EncryptedValue, KeyIndex: ep.KeyIndex, IV: ep.Iv}) + if err != nil { + return nil, err + } + + payload := &livekit.EncryptedPacketPayload{} + if err := proto.Unmarshal(plaintext, payload); err != nil { + return nil, fmt.Errorf("unmarshal payload: %w", err) } + return payload, nil +} - block, err := dc.getCipherBlock(ep.KeyIndex) +// EncryptedPayload is a payload sealed by EncryptPayload together with what DecryptPayload needs +// to open it. +type EncryptedPayload struct { + Ciphertext []byte + KeyIndex uint32 + IV []byte +} + +// EncryptPayload seals plaintext with the current key, using a random IV and no AAD. +func (dc *DataCryptor) EncryptPayload(plaintext []byte) (EncryptedPayload, error) { + keyIndex := dc.keyProvider.CurrentKeyIndex() + aesGCM, err := dc.gcm(keyIndex, types.IVLength) if err != nil { - return nil, fmt.Errorf("get cipher for index %d: %w", ep.KeyIndex, err) + return EncryptedPayload{}, err + } + + iv := make([]byte, types.IVLength) + if _, err := io.ReadFull(rand.Reader, iv); err != nil { + return EncryptedPayload{}, fmt.Errorf("generate IV: %w", err) + } + + return EncryptedPayload{ + Ciphertext: aesGCM.Seal(nil, iv, plaintext, nil), + KeyIndex: keyIndex, + IV: iv, + }, nil +} + +// DecryptPayload opens a payload sealed with the key at its key index. +func (dc *DataCryptor) DecryptPayload(payload EncryptedPayload) ([]byte, error) { + if len(payload.IV) == 0 || len(payload.Ciphertext) == 0 { + return nil, fmt.Errorf("empty IV or ciphertext") } - aesGCM, err := cipher.NewGCMWithNonceSize(block, len(ep.Iv)) + aesGCM, err := dc.gcm(payload.KeyIndex, len(payload.IV)) if err != nil { return nil, err } - plaintext, err := aesGCM.Open(nil, ep.Iv, ep.EncryptedValue, nil) + plaintext, err := aesGCM.Open(nil, payload.IV, payload.Ciphertext, nil) if err != nil { return nil, fmt.Errorf("decrypt: %w", err) } + return plaintext, nil +} - payload := &livekit.EncryptedPacketPayload{} - if err := proto.Unmarshal(plaintext, payload); err != nil { - return nil, fmt.Errorf("unmarshal payload: %w", err) +func (dc *DataCryptor) gcm(keyIndex uint32, nonceSize int) (cipher.AEAD, error) { + block, err := dc.getCipherBlock(keyIndex) + if err != nil { + return nil, fmt.Errorf("get cipher for index %d: %w", keyIndex, err) } - return payload, nil + return cipher.NewGCMWithNonceSize(block, nonceSize) } // getCipherBlock returns an AES cipher block for the given key index. If the diff --git a/e2ee/datacryptor_test.go b/e2ee/datacryptor_test.go index d5771e3d..22f75de4 100644 --- a/e2ee/datacryptor_test.go +++ b/e2ee/datacryptor_test.go @@ -6,6 +6,7 @@ import ( "github.com/stretchr/testify/require" "github.com/livekit/protocol/livekit" + "google.golang.org/protobuf/proto" "github.com/livekit/server-sdk-go/v2/e2ee" ) @@ -176,3 +177,60 @@ func newTestDataCryptor(t *testing.T) *e2ee.DataCryptor { require.NoError(t, kp.SetKeyFromPassphrase("12345", 0)) return e2ee.NewDataCryptor(kp) } + +func TestDataCryptorPayloadRoundTrip(t *testing.T) { + dc := newTestDataCryptor(t) + + encrypted, err := dc.EncryptPayload([]byte("hello encrypted world")) + require.NoError(t, err) + require.Len(t, encrypted.IV, 12) + require.NotEqual(t, []byte("hello encrypted world"), encrypted.Ciphertext) + + plaintext, err := dc.DecryptPayload(encrypted) + require.NoError(t, err) + require.Equal(t, []byte("hello encrypted world"), plaintext) +} + +func TestDataCryptorPayloadWrongKeyFails(t *testing.T) { + kpA := e2ee.NewExternalKeyProvider() + require.NoError(t, kpA.SetRawKey(bytes16(0x11), 0)) + dcA := e2ee.NewDataCryptor(kpA) + + kpB := e2ee.NewExternalKeyProvider() + require.NoError(t, kpB.SetRawKey(bytes16(0x22), 0)) + dcB := e2ee.NewDataCryptor(kpB) + + encrypted, err := dcA.EncryptPayload([]byte("secret")) + require.NoError(t, err) + + _, err = dcB.DecryptPayload(encrypted) + require.Error(t, err) +} + +func TestDataCryptorPayloadEmptyRejected(t *testing.T) { + dc := newTestDataCryptor(t) + + _, err := dc.DecryptPayload(e2ee.EncryptedPayload{Ciphertext: []byte("x")}) + require.Error(t, err) + + _, err = dc.DecryptPayload(e2ee.EncryptedPayload{IV: []byte("iv")}) + require.Error(t, err) +} + +// A DataPacket sealed by Encrypt opens with DecryptPayload: both paths share one primitive. +func TestDataCryptorPacketUsesPayloadEncryption(t *testing.T) { + dc := newTestDataCryptor(t) + + encrypted, err := dc.Encrypt(&livekit.DataPacket{ + Value: &livekit.DataPacket_User{User: &livekit.UserPacket{Payload: []byte("shared primitive")}}, + }) + require.NoError(t, err) + ep := encrypted.Value.(*livekit.DataPacket_EncryptedPacket).EncryptedPacket + + plaintext, err := dc.DecryptPayload(e2ee.EncryptedPayload{Ciphertext: ep.EncryptedValue, KeyIndex: ep.KeyIndex, IV: ep.Iv}) + require.NoError(t, err) + + payload := &livekit.EncryptedPacketPayload{} + require.NoError(t, proto.Unmarshal(plaintext, payload)) + require.Equal(t, []byte("shared primitive"), payload.GetUser().GetPayload()) +} diff --git a/e2ee/frameencryptor.go b/e2ee/frameencryptor.go index f4cb22aa..5486410f 100644 --- a/e2ee/frameencryptor.go +++ b/e2ee/frameencryptor.go @@ -75,6 +75,9 @@ func NewGCMFrameEncryptor(keyProvider types.KeyProvider, encryptFn EncryptFunc) // EncryptFrame encrypts a complete media frame. func (e *GCMFrameEncryptor) EncryptFrame(payload []byte) ([]byte, error) { idx := e.keyProvider.CurrentKeyIndex() + if idx > types.MaxKeyIndex { + return nil, types.ErrKeyIndexOutOfRange + } st := e.state.Load() // Fast path: key index unchanged, use cached cipher block (no lock). diff --git a/e2ee/frameencryptor_test.go b/e2ee/frameencryptor_test.go index e7118162..6aa62ba2 100644 --- a/e2ee/frameencryptor_test.go +++ b/e2ee/frameencryptor_test.go @@ -8,6 +8,7 @@ import ( "github.com/stretchr/testify/require" "github.com/livekit/server-sdk-go/v2/e2ee" + "github.com/livekit/server-sdk-go/v2/e2ee/types" ) // stubEncryptFunc writes [KID, payload...] — lets us assert which KID was used @@ -172,3 +173,21 @@ func TestGCMFrameEncryptorUsesRealCipher(t *testing.T) { _, err = e2ee.NewGCMFrameEncryptor(kp, stubEncryptFunc) require.NoError(t, err) } + +// unboundedKeyProvider stands in for a custom provider that does not enforce the index bound. +type unboundedKeyProvider struct { + index uint32 +} + +func (p *unboundedKeyProvider) GetKey(uint32) ([]byte, error) { return bytes16(0x11), nil } +func (p *unboundedKeyProvider) CurrentKeyIndex() uint32 { return p.index } + +func TestGCMFrameEncryptorRejectsOutOfRangeIndex(t *testing.T) { + kp := &unboundedKeyProvider{} + enc, err := e2ee.NewGCMFrameEncryptor(kp, stubEncryptFunc) + require.NoError(t, err) + + kp.index = 256 + _, err = enc.EncryptFrame([]byte{0x01, 0x02}) + require.ErrorIs(t, err, types.ErrKeyIndexOutOfRange) +} diff --git a/e2ee/keyprovider.go b/e2ee/keyprovider.go index 58b88641..ad2103c7 100644 --- a/e2ee/keyprovider.go +++ b/e2ee/keyprovider.go @@ -46,6 +46,9 @@ func (p *ExternalKeyProvider) SetKeyFromPassphrase(passphrase string, index uint if passphrase == "" { return fmt.Errorf("passphrase cannot be empty") } + if index > types.MaxKeyIndex { + return types.ErrKeyIndexOutOfRange + } derived := pbkdf2.Key( []byte(passphrase), []byte(types.SDKSalt), @@ -66,6 +69,9 @@ func (p *ExternalKeyProvider) SetRawKey(key []byte, index uint32) error { if len(key) != types.KeySizeBytes { return types.ErrIncorrectKeyLength } + if index > types.MaxKeyIndex { + return types.ErrKeyIndexOutOfRange + } p.mu.Lock() defer p.mu.Unlock() p.keys[index] = key diff --git a/e2ee/keyprovider_test.go b/e2ee/keyprovider_test.go index 6448b392..f5cfb9c4 100644 --- a/e2ee/keyprovider_test.go +++ b/e2ee/keyprovider_test.go @@ -82,3 +82,13 @@ func bytes16(b byte) []byte { } return out } + +func TestKeyIndexOutOfRangeRejected(t *testing.T) { + kp := e2ee.NewExternalKeyProvider() + + require.ErrorIs(t, kp.SetRawKey(bytes16(0x11), 256), types.ErrKeyIndexOutOfRange) + require.ErrorIs(t, kp.SetKeyFromPassphrase("12345", 256), types.ErrKeyIndexOutOfRange) + + require.NoError(t, kp.SetRawKey(bytes16(0x11), 255)) + require.Equal(t, uint32(255), kp.CurrentKeyIndex()) +} diff --git a/e2ee/types/constants.go b/e2ee/types/constants.go index dd9c6b18..a2cb1369 100644 --- a/e2ee/types/constants.go +++ b/e2ee/types/constants.go @@ -22,6 +22,8 @@ const ( PBKDFIterations = 100000 KeySizeBytes = 16 HKDFInfoBytes = 128 + // MaxKeyIndex is the largest key index the one-byte frame trailer can carry. + MaxKeyIndex = 255 ) var ( @@ -29,4 +31,5 @@ var ( ErrUnableGenerateIV = errors.New("unable to generate iv for encryption") ErrIncorrectIVLength = errors.New("incorrect iv length") ErrBlockCipherRequired = errors.New("input block cipher cannot be nil") + ErrKeyIndexOutOfRange = errors.New("key index must not exceed 255") ) diff --git a/e2ee/types/interface.go b/e2ee/types/interface.go index ffe0dad5..b93a967e 100644 --- a/e2ee/types/interface.go +++ b/e2ee/types/interface.go @@ -16,7 +16,8 @@ package types // KeyProvider manages encryption keys for E2EE. type KeyProvider interface { - // GetKey returns the derived AES key for the given index. + // GetKey returns the derived AES key for the given index. It is called on every encrypt and + // decrypt, so implementations should return pre-derived keys rather than derive on demand. GetKey(keyIndex uint32) ([]byte, error) // CurrentKeyIndex returns the active key index for encryption. CurrentKeyIndex() uint32