mod common;
use async_trait::async_trait;
use common::{MockConnectionFactory, Rng};
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use sz_orm_core::TransactOptions;
use sz_orm_core::Transaction;
use sz_orm_core::{Connection, ConnectionFactory, DbError, Pool, PoolConfigBuilder};
use tokio::sync::Mutex;
struct DisconnectableConnection {
connected: Arc<AtomicBool>,
closed: AtomicBool,
}
impl DisconnectableConnection {
fn new(connected: Arc<AtomicBool>) -> Self {
Self {
connected,
closed: AtomicBool::new(false),
}
}
fn is_alive(&self) -> bool {
self.connected.load(Ordering::SeqCst) && !self.closed.load(Ordering::SeqCst)
}
}
impl Connection for DisconnectableConnection {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
if !self.is_alive() {
return Err(DbError::ConnectionError("network partition".to_string()));
}
Ok(1)
})
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, sz_orm_core::Value>>,
DbError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async move {
if !self.is_alive() {
return Err(DbError::ConnectionError("network partition".to_string()));
}
Ok(vec![])
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if !self.is_alive() {
return Err(DbError::ConnectionError("network partition".to_string()));
}
Ok(())
})
}
fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if !self.is_alive() {
return Err(DbError::ConnectionError("network partition".to_string()));
}
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
self.is_alive()
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { self.is_alive() })
}
fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
self.closed.store(true, Ordering::SeqCst);
Ok(())
})
}
}
struct DiskFullConnection {
connected: bool,
fail_after: u32,
execute_count: AtomicU32,
}
impl DiskFullConnection {
fn new(fail_after: u32) -> Self {
Self {
connected: true,
fail_after,
execute_count: AtomicU32::new(0),
}
}
}
impl Connection for DiskFullConnection {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
let n = self.execute_count.fetch_add(1, Ordering::SeqCst);
if n >= self.fail_after {
return Err(DbError::IoError("disk full".to_string()));
}
Ok(1)
})
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, sz_orm_core::Value>>,
DbError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async move {
let n = self.execute_count.fetch_add(1, Ordering::SeqCst);
if n >= self.fail_after {
return Err(DbError::IoError("disk full".to_string()));
}
Ok(vec![])
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
let n = self.execute_count.load(Ordering::SeqCst);
if n >= self.fail_after {
return Err(DbError::IoError("disk full on commit".to_string()));
}
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
self.connected
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { self.connected })
}
fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
self.connected = false;
Ok(())
})
}
}
struct FailoverFactory {
primary_available: Arc<AtomicBool>,
failover_count: AtomicU32,
}
impl FailoverFactory {
fn new(primary_available: Arc<AtomicBool>) -> Self {
Self {
primary_available,
failover_count: AtomicU32::new(0),
}
}
fn failover_count(&self) -> u32 {
self.failover_count.load(Ordering::SeqCst)
}
}
#[async_trait]
impl ConnectionFactory for FailoverFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
if !self.primary_available.load(Ordering::SeqCst) {
self.failover_count.fetch_add(1, Ordering::SeqCst);
}
let connected = Arc::new(AtomicBool::new(true));
Ok(Box::new(DisconnectableConnection::new(connected)))
}
}
struct FlakyFactory {
available: Arc<AtomicBool>,
}
impl FlakyFactory {
fn new(available: Arc<AtomicBool>) -> Self {
Self { available }
}
}
#[async_trait]
impl ConnectionFactory for FlakyFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
if !self.available.load(Ordering::SeqCst) {
return Err(DbError::ConnectionRefused(
"service unavailable".to_string(),
));
}
let connected = Arc::new(AtomicBool::new(true));
Ok(Box::new(DisconnectableConnection::new(connected)))
}
}
struct NetworkPartitionFactory {
conns: std::sync::Mutex<Vec<Arc<AtomicBool>>>,
}
impl NetworkPartitionFactory {
fn new() -> Self {
Self {
conns: std::sync::Mutex::new(Vec::new()),
}
}
fn partition_all(&self) {
let conns = self.conns.lock().unwrap();
for c in conns.iter() {
c.store(false, Ordering::SeqCst);
}
}
}
#[async_trait]
impl ConnectionFactory for NetworkPartitionFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
let flag = Arc::new(AtomicBool::new(true));
self.conns.lock().unwrap().push(flag.clone());
Ok(Box::new(DisconnectableConnection::new(flag)))
}
}
struct DiskFullFactory {
fail_after: u32,
}
impl DiskFullFactory {
fn new(fail_after: u32) -> Self {
Self { fail_after }
}
}
#[async_trait]
impl ConnectionFactory for DiskFullFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
Ok(Box::new(DiskFullConnection::new(self.fail_after)))
}
}
fn make_pool_with_factory(
max_size: u32,
acquire_timeout_secs: u64,
factory: Arc<dyn ConnectionFactory>,
) -> &'static Pool {
let config = PoolConfigBuilder::new()
.max_size(max_size)
.min_idle(0)
.acquire_timeout(acquire_timeout_secs)
.build()
.unwrap();
let pool = Pool::new(config, factory).unwrap();
Box::leak(Box::new(pool))
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn chaos_network_partition_idle_connections_dropped() {
let factory = Arc::new(NetworkPartitionFactory::new());
let pool = make_pool_with_factory(10, 5, factory.clone());
let mut conns = Vec::new();
for _ in 0..3 {
conns.push(pool.acquire().await.unwrap());
}
for conn in conns {
pool.release(conn).await;
}
let status = pool.status().await;
assert_eq!(status.idle, 3, "should have 3 idle connections");
factory.partition_all();
let conn = pool.acquire().await.unwrap();
assert!(conn.is_connected(), "acquire should return healthy conn");
let status = pool.status().await;
assert_eq!(
status.idle, 0,
"idle should be 0 after dropping disconnected"
);
assert!(
status.active <= 10,
"active must not exceed max_size: {}",
status.active
);
}
#[tokio::test]
async fn chaos_network_partition_mid_operation_rollback() {
let connected = Arc::new(AtomicBool::new(true));
let conn = DisconnectableConnection::new(connected.clone());
let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
assert!(tx.is_active());
tx.execute("INSERT 1").await.unwrap();
connected.store(false, Ordering::SeqCst);
let result = tx.execute("INSERT 2").await;
assert!(result.is_err(), "execute during partition should fail");
let rollback_result = tx.rollback().await;
assert!(rollback_result.is_ok(), "rollback should succeed");
assert!(!tx.is_active());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn chaos_network_partition_concurrent_consistency() {
let factory = Arc::new(NetworkPartitionFactory::new());
let pool = make_pool_with_factory(20, 2, factory.clone());
let ops_completed = Arc::new(AtomicU64::new(0));
let max_active = Arc::new(AtomicU32::new(0));
let mut handles = Vec::new();
for task_id in 0..8 {
let ops = ops_completed.clone();
let max_act = max_active.clone();
let factory_c = factory.clone();
handles.push(tokio::spawn(async move {
for i in 0..50 {
if i == 25 && task_id == 0 {
factory_c.partition_all();
tokio::time::sleep(Duration::from_millis(5)).await;
}
let conn = match pool.acquire().await {
Ok(c) => c,
Err(_) => continue,
};
let status = pool.status().await;
max_act.fetch_max(status.active, Ordering::Relaxed);
tokio::task::yield_now().await;
pool.release(conn).await;
ops.fetch_add(1, Ordering::Relaxed);
}
}));
}
for h in handles {
h.await.unwrap();
}
assert!(ops_completed.load(Ordering::Relaxed) > 0);
assert!(
max_active.load(Ordering::Relaxed) <= 20,
"active must never exceed max_size: {}",
max_active.load(Ordering::Relaxed)
);
let status = pool.status().await;
assert_eq!(
status.active, status.idle,
"after all ops: active={}, idle={}",
status.active, status.idle
);
}
#[tokio::test]
async fn chaos_disk_full_write_fails() {
let conn = DiskFullConnection::new(1); let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
tx.execute("INSERT 1").await.unwrap();
let result = tx.execute("INSERT 2").await;
assert!(result.is_err(), "second execute should fail with disk full");
if let Err(e) = result {
let msg = format!("{}", e);
assert!(
msg.contains("disk full") || msg.contains("Io"),
"error should mention disk full: {}",
msg
);
}
tx.rollback().await.unwrap();
assert!(!tx.is_active());
}
#[tokio::test]
async fn chaos_disk_full_commit_fails_rollback_succeeds() {
let conn = DiskFullConnection::new(5); let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
for _ in 0..5 {
tx.execute("INSERT").await.unwrap();
}
let commit_result = tx.commit().await;
assert!(commit_result.is_err(), "commit should fail with disk full");
let rollback_result = tx.rollback().await;
assert!(
rollback_result.is_ok(),
"rollback after failed commit should succeed"
);
assert!(!tx.is_active());
}
#[tokio::test]
async fn chaos_disk_full_pool_recovers() {
let factory = Arc::new(DiskFullFactory::new(1)); let pool = make_pool_with_factory(5, 3, factory);
let conn1 = pool.acquire().await.unwrap();
let mut tx = Transaction::new(conn1.into_inner(), TransactOptions::default());
tx.execute("INSERT 1").await.unwrap();
let _ = tx.execute("INSERT 2").await; let _ = tx.rollback().await;
let conn2 = pool.acquire().await;
assert!(conn2.is_ok(), "pool should still serve new connections");
let mut conn2 = conn2.unwrap();
let exec_result = conn2.execute("SELECT 1").await;
assert!(
exec_result.is_ok(),
"new connection should execute successfully, got: {:?}",
exec_result
);
let status = pool.status().await;
assert!(
status.active <= 5,
"active must not exceed max_size: {}",
status.active
);
}
#[tokio::test]
async fn chaos_clock_drift_connection_expired() {
let db = Arc::new(Mutex::new(common::InMemoryDb::new()));
let factory = Arc::new(MockConnectionFactory::new(db));
let config = PoolConfigBuilder::new()
.max_size(5)
.min_idle(0)
.max_lifetime(1) .idle_timeout(1)
.acquire_timeout(5)
.build()
.unwrap();
let pool: &'static Pool = Box::leak(Box::new(Pool::new(config, factory).unwrap()));
let conn = pool.acquire().await.unwrap();
pool.release(conn).await;
tokio::time::sleep(Duration::from_millis(1200)).await;
let conn = pool.acquire().await.unwrap();
assert!(conn.is_connected(), "should get fresh connection");
let status = pool.status().await;
assert!(
status.active <= 5,
"active must not exceed max_size: {}",
status.active
);
}
#[tokio::test]
async fn chaos_clock_drift_reap_idle_expires_connections() {
let db = Arc::new(Mutex::new(common::InMemoryDb::new()));
let factory = Arc::new(MockConnectionFactory::new(db));
let config = PoolConfigBuilder::new()
.max_size(5)
.min_idle(0)
.max_lifetime(1)
.idle_timeout(1)
.acquire_timeout(5)
.build()
.unwrap();
let pool: &'static Pool = Box::leak(Box::new(Pool::new(config, factory).unwrap()));
let mut conns = Vec::new();
for _ in 0..3 {
conns.push(pool.acquire().await.unwrap());
}
let status = pool.status().await;
assert_eq!(status.active, 3, "should have 3 active (borrowed)");
assert_eq!(status.idle, 0, "should have 0 idle while all borrowed");
for conn in conns {
pool.release(conn).await;
}
let status = pool.status().await;
assert_eq!(status.idle, 3, "should have 3 idle after release");
tokio::time::sleep(Duration::from_millis(1200)).await;
pool.reap_idle().await;
let status = pool.status().await;
assert_eq!(status.idle, 0, "all expired connections should be reaped");
assert_eq!(status.active, 0, "active should be 0 after reaping all");
}
#[tokio::test]
async fn chaos_master_slave_failover_to_secondary() {
let primary_available = Arc::new(AtomicBool::new(true));
let factory = Arc::new(FailoverFactory::new(primary_available.clone()));
let pool = make_pool_with_factory(5, 3, factory.clone());
let conn1 = pool.acquire().await.unwrap();
assert!(conn1.is_connected());
assert_eq!(factory.failover_count(), 0, "no failover yet");
primary_available.store(false, Ordering::SeqCst);
let conn2 = pool.acquire().await.unwrap();
assert!(
conn2.is_connected(),
"secondary connection should be healthy"
);
assert_eq!(
factory.failover_count(),
1,
"failover_count should be 1 after one secondary create"
);
primary_available.store(true, Ordering::SeqCst);
let conn3 = pool.acquire().await.unwrap();
assert!(conn3.is_connected());
assert_eq!(
factory.failover_count(),
1,
"no new failover after primary recovered"
);
pool.release(conn1).await;
pool.release(conn2).await;
pool.release(conn3).await;
let status = pool.status().await;
assert!(
status.active <= 5,
"active must not exceed max_size: {}",
status.active
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn chaos_master_slave_pool_recovery_after_outage() {
let available = Arc::new(AtomicBool::new(true));
let factory = Arc::new(FlakyFactory::new(available.clone()));
let pool = make_pool_with_factory(5, 1, factory);
let conn = pool.acquire().await.unwrap();
pool.release(conn).await;
available.store(false, Ordering::SeqCst);
let mut held = Vec::new();
for _ in 0..1 {
if let Ok(c) = pool.acquire().await {
held.push(c);
}
}
let result = tokio::time::timeout(Duration::from_secs(2), pool.acquire()).await;
assert!(
result.is_err() || result.unwrap().is_err(),
"acquire during outage should fail or timeout"
);
for conn in held {
pool.release(conn).await;
}
available.store(true, Ordering::SeqCst);
let conn = pool.acquire().await;
assert!(conn.is_ok(), "pool should recover after outage");
let status = pool.status().await;
assert!(
status.active <= 5,
"active must not exceed max_size: {}",
status.active
);
}
#[tokio::test]
async fn chaos_close_all_with_borrowed_connections() {
let db = Arc::new(Mutex::new(common::InMemoryDb::new()));
let factory = Arc::new(MockConnectionFactory::new(db));
let pool = make_pool_with_factory(5, 3, factory);
let conn1 = pool.acquire().await.unwrap();
let conn2 = pool.acquire().await.unwrap();
let conn3 = pool.acquire().await.unwrap();
let status = pool.status().await;
assert_eq!(status.active, 3);
assert_eq!(status.idle, 0);
pool.close_all().await;
pool.release(conn1).await;
pool.release(conn2).await;
pool.release(conn3).await;
let status = pool.status().await;
assert_eq!(status.active, 0, "active should be 0 after releasing all");
assert_eq!(status.idle, 0, "idle should be 0 after close_all");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn chaos_factory_persistent_failure_no_leak() {
let available = Arc::new(AtomicBool::new(false)); let factory = Arc::new(FlakyFactory::new(available.clone()));
let pool = make_pool_with_factory(10, 1, factory);
let error_count = Arc::new(AtomicU32::new(0));
let mut handles = Vec::new();
for _ in 0..4 {
let err_count = error_count.clone();
handles.push(tokio::spawn(async move {
for _ in 0..10 {
if pool.acquire().await.is_err() {
err_count.fetch_add(1, Ordering::Relaxed);
}
}
}));
}
for h in handles {
h.await.unwrap();
}
assert_eq!(
error_count.load(Ordering::Relaxed),
40,
"all 40 acquires should fail"
);
let status = pool.status().await;
assert_eq!(
status.active, 0,
"active_count must be 0 after all failures, got {}",
status.active
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn chaos_cascading_failures_consistency() {
let factory = Arc::new(NetworkPartitionFactory::new());
let config = PoolConfigBuilder::new()
.max_size(8)
.min_idle(0)
.max_lifetime(2) .idle_timeout(2)
.acquire_timeout(2)
.build()
.unwrap();
let pool: &'static Pool = Box::leak(Box::new(Pool::new(config, factory.clone()).unwrap()));
let ops = Arc::new(AtomicU64::new(0));
let max_active = Arc::new(AtomicU32::new(0));
let mut handles = Vec::new();
for task_id in 0..4 {
let ops_c = ops.clone();
let max_c = max_active.clone();
let factory_c = factory.clone();
handles.push(tokio::spawn(async move {
for i in 0..30 {
if i == 15 && task_id == 0 {
factory_c.partition_all();
tokio::time::sleep(Duration::from_millis(20)).await;
}
if let Ok(conn) = pool.acquire().await {
let status = pool.status().await;
max_c.fetch_max(status.active, Ordering::Relaxed);
tokio::task::yield_now().await;
pool.release(conn).await;
ops_c.fetch_add(1, Ordering::Relaxed);
}
}
}));
}
for h in handles {
h.await.unwrap();
}
assert!(ops.load(Ordering::Relaxed) > 0);
assert!(
max_active.load(Ordering::Relaxed) <= 8,
"active must never exceed max_size: {}",
max_active.load(Ordering::Relaxed)
);
let status = pool.status().await;
assert_eq!(
status.active, status.idle,
"after cascading failures: active={}, idle={}",
status.active, status.idle
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn chaos_random_failure_injection() {
let available = Arc::new(AtomicBool::new(true));
let factory = Arc::new(FlakyFactory::new(available.clone()));
let pool = make_pool_with_factory(15, 2, factory);
let ops = Arc::new(AtomicU64::new(0));
let max_active = Arc::new(AtomicU32::new(0));
let mut rng = Rng::new(42);
let mut handles = Vec::new();
for task_id in 0..6 {
let ops_c = ops.clone();
let max_c = max_active.clone();
let avail = available.clone();
let seed = rng.next_u64();
handles.push(tokio::spawn(async move {
let mut local_rng = Rng::new(seed);
for _ in 0..40 {
if task_id == 0 && local_rng.next_f64() < 0.05 {
avail.store(false, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(10)).await;
avail.store(true, Ordering::SeqCst);
}
if let Ok(conn) = pool.acquire().await {
let status = pool.status().await;
max_c.fetch_max(status.active, Ordering::Relaxed);
let hold_ms = local_rng.next_usize(5);
tokio::time::sleep(Duration::from_millis(hold_ms as u64)).await;
pool.release(conn).await;
ops_c.fetch_add(1, Ordering::Relaxed);
}
}
}));
}
for h in handles {
h.await.unwrap();
}
assert!(ops.load(Ordering::Relaxed) > 0);
assert!(
max_active.load(Ordering::Relaxed) <= 15,
"active must never exceed max_size: {}",
max_active.load(Ordering::Relaxed)
);
let status = pool.status().await;
assert_eq!(
status.active, status.idle,
"after random failures: active={}, idle={}",
status.active, status.idle
);
}
#[tokio::test]
async fn chaos_transaction_cascading_failure_state_machine() {
let conn = DiskFullConnection::new(4); let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
for _ in 0..4 {
tx.execute("INSERT").await.unwrap();
}
let result = tx.execute("INSERT").await;
assert!(result.is_err(), "5th execute should fail");
assert!(
tx.is_active(),
"tx should still be active after execute failure"
);
let commit_result = tx.commit().await;
assert!(commit_result.is_err(), "commit should fail");
assert!(
tx.is_active(),
"tx should still be active after failed commit"
);
tx.rollback().await.unwrap();
assert!(!tx.is_active(), "tx should be inactive after rollback");
}
#[tokio::test]
async fn chaos_transaction_terminal_state_rejects_operations() {
let db = Arc::new(Mutex::new(common::InMemoryDb::new()));
let conn = common::MockConnection::new(db);
let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
tx.execute("INSERT").await.unwrap();
tx.rollback().await.unwrap();
assert!(tx.execute("SELECT").await.is_err());
assert!(tx.commit().await.is_err());
assert!(tx.rollback().await.is_err());
}