Skip to main content

xmtp_id/key_package/
welcome_pointers.rs

1use openmls::prelude::UnknownExtension;
2use openmls::prelude::{AeadType, Extension};
3use prost::{EncodeError, Message};
4use xmtp_configuration::WELCOME_POINTEE_ENCRYPTION_AEAD_TYPES_EXTENSION_ID;
5use xmtp_proto::ConversionError;
6use xmtp_proto::xmtp::mls::message_contents::{
7    WelcomePointeeEncryptionAeadType as WelcomePointeeEncryptionAeadTypeProto,
8    WelcomePointeeEncryptionAeadTypesExtension as WelcomePointeeEncryptionAeadTypesExtensionProto,
9};
10
11#[derive(Debug, Clone)]
12pub struct WelcomePointersExtension {
13    pub supported_aead_types: Vec<AeadType>,
14}
15
16impl WelcomePointersExtension {
17    pub fn new(supported_aead_types: Vec<AeadType>) -> Self {
18        Self {
19            supported_aead_types,
20        }
21    }
22    pub fn available_types() -> Self {
23        Self::new(vec![Self::preferred_type()])
24    }
25    pub fn empty() -> Self {
26        Self::new(vec![])
27    }
28    pub const fn preferred_type() -> AeadType {
29        AeadType::ChaCha20Poly1305
30    }
31    pub fn compatible(&self) -> bool {
32        self.supported_aead_types.contains(&Self::preferred_type())
33    }
34}
35
36impl TryFrom<WelcomePointersExtension> for Extension {
37    type Error = EncodeError;
38
39    fn try_from(value: WelcomePointersExtension) -> Result<Self, Self::Error> {
40        let proto_val: WelcomePointeeEncryptionAeadTypesExtensionProto = value.into();
41        let mut buf = Vec::with_capacity(proto_val.encoded_len());
42        proto_val.encode(&mut buf)?;
43
44        Ok(Extension::Unknown(
45            WELCOME_POINTEE_ENCRYPTION_AEAD_TYPES_EXTENSION_ID,
46            UnknownExtension(buf),
47        ))
48    }
49}
50
51impl TryFrom<&UnknownExtension> for WelcomePointersExtension {
52    type Error = ConversionError;
53
54    fn try_from(value: &UnknownExtension) -> Result<Self, Self::Error> {
55        value.0.as_slice().try_into()
56    }
57}
58
59impl TryFrom<&[u8]> for WelcomePointersExtension {
60    type Error = ConversionError;
61
62    fn try_from(value: &[u8]) -> Result<Self, Self::Error> {
63        let proto = WelcomePointeeEncryptionAeadTypesExtensionProto::decode(value)?;
64        let supported_aead_types: Vec<AeadType> = proto
65            .supported_aead_types
66            .iter()
67            .copied()
68            .map(|aead_type| {
69                WelcomePointeeEncryptionAeadTypeProto::try_from(aead_type)
70                    .map_err(ConversionError::UnknownEnumValue)
71                    .and_then(TryInto::try_into)
72            })
73            .collect::<Result<Vec<_>, _>>()?;
74        Ok(WelcomePointersExtension {
75            supported_aead_types,
76        })
77    }
78}
79
80impl From<WelcomePointeeEncryptionAeadTypesExtensionProto> for WelcomePointersExtension {
81    fn from(value: WelcomePointeeEncryptionAeadTypesExtensionProto) -> Self {
82        Self {
83            supported_aead_types: value
84                .supported_aead_types
85                .into_iter()
86                // Ignore any values that are not valid because they cannot be used
87                .filter_map(|aead_type| {
88                    WelcomePointeeEncryptionAeadTypeProto::try_from(aead_type).ok()
89                })
90                .filter_map(|aead_type| aead_type.try_into().ok())
91                .collect(),
92        }
93    }
94}
95
96impl From<WelcomePointersExtension> for WelcomePointeeEncryptionAeadTypesExtensionProto {
97    fn from(value: WelcomePointersExtension) -> Self {
98        Self {
99            supported_aead_types: value
100                .supported_aead_types
101                .into_iter()
102                // Ignore any values that are not valid because they cannot be used
103                .filter_map(|aead_type| {
104                    WelcomePointeeEncryptionAeadTypeProto::try_from(aead_type)
105                        .map(Into::into)
106                        .ok()
107                })
108                .collect(),
109        }
110    }
111}
112
113#[cfg(test)]
114mod tests {
115    use super::*;
116
117    #[xmtp_common::test]
118    fn test_serialization() {
119        let aead_type = AeadType::ChaCha20Poly1305;
120
121        let extension = WelcomePointersExtension::available_types();
122
123        let mls_extension: Extension = extension.try_into().unwrap();
124
125        let Extension::Unknown(id, unknown_extension) = mls_extension else {
126            panic!("Expected unknown extension");
127        };
128
129        assert_eq!(id, WELCOME_POINTEE_ENCRYPTION_AEAD_TYPES_EXTENSION_ID);
130
131        let deserialized: WelcomePointersExtension = (&unknown_extension).try_into().unwrap();
132
133        assert_eq!(deserialized.supported_aead_types, vec![aead_type]);
134    }
135}