Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
88 changes: 36 additions & 52 deletions go/crypto.go
Original file line number Diff line number Diff line change
Expand Up @@ -296,29 +296,12 @@ func buildCapsules(contentKey ContentKey, members []Pubkey, ephScalar *ristretto
return Capsules{out}
}

func scanCapsules(data []byte, ephPubkey EphPubkey, myScalar ViewScalar, nonce Nonce) (int, ContentKey, bool) {
epb := ephPubkey.b
shared, err := ecdhSharedSecret(viewScalarToRistretto(myScalar), epb[:])
if err != nil {
return 0, ContentKey{}, false
}
myTag := deriveViewTagByte(shared)
kek := deriveKeyWrap(shared, nonce)

offset := 0
idx := 0
for offset+CapsuleSize <= len(data) {
if data[offset] == myTag {
var ck [32]byte
for i := 0; i < 32; i++ {
ck[i] = data[offset+1+i] ^ kek[i]
}
return idx, ContentKey{ck}, true
}
offset += CapsuleSize
idx++
func contentKeyFromCapsule(data []byte, offset int, kek []byte) ContentKey {
var ck [32]byte
for i := 0; i < 32; i++ {
ck[i] = data[offset+1+i] ^ kek[i]
}
return 0, ContentKey{}, false
return ContentKey{ck}
}

func EncryptForGroup(plaintext Plaintext, members []Pubkey, nonce Nonce, senderSeed Seed) (EphPubkey, Capsules, Ciphertext, error) {
Expand Down Expand Up @@ -366,40 +349,41 @@ func DecryptFromGroup(content []byte, myScalar ViewScalar, nonce Nonce, knownN i
ephPubkey := EphPubkey{ephArr}
afterEph := content[32:]

capsuleIdx, contentKey, found := scanCapsules(afterEph, ephPubkey, myScalar, nonce)
if !found {
return Plaintext{}, ErrDecryptionFailed
}
ckRaw := contentKey.b

aead, err := chacha20poly1305.New(ckRaw[:])
epb := ephPubkey.b
shared, err := ecdhSharedSecret(viewScalarToRistretto(myScalar), epb[:])
if err != nil {
return Plaintext{}, err
}

if knownN > 0 {
ctStart := knownN * CapsuleSize
if ctStart > len(afterEph) {
return Plaintext{}, ErrInsufficientData
}
pt, err := aead.Open(nil, nonce.chachaNonce(), afterEph[ctStart:], nil)
if err != nil {
return Plaintext{}, ErrDecryptionFailed
}
return Plaintext{pt}, nil
return Plaintext{}, ErrDecryptionFailed
}
myTag := deriveViewTagByte(shared)
kek := deriveKeyWrap(shared, nonce)

minN := capsuleIdx + 1
maxN := (len(afterEph) - 16) / CapsuleSize
for n := minN; n <= maxN; n++ {
ctStart := n * CapsuleSize
if ctStart >= len(afterEph) {
break
}
pt, err := aead.Open(nil, nonce.chachaNonce(), afterEph[ctStart:], nil)
if err == nil {
return Plaintext{pt}, nil
offset := 0
capsuleIdx := 0
for offset+CapsuleSize <= len(afterEph) {
if afterEph[offset] == myTag {
contentKey := contentKeyFromCapsule(afterEph, offset, kek)
ckRaw := contentKey.b
aead, _ := chacha20poly1305.New(ckRaw[:])
if knownN > 0 {
ctStart := knownN * CapsuleSize
if ctStart > len(afterEph) {
return Plaintext{}, ErrInsufficientData
}
if pt, err := aead.Open(nil, nonce.chachaNonce(), afterEph[ctStart:], nil); err == nil {
return Plaintext{pt}, nil
}
} else {
maxN := (len(afterEph) - 16) / CapsuleSize
for n := capsuleIdx + 1; n <= maxN; n++ {
ctStart := n * CapsuleSize
if pt, err := aead.Open(nil, nonce.chachaNonce(), afterEph[ctStart:], nil); err == nil {
return Plaintext{pt}, nil
}
}
}
}
offset += CapsuleSize
capsuleIdx++
}
return Plaintext{}, ErrDecryptionFailed
}
73 changes: 44 additions & 29 deletions go/crypto_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,22 @@ func randomNonce(t *testing.T) Nonce {
return NonceFromBytes(b)
}

func fixedSeed(value byte) Seed {
var b [32]byte
for i := range b {
b[i] = value
}
return SeedFromBytes(b)
}

func fixedNonce(value byte) Nonce {
var b [12]byte
for i := range b {
b[i] = value
}
return NonceFromBytes(b)
}

func TestRandomNonceReadsTwelveBytesFromRandomSource(t *testing.T) {
reader := &deterministicReader{}
withRandomReader(t, reader)
Expand Down Expand Up @@ -311,35 +327,6 @@ func TestBuildCapsulesWithInvalidMemberPoint(t *testing.T) {
require.Equal(t, expected, capsules.Bytes())
}

func TestScanCapsulesNoMatchReturnsNotFound(t *testing.T) {
// Construct capsule data where all view tags are 0xAA.
capsuleData := make([]byte, CapsuleSize*2)
for i := range capsuleData {
capsuleData[i] = 0xAA
}
// Use a real ephemeral pubkey from a seed.
seed := randomSeed(t)
scalar := Sr25519SigningScalar(seed)
pub := PublicFromSeed(seed)
nonce := randomNonce(t)

_, _, found := scanCapsules(capsuleData, EphPubkeyFromBytes(pub.Bytes()), scalar, nonce)
// Either the tag happens to match (1/256) or not -- just exercise the path.
_ = found
}

func TestScanCapsulesInvalidEphPubkey(t *testing.T) {
capsuleData := make([]byte, CapsuleSize)
var badEph [32]byte
for i := range badEph {
badEph[i] = 0xFF
}
scalar := Sr25519SigningScalar(randomSeed(t))
nonce := randomNonce(t)
_, _, found := scanCapsules(capsuleData, EphPubkeyFromBytes(badEph), scalar, nonce)
require.False(t, found)
}

func TestDecryptFromGroupCorruptedCiphertextBody(t *testing.T) {
sender := randomSeed(t)
nonce := randomNonce(t)
Expand Down Expand Up @@ -464,6 +451,34 @@ func TestEncryptForGroupTrialDecryptWithoutKnownN(t *testing.T) {
require.Equal(t, msg.Bytes(), pt.Bytes())
}

func TestDecryptFromGroupSkipsCollidingViewTagCapsules(t *testing.T) {
sender := fixedSeed(0xAA)
alice := fixedSeed(0xAA)
bob := fixedSeed(0xBB)
nonce := fixedNonce(0xF1)
members := []Pubkey{PublicFromSeed(alice), PublicFromSeed(bob)}
msg := PlaintextFromBytes([]byte("view tag collision"))

ephPub, capsules, ct, err := EncryptForGroup(msg, members, nonce, sender)
require.NoError(t, err)

content := make([]byte, 0)
epb := ephPub.Bytes()
content = append(content, epb[:]...)
content = append(content, capsules.Bytes()...)
content = append(content, ct.Bytes()...)
content[32] = content[32+CapsuleSize]

scalar := Sr25519SigningScalar(bob)
pt, err := DecryptFromGroup(content, scalar, nonce, len(members))
require.NoError(t, err)
require.Equal(t, msg.Bytes(), pt.Bytes())

pt, err = DecryptFromGroup(content, scalar, nonce, 0)
require.NoError(t, err)
require.Equal(t, msg.Bytes(), pt.Bytes())
}

func TestDecryptFromGroupInvalidEphPubkeyPoint(t *testing.T) {
nonce := randomNonce(t)
scalar := Sr25519SigningScalar(randomSeed(t))
Expand Down
73 changes: 44 additions & 29 deletions python/samp-crypto/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -329,6 +329,12 @@ fn xor32(a: &[u8; 32], b: &[u8; 32]) -> [u8; 32] {
out
}

fn content_key_from_capsule(data: &[u8], offset: usize, kek: &[u8; 32]) -> [u8; 32] {
let mut wrapped = [0u8; 32];
wrapped.copy_from_slice(&data[offset + 1..offset + CAPSULE_SIZE]);
xor32(&wrapped, kek)
}

fn derive_view_tag(shared_secret: &[u8; 32]) -> u8 {
let hk = Hkdf::<Sha256>::new(None, shared_secret);
let mut tag = [0u8; 1];
Expand Down Expand Up @@ -449,9 +455,20 @@ fn encrypt_for_group(plaintext: &[u8], member_pubkeys: Vec<Vec<u8>>, nonce: &[u8
}

#[pyfunction]
fn decrypt_from_group(content: &[u8], my_scalar: &[u8], nonce: &[u8], known_n: Option<usize>) -> PyResult<Vec<u8>> {
let ms = Scalar::from_bytes_mod_order(my_scalar.try_into().map_err(|_| err("my_scalar must be 32 bytes"))?);
let n: [u8; 12] = nonce.try_into().map_err(|_| err("nonce must be 12 bytes"))?;
fn decrypt_from_group(
content: &[u8],
my_scalar: &[u8],
nonce: &[u8],
known_n: Option<usize>,
) -> PyResult<Vec<u8>> {
let ms = Scalar::from_bytes_mod_order(
my_scalar
.try_into()
.map_err(|_| err("my_scalar must be 32 bytes"))?,
);
let n: [u8; 12] = nonce
.try_into()
.map_err(|_| err("nonce must be 12 bytes"))?;
if content.len() < 32 {
return Err(err("content too short"));
}
Expand All @@ -464,38 +481,36 @@ fn decrypt_from_group(content: &[u8], my_scalar: &[u8], nonce: &[u8], known_n: O

let mut offset = 0;
let mut capsule_idx = 0;
let mut content_key: Option<[u8; 32]> = None;
while offset + CAPSULE_SIZE <= after_eph.len() {
if after_eph[offset] == my_tag {
let mut wrapped = [0u8; 32];
wrapped.copy_from_slice(&after_eph[offset + 1..offset + 33]);
content_key = Some(xor32(&wrapped, &kek));
break;
let ck = content_key_from_capsule(after_eph, offset, &kek);
let cipher = ChaCha20Poly1305::new((&ck).into());
if let Some(member_count) = known_n {
let ct_start = member_count
.checked_mul(CAPSULE_SIZE)
.ok_or_else(|| err("content too short"))?;
if ct_start > after_eph.len() {
return Err(err("content too short"));
}
if let Ok(plaintext) = cipher.decrypt(Nonce::from_slice(&n), &after_eph[ct_start..])
{
return Ok(plaintext);
}
} else {
let max_n = after_eph.len().saturating_sub(16) / CAPSULE_SIZE;
for trial_n in capsule_idx + 1..=max_n {
let ct_start = trial_n * CAPSULE_SIZE;
if let Ok(plaintext) =
cipher.decrypt(Nonce::from_slice(&n), &after_eph[ct_start..])
{
return Ok(plaintext);
}
}
}
}
offset += CAPSULE_SIZE;
capsule_idx += 1;
}
let ck = content_key.ok_or_else(|| err("decryption failed"))?;
let cipher = ChaCha20Poly1305::new((&ck).into());

if let Some(member_count) = known_n {
let ct_start = member_count * CAPSULE_SIZE;
if ct_start > after_eph.len() {
return Err(err("content too short"));
}
return cipher.decrypt(Nonce::from_slice(&n), &after_eph[ct_start..])
.map_err(|_| err("decryption failed"));
}

let min_n = capsule_idx + 1;
let max_n = after_eph.len().saturating_sub(16) / CAPSULE_SIZE;
for trial_n in min_n..=max_n {
let ct_start = trial_n * CAPSULE_SIZE;
if ct_start >= after_eph.len() { break; }
if let Ok(plaintext) = cipher.decrypt(Nonce::from_slice(&n), &after_eph[ct_start..]) {
return Ok(plaintext);
}
}
Err(err("decryption failed"))
}

Expand Down
21 changes: 21 additions & 0 deletions python/tests/test_encryption.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,27 @@ def test_encrypt_for_group_random_returns_nonce_needed_for_decrypt() -> None:
assert decrypted == plaintext


def test_decrypt_from_group_skips_colliding_view_tag_capsules() -> None:
sender = samp.Seed.from_bytes(bytes([0xAA] * 32))
alice = samp.Seed.from_bytes(bytes([0xAA] * 32))
bob = samp.Seed.from_bytes(bytes([0xBB] * 32))
nonce = samp.nonce_from_bytes(bytes([0xF1] * 12))
plaintext = samp.plaintext_from_bytes(b"view tag collision")
eph, capsules, ciphertext = samp.encrypt_for_group(
plaintext,
[samp.public_from_seed(alice), samp.public_from_seed(bob)],
nonce,
sender,
)

content = bytearray(bytes(eph) + bytes(capsules) + bytes(ciphertext))
content[32] = content[32 + samp.CAPSULE_SIZE]
scalar = samp.sr25519_signing_scalar(bob)

assert samp.decrypt_from_group(bytes(content), scalar, nonce, 2) == plaintext
assert samp.decrypt_from_group(bytes(content), scalar, nonce) == plaintext


def test_derive_group_ephemeral_returns_bytes() -> None:
seed = samp.Seed.from_bytes(bytes([0xAA] * 32))
nonce = samp.nonce_from_bytes(bytes([0x01] * 12))
Expand Down
Loading
Loading