xmtp_id/key_package/
welcome_pointers.rs1use 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 .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 .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}