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
19enum 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
36fn 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#[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 if database_url == ":memory:" {
128 return Err(StoreError::InvalidConfig(
129 "Snapshot not supported for in-memory databases".to_string(),
130 ));
131 }
132
133 let path = database_url
135 .split(['?', '#'])
136 .next()
137 .unwrap_or(database_url);
138
139 let path = path.trim_start_matches("sqlite://");
141
142 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 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 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 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 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 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 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(¤t_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 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 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 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 let target_path = format!("{}.snapshot-{}-{}", db_path, timestamp, sanitized_name);
2573
2574 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 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 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 store
2787 .set_sender_key_status(group, &[("user1:5@lid", false)])
2788 .await
2789 .expect("set failed");
2790
2791 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 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 #[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 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 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 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 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}