Skip to main content

mls_spec/
tree.rs

1pub mod hashes;
2pub mod leaf_node;
3
4use crate::{
5    SensitiveBytes,
6    crypto::{HpkeCiphertext, HpkePublicKey},
7    defs::LeafIndex,
8    tree::{hashes::ParentNodeHash, leaf_node::LeafNode},
9};
10
11#[derive(
12    Debug,
13    Clone,
14    PartialEq,
15    Eq,
16    Hash,
17    Default,
18    tls_codec::TlsDeserialize,
19    tls_codec::TlsSerialize,
20    tls_codec::TlsSize,
21)]
22#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23pub struct RatchetTree(Vec<Option<TreeNode>>);
24
25impl RatchetTree {
26    pub fn into_inner(self) -> Vec<Option<TreeNode>> {
27        self.0
28    }
29}
30
31impl From<Vec<Option<TreeNode>>> for RatchetTree {
32    fn from(value: Vec<Option<TreeNode>>) -> Self {
33        Self(value)
34    }
35}
36
37impl std::ops::Deref for RatchetTree {
38    type Target = [Option<TreeNode>];
39
40    fn deref(&self) -> &Self::Target {
41        self.0.as_slice()
42    }
43}
44
45pub type TreeHash = SensitiveBytes;
46
47#[derive(
48    Debug,
49    Clone,
50    PartialEq,
51    Eq,
52    Hash,
53    tls_codec::TlsSerialize,
54    tls_codec::TlsDeserialize,
55    tls_codec::TlsSize,
56)]
57#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
58pub struct ParentNode {
59    pub encryption_key: HpkePublicKey,
60    pub parent_hash: ParentNodeHash,
61    pub unmerged_leaves: Vec<LeafIndex>,
62}
63
64#[derive(
65    Debug,
66    Clone,
67    PartialEq,
68    Eq,
69    Hash,
70    tls_codec::TlsSerialize,
71    tls_codec::TlsDeserialize,
72    tls_codec::TlsSize,
73)]
74#[cfg_attr(
75    feature = "serde",
76    derive(serde_repr::Serialize_repr, serde_repr::Deserialize_repr)
77)]
78#[repr(u8)]
79pub enum NodeType {
80    Reserved = 0x00,
81    Leaf = 0x01,
82    Parent = 0x02,
83}
84
85#[derive(
86    Debug,
87    Clone,
88    PartialEq,
89    Eq,
90    Hash,
91    tls_codec::TlsSerialize,
92    tls_codec::TlsDeserialize,
93    tls_codec::TlsSize,
94)]
95#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
96#[repr(u8)]
97#[allow(clippy::large_enum_variant)]
98pub enum TreeNode {
99    #[tls_codec(discriminant = "NodeType::Leaf")]
100    LeafNode(LeafNode),
101    #[tls_codec(discriminant = "NodeType::Parent")]
102    ParentNode(ParentNode),
103}
104
105impl From<LeafNode> for TreeNode {
106    fn from(value: LeafNode) -> Self {
107        Self::LeafNode(value)
108    }
109}
110
111impl From<ParentNode> for TreeNode {
112    fn from(value: ParentNode) -> Self {
113        Self::ParentNode(value)
114    }
115}
116
117impl TreeNode {
118    pub fn as_leaf_node(&self) -> Option<&LeafNode> {
119        if let Self::LeafNode(leaf_node) = &self {
120            Some(leaf_node)
121        } else {
122            None
123        }
124    }
125
126    pub fn as_leaf_node_mut(&mut self) -> Option<&mut LeafNode> {
127        if let Self::LeafNode(leaf_node) = self {
128            Some(leaf_node)
129        } else {
130            None
131        }
132    }
133
134    pub fn as_parent_node(&self) -> Option<&ParentNode> {
135        if let Self::ParentNode(parent_node) = &self {
136            Some(parent_node)
137        } else {
138            None
139        }
140    }
141
142    pub fn as_parent_node_mut(&mut self) -> Option<&mut ParentNode> {
143        if let Self::ParentNode(parent_node) = self {
144            Some(parent_node)
145        } else {
146            None
147        }
148    }
149}
150
151#[derive(Debug, Clone, PartialEq, Eq, tls_codec::TlsSerialize, tls_codec::TlsSize)]
152#[repr(u8)]
153pub enum TreeNodeRef<'a> {
154    #[tls_codec(discriminant = "NodeType::Leaf")]
155    LeafNode(&'a LeafNode),
156    #[tls_codec(discriminant = "NodeType::Parent")]
157    ParentNode(&'a ParentNode),
158}
159
160#[derive(
161    Debug,
162    Clone,
163    PartialEq,
164    Eq,
165    tls_codec::TlsSerialize,
166    tls_codec::TlsDeserialize,
167    tls_codec::TlsSize,
168)]
169#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
170pub struct UpdatePathNode {
171    pub encryption_key: HpkePublicKey,
172    pub encrypted_path_secret: Vec<HpkeCiphertext>,
173}
174
175#[derive(
176    Debug,
177    Clone,
178    PartialEq,
179    Eq,
180    tls_codec::TlsSerialize,
181    tls_codec::TlsDeserialize,
182    tls_codec::TlsSize,
183)]
184#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
185pub struct UpdatePath {
186    pub leaf_node: LeafNode,
187    pub nodes: Vec<UpdatePathNode>,
188}