twonly-app/rust/src/context.rs

516 lines
19 KiB
Rust

/*
* Copyright (c) 2026, Tobias Müller git@tsmr.eu
*
*/
use crate::api::runtime::{ApiClient, ApiRuntime};
use crate::bridge::InitConfig;
use crate::database::app::{AppDatabase, APP_DATABASE_FILE};
use crate::database::signal::Database;
use crate::error::Result;
use crate::error::TwonlyError;
use crate::keys::DatabaseKey;
use crate::keys::KeyManager;
#[cfg(not(test))]
use crate::log::init_tracing;
use crate::signal::engine::RustSignalEngine;
use crate::user_discovery::UserDiscovery;
use crate::utils::Shared;
use crate::{bridge::TwonlyFlutter, secure_storage::SecureStorage};
use libsignal_protocol::IdentityKey;
use libsignal_protocol::IdentityKeyPair;
use std::{path::PathBuf, sync::Arc};
use tokio::sync::{Mutex, OnceCell, RwLock};
#[cfg(not(test))]
use zeroize::Zeroize;
#[cfg(not(test))]
static GLOBAL_CONTEXT: OnceCell<Arc<Context>> = OnceCell::const_new();
pub struct TwonlyStandalone {
#[allow(dead_code)]
pub(crate) config: InitConfig,
#[allow(dead_code)]
pub(crate) rust_db: Arc<RwLock<Arc<Database>>>,
pub(crate) app_db: Arc<RwLock<Arc<AppDatabase>>>,
#[allow(dead_code)]
pub(crate) secure_storage: SecureStorage,
pub(crate) key_manager: Arc<Mutex<KeyManager>>,
pub(crate) user_discovery: Shared<UserDiscovery>,
pub(crate) signal_engine: Arc<Mutex<Option<RustSignalEngine>>>,
pub(crate) api_client: OnceCell<RwLock<Arc<ApiClient>>>,
}
#[allow(private_interfaces)] // for the test
pub enum Context {
Flutter(TwonlyFlutter),
Standalone(TwonlyStandalone),
}
impl Context {
pub(crate) fn get_api_client(&self) -> &OnceCell<RwLock<Arc<ApiClient>>> {
match self {
Self::Flutter(value) => &value.api_client,
Self::Standalone(value) => &value.api_client,
}
}
pub(crate) fn data_dir(&self) -> &str {
match self {
Self::Flutter(value) => &value.config.data_dir,
Self::Standalone(value) => &value.config.data_dir,
}
}
pub(crate) fn get_user_discovery(&self) -> &Shared<UserDiscovery> {
match self {
Self::Flutter(value) => &value.user_discovery,
Self::Standalone(value) => &value.user_discovery,
}
}
pub(crate) fn get_signal_engine(&self) -> &Arc<Mutex<Option<RustSignalEngine>>> {
match self {
Self::Flutter(value) => &value.signal_engine,
Self::Standalone(value) => &value.signal_engine,
}
}
pub fn from_standalone(standalone: TwonlyStandalone) -> Self {
Self::Standalone(standalone)
}
pub(crate) async fn init_flutter(config: InitConfig) -> Result<()> {
Self::init_common(config, true).await
}
#[allow(dead_code)]
pub(crate) async fn init_standalone(config: InitConfig) -> Result<()> {
Self::init_common(config, false).await
}
pub async fn init_for_testing(
database_dir: PathBuf,
data_dir: PathBuf,
) -> Result<Arc<Context>> {
std::fs::create_dir_all(&database_dir)?;
std::fs::create_dir_all(&data_dir)?;
let config = InitConfig {
database_dir: database_dir.display().to_string(),
data_dir: data_dir.display().to_string(),
};
// Initialize tracing and secure storage if not already done
let _ = SecureStorage::init();
let secure_storage = SecureStorage::new("eu.twonly.testing");
let key_manager = KeyManager::generate()?;
key_manager.store_to_keychain(&secure_storage)?;
let rust_db_path = database_dir.join("rust_db.sqlite");
let rust_db = Database::new(
&rust_db_path.display().to_string(),
Some(&key_manager.main_key.get_database_key(DatabaseKey::RustDb)),
false,
)
.await?;
rust_db.run_migrations().await?;
let rust_db = Arc::new(rust_db);
let app_db_path = database_dir.join(APP_DATABASE_FILE);
let app_db = AppDatabase::new(
&app_db_path.display().to_string(),
Some(&key_manager.main_key.get_database_key(DatabaseKey::AppDb)),
false,
)
.await?;
app_db.run_migrations().await?;
let app_db = Arc::new(RwLock::new(Arc::new(app_db)));
let rust_db = Arc::new(RwLock::new(rust_db));
let key_manager = Arc::new(Mutex::new(key_manager));
let user_discovery = Shared::new(UserDiscovery::new(
data_dir.to_str().unwrap(),
key_manager.clone(),
rust_db.clone(),
)?);
let ctx = Arc::new(Context::from_standalone(TwonlyStandalone {
config,
rust_db,
app_db,
secure_storage,
key_manager,
user_discovery,
signal_engine: Arc::new(Mutex::new(None)),
api_client: OnceCell::const_new(),
}));
ApiRuntime::initialize(&ctx).await?;
ApiRuntime::connect(&ctx).await?;
Ok(ctx)
}
#[doc(hidden)]
#[cfg(any(test, debug_assertions))]
pub async fn inject_test_signal_identity(
&self,
identity_key_pair_structure: Vec<u8>,
registration_id: i64,
pre_key_store: std::collections::HashMap<i64, Vec<u8>>,
) -> Result<()> {
let mut key_manager = self.get_key_manager().await?;
key_manager.signal_identity = Some(crate::keys::SignalIdentityKey {
identity_key_pair_structure,
registration_id,
pre_key_store,
});
Ok(())
}
#[doc(hidden)]
#[cfg(any(test, debug_assertions))]
pub async fn inject_test_user_id(&self, user_id: i64) -> Result<()> {
let mut key_manager = self.get_key_manager().await?;
key_manager.user_id = Some(user_id);
key_manager.store_to_keychain(self.get_secure_storage())?;
let signal_identity = key_manager.signal_identity.as_ref().map(|identity| {
(
identity.identity_key_pair_structure.clone(),
identity.registration_id,
)
});
drop(key_manager);
if let Some((identity_key_pair_structure, registration_id)) = signal_identity {
let database = self.get_rust_db().await;
*self.get_signal_engine().lock().await = Some(RustSignalEngine::new_with_pool(
database.pool.clone(),
identity_key_pair_structure,
registration_id as u32,
user_id.to_string(),
)?);
}
self.initialize_user_discovery_from_config().await?;
Ok(())
}
pub(crate) async fn initialize_user_discovery_from_config(&self) -> Result<()> {
let Some(config) = crate::user_config::UserConfig::load_from(self)? else {
return Ok(());
};
if !config.is_user_discovery_enabled {
return Ok(());
}
let discovery_config_path =
PathBuf::from(self.data_dir()).join("user_discovery_config.json");
let settings_are_current = std::fs::read_to_string(&discovery_config_path)
.ok()
.and_then(|value| serde_json::from_str::<serde_json::Value>(&value).ok())
.is_some_and(|value| {
value.get("threshold").and_then(serde_json::Value::as_u64)
== Some(u64::from(config.user_discovery_threshold))
&& value
.get("share_promotion")
.and_then(serde_json::Value::as_bool)
== Some(config.user_discovery_share_promotion)
});
if settings_are_current {
let database = self.get_app_database().await;
let has_shares =
sqlx::query_scalar!("SELECT EXISTS(SELECT 1 FROM user_discovery_shares LIMIT 1)")
.fetch_one(&database.pool)
.await?
!= 0;
if has_shares {
return Ok(());
}
}
let key_manager = self.get_key_manager().await?;
let user_id = key_manager.user_id.ok_or_else(|| {
TwonlyError::Generic("cannot initialize user discovery without user ID".into())
})?;
let identity = key_manager
.signal_identity
.as_ref()
.ok_or(TwonlyError::SignalIdentityNotFound)?;
let identity = IdentityKeyPair::try_from(identity.identity_key_pair_structure.as_slice())
.map_err(|error| {
TwonlyError::Generic(format!("invalid Signal identity: {error}"))
})?;
let public_key = identity.identity_key().serialize().to_vec();
drop(key_manager);
let database = self.get_app_database().await;
let mut transaction = database.pool.begin().await?;
self.get_user_discovery()
.get()
.await
.initialize_or_update(
config.user_discovery_threshold,
user_id,
public_key,
config.user_discovery_share_promotion,
&mut transaction,
)
.await?;
transaction.commit().await?;
database.notify_committed(["user_discovery_shares"]);
Ok(())
}
#[cfg(not(test))]
async fn init_common(config: InitConfig, is_flutter: bool) -> Result<()> {
if GLOBAL_CONTEXT.initialized() {
tracing::info!("twonly already initialized. Ensuring storage directories exist.");
std::fs::create_dir_all(&config.database_dir)?;
std::fs::create_dir_all(&config.data_dir)?;
return Ok(());
}
std::fs::create_dir_all(&config.database_dir)?;
std::fs::create_dir_all(&config.data_dir)?;
let log_dir = PathBuf::from(&config.data_dir).join("log");
init_tracing(&log_dir, is_flutter).await;
SecureStorage::init()?;
let secure_storage = SecureStorage::new("eu.twonly");
let database_dir = PathBuf::from(&config.database_dir.clone());
let rust_db_path = database_dir.join("rust_db.sqlite");
let app_db_path = database_dir.join(APP_DATABASE_FILE);
tracing::info!("Initialized twonly workspace.");
let res: Result<&'static Arc<Context>> = GLOBAL_CONTEXT
.get_or_try_init(|| async {
let key_manager = match KeyManager::try_from_keychain(&secure_storage) {
Ok(key) => key,
Err(err) => {
tracing::error!("{err}");
if rust_db_path.exists() {
tracing::error!("Rust Database exists, while the key manager not. This must be a secure storage error.");
return Err(TwonlyError::SecureStorageError);
}
tracing::info!("Generating a new key manager.");
let new = KeyManager::generate()?;
new.store_to_keychain(&secure_storage)?;
new
}
};
let mut rust_db_key = key_manager.main_key.get_database_key(DatabaseKey::RustDb);
let rust_db = Database::new(
&rust_db_path.display().to_string(),
Some(rust_db_key.as_str()),
false,
)
.await?;
rust_db.run_migrations().await?;
let rust_db = Arc::new(rust_db);
let rust_db_handle = Arc::new(RwLock::new(rust_db));
let mut app_db_key = key_manager.main_key.get_database_key(DatabaseKey::AppDb);
let app_db = AppDatabase::new(
&app_db_path.display().to_string(),
Some(app_db_key.as_str()),
false,
)
.await?;
app_db.run_migrations().await?;
let app_db = Arc::new(RwLock::new(Arc::new(app_db)));
app_db_key.zeroize();
rust_db_key.zeroize();
if is_flutter {
let key_manager = Arc::new(Mutex::new(key_manager));
let signal_engine = {
let key_manager_guard = key_manager.lock().await;
let engine = match (
key_manager_guard.user_id,
&key_manager_guard.signal_identity,
) {
(Some(user_id), Some(signal_identity)) => {
Some(RustSignalEngine::new_with_pool(
rust_db_handle.read().await.pool.clone(),
signal_identity.identity_key_pair_structure.clone(),
signal_identity.registration_id as u32,
user_id.to_string(),
)?)
}
_ => None,
};
Arc::new(Mutex::new(engine))
};
let user_discovery = Shared::new(UserDiscovery::new(
&config.data_dir,
key_manager.clone(),
rust_db_handle.clone(),
)?);
let ctx = Arc::new(Context::Flutter(TwonlyFlutter {
config,
secure_storage,
rust_db: rust_db_handle,
app_db,
key_manager,
user_discovery,
signal_engine,
api_client: OnceCell::const_new(),
}));
if let Err(error) = ctx.initialize_user_discovery_from_config().await {
tracing::warn!("failed to initialize user discovery: {error}");
}
ApiRuntime::initialize(&ctx).await?;
Ok(ctx)
} else {
let key_manager = Arc::new(Mutex::new(key_manager));
let signal_engine = {
let key_manager_guard = key_manager.lock().await;
let engine = match (key_manager_guard.user_id, &key_manager_guard.signal_identity) {
(Some(user_id), Some(identity)) => Some(RustSignalEngine::new_with_pool(
rust_db_handle.read().await.pool.clone(),
identity.identity_key_pair_structure.clone(),
identity.registration_id as u32,
user_id.to_string(),
)?),
_ => None,
};
Arc::new(Mutex::new(engine))
};
let user_discovery = Shared::new(UserDiscovery::new(
&config.data_dir,
key_manager.clone(),
rust_db_handle.clone(),
)?);
let ctx = Arc::new(Context::Standalone(TwonlyStandalone {
config,
rust_db: rust_db_handle,
app_db,
key_manager,
secure_storage,
user_discovery,
signal_engine,
api_client: OnceCell::const_new(),
}));
if let Err(error) = ctx.initialize_user_discovery_from_config().await {
tracing::warn!("failed to initialize user discovery: {error}");
}
ApiRuntime::initialize(&ctx).await?;
Ok(ctx)
}
})
.await;
let ctx = res?;
ApiRuntime::connect(ctx).await?;
Ok(())
}
#[cfg(test)]
async fn init_common(_config: InitConfig, _is_flutter: bool) -> Result<()> {
Err(TwonlyError::Initialization)
}
#[cfg(not(test))]
pub(super) fn get_static() -> Result<&'static Arc<Context>> {
GLOBAL_CONTEXT.get().ok_or(TwonlyError::Initialization)
}
#[cfg(test)]
pub(super) fn get_static() -> Result<&'static Arc<Context>> {
Err(TwonlyError::Initialization)
}
pub(crate) fn get_secure_storage(&self) -> &SecureStorage {
match self {
Self::Flutter(twonly) => &twonly.secure_storage,
Self::Standalone(twonly) => &twonly.secure_storage,
}
}
pub(crate) fn get_config(&self) -> Result<&InitConfig> {
match self {
Self::Flutter(twonly) => Ok(&twonly.config),
Self::Standalone(twonly) => Ok(&twonly.config),
}
}
pub(crate) async fn get_key_manager(&self) -> Result<tokio::sync::MutexGuard<'_, KeyManager>> {
match self {
Self::Flutter(twonly) => Ok(twonly.key_manager.lock().await),
Self::Standalone(twonly) => Ok(twonly.key_manager.lock().await),
}
}
pub(crate) async fn user_id(&self) -> Result<i64> {
self.get_key_manager()
.await?
.user_id
.ok_or_else(|| TwonlyError::Generic("local user ID is missing".into()))
}
pub async fn get_app_database(&self) -> Arc<AppDatabase> {
match self {
Self::Flutter(twonly) => twonly.app_db.read().await.clone(),
Self::Standalone(twonly) => twonly.app_db.read().await.clone(),
}
}
pub(crate) async fn get_rust_db(&self) -> Arc<Database> {
match self {
Self::Flutter(twonly) => twonly.rust_db.read().await.clone(),
Self::Standalone(twonly) => twonly.rust_db.read().await.clone(),
}
}
pub(crate) async fn get_identity(&self, user_id: i64) -> Result<Option<IdentityKey>> {
let database = self.get_rust_db().await;
let user_id = user_id.to_string();
let identity_key = sqlx::query_scalar!(
r#"SELECT identity_key FROM signal_identities WHERE name = ?"#,
user_id,
)
.fetch_optional(&database.pool)
.await?;
identity_key
.map(|bytes| {
IdentityKey::decode(&bytes).map_err(|error| TwonlyError::Signal(error.to_string()))
})
.transpose()
}
pub(crate) async fn replace_rust_database(
&self,
database: Database,
key_manager: &KeyManager,
) -> Result<()> {
let database = Arc::new(database);
match self {
Self::Flutter(twonly) => {
*twonly.rust_db.write().await = database.clone();
let engine = match (key_manager.user_id, &key_manager.signal_identity) {
(Some(user_id), Some(identity)) => Some(RustSignalEngine::new_with_pool(
database.pool.clone(),
identity.identity_key_pair_structure.clone(),
identity.registration_id as u32,
user_id.to_string(),
)?),
_ => None,
};
*twonly.signal_engine.lock().await = engine;
}
Self::Standalone(twonly) => *twonly.rust_db.write().await = database,
}
Ok(())
}
pub(crate) async fn replace_app_database(&self, database: AppDatabase) {
match self {
Self::Flutter(twonly) => *twonly.app_db.write().await = Arc::new(database),
Self::Standalone(twonly) => *twonly.app_db.write().await = Arc::new(database),
}
}
}