Skip to main content

mls_spec/tree/
leaf_node.rs

1use crate::{
2    SensitiveBytes,
3    credential::Credential,
4    crypto::{HpkePublicKey, HpkePublicKeyRef, SignaturePublicKey, SignaturePublicKeyRef},
5    defs::{Capabilities, LeafIndex},
6    group::{GroupIdRef, KeyPackageLifetime, RequiredCapabilities, extensions::Extension},
7};
8
9#[derive(
10    Debug,
11    Clone,
12    Copy,
13    PartialEq,
14    Eq,
15    tls_codec::TlsSize,
16    tls_codec::TlsDeserialize,
17    tls_codec::TlsSerialize,
18    strum::Display,
19)]
20#[cfg_attr(
21    feature = "serde",
22    derive(serde_repr::Serialize_repr, serde_repr::Deserialize_repr)
23)]
24#[repr(u8)]
25pub enum LeafNodeSourceType {
26    Reserved = 0x00,
27    KeyPackage = 0x01,
28    Update = 0x02,
29    Commit = 0x03,
30}
31
32#[derive(
33    Debug,
34    Clone,
35    PartialEq,
36    Eq,
37    Hash,
38    tls_codec::TlsSerialize,
39    tls_codec::TlsDeserialize,
40    tls_codec::TlsSize,
41)]
42#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
43#[repr(u8)]
44pub enum LeafNodeSource {
45    #[tls_codec(discriminant = "LeafNodeSourceType::KeyPackage")]
46    KeyPackage { lifetime: KeyPackageLifetime },
47    #[tls_codec(discriminant = "LeafNodeSourceType::Update")]
48    Update,
49    #[tls_codec(discriminant = "LeafNodeSourceType::Commit")]
50    Commit { parent_hash: SensitiveBytes },
51}
52
53impl From<&LeafNodeSource> for LeafNodeSourceType {
54    fn from(value: &LeafNodeSource) -> Self {
55        match value {
56            LeafNodeSource::KeyPackage { .. } => Self::KeyPackage,
57            LeafNodeSource::Update => Self::Update,
58            LeafNodeSource::Commit { .. } => Self::Commit,
59        }
60    }
61}
62
63#[derive(Debug, Copy, Clone, PartialEq, Eq, tls_codec::TlsSerialize, tls_codec::TlsSize)]
64#[cfg_attr(feature = "serde", derive(serde::Serialize))]
65pub struct LeafNodeMemberInfo<'a> {
66    #[tls_codec(with = "crate::tlspl::bytes")]
67    pub group_id: GroupIdRef<'a>,
68    pub leaf_index: LeafIndex,
69}
70
71#[derive(
72    Debug,
73    Clone,
74    PartialEq,
75    Eq,
76    Hash,
77    tls_codec::TlsSerialize,
78    tls_codec::TlsDeserialize,
79    tls_codec::TlsSize,
80)]
81#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
82pub struct LeafNode {
83    pub encryption_key: HpkePublicKey,
84    pub signature_key: SignaturePublicKey,
85    pub credential: Credential,
86    pub capabilities: Capabilities,
87    pub source: LeafNodeSource,
88    pub extensions: Vec<Extension>,
89    pub signature: SensitiveBytes,
90}
91
92impl LeafNode {
93    #[inline]
94    pub fn requires_member_info(&self) -> bool {
95        matches!(
96            self.source,
97            LeafNodeSource::Update | LeafNodeSource::Commit { .. }
98        )
99    }
100
101    pub fn parent_hash(&self) -> Option<&[u8]> {
102        match &self.source {
103            LeafNodeSource::Commit { parent_hash } => Some(parent_hash),
104            _ => None,
105        }
106    }
107
108    pub fn to_tbs<'a>(
109        &'a self,
110        member_info: Option<LeafNodeMemberInfo<'a>>,
111    ) -> Option<LeafNodeTBS<'a>> {
112        Some(LeafNodeTBS {
113            encryption_key: &self.encryption_key,
114            signature_key: &self.signature_key,
115            credential: &self.credential,
116            capabilities: &self.capabilities,
117            source: &self.source,
118            extensions: &self.extensions,
119            member_info: if self.requires_member_info() {
120                // Invalid because in those context we should have a valid member_info
121                Some(member_info?)
122            } else {
123                None
124            },
125        })
126    }
127
128    pub fn application_id(&self) -> Option<&[u8]> {
129        self.extensions.iter().find_map(|ext| {
130            if let Extension::ApplicationId(app_id) = ext {
131                Some(app_id.as_slice())
132            } else {
133                None
134            }
135        })
136    }
137
138    pub fn supports_required_capabilities(&self, required_caps: &RequiredCapabilities) -> bool {
139        if !required_caps.extension_types.iter().all(|req_ext| {
140            req_ext.is_grease_value()
141                || req_ext.is_spec_default()
142                || self.capabilities.extensions.contains(req_ext)
143        }) {
144            return false;
145        }
146
147        if !required_caps.proposal_types.iter().all(|req_prop| {
148            req_prop.is_grease_value()
149                || req_prop.is_spec_default()
150                || self.capabilities.proposals.contains(req_prop)
151        }) {
152            return false;
153        }
154
155        if !required_caps.credential_types.iter().all(|req_cred| {
156            req_cred.is_grease_value() || self.capabilities.credentials.contains(req_cred)
157        }) {
158            return false;
159        }
160
161        true
162    }
163}
164
165#[derive(Debug, PartialEq, Eq)]
166pub struct LeafNodeTBS<'a> {
167    pub encryption_key: HpkePublicKeyRef<'a>,
168    pub signature_key: SignaturePublicKeyRef<'a>,
169    pub credential: &'a Credential,
170    pub capabilities: &'a Capabilities,
171    pub source: &'a LeafNodeSource,
172    pub extensions: &'a Vec<Extension>,
173    pub member_info: Option<LeafNodeMemberInfo<'a>>,
174}
175
176impl tls_codec::Size for LeafNodeTBS<'_> {
177    fn tls_serialized_len(&self) -> usize {
178        self.encryption_key.tls_serialized_len()
179            + self.signature_key.tls_serialized_len()
180            + self.credential.tls_serialized_len()
181            + self.capabilities.tls_serialized_len()
182            + self.source.tls_serialized_len()
183            + self.extensions.tls_serialized_len()
184            + self.member_info.map_or(0, |mi| mi.tls_serialized_len())
185    }
186}
187
188impl tls_codec::Serialize for LeafNodeTBS<'_> {
189    fn tls_serialize<W: std::io::Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
190        let mut written = 0;
191        written += crate::tlspl::bytes::tls_serialize(self.encryption_key, writer)?;
192        written += crate::tlspl::bytes::tls_serialize(self.signature_key, writer)?;
193        written += self.credential.tls_serialize(writer)?;
194        written += self.capabilities.tls_serialize(writer)?;
195        written += self.source.tls_serialize(writer)?;
196        written += self.extensions.tls_serialize(writer)?;
197        if let Some(member_info) = self.member_info {
198            written += member_info.tls_serialize(writer)?;
199        }
200
201        Ok(written)
202    }
203}