diff --git a/src/bin/bal-pusher.rs b/src/bin/bal-pusher.rs index 5c51813..4d95899 100644 --- a/src/bin/bal-pusher.rs +++ b/src/bin/bal-pusher.rs @@ -312,7 +312,9 @@ async fn main_result(cfg: &MyConfig, network_params: &NetworkParams) -> Result<( let _ = calculate_stats(&db, network_params.db_field.clone()).await; } Err(erx) => { - panic!("impossible to get client {}", erx) + error!("impossible to get client: {}, retrying on next block", erx); + thread::sleep(Duration::from_secs(5)); + return Ok(()); } } Ok(()) @@ -622,17 +624,38 @@ async fn main() -> std::io::Result<()> { let zmq_address = network_params.zmq_listener.clone(); info!("zmq listening on: {}", zmq_address); - socket.connect(&zmq_address).unwrap(); + loop { + match socket.connect(&zmq_address) { + Ok(_) => break, + Err(e) => { + error!("ZMQ connect failed: {}, retrying in 5s...", e); + thread::sleep(Duration::from_secs(5)); + } + } + } - socket.set_subscribe(b"").unwrap(); + match socket.set_subscribe(b"") { + Ok(_) => {}, + Err(e) => { + error!("ZMQ subscribe failed: {}, exiting", e); + return Ok(()); + } + } let _ = main_result(&cfg, network_params).await; info!("waiting new blocks.."); let mut last_seq: Vec = [0; 4].to_vec(); let mut counter = 0; let max = 100; + socket.set_rcvtimeo(5000).unwrap(); // 5 seconds timeout loop { - let message = socket.recv_multipart(0).unwrap(); + let message = match socket.recv_multipart(0) { + Ok(m) => m, + Err(e) => { + warn!("ZMQ recv timeout or error: {}, retrying...", e); + continue; + } + }; let topic = message[0].clone(); let body = message[1].clone(); let seq = message[2].clone(); diff --git a/src/bin/bal-server.rs b/src/bin/bal-server.rs index def0906..ed27bfd 100644 --- a/src/bin/bal-server.rs +++ b/src/bin/bal-server.rs @@ -142,6 +142,7 @@ async fn echo_pub_key( async fn echo_stats( param: &str, cfg: &MyConfig, + db: &Arc>, ) -> Result>, hyper::Error> { info!("echo stats!!! {} - {}", param, cfg.expose_stats); let netconfig = MyConfig::get_net_config(cfg, param); @@ -165,39 +166,45 @@ async fn echo_stats( netconfig.name ); let mut stats: Vec = vec![]; - let db = sqlite::open(&cfg.db_file).unwrap(); + let db = match db.lock() { + Ok(g) => g, + Err(p) => { + error!("DB mutex poisoned in echo_stats"); + p.into_inner() + } + }; let _ = db.iterate(&sql, |pairs| { let row: HashMap<_, _> = pairs .into_iter() .map(|(k, v)| (k.to_string(), v.map(|s| s))) .collect(); //let row:HashMap<_,_>= pairs.into_iter().collect(); - println!("row report date {}", row["report_date"].clone().unwrap()); + println!("row report date {}", row["report_date"].clone().unwrap_or("0")); dbg!(&row); stats.push(StatsResponse { - report_date: row["report_date"].clone().unwrap().to_string(), - chain: row["chain"].clone().unwrap().to_string(), - totals: row["totals"].clone().unwrap().parse::().unwrap(), - waiting: row["waiting"].clone().unwrap().parse::().unwrap(), - sent: row["sent"].clone().unwrap().parse::().unwrap(), - failed: row["failed"].clone().unwrap().parse::().unwrap(), + report_date: row["report_date"].clone().unwrap_or("0").to_string(), + chain: row["chain"].clone().unwrap_or("?").to_string(), + totals: row["totals"].clone().unwrap_or("0").parse::().unwrap_or(0), + waiting: row["waiting"].clone().unwrap_or("0").parse::().unwrap_or(0), + sent: row["sent"].clone().unwrap_or("0").parse::().unwrap_or(0), + failed: row["failed"].clone().unwrap_or("0").parse::().unwrap_or(0), waiting_profit: row["waiting_profit"] .clone() - .unwrap() + .unwrap_or("0") .parse::() - .unwrap(), - sent_profit: row["sent_profit"].clone().unwrap().parse::().unwrap(), + .unwrap_or(0), + sent_profit: row["sent_profit"].clone().unwrap_or("0").parse::().unwrap_or(0), missed_profit: row["missed_profit"] .clone() - .unwrap() + .unwrap_or("0") .parse::() - .unwrap(), + .unwrap_or(0), unique_inputs: row["unique_inputs"] .clone() - .unwrap() + .unwrap_or("0") .parse::() - .unwrap(), + .unwrap_or(0), }); true }); @@ -214,6 +221,7 @@ async fn echo_info( param: &str, cfg: &MyConfig, remote_addr: &String, + db: &Arc>, ) -> Result>, hyper::Error> { info!("echo info!!!{}", param); let netconfig = MyConfig::get_net_config(cfg, param); @@ -228,7 +236,13 @@ async fn echo_info( address } true => { - let db = sqlite::open(&cfg.db_file).unwrap(); + let db = match db.lock() { + Ok(g) => g, + Err(p) => { + error!("DB mutex poisoned in echo_info"); + p.into_inner() + } + }; match get_last_used_address_by_ip( &db, &netconfig.name, @@ -268,15 +282,27 @@ async fn echo_info( async fn echo_search( whole_body: &Bytes, cfg: &MyConfig, + db: &Arc>, ) -> Result>, hyper::Error> { info!("echo search!!!"); - let strbody = std::str::from_utf8(whole_body).unwrap(); + let strbody = match std::str::from_utf8(whole_body) { + Ok(s) => s, + Err(_) => { + return Ok(Response::new(full("Invalid UTF-8 body"))); + } + }; info!("{}", strbody); let mut response = Response::new(full("Bad data received".to_owned())); *response.status_mut() = StatusCode::BAD_REQUEST; if !strbody.is_empty() && strbody.len() <= 70 { - let db = sqlite::open(&cfg.db_file).unwrap(); + let db = match db.lock() { + Ok(g) => g, + Err(p) => { + error!("DB mutex poisoned in echo_search"); + p.into_inner() + } + }; let mut statement = db .prepare("SELECT * FROM tbl_tx WHERE txid = ? LIMIT 1") .unwrap(); @@ -344,10 +370,16 @@ async fn echo_push( whole_body: &Bytes, cfg: &MyConfig, param: &str, + db: &Arc>, ) -> Result>, hyper::Error> { //let whole_body = req.collect().await?.to_bytes(); trace!("echo_push"); - let strbody = std::str::from_utf8(whole_body).unwrap(); + let strbody = match std::str::from_utf8(whole_body) { + Ok(s) => s, + Err(_) => { + return Ok(Response::new(full("Invalid UTF-8 body"))); + } + }; let mut response = Response::new(full("Bad data received".to_owned())); let mut response_not_enable = Response::new(full("Network not enabled".to_owned())); *response.status_mut() = StatusCode::BAD_REQUEST; @@ -357,9 +389,20 @@ async fn echo_push( trace!("network not enabled {}", &netconfig.name); return Ok(response_not_enable); } - let req_time = Utc::now().timestamp_nanos_opt().unwrap(); // Returns i64 - - let db = sqlite::open(&cfg.db_file).unwrap(); + let req_time = match Utc::now().timestamp_nanos_opt() { + Some(t) => t, + None => { + error!("Invalid timestamp"); + return Ok(response); + } + }; // Returns i64 + let db = match db.lock() { + Ok(g) => g, + Err(p) => { + error!("DB mutex poisoned in echo_push"); + p.into_inner() + } + }; let lines = strbody.split("\n"); let sqltxshead = "INSERT INTO tbl_tx (txid, wtxid, ntxid, tx, locktime, reqid, network, our_address, our_fees)".to_string(); @@ -532,6 +575,7 @@ async fn echo( req: Request, cfg: &MyConfig, ip: &String, + db: &Arc>, ) -> Result>, hyper::Error> { let mut not_found = Response::new(empty()); *not_found.status_mut() = StatusCode::NOT_FOUND; @@ -553,20 +597,20 @@ async fn echo( let whole_body = req.collect().await?.to_bytes(); if let Some(param) = match_uri(r"^?/?(?P[^/]?+)?/pushtxs$", uri.as_str()) { //let whole_body = collect_body(req,512_000).await?; - ret = echo_push(&whole_body, cfg, param).await; + ret = echo_push(&whole_body, cfg, param, db).await; } if uri == "/searchtx" { //let whole_body = collect_body(req,64).await?; - ret = echo_search(&whole_body, cfg).await; + ret = echo_search(&whole_body, cfg, db).await; } ret } Method::GET => { if let Some(param) = match_uri(r"^?/?(?P[^/]?+)?/stats$", uri.as_str()) { - ret = echo_stats(param, cfg).await; + ret = echo_stats(param, cfg, db).await; } if let Some(param) = match_uri(r"^?/?(?P[^/]?+)?/info$", uri.as_str()) { - ret = echo_info(param, cfg, &remote_addr).await; + ret = echo_info(param, cfg, &remote_addr, db).await; } if uri == "/version" { ret = echo_version().await; @@ -600,7 +644,13 @@ fn parse_env(cfg: &Arc>) { //for (key, value) in std::env::vars() { // debug!("ENVIRONMENT {key}: {value}"); //} - let mut cfg_lock = cfg.lock().unwrap(); + let mut cfg_lock = match cfg.lock() { + Ok(g) => g, + Err(p) => { + error!("Config mutex poisoned in parse_env, recovering"); + p.into_inner() + } + }; if let Ok(value) = env::var("BAL_SERVER_DB_FILE") { debug!("BAL_SERVER_DB_FILE: {}", value); cfg_lock.db_file = value; @@ -675,11 +725,19 @@ async fn main() -> Result<(), Box> { let cfg: Arc> = Arc::>::default(); parse_env(&cfg); - let cfg_lock = cfg.lock().unwrap(); + let cfg_lock = match cfg.lock() { + Ok(g) => g, + Err(p) => { + error!("Config mutex poisoned at startup, recovering"); + p.into_inner() + } + }; - let db = sqlite::open(&cfg_lock.db_file).unwrap(); - create_database(&db); - init_network(&db, &cfg_lock); + let db = Arc::new(Mutex::new(sqlite::open(&cfg_lock.db_file).unwrap())); + let db_guard = db.lock().unwrap(); + create_database(&*db_guard); + init_network(&*db_guard, &cfg_lock); + drop(db_guard); let addr = cfg_lock.bind_address.to_string(); let addr: IpAddr = addr.parse()?; @@ -692,7 +750,7 @@ async fn main() -> Result<(), Box> { let ip = stream .peer_addr()? .to_string() - .split(":") + .split(':') .next() .unwrap() .to_string(); @@ -700,12 +758,13 @@ async fn main() -> Result<(), Box> { tokio::task::spawn({ let cfg = cfg_lock.clone(); + let db = db.clone(); async move { if let Err(err) = http1::Builder::new() .serve_connection( io, service_fn(|req: Request| async { - echo(req, &cfg, &ip).await + echo(req, &cfg, &ip, &db).await }), ) .await diff --git a/tests/panic_regression_tests.rs b/tests/panic_regression_tests.rs new file mode 100644 index 0000000..70f277b --- /dev/null +++ b/tests/panic_regression_tests.rs @@ -0,0 +1,43 @@ +use std::sync::{Arc, Mutex}; +use std::thread; +use std::collections::HashMap; +use sqlite::Connection; + +#[test] +fn test_mutex_poisoning_recovery() { + let data = Arc::new(Mutex::new(0)); + let c = data.clone(); + let handle = thread::spawn(move || { + let _guard = c.lock(); // Acquire lock + panic!("test panic"); // Panic while holding the lock + // _guard is dropped during panic unwinding, poisoning the mutex + }); + let result = handle.join(); + assert!(result.is_err()); // Thread panicked + + // Recovery: the same pattern used in bal-server.rs + let guard = match data.lock() { + Ok(g) => g, + Err(p) => { + p.into_inner() // Should not panic + } + }; + assert_eq!(*guard, 0); +} + +#[test] +fn test_db_null_unwrap_or() { + let db = Connection::open(":memory:").unwrap(); + let _ = db.execute("CREATE TABLE test_stats (report_date TEXT, chain TEXT, totals TEXT, waiting TEXT);"); + let _ = db.execute("INSERT INTO test_stats (report_date, chain) VALUES ('2024-01-01', 'testnet');"); + + let mut found_value = None; + let _ = db.iterate("SELECT * FROM test_stats;", |pairs| { + let row: HashMap<_, _> = pairs.into_iter().map(|(k,v)| (k.to_string(), v.map(|s| s))).collect(); + let totals = row["totals"].clone().unwrap_or("0").to_string(); + found_value = Some(totals); + true + }); + + assert_eq!(found_value.unwrap(), "0"); +}