Skip to main content

whatsapp_rust_sqlite_storage/
sqlite_store.rs

1use crate::schema::*;
2use async_trait::async_trait;
3use bytes::Bytes;
4use diesel::prelude::*;
5use diesel::r2d2::{ConnectionManager, Pool};
6use diesel::result::{DatabaseErrorKind, Error as DieselError};
7use diesel::sqlite::SqliteConnection;
8use diesel::upsert::excluded;
9use diesel_migrations::{EmbeddedMigrations, MigrationHarness, embed_migrations};
10use log::warn;
11use std::sync::Arc;
12use wacore::appstate::hash::HashState;
13use wacore::appstate::processor::AppStateMutationMAC;
14use wacore::libsignal::protocol::{KeyPair, PrivateKey, PublicKey};
15use wacore::store::Device as CoreDevice;
16use wacore::store::error::{Result, StoreError};
17use wacore::store::traits::*;
18
19/// Internal error type that preserves the Diesel error for structured matching
20/// before converting to `StoreError`. Used in retry loops where we need to
21/// distinguish retriable SQLite lock errors from other failures.
22enum DieselOrStore {
23    Diesel(DieselError),
24    Store(StoreError),
25}
26
27impl From<DieselOrStore> for StoreError {
28    fn from(e: DieselOrStore) -> Self {
29        match e {
30            DieselOrStore::Diesel(e) => StoreError::Database(Box::new(e)),
31            DieselOrStore::Store(e) => e,
32        }
33    }
34}
35
36/// Check if a Diesel error represents a retriable SQLite lock contention.
37///
38/// SQLite BUSY (error code 5) and LOCKED (error code 6) both map to
39/// `DatabaseError(Unknown, _)` in Diesel. We inspect the error message
40/// from `sqlite3_errmsg()` to distinguish them from other unknown errors.
41fn is_retriable_sqlite_error(error: &DieselError) -> bool {
42    match error {
43        DieselError::DatabaseError(DatabaseErrorKind::Unknown, info) => {
44            let msg = info.message();
45            msg.contains("locked") || msg.contains("busy")
46        }
47        _ => false,
48    }
49}
50
51const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations");
52
53type SqlitePool = Pool<ConnectionManager<SqliteConnection>>;
54
55/// Row representation for the `device` table.
56///
57/// Field order must match the column order in `schema::device`.
58/// Using a named struct instead of a positional tuple so fields are
59/// accessed by name, reducing the risk of mix-ups when columns are added.
60#[derive(Queryable, Selectable)]
61#[diesel(table_name = crate::schema::device)]
62#[allow(dead_code)]
63struct DeviceRow {
64    id: i32,
65    lid: String,
66    pn: String,
67    registration_id: i32,
68    noise_key: Vec<u8>,
69    identity_key: Vec<u8>,
70    signed_pre_key: Vec<u8>,
71    signed_pre_key_id: i32,
72    signed_pre_key_signature: Vec<u8>,
73    adv_secret_key: Vec<u8>,
74    account: Option<Vec<u8>>,
75    push_name: String,
76    app_version_primary: i32,
77    app_version_secondary: i32,
78    app_version_tertiary: i64,
79    app_version_last_fetched_ms: i64,
80    edge_routing_info: Option<Vec<u8>>,
81    props_hash: Option<String>,
82    next_pre_key_id: i32,
83    nct_salt: Option<Vec<u8>>,
84    server_has_prekeys: bool,
85    server_cert_chain: Option<Vec<u8>>,
86}
87
88#[derive(Clone)]
89pub struct SqliteStore {
90    pub(crate) pool: SqlitePool,
91    pub(crate) db_semaphore: Arc<tokio::sync::Semaphore>,
92    pub(crate) database_path: String,
93    device_id: i32,
94}
95
96#[derive(Debug, Clone, Copy)]
97struct ConnectionOptions;
98
99impl diesel::r2d2::CustomizeConnection<SqliteConnection, diesel::r2d2::Error>
100    for ConnectionOptions
101{
102    fn on_acquire(
103        &self,
104        conn: &mut SqliteConnection,
105    ) -> std::result::Result<(), diesel::r2d2::Error> {
106        diesel::sql_query("PRAGMA busy_timeout = 30000;")
107            .execute(conn)
108            .map_err(diesel::r2d2::Error::QueryError)?;
109        diesel::sql_query("PRAGMA synchronous = NORMAL;")
110            .execute(conn)
111            .map_err(diesel::r2d2::Error::QueryError)?;
112        diesel::sql_query("PRAGMA cache_size = 512;")
113            .execute(conn)
114            .map_err(diesel::r2d2::Error::QueryError)?;
115        diesel::sql_query("PRAGMA temp_store = memory;")
116            .execute(conn)
117            .map_err(diesel::r2d2::Error::QueryError)?;
118        diesel::sql_query("PRAGMA foreign_keys = ON;")
119            .execute(conn)
120            .map_err(diesel::r2d2::Error::QueryError)?;
121        Ok(())
122    }
123}
124
125fn parse_database_path(database_url: &str) -> Result<String> {
126    // Reject in-memory databases
127    if database_url == ":memory:" {
128        return Err(StoreError::InvalidConfig(
129            "Snapshot not supported for in-memory databases".to_string(),
130        ));
131    }
132
133    // Strip query string and fragment
134    let path = database_url
135        .split(['?', '#'])
136        .next()
137        .unwrap_or(database_url);
138
139    // Remove sqlite:// prefix if present
140    let path = path.trim_start_matches("sqlite://");
141
142    // Check if the resulting path looks like an in-memory marker
143    if path == ":memory:" || path.starts_with(":memory:?") {
144        return Err(StoreError::InvalidConfig(
145            "Snapshot not supported for in-memory databases".to_string(),
146        ));
147    }
148
149    Ok(path.to_string())
150}
151
152impl SqliteStore {
153    pub async fn new(database_url: &str) -> std::result::Result<Self, StoreError> {
154        let manager = ConnectionManager::<SqliteConnection>::new(database_url);
155
156        let pool_size = 2;
157
158        let pool = Pool::builder()
159            .max_size(pool_size)
160            .connection_customizer(Box::new(ConnectionOptions))
161            .build(manager)
162            .map_err(|e| StoreError::Connection(Box::new(e)))?;
163
164        let pool_clone = pool.clone();
165        tokio::task::spawn_blocking(move || -> std::result::Result<(), StoreError> {
166            let mut conn = pool_clone
167                .get()
168                .map_err(|e| StoreError::Connection(Box::new(e)))?;
169
170            diesel::sql_query("PRAGMA journal_mode = WAL;")
171                .execute(&mut conn)
172                .map_err(|e| StoreError::Database(Box::new(e)))?;
173
174            conn.run_pending_migrations(MIGRATIONS)
175                .map_err(StoreError::Migration)?;
176
177            Ok(())
178        })
179        .await
180        .map_err(|e| StoreError::Database(Box::new(e)))??;
181
182        let database_path = parse_database_path(database_url)?;
183
184        Ok(Self {
185            pool,
186            db_semaphore: Arc::new(tokio::sync::Semaphore::new(1)),
187            database_path,
188            device_id: 1,
189        })
190    }
191
192    pub async fn new_for_device(
193        database_url: &str,
194        device_id: i32,
195    ) -> std::result::Result<Self, StoreError> {
196        let mut store = Self::new(database_url).await?;
197        store.device_id = device_id;
198        Ok(store)
199    }
200
201    pub fn device_id(&self) -> i32 {
202        self.device_id
203    }
204
205    async fn with_semaphore<F, T>(&self, f: F) -> Result<T>
206    where
207        F: FnOnce() -> Result<T> + Send + 'static,
208        T: Send + 'static,
209    {
210        let permit = self
211            .db_semaphore
212            .clone()
213            .acquire_owned()
214            .await
215            .map_err(|e| StoreError::Database(Box::new(e)))?;
216        let result = tokio::task::spawn_blocking(move || {
217            let res = f();
218            drop(permit);
219            res
220        })
221        .await
222        .map_err(|e| StoreError::Database(Box::new(e)))??;
223        Ok(result)
224    }
225
226    /// Execute a database operation with semaphore serialization and retry on
227    /// transient SQLite lock/busy errors. Mirrors WhatsApp Web's PromiseQueue
228    /// pattern that serializes database commits to avoid concurrent write contention.
229    async fn with_retry<F, T>(&self, op_name: &str, make_op: F) -> Result<T>
230    where
231        F: Fn() -> Box<
232            dyn FnOnce(&mut SqliteConnection) -> std::result::Result<T, DieselError> + Send,
233        >,
234        T: Send + 'static,
235    {
236        const MAX_RETRIES: u32 = 5;
237
238        for attempt in 0..=MAX_RETRIES {
239            let permit = self
240                .db_semaphore
241                .clone()
242                .acquire_owned()
243                .await
244                .map_err(|e| StoreError::Database(Box::new(e)))?;
245
246            let pool = self.pool.clone();
247            let op = make_op();
248
249            let result =
250                tokio::task::spawn_blocking(move || -> std::result::Result<T, DieselOrStore> {
251                    let _permit = permit;
252                    let mut conn = pool
253                        .get()
254                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
255                    op(&mut conn).map_err(DieselOrStore::Diesel)
256                })
257                .await;
258
259            match result {
260                Ok(Ok(val)) => return Ok(val),
261                Ok(Err(DieselOrStore::Diesel(ref e)))
262                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
263                {
264                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
265                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
266                }
267                Ok(Err(e)) => return Err(e.into()),
268                Err(e) => return Err(StoreError::Database(Box::new(e))),
269            }
270        }
271
272        Err(StoreError::RetriesExhausted {
273            op: op_name.to_string(),
274        })
275    }
276
277    fn serialize_keypair(&self, key_pair: &KeyPair) -> Result<Vec<u8>> {
278        let mut bytes = Vec::with_capacity(64);
279        bytes.extend_from_slice(key_pair.private_key.serialize());
280        bytes.extend_from_slice(key_pair.public_key.public_key_bytes());
281        Ok(bytes)
282    }
283
284    fn deserialize_keypair(&self, bytes: &[u8]) -> Result<KeyPair> {
285        if bytes.len() != 64 {
286            return Err(StoreError::Validation(format!(
287                "Invalid KeyPair length: {}",
288                bytes.len()
289            )));
290        }
291
292        let private_key = PrivateKey::deserialize(&bytes[0..32])
293            .map_err(|e| StoreError::Serialization(Box::new(e)))?;
294        let public_key = PublicKey::from_djb_public_key_bytes(&bytes[32..64])
295            .map_err(|e| StoreError::Serialization(Box::new(e)))?;
296
297        Ok(KeyPair::new(public_key, private_key))
298    }
299
300    pub async fn save_device_data_for_device(
301        &self,
302        device_id: i32,
303        device_data: &CoreDevice,
304    ) -> Result<()> {
305        // Use Arc so retry clones are just atomic increments, not deep copies.
306        let noise_key_data: Arc<[u8]> = self.serialize_keypair(&device_data.noise_key)?.into();
307        let identity_key_data: Arc<[u8]> =
308            self.serialize_keypair(&device_data.identity_key)?.into();
309        let signed_pre_key_data: Arc<[u8]> =
310            self.serialize_keypair(&device_data.signed_pre_key)?.into();
311        let account_data: Option<Arc<[u8]>> = device_data
312            .account
313            .as_ref()
314            .map(|a| Arc::from(wacore::store::device::account_serde::to_bytes(a)));
315        let registration_id = device_data.registration_id as i32;
316        let signed_pre_key_id = device_data.signed_pre_key_id as i32;
317        let signed_pre_key_signature: Arc<[u8]> =
318            Arc::from(&device_data.signed_pre_key_signature[..]);
319        let adv_secret_key: Arc<[u8]> = Arc::from(&device_data.adv_secret_key[..]);
320        let push_name: Arc<str> = Arc::from(device_data.push_name.as_str());
321        let app_version_primary = device_data.app_version_primary as i32;
322        let app_version_secondary = device_data.app_version_secondary as i32;
323        let app_version_tertiary = device_data.app_version_tertiary as i64;
324        let app_version_last_fetched_ms = device_data.app_version_last_fetched_ms;
325        let edge_routing_info: Option<Arc<[u8]>> =
326            device_data.edge_routing_info.as_deref().map(Arc::from);
327        let props_hash: Option<Arc<str>> = device_data.props_hash.as_deref().map(Arc::from);
328        let next_pre_key_id = device_data.next_pre_key_id as i32;
329        let server_has_prekeys = device_data.server_has_prekeys;
330        let nct_salt: Option<Arc<[u8]>> = device_data.nct_salt.as_deref().map(Arc::from);
331        let server_cert_chain: Option<Arc<[u8]>> = device_data
332            .server_cert_chain
333            .as_ref()
334            .map(|chain| {
335                bincode::serde::encode_to_vec(chain, bincode::config::standard())
336                    .map(Arc::from)
337                    .map_err(|e| StoreError::Serialization(Box::new(e)))
338            })
339            .transpose()?;
340        let new_lid: Arc<str> = Arc::from(
341            device_data
342                .lid
343                .as_ref()
344                .map(|j| j.to_string())
345                .unwrap_or_default()
346                .as_str(),
347        );
348        let new_pn: Arc<str> = Arc::from(
349            device_data
350                .pn
351                .as_ref()
352                .map(|j| j.to_string())
353                .unwrap_or_default()
354                .as_str(),
355        );
356
357        self.with_retry("save_device_data", || {
358            let noise_key_data = Arc::clone(&noise_key_data);
359            let identity_key_data = Arc::clone(&identity_key_data);
360            let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
361            let account_data = account_data.clone();
362            let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
363            let adv_secret_key = Arc::clone(&adv_secret_key);
364            let push_name = Arc::clone(&push_name);
365            let edge_routing_info = edge_routing_info.clone();
366            let props_hash = props_hash.clone();
367            let nct_salt = nct_salt.clone();
368            let server_cert_chain = server_cert_chain.clone();
369            let new_lid = Arc::clone(&new_lid);
370            let new_pn = Arc::clone(&new_pn);
371
372            Box::new(move |conn: &mut SqliteConnection| {
373                diesel::insert_into(device::table)
374                    .values((
375                        device::id.eq(device_id),
376                        device::lid.eq(&*new_lid),
377                        device::pn.eq(&*new_pn),
378                        device::registration_id.eq(registration_id),
379                        device::noise_key.eq(&*noise_key_data),
380                        device::identity_key.eq(&*identity_key_data),
381                        device::signed_pre_key.eq(&*signed_pre_key_data),
382                        device::signed_pre_key_id.eq(signed_pre_key_id),
383                        device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
384                        device::adv_secret_key.eq(&*adv_secret_key),
385                        device::account.eq(account_data.as_deref()),
386                        device::push_name.eq(&*push_name),
387                        device::app_version_primary.eq(app_version_primary),
388                        device::app_version_secondary.eq(app_version_secondary),
389                        device::app_version_tertiary.eq(app_version_tertiary),
390                        device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
391                        device::edge_routing_info.eq(edge_routing_info.as_deref()),
392                        device::props_hash.eq(props_hash.as_deref()),
393                        device::next_pre_key_id.eq(next_pre_key_id),
394                        device::server_has_prekeys.eq(server_has_prekeys),
395                        device::nct_salt.eq(nct_salt.as_deref()),
396                        device::server_cert_chain.eq(server_cert_chain.as_deref()),
397                    ))
398                    .on_conflict(device::id)
399                    .do_update()
400                    .set((
401                        device::lid.eq(excluded(device::lid)),
402                        device::pn.eq(excluded(device::pn)),
403                        device::registration_id.eq(excluded(device::registration_id)),
404                        device::noise_key.eq(excluded(device::noise_key)),
405                        device::identity_key.eq(excluded(device::identity_key)),
406                        device::signed_pre_key.eq(excluded(device::signed_pre_key)),
407                        device::signed_pre_key_id.eq(excluded(device::signed_pre_key_id)),
408                        device::signed_pre_key_signature
409                            .eq(excluded(device::signed_pre_key_signature)),
410                        device::adv_secret_key.eq(excluded(device::adv_secret_key)),
411                        device::account.eq(excluded(device::account)),
412                        device::push_name.eq(excluded(device::push_name)),
413                        device::app_version_primary.eq(excluded(device::app_version_primary)),
414                        device::app_version_secondary.eq(excluded(device::app_version_secondary)),
415                        device::app_version_tertiary.eq(excluded(device::app_version_tertiary)),
416                        device::app_version_last_fetched_ms
417                            .eq(excluded(device::app_version_last_fetched_ms)),
418                        device::edge_routing_info.eq(excluded(device::edge_routing_info)),
419                        device::props_hash.eq(excluded(device::props_hash)),
420                        device::next_pre_key_id.eq(excluded(device::next_pre_key_id)),
421                        device::server_has_prekeys.eq(excluded(device::server_has_prekeys)),
422                        device::nct_salt.eq(excluded(device::nct_salt)),
423                        device::server_cert_chain.eq(excluded(device::server_cert_chain)),
424                    ))
425                    .execute(conn)
426                    .map(|_| ())
427            })
428        })
429        .await
430    }
431
432    pub async fn create_new_device(&self) -> Result<i32> {
433        let device_id = self.device_id;
434        let new_device = wacore::store::Device::new();
435
436        let noise_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.noise_key)?.into();
437        let identity_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.identity_key)?.into();
438        let signed_pre_key_data: Arc<[u8]> =
439            self.serialize_keypair(&new_device.signed_pre_key)?.into();
440        let registration_id = new_device.registration_id as i32;
441        let signed_pre_key_id = new_device.signed_pre_key_id as i32;
442        let signed_pre_key_signature: Arc<[u8]> =
443            Arc::from(&new_device.signed_pre_key_signature[..]);
444        let adv_secret_key: Arc<[u8]> = Arc::from(&new_device.adv_secret_key[..]);
445        let push_name: Arc<str> = Arc::from(new_device.push_name.as_str());
446        let app_version_primary = new_device.app_version_primary as i32;
447        let app_version_secondary = new_device.app_version_secondary as i32;
448        let app_version_tertiary = new_device.app_version_tertiary as i64;
449        let app_version_last_fetched_ms = new_device.app_version_last_fetched_ms;
450        let next_pre_key_id = new_device.next_pre_key_id as i32;
451        let server_has_prekeys = new_device.server_has_prekeys;
452
453        self.with_retry("create_new_device", || {
454            let noise_key_data = Arc::clone(&noise_key_data);
455            let identity_key_data = Arc::clone(&identity_key_data);
456            let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
457            let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
458            let adv_secret_key = Arc::clone(&adv_secret_key);
459            let push_name = Arc::clone(&push_name);
460
461            Box::new(move |conn: &mut SqliteConnection| {
462                diesel::insert_into(device::table)
463                    .values((
464                        device::id.eq(device_id),
465                        device::lid.eq(""),
466                        device::pn.eq(""),
467                        device::registration_id.eq(registration_id),
468                        device::noise_key.eq(&*noise_key_data),
469                        device::identity_key.eq(&*identity_key_data),
470                        device::signed_pre_key.eq(&*signed_pre_key_data),
471                        device::signed_pre_key_id.eq(signed_pre_key_id),
472                        device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
473                        device::adv_secret_key.eq(&*adv_secret_key),
474                        device::account.eq(None::<&[u8]>),
475                        device::push_name.eq(&*push_name),
476                        device::app_version_primary.eq(app_version_primary),
477                        device::app_version_secondary.eq(app_version_secondary),
478                        device::app_version_tertiary.eq(app_version_tertiary),
479                        device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
480                        device::edge_routing_info.eq(None::<&[u8]>),
481                        device::props_hash.eq(None::<&str>),
482                        device::next_pre_key_id.eq(next_pre_key_id),
483                        device::server_has_prekeys.eq(server_has_prekeys),
484                        device::nct_salt.eq(None::<&[u8]>),
485                        device::server_cert_chain.eq(None::<&[u8]>),
486                    ))
487                    .execute(conn)
488                    .map(|_| device_id)
489            })
490        })
491        .await
492    }
493
494    pub async fn device_exists(&self, device_id: i32) -> Result<bool> {
495        use crate::schema::device;
496
497        let pool = self.pool.clone();
498        tokio::task::spawn_blocking(move || -> Result<bool> {
499            let mut conn = pool
500                .get()
501                .map_err(|e| StoreError::Connection(Box::new(e)))?;
502
503            let count: i64 = device::table
504                .filter(device::id.eq(device_id))
505                .count()
506                .get_result(&mut conn)
507                .map_err(|e| StoreError::Database(Box::new(e)))?;
508
509            Ok(count > 0)
510        })
511        .await
512        .map_err(|e| StoreError::Database(Box::new(e)))?
513    }
514
515    pub async fn load_device_data_for_device(&self, device_id: i32) -> Result<Option<CoreDevice>> {
516        use crate::schema::device;
517
518        let pool = self.pool.clone();
519        let row = tokio::task::spawn_blocking(move || -> Result<Option<DeviceRow>> {
520            let mut conn = pool
521                .get()
522                .map_err(|e| StoreError::Connection(Box::new(e)))?;
523            let result = device::table
524                .filter(device::id.eq(device_id))
525                .first::<DeviceRow>(&mut conn)
526                .optional()
527                .map_err(|e| StoreError::Database(Box::new(e)))?;
528            Ok(result)
529        })
530        .await
531        .map_err(|e| StoreError::Database(Box::new(e)))??;
532
533        if let Some(row) = row {
534            let pn = if !row.pn.is_empty() {
535                row.pn.parse().ok()
536            } else {
537                None
538            };
539            let lid = if !row.lid.is_empty() {
540                row.lid.parse().ok()
541            } else {
542                None
543            };
544
545            let noise_key = self.deserialize_keypair(&row.noise_key)?;
546            let identity_key = self.deserialize_keypair(&row.identity_key)?;
547            let signed_pre_key = self.deserialize_keypair(&row.signed_pre_key)?;
548
549            let signed_pre_key_signature: [u8; 64] =
550                row.signed_pre_key_signature.try_into().map_err(|_| {
551                    StoreError::Validation("Invalid signed_pre_key_signature length".to_string())
552                })?;
553
554            let adv_secret_key: [u8; 32] = row
555                .adv_secret_key
556                .try_into()
557                .map_err(|_| StoreError::Validation("Invalid adv_secret_key length".to_string()))?;
558
559            let account = row
560                .account
561                .map(|data| {
562                    wacore::store::device::account_serde::from_bytes(&data)
563                        .map_err(|e| StoreError::Serialization(Box::new(e)))
564                })
565                .transpose()?;
566
567            Ok(Some(CoreDevice {
568                pn,
569                lid,
570                registration_id: row.registration_id as u32,
571                noise_key,
572                identity_key,
573                signed_pre_key,
574                signed_pre_key_id: row.signed_pre_key_id as u32,
575                signed_pre_key_signature,
576                adv_secret_key,
577                account,
578                push_name: row.push_name,
579                app_version_primary: row.app_version_primary as u32,
580                app_version_secondary: row.app_version_secondary as u32,
581                app_version_tertiary: row.app_version_tertiary.try_into().unwrap_or(0u32),
582                app_version_last_fetched_ms: row.app_version_last_fetched_ms,
583                device_props: wacore::store::device::DEVICE_PROPS.clone(),
584                client_profile: wacore::client_profile::ClientProfile::web(),
585                edge_routing_info: row.edge_routing_info,
586                props_hash: row.props_hash,
587                next_pre_key_id: row.next_pre_key_id as u32,
588                server_has_prekeys: row.server_has_prekeys,
589                nct_salt: row.nct_salt,
590                nct_salt_sync_seen: false,
591                server_cert_chain: row
592                    .server_cert_chain
593                    .as_deref()
594                    .and_then(|bytes| {
595                        // The cert chain is a perf cache, not load-bearing
596                        // identity. A corrupt blob (truncated row, format
597                        // change between versions) must NOT block startup —
598                        // log it and degrade to None so the next connect
599                        // simply pays one XX handshake to repopulate.
600                        match bincode::serde::decode_from_slice(
601                            bytes,
602                            bincode::config::standard(),
603                        ) {
604                            Ok((chain, _)) => Some(chain),
605                            Err(e) => {
606                                log::warn!(
607                                    "device {} server_cert_chain blob ({} bytes) failed to decode: {e}; \
608                                     dropping cache, next connect will use XX",
609                                    self.device_id,
610                                    bytes.len(),
611                                );
612                                None
613                            }
614                        }
615                    }),
616            }))
617        } else {
618            Ok(None)
619        }
620    }
621
622    pub async fn put_identity_for_device(
623        &self,
624        address: &str,
625        key: [u8; 32],
626        device_id: i32,
627    ) -> Result<()> {
628        let pool = self.pool.clone();
629        let db_semaphore = self.db_semaphore.clone();
630        let address_owned = address.to_string();
631        let key_vec = key.to_vec();
632
633        const MAX_RETRIES: u32 = 5;
634
635        for attempt in 0..=MAX_RETRIES {
636            let permit = db_semaphore
637                .clone()
638                .acquire_owned()
639                .await
640                .map_err(|e| StoreError::Database(Box::new(e)))?;
641
642            let pool_clone = pool.clone();
643            let address_clone = address_owned.clone();
644            let key_clone = key_vec.clone();
645
646            let result =
647                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
648                    let mut conn = pool_clone
649                        .get()
650                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
651                    diesel::insert_into(identities::table)
652                        .values((
653                            identities::address.eq(address_clone),
654                            identities::key.eq(&key_clone[..]),
655                            identities::device_id.eq(device_id),
656                        ))
657                        .on_conflict((identities::address, identities::device_id))
658                        .do_update()
659                        .set(identities::key.eq(&key_clone[..]))
660                        .execute(&mut conn)
661                        .map_err(DieselOrStore::Diesel)?;
662                    Ok(())
663                })
664                .await;
665
666            drop(permit);
667
668            match result {
669                Ok(Ok(())) => return Ok(()),
670                Ok(Err(DieselOrStore::Diesel(ref e)))
671                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
672                {
673                    let delay_ms = 10 * 2u64.pow(attempt);
674                    warn!(
675                        "Identity write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
676                        attempt + 1,
677                        MAX_RETRIES + 1,
678                    );
679                    tokio::time::sleep(std::time::Duration::from_millis(delay_ms)).await;
680                    continue;
681                }
682                Ok(Err(e)) => return Err(e.into()),
683                Err(e) => return Err(StoreError::Database(Box::new(e))),
684            }
685        }
686
687        Err(StoreError::RetriesExhausted {
688            op: format!("identity_write (after {} attempts)", MAX_RETRIES + 1),
689        })
690    }
691
692    pub async fn delete_identity_for_device(&self, address: &str, device_id: i32) -> Result<()> {
693        let pool = self.pool.clone();
694        let address_owned = address.to_string();
695
696        tokio::task::spawn_blocking(move || -> Result<()> {
697            let mut conn = pool
698                .get()
699                .map_err(|e| StoreError::Connection(Box::new(e)))?;
700            diesel::delete(
701                identities::table
702                    .filter(identities::address.eq(address_owned))
703                    .filter(identities::device_id.eq(device_id)),
704            )
705            .execute(&mut conn)
706            .map_err(|e| StoreError::Database(Box::new(e)))?;
707            Ok(())
708        })
709        .await
710        .map_err(|e| StoreError::Database(Box::new(e)))??;
711
712        Ok(())
713    }
714
715    pub async fn load_identity_for_device(
716        &self,
717        address: &str,
718        device_id: i32,
719    ) -> Result<Option<Vec<u8>>> {
720        let pool = self.pool.clone();
721        let address = address.to_string();
722        let result = self
723            .with_semaphore(move || -> Result<Option<Vec<u8>>> {
724                let mut conn = pool
725                    .get()
726                    .map_err(|e| StoreError::Connection(Box::new(e)))?;
727                let res: Option<Vec<u8>> = identities::table
728                    .select(identities::key)
729                    .filter(identities::address.eq(address))
730                    .filter(identities::device_id.eq(device_id))
731                    .first(&mut conn)
732                    .optional()
733                    .map_err(|e| StoreError::Database(Box::new(e)))?;
734                Ok(res)
735            })
736            .await?;
737
738        Ok(result)
739    }
740
741    pub async fn get_session_for_device(
742        &self,
743        address: &str,
744        device_id: i32,
745    ) -> Result<Option<Vec<u8>>> {
746        let pool = self.pool.clone();
747        let address_for_query = address.to_string();
748        let result = self
749            .with_semaphore(move || -> Result<Option<Vec<u8>>> {
750                let mut conn = pool
751                    .get()
752                    .map_err(|e| StoreError::Connection(Box::new(e)))?;
753                let res: Option<Vec<u8>> = sessions::table
754                    .select(sessions::record)
755                    .filter(sessions::address.eq(address_for_query.clone()))
756                    .filter(sessions::device_id.eq(device_id))
757                    .first(&mut conn)
758                    .optional()
759                    .map_err(|e| StoreError::Database(Box::new(e)))?;
760
761                Ok(res)
762            })
763            .await?;
764
765        Ok(result)
766    }
767
768    pub async fn put_session_for_device(
769        &self,
770        address: &str,
771        session: &[u8],
772        device_id: i32,
773    ) -> Result<()> {
774        let pool = self.pool.clone();
775        let db_semaphore = self.db_semaphore.clone();
776        let address_owned = address.to_string();
777        let session_vec = session.to_vec();
778
779        const MAX_RETRIES: u32 = 5;
780
781        for attempt in 0..=MAX_RETRIES {
782            let permit = db_semaphore
783                .clone()
784                .acquire_owned()
785                .await
786                .map_err(|e| StoreError::Database(Box::new(e)))?;
787
788            let pool_clone = pool.clone();
789            let address_clone = address_owned.clone();
790            let session_clone = session_vec.clone();
791
792            let result =
793                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
794                    let mut conn = pool_clone
795                        .get()
796                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
797                    diesel::insert_into(sessions::table)
798                        .values((
799                            sessions::address.eq(address_clone),
800                            sessions::record.eq(&session_clone),
801                            sessions::device_id.eq(device_id),
802                        ))
803                        .on_conflict((sessions::address, sessions::device_id))
804                        .do_update()
805                        .set(sessions::record.eq(&session_clone))
806                        .execute(&mut conn)
807                        .map_err(DieselOrStore::Diesel)?;
808                    Ok(())
809                })
810                .await;
811
812            drop(permit);
813
814            match result {
815                Ok(Ok(())) => return Ok(()),
816                Ok(Err(DieselOrStore::Diesel(ref e)))
817                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
818                {
819                    let delay_ms = 10 * 2u64.pow(attempt);
820                    warn!(
821                        "Session write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
822                        attempt + 1,
823                        MAX_RETRIES + 1,
824                    );
825                    tokio::time::sleep(std::time::Duration::from_millis(delay_ms)).await;
826                    continue;
827                }
828                Ok(Err(e)) => return Err(e.into()),
829                Err(e) => return Err(StoreError::Database(Box::new(e))),
830            }
831        }
832
833        Err(StoreError::RetriesExhausted {
834            op: format!("session_write (after {} attempts)", MAX_RETRIES + 1),
835        })
836    }
837
838    pub async fn delete_session_for_device(&self, address: &str, device_id: i32) -> Result<()> {
839        let pool = self.pool.clone();
840        let address_owned = address.to_string();
841
842        tokio::task::spawn_blocking(move || -> Result<()> {
843            let mut conn = pool
844                .get()
845                .map_err(|e| StoreError::Connection(Box::new(e)))?;
846            diesel::delete(
847                sessions::table
848                    .filter(sessions::address.eq(address_owned))
849                    .filter(sessions::device_id.eq(device_id)),
850            )
851            .execute(&mut conn)
852            .map_err(|e| StoreError::Database(Box::new(e)))?;
853            Ok(())
854        })
855        .await
856        .map_err(|e| StoreError::Database(Box::new(e)))??;
857
858        Ok(())
859    }
860
861    pub async fn put_sender_key_for_device(
862        &self,
863        address: &str,
864        record: &[u8],
865        device_id: i32,
866    ) -> Result<()> {
867        let pool = self.pool.clone();
868        let address = address.to_string();
869        let record_vec = record.to_vec();
870        tokio::task::spawn_blocking(move || -> Result<()> {
871            let mut conn = pool
872                .get()
873                .map_err(|e| StoreError::Connection(Box::new(e)))?;
874            diesel::insert_into(sender_keys::table)
875                .values((
876                    sender_keys::address.eq(address),
877                    sender_keys::record.eq(&record_vec),
878                    sender_keys::device_id.eq(device_id),
879                ))
880                .on_conflict((sender_keys::address, sender_keys::device_id))
881                .do_update()
882                .set(sender_keys::record.eq(&record_vec))
883                .execute(&mut conn)
884                .map_err(|e| StoreError::Database(Box::new(e)))?;
885            Ok(())
886        })
887        .await
888        .map_err(|e| StoreError::Database(Box::new(e)))??;
889        Ok(())
890    }
891
892    pub async fn get_sender_key_for_device(
893        &self,
894        address: &str,
895        device_id: i32,
896    ) -> Result<Option<Vec<u8>>> {
897        let pool = self.pool.clone();
898        let address = address.to_string();
899        tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
900            let mut conn = pool
901                .get()
902                .map_err(|e| StoreError::Connection(Box::new(e)))?;
903            let res: Option<Vec<u8>> = sender_keys::table
904                .select(sender_keys::record)
905                .filter(sender_keys::address.eq(address))
906                .filter(sender_keys::device_id.eq(device_id))
907                .first(&mut conn)
908                .optional()
909                .map_err(|e| StoreError::Database(Box::new(e)))?;
910            Ok(res)
911        })
912        .await
913        .map_err(|e| StoreError::Database(Box::new(e)))?
914    }
915
916    pub async fn delete_sender_key_for_device(&self, address: &str, device_id: i32) -> Result<()> {
917        let pool = self.pool.clone();
918        let address = address.to_string();
919        tokio::task::spawn_blocking(move || -> Result<()> {
920            let mut conn = pool
921                .get()
922                .map_err(|e| StoreError::Connection(Box::new(e)))?;
923            diesel::delete(
924                sender_keys::table
925                    .filter(sender_keys::address.eq(address))
926                    .filter(sender_keys::device_id.eq(device_id)),
927            )
928            .execute(&mut conn)
929            .map_err(|e| StoreError::Database(Box::new(e)))?;
930            Ok(())
931        })
932        .await
933        .map_err(|e| StoreError::Database(Box::new(e)))??;
934        Ok(())
935    }
936
937    pub async fn get_app_state_sync_key_for_device(
938        &self,
939        key_id: &[u8],
940        device_id: i32,
941    ) -> Result<Option<AppStateSyncKey>> {
942        let pool = self.pool.clone();
943        let key_id = key_id.to_vec();
944        let res: Option<Vec<u8>> =
945            tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
946                let mut conn = pool
947                    .get()
948                    .map_err(|e| StoreError::Connection(Box::new(e)))?;
949                let res: Option<Vec<u8>> = app_state_keys::table
950                    .select(app_state_keys::key_data)
951                    .filter(app_state_keys::key_id.eq(&key_id))
952                    .filter(app_state_keys::device_id.eq(device_id))
953                    .first(&mut conn)
954                    .optional()
955                    .map_err(|e| StoreError::Database(Box::new(e)))?;
956                Ok(res)
957            })
958            .await
959            .map_err(|e| StoreError::Database(Box::new(e)))??;
960
961        if let Some(data) = res {
962            let (key, _) = bincode::serde::decode_from_slice(&data, bincode::config::standard())
963                .map_err(|e| StoreError::Serialization(Box::new(e)))?;
964            Ok(Some(key))
965        } else {
966            Ok(None)
967        }
968    }
969
970    pub async fn set_app_state_sync_key_for_device(
971        &self,
972        key_id: &[u8],
973        key: AppStateSyncKey,
974        device_id: i32,
975    ) -> Result<()> {
976        let pool = self.pool.clone();
977        let key_id = key_id.to_vec();
978        let data = bincode::serde::encode_to_vec(&key, bincode::config::standard())
979            .map_err(|e| StoreError::Serialization(Box::new(e)))?;
980        tokio::task::spawn_blocking(move || -> Result<()> {
981            let mut conn = pool
982                .get()
983                .map_err(|e| StoreError::Connection(Box::new(e)))?;
984            diesel::insert_into(app_state_keys::table)
985                .values((
986                    app_state_keys::key_id.eq(&key_id),
987                    app_state_keys::key_data.eq(&data),
988                    app_state_keys::device_id.eq(device_id),
989                ))
990                .on_conflict((app_state_keys::key_id, app_state_keys::device_id))
991                .do_update()
992                .set(app_state_keys::key_data.eq(&data))
993                .execute(&mut conn)
994                .map_err(|e| StoreError::Database(Box::new(e)))?;
995            Ok(())
996        })
997        .await
998        .map_err(|e| StoreError::Database(Box::new(e)))??;
999        Ok(())
1000    }
1001
1002    pub async fn get_latest_app_state_sync_key_id_for_device(
1003        &self,
1004        device_id: i32,
1005    ) -> Result<Option<Vec<u8>>> {
1006        let pool = self.pool.clone();
1007        let res: Option<Vec<u8>> =
1008            tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1009                let mut conn = pool
1010                    .get()
1011                    .map_err(|e| StoreError::Connection(Box::new(e)))?;
1012                let res: Option<Vec<u8>> = app_state_keys::table
1013                    .select(app_state_keys::key_id)
1014                    .filter(app_state_keys::device_id.eq(device_id))
1015                    .order(app_state_keys::key_id.desc())
1016                    .first(&mut conn)
1017                    .optional()
1018                    .map_err(|e| StoreError::Database(Box::new(e)))?;
1019                Ok(res)
1020            })
1021            .await
1022            .map_err(|e| StoreError::Database(Box::new(e)))??;
1023        Ok(res)
1024    }
1025
1026    pub async fn get_app_state_version_for_device(
1027        &self,
1028        name: &str,
1029        device_id: i32,
1030    ) -> Result<HashState> {
1031        let pool = self.pool.clone();
1032        let name = name.to_string();
1033        let res: Option<Vec<u8>> =
1034            tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1035                let mut conn = pool
1036                    .get()
1037                    .map_err(|e| StoreError::Connection(Box::new(e)))?;
1038                let res: Option<Vec<u8>> = app_state_versions::table
1039                    .select(app_state_versions::state_data)
1040                    .filter(app_state_versions::name.eq(name))
1041                    .filter(app_state_versions::device_id.eq(device_id))
1042                    .first(&mut conn)
1043                    .optional()
1044                    .map_err(|e| StoreError::Database(Box::new(e)))?;
1045                Ok(res)
1046            })
1047            .await
1048            .map_err(|e| StoreError::Database(Box::new(e)))??;
1049
1050        if let Some(data) = res {
1051            let (state, _) = bincode::serde::decode_from_slice(&data, bincode::config::standard())
1052                .map_err(|e| StoreError::Serialization(Box::new(e)))?;
1053            Ok(state)
1054        } else {
1055            Ok(HashState::default())
1056        }
1057    }
1058
1059    pub async fn set_app_state_version_for_device(
1060        &self,
1061        name: &str,
1062        state: HashState,
1063        device_id: i32,
1064    ) -> Result<()> {
1065        let name = name.to_string();
1066        let data = bincode::serde::encode_to_vec(&state, bincode::config::standard())
1067            .map_err(|e| StoreError::Serialization(Box::new(e)))?;
1068        self.with_retry("set_app_state_version", || {
1069            let name = name.clone();
1070            let data = data.clone();
1071            Box::new(move |conn: &mut SqliteConnection| {
1072                diesel::insert_into(app_state_versions::table)
1073                    .values((
1074                        app_state_versions::name.eq(&name),
1075                        app_state_versions::state_data.eq(&data),
1076                        app_state_versions::device_id.eq(device_id),
1077                    ))
1078                    .on_conflict((app_state_versions::name, app_state_versions::device_id))
1079                    .do_update()
1080                    .set(app_state_versions::state_data.eq(&data))
1081                    .execute(conn)?;
1082                Ok(())
1083            })
1084        })
1085        .await
1086    }
1087
1088    pub async fn put_app_state_mutation_macs_for_device(
1089        &self,
1090        name: &str,
1091        version: u64,
1092        mutations: &[AppStateMutationMAC],
1093        device_id: i32,
1094    ) -> Result<()> {
1095        if mutations.is_empty() {
1096            return Ok(());
1097        }
1098        let name = name.to_string();
1099        let mutations: Vec<AppStateMutationMAC> = mutations.to_vec();
1100        self.with_retry("put_app_state_mutation_macs", || {
1101            let name = name.clone();
1102            let mutations = mutations.clone();
1103            Box::new(move |conn: &mut SqliteConnection| {
1104                let records: Vec<_> = mutations
1105                    .iter()
1106                    .map(|m| {
1107                        (
1108                            app_state_mutation_macs::name.eq(&name),
1109                            app_state_mutation_macs::version.eq(version as i64),
1110                            app_state_mutation_macs::index_mac.eq(&m.index_mac),
1111                            app_state_mutation_macs::value_mac.eq(&m.value_mac),
1112                            app_state_mutation_macs::device_id.eq(device_id),
1113                        )
1114                    })
1115                    .collect();
1116
1117                // SQLite variable limit is typically 999 or 32766.
1118                // Each row has 5 columns. 100 rows * 5 = 500 params, which is safe.
1119                const CHUNK_SIZE: usize = 100;
1120
1121                for chunk in records.chunks(CHUNK_SIZE) {
1122                    diesel::insert_into(app_state_mutation_macs::table)
1123                        .values(chunk)
1124                        .on_conflict((
1125                            app_state_mutation_macs::name,
1126                            app_state_mutation_macs::index_mac,
1127                            app_state_mutation_macs::device_id,
1128                        ))
1129                        .do_update()
1130                        .set((
1131                            app_state_mutation_macs::version
1132                                .eq(excluded(app_state_mutation_macs::version)),
1133                            app_state_mutation_macs::value_mac
1134                                .eq(excluded(app_state_mutation_macs::value_mac)),
1135                        ))
1136                        .execute(conn)?;
1137                }
1138                Ok(())
1139            })
1140        })
1141        .await
1142    }
1143
1144    pub async fn delete_app_state_mutation_macs_for_device(
1145        &self,
1146        name: &str,
1147        index_macs: &[Vec<u8>],
1148        device_id: i32,
1149    ) -> Result<()> {
1150        if index_macs.is_empty() {
1151            return Ok(());
1152        }
1153        let name = name.to_string();
1154        let index_macs: Vec<Vec<u8>> = index_macs.to_vec();
1155        self.with_retry("delete_app_state_mutation_macs", || {
1156            let name = name.clone();
1157            let index_macs = index_macs.clone();
1158            Box::new(move |conn: &mut SqliteConnection| {
1159                // SQLite variable limit is usually 999 or higher.
1160                // We use a safe chunk size to stay well within limits.
1161                const CHUNK_SIZE: usize = 500;
1162
1163                for chunk in index_macs.chunks(CHUNK_SIZE) {
1164                    diesel::delete(
1165                        app_state_mutation_macs::table.filter(
1166                            app_state_mutation_macs::name
1167                                .eq(&name)
1168                                .and(app_state_mutation_macs::index_mac.eq_any(chunk))
1169                                .and(app_state_mutation_macs::device_id.eq(device_id)),
1170                        ),
1171                    )
1172                    .execute(conn)?;
1173                }
1174                Ok(())
1175            })
1176        })
1177        .await
1178    }
1179
1180    pub async fn get_app_state_mutation_mac_for_device(
1181        &self,
1182        name: &str,
1183        index_mac: &[u8],
1184        device_id: i32,
1185    ) -> Result<Option<Vec<u8>>> {
1186        let pool = self.pool.clone();
1187        let name = name.to_string();
1188        let index_mac = index_mac.to_vec();
1189        tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1190            let mut conn = pool
1191                .get()
1192                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1193            let res: Option<Vec<u8>> = app_state_mutation_macs::table
1194                .select(app_state_mutation_macs::value_mac)
1195                .filter(app_state_mutation_macs::name.eq(&name))
1196                .filter(app_state_mutation_macs::index_mac.eq(&index_mac))
1197                .filter(app_state_mutation_macs::device_id.eq(device_id))
1198                .first(&mut conn)
1199                .optional()
1200                .map_err(|e| StoreError::Database(Box::new(e)))?;
1201            Ok(res)
1202        })
1203        .await
1204        .map_err(|e| StoreError::Database(Box::new(e)))?
1205    }
1206}
1207
1208#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
1209#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
1210impl SignalStore for SqliteStore {
1211    async fn put_identity(&self, address: &str, key: [u8; 32]) -> Result<()> {
1212        self.put_identity_for_device(address, key, self.device_id)
1213            .await
1214    }
1215
1216    async fn load_identity(&self, address: &str) -> Result<Option<[u8; 32]>> {
1217        let blob = self
1218            .load_identity_for_device(address, self.device_id)
1219            .await?;
1220        match blob {
1221            None => Ok(None),
1222            Some(v) => Ok(Some(v.try_into().map_err(|v: Vec<u8>| {
1223                StoreError::Validation(format!(
1224                    "identity key for '{}' has invalid length {} (expected 32)",
1225                    address,
1226                    v.len()
1227                ))
1228            })?)),
1229        }
1230    }
1231
1232    async fn delete_identity(&self, address: &str) -> Result<()> {
1233        self.delete_identity_for_device(address, self.device_id)
1234            .await
1235    }
1236
1237    async fn get_session(&self, address: &str) -> Result<Option<bytes::Bytes>> {
1238        Ok(self
1239            .get_session_for_device(address, self.device_id)
1240            .await?
1241            .map(bytes::Bytes::from))
1242    }
1243
1244    async fn has_session(&self, address: &str) -> Result<bool> {
1245        let pool = self.pool.clone();
1246        let device_id = self.device_id;
1247        let address_owned = address.to_string();
1248        self.with_semaphore(move || -> Result<bool> {
1249            let mut conn = pool
1250                .get()
1251                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1252            let exists = diesel::select(diesel::dsl::exists(
1253                sessions::table
1254                    .filter(sessions::address.eq(&address_owned))
1255                    .filter(sessions::device_id.eq(device_id)),
1256            ))
1257            .get_result(&mut conn)
1258            .map_err(|e| StoreError::Database(Box::new(e)))?;
1259            Ok(exists)
1260        })
1261        .await
1262    }
1263
1264    async fn put_session(&self, address: &str, session: &[u8]) -> Result<()> {
1265        self.put_session_for_device(address, session, self.device_id)
1266            .await
1267    }
1268
1269    async fn delete_session(&self, address: &str) -> Result<()> {
1270        self.delete_session_for_device(address, self.device_id)
1271            .await
1272    }
1273
1274    async fn store_prekey(&self, id: u32, record: &[u8], uploaded: bool) -> Result<()> {
1275        let pool = self.pool.clone();
1276        let db_semaphore = self.db_semaphore.clone();
1277        let device_id = self.device_id;
1278        let record = record.to_vec();
1279
1280        const MAX_RETRIES: u32 = 5;
1281
1282        for attempt in 0..=MAX_RETRIES {
1283            let permit = db_semaphore
1284                .clone()
1285                .acquire_owned()
1286                .await
1287                .map_err(|e| StoreError::Database(Box::new(e)))?;
1288
1289            let pool_clone = pool.clone();
1290            let record_clone = record.clone();
1291
1292            let result =
1293                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1294                    let mut conn = pool_clone
1295                        .get()
1296                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1297                    diesel::insert_into(prekeys::table)
1298                        .values((
1299                            prekeys::id.eq(id as i32),
1300                            prekeys::key.eq(&record_clone),
1301                            prekeys::uploaded.eq(uploaded),
1302                            prekeys::device_id.eq(device_id),
1303                        ))
1304                        .on_conflict((prekeys::id, prekeys::device_id))
1305                        .do_update()
1306                        .set((
1307                            prekeys::key.eq(&record_clone),
1308                            prekeys::uploaded.eq(uploaded),
1309                        ))
1310                        .execute(&mut conn)
1311                        .map_err(DieselOrStore::Diesel)?;
1312                    Ok(())
1313                })
1314                .await;
1315
1316            drop(permit);
1317
1318            match result {
1319                Ok(Ok(())) => return Ok(()),
1320                Ok(Err(DieselOrStore::Diesel(ref e)))
1321                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1322                {
1323                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
1324                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
1325                }
1326                Ok(Err(e)) => return Err(e.into()),
1327                Err(e) => return Err(StoreError::Database(Box::new(e))),
1328            }
1329        }
1330
1331        Err(StoreError::RetriesExhausted {
1332            op: "store_prekey".to_string(),
1333        })
1334    }
1335
1336    async fn store_prekeys_batch(&self, keys: &[(u32, Bytes)], uploaded: bool) -> Result<()> {
1337        if keys.is_empty() {
1338            return Ok(());
1339        }
1340
1341        let pool = self.pool.clone();
1342        let db_semaphore = self.db_semaphore.clone();
1343        let device_id = self.device_id;
1344        let keys: Vec<(u32, Bytes)> = keys.to_vec();
1345
1346        const MAX_RETRIES: u32 = 5;
1347
1348        for attempt in 0..=MAX_RETRIES {
1349            let permit = db_semaphore
1350                .clone()
1351                .acquire_owned()
1352                .await
1353                .map_err(|e| StoreError::Database(Box::new(e)))?;
1354
1355            let pool_clone = pool.clone();
1356            let keys_clone = keys.clone();
1357
1358            let result =
1359                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1360                    let mut conn = pool_clone
1361                        .get()
1362                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1363
1364                    conn.transaction(|conn| {
1365                        for (id, record) in &keys_clone {
1366                            diesel::insert_into(prekeys::table)
1367                                .values((
1368                                    prekeys::id.eq(*id as i32),
1369                                    prekeys::key.eq(record.as_ref()),
1370                                    prekeys::uploaded.eq(uploaded),
1371                                    prekeys::device_id.eq(device_id),
1372                                ))
1373                                .on_conflict((prekeys::id, prekeys::device_id))
1374                                .do_update()
1375                                .set((
1376                                    prekeys::key.eq(record.as_ref()),
1377                                    prekeys::uploaded.eq(uploaded),
1378                                ))
1379                                .execute(conn)?;
1380                        }
1381                        Ok::<(), diesel::result::Error>(())
1382                    })
1383                    .map_err(DieselOrStore::Diesel)
1384                })
1385                .await;
1386
1387            drop(permit);
1388
1389            match result {
1390                Ok(Ok(())) => return Ok(()),
1391                Ok(Err(DieselOrStore::Diesel(ref e)))
1392                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1393                {
1394                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
1395                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
1396                }
1397                Ok(Err(e)) => return Err(e.into()),
1398                Err(e) => return Err(StoreError::Database(Box::new(e))),
1399            }
1400        }
1401
1402        Err(StoreError::RetriesExhausted {
1403            op: "store_prekeys_batch".to_string(),
1404        })
1405    }
1406
1407    async fn load_prekey(&self, id: u32) -> Result<Option<Bytes>> {
1408        let pool = self.pool.clone();
1409        let device_id = self.device_id;
1410        tokio::task::spawn_blocking(move || -> Result<Option<Bytes>> {
1411            let mut conn = pool
1412                .get()
1413                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1414            let res: Option<Vec<u8>> = prekeys::table
1415                .select(prekeys::key)
1416                .filter(prekeys::id.eq(id as i32))
1417                .filter(prekeys::device_id.eq(device_id))
1418                .first(&mut conn)
1419                .optional()
1420                .map_err(|e| StoreError::Database(Box::new(e)))?;
1421            Ok(res.map(Bytes::from))
1422        })
1423        .await
1424        .map_err(|e| StoreError::Database(Box::new(e)))?
1425    }
1426
1427    async fn load_prekeys_batch(&self, ids: &[u32]) -> Result<Vec<(u32, Bytes)>> {
1428        if ids.is_empty() {
1429            return Ok(Vec::new());
1430        }
1431        let pool = self.pool.clone();
1432        let device_id = self.device_id;
1433        let ids: Vec<i32> = ids.iter().map(|&id| id as i32).collect();
1434        self.with_semaphore(move || -> Result<Vec<(u32, Bytes)>> {
1435            let mut conn = pool
1436                .get()
1437                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1438            let rows: Vec<(i32, Vec<u8>)> = prekeys::table
1439                .select((prekeys::id, prekeys::key))
1440                .filter(prekeys::id.eq_any(&ids))
1441                .filter(prekeys::device_id.eq(device_id))
1442                .load(&mut conn)
1443                .map_err(|e| StoreError::Database(Box::new(e)))?;
1444            Ok(rows
1445                .into_iter()
1446                .map(|(id, key)| (id as u32, Bytes::from(key)))
1447                .collect())
1448        })
1449        .await
1450    }
1451
1452    async fn remove_prekey(&self, id: u32) -> Result<()> {
1453        let pool = self.pool.clone();
1454        let db_semaphore = self.db_semaphore.clone();
1455        let device_id = self.device_id;
1456
1457        const MAX_RETRIES: u32 = 5;
1458
1459        for attempt in 0..=MAX_RETRIES {
1460            let permit = db_semaphore
1461                .clone()
1462                .acquire_owned()
1463                .await
1464                .map_err(|e| StoreError::Database(Box::new(e)))?;
1465
1466            let pool_clone = pool.clone();
1467
1468            let result =
1469                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1470                    let mut conn = pool_clone
1471                        .get()
1472                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1473                    diesel::delete(
1474                        prekeys::table
1475                            .filter(prekeys::id.eq(id as i32))
1476                            .filter(prekeys::device_id.eq(device_id)),
1477                    )
1478                    .execute(&mut conn)
1479                    .map_err(DieselOrStore::Diesel)?;
1480                    Ok(())
1481                })
1482                .await;
1483
1484            drop(permit);
1485
1486            match result {
1487                Ok(Ok(())) => return Ok(()),
1488                Ok(Err(DieselOrStore::Diesel(ref e)))
1489                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1490                {
1491                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
1492                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
1493                }
1494                Ok(Err(e)) => return Err(e.into()),
1495                Err(e) => return Err(StoreError::Database(Box::new(e))),
1496            }
1497        }
1498
1499        Err(StoreError::RetriesExhausted {
1500            op: "remove_prekey".to_string(),
1501        })
1502    }
1503
1504    async fn get_max_prekey_id(&self) -> Result<u32> {
1505        let pool = self.pool.clone();
1506        let device_id = self.device_id;
1507        let db_semaphore = self.db_semaphore.clone();
1508        let _permit = db_semaphore
1509            .acquire()
1510            .await
1511            .map_err(|e| StoreError::Database(Box::new(e)))?;
1512
1513        tokio::task::spawn_blocking(move || -> Result<u32> {
1514            let mut conn = pool
1515                .get()
1516                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1517            use diesel::dsl::max;
1518            let result: Option<i32> = prekeys::table
1519                .filter(prekeys::device_id.eq(device_id))
1520                .select(max(prekeys::id))
1521                .first(&mut conn)
1522                .map_err(|e| StoreError::Database(Box::new(e)))?;
1523            Ok(result.unwrap_or(0) as u32)
1524        })
1525        .await
1526        .map_err(|e| StoreError::Database(Box::new(e)))?
1527    }
1528
1529    async fn store_signed_prekey(&self, id: u32, record: &[u8]) -> Result<()> {
1530        let pool = self.pool.clone();
1531        let db_semaphore = self.db_semaphore.clone();
1532        let device_id = self.device_id;
1533        let record = record.to_vec();
1534
1535        const MAX_RETRIES: u32 = 5;
1536
1537        for attempt in 0..=MAX_RETRIES {
1538            let permit = db_semaphore
1539                .clone()
1540                .acquire_owned()
1541                .await
1542                .map_err(|e| StoreError::Database(Box::new(e)))?;
1543
1544            let pool_clone = pool.clone();
1545            let record_clone = record.clone();
1546
1547            let result =
1548                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1549                    let mut conn = pool_clone
1550                        .get()
1551                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1552                    diesel::insert_into(signed_prekeys::table)
1553                        .values((
1554                            signed_prekeys::id.eq(id as i32),
1555                            signed_prekeys::record.eq(&record_clone),
1556                            signed_prekeys::device_id.eq(device_id),
1557                        ))
1558                        .on_conflict((signed_prekeys::id, signed_prekeys::device_id))
1559                        .do_update()
1560                        .set(signed_prekeys::record.eq(&record_clone))
1561                        .execute(&mut conn)
1562                        .map_err(DieselOrStore::Diesel)?;
1563                    Ok(())
1564                })
1565                .await;
1566
1567            drop(permit);
1568
1569            match result {
1570                Ok(Ok(())) => return Ok(()),
1571                Ok(Err(DieselOrStore::Diesel(ref e)))
1572                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1573                {
1574                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
1575                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
1576                }
1577                Ok(Err(e)) => return Err(e.into()),
1578                Err(e) => return Err(StoreError::Database(Box::new(e))),
1579            }
1580        }
1581
1582        Err(StoreError::RetriesExhausted {
1583            op: "store_signed_prekey".to_string(),
1584        })
1585    }
1586
1587    async fn load_signed_prekey(&self, id: u32) -> Result<Option<Vec<u8>>> {
1588        let pool = self.pool.clone();
1589        let device_id = self.device_id;
1590        tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1591            let mut conn = pool
1592                .get()
1593                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1594            let res: Option<Vec<u8>> = signed_prekeys::table
1595                .select(signed_prekeys::record)
1596                .filter(signed_prekeys::id.eq(id as i32))
1597                .filter(signed_prekeys::device_id.eq(device_id))
1598                .first(&mut conn)
1599                .optional()
1600                .map_err(|e| StoreError::Database(Box::new(e)))?;
1601            Ok(res)
1602        })
1603        .await
1604        .map_err(|e| StoreError::Database(Box::new(e)))?
1605    }
1606
1607    async fn load_all_signed_prekeys(&self) -> Result<Vec<(u32, Vec<u8>)>> {
1608        let pool = self.pool.clone();
1609        let device_id = self.device_id;
1610        tokio::task::spawn_blocking(move || -> Result<Vec<(u32, Vec<u8>)>> {
1611            let mut conn = pool
1612                .get()
1613                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1614            let results: Vec<(i32, Vec<u8>)> = signed_prekeys::table
1615                .select((signed_prekeys::id, signed_prekeys::record))
1616                .filter(signed_prekeys::device_id.eq(device_id))
1617                .load(&mut conn)
1618                .map_err(|e| StoreError::Database(Box::new(e)))?;
1619            Ok(results
1620                .into_iter()
1621                .map(|(id, record)| (id as u32, record))
1622                .collect())
1623        })
1624        .await
1625        .map_err(|e| StoreError::Database(Box::new(e)))?
1626    }
1627
1628    async fn remove_signed_prekey(&self, id: u32) -> Result<()> {
1629        let pool = self.pool.clone();
1630        let db_semaphore = self.db_semaphore.clone();
1631        let device_id = self.device_id;
1632
1633        const MAX_RETRIES: u32 = 5;
1634
1635        for attempt in 0..=MAX_RETRIES {
1636            let permit = db_semaphore
1637                .clone()
1638                .acquire_owned()
1639                .await
1640                .map_err(|e| StoreError::Database(Box::new(e)))?;
1641
1642            let pool_clone = pool.clone();
1643
1644            let result =
1645                tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1646                    let mut conn = pool_clone
1647                        .get()
1648                        .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1649                    diesel::delete(
1650                        signed_prekeys::table
1651                            .filter(signed_prekeys::id.eq(id as i32))
1652                            .filter(signed_prekeys::device_id.eq(device_id)),
1653                    )
1654                    .execute(&mut conn)
1655                    .map_err(DieselOrStore::Diesel)?;
1656                    Ok(())
1657                })
1658                .await;
1659
1660            drop(permit);
1661
1662            match result {
1663                Ok(Ok(())) => return Ok(()),
1664                Ok(Err(DieselOrStore::Diesel(ref e)))
1665                    if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1666                {
1667                    let delay_ms = 10u64 * (1u64 << attempt.min(4));
1668                    tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
1669                }
1670                Ok(Err(e)) => return Err(e.into()),
1671                Err(e) => return Err(StoreError::Database(Box::new(e))),
1672            }
1673        }
1674
1675        Err(StoreError::RetriesExhausted {
1676            op: "remove_signed_prekey".to_string(),
1677        })
1678    }
1679
1680    async fn put_sender_key(&self, address: &str, record: &[u8]) -> Result<()> {
1681        self.put_sender_key_for_device(address, record, self.device_id)
1682            .await
1683    }
1684
1685    async fn get_sender_key(&self, address: &str) -> Result<Option<Vec<u8>>> {
1686        self.get_sender_key_for_device(address, self.device_id)
1687            .await
1688    }
1689
1690    async fn delete_sender_key(&self, address: &str) -> Result<()> {
1691        self.delete_sender_key_for_device(address, self.device_id)
1692            .await
1693    }
1694}
1695
1696#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
1697#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
1698impl AppSyncStore for SqliteStore {
1699    async fn get_sync_key(&self, key_id: &[u8]) -> Result<Option<AppStateSyncKey>> {
1700        self.get_app_state_sync_key_for_device(key_id, self.device_id)
1701            .await
1702    }
1703
1704    async fn set_sync_key(&self, key_id: &[u8], key: AppStateSyncKey) -> Result<()> {
1705        self.set_app_state_sync_key_for_device(key_id, key, self.device_id)
1706            .await
1707    }
1708
1709    async fn get_version(&self, name: &str) -> Result<HashState> {
1710        self.get_app_state_version_for_device(name, self.device_id)
1711            .await
1712    }
1713
1714    async fn set_version(&self, name: &str, state: HashState) -> Result<()> {
1715        self.set_app_state_version_for_device(name, state, self.device_id)
1716            .await
1717    }
1718
1719    async fn put_mutation_macs(
1720        &self,
1721        name: &str,
1722        version: u64,
1723        mutations: &[AppStateMutationMAC],
1724    ) -> Result<()> {
1725        self.put_app_state_mutation_macs_for_device(name, version, mutations, self.device_id)
1726            .await
1727    }
1728
1729    async fn get_mutation_mac(&self, name: &str, index_mac: &[u8]) -> Result<Option<Vec<u8>>> {
1730        self.get_app_state_mutation_mac_for_device(name, index_mac, self.device_id)
1731            .await
1732    }
1733
1734    async fn delete_mutation_macs(&self, name: &str, index_macs: &[Vec<u8>]) -> Result<()> {
1735        self.delete_app_state_mutation_macs_for_device(name, index_macs, self.device_id)
1736            .await
1737    }
1738
1739    async fn get_latest_sync_key_id(&self) -> Result<Option<Vec<u8>>> {
1740        self.get_latest_app_state_sync_key_id_for_device(self.device_id)
1741            .await
1742    }
1743}
1744
1745#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
1746#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
1747impl ProtocolStore for SqliteStore {
1748    async fn get_sender_key_devices(&self, group_jid: &str) -> Result<Vec<(String, bool)>> {
1749        let pool = self.pool.clone();
1750        let device_id = self.device_id;
1751        let group_jid = group_jid.to_string();
1752        tokio::task::spawn_blocking(move || -> Result<Vec<(String, bool)>> {
1753            let mut conn = pool
1754                .get()
1755                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1756            let rows: Vec<(String, i32)> = sender_key_devices::table
1757                .select((sender_key_devices::device_jid, sender_key_devices::has_key))
1758                .filter(sender_key_devices::group_jid.eq(&group_jid))
1759                .filter(sender_key_devices::device_id.eq(device_id))
1760                .load(&mut conn)
1761                .map_err(|e| StoreError::Database(Box::new(e)))?;
1762            Ok(rows
1763                .into_iter()
1764                .map(|(jid, has_key)| (jid, has_key != 0))
1765                .collect())
1766        })
1767        .await
1768        .map_err(|e| StoreError::Database(Box::new(e)))?
1769    }
1770
1771    async fn set_sender_key_status(&self, group_jid: &str, entries: &[(&str, bool)]) -> Result<()> {
1772        if entries.is_empty() {
1773            return Ok(());
1774        }
1775        let device_id = self.device_id;
1776        let group_jid = group_jid.to_string();
1777        let owned_entries: Arc<Vec<(String, bool)>> = Arc::new(
1778            entries
1779                .iter()
1780                .map(|(jid, has_key)| (jid.to_string(), *has_key))
1781                .collect(),
1782        );
1783        let now = wacore::time::now_secs();
1784        self.with_retry("set_sender_key_status", || {
1785            let group_jid = group_jid.clone();
1786            let owned_entries = Arc::clone(&owned_entries);
1787            Box::new(move |conn: &mut SqliteConnection| {
1788                let values: Vec<_> = owned_entries
1789                    .iter()
1790                    .map(|(device_jid, has_key)| {
1791                        (
1792                            sender_key_devices::group_jid.eq(&group_jid),
1793                            sender_key_devices::device_jid.eq(device_jid),
1794                            sender_key_devices::has_key.eq(i32::from(*has_key)),
1795                            sender_key_devices::device_id.eq(device_id),
1796                            sender_key_devices::updated_at.eq(now),
1797                        )
1798                    })
1799                    .collect();
1800
1801                const CHUNK_SIZE: usize = 190;
1802
1803                for chunk in values.chunks(CHUNK_SIZE) {
1804                    diesel::insert_into(sender_key_devices::table)
1805                        .values(chunk)
1806                        .on_conflict((
1807                            sender_key_devices::group_jid,
1808                            sender_key_devices::device_jid,
1809                            sender_key_devices::device_id,
1810                        ))
1811                        .do_update()
1812                        .set((
1813                            sender_key_devices::has_key
1814                                .eq(diesel::upsert::excluded(sender_key_devices::has_key)),
1815                            sender_key_devices::updated_at.eq(now),
1816                        ))
1817                        .execute(conn)?;
1818                }
1819                Ok(())
1820            })
1821        })
1822        .await
1823    }
1824
1825    async fn clear_sender_key_devices(&self, group_jid: &str) -> Result<()> {
1826        let device_id = self.device_id;
1827        let group_jid = group_jid.to_string();
1828        self.with_retry("clear_sender_key_devices", || {
1829            let group_jid = group_jid.clone();
1830            Box::new(move |conn: &mut SqliteConnection| {
1831                diesel::delete(
1832                    sender_key_devices::table
1833                        .filter(sender_key_devices::group_jid.eq(&group_jid))
1834                        .filter(sender_key_devices::device_id.eq(device_id)),
1835                )
1836                .execute(conn)?;
1837                Ok(())
1838            })
1839        })
1840        .await
1841    }
1842
1843    async fn clear_all_sender_key_devices(&self) -> Result<()> {
1844        let device_id = self.device_id;
1845        self.with_retry("clear_all_sender_key_devices", || {
1846            Box::new(move |conn: &mut SqliteConnection| {
1847                diesel::delete(
1848                    sender_key_devices::table.filter(sender_key_devices::device_id.eq(device_id)),
1849                )
1850                .execute(conn)?;
1851                Ok(())
1852            })
1853        })
1854        .await
1855    }
1856
1857    async fn delete_sender_key_device_rows(&self, device_jids: &[&str]) -> Result<()> {
1858        if device_jids.is_empty() {
1859            return Ok(());
1860        }
1861        let device_id = self.device_id;
1862        let owned: Arc<Vec<String>> = Arc::new(device_jids.iter().map(|s| s.to_string()).collect());
1863        self.with_retry("delete_sender_key_device_rows", || {
1864            let owned = Arc::clone(&owned);
1865            Box::new(move |conn: &mut SqliteConnection| {
1866                const CHUNK: usize = 190;
1867                for chunk in owned.chunks(CHUNK) {
1868                    diesel::delete(
1869                        sender_key_devices::table
1870                            .filter(sender_key_devices::device_jid.eq_any(chunk))
1871                            .filter(sender_key_devices::device_id.eq(device_id)),
1872                    )
1873                    .execute(conn)?;
1874                }
1875                Ok(())
1876            })
1877        })
1878        .await
1879    }
1880
1881    async fn get_lid_mapping(&self, lid: &str) -> Result<Option<LidPnMappingEntry>> {
1882        let pool = self.pool.clone();
1883        let device_id = self.device_id;
1884        let lid = lid.to_string();
1885        tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
1886            let mut conn = pool
1887                .get()
1888                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1889            let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
1890                .select((
1891                    lid_pn_mapping::lid,
1892                    lid_pn_mapping::phone_number,
1893                    lid_pn_mapping::created_at,
1894                    lid_pn_mapping::learning_source,
1895                    lid_pn_mapping::updated_at,
1896                ))
1897                .filter(lid_pn_mapping::lid.eq(&lid))
1898                .filter(lid_pn_mapping::device_id.eq(device_id))
1899                .first(&mut conn)
1900                .optional()
1901                .map_err(|e| StoreError::Database(Box::new(e)))?;
1902            Ok(row.map(
1903                |(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
1904                    lid,
1905                    phone_number,
1906                    created_at,
1907                    updated_at,
1908                    learning_source,
1909                },
1910            ))
1911        })
1912        .await
1913        .map_err(|e| StoreError::Database(Box::new(e)))?
1914    }
1915
1916    async fn get_pn_mapping(&self, phone: &str) -> Result<Option<LidPnMappingEntry>> {
1917        let pool = self.pool.clone();
1918        let device_id = self.device_id;
1919        let phone = phone.to_string();
1920        tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
1921            let mut conn = pool
1922                .get()
1923                .map_err(|e| StoreError::Connection(Box::new(e)))?;
1924            let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
1925                .select((
1926                    lid_pn_mapping::lid,
1927                    lid_pn_mapping::phone_number,
1928                    lid_pn_mapping::created_at,
1929                    lid_pn_mapping::learning_source,
1930                    lid_pn_mapping::updated_at,
1931                ))
1932                .filter(lid_pn_mapping::phone_number.eq(&phone))
1933                .filter(lid_pn_mapping::device_id.eq(device_id))
1934                .order(lid_pn_mapping::updated_at.desc())
1935                .first(&mut conn)
1936                .optional()
1937                .map_err(|e| StoreError::Database(Box::new(e)))?;
1938            Ok(row.map(
1939                |(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
1940                    lid,
1941                    phone_number,
1942                    created_at,
1943                    updated_at,
1944                    learning_source,
1945                },
1946            ))
1947        })
1948        .await
1949        .map_err(|e| StoreError::Database(Box::new(e)))?
1950    }
1951
1952    async fn put_lid_mapping(&self, entry: &LidPnMappingEntry) -> Result<()> {
1953        self.put_lid_mappings(std::slice::from_ref(entry)).await
1954    }
1955
1956    async fn put_lid_mappings(&self, entries: &[LidPnMappingEntry]) -> Result<()> {
1957        if entries.is_empty() {
1958            return Ok(());
1959        }
1960        let device_id = self.device_id;
1961        // Share the batch across retry attempts via Arc so no retry re-clones
1962        // the Vec. `with_retry` invokes `make_op` once per attempt; we only
1963        // bump the Arc refcount.
1964        let entries: std::sync::Arc<Vec<LidPnMappingEntry>> = std::sync::Arc::new(entries.to_vec());
1965        self.with_retry("put_lid_mappings", move || {
1966            let entries = std::sync::Arc::clone(&entries);
1967            Box::new(move |conn: &mut SqliteConnection| {
1968                conn.transaction::<_, DieselError, _>(|conn| {
1969                    for entry in entries.iter() {
1970                        diesel::insert_into(lid_pn_mapping::table)
1971                            .values((
1972                                lid_pn_mapping::lid.eq(&entry.lid),
1973                                lid_pn_mapping::phone_number.eq(&entry.phone_number),
1974                                lid_pn_mapping::created_at.eq(entry.created_at),
1975                                lid_pn_mapping::learning_source.eq(&entry.learning_source),
1976                                lid_pn_mapping::updated_at.eq(entry.updated_at),
1977                                lid_pn_mapping::device_id.eq(device_id),
1978                            ))
1979                            .on_conflict((lid_pn_mapping::lid, lid_pn_mapping::device_id))
1980                            .do_update()
1981                            .set((
1982                                lid_pn_mapping::phone_number.eq(&entry.phone_number),
1983                                lid_pn_mapping::learning_source.eq(&entry.learning_source),
1984                                lid_pn_mapping::updated_at.eq(entry.updated_at),
1985                            ))
1986                            .execute(conn)?;
1987                    }
1988                    Ok(())
1989                })
1990            })
1991        })
1992        .await
1993    }
1994
1995    async fn get_all_lid_mappings(&self) -> Result<Vec<LidPnMappingEntry>> {
1996        let pool = self.pool.clone();
1997        let device_id = self.device_id;
1998        tokio::task::spawn_blocking(move || -> Result<Vec<LidPnMappingEntry>> {
1999            let mut conn = pool
2000                .get()
2001                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2002            let rows: Vec<(String, String, i64, String, i64)> = lid_pn_mapping::table
2003                .select((
2004                    lid_pn_mapping::lid,
2005                    lid_pn_mapping::phone_number,
2006                    lid_pn_mapping::created_at,
2007                    lid_pn_mapping::learning_source,
2008                    lid_pn_mapping::updated_at,
2009                ))
2010                .filter(lid_pn_mapping::device_id.eq(device_id))
2011                .load(&mut conn)
2012                .map_err(|e| StoreError::Database(Box::new(e)))?;
2013            Ok(rows
2014                .into_iter()
2015                .map(
2016                    |(lid, phone_number, created_at, learning_source, updated_at)| {
2017                        LidPnMappingEntry {
2018                            lid,
2019                            phone_number,
2020                            created_at,
2021                            updated_at,
2022                            learning_source,
2023                        }
2024                    },
2025                )
2026                .collect())
2027        })
2028        .await
2029        .map_err(|e| StoreError::Database(Box::new(e)))?
2030    }
2031
2032    async fn save_base_key(&self, address: &str, message_id: &str, base_key: &[u8]) -> Result<()> {
2033        let pool = self.pool.clone();
2034        let device_id = self.device_id;
2035        let address = address.to_string();
2036        let message_id = message_id.to_string();
2037        let base_key = base_key.to_vec();
2038        let now = wacore::time::now_secs() as i32;
2039        tokio::task::spawn_blocking(move || -> Result<()> {
2040            let mut conn = pool
2041                .get()
2042                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2043            diesel::insert_into(base_keys::table)
2044                .values((
2045                    base_keys::address.eq(&address),
2046                    base_keys::message_id.eq(&message_id),
2047                    base_keys::base_key.eq(&base_key),
2048                    base_keys::device_id.eq(device_id),
2049                    base_keys::created_at.eq(now),
2050                ))
2051                .on_conflict((
2052                    base_keys::address,
2053                    base_keys::message_id,
2054                    base_keys::device_id,
2055                ))
2056                .do_update()
2057                .set(base_keys::base_key.eq(&base_key))
2058                .execute(&mut conn)
2059                .map_err(|e| StoreError::Database(Box::new(e)))?;
2060            Ok(())
2061        })
2062        .await
2063        .map_err(|e| StoreError::Database(Box::new(e)))??;
2064        Ok(())
2065    }
2066
2067    async fn has_same_base_key(
2068        &self,
2069        address: &str,
2070        message_id: &str,
2071        current_base_key: &[u8],
2072    ) -> Result<bool> {
2073        let pool = self.pool.clone();
2074        let device_id = self.device_id;
2075        let address = address.to_string();
2076        let message_id = message_id.to_string();
2077        let current_base_key = current_base_key.to_vec();
2078        tokio::task::spawn_blocking(move || -> Result<bool> {
2079            let mut conn = pool
2080                .get()
2081                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2082            let stored_key: Option<Vec<u8>> = base_keys::table
2083                .select(base_keys::base_key)
2084                .filter(base_keys::address.eq(&address))
2085                .filter(base_keys::message_id.eq(&message_id))
2086                .filter(base_keys::device_id.eq(device_id))
2087                .first(&mut conn)
2088                .optional()
2089                .map_err(|e| StoreError::Database(Box::new(e)))?;
2090            Ok(stored_key.as_ref() == Some(&current_base_key))
2091        })
2092        .await
2093        .map_err(|e| StoreError::Database(Box::new(e)))?
2094    }
2095
2096    async fn delete_base_key(&self, address: &str, message_id: &str) -> Result<()> {
2097        let pool = self.pool.clone();
2098        let device_id = self.device_id;
2099        let address = address.to_string();
2100        let message_id = message_id.to_string();
2101        tokio::task::spawn_blocking(move || -> Result<()> {
2102            let mut conn = pool
2103                .get()
2104                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2105            diesel::delete(
2106                base_keys::table
2107                    .filter(base_keys::address.eq(&address))
2108                    .filter(base_keys::message_id.eq(&message_id))
2109                    .filter(base_keys::device_id.eq(device_id)),
2110            )
2111            .execute(&mut conn)
2112            .map_err(|e| StoreError::Database(Box::new(e)))?;
2113            Ok(())
2114        })
2115        .await
2116        .map_err(|e| StoreError::Database(Box::new(e)))??;
2117        Ok(())
2118    }
2119
2120    async fn update_device_list(&self, record: DeviceListRecord) -> Result<()> {
2121        let pool = self.pool.clone();
2122        let device_id = self.device_id;
2123        let devices_json = serde_json::to_string(&record.devices)
2124            .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2125        let now = wacore::time::now_secs() as i32;
2126        tokio::task::spawn_blocking(move || -> Result<()> {
2127            let mut conn = pool
2128                .get()
2129                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2130            let raw_id_i32 = record.raw_id.map(|r| r as i32);
2131            diesel::insert_into(device_registry::table)
2132                .values((
2133                    device_registry::user_id.eq(&record.user),
2134                    device_registry::devices_json.eq(&devices_json),
2135                    device_registry::timestamp.eq(record.timestamp as i32),
2136                    device_registry::phash.eq(&record.phash),
2137                    device_registry::device_id.eq(device_id),
2138                    device_registry::updated_at.eq(now),
2139                    device_registry::raw_id.eq(raw_id_i32),
2140                ))
2141                .on_conflict((device_registry::user_id, device_registry::device_id))
2142                .do_update()
2143                .set((
2144                    device_registry::devices_json.eq(&devices_json),
2145                    device_registry::timestamp.eq(record.timestamp as i32),
2146                    device_registry::phash.eq(&record.phash),
2147                    device_registry::updated_at.eq(now),
2148                    device_registry::raw_id.eq(raw_id_i32),
2149                ))
2150                .execute(&mut conn)
2151                .map_err(|e| StoreError::Database(Box::new(e)))?;
2152            Ok(())
2153        })
2154        .await
2155        .map_err(|e| StoreError::Database(Box::new(e)))??;
2156        Ok(())
2157    }
2158
2159    async fn update_device_lists(&self, records: Vec<DeviceListRecord>) -> Result<()> {
2160        if records.is_empty() {
2161            return Ok(());
2162        }
2163        let device_id = self.device_id;
2164        let now = wacore::time::now_secs() as i32;
2165
2166        // Pre-serialize devices_json once (outside the retry loop and outside
2167        // spawn_blocking) so retries are zero-allocation. Each row carries its
2168        // own json+raw_id alongside the record.
2169        struct PreparedRow {
2170            user: String,
2171            devices_json: String,
2172            timestamp: i32,
2173            phash: Option<String>,
2174            raw_id: Option<i32>,
2175        }
2176
2177        let prepared: Vec<PreparedRow> = records
2178            .into_iter()
2179            .map(|r| {
2180                let devices_json = serde_json::to_string(&r.devices)
2181                    .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2182                Ok(PreparedRow {
2183                    user: r.user,
2184                    devices_json,
2185                    timestamp: r.timestamp as i32,
2186                    phash: r.phash,
2187                    raw_id: r.raw_id.map(|v| v as i32),
2188                })
2189            })
2190            .collect::<Result<Vec<_>>>()?;
2191        let prepared = std::sync::Arc::new(prepared);
2192
2193        self.with_retry("update_device_lists", move || {
2194            let prepared = std::sync::Arc::clone(&prepared);
2195            Box::new(move |conn: &mut SqliteConnection| {
2196                conn.transaction::<_, DieselError, _>(|conn| {
2197                    for row in prepared.iter() {
2198                        diesel::insert_into(device_registry::table)
2199                            .values((
2200                                device_registry::user_id.eq(&row.user),
2201                                device_registry::devices_json.eq(&row.devices_json),
2202                                device_registry::timestamp.eq(row.timestamp),
2203                                device_registry::phash.eq(&row.phash),
2204                                device_registry::device_id.eq(device_id),
2205                                device_registry::updated_at.eq(now),
2206                                device_registry::raw_id.eq(row.raw_id),
2207                            ))
2208                            .on_conflict((device_registry::user_id, device_registry::device_id))
2209                            .do_update()
2210                            .set((
2211                                device_registry::devices_json.eq(&row.devices_json),
2212                                device_registry::timestamp.eq(row.timestamp),
2213                                device_registry::phash.eq(&row.phash),
2214                                device_registry::updated_at.eq(now),
2215                                device_registry::raw_id.eq(row.raw_id),
2216                            ))
2217                            .execute(conn)?;
2218                    }
2219                    Ok(())
2220                })
2221            })
2222        })
2223        .await
2224    }
2225
2226    async fn get_devices(&self, user: &str) -> Result<Option<DeviceListRecord>> {
2227        let pool = self.pool.clone();
2228        let device_id = self.device_id;
2229        let user = user.to_string();
2230        tokio::task::spawn_blocking(move || -> Result<Option<DeviceListRecord>> {
2231            let mut conn = pool
2232                .get()
2233                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2234            let row: Option<(String, String, i32, Option<String>, Option<i32>)> =
2235                device_registry::table
2236                    .select((
2237                        device_registry::user_id,
2238                        device_registry::devices_json,
2239                        device_registry::timestamp,
2240                        device_registry::phash,
2241                        device_registry::raw_id,
2242                    ))
2243                    .filter(device_registry::user_id.eq(&user))
2244                    .filter(device_registry::device_id.eq(device_id))
2245                    .first(&mut conn)
2246                    .optional()
2247                    .map_err(|e| StoreError::Database(Box::new(e)))?;
2248            match row {
2249                Some((user, devices_json, timestamp, phash, raw_id)) => {
2250                    let devices: Vec<DeviceInfo> = serde_json::from_str(&devices_json)
2251                        .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2252                    Ok(Some(DeviceListRecord {
2253                        user,
2254                        devices,
2255                        timestamp: timestamp as i64,
2256                        phash,
2257                        raw_id: raw_id.map(|r| r as u32),
2258                    }))
2259                }
2260                None => Ok(None),
2261            }
2262        })
2263        .await
2264        .map_err(|e| StoreError::Database(Box::new(e)))?
2265    }
2266
2267    async fn delete_devices(&self, user: &str) -> Result<()> {
2268        let pool = self.pool.clone();
2269        let device_id = self.device_id;
2270        let user = user.to_string();
2271        tokio::task::spawn_blocking(move || -> Result<()> {
2272            let mut conn = pool
2273                .get()
2274                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2275            diesel::delete(
2276                device_registry::table
2277                    .filter(device_registry::user_id.eq(&user))
2278                    .filter(device_registry::device_id.eq(device_id)),
2279            )
2280            .execute(&mut conn)
2281            .map_err(|e| StoreError::Database(Box::new(e)))?;
2282            Ok(())
2283        })
2284        .await
2285        .map_err(|e| StoreError::Database(Box::new(e)))??;
2286        Ok(())
2287    }
2288
2289    async fn get_tc_token(&self, jid: &str) -> Result<Option<TcTokenEntry>> {
2290        let pool = self.pool.clone();
2291        let device_id = self.device_id;
2292        let jid = jid.to_string();
2293        tokio::task::spawn_blocking(move || -> Result<Option<TcTokenEntry>> {
2294            let mut conn = pool
2295                .get()
2296                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2297            let row: Option<(Vec<u8>, i64, Option<i64>)> = tc_tokens::table
2298                .select((
2299                    tc_tokens::token,
2300                    tc_tokens::token_timestamp,
2301                    tc_tokens::sender_timestamp,
2302                ))
2303                .filter(tc_tokens::jid.eq(&jid))
2304                .filter(tc_tokens::device_id.eq(device_id))
2305                .first(&mut conn)
2306                .optional()
2307                .map_err(|e| StoreError::Database(Box::new(e)))?;
2308            Ok(
2309                row.map(|(token, token_timestamp, sender_timestamp)| TcTokenEntry {
2310                    token,
2311                    token_timestamp,
2312                    sender_timestamp,
2313                }),
2314            )
2315        })
2316        .await
2317        .map_err(|e| StoreError::Database(Box::new(e)))?
2318    }
2319
2320    async fn put_tc_token(&self, jid: &str, entry: &TcTokenEntry) -> Result<()> {
2321        let pool = self.pool.clone();
2322        let device_id = self.device_id;
2323        let jid = jid.to_string();
2324        let entry = entry.clone();
2325        let now = wacore::time::now_secs();
2326        tokio::task::spawn_blocking(move || -> Result<()> {
2327            let mut conn = pool
2328                .get()
2329                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2330            diesel::insert_into(tc_tokens::table)
2331                .values((
2332                    tc_tokens::jid.eq(&jid),
2333                    tc_tokens::token.eq(&entry.token),
2334                    tc_tokens::token_timestamp.eq(entry.token_timestamp),
2335                    tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
2336                    tc_tokens::device_id.eq(device_id),
2337                    tc_tokens::updated_at.eq(now),
2338                ))
2339                .on_conflict((tc_tokens::jid, tc_tokens::device_id))
2340                .do_update()
2341                .set((
2342                    tc_tokens::token.eq(&entry.token),
2343                    tc_tokens::token_timestamp.eq(entry.token_timestamp),
2344                    tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
2345                    tc_tokens::updated_at.eq(now),
2346                ))
2347                .execute(&mut conn)
2348                .map_err(|e| StoreError::Database(Box::new(e)))?;
2349            Ok(())
2350        })
2351        .await
2352        .map_err(|e| StoreError::Database(Box::new(e)))??;
2353        Ok(())
2354    }
2355
2356    async fn delete_tc_token(&self, jid: &str) -> Result<()> {
2357        let pool = self.pool.clone();
2358        let device_id = self.device_id;
2359        let jid = jid.to_string();
2360        tokio::task::spawn_blocking(move || -> Result<()> {
2361            let mut conn = pool
2362                .get()
2363                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2364            diesel::delete(
2365                tc_tokens::table
2366                    .filter(tc_tokens::jid.eq(&jid))
2367                    .filter(tc_tokens::device_id.eq(device_id)),
2368            )
2369            .execute(&mut conn)
2370            .map_err(|e| StoreError::Database(Box::new(e)))?;
2371            Ok(())
2372        })
2373        .await
2374        .map_err(|e| StoreError::Database(Box::new(e)))??;
2375        Ok(())
2376    }
2377
2378    async fn get_all_tc_token_jids(&self) -> Result<Vec<String>> {
2379        let pool = self.pool.clone();
2380        let device_id = self.device_id;
2381        tokio::task::spawn_blocking(move || -> Result<Vec<String>> {
2382            let mut conn = pool
2383                .get()
2384                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2385            let jids: Vec<String> = tc_tokens::table
2386                .select(tc_tokens::jid)
2387                .filter(tc_tokens::device_id.eq(device_id))
2388                .load(&mut conn)
2389                .map_err(|e| StoreError::Database(Box::new(e)))?;
2390            Ok(jids)
2391        })
2392        .await
2393        .map_err(|e| StoreError::Database(Box::new(e)))?
2394    }
2395
2396    async fn delete_expired_tc_tokens(&self, cutoff_timestamp: i64) -> Result<u32> {
2397        let pool = self.pool.clone();
2398        let device_id = self.device_id;
2399        tokio::task::spawn_blocking(move || -> Result<u32> {
2400            let mut conn = pool
2401                .get()
2402                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2403            let deleted = diesel::delete(
2404                tc_tokens::table
2405                    .filter(tc_tokens::token_timestamp.lt(cutoff_timestamp))
2406                    .filter(tc_tokens::device_id.eq(device_id)),
2407            )
2408            .execute(&mut conn)
2409            .map_err(|e| StoreError::Database(Box::new(e)))?;
2410            Ok(deleted as u32)
2411        })
2412        .await
2413        .map_err(|e| StoreError::Database(Box::new(e)))?
2414    }
2415
2416    async fn store_sent_message(
2417        &self,
2418        chat_jid: &str,
2419        message_id: &str,
2420        payload: &[u8],
2421    ) -> Result<()> {
2422        let chat_jid = chat_jid.to_string();
2423        let message_id = message_id.to_string();
2424        // Arc avoids cloning the full payload bytes on each retry iteration
2425        let payload: Arc<Vec<u8>> = Arc::new(payload.to_vec());
2426        let device_id = self.device_id;
2427        self.with_retry("store_sent_message", || {
2428            let chat_jid = chat_jid.clone();
2429            let message_id = message_id.clone();
2430            let payload = Arc::clone(&payload);
2431            Box::new(move |conn: &mut SqliteConnection| {
2432                diesel::replace_into(sent_messages::table)
2433                    .values((
2434                        sent_messages::chat_jid.eq(&chat_jid),
2435                        sent_messages::message_id.eq(&message_id),
2436                        sent_messages::payload.eq(payload.as_slice()),
2437                        sent_messages::device_id.eq(device_id),
2438                    ))
2439                    .execute(conn)?;
2440                Ok(())
2441            })
2442        })
2443        .await
2444    }
2445
2446    async fn take_sent_message(&self, chat_jid: &str, message_id: &str) -> Result<Option<Vec<u8>>> {
2447        let chat_jid = chat_jid.to_string();
2448        let message_id = message_id.to_string();
2449        let device_id = self.device_id;
2450        // Atomic SELECT+DELETE with retry for SQLITE_BUSY resilience.
2451        self.with_retry("take_sent_message", || {
2452            let chat_jid = chat_jid.clone();
2453            let message_id = message_id.clone();
2454            Box::new(move |conn: &mut SqliteConnection| {
2455                conn.immediate_transaction(|conn| {
2456                    let row: Option<Vec<u8>> = sent_messages::table
2457                        .select(sent_messages::payload)
2458                        .filter(sent_messages::chat_jid.eq(&chat_jid))
2459                        .filter(sent_messages::message_id.eq(&message_id))
2460                        .filter(sent_messages::device_id.eq(device_id))
2461                        .first(conn)
2462                        .optional()?;
2463                    if row.is_some() {
2464                        diesel::delete(
2465                            sent_messages::table
2466                                .filter(sent_messages::chat_jid.eq(&chat_jid))
2467                                .filter(sent_messages::message_id.eq(&message_id))
2468                                .filter(sent_messages::device_id.eq(device_id)),
2469                        )
2470                        .execute(conn)?;
2471                    }
2472                    Ok(row)
2473                })
2474            })
2475        })
2476        .await
2477    }
2478
2479    async fn delete_expired_sent_messages(&self, cutoff_timestamp: i64) -> Result<u32> {
2480        let pool = self.pool.clone();
2481        let device_id = self.device_id;
2482        tokio::task::spawn_blocking(move || -> Result<u32> {
2483            let mut conn = pool
2484                .get()
2485                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2486            let deleted = diesel::delete(
2487                sent_messages::table
2488                    .filter(sent_messages::created_at.lt(cutoff_timestamp))
2489                    .filter(sent_messages::device_id.eq(device_id)),
2490            )
2491            .execute(&mut conn)
2492            .map_err(|e| StoreError::Database(Box::new(e)))?;
2493            Ok(deleted as u32)
2494        })
2495        .await
2496        .map_err(|e| StoreError::Database(Box::new(e)))?
2497    }
2498}
2499
2500#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
2501#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
2502impl DeviceStore for SqliteStore {
2503    async fn save(&self, device: &CoreDevice) -> Result<()> {
2504        SqliteStore::save_device_data_for_device(self, self.device_id, device).await
2505    }
2506
2507    async fn load(&self) -> Result<Option<CoreDevice>> {
2508        SqliteStore::load_device_data_for_device(self, self.device_id).await
2509    }
2510
2511    async fn exists(&self) -> Result<bool> {
2512        SqliteStore::device_exists(self, self.device_id).await
2513    }
2514
2515    async fn create(&self) -> Result<i32> {
2516        SqliteStore::create_new_device(self).await
2517    }
2518
2519    async fn snapshot_db(&self, name: &str, extra_content: Option<&[u8]>) -> Result<()> {
2520        fn sanitize_snapshot_name(name: &str) -> Result<String> {
2521            const MAX_LENGTH: usize = 100;
2522
2523            let sanitized: String = name
2524                .chars()
2525                .map(|c| {
2526                    if c.is_ascii_alphanumeric() || c == '_' || c == '-' || c == '.' {
2527                        c
2528                    } else {
2529                        '_'
2530                    }
2531                })
2532                .collect();
2533
2534            let sanitized = sanitized
2535                .split('.')
2536                .filter(|part| !part.is_empty() && *part != "..")
2537                .collect::<Vec<_>>()
2538                .join(".");
2539
2540            let sanitized = sanitized.trim_matches(['/', '\\', '.']);
2541
2542            if sanitized.is_empty() {
2543                return Err(StoreError::InvalidConfig(
2544                    "Snapshot name cannot be empty after sanitization".to_string(),
2545                ));
2546            }
2547
2548            if sanitized.len() > MAX_LENGTH {
2549                return Err(StoreError::InvalidConfig(format!(
2550                    "Snapshot name exceeds maximum length of {} characters",
2551                    MAX_LENGTH
2552                )));
2553            }
2554
2555            Ok(sanitized.to_string())
2556        }
2557
2558        let sanitized_name = sanitize_snapshot_name(name)?;
2559
2560        let pool = self.pool.clone();
2561        let db_path = self.database_path.clone();
2562        let extra_data = extra_content.map(|b| b.to_vec());
2563
2564        tokio::task::spawn_blocking(move || -> Result<()> {
2565            let mut conn = pool
2566                .get()
2567                .map_err(|e| StoreError::Connection(Box::new(e)))?;
2568
2569            let timestamp = wacore::time::now_secs();
2570
2571            // Construct target path: db_path.snapshot-TIMESTAMP-SANITIZED_NAME
2572            let target_path = format!("{}.snapshot-{}-{}", db_path, timestamp, sanitized_name);
2573
2574            // Use VACUUM INTO to create a consistent backup
2575            // Note: We escape single quotes in the path just in case
2576            let query = format!("VACUUM INTO '{}'", target_path.replace("'", "''"));
2577
2578            diesel::sql_query(query)
2579                .execute(&mut conn)
2580                .map_err(|e| StoreError::Database(Box::new(e)))?;
2581
2582            // Save extra content if provided
2583            if let Some(data) = extra_data {
2584                let extra_path = format!("{}.json", target_path);
2585                std::fs::write(&extra_path, data)?;
2586            }
2587
2588            Ok(())
2589        })
2590        .await
2591        .map_err(|e| StoreError::Database(Box::new(e)))??;
2592
2593        Ok(())
2594    }
2595}
2596
2597#[cfg(test)]
2598mod tests {
2599    use super::*;
2600
2601    async fn create_test_store() -> SqliteStore {
2602        use portable_atomic::AtomicU64;
2603        use std::sync::atomic::Ordering;
2604        static COUNTER: AtomicU64 = AtomicU64::new(0);
2605        let id = COUNTER.fetch_add(1, Ordering::Relaxed);
2606        let db_name = format!(
2607            "file:memdb_test_{}_{}?mode=memory&cache=shared",
2608            std::process::id(),
2609            id
2610        );
2611        SqliteStore::new(&db_name)
2612            .await
2613            .expect("Failed to create test store")
2614    }
2615
2616    #[test]
2617    fn test_parse_database_path_regular_path() {
2618        let path = "/var/lib/whatsapp/database.db";
2619        let result = parse_database_path(path).unwrap();
2620        assert_eq!(result, "/var/lib/whatsapp/database.db");
2621    }
2622
2623    #[test]
2624    fn test_parse_database_path_with_sqlite_prefix() {
2625        let path = "sqlite:///var/lib/whatsapp/database.db";
2626        let result = parse_database_path(path).unwrap();
2627        assert_eq!(result, "/var/lib/whatsapp/database.db");
2628    }
2629
2630    #[test]
2631    fn test_parse_database_path_with_query_params() {
2632        let path = "file:database.db?mode=memory&cache=shared";
2633        let result = parse_database_path(path).unwrap();
2634        assert_eq!(result, "file:database.db");
2635    }
2636
2637    #[test]
2638    fn test_parse_database_path_with_fragment() {
2639        let path = "file:database.db#fragment";
2640        let result = parse_database_path(path).unwrap();
2641        assert_eq!(result, "file:database.db");
2642    }
2643
2644    #[test]
2645    fn test_parse_database_path_with_both_query_and_fragment() {
2646        let path = "sqlite:///var/lib/database.db?mode=ro#backup";
2647        let result = parse_database_path(path).unwrap();
2648        assert_eq!(result, "/var/lib/database.db");
2649    }
2650
2651    #[test]
2652    fn test_parse_database_path_in_memory_rejected() {
2653        let result = parse_database_path(":memory:");
2654        assert!(result.is_err());
2655        assert!(result.unwrap_err().to_string().contains("not supported"));
2656    }
2657
2658    #[test]
2659    fn test_parse_database_path_in_memory_with_query_rejected() {
2660        let result = parse_database_path(":memory:?cache=shared");
2661        assert!(result.is_err());
2662        assert!(result.unwrap_err().to_string().contains("not supported"));
2663    }
2664
2665    #[tokio::test]
2666    async fn test_device_registry_save_and_get() {
2667        let store = create_test_store().await;
2668
2669        let record = DeviceListRecord {
2670            user: "1234567890".to_string(),
2671            devices: vec![
2672                DeviceInfo {
2673                    device_id: 0,
2674                    key_index: None,
2675                },
2676                DeviceInfo {
2677                    device_id: 1,
2678                    key_index: Some(42),
2679                },
2680            ],
2681            timestamp: 1234567890,
2682            phash: Some("2:abcdef".to_string()),
2683            raw_id: None,
2684        };
2685
2686        store.update_device_list(record).await.expect("save failed");
2687        let loaded = store
2688            .get_devices("1234567890")
2689            .await
2690            .expect("get failed")
2691            .expect("record should exist");
2692
2693        assert_eq!(loaded.user, "1234567890");
2694        assert_eq!(loaded.devices.len(), 2);
2695        assert_eq!(loaded.devices[0].device_id, 0);
2696        assert_eq!(loaded.devices[1].device_id, 1);
2697        assert_eq!(loaded.devices[1].key_index, Some(42));
2698        assert_eq!(loaded.phash, Some("2:abcdef".to_string()));
2699    }
2700
2701    #[tokio::test]
2702    async fn test_device_registry_update_existing() {
2703        let store = create_test_store().await;
2704
2705        let record1 = DeviceListRecord {
2706            user: "1234567890".to_string(),
2707            devices: vec![DeviceInfo {
2708                device_id: 0,
2709                key_index: None,
2710            }],
2711            timestamp: 1000,
2712            phash: Some("2:old".to_string()),
2713            raw_id: None,
2714        };
2715        store
2716            .update_device_list(record1)
2717            .await
2718            .expect("save1 failed");
2719
2720        let record2 = DeviceListRecord {
2721            user: "1234567890".to_string(),
2722            devices: vec![
2723                DeviceInfo {
2724                    device_id: 0,
2725                    key_index: None,
2726                },
2727                DeviceInfo {
2728                    device_id: 2,
2729                    key_index: None,
2730                },
2731            ],
2732            timestamp: 2000,
2733            phash: Some("2:new".to_string()),
2734            raw_id: None,
2735        };
2736        store
2737            .update_device_list(record2)
2738            .await
2739            .expect("save2 failed");
2740
2741        let loaded = store
2742            .get_devices("1234567890")
2743            .await
2744            .expect("get failed")
2745            .expect("record should exist");
2746
2747        assert_eq!(loaded.devices.len(), 2);
2748        assert_eq!(loaded.phash, Some("2:new".to_string()));
2749    }
2750
2751    #[tokio::test]
2752    async fn test_device_registry_get_nonexistent() {
2753        let store = create_test_store().await;
2754        let result = store.get_devices("nonexistent").await.expect("get failed");
2755        assert!(result.is_none());
2756    }
2757
2758    #[tokio::test]
2759    async fn test_sender_key_devices_set_and_get() {
2760        let store = create_test_store().await;
2761
2762        let group = "group123@g.us";
2763
2764        // Set two devices: one has key, one needs SKDM
2765        store
2766            .set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", false)])
2767            .await
2768            .expect("set failed");
2769
2770        let devices = store
2771            .get_sender_key_devices(group)
2772            .await
2773            .expect("get failed");
2774        assert_eq!(devices.len(), 2);
2775        assert!(devices.contains(&("user1:5@lid".to_string(), true)));
2776        assert!(devices.contains(&("user2:3@lid".to_string(), false)));
2777    }
2778
2779    #[tokio::test]
2780    async fn test_sender_key_devices_upsert_overwrites() {
2781        let store = create_test_store().await;
2782
2783        let group = "group123@g.us";
2784
2785        // Initially mark as needing SKDM
2786        store
2787            .set_sender_key_status(group, &[("user1:5@lid", false)])
2788            .await
2789            .expect("set failed");
2790
2791        // Then mark as having key (simulates successful SKDM delivery)
2792        store
2793            .set_sender_key_status(group, &[("user1:5@lid", true)])
2794            .await
2795            .expect("set failed");
2796
2797        let devices = store
2798            .get_sender_key_devices(group)
2799            .await
2800            .expect("get failed");
2801        assert_eq!(devices.len(), 1);
2802        assert_eq!(devices[0], ("user1:5@lid".to_string(), true));
2803    }
2804
2805    #[tokio::test]
2806    async fn test_sender_key_devices_clear() {
2807        let store = create_test_store().await;
2808
2809        let group = "group123@g.us";
2810
2811        store
2812            .set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", true)])
2813            .await
2814            .expect("set failed");
2815
2816        store
2817            .clear_sender_key_devices(group)
2818            .await
2819            .expect("clear failed");
2820
2821        let devices = store
2822            .get_sender_key_devices(group)
2823            .await
2824            .expect("get failed");
2825        assert!(devices.is_empty());
2826    }
2827
2828    #[tokio::test]
2829    async fn test_tc_token_put_and_get() {
2830        let store = create_test_store().await;
2831
2832        let entry = TcTokenEntry {
2833            token: vec![1, 2, 3, 4, 5],
2834            token_timestamp: 1707000000,
2835            sender_timestamp: Some(1707000100),
2836        };
2837
2838        store
2839            .put_tc_token("user@lid", &entry)
2840            .await
2841            .expect("put failed");
2842
2843        let loaded = store
2844            .get_tc_token("user@lid")
2845            .await
2846            .expect("get failed")
2847            .expect("should exist");
2848
2849        assert_eq!(loaded.token, vec![1, 2, 3, 4, 5]);
2850        assert_eq!(loaded.token_timestamp, 1707000000);
2851        assert_eq!(loaded.sender_timestamp, Some(1707000100));
2852    }
2853
2854    #[tokio::test]
2855    async fn test_tc_token_upsert() {
2856        let store = create_test_store().await;
2857
2858        let entry1 = TcTokenEntry {
2859            token: vec![1, 2, 3],
2860            token_timestamp: 1000,
2861            sender_timestamp: None,
2862        };
2863        store.put_tc_token("user@lid", &entry1).await.unwrap();
2864
2865        let entry2 = TcTokenEntry {
2866            token: vec![4, 5, 6],
2867            token_timestamp: 2000,
2868            sender_timestamp: Some(1500),
2869        };
2870        store.put_tc_token("user@lid", &entry2).await.unwrap();
2871
2872        let loaded = store.get_tc_token("user@lid").await.unwrap().unwrap();
2873        assert_eq!(loaded.token, vec![4, 5, 6]);
2874        assert_eq!(loaded.token_timestamp, 2000);
2875        assert_eq!(loaded.sender_timestamp, Some(1500));
2876    }
2877
2878    #[tokio::test]
2879    async fn test_tc_token_delete() {
2880        let store = create_test_store().await;
2881
2882        let entry = TcTokenEntry {
2883            token: vec![1, 2, 3],
2884            token_timestamp: 1000,
2885            sender_timestamp: None,
2886        };
2887        store.put_tc_token("user@lid", &entry).await.unwrap();
2888        store.delete_tc_token("user@lid").await.unwrap();
2889
2890        let result = store.get_tc_token("user@lid").await.unwrap();
2891        assert!(result.is_none());
2892    }
2893
2894    #[tokio::test]
2895    async fn test_tc_token_get_all_jids() {
2896        let store = create_test_store().await;
2897
2898        let entry = TcTokenEntry {
2899            token: vec![1],
2900            token_timestamp: 1000,
2901            sender_timestamp: None,
2902        };
2903        store.put_tc_token("user1@lid", &entry).await.unwrap();
2904        store.put_tc_token("user2@lid", &entry).await.unwrap();
2905        store.put_tc_token("user3@lid", &entry).await.unwrap();
2906
2907        let mut jids = store.get_all_tc_token_jids().await.unwrap();
2908        jids.sort();
2909        assert_eq!(jids, vec!["user1@lid", "user2@lid", "user3@lid"]);
2910    }
2911
2912    #[tokio::test]
2913    async fn test_tc_token_delete_expired() {
2914        let store = create_test_store().await;
2915
2916        let old = TcTokenEntry {
2917            token: vec![1],
2918            token_timestamp: 1000,
2919            sender_timestamp: None,
2920        };
2921        let recent = TcTokenEntry {
2922            token: vec![2],
2923            token_timestamp: 5000,
2924            sender_timestamp: None,
2925        };
2926        store.put_tc_token("old@lid", &old).await.unwrap();
2927        store.put_tc_token("recent@lid", &recent).await.unwrap();
2928
2929        let deleted = store.delete_expired_tc_tokens(3000).await.unwrap();
2930        assert_eq!(deleted, 1);
2931
2932        assert!(store.get_tc_token("old@lid").await.unwrap().is_none());
2933        assert!(store.get_tc_token("recent@lid").await.unwrap().is_some());
2934    }
2935
2936    #[tokio::test]
2937    async fn test_tc_token_get_nonexistent() {
2938        let store = create_test_store().await;
2939        let result = store.get_tc_token("nonexistent@lid").await.unwrap();
2940        assert!(result.is_none());
2941    }
2942
2943    #[tokio::test]
2944    async fn test_sender_key_devices_different_groups() {
2945        let store = create_test_store().await;
2946
2947        let group1 = "group1@g.us";
2948        let group2 = "group2@g.us";
2949
2950        store
2951            .set_sender_key_status(group1, &[("user:5@lid", true)])
2952            .await
2953            .expect("set failed");
2954
2955        let g1 = store.get_sender_key_devices(group1).await.unwrap();
2956        assert_eq!(g1.len(), 1);
2957
2958        let g2 = store.get_sender_key_devices(group2).await.unwrap();
2959        assert!(g2.is_empty());
2960    }
2961
2962    #[tokio::test]
2963    async fn test_create_new_device_uses_configured_device_id() {
2964        use portable_atomic::AtomicU64;
2965        use std::sync::atomic::Ordering;
2966        static COUNTER: AtomicU64 = AtomicU64::new(100);
2967        let id = COUNTER.fetch_add(1, Ordering::Relaxed);
2968        let db_name = format!(
2969            "file:memdb_devid_{}_{}?mode=memory&cache=shared",
2970            std::process::id(),
2971            id
2972        );
2973
2974        let device_id = 42;
2975        let store = SqliteStore::new_for_device(&db_name, device_id)
2976            .await
2977            .expect("Failed to create test store");
2978
2979        assert!(!store.device_exists(device_id).await.unwrap());
2980        let returned_id = store.create_new_device().await.unwrap();
2981        assert_eq!(returned_id, device_id);
2982        assert!(store.device_exists(device_id).await.unwrap());
2983
2984        // Row 1 should NOT exist (would if auto-increment was used)
2985        if device_id != 1 {
2986            assert!(!store.device_exists(1).await.unwrap());
2987        }
2988
2989        let loaded = store.load_device_data_for_device(device_id).await.unwrap();
2990        assert!(
2991            loaded.is_some(),
2992            "device data should be loadable by configured id"
2993        );
2994    }
2995
2996    /// Round-trips a `CachedServerCertChain` through the SQLite schema:
2997    /// save → close store → reopen on the same db_name → load. Exercises
2998    /// the `2026-04-26-000000_add_server_cert_chain` migration plus the
2999    /// bincode encode/decode path in `save_device_data_for_device` /
3000    /// `load_device_data_for_device` (the part that the in-memory backend
3001    /// integration tests don't reach).
3002    #[tokio::test]
3003    async fn test_server_cert_chain_survives_save_load_roundtrip() {
3004        use portable_atomic::AtomicU64;
3005        use std::sync::atomic::Ordering;
3006        use wacore::store::device::{CachedNoiseCert, CachedServerCertChain};
3007
3008        static COUNTER: AtomicU64 = AtomicU64::new(200);
3009        let id = COUNTER.fetch_add(1, Ordering::Relaxed);
3010        // shared-cache so a second SqliteStore opened on the same name
3011        // sees the same on-disk state — the closest we can get to a real
3012        // process restart inside a single test run.
3013        let db_name = format!(
3014            "file:memdb_certchain_{}_{}?mode=memory&cache=shared",
3015            std::process::id(),
3016            id
3017        );
3018
3019        let device_id = 7;
3020        let chain = CachedServerCertChain {
3021            intermediate: CachedNoiseCert {
3022                key: [0xAB; 32],
3023                not_before: 1_700_000_000,
3024                not_after: 1_900_000_000,
3025            },
3026            leaf: CachedNoiseCert {
3027                key: [0xCD; 32],
3028                not_before: 1_700_000_500,
3029                not_after: 1_899_999_500,
3030            },
3031        };
3032
3033        // First store: create + populate. Keep it alive until after the
3034        // second store opens — `cache=shared` only persists the in-memory
3035        // database while at least one connection is open. Dropping the
3036        // first store would also drop the schema before the second can
3037        // see it.
3038        let _writer = SqliteStore::new_for_device(&db_name, device_id)
3039            .await
3040            .expect("create store");
3041        _writer.create_new_device().await.expect("create device");
3042
3043        let mut device = _writer
3044            .load_device_data_for_device(device_id)
3045            .await
3046            .expect("load")
3047            .expect("device should exist after create");
3048        device.server_cert_chain = Some(chain.clone());
3049        _writer
3050            .save_device_data_for_device(device_id, &device)
3051            .await
3052            .expect("save with cert chain");
3053
3054        // Second store on the SAME shared-cache db: this exercises the
3055        // exact path a fresh-process load would take — schema migration
3056        // already applied, BLOB column present, and the bincode-encoded
3057        // chain decoded by the load path.
3058        let store = SqliteStore::new_for_device(&db_name, device_id)
3059            .await
3060            .expect("reopen store");
3061        let loaded = store
3062            .load_device_data_for_device(device_id)
3063            .await
3064            .expect("load")
3065            .expect("device should exist after reopen");
3066        assert_eq!(
3067            loaded.server_cert_chain.as_ref(),
3068            Some(&chain),
3069            "server_cert_chain must survive a save/load roundtrip"
3070        );
3071
3072        // Sanity: clearing the chain and saving leaves the column as NULL,
3073        // not as an empty serialized struct.
3074        let mut device = loaded;
3075        device.server_cert_chain = None;
3076        store
3077            .save_device_data_for_device(device_id, &device)
3078            .await
3079            .expect("save with cleared cert chain");
3080
3081        let reloaded = store
3082            .load_device_data_for_device(device_id)
3083            .await
3084            .expect("reload")
3085            .expect("device should exist");
3086        assert!(
3087            reloaded.server_cert_chain.is_none(),
3088            "cleared chain must round-trip as None"
3089        );
3090    }
3091}