Skip to main content

xmtp_archive/
importer.rs

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                    // Attempt to decrypt using a decremented nonce to support legacy archives.
64                    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}