[Fix] FS Operations

This commit is contained in:
Alex 2026-09-10 13:48:02 +02:00
commit 68cedff1d9
Signed by: alex
SSH key fingerprint: SHA256:D1+Ub8o0v4K5y1JNivW8IxEOelqLSvPmUzBbDIoZkRQ
12 changed files with 406 additions and 283 deletions

View file

@ -33,10 +33,8 @@ fn now_millis() -> i64 {
.as_millis() as i64
}
pub fn add_user(user: UserProfile) {
if let Err(e) = try_add_user(user) {
eprintln!("Failed to add_user: {}", e);
}
pub fn add_user(user: UserProfile) -> Result<(), crate::storage_error::StorageError> {
try_add_user(user)
}
pub fn try_add_user(user: UserProfile) -> Result<(), crate::storage_error::StorageError> {
@ -89,12 +87,14 @@ pub fn try_add_user_with_credential_origin(
})
}
pub fn update_user(user: UserProfile) {
add_user(user);
pub fn update_user(user: UserProfile) -> Result<(), crate::storage_error::StorageError> {
try_add_user(user)
}
pub fn get_user_by_username(username: &str) -> Option<UserProfile> {
match db::with_db(|conn| {
pub fn get_user_by_username(
username: &str,
) -> Result<Option<UserProfile>, crate::storage_error::StorageError> {
let user = db::with_db(|conn| {
match conn.query_row(
"SELECT user_id, username, public_key, private_key_hash, reset_token, created_at, display_name FROM users WHERE username = ?1 LIMIT 1",
params![username],
@ -116,13 +116,12 @@ pub fn get_user_by_username(username: &str) -> Option<UserProfile> {
Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
Err(e) => Err(e.into()),
}
}) {
Ok(opt) => opt,
Err(e) => {
eprintln!("Error querying user by username: {}", e);
None
}
}
})?;
user.map(|mut user| {
user.trusted_apps = load_trusted_apps(user.user_id)?;
Ok(user)
})
.transpose()
}
pub fn get_user(user_id: i64) -> Result<Option<UserProfile>, crate::storage_error::StorageError> {
@ -156,8 +155,8 @@ pub fn get_user(user_id: i64) -> Result<Option<UserProfile>, crate::storage_erro
.transpose()
}
pub fn get_users() -> Vec<UserProfile> {
match db::with_db(|conn| {
pub fn get_users() -> Result<Vec<UserProfile>, crate::storage_error::StorageError> {
db::with_db(|conn| {
let mut stmt = conn.prepare(
r#"
SELECT user_id, username, public_key, private_key_hash, reset_token, created_at, display_name
@ -189,41 +188,36 @@ pub fn get_users() -> Vec<UserProfile> {
let mut out = Vec::new();
for row in rows {
match row {
Ok(mut user) => {
user.trusted_apps = load_trusted_apps(user.user_id)?;
out.push(user);
}
Err(e) => eprintln!("Failed to read user row: {}", e),
}
let mut user = row?;
user.trusted_apps = load_trusted_apps_from(conn, user.user_id)?;
out.push(user);
}
Ok(out)
}) {
Ok(v) => v,
Err(e) => {
eprintln!("Failed to query users: {}", e);
Vec::new()
}
}
})
}
fn load_trusted_apps(
user_id: i64,
) -> Result<std::collections::HashMap<String, String>, crate::storage_error::StorageError> {
db::with_db(|conn| {
let mut stmt =
conn.prepare("SELECT app_id, app_secret FROM trusted_apps WHERE user_id = ?1")?;
let rows = stmt.query_map(params![user_id], |r| {
Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?))
})?;
db::with_db(|conn| load_trusted_apps_from(conn, user_id))
}
let mut map = std::collections::HashMap::new();
for row in rows {
let (key, value) = row?;
map.insert(key, value);
}
Ok(map)
})
fn load_trusted_apps_from(
conn: &rusqlite::Connection,
user_id: i64,
) -> Result<std::collections::HashMap<String, String>, crate::storage_error::StorageError> {
let mut stmt =
conn.prepare("SELECT app_id, app_secret FROM trusted_apps WHERE user_id = ?1")?;
let rows = stmt.query_map(params![user_id], |r| {
Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?))
})?;
let mut map = std::collections::HashMap::new();
for row in rows {
let (key, value) = row?;
map.insert(key, value);
}
Ok(map)
}
pub fn revoke_trusted_app(
@ -249,17 +243,15 @@ pub fn revoke_all_trusted_apps(user_id: i64) -> Result<usize, crate::storage_err
})
}
pub fn remove_user(user_id: i64) {
if let Err(e) = db::with_db(|conn| {
pub fn remove_user(user_id: i64) -> Result<(), crate::storage_error::StorageError> {
db::with_immediate_transaction(|conn| {
conn.execute(
"DELETE FROM trusted_apps WHERE user_id = ?1",
params![user_id],
)?;
conn.execute("DELETE FROM users WHERE user_id = ?1", params![user_id])?;
Ok(())
}) {
eprintln!("Failed to remove_user: {}", e);
}
})
}
/// Remove only local management authority. Hosted content is intentionally
@ -281,15 +273,12 @@ pub fn finalize_local_release(
user_id: i64,
username_hint: Option<&str>,
) -> Result<(), crate::storage_error::StorageError> {
let username = get_user(user_id)?
.map(|user| user.username)
.or_else(|| {
get_residency_by_id(user_id)
.ok()
.flatten()
.map(|residency| residency.username)
})
.or_else(|| username_hint.map(str::to_owned));
let username = match get_user(user_id)? {
Some(user) => Some(user.username),
None => get_residency_by_id(user_id)?
.map(|residency| residency.username)
.or_else(|| username_hint.map(str::to_owned)),
};
let Some(username) = username else {
return Err(crate::storage_error::StorageError::Other(
"user residency was not found".into(),
@ -553,12 +542,13 @@ pub fn purge_user_data(user_id: i64) -> Result<(), crate::storage_error::Storage
/// Complete local erasure is idempotent and is the target for a durable
/// Omega-hosted erasure request after account deletion.
pub fn erase_user_locally(user_id: i64) -> Result<(), crate::storage_error::StorageError> {
let username = get_user(user_id)?.map(|user| user.username).or_else(|| {
get_residency()
let username = match get_user(user_id)? {
Some(user) => Some(user.username),
None => get_residency()?
.into_iter()
.find(|entry| entry.user_id == user_id)
.map(|entry| entry.username)
});
.map(|entry| entry.username),
};
purge_user_data(user_id)?;
db::with_db(|conn| {
conn.execute(
@ -576,20 +566,29 @@ pub fn erase_user_locally(user_id: i64) -> Result<(), crate::storage_error::Stor
.map_err(|error| crate::storage_error::StorageError::Other(error.to_string()))
}
pub fn get_residency() -> Vec<UserResidency> {
pub fn get_residency() -> Result<Vec<UserResidency>, crate::storage_error::StorageError> {
db::with_db(|conn| {
let mut stmt = conn.prepare("SELECT user_id, username, lifecycle_state, data_state, credential_origin FROM user_residency ORDER BY username")?;
let rows = stmt.query_map([], |row| {
let lifecycle: String = row.get(2)?;
Ok(UserResidency {
user_id: row.get(0)?, username: row.get(1)?,
state: if lifecycle == "managed" { LocalUserState::Managed } else { LocalUserState::Released },
user_id: row.get(0)?,
username: row.get(1)?,
state: if lifecycle == "managed" {
LocalUserState::Managed
} else {
LocalUserState::Released
},
data_present: row.get::<_, String>(3)? == "present",
credential_origin: if row.get::<_, String>(4)? == "external" { CredentialOrigin::External } else { CredentialOrigin::Local },
credential_origin: if row.get::<_, String>(4)? == "external" {
CredentialOrigin::External
} else {
CredentialOrigin::Local
},
})
})?;
rows.collect::<Result<Vec<_>, _>>().map_err(Into::into)
}).unwrap_or_default()
})
}
pub fn get_residency_by_id(
@ -622,15 +621,13 @@ pub fn get_residency_by_id(
})
}
pub fn clear() {
if let Err(e) = db::with_db(|conn| {
pub fn clear() -> Result<(), crate::storage_error::StorageError> {
db::with_immediate_transaction(|conn| {
conn.execute_batch(
"DELETE FROM trusted_apps; DELETE FROM users; DELETE FROM user_residency;",
)?;
Ok(())
}) {
eprintln!("Failed to clear users: {}", e);
}
})
}
pub fn save_users() {
@ -648,7 +645,7 @@ pub fn load_users_sync() -> std::io::Result<()> {
if let json::JsonValue::Array(arr) = parsed {
for j in arr.iter() {
if let Some(up) = UserProfile::from_json(j) {
add_user(up);
try_add_user(up).map_err(|error| std::io::Error::other(error.to_string()))?;
}
}
}

View file

@ -2,7 +2,6 @@ use crate::storage_error::StorageError;
use crate::util::db;
use crate::util::message_storage_policy::{self, MessageRetention};
use crate::util::sync::{self, EntityType, Operation};
use iota_logger::log;
use rusqlite::{OptionalExtension, Transaction, params};
pub const MAX_UNIQUE_REACTIONS_PER_MESSAGE: usize = 10;
@ -907,9 +906,9 @@ pub fn change_message_state(
fn load_reactions(
conn: &rusqlite::Connection,
msg_ids: &[i64],
) -> std::collections::HashMap<i64, Vec<StoredReaction>> {
) -> Result<std::collections::HashMap<i64, Vec<StoredReaction>>, StorageError> {
if msg_ids.is_empty() {
return std::collections::HashMap::new();
return Ok(std::collections::HashMap::new());
}
let placeholders: Vec<String> = msg_ids
@ -924,34 +923,33 @@ fn load_reactions(
let mut map: std::collections::HashMap<i64, Vec<StoredReaction>> =
std::collections::HashMap::new();
if let Ok(mut stmt) = conn.prepare(&query) {
let params: Vec<&dyn rusqlite::types::ToSql> = msg_ids
let mut stmt = conn.prepare(&query)?;
let params: Vec<&dyn rusqlite::types::ToSql> = msg_ids
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql)
.collect();
let rows = stmt.query_map(params.as_slice(), |row| {
Ok((
row.get::<_, i64>(0)?,
StoredReaction {
reaction: row.get(1)?,
user_id: row.get(2)?,
},
))
})?;
for row in rows {
let row = row?;
let reactions = map.entry(row.0).or_default();
if reactions
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql)
.collect();
if let Ok(rows) = stmt.query_map(params.as_slice(), |row| {
Ok((
row.get::<_, i64>(0)?,
StoredReaction {
reaction: row.get(1)?,
user_id: row.get(2)?,
},
))
}) {
for row in rows.flatten() {
let reactions = map.entry(row.0).or_default();
if reactions
.iter()
.any(|stored: &StoredReaction| stored.reaction == row.1.reaction)
{
reactions.push(row.1);
} else if reactions.len() < MAX_UNIQUE_REACTIONS_PER_MESSAGE {
reactions.push(row.1);
}
}
.any(|stored: &StoredReaction| stored.reaction == row.1.reaction)
{
reactions.push(row.1);
} else if reactions.len() < MAX_UNIQUE_REACTIONS_PER_MESSAGE {
reactions.push(row.1);
}
}
map
Ok(map)
}
pub fn get_messages(
@ -959,12 +957,12 @@ pub fn get_messages(
external_user: i64,
loaded_messages: i64,
amount: i64,
) -> Vec<StoredMessage> {
) -> Result<Vec<StoredMessage>, StorageError> {
if amount <= 0 || loaded_messages < 0 {
return Vec::new();
return Ok(Vec::new());
}
match db::with_db(|conn| {
db::with_db(|conn| {
let mut stmt = conn.prepare(
r#"
SELECT id, relay_signer_id, relay_message_id, message_time, authored_at,
@ -1000,10 +998,10 @@ pub fn get_messages(
content: row.get(13)?,
sent_by_self: row.get::<_, i64>(14)? != 0,
message_state: row.get(15)?,
height: row.get(16).unwrap_or(0),
key_version: row.get(17).unwrap_or(1),
reply_to: row.get(18).ok().flatten(),
edited: row.get::<_, i64>(19).unwrap_or(0) > 0,
height: row.get(16)?,
key_version: row.get(17)?,
reply_to: row.get(18)?,
edited: row.get::<_, i64>(19)? > 0,
reactions: Vec::new(),
})
},
@ -1011,26 +1009,17 @@ pub fn get_messages(
let mut out = Vec::new();
for row in rows {
match row {
Ok(msg) => out.push(msg),
Err(e) => log!("Failed to read row from sqlite: {}", e),
}
out.push(row?);
}
let msg_ids: Vec<i64> = out.iter().map(|m| m.id).collect();
let reaction_map = load_reactions(conn, &msg_ids);
let reaction_map = load_reactions(conn, &msg_ids)?;
for msg in &mut out {
msg.reactions = reaction_map.get(&msg.id).cloned().unwrap_or_default();
}
Ok(out)
}) {
Ok(v) => v,
Err(e) => {
log!("Failed to query messages: {}", e);
Vec::new()
}
}
})
}
pub fn get_message(
@ -1074,10 +1063,10 @@ pub fn get_message(
content: row.get(13)?,
sent_by_self: row.get::<_, i64>(14)? != 0,
message_state: row.get(15)?,
height: row.get(16).unwrap_or(0),
key_version: row.get(17).unwrap_or(1),
reply_to: row.get(18).ok().flatten(),
edited: row.get::<_, i64>(19).unwrap_or(0) > 0,
height: row.get(16)?,
key_version: row.get(17)?,
reply_to: row.get(18)?,
edited: row.get::<_, i64>(19)? > 0,
external_user: row.get(20)?,
reactions: Vec::new(),
})
@ -1099,7 +1088,7 @@ pub fn get_message(
}
let mut message = messages.into_iter().next().expect("checked non-empty");
let reaction_map = load_reactions(conn, &[message.id]);
let reaction_map = load_reactions(conn, &[message.id])?;
message.reactions = reaction_map.get(&message.id).cloned().unwrap_or_default();
Ok(Some(message))
})
@ -1142,14 +1131,17 @@ pub fn get_message_with_offset(
Ok(Some((message, offset)))
}
pub fn get_messages_by_ids(storage_owner: i64, ids: &[i64]) -> Vec<StoredMessage> {
pub fn get_messages_by_ids(
storage_owner: i64,
ids: &[i64],
) -> Result<Vec<StoredMessage>, StorageError> {
if ids.is_empty() {
return Vec::new();
return Ok(Vec::new());
}
let wanted: std::collections::HashSet<i64> = ids.iter().copied().collect();
// A journal id uniquely identifies a row. Load all messages for this owner and retain only
// those ids; this keeps reaction hydration identical to normal message loading.
match db::with_db(|conn| {
db::with_db(|conn| {
let mut stmt = conn.prepare("SELECT id, relay_signer_id, relay_message_id, message_time, authored_at, origin_iota_received_at, destination_iota_received_at, client_received_at, client_received_recorded_at, read_at, read_recorded_at, delivery_failed_at, delivery_failure, content, sent_by_self, message_state, height, key_version, reply_to, edited_count, external_user FROM messages WHERE storage_owner = ?1 AND deleted_by_external = 0 AND history_deleted = 0")?;
let rows = stmt.query_map([storage_owner], |row| {
let external_user: i64 = row.get(20)?;
@ -1171,10 +1163,10 @@ pub fn get_messages_by_ids(storage_owner: i64, ids: &[i64]) -> Vec<StoredMessage
content: row.get(13)?,
sent_by_self: row.get::<_, i64>(14)? != 0,
message_state: row.get(15)?,
height: row.get(16).unwrap_or(0),
key_version: row.get(17).unwrap_or(1),
reply_to: row.get(18).ok().flatten(),
edited: row.get::<_, i64>(19).unwrap_or(0) > 0,
height: row.get(16)?,
key_version: row.get(17)?,
reply_to: row.get(18)?,
edited: row.get::<_, i64>(19)? > 0,
reactions: Vec::new(),
})
})?;
@ -1185,35 +1177,24 @@ pub fn get_messages_by_ids(storage_owner: i64, ids: &[i64]) -> Vec<StoredMessage
messages.push(message);
}
}
let reaction_map = load_reactions(conn, &messages.iter().map(|m| m.id).collect::<Vec<_>>());
let reaction_map =
load_reactions(conn, &messages.iter().map(|m| m.id).collect::<Vec<_>>())?;
for message in &mut messages {
message.reactions = reaction_map.get(&message.id).cloned().unwrap_or_default();
}
Ok(messages)
}) {
Ok(messages) => messages,
Err(e) => {
log!("Failed to query messages by id: {}", e);
Vec::new()
}
}
})
}
pub fn get_all_messages(storage_owner: i64) -> Vec<StoredMessage> {
let ids = match db::with_db(|conn| {
pub fn get_all_messages(storage_owner: i64) -> Result<Vec<StoredMessage>, StorageError> {
let ids = db::with_db(|conn| {
let mut stmt = conn.prepare(
"SELECT id FROM messages WHERE storage_owner = ?1 AND deleted_by_external = 0 AND history_deleted = 0",
)?;
Ok(stmt
.query_map([storage_owner], |row| row.get::<_, i64>(0))?
.collect::<Result<Vec<_>, _>>()?)
}) {
Ok(ids) => ids,
Err(e) => {
log!("Failed to query all messages: {}", e);
return Vec::new();
}
};
})?;
get_messages_by_ids(storage_owner, &ids)
}

View file

@ -23,8 +23,13 @@ impl CommunitiesUtil {
})
}
pub fn add_community(storage_owner: i64, address: String, title: String, position: String) {
if let Err(e) = db::with_db(|conn| {
pub fn add_community(
storage_owner: i64,
address: String,
title: String,
position: String,
) -> Result<(), StorageError> {
db::with_db(|conn| {
conn.execute(
r#"
INSERT INTO communities (storage_owner, address, title, position)
@ -36,9 +41,7 @@ impl CommunitiesUtil {
params![storage_owner, address, title, position],
)?;
Ok(())
}) {
eprintln!("Failed to add_community: {}", e);
}
})
}
pub fn remove_community(
@ -59,8 +62,8 @@ impl CommunitiesUtil {
})
}
pub fn get_communities(storage_owner: i64) -> Vec<StoredCommunity> {
match db::with_db(|conn| {
pub fn get_communities(storage_owner: i64) -> Result<Vec<StoredCommunity>, StorageError> {
db::with_db(|conn| {
let mut stmt = conn.prepare(
r#"
SELECT address, title, position
@ -77,20 +80,7 @@ impl CommunitiesUtil {
})
})?;
let mut out = Vec::new();
for row in rows {
match row {
Ok(community) => out.push(community),
Err(e) => eprintln!("Failed to read community row: {}", e),
}
}
Ok(out)
}) {
Ok(v) => v,
Err(e) => {
eprintln!("Failed to query communities in get_communities: {}", e);
Vec::new()
}
}
rows.collect::<Result<Vec<_>, _>>().map_err(Into::into)
})
}
}

View file

@ -5,11 +5,37 @@ use std::fs;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::OnceLock;
use thiserror::Error;
pub static CONFIG: Lazy<ArcSwap<IotaConfig>> =
Lazy::new(|| ArcSwap::new(Arc::new(IotaConfig::default())));
#[derive(Debug, Error)]
pub enum ConfigError {
#[error("cannot read {path}: {source}")]
Read {
path: PathBuf,
#[source]
source: std::io::Error,
},
#[error("cannot parse {path}: {source}")]
Parse {
path: PathBuf,
#[source]
source: serde_yaml::Error,
},
#[error("invalid web.bind {bind:?}: {source}")]
InvalidWebBind {
bind: String,
#[source]
source: std::net::AddrParseError,
},
#[error("max_ipc_clients must be greater than zero")]
InvalidMaxIpcClients,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct IotaConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub iota_id: Option<u64>,
@ -49,6 +75,7 @@ impl Default for WebMode {
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct WebSettings {
#[serde(default)]
pub mode: WebMode,
@ -113,30 +140,54 @@ impl Default for IotaConfig {
}
}
pub fn load_config() {
load_config_from(&default_config_path());
pub fn load_config() -> Result<(), ConfigError> {
load_config_from(&default_config_path())
}
/// Loading is intentionally side-effect free: a missing configuration means
/// documented defaults, not a newly-created file.
pub fn load_config_from(path: &Path) {
pub fn load_config_from(path: &Path) -> Result<(), ConfigError> {
let s = match fs::read_to_string(path) {
Ok(contents) => contents,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return,
Err(error) => {
eprintln!("Failed to read {}: {error}", path.display());
return;
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
CONFIG.store(Arc::new(IotaConfig::default()));
return Ok(());
}
Err(source) => {
return Err(ConfigError::Read {
path: path.into(),
source,
});
}
};
match serde_yaml::from_str::<IotaConfig>(&s) {
Ok(parsed) => {
CONFIG.store(Arc::new(parsed));
}
Err(e) => {
eprintln!("Failed to parse {}: {e}", path.display());
}
let parsed = parse_config(path, &s)?;
CONFIG.store(Arc::new(parsed));
Ok(())
}
fn parse_config(path: &Path, yaml: &str) -> Result<IotaConfig, ConfigError> {
let parsed = serde_yaml::from_str::<IotaConfig>(yaml).map_err(|source| ConfigError::Parse {
path: path.into(),
source,
})?;
validate_config(&parsed)?;
Ok(parsed)
}
pub fn validate_config(config: &IotaConfig) -> Result<(), ConfigError> {
config
.web
.bind
.parse::<std::net::IpAddr>()
.map_err(|source| ConfigError::InvalidWebBind {
bind: config.web.bind.clone(),
source,
})?;
if config.max_ipc_clients == 0 {
return Err(ConfigError::InvalidMaxIpcClients);
}
Ok(())
}
pub fn clear_config() {
@ -232,6 +283,9 @@ pub fn modify_config_value(key: &str, value: &str) -> Result<(), &'static str> {
Ok(())
}
"web.bind" => {
value
.parse::<std::net::IpAddr>()
.map_err(|_| "invalid web.bind")?;
let bind = value.to_string();
modify_config(|cfg| cfg.web.bind = bind);
Ok(())
@ -244,3 +298,45 @@ static CONFIG_PATH: OnceLock<PathBuf> = OnceLock::new();
pub fn configure_config_path(path: PathBuf) {
let _ = CONFIG_PATH.set(path);
}
#[cfg(test)]
mod tests {
use super::{ConfigError, IotaConfig, parse_config, validate_config};
use std::path::Path;
#[test]
fn malformed_yaml_is_rejected() {
assert!(matches!(
parse_config(Path::new("config.yaml"), "web: ["),
Err(ConfigError::Parse { .. })
));
}
#[test]
fn unknown_explicit_fields_are_rejected() {
assert!(matches!(
parse_config(Path::new("config.yaml"), "unexpected: true\n"),
Err(ConfigError::Parse { .. })
));
}
#[test]
fn invalid_explicit_bind_is_rejected() {
let mut config = IotaConfig::default();
config.web.bind = "localhost:1984".into();
assert!(matches!(
validate_config(&config),
Err(ConfigError::InvalidWebBind { .. })
));
}
#[test]
fn zero_ipc_capacity_is_rejected() {
let mut config = IotaConfig::default();
config.max_ipc_clients = 0;
assert!(matches!(
validate_config(&config),
Err(ConfigError::InvalidMaxIpcClients)
));
}
}

View file

@ -51,19 +51,42 @@ pub fn with_db<T, F>(f: F) -> Result<T, StorageError>
where
F: FnOnce(&Connection) -> Result<T, StorageError>,
{
let conn = POOL.get().map_err(|e| StorageError::Pool(e.to_string()))?;
f(&conn)
blocking_region(|| {
let conn = POOL.get().map_err(|e| StorageError::Pool(e.to_string()))?;
f(&conn)
})
}
pub fn with_immediate_transaction<T, F>(f: F) -> Result<T, StorageError>
where
F: FnOnce(&Transaction<'_>) -> Result<T, StorageError>,
{
let mut conn = POOL.get().map_err(|e| StorageError::Pool(e.to_string()))?;
let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Immediate)?;
let value = f(&tx)?;
tx.commit()?;
Ok(value)
blocking_region(|| {
let mut conn = POOL.get().map_err(|e| StorageError::Pool(e.to_string()))?;
let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Immediate)?;
let value = f(&tx)?;
tx.commit()?;
Ok(value)
})
}
/* SQLite remains synchronous, but daemon transports run on Tokio. Entering a
* blocking region lets Tokio replace the worker that is waiting on the bounded
* eight-connection pool, so database latency cannot stall unrelated I/O. A
* current-thread runtime (used by unit tests and small tools) has no worker to
* hand off and therefore executes directly. */
fn blocking_region<T>(f: impl FnOnce() -> T) -> T {
match tokio::runtime::Handle::try_current() {
Ok(handle)
if matches!(
handle.runtime_flavor(),
tokio::runtime::RuntimeFlavor::MultiThread
) =>
{
tokio::task::block_in_place(f)
}
_ => f(),
}
}
/* Verify the persistent database before the pool is initialized. A corrupt
@ -156,9 +179,7 @@ fn add_table_column_if_missing(
}
fn run_migrations_on_connection(conn: &Connection) -> Result<(), StorageError> {
let current_version: i64 = conn
.pragma_query_value(None, "user_version", |r| r.get(0))
.unwrap_or(0);
let current_version: i64 = conn.pragma_query_value(None, "user_version", |r| r.get(0))?;
if current_version < 1 {
conn.execute_batch(
@ -272,13 +293,11 @@ fn run_migrations_on_connection(conn: &Connection) -> Result<(), StorageError> {
)?;
}
let messages_exist: bool = conn
.query_row(
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'messages')",
[],
|row| row.get(0),
)
.unwrap_or(false);
let messages_exist: bool = conn.query_row(
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'messages')",
[],
|row| row.get(0),
)?;
if messages_exist {
add_column_if_missing(
conn,
@ -926,19 +945,12 @@ pub fn with_conn<T, F>(shared: &Arc<std::sync::Mutex<Connection>>, f: F) -> Resu
where
F: FnOnce(&Connection) -> Result<T, rusqlite::Error>,
{
if tokio::runtime::Handle::try_current().is_ok() {
tokio::task::block_in_place(|| {
let guard = shared
.lock()
.map_err(|e| format!("DB mutex poisoned: {:?}", e))?;
f(&*guard).map_err(|e| e.to_string())
})
} else {
blocking_region(|| {
let guard = shared
.lock()
.map_err(|e| format!("DB mutex poisoned: {:?}", e))?;
f(&*guard).map_err(|e| e.to_string())
}
})
}
/// Legacy - kept for e2ee_storage which uses its own DB.