radrootsd

JSON-RPC bridge for Radroots event publishing
git clone https://radroots.dev/git/radrootsd.git
Log | Files | Refs | README | LICENSE

commit 350b4c16a32572ff65e644ce2db2556135fc0426
parent d9116fd34401367f6d6503859b25a3a7829bda92
Author: triesap <tyson@radroots.org>
Date:   Wed, 15 Jul 2026 03:50:00 +0000

sqlite: port transport publish store to SQLx

- replace rusqlite transport publish storage with bundled SQLx SQLite
- preserve schema validation, idempotency, and recovery tests on the SQLx path
- remove rusqlite from the dependency graph and lockfile
- close Rust 1.97 clippy warning hygiene for the workspace

Diffstat:
MCargo.lock | 253++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++-------------
MCargo.toml | 3++-
Msrc/app/config.rs | 30++++++++++++++++++++----------
Msrc/app/runtime.rs | 46+++++++++++++++++++++++-----------------------
Msrc/core/nip46/session.rs | 6++++++
Msrc/core/transport_publish.rs | 1076+++++++++++++++++++++++++++++++++++++++++++++----------------------------------
Msrc/main.rs | 14++++++--------
Msrc/transport/jsonrpc/methods/nip46/connect.rs | 13+++++--------
Msrc/transport/jsonrpc/methods/nip46/get_public_key.rs | 11+++++------
Msrc/transport/jsonrpc/nip46/client.rs | 4++--
Msrc/transport/jsonrpc/server.rs | 90++++++++++++++++++++++++++++++++++++++++----------------------------------------
Msrc/transport/nostr/listener.rs | 12++++++------
12 files changed, 950 insertions(+), 608 deletions(-)

