Skip to main content

openmls/framing/
message_out.rs

1//! MLS Message (Output)
2//!
3//! This module defines the [`MlsMessageOut`] structs which implements the
4//! `MLSMessage` struct as defined by the MLS specification, but is used
5//! exclusively as output of the [`MlsGroup`] API. [`MlsMessageIn`] also
6//! implements `MLSMessage`, but for inputs.
7//!
8//! The [`MlsMessageOut`] struct is meant to be serialized upon its return from
9//! a function of the `MlsGroup` API so that it can be sent to the DS.
10use tls_codec::Serialize;
11
12use super::*;
13
14use crate::{
15    key_packages::KeyPackage, messages::group_info::GroupInfo, prelude::KeyPackageBundle,
16    versions::ProtocolVersion,
17};
18
19#[cfg(any(feature = "test-utils", test))]
20use crate::messages::group_info::VerifiableGroupInfo;
21
22/// An [`MlsMessageOut`] is typically returned from an [`MlsGroup`] function and
23/// meant to be serialized and sent to the DS.
24#[derive(Debug, Clone, PartialEq, TlsSerialize, TlsSize, serde::Serialize, serde::Deserialize)]
25pub struct MlsMessageOut {
26    pub(crate) version: ProtocolVersion,
27    pub(crate) body: MlsMessageBodyOut,
28}
29
30/// MLSMessage (Body)
31///
32/// Note: Because [MlsMessageBodyOut] already discriminates between
33/// `public_message`, `private_message`, etc., we don't use the
34/// `wire_format` field. This prevents inconsistent assignments
35/// where `wire_format` contradicts the variant given in `body`.
36#[derive(Debug, PartialEq, Clone, TlsSerialize, TlsSize, serde::Serialize, serde::Deserialize)]
37#[repr(u16)]
38pub enum MlsMessageBodyOut {
39    /// Plaintext message
40    #[tls_codec(discriminant = 1)]
41    PublicMessage(PublicMessage),
42
43    /// Ciphertext message
44    #[tls_codec(discriminant = 2)]
45    PrivateMessage(PrivateMessage),
46
47    /// Welcome message
48    #[tls_codec(discriminant = 3)]
49    Welcome(Welcome),
50
51    /// Group information
52    #[tls_codec(discriminant = 4)]
53    GroupInfo(GroupInfo),
54
55    /// KeyPackage
56    #[tls_codec(discriminant = 5)]
57    #[allow(dead_code)]
58    KeyPackage(KeyPackage),
59
60    /// Targeted message (draft-ietf-mls-targeted-messages)
61    #[cfg(feature = "targeted-messages-draft")]
62    #[cfg_attr(docsrs, doc(cfg(feature = "targeted-messages-draft")))]
63    #[tls_codec(discriminant = 6)]
64    TargetedMessage(crate::targeted_messages::TargetedMessage),
65}
66
67impl From<PublicMessage> for MlsMessageOut {
68    fn from(public_message: PublicMessage) -> Self {
69        Self {
70            version: ProtocolVersion::default(),
71            body: MlsMessageBodyOut::PublicMessage(public_message),
72        }
73    }
74}
75
76impl From<PrivateMessage> for MlsMessageOut {
77    fn from(private_message: PrivateMessage) -> Self {
78        Self {
79            version: ProtocolVersion::default(),
80            body: MlsMessageBodyOut::PrivateMessage(private_message),
81        }
82    }
83}
84
85impl From<GroupInfo> for MlsMessageOut {
86    fn from(group_info: GroupInfo) -> Self {
87        Self {
88            version: group_info.group_context().protocol_version(),
89            body: MlsMessageBodyOut::GroupInfo(group_info),
90        }
91    }
92}
93
94impl From<KeyPackage> for MlsMessageOut {
95    fn from(key_package: KeyPackage) -> Self {
96        Self {
97            version: key_package.protocol_version(),
98            body: MlsMessageBodyOut::KeyPackage(key_package),
99        }
100    }
101}
102
103impl From<KeyPackageBundle> for MlsMessageOut {
104    fn from(key_package: KeyPackageBundle) -> Self {
105        Self {
106            version: key_package.key_package().protocol_version(),
107            body: MlsMessageBodyOut::KeyPackage(key_package.key_package),
108        }
109    }
110}
111
112#[cfg(feature = "targeted-messages-draft")]
113#[cfg_attr(docsrs, doc(cfg(feature = "targeted-messages-draft")))]
114impl From<crate::targeted_messages::TargetedMessage> for MlsMessageOut {
115    fn from(targeted_message: crate::targeted_messages::TargetedMessage) -> Self {
116        Self {
117            version: ProtocolVersion::default(),
118            body: MlsMessageBodyOut::TargetedMessage(targeted_message),
119        }
120    }
121}
122
123impl MlsMessageOut {
124    /// Create an [`MlsMessageOut`] from a [`PrivateMessage`], as well as the
125    /// currently used [`ProtocolVersion`].
126    pub(crate) fn from_private_message(
127        private_message: PrivateMessage,
128        version: ProtocolVersion,
129    ) -> Self {
130        Self {
131            version,
132            body: MlsMessageBodyOut::PrivateMessage(private_message),
133        }
134    }
135
136    /// Create an [`MlsMessageOut`] from a [`Welcome`] message and the currently
137    /// used [`ProtocolVersion`].
138    pub fn from_welcome(welcome: Welcome, version: ProtocolVersion) -> Self {
139        MlsMessageOut {
140            version,
141            body: MlsMessageBodyOut::Welcome(welcome),
142        }
143    }
144
145    /// Serializes the message to a byte vector. Returns [`MlsMessageError::UnableToEncode`] on failure.
146    pub fn to_bytes(&self) -> Result<Vec<u8>, MlsMessageError> {
147        self.tls_serialize_detached()
148            .map_err(|_| MlsMessageError::UnableToEncode)
149    }
150
151    /// Returns a reference to the contents of this [`MlsMessageOut`].
152    pub fn body(&self) -> &MlsMessageBodyOut {
153        &self.body
154    }
155}
156
157// Convenience functions for tests and test-utils
158
159#[cfg(any(feature = "test-utils", test))]
160impl MlsMessageOut {
161    /// Turn an [`MlsMessageOut`] into a [`Welcome`].
162    #[cfg(any(feature = "test-utils", test))]
163    pub fn into_welcome(self) -> Option<Welcome> {
164        match self.body {
165            MlsMessageBodyOut::Welcome(w) => Some(w),
166            _ => None,
167        }
168    }
169
170    #[cfg(any(feature = "test-utils", test))]
171    pub fn into_protocol_message(self) -> Option<ProtocolMessage> {
172        let mls_message_in: MlsMessageIn = self.into();
173
174        match mls_message_in.extract() {
175            MlsMessageBodyIn::PublicMessage(pm) => Some(pm.into()),
176            MlsMessageBodyIn::PrivateMessage(pm) => Some(pm.into()),
177            _ => None,
178        }
179    }
180
181    #[cfg(any(feature = "test-utils", test))]
182    pub fn into_verifiable_group_info(self) -> Option<VerifiableGroupInfo> {
183        match self.body {
184            MlsMessageBodyOut::GroupInfo(group_info) => {
185                Some(group_info.into_verifiable_group_info())
186            }
187            _ => None,
188        }
189    }
190}
191
192impl From<MlsMessageBodyOut> for MlsMessageBodyIn {
193    fn from(value: MlsMessageBodyOut) -> Self {
194        match value {
195            MlsMessageBodyOut::PublicMessage(pm) => MlsMessageBodyIn::PublicMessage(pm.into()),
196            MlsMessageBodyOut::PrivateMessage(pm) => MlsMessageBodyIn::PrivateMessage(pm.into()),
197            MlsMessageBodyOut::Welcome(w) => MlsMessageBodyIn::Welcome(w),
198            MlsMessageBodyOut::GroupInfo(gi) => {
199                MlsMessageBodyIn::GroupInfo(gi.into_verifiable_group_info())
200            }
201            MlsMessageBodyOut::KeyPackage(kp) => MlsMessageBodyIn::KeyPackage(kp.into()),
202            #[cfg(feature = "targeted-messages-draft")]
203            MlsMessageBodyOut::TargetedMessage(tm) => MlsMessageBodyIn::TargetedMessage(tm.into()),
204        }
205    }
206}
207
208impl From<MlsMessageOut> for MlsMessageIn {
209    fn from(mls_message_out: MlsMessageOut) -> Self {
210        let MlsMessageOut { version, body } = mls_message_out;
211        Self {
212            version,
213            body: body.into(),
214        }
215    }
216}
217
218// The following two `From` implementations break abstraction layers and MUST
219// NOT be made available outside of tests or "test-utils".
220
221#[cfg(any(feature = "test-utils", test))]
222impl From<MlsMessageIn> for MlsMessageOut {
223    fn from(mls_message: MlsMessageIn) -> Self {
224        let version = mls_message.version;
225        let body = match mls_message.body {
226            MlsMessageBodyIn::Welcome(w) => MlsMessageBodyOut::Welcome(w),
227            MlsMessageBodyIn::GroupInfo(gi) => MlsMessageBodyOut::GroupInfo(gi.into()),
228            MlsMessageBodyIn::KeyPackage(kp) => MlsMessageBodyOut::KeyPackage(kp.into()),
229            MlsMessageBodyIn::PublicMessage(pm) => MlsMessageBodyOut::PublicMessage(pm.into()),
230            MlsMessageBodyIn::PrivateMessage(pm) => MlsMessageBodyOut::PrivateMessage(pm.into()),
231            #[cfg(feature = "targeted-messages-draft")]
232            MlsMessageBodyIn::TargetedMessage(tm) => MlsMessageBodyOut::TargetedMessage(tm.into()),
233        };
234        Self { version, body }
235    }
236}