1use super::{ArchiveError, BackupMetadata};
2use crate::{NONCE_SIZE, util::GenericArrayExt};
3use aes_gcm::{Aes256Gcm, AesGcm, KeyInit, aead::Aead, aes::Aes256};
4use async_compression::futures::bufread::ZstdDecoder;
5use futures::{FutureExt, Stream, StreamExt};
6use futures_util::{AsyncBufRead, AsyncReadExt};
7use prost::Message;
8#[allow(deprecated)]
9use sha2::digest::{generic_array::GenericArray, typenum};
10use std::{pin::Pin, task::Poll};
11use xmtp_common::{if_native, if_wasm};
12use xmtp_proto::xmtp::device_sync::{BackupElement, backup_element::Element};
13
14if_native! {
15 mod file_import;
16 type AsyncReader = Pin<Box<dyn AsyncBufRead + Send>>;
17}
18if_wasm! {
19 type AsyncReader = Pin<Box<dyn AsyncBufRead>>;
20}
21
22pub struct ArchiveImporter {
23 pub metadata: BackupMetadata,
24 decoded: Vec<u8>,
25 decoder: ZstdDecoder<AsyncReader>,
26
27 cipher: AesGcm<Aes256, typenum::U12, typenum::U16>,
28 #[allow(deprecated)]
29 nonce: GenericArray<u8, typenum::U12>,
30}
31
32impl Stream for ArchiveImporter {
33 type Item = Result<BackupElement, ArchiveError>;
34
35 fn poll_next(
36 self: Pin<&mut Self>,
37 cx: &mut std::task::Context<'_>,
38 ) -> std::task::Poll<Option<Self::Item>> {
39 let this = self.get_mut();
40
41 let mut buffer = [0u8; 1024];
42 let mut element_len = 0;
43 loop {
44 let amount = match this.decoder.read(&mut buffer).poll_unpin(cx) {
45 Poll::Ready(Ok(amt)) => amt,
46 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)?),
47 Poll::Pending => return Poll::Pending,
48 };
49 this.decoded.extend_from_slice(&buffer[..amount]);
50
51 if element_len == 0 && this.decoded.len() >= 4 {
52 let bytes = this.decoded.drain(..4).collect::<Vec<_>>();
53 element_len = u32::from_le_bytes(bytes.try_into().expect("is 4 bytes")) as usize;
54 }
55
56 if element_len != 0 && this.decoded.len() >= element_len {
57 let decrypted_result = this
58 .cipher
59 .decrypt(&this.nonce, &this.decoded[..element_len]);
60
61 let decrypted = match decrypted_result {
62 Ok(decrypted) => decrypted,
63 Err(_) => {
65 this.nonce.decrement();
66 this.cipher
67 .decrypt(&this.nonce, &this.decoded[..element_len])
68 .inspect_err(|_| this.nonce.increment())?
69 }
70 };
71
72 let element = BackupElement::decode(&*decrypted);
73 this.decoded.drain(..element_len);
74 this.nonce.increment();
75 return Poll::Ready(Some(element.map_err(ArchiveError::from)));
76 }
77
78 if amount == 0 && this.decoded.is_empty() {
79 break;
80 }
81 }
82
83 Poll::Ready(None)
84 }
85}
86
87impl ArchiveImporter {
88 pub async fn load(mut reader: AsyncReader, key: &[u8]) -> Result<Self, ArchiveError> {
89 let mut version = [0; 2];
90 reader.read_exact(&mut version).await?;
91 let version = u16::from_le_bytes(version);
92
93 let mut nonce = [0; NONCE_SIZE];
94 reader.read_exact(&mut nonce).await?;
95
96 let mut importer = Self {
97 decoder: ZstdDecoder::new(reader),
98 decoded: vec![],
99 metadata: BackupMetadata::default(),
100
101 #[allow(deprecated)]
102 cipher: Aes256Gcm::new(GenericArray::from_slice(key)),
103 #[allow(deprecated)]
104 nonce: GenericArray::from(nonce),
105 };
106
107 let Some(Ok(BackupElement {
108 element: Some(Element::Metadata(metadata)),
109 })) = importer.next().await
110 else {
111 return Err(ArchiveError::MissingMetadata)?;
112 };
113
114 importer.metadata = BackupMetadata::from_metadata_save(metadata, version);
115 Ok(importer)
116 }
117
118 pub fn metadata(&self) -> &BackupMetadata {
119 &self.metadata
120 }
121}