use std::collections::BTreeMap;
use openmls_traits::{crypto::OpenMlsCrypto, types::Ciphersuite};
use serde::{Deserialize, Serialize};
use tls_codec::{Serialize as _, TlsSerialize, TlsSize, VLByteVec};
use crate::{
binary_tree::{
array_representation::{
direct_path, left, right, root, ParentNodeIndex, TreeNodeIndex, TreeSize,
},
LeafNodeIndex,
},
ciphersuite::Secret,
components::vc_derivation_info::{
EpochId, OperationSecret, VirtualClientOperationType, VirtualClientsError,
},
tree::secret_tree::derive_child_secrets,
utils::vector_converter,
};
const OPERATION_RATCHET_INIT_LABEL: &str = "vc operation init";
const OPERATION_GENERATION_LABEL: &str = "VC Operation Secret";
const OPERATION_RATCHET_ADVANCE_LABEL: &str = "VC Operation Ratchet";
const OPERATION_SECRET_LABEL: &str = "vc operation";
const MAXIMUM_FORWARD_DISTANCE: u32 = 1024;
const OUT_OF_ORDER_TOLERANCE: usize = 32;
#[derive(Debug, Serialize, Deserialize)]
pub struct OperationSecretTree {
leaf_nodes: Vec<Option<Secret>>,
parent_nodes: Vec<Option<Secret>>,
operation_ratchets: Vec<Option<LeafOperationRatchets>>,
size: TreeSize,
}
impl OperationSecretTree {
pub(crate) fn new(epoch_base_secret: Secret, size: TreeSize) -> Self {
let leaf_count = size.leaf_count() as usize;
let mut tree = Self {
leaf_nodes: std::iter::repeat_with(|| None).take(leaf_count).collect(),
parent_nodes: std::iter::repeat_with(|| None).take(leaf_count).collect(),
operation_ratchets: std::iter::repeat_with(|| None).take(leaf_count).collect(),
size,
};
let _ = tree.set_node(root(size), Some(epoch_base_secret));
tree
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn derive_operation_secret(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
epoch_id: &EpochId,
leaf_index: LeafNodeIndex,
operation_type: VirtualClientOperationType,
generation: u32,
operation_context: &[u8],
) -> Result<OperationSecret, VirtualClientsError> {
let ratchet = self.ratchet_mut(crypto, ciphersuite, leaf_index, operation_type)?;
let operation_generation_secret =
ratchet.generation_secret(crypto, ciphersuite, generation)?;
let context = OperationContext {
epoch_id: epoch_id.clone(),
leaf_index,
generation,
operation_type,
operation_context: operation_context.to_vec().into(),
};
context.expand_operation_secret(crypto, ciphersuite, &operation_generation_secret)
}
pub(crate) fn next_operation_secret(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
epoch_id: &EpochId,
own_leaf_index: LeafNodeIndex,
operation_type: VirtualClientOperationType,
operation_context: &[u8],
) -> Result<(u32, OperationSecret), VirtualClientsError> {
let generation = self
.ratchet_mut(crypto, ciphersuite, own_leaf_index, operation_type)?
.head_generation();
let operation_secret = self.derive_operation_secret(
crypto,
ciphersuite,
epoch_id,
own_leaf_index,
operation_type,
generation,
operation_context,
)?;
Ok((generation, operation_secret))
}
fn ratchet_mut(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
leaf_index: LeafNodeIndex,
operation_type: VirtualClientOperationType,
) -> Result<&mut OperationRatchet, VirtualClientsError> {
if leaf_index.u32() >= self.size.leaf_count() {
log::error!("vc: leaf index is larger than the operation secret tree size.");
return Err(VirtualClientsError::IndexOutOfBounds);
}
if self
.operation_ratchets
.get(leaf_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)?
.is_none()
{
self.initialize_leaf_ratchets(crypto, ciphersuite, leaf_index)?;
}
let ratchets = self
.operation_ratchets
.get_mut(leaf_index.usize())
.and_then(|ratchets| ratchets.as_mut())
.ok_or(VirtualClientsError::LibraryError)?;
Ok(ratchets.ratchet_mut(operation_type))
}
fn initialize_leaf_ratchets(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
leaf_index: LeafNodeIndex,
) -> Result<(), VirtualClientsError> {
if self.get_node(leaf_index.into())?.is_none() {
let mut empty_nodes: Vec<ParentNodeIndex> = Vec::new();
for parent_node in direct_path(leaf_index, self.size) {
empty_nodes.push(parent_node);
if self.get_node(parent_node.into())?.is_some() {
break;
}
}
empty_nodes.reverse();
for parent_node in empty_nodes {
self.derive_down(crypto, ciphersuite, parent_node)?;
}
}
let leaf_secret = self
.leaf_nodes
.get_mut(leaf_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)?
.take()
.ok_or(VirtualClientsError::LibraryError)?;
let ratchets = LeafOperationRatchets::initialize(crypto, ciphersuite, leaf_secret)?;
*self
.operation_ratchets
.get_mut(leaf_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)? = Some(ratchets);
Ok(())
}
fn derive_down(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
parent_index: ParentNodeIndex,
) -> Result<(), VirtualClientsError> {
let parent_secret = self
.parent_nodes
.get_mut(parent_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)?
.take()
.ok_or(VirtualClientsError::LibraryError)?;
let (left_secret, right_secret) =
derive_child_secrets(&parent_secret, crypto, ciphersuite)?;
self.set_node(left(parent_index), Some(left_secret))?;
self.set_node(right(parent_index), Some(right_secret))?;
Ok(())
}
fn get_node(&self, index: TreeNodeIndex) -> Result<Option<&Secret>, VirtualClientsError> {
match index {
TreeNodeIndex::Leaf(leaf_index) => Ok(self
.leaf_nodes
.get(leaf_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)?
.as_ref()),
TreeNodeIndex::Parent(parent_index) => Ok(self
.parent_nodes
.get(parent_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)?
.as_ref()),
}
}
fn set_node(
&mut self,
index: TreeNodeIndex,
secret: Option<Secret>,
) -> Result<(), VirtualClientsError> {
match index {
TreeNodeIndex::Leaf(leaf_index) => {
*self
.leaf_nodes
.get_mut(leaf_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)? = secret;
}
TreeNodeIndex::Parent(parent_index) => {
*self
.parent_nodes
.get_mut(parent_index.usize())
.ok_or(VirtualClientsError::IndexOutOfBounds)? = secret;
}
}
Ok(())
}
}
#[derive(Debug, Serialize, Deserialize)]
struct LeafOperationRatchets {
key_package: OperationRatchet,
leaf_node: OperationRatchet,
application: OperationRatchet,
}
impl LeafOperationRatchets {
fn initialize(
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
leaf_secret: Secret,
) -> Result<Self, VirtualClientsError> {
let initial_ratchet_secret =
|operation_type: VirtualClientOperationType| -> Result<Secret, VirtualClientsError> {
let context = operation_type.tls_serialize_detached()?;
Ok(leaf_secret.kdf_expand_label(
crypto,
ciphersuite,
OPERATION_RATCHET_INIT_LABEL,
&context,
ciphersuite.hash_length(),
)?)
};
Ok(Self {
key_package: OperationRatchet::new(initial_ratchet_secret(
VirtualClientOperationType::KeyPackage,
)?),
leaf_node: OperationRatchet::new(initial_ratchet_secret(
VirtualClientOperationType::LeafNode,
)?),
application: OperationRatchet::new(initial_ratchet_secret(
VirtualClientOperationType::Application,
)?),
})
}
fn ratchet_mut(&mut self, operation_type: VirtualClientOperationType) -> &mut OperationRatchet {
match operation_type {
VirtualClientOperationType::KeyPackage => &mut self.key_package,
VirtualClientOperationType::LeafNode => &mut self.leaf_node,
VirtualClientOperationType::Application => &mut self.application,
}
}
}
#[derive(Debug, Serialize, Deserialize)]
struct OperationRatchet {
ratchet_secret: Secret,
next_generation: u32,
#[serde(with = "vector_converter")]
retained_generation_secrets: BTreeMap<u32, Secret>,
}
impl OperationRatchet {
fn new(initial_ratchet_secret: Secret) -> Self {
Self {
ratchet_secret: initial_ratchet_secret,
next_generation: 0,
retained_generation_secrets: BTreeMap::new(),
}
}
fn head_generation(&self) -> u32 {
self.next_generation
}
fn generation_secret(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
generation: u32,
) -> Result<Secret, VirtualClientsError> {
if generation < self.next_generation {
return self
.retained_generation_secrets
.remove(&generation)
.ok_or(VirtualClientsError::OperationGenerationConsumed);
}
if self.next_generation < u32::MAX - MAXIMUM_FORWARD_DISTANCE
&& generation > self.next_generation + MAXIMUM_FORWARD_DISTANCE
{
log::error!(
"vc: requested operation generation {generation} is more than \
{MAXIMUM_FORWARD_DISTANCE} beyond the ratchet head {}.",
self.next_generation
);
return Err(VirtualClientsError::OperationGenerationTooDistant);
}
while self.next_generation < generation {
let skipped_generation = self.next_generation;
let skipped_secret = self.advance(crypto, ciphersuite)?;
self.retained_generation_secrets
.insert(skipped_generation, skipped_secret);
}
while self.retained_generation_secrets.len() > OUT_OF_ORDER_TOLERANCE {
self.retained_generation_secrets.pop_first();
}
self.advance(crypto, ciphersuite)
}
fn advance(
&mut self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
) -> Result<Secret, VirtualClientsError> {
if self.next_generation == u32::MAX {
return Err(VirtualClientsError::OperationRatchetTooLong);
}
let operation_generation_secret =
self.ratchet_secret
.derive_secret(crypto, ciphersuite, OPERATION_GENERATION_LABEL)?;
self.ratchet_secret = self.ratchet_secret.derive_secret(
crypto,
ciphersuite,
OPERATION_RATCHET_ADVANCE_LABEL,
)?;
self.next_generation += 1;
Ok(operation_generation_secret)
}
}
#[derive(Debug, TlsSize, TlsSerialize)]
struct OperationContext {
epoch_id: EpochId,
leaf_index: LeafNodeIndex,
generation: u32,
operation_type: VirtualClientOperationType,
operation_context: VLByteVec,
}
impl OperationContext {
fn expand_operation_secret(
&self,
crypto: &impl OpenMlsCrypto,
ciphersuite: Ciphersuite,
operation_generation_secret: &Secret,
) -> Result<OperationSecret, VirtualClientsError> {
let context = self.tls_serialize_detached()?;
let operation_secret = operation_generation_secret.kdf_expand_label(
crypto,
ciphersuite,
OPERATION_SECRET_LABEL,
&context,
ciphersuite.hash_length(),
)?;
Ok(OperationSecret::from(operation_secret))
}
}
#[cfg(test)]
mod tests {
use openmls_rust_crypto::OpenMlsRustCrypto;
use openmls_traits::{random::OpenMlsRand, OpenMlsProvider};
use super::*;
use crate::components::vc_derivation_info::EmulatorEpochSecret;
const CIPHERSUITE: Ciphersuite = Ciphersuite::MLS_128_DHKEMX25519_AES128GCM_SHA256_Ed25519;
fn setup(
leaf_count: u32,
) -> (
OpenMlsRustCrypto,
EpochId,
OperationSecretTree,
OperationSecretTree,
) {
let provider = OpenMlsRustCrypto::default();
let emulator = EmulatorEpochSecret::new(
&provider
.rand()
.random_vec(CIPHERSUITE.hash_length())
.expect("randomness"),
);
let epoch_id = emulator
.derive_epoch_id(provider.crypto(), CIPHERSUITE)
.expect("derive epoch id");
let epoch_base_secret = emulator
.derive_epoch_base_secret(provider.crypto(), CIPHERSUITE)
.expect("derive epoch base secret");
let size = TreeSize::from_leaf_count(leaf_count);
let tree_a = OperationSecretTree::new(epoch_base_secret.clone(), size);
let tree_b = OperationSecretTree::new(epoch_base_secret, size);
(provider, epoch_id, tree_a, tree_b)
}
#[test]
fn cross_instance_agreement() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(8);
let secret_a = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
LeafNodeIndex::new(2),
VirtualClientOperationType::LeafNode,
3,
b"commit context",
)
.expect("derive on tree a");
for generation in 0..3 {
tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
LeafNodeIndex::new(2),
VirtualClientOperationType::LeafNode,
generation,
b"earlier context",
)
.expect("derive earlier generation on tree b");
}
let secret_b = tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
LeafNodeIndex::new(2),
VirtualClientOperationType::LeafNode,
3,
b"commit context",
)
.expect("derive on tree b");
assert_eq!(secret_a.as_slice(), secret_b.as_slice());
}
#[test]
fn coordinates_and_context_bind_the_secret() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(8);
let derive = |tree: &mut OperationSecretTree,
leaf: u32,
operation_type: VirtualClientOperationType,
generation: u32,
context: &[u8]| {
tree.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
LeafNodeIndex::new(leaf),
operation_type,
generation,
context,
)
.expect("derive operation secret")
};
let leaf_node = VirtualClientOperationType::LeafNode;
let baseline = derive(&mut tree_a, 0, leaf_node, 0, b"ctx");
let other_generation = derive(&mut tree_a, 0, leaf_node, 1, b"ctx");
let other_leaf = derive(&mut tree_a, 1, leaf_node, 0, b"ctx");
let other_type = derive(
&mut tree_a,
0,
VirtualClientOperationType::KeyPackage,
0,
b"ctx",
);
let other_context = derive(&mut tree_b, 0, leaf_node, 0, b"other ctx");
let secrets = [
baseline.as_slice(),
other_generation.as_slice(),
other_leaf.as_slice(),
other_type.as_slice(),
other_context.as_slice(),
];
for (i, secret) in secrets.iter().enumerate() {
for other in &secrets[i + 1..] {
assert_ne!(secret, other);
}
}
}
#[test]
fn out_of_order_derivation_and_consumption() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(4);
let leaf = LeafNodeIndex::new(0);
let operation_type = VirtualClientOperationType::Application;
let context_for = |generation: u32| format!("operation {generation}").into_bytes();
let in_order: Vec<_> = (0..=5)
.map(|generation| {
tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
&context_for(generation),
)
.expect("in-order derivation")
})
.collect();
let skipped_ahead = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
5,
&context_for(5),
)
.expect("derive generation 5");
assert_eq!(skipped_ahead.as_slice(), in_order[5].as_slice());
for generation in 0..5 {
let retained = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
&context_for(generation),
)
.expect("derive retained generation");
assert_eq!(
retained.as_slice(),
in_order[generation as usize].as_slice()
);
}
for generation in 0..=5 {
let err = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
&context_for(generation),
)
.expect_err("consumed generation must fail");
assert_eq!(err, VirtualClientsError::OperationGenerationConsumed);
}
let out_of_bounds = LeafNodeIndex::new(TreeSize::from_leaf_count(4).leaf_count());
let err = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
out_of_bounds,
operation_type,
0,
b"ctx",
)
.expect_err("out-of-bounds leaf index must fail");
assert_eq!(err, VirtualClientsError::IndexOutOfBounds);
}
#[test]
fn next_operation_secret_advances_sequentially() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(4);
let leaf = LeafNodeIndex::new(1);
let operation_type = VirtualClientOperationType::KeyPackage;
for expected_generation in 0..3 {
let context = format!("key package {expected_generation}").into_bytes();
let (generation, own_secret) = tree_a
.next_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
&context,
)
.expect("next operation secret");
assert_eq!(generation, expected_generation);
let positional = tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
&context,
)
.expect("positional derivation");
assert_eq!(own_secret.as_slice(), positional.as_slice());
}
}
#[test]
fn serde_roundtrip_preserves_ratchet_state() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(4);
let leaf = LeafNodeIndex::new(2);
let operation_type = VirtualClientOperationType::LeafNode;
tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
4,
b"four",
)
.expect("derive generation 4");
tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
1,
b"one",
)
.expect("derive retained generation 1");
let serialized = serde_json::to_vec(&tree_a).expect("serialize tree");
let mut restored: OperationSecretTree =
serde_json::from_slice(&serialized).expect("deserialize tree");
let restored_secret = restored
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
2,
b"two",
)
.expect("derive retained generation after round-trip");
let positional = tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
2,
b"two",
)
.expect("positional derivation");
assert_eq!(restored_secret.as_slice(), positional.as_slice());
for (generation, context) in [(1, b"one".as_slice()), (4, b"four".as_slice())] {
let err = restored
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
context,
)
.expect_err("consumed generation must fail after round-trip");
assert_eq!(err, VirtualClientsError::OperationGenerationConsumed);
}
let (generation, _secret) = restored
.next_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
b"five",
)
.expect("next operation secret after round-trip");
assert_eq!(generation, 5);
}
#[test]
fn forward_distance_bound_rejects_without_advancing() {
let (provider, epoch_id, mut tree_a, _tree_b) = setup(4);
let leaf = LeafNodeIndex::new(0);
let operation_type = VirtualClientOperationType::LeafNode;
let err = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
MAXIMUM_FORWARD_DISTANCE + 1,
b"ctx",
)
.expect_err("generation beyond the forward distance must fail");
assert_eq!(err, VirtualClientsError::OperationGenerationTooDistant);
let (generation, _secret) = tree_a
.next_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
b"ctx",
)
.expect("next operation secret after rejected request");
assert_eq!(generation, 0);
}
#[test]
fn forward_distance_boundary_succeeds() {
let (provider, epoch_id, mut tree_a, _tree_b) = setup(4);
let leaf = LeafNodeIndex::new(0);
let operation_type = VirtualClientOperationType::LeafNode;
tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
MAXIMUM_FORWARD_DISTANCE,
b"ctx",
)
.expect("generation at the forward-distance limit must succeed");
let (generation, _secret) = tree_a
.next_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
b"ctx",
)
.expect("next operation secret after skipping to the limit");
assert_eq!(generation, MAXIMUM_FORWARD_DISTANCE + 1);
}
#[test]
fn skipping_beyond_tolerance_evicts_oldest() {
let (provider, epoch_id, mut tree_a, mut tree_b) = setup(4);
let leaf = LeafNodeIndex::new(0);
let operation_type = VirtualClientOperationType::Application;
let tolerance = OUT_OF_ORDER_TOLERANCE as u32;
let skip_to = tolerance + 1;
tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
skip_to,
b"ctx",
)
.expect("skipping derivation");
let err = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
0,
b"ctx",
)
.expect_err("evicted generation must fail");
assert_eq!(err, VirtualClientsError::OperationGenerationConsumed);
let mut in_order = Vec::new();
for generation in 0..=tolerance {
let secret = tree_b
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
b"ctx",
)
.expect("in-order derivation");
in_order.push(secret);
}
for generation in [1, tolerance] {
let retained = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
b"ctx",
)
.expect("retained generation within the window");
assert_eq!(
retained.as_slice(),
in_order[generation as usize].as_slice()
);
let err = tree_a
.derive_operation_secret(
provider.crypto(),
CIPHERSUITE,
&epoch_id,
leaf,
operation_type,
generation,
b"ctx",
)
.expect_err("second request for the same generation must fail");
assert_eq!(err, VirtualClientsError::OperationGenerationConsumed);
}
}
}