1use std::borrow::Cow;
6
7use ruma_common::{
8 OwnedTransactionId,
9 serde::{Base64, JsonObject},
10};
11use ruma_macros::EventContent;
12use serde::{Deserialize, Serialize};
13use serde_json::{Value as JsonValue, from_value as from_json_value};
14
15use super::{
16 HashAlgorithm, KeyAgreementProtocol, MessageAuthenticationCode, ShortAuthenticationString,
17};
18use crate::relation::Reference;
19
20#[derive(Clone, Debug, Deserialize, Serialize, EventContent)]
24#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
25#[ruma_event(type = "m.key.verification.accept", kind = ToDevice)]
26pub struct ToDeviceKeyVerificationAcceptEventContent {
27 pub transaction_id: OwnedTransactionId,
31
32 #[serde(flatten)]
34 pub method: AcceptMethod,
35}
36
37impl ToDeviceKeyVerificationAcceptEventContent {
38 pub fn new(transaction_id: OwnedTransactionId, method: AcceptMethod) -> Self {
41 Self { transaction_id, method }
42 }
43}
44
45#[derive(Clone, Debug, Deserialize, Serialize, EventContent)]
49#[ruma_event(type = "m.key.verification.accept", kind = MessageLike)]
50#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
51pub struct KeyVerificationAcceptEventContent {
52 #[serde(flatten)]
54 pub method: AcceptMethod,
55
56 #[serde(rename = "m.relates_to")]
58 pub relates_to: Reference,
59}
60
61impl KeyVerificationAcceptEventContent {
62 pub fn new(method: AcceptMethod, relates_to: Reference) -> Self {
65 Self { method, relates_to }
66 }
67}
68
69#[derive(Clone, Debug, Serialize)]
71#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
72#[serde(untagged)]
73pub enum AcceptMethod {
74 SasV1(SasV1Content),
76
77 #[doc(hidden)]
79 _Custom(_CustomAcceptMethodContent),
80}
81
82impl AcceptMethod {
83 pub fn data(&self) -> Cow<'_, JsonObject> {
88 fn serialize<T: Serialize>(obj: T) -> JsonObject {
89 match serde_json::to_value(obj).expect("accept method serialization to succeed") {
90 JsonValue::Object(mut obj) => {
91 obj.remove("method");
92 obj
93 }
94 _ => panic!("all accept method variants must serialize to objects"),
95 }
96 }
97
98 match self {
99 Self::SasV1(c) => Cow::Owned(serialize(c)),
100 Self::_Custom(c) => Cow::Borrowed(&c.data),
101 }
102 }
103}
104
105impl<'de> Deserialize<'de> for AcceptMethod {
106 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
107 where
108 D: serde::Deserializer<'de>,
109 {
110 let data = JsonObject::deserialize(deserializer)?;
111
112 Ok(match from_json_value(data.clone().into()) {
113 Ok(sas_v1_content) => AcceptMethod::SasV1(sas_v1_content),
114 Err(_) => AcceptMethod::_Custom(_CustomAcceptMethodContent { data }),
115 })
116 }
117}
118
119#[doc(hidden)]
121#[derive(Clone, Debug, Serialize)]
122pub struct _CustomAcceptMethodContent {
123 #[serde(flatten)]
125 data: JsonObject,
126}
127
128#[derive(Clone, Debug, Deserialize, Serialize)]
130#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
131pub struct SasV1Content {
132 pub key_agreement_protocol: KeyAgreementProtocol,
135
136 pub hash: HashAlgorithm,
139
140 pub message_authentication_code: MessageAuthenticationCode,
143
144 pub short_authentication_string: Vec<ShortAuthenticationString>,
150
151 pub commitment: Base64,
155}
156
157#[derive(Debug)]
159#[allow(clippy::exhaustive_structs)]
160pub struct SasV1ContentInit {
161 pub key_agreement_protocol: KeyAgreementProtocol,
164
165 pub hash: HashAlgorithm,
168
169 pub message_authentication_code: MessageAuthenticationCode,
171
172 pub short_authentication_string: Vec<ShortAuthenticationString>,
178
179 pub commitment: Base64,
183}
184
185impl From<SasV1ContentInit> for SasV1Content {
186 fn from(init: SasV1ContentInit) -> Self {
188 SasV1Content {
189 hash: init.hash,
190 key_agreement_protocol: init.key_agreement_protocol,
191 message_authentication_code: init.message_authentication_code,
192 short_authentication_string: init.short_authentication_string,
193 commitment: init.commitment,
194 }
195 }
196}
197
198#[cfg(test)]
199mod tests {
200 use assert_matches::assert_matches;
201 use ruma_common::{
202 canonical_json::assert_to_canonical_json_eq,
203 event_id,
204 serde::{Base64, Raw},
205 };
206 use serde_json::{from_value as from_json_value, json};
207 use strass::assert_let;
208
209 use super::{
210 AcceptMethod, HashAlgorithm, KeyAgreementProtocol, KeyVerificationAcceptEventContent,
211 MessageAuthenticationCode, SasV1Content, ShortAuthenticationString,
212 ToDeviceKeyVerificationAcceptEventContent,
213 };
214 use crate::{ToDeviceEvent, relation::Reference};
215
216 #[test]
217 fn to_device_serialization() {
218 let key_verification_accept_content = ToDeviceKeyVerificationAcceptEventContent {
219 transaction_id: "456".into(),
220 method: AcceptMethod::SasV1(SasV1Content {
221 hash: HashAlgorithm::Sha256,
222 key_agreement_protocol: KeyAgreementProtocol::Curve25519,
223 message_authentication_code: MessageAuthenticationCode::HkdfHmacSha256V2,
224 short_authentication_string: vec![ShortAuthenticationString::Decimal],
225 commitment: Base64::new(b"hello".to_vec()),
226 }),
227 };
228
229 assert_to_canonical_json_eq!(
230 key_verification_accept_content,
231 json!({
232 "transaction_id": "456",
233 "commitment": "aGVsbG8",
234 "key_agreement_protocol": "curve25519",
235 "hash": "sha256",
236 "message_authentication_code": "hkdf-hmac-sha256.v2",
237 "short_authentication_string": ["decimal"],
238 }),
239 );
240 }
241
242 #[test]
243 fn in_room_serialization() {
244 let event_id = event_id!("$1598361704261elfgc:localhost");
245
246 let key_verification_accept_content = KeyVerificationAcceptEventContent {
247 relates_to: Reference { event_id: event_id.to_owned() },
248 method: AcceptMethod::SasV1(SasV1Content {
249 hash: HashAlgorithm::Sha256,
250 key_agreement_protocol: KeyAgreementProtocol::Curve25519,
251 message_authentication_code: MessageAuthenticationCode::HkdfHmacSha256V2,
252 short_authentication_string: vec![ShortAuthenticationString::Decimal],
253 commitment: Base64::new(b"hello".to_vec()),
254 }),
255 };
256
257 assert_to_canonical_json_eq!(
258 key_verification_accept_content,
259 json!({
260 "commitment": "aGVsbG8",
261 "key_agreement_protocol": "curve25519",
262 "hash": "sha256",
263 "message_authentication_code": "hkdf-hmac-sha256.v2",
264 "short_authentication_string": ["decimal"],
265 "m.relates_to": {
266 "rel_type": "m.reference",
267 "event_id": event_id,
268 },
269 }),
270 );
271 }
272
273 #[test]
274 fn to_device_deserialization() {
275 let json = json!({
276 "transaction_id": "456",
277 "commitment": "aGVsbG8",
278 "hash": "sha256",
279 "key_agreement_protocol": "curve25519",
280 "message_authentication_code": "hkdf-hmac-sha256.v2",
281 "short_authentication_string": ["decimal"]
282 });
283
284 let content = from_json_value::<ToDeviceKeyVerificationAcceptEventContent>(json).unwrap();
286 assert_eq!(content.transaction_id, "456");
287
288 assert_let!(AcceptMethod::SasV1(sas) = content.method);
289 assert_eq!(sas.commitment.encode(), "aGVsbG8");
290 assert_eq!(sas.hash, HashAlgorithm::Sha256);
291 assert_eq!(sas.key_agreement_protocol, KeyAgreementProtocol::Curve25519);
292 assert_eq!(sas.message_authentication_code, MessageAuthenticationCode::HkdfHmacSha256V2);
293 assert_eq!(sas.short_authentication_string, vec![ShortAuthenticationString::Decimal]);
294
295 let json = json!({
296 "content": {
297 "commitment": "aGVsbG8",
298 "transaction_id": "456",
299 "key_agreement_protocol": "curve25519",
300 "hash": "sha256",
301 "message_authentication_code": "hkdf-hmac-sha256.v2",
302 "short_authentication_string": ["decimal"]
303 },
304 "type": "m.key.verification.accept",
305 "sender": "@example:localhost",
306 });
307
308 let ev = from_json_value::<ToDeviceEvent<ToDeviceKeyVerificationAcceptEventContent>>(json)
309 .unwrap();
310 assert_eq!(ev.content.transaction_id, "456");
311 assert_eq!(ev.sender, "@example:localhost");
312
313 assert_let!(AcceptMethod::SasV1(sas) = ev.content.method);
314 assert_eq!(sas.commitment.encode(), "aGVsbG8");
315 assert_eq!(sas.hash, HashAlgorithm::Sha256);
316 assert_eq!(sas.key_agreement_protocol, KeyAgreementProtocol::Curve25519);
317 assert_eq!(sas.message_authentication_code, MessageAuthenticationCode::HkdfHmacSha256V2);
318 assert_eq!(sas.short_authentication_string, vec![ShortAuthenticationString::Decimal]);
319 }
320
321 #[test]
322 fn in_room_deserialization() {
323 let json = json!({
324 "commitment": "aGVsbG8",
325 "hash": "sha256",
326 "key_agreement_protocol": "curve25519",
327 "message_authentication_code": "hkdf-hmac-sha256.v2",
328 "short_authentication_string": ["decimal"],
329 "m.relates_to": {
330 "rel_type": "m.reference",
331 "event_id": "$1598361704261elfgc:localhost",
332 }
333 });
334
335 let content = from_json_value::<KeyVerificationAcceptEventContent>(json).unwrap();
337 assert_eq!(content.relates_to.event_id, "$1598361704261elfgc:localhost");
338
339 assert_let!(AcceptMethod::SasV1(sas) = content.method);
340 assert_eq!(sas.commitment.encode(), "aGVsbG8");
341 assert_eq!(sas.hash, HashAlgorithm::Sha256);
342 assert_eq!(sas.key_agreement_protocol, KeyAgreementProtocol::Curve25519);
343 assert_eq!(sas.message_authentication_code, MessageAuthenticationCode::HkdfHmacSha256V2);
344 assert_eq!(sas.short_authentication_string, vec![ShortAuthenticationString::Decimal]);
345 }
346
347 #[test]
348 fn in_room_serialization_roundtrip() {
349 let event_id = event_id!("$1598361704261elfgc:localhost");
350
351 let content = KeyVerificationAcceptEventContent {
352 relates_to: Reference { event_id: event_id.to_owned() },
353 method: AcceptMethod::SasV1(SasV1Content {
354 hash: HashAlgorithm::Sha256,
355 key_agreement_protocol: KeyAgreementProtocol::Curve25519,
356 message_authentication_code: MessageAuthenticationCode::HkdfHmacSha256V2,
357 short_authentication_string: vec![ShortAuthenticationString::Decimal],
358 commitment: Base64::new(b"hello".to_vec()),
359 }),
360 };
361
362 let json_content = Raw::new(&content).unwrap();
363 let deser_content = json_content.deserialize().unwrap();
364
365 assert_matches!(deser_content.method, AcceptMethod::SasV1(_));
366 assert_eq!(deser_content.relates_to.event_id, event_id);
367 }
368
369 #[test]
370 fn custom_to_device_serialization_roundtrip() {
371 let json = json!({
372 "transaction_id": "456",
373 "test": "field",
374 });
375
376 let content =
377 from_json_value::<ToDeviceKeyVerificationAcceptEventContent>(json.clone()).unwrap();
378
379 assert_eq!(content.transaction_id, "456");
380 let data = &*content.method.data();
381 assert_eq!(data.len(), 1);
382 assert_eq!(data["test"], "field");
383
384 assert_to_canonical_json_eq!(content, json);
385 }
386}