diff --git a/Cargo.lock b/Cargo.lock @@ -170,6 +170,15 @@ dependencies = [ ] [[package]] +name = "atoi" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f28d99ec8bfea296261ca1af174f24225171fea9664ba9003cbebee704810528" +dependencies = [ + "num-traits", +] + +[[package]] name = "atomic-destructor" version = "0.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -404,6 +413,15 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" [[package]] +name = "concurrent-queue" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" +dependencies = [ + "crossbeam-utils", +] + +[[package]] name = "config" version = "0.14.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -467,6 +485,21 @@ dependencies = [ ] [[package]] +name = "crc" +version = "3.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5eb8a2a1cd12ab0d987a5d5e825195d372001a4094a0376319d5a0ad71c1ba0d" +dependencies = [ + "crc-catalog", +] + +[[package]] +name = "crc-catalog" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "217698eaf96b4a3f0bc4f3662aaa55bdf913cd54d7204591faa790070c6d0853" + +[[package]] name = "crossbeam-channel" version = "0.5.15" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -476,6 +509,15 @@ dependencies = [ ] [[package]] +name = "crossbeam-queue" +version = "0.3.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "803d13fb3b09d88be9f4dbc29062c66b19bf7170867ceb746d2a8689bf6c7a26" +dependencies = [ + "crossbeam-utils", +] + +[[package]] name = "crossbeam-utils" version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -545,10 +587,19 @@ dependencies = [ ] [[package]] +name = "dotenvy" +version = "0.15.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b" + +[[package]] name = "either" version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +dependencies = [ + "serde", +] [[package]] name = "encoding_rs" @@ -576,16 +627,15 @@ dependencies = [ ] [[package]] -name = "fallible-iterator" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" - -[[package]] -name = "fallible-streaming-iterator" -version = "0.1.9" +name = "event-listener" +version = "5.4.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" +checksum = "e13b66accf52311f30a0db42147dadea9850cb48cd070028831ae5f5d4b856ab" +dependencies = [ + "concurrent-queue", + "parking", + "pin-project-lite", +] [[package]] name = "fastrand" @@ -600,6 +650,17 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" [[package]] +name = "flume" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e139bc46ca777eb5efaf62df0ab8cc5fd400866427e56c68b22e414e53bd3be" +dependencies = [ + "futures-core", + "futures-sink", + "spin", +] + +[[package]] name = "fnv" version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -669,6 +730,17 @@ dependencies = [ ] [[package]] +name = "futures-intrusive" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d930c203dd0b6ff06e0201a4a2fe9149b43c684fd4420555b26d21b1a02956f" +dependencies = [ + "futures-core", + "lock_api", + "parking_lot", +] + +[[package]] name = "futures-io" version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -820,6 +892,8 @@ version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" dependencies = [ + "allocator-api2", + "equivalent", "foldhash 0.2.0", ] @@ -1293,6 +1367,8 @@ name = "libsqlite3-sys" version = "0.37.0" dependencies = [ "cc", + "pkg-config", + "vcpkg", ] [[package]] @@ -1513,6 +1589,12 @@ dependencies = [ ] [[package]] +name = "parking" +version = "2.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" + +[[package]] name = "parking_lot" version = "0.12.5" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1644,6 +1726,12 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" [[package]] +name = "pkg-config" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" + +[[package]] name = "poly1305" version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1929,6 +2017,7 @@ version = "0.1.0" dependencies = [ "anyhow", "clap", + "futures-executor", "jsonrpsee", "nostr", "radroots_event", @@ -1940,11 +2029,11 @@ dependencies = [ "radroots_transport_nostr", "radroots_transport_publish_protocol", "rand 0.9.2", - "rusqlite", "serde", "serde_json", "serde_qs", "sha2", + "sqlx", "tempfile", "thiserror 2.0.18", "tokio", @@ -2110,31 +2199,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "afab94fb28594581f62d981211a9a4d53cc8130bbcbbb89a0440d9b8e81a7746" [[package]] -name = "rsqlite-vfs" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c51c9ae4df8a7fba42103df5c621fa3c37eccf3a3c650879e90fc48b11cc192c" -dependencies = [ - "hashbrown 0.16.1", - "thiserror 2.0.18", -] - -[[package]] -name = "rusqlite" -version = "0.39.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a0d2b0146dd9661bf67bb107c0bb2a55064d556eeb3fc314151b957f313bcd4e" -dependencies = [ - "bitflags", - "fallible-iterator", - "fallible-streaming-iterator", - "hashlink 0.11.1", - "libsqlite3-sys", - "smallvec", - "sqlite-wasm-rs", -] - -[[package]] name = "rust-ini" version = "0.20.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2446,15 +2510,113 @@ dependencies = [ ] [[package]] -name = "sqlite-wasm-rs" -version = "0.5.5" +name = "spin" +version = "0.9.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dc3efc0da82635d7e1ced0053bbbfa8c7ab9645d0bf36ceb4f7127bb85315d75" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" dependencies = [ - "cc", - "js-sys", - "rsqlite-vfs", - "wasm-bindgen", + "lock_api", +] + +[[package]] +name = "sqlx" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "378620ccc25c62c89d8be1c819e76a88d59bdcc3304733330788948e619bfd71" +dependencies = [ + "sqlx-core", + "sqlx-macros", + "sqlx-sqlite", +] + +[[package]] +name = "sqlx-core" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05b44e85bf579a8eeb4ceaa77a3a523baf2bf0e9bac7e40f405d537b5d2d5ccb" +dependencies = [ + "base64 0.22.1", + "bytes", + "cfg-if", + "crc", + "crossbeam-queue", + "either", + "event-listener", + "futures-core", + "futures-intrusive", + "futures-io", + "futures-util", + "hashbrown 0.16.1", + "hashlink 0.11.1", + "indexmap", + "log", + "memchr", + "percent-encoding", + "serde", + "sha2", + "smallvec", + "thiserror 2.0.18", + "tracing", + "url", +] + +[[package]] +name = "sqlx-macros" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bd2b84f2bc39a5705ef27ec785a11c934a41bbd4a24941e257927cddc26b60bf" +dependencies = [ + "proc-macro2", + "quote", + "sqlx-core", + "sqlx-macros-core", + "syn", +] + +[[package]] +name = "sqlx-macros-core" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb8d96de5fdc85a5c4ec813432b523ec637e80ba98f046555f75f7908ddac7c3" +dependencies = [ + "cfg-if", + "dotenvy", + "either", + "heck", + "hex", + "proc-macro2", + "quote", + "serde", + "serde_json", + "sha2", + "sqlx-core", + "sqlx-sqlite", + "syn", + "url", +] + +[[package]] +name = "sqlx-sqlite" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "488e99c397a62007e4229aec669a179816339afc6d2620ca6fa420dbee2e982c" +dependencies = [ + "atoi", + "flume", + "form_urlencoded", + "futures-channel", + "futures-core", + "futures-executor", + "futures-intrusive", + "futures-util", + "libsqlite3-sys", + "log", + "percent-encoding", + "serde", + "sqlx-core", + "thiserror 2.0.18", + "tracing", + "url", ] [[package]] @@ -2817,6 +2979,7 @@ version = "0.1.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" dependencies = [ + "log", "pin-project-lite", "tracing-attributes", "tracing-core", @@ -3013,6 +3176,12 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" [[package]] +name = "vcpkg" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" + +[[package]] name = "version_check" version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" diff --git a/Cargo.toml b/Cargo.toml @@ -36,12 +36,13 @@ nostr = { version = "0.44.2", features = ["nip46"] } anyhow = { version = "1" } clap = { version = "4", features = ["derive"] } jsonrpsee = { version = "0.26", features = ["server"] } +futures-executor = { version = "0.3" } rand = { version = "0.9" } -rusqlite = { version = "0.39.0", features = ["bundled"] } serde = { version = "1", default-features = false, features = ["derive"] } serde_json = { version = "1", default-features = false } serde_qs = { version = "1.0" } sha2 = { version = "0.10" } +sqlx = { version = "0.9.0", default-features = false, features = ["derive", "sqlite-bundled"] } tokio = { version = "1", features = ["full"] } thiserror = { version = "2" } tower = { version = "0.5.3", features = ["util"] } diff --git a/src/app/config.rs b/src/app/config.rs @@ -504,20 +504,30 @@ mod tests { #[test] fn transport_publish_validation_rejects_zero_limits() { - let mut cfg = TransportPublishConfig::default(); - cfg.max_event_bytes = 0; + let cfg = TransportPublishConfig { + max_event_bytes: 0, + ..TransportPublishConfig::default() + }; assert!(cfg.validate().is_err()); - let mut cfg = TransportPublishConfig::default(); - cfg.max_targets_per_request = 0; + let cfg = TransportPublishConfig { + max_targets_per_request: 0, + ..TransportPublishConfig::default() + }; assert!(cfg.validate().is_err()); - let mut cfg = TransportPublishConfig::default(); - cfg.job_list_limit = 0; + let cfg = TransportPublishConfig { + job_list_limit: 0, + ..TransportPublishConfig::default() + }; assert!(cfg.validate().is_err()); - let mut cfg = TransportPublishConfig::default(); - cfg.max_concurrent_publish_jobs = 0; + let cfg = TransportPublishConfig { + max_concurrent_publish_jobs: 0, + ..TransportPublishConfig::default() + }; assert!(cfg.validate().is_err()); - let mut cfg = TransportPublishConfig::default(); - cfg.connect_timeout_secs = 0; + let cfg = TransportPublishConfig { + connect_timeout_secs: 0, + ..TransportPublishConfig::default() + }; assert!(cfg.validate().is_err()); } diff --git a/src/app/runtime.rs b/src/app/runtime.rs @@ -27,9 +27,13 @@ use radroots_nostr::prelude::{ use std::path::PathBuf; #[cfg(test)] -static RUN_LOAD_HOOK: std::sync::OnceLock< - std::sync::Mutex<Option<Result<(cli::Args, config::Settings), String>>>, -> = std::sync::OnceLock::new(); +type RunLoadHookValue = Result<(cli::Args, config::Settings), String>; + +#[cfg(test)] +type RunLoadHook = std::sync::Mutex<Option<RunLoadHookValue>>; + +#[cfg(test)] +static RUN_LOAD_HOOK: std::sync::OnceLock<RunLoadHook> = std::sync::OnceLock::new(); #[cfg(test)] static RUN_BOOTSTRAP_HOOK: std::sync::OnceLock<std::sync::Mutex<Option<Result<(), String>>>> = @@ -71,8 +75,7 @@ struct RadrootsdRuntimeStartupReport { } #[cfg(test)] -fn run_load_hook() --> &'static std::sync::Mutex<Option<Result<(cli::Args, config::Settings), String>>> { +fn run_load_hook() -> &'static RunLoadHook { RUN_LOAD_HOOK.get_or_init(|| std::sync::Mutex::new(None)) } @@ -134,7 +137,7 @@ fn load_args_and_settings() -> Result<(cli::Args, config::Settings)> { if let Some(result) = take_load_hook_result() { return result.map_err(anyhow::Error::msg); } - return Err(anyhow::anyhow!("run loader hook not set")); + Err(anyhow::anyhow!("run loader hook not set")) } #[cfg(not(test))] @@ -322,7 +325,6 @@ async fn maybe_publish_service_presence( } else { info!("Published service presence on startup"); } - return; } #[cfg(not(test))] @@ -526,14 +528,12 @@ mod tests { use radroots_nostr::prelude::RadrootsNostrMetadata; use std::path::Path; use std::path::PathBuf; - use std::sync::{Mutex, MutexGuard}; + use tokio::sync::{Mutex, MutexGuard}; - static TEST_LOCK: Mutex<()> = Mutex::new(()); + static TEST_LOCK: Mutex<()> = Mutex::const_new(()); - fn test_guard() -> MutexGuard<'static, ()> { - let guard = TEST_LOCK - .lock() - .unwrap_or_else(std::sync::PoisonError::into_inner); + async fn test_guard() -> MutexGuard<'static, ()> { + let guard = TEST_LOCK.lock().await; *run_load_hook() .lock() .unwrap_or_else(std::sync::PoisonError::into_inner) = None; @@ -652,7 +652,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_hook_is_missing() { - let _guard = test_guard(); + let _guard = test_guard().await; let err = run().await.expect_err("missing loader hook should error"); let msg = format!("{err:#}"); assert!(msg.contains("run loader hook not set")); @@ -660,7 +660,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_identity_missing() { - let _guard = test_guard(); + let _guard = test_guard().await; let args = args_for_identity(PathBuf::from("/tmp/radrootsd-missing.secret.json"), false); let settings = settings_with_relays(Vec::new()); *run_load_hook() @@ -673,7 +673,7 @@ mod tests { #[tokio::test] async fn run_covers_shutdown_path_and_presence_success() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("shutdown"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(vec!["wss://relay.example.com".to_string()]); @@ -697,7 +697,7 @@ mod tests { #[tokio::test] async fn run_covers_stopped_path_and_presence_failure() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("stopped"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(vec!["wss://relay.example.com".to_string()]); @@ -722,7 +722,7 @@ mod tests { #[tokio::test] async fn run_skips_presence_when_relays_empty() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("empty"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(Vec::new()); @@ -743,7 +743,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_relay_is_invalid() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("invalid-relay"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(vec!["not-a-relay".to_string()]); @@ -758,7 +758,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_rpc_addr_is_invalid() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("invalid-rpc-addr"); let args = args_for_identity(path.clone(), true); let mut settings = settings_with_relays(Vec::new()); @@ -774,7 +774,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_rpc_start_fails() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("rpc-start-fail"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(Vec::new()); @@ -793,7 +793,7 @@ mod tests { #[tokio::test] async fn run_waits_for_stopped_when_wait_hook_is_not_set() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("wait-no-hook"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(Vec::new()); @@ -811,7 +811,7 @@ mod tests { #[tokio::test] async fn run_starts_rpc_when_start_hook_is_not_set() { - let _guard = test_guard(); + let _guard = test_guard().await; let path = unique_identity_path("start-rpc-real"); let args = args_for_identity(path.clone(), true); let settings = settings_with_relays(Vec::new()); diff --git a/src/core/nip46/session.rs b/src/core/nip46/session.rs @@ -201,6 +201,12 @@ impl Nip46SessionStore { } } +impl Default for Nip46SessionStore { + fn default() -> Self { + Self::new() + } +} + impl Nip46Session { pub fn normalize_authority( authority: Option<Nip46SessionAuthority>, diff --git a/src/core/transport_publish.rs b/src/core/transport_publish.rs @@ -33,10 +33,10 @@ use radroots_transport_publish_protocol::{ TransportPublishTargetOutcome, TransportPublishTargetPolicy, TransportPublishTargetPolicyName, TransportPublishTargetSource, }; -use rusqlite::types::Type; -use rusqlite::{Connection, OptionalExtension, Row, params}; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; +use sqlx::sqlite::{SqliteConnectOptions, SqliteConnection, SqliteRow}; +use sqlx::{Connection as _, Row}; use thiserror::Error; use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use uuid::Uuid; @@ -135,7 +135,7 @@ const TRANSPORT_PUBLISH_TABLES: &[&str] = &[ #[derive(Debug, Error)] pub enum TransportPublishError { #[error("transport publish storage error: {0}")] - Sqlite(#[from] rusqlite::Error), + Sqlite(#[from] sqlx::Error), #[error("transport publish json error: {0}")] Json(#[from] serde_json::Error), #[error("transport publish io error: {0}")] @@ -687,7 +687,7 @@ impl TransportPublish { #[derive(Clone)] pub struct TransportPublishStore { - inner: Arc<Mutex<Connection>>, + inner: Arc<Mutex<SqliteConnection>>, } #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] @@ -1015,28 +1015,35 @@ impl TransportPublishStore { { std::fs::create_dir_all(parent)?; } - let connection = Connection::open(path)?; + let connection = connect_sqlite( + SqliteConnectOptions::new() + .filename(path) + .create_if_missing(true), + )?; Self::from_connection(connection) } pub fn memory() -> Result<Self, TransportPublishError> { - Self::from_connection(Connection::open_in_memory()?) + Self::from_connection(connect_sqlite(SqliteConnectOptions::new().in_memory(true))?) } - fn from_connection(connection: Connection) -> Result<Self, TransportPublishError> { - connection.execute_batch("PRAGMA foreign_keys = ON;")?; - match transport_publish_schema_state(&connection)? { + fn from_connection(mut connection: SqliteConnection) -> Result<Self, TransportPublishError> { + execute_sql(&mut connection, "PRAGMA foreign_keys = ON")?; + match transport_publish_schema_state(&mut connection)? { TransportPublishSchemaState::Fresh => { - connection.execute_batch(TRANSPORT_PUBLISH_SCHEMA_SQL)?; - connection.pragma_update(None, "user_version", SCHEMA_VERSION)?; - validate_transport_publish_schema(&connection)?; + execute_raw_sql(&mut connection, TRANSPORT_PUBLISH_SCHEMA_SQL)?; + execute_sql( + &mut connection, + format!("PRAGMA user_version = {SCHEMA_VERSION}").as_str(), + )?; + validate_transport_publish_schema(&mut connection)?; } TransportPublishSchemaState::Existing => { - validate_transport_publish_schema_version(&connection)?; - validate_transport_publish_schema(&connection)?; + validate_transport_publish_schema_version(&mut connection)?; + validate_transport_publish_schema(&mut connection)?; } } - recover_interrupted_publish_jobs(&connection)?; + recover_interrupted_publish_jobs(&mut connection)?; Ok(Self { inner: Arc::new(Mutex::new(connection)), }) @@ -1053,8 +1060,10 @@ impl TransportPublishStore { .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection.execute( - r#" + let mut connection = connection; + block_on_sqlite( + sqlx::query( + r#" INSERT INTO transport_publish_principals ( principal_id, label, @@ -1072,20 +1081,22 @@ impl TransportPublishStore { ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, NULL, ?12) "#, - params![ - principal_id, - input.label.trim(), - input.token_hash, - serde_json::to_string(&input.allowed_pubkeys)?, - serde_json::to_string(&input.allowed_kinds)?, - serde_json::to_string(&input.allowed_target_policies)?, - serde_json::to_string(&input.allowed_explicit_transport_kinds)?, - serde_json::to_string(&input.allowed_nostr_source_policies)?, - input.allow_request_targets, - input.job_visibility.to_string(), - input.expires_at_unix, - now, - ], + ) + .bind(principal_id.as_str()) + .bind(input.label.trim()) + .bind(input.token_hash.as_str()) + .bind(serde_json::to_string(&input.allowed_pubkeys)?) + .bind(serde_json::to_string(&input.allowed_kinds)?) + .bind(serde_json::to_string(&input.allowed_target_policies)?) + .bind(serde_json::to_string( + &input.allowed_explicit_transport_kinds, + )?) + .bind(serde_json::to_string(&input.allowed_nostr_source_policies)?) + .bind(input.allow_request_targets) + .bind(input.job_visibility.to_string()) + .bind(input.expires_at_unix) + .bind(now) + .execute(&mut *connection), )?; drop(connection); self.principal_by_id(principal_id.as_str())?.ok_or_else(|| { @@ -1098,12 +1109,12 @@ impl TransportPublishStore { token_hash: &str, ) -> Result<Option<PublishPrincipal>, TransportPublishError> { let now = current_unix_secs(); - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let principal = connection - .query_row( + let principal = block_on_sqlite( + sqlx::query( r#" SELECT principal_id, @@ -1121,10 +1132,13 @@ impl TransportPublishStore { AND revoked_at_unix IS NULL AND (expires_at_unix IS NULL OR expires_at_unix > ?2) "#, - params![token_hash, now], - principal_from_row, ) - .optional()?; + .bind(token_hash) + .bind(now) + .fetch_optional(&mut *connection), + )? + .map(|row| principal_from_row(&row)) + .transpose()?; Ok(principal) } @@ -1132,12 +1146,12 @@ impl TransportPublishStore { &self, principal_id: &str, ) -> Result<Option<PublishPrincipal>, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let principal = connection - .query_row( + let principal = block_on_sqlite( + sqlx::query( r#" SELECT principal_id, @@ -1153,10 +1167,12 @@ impl TransportPublishStore { FROM transport_publish_principals WHERE principal_id = ?1 "#, - params![principal_id], - principal_from_row, ) - .optional()?; + .bind(principal_id) + .fetch_optional(&mut *connection), + )? + .map(|row| principal_from_row(&row)) + .transpose()?; Ok(principal) } @@ -1169,20 +1185,19 @@ impl TransportPublishStore { "publish job target snapshot count must match effective target count".to_owned(), )); } - if let Some(idempotency_key) = insert.idempotency_key.as_deref() { - if let Some(existing) = + if let Some(idempotency_key) = insert.idempotency_key.as_deref() + && let Some(existing) = self.job_for_principal_id_and_key(insert.principal_id.as_str(), idempotency_key)? - { - if existing.request_fingerprint != insert.request_fingerprint { - return Err(TransportPublishError::IdempotencyConflict( - idempotency_key.to_owned(), - )); - } - return Ok(TransportPublishEventResponse { - deduplicated: true, - job: existing.view, - }); + { + if existing.request_fingerprint != insert.request_fingerprint { + return Err(TransportPublishError::IdempotencyConflict( + idempotency_key.to_owned(), + )); } + return Ok(TransportPublishEventResponse { + deduplicated: true, + job: existing.view, + }); } let job_id = Uuid::new_v4().to_string(); @@ -1198,59 +1213,76 @@ impl TransportPublishStore { .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let transaction = connection.transaction()?; - let insert_result = transaction.execute( - r#" - INSERT INTO transport_publish_jobs ( - job_id, - principal_id, - idempotency_key, - request_fingerprint, - status, - event_id, - event_pubkey, - event_kind, - target_policy_json, - delivery_policy_json, - requested_target_count, - effective_target_count, - request_json, - requested_at_ms, - updated_at_ms - ) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15) - "#, - params![ - job_id, - insert.principal_id, - insert.idempotency_key, - insert.request_fingerprint, - serde_json::to_string(&TransportPublishJobStatus::Publishing)?, - insert.event.event_id, - insert.event.pubkey, - insert.event.kind, - serde_json::to_string(&insert.request.target_policy)?, - serde_json::to_string(&insert.request.delivery_policy)?, - requested_target_count, - effective_target_count, - request_json, - now, + execute_sql(&mut connection, "BEGIN")?; + let transaction_result = (|| -> Result<(), TransportPublishError> { + let insert_result = block_on_sqlite( + sqlx::query( + r#" + INSERT INTO transport_publish_jobs ( + job_id, + principal_id, + idempotency_key, + request_fingerprint, + status, + event_id, + event_pubkey, + event_kind, + target_policy_json, + delivery_policy_json, + requested_target_count, + effective_target_count, + request_json, + requested_at_ms, + updated_at_ms + ) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15) + "#, + ) + .bind(job_id.as_str()) + .bind(insert.principal_id.as_str()) + .bind(insert.idempotency_key.as_deref()) + .bind(insert.request_fingerprint.as_str()) + .bind(serde_json::to_string( + &TransportPublishJobStatus::Publishing, + )?) + .bind(insert.event.event_id.as_str()) + .bind(insert.event.pubkey.as_str()) + .bind(i64::from(insert.event.kind)) + .bind(serde_json::to_string(&insert.request.target_policy)?) + .bind(serde_json::to_string(&insert.request.delivery_policy)?) + .bind(requested_target_count) + .bind(effective_target_count) + .bind(request_json.as_str()) + .bind(now) + .bind(now) + .execute(&mut *connection), + ); + match insert_result { + Ok(_) => {} + Err(error) if is_sqlite_constraint_error(&error) => { + return Err(TransportPublishError::IdempotencyConflict( + "idempotency key conflicts with an existing publish job".to_owned(), + )); + } + Err(error) => return Err(error), + } + insert_target_snapshots( + &mut connection, + job_id.as_str(), + &insert.target_snapshots, now, - ], - ); - match insert_result { - Ok(_) => {} - Err(rusqlite::Error::SqliteFailure(error, _)) - if error.code == rusqlite::ErrorCode::ConstraintViolation => - { - return Err(TransportPublishError::IdempotencyConflict( - "idempotency key conflicts with an existing publish job".to_owned(), - )); + )?; + Ok(()) + })(); + match transaction_result { + Ok(()) => { + execute_sql(&mut connection, "COMMIT")?; + } + Err(error) => { + let _ = execute_sql(&mut connection, "ROLLBACK"); + return Err(error); } - Err(error) => return Err(error.into()), } - insert_target_snapshots(&transaction, job_id.as_str(), &insert.target_snapshots, now)?; - transaction.commit()?; drop(connection); let job = self.job_by_id(job_id.as_str())?; Ok(TransportPublishEventResponse { @@ -1264,14 +1296,18 @@ impl TransportPublishStore { job_id: &str, principal: &PublishPrincipal, ) -> Result<Option<TransportPublishJobView>, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); let sql = job_select_sql("WHERE job_id = ?1"); - let row = connection - .query_row(sql.as_str(), params![job_id], job_from_row) - .optional()?; + let row = block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(job_id) + .fetch_optional(&mut *connection), + )? + .map(|row| job_from_row(&row)) + .transpose()?; drop(connection); let Some(job) = row else { return Ok(None); @@ -1289,7 +1325,7 @@ impl TransportPublishStore { limit: usize, ) -> Result<Vec<TransportPublishJobView>, TransportPublishError> { let limit = i64::try_from(limit.clamp(1, 200)).unwrap_or(200); - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); @@ -1300,15 +1336,24 @@ impl TransportPublishStore { "WHERE principal_id = ?1 ORDER BY requested_at_ms DESC, job_id DESC LIMIT ?2", ) }; - let mut stmt = connection.prepare(sql.as_str())?; let rows = if principal.job_visibility == PublishJobVisibility::Admin { - stmt.query_map(params![limit], job_from_row)? - .collect::<Result<Vec<_>, _>>()? + block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(limit) + .fetch_all(&mut *connection), + )? } else { - stmt.query_map(params![principal.principal_id, limit], job_from_row)? - .collect::<Result<Vec<_>, _>>()? + block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(principal.principal_id.as_str()) + .bind(limit) + .fetch_all(&mut *connection), + )? }; - drop(stmt); + let rows = rows + .iter() + .map(job_from_row) + .collect::<Result<Vec<_>, _>>()?; drop(connection); rows.into_iter() @@ -1324,18 +1369,19 @@ impl TransportPublishStore { principal_id: &str, idempotency_key: &str, ) -> Result<Option<PublishJobRow>, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); let sql = job_select_sql("WHERE principal_id = ?1 AND idempotency_key = ?2"); - let row = connection - .query_row( - sql.as_str(), - params![principal_id, idempotency_key], - job_from_row, - ) - .optional()?; + let row = block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(principal_id) + .bind(idempotency_key) + .fetch_optional(&mut *connection), + )? + .map(|row| job_from_row(&row)) + .transpose()?; drop(connection); let Some(job) = row else { return Ok(None); @@ -1348,14 +1394,18 @@ impl TransportPublishStore { &self, job_id: &str, ) -> Result<TransportPublishJobView, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); let sql = job_select_sql("WHERE job_id = ?1"); - let row = connection - .query_row(sql.as_str(), params![job_id], job_from_row) - .optional()?; + let row = block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(job_id) + .fetch_optional(&mut *connection), + )? + .map(|row| job_from_row(&row)) + .transpose()?; drop(connection); let Some(job) = row else { return Err(TransportPublishError::InvalidScope( @@ -1375,12 +1425,13 @@ impl TransportPublishStore { ) -> Result<(), TransportPublishError> { let now = current_unix_millis(); let target_count = storage_count_i64(outcomes.len(), "effective_target_count")?; - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" UPDATE transport_publish_jobs SET status = ?2, updated_at_ms = ?3, @@ -1389,16 +1440,16 @@ impl TransportPublishStore { effective_target_count = ?6 WHERE job_id = ?1 "#, - params![ - job_id, - serde_json::to_string(&status)?, - now, - now, - last_error, - target_count, - ], + ) + .bind(job_id) + .bind(serde_json::to_string(&status)?) + .bind(now) + .bind(now) + .bind(last_error.as_deref()) + .bind(target_count) + .execute(&mut *connection), )?; - replace_target_outcomes(&connection, job_id, &outcomes, now)?; + replace_target_outcomes(&mut connection, job_id, &outcomes, now)?; Ok(()) } @@ -1406,17 +1457,17 @@ impl TransportPublishStore { &self, pubkey: &str, ) -> Result<Vec<String>, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let relays_json = connection - .query_row( + let relays_json = block_on_sqlite( + sqlx::query_scalar::<_, String>( "SELECT relays_json FROM transport_publish_nostr_author_cache WHERE pubkey = ?1", - params![pubkey], - |row| row.get::<_, String>(0), ) - .optional()?; + .bind(pubkey) + .fetch_optional(&mut *connection), + )?; relays_json .map(|value| serde_json::from_str(value.as_str()).map_err(TransportPublishError::from)) .unwrap_or_else(|| Ok(Vec::new())) @@ -1428,19 +1479,24 @@ impl TransportPublishStore { relays: &[String], ) -> Result<(), TransportPublishError> { let now = current_unix_millis(); - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" INSERT INTO transport_publish_nostr_author_cache (pubkey, relays_json, updated_at_ms) VALUES (?1, ?2, ?3) ON CONFLICT(pubkey) DO UPDATE SET relays_json = excluded.relays_json, updated_at_ms = excluded.updated_at_ms "#, - params![pubkey, serde_json::to_string(relays)?, now], + ) + .bind(pubkey) + .bind(serde_json::to_string(relays)?) + .bind(now) + .execute(&mut *connection), )?; Ok(()) } @@ -1449,20 +1505,25 @@ impl TransportPublishStore { &self, job_id: &str, ) -> Result<Vec<TransportPublishTargetOutcome>, TransportPublishError> { - let connection = self + let mut connection = self .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let mut stmt = connection.prepare( + let rows = block_on_sqlite( + sqlx::query( r#" SELECT transport_kind, endpoint_uri, target_scope, target_label, source, attempted, outcome_kind, message, latency_ms FROM transport_publish_target_results WHERE job_id = ?1 ORDER BY transport_kind, endpoint_uri, target_scope "#, + ) + .bind(job_id) + .fetch_all(&mut *connection), )?; - let outcomes = stmt - .query_map(params![job_id], target_outcome_from_row)? + let outcomes = rows + .iter() + .map(target_outcome_from_row) .collect::<Result<Vec<_>, _>>()?; Ok(outcomes) } @@ -1491,14 +1552,62 @@ enum TransportPublishSchemaState { Existing, } +fn connect_sqlite( + options: SqliteConnectOptions, +) -> Result<SqliteConnection, TransportPublishError> { + block_on_sqlite(SqliteConnection::connect_with(&options)) +} + +fn block_on_sqlite<T>( + future: impl Future<Output = Result<T, sqlx::Error>>, +) -> Result<T, TransportPublishError> { + Ok(futures_executor::block_on(future)?) +} + +fn execute_sql(connection: &mut SqliteConnection, sql: &str) -> Result<u64, TransportPublishError> { + Ok(block_on_sqlite(sqlx::query(sqlx::AssertSqlSafe(sql)).execute(connection))?.rows_affected()) +} + +fn execute_raw_sql( + connection: &mut SqliteConnection, + sql: &str, +) -> Result<(), TransportPublishError> { + block_on_sqlite(sqlx::raw_sql(sqlx::AssertSqlSafe(sql)).execute(connection))?; + Ok(()) +} + +fn fetch_all_sql( + connection: &mut SqliteConnection, + sql: &str, +) -> Result<Vec<SqliteRow>, TransportPublishError> { + block_on_sqlite(sqlx::query(sqlx::AssertSqlSafe(sql)).fetch_all(connection)) +} + +fn is_sqlite_constraint_error(error: &TransportPublishError) -> bool { + match error { + TransportPublishError::Sqlite(sqlx::Error::Database(error)) => { + error.is_unique_violation() + || error + .code() + .as_deref() + .is_some_and(|code| matches!(code, "1555" | "2067" | "19")) + } + _ => false, + } +} + fn transport_publish_schema_state( - connection: &Connection, + connection: &mut SqliteConnection, ) -> Result<TransportPublishSchemaState, TransportPublishError> { - let names = connection - .prepare( + let rows = block_on_sqlite( + sqlx::query( "SELECT name FROM sqlite_schema WHERE type = 'table' AND name LIKE 'transport_publish_%'", - )? - .query_map([], |row| row.get::<_, String>(0))? + ) + .fetch_all(&mut *connection), + )?; + let names = rows + .iter() + .map(|row| row.try_get::<String, _>(0)) .collect::<Result<BTreeSet<_>, _>>()?; if names.is_empty() { let version = transport_publish_schema_version(connection)?; @@ -1517,12 +1626,14 @@ fn transport_publish_schema_state( } } -fn transport_publish_schema_version(connection: &Connection) -> Result<i64, TransportPublishError> { - Ok(connection.pragma_query_value(None, "user_version", |row| row.get::<_, i64>(0))?) +fn transport_publish_schema_version( + connection: &mut SqliteConnection, +) -> Result<i64, TransportPublishError> { + block_on_sqlite(sqlx::query_scalar::<_, i64>("PRAGMA user_version").fetch_one(connection)) } fn validate_transport_publish_schema_version( - connection: &Connection, + connection: &mut SqliteConnection, ) -> Result<(), TransportPublishError> { let version = transport_publish_schema_version(connection)?; if version == SCHEMA_VERSION { @@ -1535,7 +1646,9 @@ fn validate_transport_publish_schema_version( } } -fn validate_transport_publish_schema(connection: &Connection) -> Result<(), TransportPublishError> { +fn validate_transport_publish_schema( + connection: &mut SqliteConnection, +) -> Result<(), TransportPublishError> { validate_foreign_keys_enabled(connection)?; validate_table_columns( connection, @@ -1736,9 +1849,11 @@ struct ForeignKeyEntry { to_columns: Vec<String>, } -fn validate_foreign_keys_enabled(connection: &Connection) -> Result<(), TransportPublishError> { +fn validate_foreign_keys_enabled( + connection: &mut SqliteConnection, +) -> Result<(), TransportPublishError> { let enabled = - connection.pragma_query_value(None, "foreign_keys", |row| row.get::<_, i64>(0))?; + block_on_sqlite(sqlx::query_scalar::<_, i64>("PRAGMA foreign_keys").fetch_one(connection))?; if enabled == 1 { Ok(()) } else { @@ -1750,17 +1865,17 @@ fn validate_foreign_keys_enabled(connection: &Connection) -> Result<(), Transpor } fn validate_table_present( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, ) -> Result<(), TransportPublishError> { - let exists = connection - .query_row( + let exists = block_on_sqlite( + sqlx::query_scalar::<_, i64>( "SELECT 1 FROM sqlite_schema WHERE type = 'table' AND name = ?1", - params![table], - |_| Ok(()), ) - .optional()? - .is_some(); + .bind(table) + .fetch_optional(connection), + )? + .is_some(); if exists { Ok(()) } else { @@ -1772,26 +1887,28 @@ fn validate_table_present( } fn table_columns( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, ) -> Result<BTreeMap<String, TableColumnInfo>, TransportPublishError> { - Ok(connection - .prepare(format!("PRAGMA table_info({table})").as_str())? - .query_map([], |row| { + let sql = format!("PRAGMA table_info({table})"); + let rows = fetch_all_sql(connection, sql.as_str())?; + Ok(rows + .iter() + .map(|row| { Ok(( - row.get::<_, String>(1)?, + row.try_get::<String, _>(1)?, TableColumnInfo { - column_type: row.get::<_, String>(2)?.to_ascii_uppercase(), - not_null: row.get::<_, i64>(3)? != 0, - primary_key_position: row.get::<_, i64>(5)?, + column_type: row.try_get::<String, _>(2)?.to_ascii_uppercase(), + not_null: row.try_get::<i64, _>(3)? != 0, + primary_key_position: row.try_get::<i64, _>(5)?, }, )) - })? - .collect::<Result<BTreeMap<_, _>, _>>()?) + }) + .collect::<Result<BTreeMap<_, _>, sqlx::Error>>()?) } fn validate_primary_key( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, expected_columns: &[&'static str], ) -> Result<(), TransportPublishError> { @@ -1819,21 +1936,23 @@ fn validate_primary_key( } fn validate_unique_index( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, expected_columns: &[&'static str], partial_where: Option<&'static str>, ) -> Result<(), TransportPublishError> { - let indexes = connection - .prepare(format!("PRAGMA index_list({table})").as_str())? - .query_map([], |row| { + let sql = format!("PRAGMA index_list({table})"); + let rows = fetch_all_sql(connection, sql.as_str())?; + let indexes = rows + .iter() + .map(|row| { Ok(( - row.get::<_, String>(1)?, - row.get::<_, i64>(2)? != 0, - row.get::<_, i64>(4)? != 0, + row.try_get::<String, _>(1)?, + row.try_get::<i64, _>(2)? != 0, + row.try_get::<i64, _>(4)? != 0, )) - })? - .collect::<Result<Vec<_>, _>>()?; + }) + .collect::<Result<Vec<_>, sqlx::Error>>()?; for (index_name, unique, partial) in indexes { if !unique { continue; @@ -1865,34 +1984,35 @@ fn validate_unique_index( } fn index_columns( - connection: &Connection, + connection: &mut SqliteConnection, index_name: &str, ) -> Result<Vec<String>, TransportPublishError> { - let mut columns = connection - .prepare(format!("PRAGMA index_info({index_name})").as_str())? - .query_map([], |row| { - Ok((row.get::<_, i64>(0)?, row.get::<_, String>(2)?)) - })? - .collect::<Result<Vec<_>, _>>()?; + let sql = format!("PRAGMA index_info({index_name})"); + let rows = fetch_all_sql(connection, sql.as_str())?; + let mut columns = rows + .iter() + .map(|row| Ok((row.try_get::<i64, _>(0)?, row.try_get::<String, _>(2)?))) + .collect::<Result<Vec<_>, sqlx::Error>>()?; columns.sort_by_key(|(position, _)| *position); Ok(columns.into_iter().map(|(_, name)| name).collect()) } fn index_sql_contains_where( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, index_name: &str, required_where: &str, ) -> Result<bool, TransportPublishError> { - let sql = connection - .query_row( + let sql = block_on_sqlite( + sqlx::query_scalar::<_, Option<String>>( "SELECT sql FROM sqlite_schema WHERE type = 'index' AND tbl_name = ?1 AND name = ?2", - params![table, index_name], - |row| row.get::<_, Option<String>>(0), ) - .optional()? - .flatten() - .unwrap_or_default(); + .bind(table) + .bind(index_name) + .fetch_optional(connection), + )? + .flatten() + .unwrap_or_default(); Ok(normalized_sql(sql.as_str()).contains(normalized_sql(required_where).as_str())) } @@ -1904,7 +2024,7 @@ fn normalized_sql(sql: &str) -> String { } fn validate_foreign_key( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, expected_from_columns: &[&'static str], expected_target_table: &'static str, @@ -1930,22 +2050,24 @@ fn validate_foreign_key( } fn foreign_keys( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, ) -> Result<Vec<ForeignKeyEntry>, TransportPublishError> { let mut groups = BTreeMap::<i64, (String, Vec<(i64, String, String)>)>::new(); - let rows = connection - .prepare(format!("PRAGMA foreign_key_list({table})").as_str())? - .query_map([], |row| { + let sql = format!("PRAGMA foreign_key_list({table})"); + let rows = fetch_all_sql(connection, sql.as_str())?; + let rows = rows + .iter() + .map(|row| { Ok(( - row.get::<_, i64>(0)?, - row.get::<_, i64>(1)?, - row.get::<_, String>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, + row.try_get::<i64, _>(0)?, + row.try_get::<i64, _>(1)?, + row.try_get::<String, _>(2)?, + row.try_get::<String, _>(3)?, + row.try_get::<String, _>(4)?, )) - })? - .collect::<Result<Vec<_>, _>>()?; + }) + .collect::<Result<Vec<_>, sqlx::Error>>()?; for (id, seq, target_table, from_column, to_column) in rows { let entry = groups .entry(id) @@ -1980,7 +2102,7 @@ fn columns_match(actual: &[String], expected: &[&'static str]) -> bool { } fn validate_table_columns( - connection: &Connection, + connection: &mut SqliteConnection, table: &'static str, required_columns: &[RequiredColumn], ) -> Result<(), TransportPublishError> { @@ -2040,20 +2162,28 @@ fn validate_table_columns( Ok(()) } -fn recover_interrupted_publish_jobs(connection: &Connection) -> Result<(), TransportPublishError> { +fn recover_interrupted_publish_jobs( + connection: &mut SqliteConnection, +) -> Result<(), TransportPublishError> { let now = current_unix_millis(); let publishing = serde_json::to_string(&TransportPublishJobStatus::Publishing)?; let sql = job_select_sql("WHERE status = ?1"); - let rows = connection - .prepare(sql.as_str())? - .query_map(params![publishing], job_from_row)? + let rows = block_on_sqlite( + sqlx::query(sqlx::AssertSqlSafe(sql.as_str())) + .bind(publishing.as_str()) + .fetch_all(&mut *connection), + )?; + let rows = rows + .iter() + .map(job_from_row) .collect::<Result<Vec<_>, _>>()?; for row in rows { let job_id = row.view.job_id.clone(); let snapshots = target_snapshot_outcomes(connection, job_id.as_str())?; if snapshots.is_empty() { - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" UPDATE transport_publish_jobs SET status = ?2, updated_at_ms = ?3, @@ -2062,13 +2192,13 @@ fn recover_interrupted_publish_jobs(connection: &Connection) -> Result<(), Trans effective_target_count = 0 WHERE job_id = ?1 "#, - params![ - job_id.as_str(), - serde_json::to_string(&TransportPublishJobStatus::Rejected)?, - now, - now, - "publish_attempt_interrupted_missing_target_snapshot", - ], + ) + .bind(job_id.as_str()) + .bind(serde_json::to_string(&TransportPublishJobStatus::Rejected)?) + .bind(now) + .bind(now) + .bind("publish_attempt_interrupted_missing_target_snapshot") + .execute(&mut *connection), )?; replace_target_outcomes(connection, job_id.as_str(), &[], now)?; continue; @@ -2080,8 +2210,9 @@ fn recover_interrupted_publish_jobs(connection: &Connection) -> Result<(), Trans } else { last_error_for_status(status).map(str::to_owned) }; - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" UPDATE transport_publish_jobs SET status = ?2, updated_at_ms = ?3, @@ -2090,14 +2221,14 @@ fn recover_interrupted_publish_jobs(connection: &Connection) -> Result<(), Trans effective_target_count = ?6 WHERE job_id = ?1 "#, - params![ - job_id.as_str(), - serde_json::to_string(&status)?, - now, - now, - last_error, - effective_target_count, - ], + ) + .bind(job_id.as_str()) + .bind(serde_json::to_string(&status)?) + .bind(now) + .bind(now) + .bind(last_error.as_deref()) + .bind(effective_target_count) + .execute(&mut *connection), )?; replace_target_outcomes(connection, job_id.as_str(), &snapshots, now)?; } @@ -2105,14 +2236,15 @@ fn recover_interrupted_publish_jobs(connection: &Connection) -> Result<(), Trans } fn insert_target_snapshots( - connection: &Connection, + connection: &mut SqliteConnection, job_id: &str, outcomes: &[TransportPublishTargetOutcome], now: i64, ) -> Result<(), TransportPublishError> { for (target_index, outcome) in outcomes.iter().enumerate() { - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" INSERT INTO transport_publish_target_snapshots ( job_id, target_index, @@ -2129,40 +2261,44 @@ fn insert_target_snapshots( ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12) "#, - params![ - job_id, - i64::try_from(target_index).unwrap_or(i64::MAX), - outcome.transport_kind, - outcome.endpoint_uri, - storage_target_scope(outcome.target_scope.as_deref()), - outcome.target_label, - serde_json::to_string(&outcome.source)?, - outcome.attempted, - serde_json::to_string(&outcome.outcome_kind)?, - outcome.message, + ) + .bind(job_id) + .bind(i64::try_from(target_index).unwrap_or(i64::MAX)) + .bind(outcome.transport_kind.as_str()) + .bind(outcome.endpoint_uri.as_str()) + .bind(storage_target_scope(outcome.target_scope.as_deref())) + .bind(outcome.target_label.as_deref()) + .bind(serde_json::to_string(&outcome.source)?) + .bind(outcome.attempted) + .bind(serde_json::to_string(&outcome.outcome_kind)?) + .bind(outcome.message.as_deref()) + .bind( outcome .latency_ms .and_then(|value| i64::try_from(value).ok()), - now, - ], + ) + .bind(now) + .execute(&mut *connection), )?; } Ok(()) } fn replace_target_outcomes( - connection: &Connection, + connection: &mut SqliteConnection, job_id: &str, outcomes: &[TransportPublishTargetOutcome], now: i64, ) -> Result<(), TransportPublishError> { - connection.execute( - "DELETE FROM transport_publish_target_results WHERE job_id = ?1", - params![job_id], + block_on_sqlite( + sqlx::query("DELETE FROM transport_publish_target_results WHERE job_id = ?1") + .bind(job_id) + .execute(&mut *connection), )?; for outcome in outcomes { - connection.execute( - r#" + block_on_sqlite( + sqlx::query( + r#" INSERT OR REPLACE INTO transport_publish_target_results ( job_id, transport_kind, @@ -2178,40 +2314,47 @@ fn replace_target_outcomes( ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11) "#, - params![ - job_id, - outcome.transport_kind, - outcome.endpoint_uri, - storage_target_scope(outcome.target_scope.as_deref()), - outcome.target_label, - serde_json::to_string(&outcome.source)?, - outcome.attempted, - serde_json::to_string(&outcome.outcome_kind)?, - outcome.message, + ) + .bind(job_id) + .bind(outcome.transport_kind.as_str()) + .bind(outcome.endpoint_uri.as_str()) + .bind(storage_target_scope(outcome.target_scope.as_deref())) + .bind(outcome.target_label.as_deref()) + .bind(serde_json::to_string(&outcome.source)?) + .bind(outcome.attempted) + .bind(serde_json::to_string(&outcome.outcome_kind)?) + .bind(outcome.message.as_deref()) + .bind( outcome .latency_ms .and_then(|value| i64::try_from(value).ok()), - now, - ], + ) + .bind(now) + .execute(&mut *connection), )?; } Ok(()) } fn target_snapshot_outcomes( - connection: &Connection, + connection: &mut SqliteConnection, job_id: &str, ) -> Result<Vec<TransportPublishTargetOutcome>, TransportPublishError> { - let mut stmt = connection.prepare( + let rows = block_on_sqlite( + sqlx::query( r#" SELECT transport_kind, endpoint_uri, target_scope, target_label, source, attempted, outcome_kind, message, latency_ms FROM transport_publish_target_snapshots WHERE job_id = ?1 ORDER BY target_index "#, + ) + .bind(job_id) + .fetch_all(connection), )?; - let outcomes = stmt - .query_map(params![job_id], target_outcome_from_row)? + let outcomes = rows + .iter() + .map(target_outcome_from_row) .collect::<Result<Vec<_>, _>>()?; Ok(outcomes) } @@ -2239,37 +2382,41 @@ fn job_select_sql(tail: &str) -> String { ) } -fn principal_from_row(row: &Row<'_>) -> Result<PublishPrincipal, rusqlite::Error> { - let visibility: String = row.get(8)?; +fn principal_from_row(row: &SqliteRow) -> Result<PublishPrincipal, TransportPublishError> { + let visibility: String = row.try_get(8)?; Ok(PublishPrincipal { - principal_id: row.get(0)?, - label: row.get(1)?, - allowed_pubkeys: json_column(row, 2)?, - allowed_kinds: json_column(row, 3)?, - allowed_target_policies: json_column(row, 4)?, - allowed_explicit_transport_kinds: json_column(row, 5)?, - allowed_nostr_source_policies: json_column(row, 6)?, - allow_request_targets: row.get(7)?, - job_visibility: PublishJobVisibility::from_str(visibility.as_str()) - .map_err(|error| conversion_error(8, error))?, - expires_at_unix: row.get(9)?, + principal_id: row.try_get(0)?, + label: row.try_get(1)?, + allowed_pubkeys: json_column(row, 2, "allowed_pubkeys_json")?, + allowed_kinds: json_column(row, 3, "allowed_kinds_json")?, + allowed_target_policies: json_column(row, 4, "allowed_target_policies_json")?, + allowed_explicit_transport_kinds: json_column( + row, + 5, + "allowed_explicit_transport_kinds_json", + )?, + allowed_nostr_source_policies: json_column(row, 6, "allowed_nostr_source_policies_json")?, + allow_request_targets: row.try_get(7)?, + job_visibility: PublishJobVisibility::from_str(visibility.as_str())?, + expires_at_unix: row.try_get(9)?, }) } -fn job_from_row(row: &Row<'_>) -> Result<PublishJobRow, rusqlite::Error> { - let status: TransportPublishJobStatus = json_text(row, 3)?; - let target_policy: TransportPublishTargetPolicy = json_text(row, 7)?; - let delivery_policy: TransportPublishDeliveryPolicy = json_text(row, 8)?; +fn job_from_row(row: &SqliteRow) -> Result<PublishJobRow, TransportPublishError> { + let status: TransportPublishJobStatus = json_text(row, 3, "status")?; + let target_policy: TransportPublishTargetPolicy = json_text(row, 7, "target_policy_json")?; + let delivery_policy: TransportPublishDeliveryPolicy = + json_text(row, 8, "delivery_policy_json")?; Ok(PublishJobRow { - principal_id: row.get(1)?, - request_fingerprint: row.get(2)?, + principal_id: row.try_get(1)?, + request_fingerprint: row.try_get(2)?, view: TransportPublishJobView { - job_id: row.get(0)?, + job_id: row.try_get(0)?, status, terminal: false, delivery_satisfied: false, - event_id: row.get(4)?, - pubkey: row.get(5)?, + event_id: row.try_get(4)?, + pubkey: row.try_get(5)?, event_kind: checked_event_kind_column(row, 6)?, target_policy, delivery_policy, @@ -2277,28 +2424,28 @@ fn job_from_row(row: &Row<'_>) -> Result<PublishJobRow, rusqlite::Error> { acknowledged_count: 0, retryable_count: 0, terminal_count: 0, - requested_at_ms: row.get(10)?, - completed_at_ms: row.get(11)?, - last_error: row.get(12)?, + requested_at_ms: row.try_get(10)?, + completed_at_ms: row.try_get(11)?, + last_error: row.try_get(12)?, targets: Vec::new(), }, }) } fn target_outcome_from_row( - row: &Row<'_>, -) -> Result<TransportPublishTargetOutcome, rusqlite::Error> { - let source: TransportPublishTargetSource = json_text(row, 4)?; - let outcome_kind: TransportPublishOutcomeKind = json_text(row, 6)?; + row: &SqliteRow, +) -> Result<TransportPublishTargetOutcome, TransportPublishError> { + let source: TransportPublishTargetSource = json_text(row, 4, "source")?; + let outcome_kind: TransportPublishOutcomeKind = json_text(row, 6, "outcome_kind")?; Ok(TransportPublishTargetOutcome { - transport_kind: row.get(0)?, - endpoint_uri: row.get(1)?, - target_scope: storage_target_scope_to_protocol(row.get::<_, String>(2)?), - target_label: row.get(3)?, + transport_kind: row.try_get(0)?, + endpoint_uri: row.try_get(1)?, + target_scope: storage_target_scope_to_protocol(row.try_get::<String, _>(2)?), + target_label: row.try_get(3)?, source, - attempted: row.get(5)?, + attempted: row.try_get(5)?, outcome_kind, - message: row.get(7)?, + message: row.try_get(7)?, latency_ms: checked_optional_u64_column(row, 8, "latency_ms")?, }) } @@ -3001,19 +3148,21 @@ fn ensure_lower_hex( } fn json_column<T: for<'de> Deserialize<'de>>( - row: &Row<'_>, + row: &SqliteRow, index: usize, -) -> Result<T, rusqlite::Error> { - let value: String = row.get(index)?; - serde_json::from_str(value.as_str()).map_err(|error| conversion_error(index, error)) + field: &'static str, +) -> Result<T, TransportPublishError> { + let value: String = row.try_get(index)?; + serde_json::from_str(value.as_str()).map_err(|error| persisted_decode_error(field, error)) } fn json_text<T: for<'de> Deserialize<'de>>( - row: &Row<'_>, + row: &SqliteRow, index: usize, -) -> Result<T, rusqlite::Error> { - let value: String = row.get(index)?; - serde_json::from_str(value.as_str()).map_err(|error| conversion_error(index, error)) + field: &'static str, +) -> Result<T, TransportPublishError> { + let value: String = row.try_get(index)?; + serde_json::from_str(value.as_str()).map_err(|error| persisted_decode_error(field, error)) } #[derive(Debug, Error)] @@ -3024,11 +3173,11 @@ struct TransportPublishStorageIntegerRangeError { target: &'static str, } -fn checked_event_kind_column(row: &Row<'_>, index: usize) -> Result<u32, rusqlite::Error> { - let value = row.get::<_, i64>(index)?; +fn checked_event_kind_column(row: &SqliteRow, index: usize) -> Result<u32, TransportPublishError> { + let value = row.try_get::<i64, _>(index)?; if !(0..=i64::from(u32::MAX)).contains(&value) { - return Err(integer_conversion_error( - index, + return Err(persisted_decode_error( + "event_kind", TransportPublishStorageIntegerRangeError { field: "event_kind", value, @@ -3037,8 +3186,8 @@ fn checked_event_kind_column(row: &Row<'_>, index: usize) -> Result<u32, rusqlit )); } u32::try_from(value).map_err(|_| { - integer_conversion_error( - index, + persisted_decode_error( + "event_kind", TransportPublishStorageIntegerRangeError { field: "event_kind", value, @@ -3049,14 +3198,14 @@ fn checked_event_kind_column(row: &Row<'_>, index: usize) -> Result<u32, rusqlit } fn checked_usize_column( - row: &Row<'_>, + row: &SqliteRow, index: usize, field: &'static str, -) -> Result<usize, rusqlite::Error> { - let value = row.get::<_, i64>(index)?; +) -> Result<usize, TransportPublishError> { + let value = row.try_get::<i64, _>(index)?; usize::try_from(value).map_err(|_| { - integer_conversion_error( - index, + persisted_decode_error( + field, TransportPublishStorageIntegerRangeError { field, value, @@ -3075,15 +3224,15 @@ fn storage_count_i64(value: usize, field: &'static str) -> Result<i64, Transport } fn checked_optional_u64_column( - row: &Row<'_>, + row: &SqliteRow, index: usize, field: &'static str, -) -> Result<Option<u64>, rusqlite::Error> { - row.get::<_, Option<i64>>(index)? +) -> Result<Option<u64>, TransportPublishError> { + row.try_get::<Option<i64>, _>(index)? .map(|value| { u64::try_from(value).map_err(|_| { - integer_conversion_error( - index, + persisted_decode_error( + field, TransportPublishStorageIntegerRangeError { field, value, @@ -3103,18 +3252,13 @@ fn storage_target_scope_to_protocol(target_scope: String) -> Option<String> { (!target_scope.is_empty()).then_some(target_scope) } -fn conversion_error<E>(index: usize, error: E) -> rusqlite::Error +fn persisted_decode_error<E>(field: &'static str, error: E) -> TransportPublishError where - E: std::error::Error + Send + Sync + 'static, + E: std::error::Error, { - rusqlite::Error::FromSqlConversionFailure(index, Type::Text, Box::new(error)) -} - -fn integer_conversion_error<E>(index: usize, error: E) -> rusqlite::Error -where - E: std::error::Error + Send + Sync + 'static, -{ - rusqlite::Error::FromSqlConversionFailure(index, Type::Integer, Box::new(error)) + TransportPublishError::InvalidPublishJobState(format!( + "{field} persisted value could not be decoded: {error}" + )) } fn current_unix_secs() -> i64 { @@ -3154,6 +3298,8 @@ mod tests { TransportPublishTargetPolicy, TransportPublishTargetPolicyName, TransportPublishTargetSource, }; + use sqlx::Row; + use sqlx::sqlite::{SqliteConnectOptions, SqliteConnection}; use std::collections::BTreeMap; use std::net::{IpAddr, Ipv4Addr}; use std::sync::Arc; @@ -3203,15 +3349,26 @@ mod tests { schema_sql: &str, user_version: Option<i64>, ) { - let connection = rusqlite::Connection::open(database_path).expect("open schema database"); - connection.execute_batch(schema_sql).expect("create schema"); + let mut connection = open_test_database(database_path); + super::execute_raw_sql(&mut connection, schema_sql).expect("create schema"); if let Some(version) = user_version { - connection - .pragma_update(None, "user_version", version) - .expect("set schema version"); + super::execute_sql( + &mut connection, + format!("PRAGMA user_version = {version}").as_str(), + ) + .expect("set schema version"); } } + fn open_test_database(database_path: &std::path::Path) -> SqliteConnection { + super::connect_sqlite( + SqliteConnectOptions::new() + .filename(database_path) + .create_if_missing(true), + ) + .expect("open schema database") + } + fn open_schema_error(database_path: &std::path::Path) -> TransportPublishError { match TransportPublishStore::open(database_path.to_path_buf()) { Ok(_) => panic!("malformed schema opened"), @@ -3237,10 +3394,17 @@ mod tests { } fn database_user_version(database_path: &std::path::Path) -> i64 { - let connection = rusqlite::Connection::open(database_path).expect("open schema database"); - connection - .pragma_query_value(None, "user_version", |row| row.get::<_, i64>(0)) - .expect("user version") + let mut connection = open_test_database(database_path); + super::transport_publish_schema_version(&mut connection).expect("user version") + } + + fn test_query_column_names(connection: &mut SqliteConnection, table: &str) -> Vec<String> { + let sql = format!("PRAGMA table_info({table})"); + super::fetch_all_sql(connection, sql.as_str()) + .expect("query schema") + .iter() + .map(|row| row.try_get::<String, _>(1).expect("column name")) + .collect() } fn signed_event(identity: &RadrootsIdentity, content: &str) -> String { @@ -3374,12 +3538,7 @@ mod tests { fn assert_storage_integer_range_error(error: TransportPublishError, expected: &str) { match error { - TransportPublishError::Sqlite(rusqlite::Error::FromSqlConversionFailure( - _, - rusqlite::types::Type::Integer, - source, - )) => { - let message = source.to_string(); + TransportPublishError::InvalidPublishJobState(message) => { assert!(message.contains(expected), "{message}"); assert!(!message.contains("rrd_tp_")); assert!(!message.contains("token")); @@ -3746,16 +3905,18 @@ mod tests { ) .expect("complete job"); { - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection - .execute( + super::block_on_sqlite( + sqlx::query( "UPDATE transport_publish_jobs SET effective_target_count = 2 WHERE job_id = ?1", - rusqlite::params![response.job.job_id.as_str()], ) - .expect("corrupt target count"); + .bind(response.job.job_id.as_str()) + .execute(&mut *connection), + ) + .expect("corrupt target count"); } assert_invalid_job_state( @@ -3822,16 +3983,19 @@ mod tests { ) .expect("complete job"); { - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection - .execute( + super::block_on_sqlite( + sqlx::query( "UPDATE transport_publish_jobs SET event_kind = ?2 WHERE job_id = ?1", - rusqlite::params![response.job.job_id.as_str(), invalid_kind], ) - .expect("corrupt event kind"); + .bind(response.job.job_id.as_str()) + .bind(invalid_kind) + .execute(&mut *connection), + ) + .expect("corrupt event kind"); } assert_storage_integer_range_error( @@ -3898,16 +4062,18 @@ mod tests { ) .expect("complete job"); { - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection - .execute( + super::block_on_sqlite( + sqlx::query( "UPDATE transport_publish_jobs SET effective_target_count = -1 WHERE job_id = ?1", - rusqlite::params![response.job.job_id.as_str()], ) - .expect("corrupt target count"); + .bind(response.job.job_id.as_str()) + .execute(&mut *connection), + ) + .expect("corrupt target count"); } assert_storage_integer_range_error( @@ -3973,16 +4139,18 @@ mod tests { ) .expect("complete job"); { - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection - .execute( + super::block_on_sqlite( + sqlx::query( "UPDATE transport_publish_target_results SET latency_ms = -5 WHERE job_id = ?1", - rusqlite::params![response.job.job_id.as_str()], ) - .expect("corrupt latency"); + .bind(response.job.job_id.as_str()) + .execute(&mut *connection), + ) + .expect("corrupt latency"); } assert_storage_integer_range_error( @@ -4310,12 +4478,12 @@ mod tests { }) .expect("principal"); let now = super::current_unix_millis(); - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - connection - .execute( + super::block_on_sqlite( + sqlx::query( r#" INSERT INTO transport_publish_jobs ( job_id, @@ -4336,30 +4504,33 @@ mod tests { ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15) "#, - rusqlite::params![ - "job-missing-snapshot", - principal.principal_id, - "idem-missing-snapshot", - "fingerprint-missing-snapshot", - serde_json::to_string(&TransportPublishJobStatus::Publishing) - .expect("status"), - "0".repeat(64), - pubkey, - 30_402_i64, - serde_json::to_string(&request.target_policy).expect("target policy"), - serde_json::to_string(&request.delivery_policy).expect("delivery policy"), - super::storage_count_i64( - request.target_policy.request_target_count(), - "requested_target_count", - ) - .expect("requested target count"), - 1_i64, - serde_json::to_string(&request).expect("request"), - now, - now, - ], ) - .expect("insert historical job"); + .bind("job-missing-snapshot") + .bind(principal.principal_id.as_str()) + .bind("idem-missing-snapshot") + .bind("fingerprint-missing-snapshot") + .bind( + serde_json::to_string(&TransportPublishJobStatus::Publishing).expect("status"), + ) + .bind("0".repeat(64)) + .bind(pubkey.as_str()) + .bind(30_402_i64) + .bind(serde_json::to_string(&request.target_policy).expect("target policy")) + .bind(serde_json::to_string(&request.delivery_policy).expect("delivery policy")) + .bind( + super::storage_count_i64( + request.target_policy.request_target_count(), + "requested_target_count", + ) + .expect("requested target count"), + ) + .bind(1_i64) + .bind(serde_json::to_string(&request).expect("request")) + .bind(now) + .bind(now) + .execute(&mut *connection), + ) + .expect("insert historical job"); } let reopened = TransportPublishStore::open(database_path).expect("reopen store"); @@ -4381,25 +4552,18 @@ mod tests { let directory = tempfile::tempdir().expect("tempdir"); let database_path = directory.path().join("publish-proxy-current.sqlite"); let store = TransportPublishStore::open(database_path).expect("open current schema"); - let connection = store + let mut connection = store .inner .lock() .unwrap_or_else(std::sync::PoisonError::into_inner); - let columns = connection - .prepare("PRAGMA table_info(transport_publish_principals)") - .expect("prepare schema") - .query_map([], |row| row.get::<_, String>(1)) - .expect("query schema") - .collect::<Result<Vec<_>, _>>() - .expect("columns"); + let columns = test_query_column_names(&mut connection, "transport_publish_principals"); assert!( columns .iter() .any(|column| column == "allowed_explicit_transport_kinds_json") ); - let version = connection - .pragma_query_value(None, "user_version", |row| row.get::<_, i64>(0)) - .expect("user version"); + let version = + super::transport_publish_schema_version(&mut connection).expect("user version"); assert_eq!(version, SCHEMA_VERSION); } @@ -4409,10 +4573,10 @@ mod tests { let database_path = directory.path().join("publish-proxy-v1.sqlite"); let token_hash = hash_bearer_token(generate_bearer_token().as_str()); { - let connection = rusqlite::Connection::open(database_path.as_path()).expect("open"); - connection - .execute_batch( - r#" + let mut connection = open_test_database(database_path.as_path()); + super::execute_raw_sql( + &mut connection, + r#" CREATE TABLE transport_publish_principals ( principal_id TEXT PRIMARY KEY NOT NULL, label TEXT NOT NULL, @@ -4428,13 +4592,15 @@ mod tests { created_at_unix INTEGER NOT NULL ); "#, - ) - .expect("schema"); - connection - .pragma_update(None, "user_version", SCHEMA_VERSION) - .expect("set schema version"); - connection - .execute( + ) + .expect("schema"); + super::execute_sql( + &mut connection, + format!("PRAGMA user_version = {SCHEMA_VERSION}").as_str(), + ) + .expect("set schema version"); + super::block_on_sqlite( + sqlx::query( r#" INSERT INTO transport_publish_principals ( principal_id, @@ -4452,24 +4618,26 @@ mod tests { ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, NULL, NULL, ?10) "#, - rusqlite::params![ - "principal-v1", - "v1", - token_hash, - serde_json::to_string(&vec!["a".repeat(64)]).expect("pubkeys"), - serde_json::to_string(&vec![30_402]).expect("kinds"), - serde_json::to_string(&vec![TransportPublishTargetPolicyName::Nostr]) - .expect("policies"), - serde_json::to_string(&vec![ - NostrPublishTargetSourcePolicy::DaemonDefaultOnly - ]) + ) + .bind("principal-v1") + .bind("v1") + .bind(token_hash.as_str()) + .bind(serde_json::to_string(&vec!["a".repeat(64)]).expect("pubkeys")) + .bind(serde_json::to_string(&vec![30_402]).expect("kinds")) + .bind( + serde_json::to_string(&vec![TransportPublishTargetPolicyName::Nostr]) + .expect("policies"), + ) + .bind( + serde_json::to_string(&vec![NostrPublishTargetSourcePolicy::DaemonDefaultOnly]) .expect("source policies"), - false, - PublishJobVisibility::Own.to_string(), - 1_i64, - ], ) - .expect("principal"); + .bind(false) + .bind(PublishJobVisibility::Own.to_string()) + .bind(1_i64) + .execute(&mut connection), + ) + .expect("principal"); } let error = match TransportPublishStore::open(database_path.clone()) { @@ -4484,14 +4652,8 @@ mod tests { error => panic!("unexpected error: {error}"), } - let connection = rusqlite::Connection::open(database_path.as_path()).expect("open legacy"); - let columns = connection - .prepare("PRAGMA table_info(transport_publish_principals)") - .expect("prepare schema") - .query_map([], |row| row.get::<_, String>(1)) - .expect("query schema") - .collect::<Result<Vec<_>, _>>() - .expect("columns"); + let mut connection = open_test_database(database_path.as_path()); + let columns = test_query_column_names(&mut connection, "transport_publish_principals"); assert!( !columns .iter() diff --git a/src/main.rs b/src/main.rs @@ -64,14 +64,12 @@ async fn run() -> Result<()> { mod tests { use super::{exit_code_from_run, main, run, run_hook}; use std::process::ExitCode; - use std::sync::Mutex; + use tokio::sync::{Mutex, MutexGuard}; - static TEST_LOCK: Mutex<()> = Mutex::new(()); + static TEST_LOCK: Mutex<()> = Mutex::const_new(()); - fn test_guard() -> std::sync::MutexGuard<'static, ()> { - let guard = TEST_LOCK - .lock() - .unwrap_or_else(std::sync::PoisonError::into_inner); + async fn test_guard() -> MutexGuard<'static, ()> { + let guard = TEST_LOCK.lock().await; *run_hook() .lock() .unwrap_or_else(std::sync::PoisonError::into_inner) = None; @@ -94,7 +92,7 @@ mod tests { #[tokio::test] async fn run_returns_error_when_hook_is_missing() { - let _guard = test_guard(); + let _guard = test_guard().await; let err = run().await.expect_err("hook missing should error"); let msg = format!("{err:#}"); assert!(msg.contains("run hook not set")); @@ -102,7 +100,7 @@ mod tests { #[tokio::test] async fn run_uses_hook_result() { - let _guard = test_guard(); + let _guard = test_guard().await; *run_hook() .lock() .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(Ok(())); diff --git a/src/transport/jsonrpc/methods/nip46/connect.rs b/src/transport/jsonrpc/methods/nip46/connect.rs @@ -107,7 +107,7 @@ async fn connect_bunker( .await; let request = NostrConnectRequest::Connect { - remote_signer_public_key: remote_signer_pubkey.clone(), + remote_signer_public_key: remote_signer_pubkey, secret: info.secret.clone(), }; let message = NostrConnectMessage::request(&request); @@ -287,12 +287,9 @@ async fn send_connect_request( remote_signer_pubkey: &RadrootsNostrPublicKey, message: NostrConnectMessage, ) -> Result<(), RpcError> { - let event = RadrootsNostrEventBuilder::nostr_connect( - client_keys, - remote_signer_pubkey.clone(), - message, - ) - .map_err(|e| RpcError::Other(format!("nip46 connect request failed: {e}")))?; + let event = + RadrootsNostrEventBuilder::nostr_connect(client_keys, *remote_signer_pubkey, message) + .map_err(|e| RpcError::Other(format!("nip46 connect request failed: {e}")))?; client .send_event_builder(event) .await @@ -307,7 +304,7 @@ fn connect_response_filter( ) -> Result<RadrootsNostrFilter, RpcError> { let filter = RadrootsNostrFilter::new() .kind(RadrootsNostrKind::NostrConnect) - .author(remote_signer_pubkey.clone()) + .author(*remote_signer_pubkey) .since(since); radroots_nostr_filter_tag(filter, "p", vec![client_pubkey.to_hex()]) .map_err(|e| RpcError::Other(format!("nip46 connect filter failed: {e}"))) diff --git a/src/transport/jsonrpc/methods/nip46/get_public_key.rs b/src/transport/jsonrpc/methods/nip46/get_public_key.rs @@ -25,15 +25,14 @@ pub fn register(m: &mut RpcModule<RpcContext>, registry: &MethodRegistry) -> Res .map_err(|e| RpcError::InvalidParams(e.to_string()))?; let session = session::get_session(ctx.as_ref(), &session_id).await?; let (pubkey, updated) = request_get_public_key(&session).await?; - if updated { - if !ctx + if updated + && !ctx .state .nip46_sessions - .set_user_pubkey(&session_id, pubkey.clone()) + .set_user_pubkey(&session_id, pubkey) .await - { - return Err(RpcError::Other("nip46 session update failed".to_string())); - } + { + return Err(RpcError::Other("nip46 session update failed".to_string())); } Ok::<Nip46GetPublicKeyResponse, RpcError>(Nip46GetPublicKeyResponse { pubkey: pubkey.to_hex(), diff --git a/src/transport/jsonrpc/nip46/client.rs b/src/transport/jsonrpc/nip46/client.rs @@ -72,7 +72,7 @@ pub async fn request( .map_err(|e| RpcError::Other(format!("nip46 {label} failed: {e}")))?; let event = RadrootsNostrEventBuilder::nostr_connect( &session.client_keys, - session.remote_signer_pubkey.clone(), + session.remote_signer_pubkey, message, ) .map_err(|e| RpcError::Other(format!("nip46 {label} failed: {e}")))?; @@ -104,7 +104,7 @@ fn response_filter( ) -> Result<RadrootsNostrFilter, RpcError> { let filter = RadrootsNostrFilter::new() .kind(RadrootsNostrKind::NostrConnect) - .author(session.remote_signer_pubkey.clone()) + .author(session.remote_signer_pubkey) .since(since); radroots_nostr_filter_tag(filter, "p", vec![session.client_pubkey.to_hex()]) .map_err(|e| RpcError::Other(format!("nip46 {label} failed: {e}"))) diff --git a/src/transport/jsonrpc/server.rs b/src/transport/jsonrpc/server.rs @@ -73,6 +73,51 @@ where } } +pub async fn start_server( + addr: SocketAddr, + rpc_cfg: &RpcConfig, + transport_publish_store: TransportPublishStore, + root: RpcModule<RpcContext>, +) -> Result<ServerHandle> { + let mut builder = ServerConfigBuilder::new() + .max_request_body_size(rpc_cfg.max_request_body_size) + .max_response_body_size(rpc_cfg.max_response_body_size) + .max_connections(rpc_cfg.max_connections) + .max_subscriptions_per_connection(rpc_cfg.max_subscriptions_per_connection) + .set_message_buffer_capacity(rpc_cfg.message_buffer_capacity); + + if let Some(limit) = rpc_cfg.batch_request_limit { + let cfg = if limit == 0 { + BatchRequestConfig::Disabled + } else { + BatchRequestConfig::Limit(limit) + }; + builder = builder.set_batch_request_config(cfg); + } + + let server_cfg = builder.build(); + let rpc_middleware = + RpcServiceBuilder::new().layer_fn(|service| RejectPublishNotifications { service }); + let server = ServerBuilder::with_config(server_cfg) + .set_rpc_middleware(rpc_middleware) + .set_http_middleware(tower::ServiceBuilder::new().map_request( + move |mut request: HttpRequest<HttpBody>| { + let transport_publish_auth = auth::authorize_transport_publish_request( + request + .headers() + .get("authorization") + .and_then(|value| value.to_str().ok()), + &transport_publish_store, + ); + request.extensions_mut().insert(transport_publish_auth); + request + }, + )) + .build(addr) + .await?; + Ok(server.start(root)) +} + #[cfg(test)] mod tests { use super::start_server; @@ -428,48 +473,3 @@ mod tests { ); } } - -pub async fn start_server( - addr: SocketAddr, - rpc_cfg: &RpcConfig, - transport_publish_store: TransportPublishStore, - root: RpcModule<RpcContext>, -) -> Result<ServerHandle> { - let mut builder = ServerConfigBuilder::new() - .max_request_body_size(rpc_cfg.max_request_body_size) - .max_response_body_size(rpc_cfg.max_response_body_size) - .max_connections(rpc_cfg.max_connections) - .max_subscriptions_per_connection(rpc_cfg.max_subscriptions_per_connection) - .set_message_buffer_capacity(rpc_cfg.message_buffer_capacity); - - if let Some(limit) = rpc_cfg.batch_request_limit { - let cfg = if limit == 0 { - BatchRequestConfig::Disabled - } else { - BatchRequestConfig::Limit(limit) - }; - builder = builder.set_batch_request_config(cfg); - } - - let server_cfg = builder.build(); - let rpc_middleware = - RpcServiceBuilder::new().layer_fn(|service| RejectPublishNotifications { service }); - let server = ServerBuilder::with_config(server_cfg) - .set_rpc_middleware(rpc_middleware) - .set_http_middleware(tower::ServiceBuilder::new().map_request( - move |mut request: HttpRequest<HttpBody>| { - let transport_publish_auth = auth::authorize_transport_publish_request( - request - .headers() - .get("authorization") - .and_then(|value| value.to_str().ok()), - &transport_publish_store, - ); - request.extensions_mut().insert(transport_publish_auth); - request - }, - )) - .build(addr) - .await?; - Ok(server.start(root)) -} diff --git a/src/transport/nostr/listener.rs b/src/transport/nostr/listener.rs @@ -129,7 +129,7 @@ pub(crate) async fn handle_request( id: session_id, client: radrootsd.client.clone(), client_keys: radrootsd.keys.clone(), - client_pubkey: client_pubkey.clone(), + client_pubkey: *client_pubkey, remote_signer_pubkey: radrootsd.pubkey, user_pubkey: Some(radrootsd.pubkey), relays: Vec::new(), @@ -193,7 +193,7 @@ pub(crate) async fn handle_request( request_id, client_pubkey, NostrConnectRequest::Nip04Encrypt { - public_key: public_key.clone(), + public_key, text: text.clone(), }, ) @@ -227,7 +227,7 @@ pub(crate) async fn handle_request( request_id, client_pubkey, NostrConnectRequest::Nip04Decrypt { - public_key: public_key.clone(), + public_key, ciphertext: ciphertext.clone(), }, ) @@ -258,7 +258,7 @@ pub(crate) async fn handle_request( request_id, client_pubkey, NostrConnectRequest::Nip44Encrypt { - public_key: public_key.clone(), + public_key, text: text.clone(), }, ) @@ -297,7 +297,7 @@ pub(crate) async fn handle_request( request_id, client_pubkey, NostrConnectRequest::Nip44Decrypt { - public_key: public_key.clone(), + public_key, ciphertext: ciphertext.clone(), }, ) @@ -349,7 +349,7 @@ async fn auth_challenge( } let pending = PendingNostrRequest { request_id: request_id.to_string(), - client_pubkey: client_pubkey.clone(), + client_pubkey: *client_pubkey, request, }; let _ = radrootsd