use rusqlite::Connection;
use std::panic::{catch_unwind, AssertUnwindSafe};
use tokio::sync::{mpsc, oneshot};
use khive_storage::error::{StorageError, WriterTaskRequestState};
use crate::error::SqliteError;
use crate::pool::ConnectionPool;
type WriteOp<R> = Box<dyn FnOnce(&Connection) -> Result<R, StorageError> + Send>;
pub struct WriteRequest<R: Send + 'static> {
op: WriteOp<R>,
reply: oneshot::Sender<Result<R, StorageError>>,
top_level: bool,
}
mod sealed {
pub trait Sealed {
fn execute_and_reply_reporting_terminal(
self: Box<Self>,
conn: &rusqlite::Connection,
) -> Option<khive_storage::error::WriterTaskRequestState>;
fn execute_and_reply_top_level_reporting_terminal(
self: Box<Self>,
conn: &rusqlite::Connection,
) -> Option<khive_storage::error::WriterTaskRequestState>;
}
}
pub trait AnyWriteRequest: sealed::Sealed + Send {
fn execute_and_reply(self: Box<Self>, conn: &Connection);
fn execute_and_reply_top_level(self: Box<Self>, conn: &Connection);
fn reply_error(self: Box<Self>, err: StorageError);
fn is_top_level(&self) -> bool;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RollbackDisposition {
RolledBack,
SideEffectsUnknown,
}
fn rollback_after_failure(conn: &Connection, failure_context: &'static str) -> RollbackDisposition {
match conn.execute_batch("ROLLBACK") {
Ok(()) if conn.is_autocommit() => RollbackDisposition::RolledBack,
Ok(()) => {
tracing::error!(
failure_context,
"writer task: ROLLBACK returned success but the connection is still in a \
transaction; request side effects are unknown"
);
RollbackDisposition::SideEffectsUnknown
}
Err(rollback_error) => {
tracing::error!(
error = %rollback_error,
failure_context,
"writer task: rollback after request failure failed; request side effects are \
unknown"
);
RollbackDisposition::SideEffectsUnknown
}
}
}
impl<R: Send + 'static> sealed::Sealed for WriteRequest<R> {
fn execute_and_reply_reporting_terminal(
self: Box<Self>,
conn: &Connection,
) -> Option<WriterTaskRequestState> {
let WriteRequest { op, reply, .. } = *self;
match catch_unwind(AssertUnwindSafe(|| op(conn))) {
Ok(Ok(value)) => match conn.execute_batch("COMMIT") {
Ok(()) if conn.is_autocommit() => {
let _ = reply.send(Ok(value));
None
}
Ok(()) => {
tracing::error!(
"writer task: COMMIT returned success but the connection is still in a \
transaction; request side effects are unknown"
);
let request_state = WriterTaskRequestState::SideEffectsUnknown;
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
Err(commit_error) => match rollback_after_failure(conn, "commit failure") {
RollbackDisposition::RolledBack => {
let _ = reply.send(Err(StorageError::Pool {
operation: "writer_task_commit".into(),
message: commit_error.to_string(),
}));
None
}
RollbackDisposition::SideEffectsUnknown => {
let request_state = WriterTaskRequestState::SideEffectsUnknown;
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
},
},
Ok(Err(operation_error)) => {
match rollback_after_failure(conn, "request operation failure") {
RollbackDisposition::RolledBack => {
let _ = reply.send(Err(operation_error));
None
}
RollbackDisposition::SideEffectsUnknown => {
let request_state = WriterTaskRequestState::SideEffectsUnknown;
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
}
}
Err(_panic_payload) => {
let request_state = match rollback_after_failure(conn, "request panic") {
RollbackDisposition::RolledBack => {
WriterTaskRequestState::TransactionRolledBack
}
RollbackDisposition::SideEffectsUnknown => {
WriterTaskRequestState::SideEffectsUnknown
}
};
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
}
}
fn execute_and_reply_top_level_reporting_terminal(
self: Box<Self>,
conn: &Connection,
) -> Option<WriterTaskRequestState> {
let WriteRequest { op, reply, .. } = *self;
match catch_unwind(AssertUnwindSafe(|| op(conn))) {
Ok(outcome) if conn.is_autocommit() => {
let _ = reply.send(outcome);
None
}
Ok(_outcome) => {
tracing::error!(
"writer task: top-level request returned with an open transaction; request \
side effects are unknown"
);
let request_state = WriterTaskRequestState::SideEffectsUnknown;
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
Err(_panic_payload) => {
let request_state = WriterTaskRequestState::SideEffectsUnknown;
let _ = reply.send(Err(writer_task_terminated(request_state)));
Some(request_state)
}
}
}
}
impl<R: Send + 'static> AnyWriteRequest for WriteRequest<R> {
fn execute_and_reply(self: Box<Self>, conn: &Connection) {
let _ = sealed::Sealed::execute_and_reply_reporting_terminal(self, conn);
}
fn execute_and_reply_top_level(self: Box<Self>, conn: &Connection) {
let _ = sealed::Sealed::execute_and_reply_top_level_reporting_terminal(self, conn);
}
fn reply_error(self: Box<Self>, err: StorageError) {
let _ = self.reply.send(Err(err));
}
fn is_top_level(&self) -> bool {
self.top_level
}
}
fn writer_task_terminated(request_state: WriterTaskRequestState) -> StorageError {
StorageError::WriterTaskTerminated { request_state }
}
#[derive(Clone, Debug)]
pub struct WriterTaskHandle {
tx: mpsc::Sender<Box<dyn AnyWriteRequest + Send>>,
}
impl WriterTaskHandle {
async fn enqueue<R, F>(
&self,
op: F,
) -> Result<oneshot::Receiver<Result<R, StorageError>>, StorageError>
where
R: Send + 'static,
F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
{
self.enqueue_inner(op, false).await
}
async fn enqueue_inner<R, F>(
&self,
op: F,
top_level: bool,
) -> Result<oneshot::Receiver<Result<R, StorageError>>, StorageError>
where
R: Send + 'static,
F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
{
let (reply_tx, reply_rx) = oneshot::channel();
let request = WriteRequest {
op: Box::new(op),
reply: reply_tx,
top_level,
};
self.tx
.send(Box::new(request))
.await
.map_err(|_| writer_task_terminated(WriterTaskRequestState::NotStarted))?;
Ok(reply_rx)
}
pub async fn send<R, F>(&self, op: F) -> Result<R, StorageError>
where
R: Send + 'static,
F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
{
let reply_rx = self.enqueue(op).await?;
reply_rx
.await
.map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
}
pub async fn send_with_timeout<R, F>(
&self,
op: F,
timeout: std::time::Duration,
) -> Result<R, StorageError>
where
R: Send + 'static,
F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
{
let reply_rx = match tokio::time::timeout(timeout, self.enqueue(op)).await {
Ok(Ok(reply_rx)) => reply_rx,
Ok(Err(e)) => return Err(e),
Err(_elapsed) => {
return Err(StorageError::WriteQueueFull {
timeout_ms: timeout.as_millis() as u64,
})
}
};
reply_rx
.await
.map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
}
pub async fn send_top_level<R, F>(&self, op: F) -> Result<R, StorageError>
where
R: Send + 'static,
F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
{
let reply_rx = self.enqueue_inner(op, true).await?;
reply_rx
.await
.map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
}
pub fn queue_depth(&self) -> usize {
self.tx.max_capacity() - self.tx.capacity()
}
pub fn capacity(&self) -> usize {
self.tx.max_capacity()
}
}
pub fn spawn(pool: &ConnectionPool, capacity: usize) -> Result<WriterTaskHandle, SqliteError> {
let conn = pool.open_standalone_writer()?;
let origin = pool.origin();
let (tx, rx) = mpsc::channel(capacity.max(1));
tokio::spawn(run_writer_task(conn, rx, origin));
Ok(WriterTaskHandle { tx })
}
async fn close_and_fail_queued_requests(rx: &mut mpsc::Receiver<Box<dyn AnyWriteRequest + Send>>) {
rx.close();
while let Some(request) = rx.recv().await {
request.reply_error(writer_task_terminated(WriterTaskRequestState::NotStarted));
}
}
async fn run_writer_task(
mut conn: Connection,
mut rx: mpsc::Receiver<Box<dyn AnyWriteRequest + Send>>,
origin: khive_storage::tx_registry::TxOrigin,
) {
while let Some(request) = rx.recv().await {
let origin = origin.clone();
let outcome = tokio::task::spawn_blocking(move || {
if !conn.is_autocommit() {
tracing::error!(
"writer task: connection is not in autocommit mode before request dispatch; \
retiring the poisoned writer without running the request"
);
let request_state = WriterTaskRequestState::NotStarted;
request.reply_error(writer_task_terminated(request_state));
return (conn, Some(request_state));
}
let terminal_state = if request.is_top_level() {
sealed::Sealed::execute_and_reply_top_level_reporting_terminal(request, &conn)
} else {
let _tx_handle = khive_storage::tx_registry::register_scoped(
Some("writer_task_tx".to_string()),
origin,
);
match conn.execute_batch("BEGIN IMMEDIATE") {
Ok(()) => sealed::Sealed::execute_and_reply_reporting_terminal(request, &conn),
Err(e) => {
tracing::warn!(
error = %e,
"writer task: BEGIN IMMEDIATE failed; replying an \
error without running the request's operation"
);
request.reply_error(StorageError::Pool {
operation: "writer_task_begin".into(),
message: e.to_string(),
});
None
}
}
};
(conn, terminal_state)
})
.await;
match outcome {
Ok((returned_conn, None)) => conn = returned_conn,
Ok((_returned_conn, Some(request_state))) => {
tracing::error!(
request_state = %request_state,
"writer task reached a terminal request or connection state; closing and \
failing the queue without restarting"
);
close_and_fail_queued_requests(&mut rx).await;
return;
}
Err(join_err) => {
tracing::error!(
error = %join_err,
"writer task blocking closure failed outside the request \
panic boundary; closing and failing the queue without restarting"
);
close_and_fail_queued_requests(&mut rx).await;
return;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pool::PoolConfig;
use rusqlite::hooks::{AuthAction, AuthContext, Authorization, TransactionOperation};
use serial_test::serial;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc as std_mpsc;
use std::sync::Arc;
use std::time::Duration;
fn file_pool(path: &std::path::Path) -> ConnectionPool {
let cfg = PoolConfig {
path: Some(path.to_path_buf()),
..PoolConfig::default()
};
ConnectionPool::new(cfg).expect("pool open")
}
fn deny_commit_and_rollback(ctx: AuthContext<'_>) -> Authorization {
match ctx.action {
AuthAction::Transaction {
operation: TransactionOperation::Unknown | TransactionOperation::Rollback,
} => Authorization::Deny,
_ => Authorization::Allow,
}
}
fn deny_commit(ctx: AuthContext<'_>) -> Authorization {
match ctx.action {
AuthAction::Transaction {
operation: TransactionOperation::Unknown,
} => Authorization::Deny,
_ => Authorization::Allow,
}
}
fn deny_rollback(ctx: AuthContext<'_>) -> Authorization {
match ctx.action {
AuthAction::Transaction {
operation: TransactionOperation::Rollback,
} => Authorization::Deny,
_ => Authorization::Allow,
}
}
fn assert_writer_task_terminal_state<T: std::fmt::Debug>(
result: Result<T, StorageError>,
expected: WriterTaskRequestState,
) {
match result {
Err(StorageError::WriterTaskTerminated { request_state }) => {
assert_eq!(request_state, expected)
}
other => panic!("expected WriterTaskTerminated({expected:?}), got {other:?}"),
}
}
#[tokio::test]
#[serial(tx_registry)]
async fn begin_immediate_failure_replies_error_without_running_op() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_begin_failure.db");
let cfg = PoolConfig {
path: Some(path.clone()),
busy_timeout: Duration::from_millis(150),
..PoolConfig::default()
};
let pool = ConnectionPool::new(cfg).unwrap();
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
let lock_holder = pool.try_writer().unwrap();
lock_holder.conn().execute_batch("BEGIN IMMEDIATE").unwrap();
let op_ran = Arc::new(AtomicBool::new(false));
let op_ran_clone = Arc::clone(&op_ran);
let result = handle
.send(move |conn| {
op_ran_clone.store(true, Ordering::SeqCst);
conn.execute("INSERT INTO t (id, v) VALUES (99, 'should-not-land')", [])
.map_err(|e| StorageError::Pool {
operation: "test_insert".into(),
message: e.to_string(),
})
})
.await;
assert!(
matches!(
&result,
Err(StorageError::Pool { operation, .. }) if operation == "writer_task_begin"
),
"expected a writer_task_begin Pool error on BEGIN IMMEDIATE \
failure, got {result:?}"
);
assert!(
!op_ran.load(Ordering::SeqCst),
"the request's operation closure must never run when BEGIN \
IMMEDIATE fails — running it would land a partial write in \
autocommit mode for a request the caller is told failed"
);
lock_holder.conn().execute_batch("ROLLBACK").unwrap();
drop(lock_holder);
let reader = pool.reader().expect("reader");
let count: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 99", [], |row| row.get(0))
.unwrap();
assert_eq!(
count, 0,
"no row must have landed from the request whose BEGIN IMMEDIATE failed"
);
}
#[tokio::test]
#[serial(tx_registry)]
async fn writer_task_executes_op_and_commits() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_commit.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
let affected = handle
.send(|conn| {
conn.execute("INSERT INTO t (id, v) VALUES (1, 'hello')", [])
.map_err(|e| StorageError::Pool {
operation: "test_insert".into(),
message: e.to_string(),
})
})
.await
.expect("op should succeed");
assert_eq!(affected, 1);
let reader = pool.reader().expect("reader");
let v: String = reader
.conn()
.query_row("SELECT v FROM t WHERE id = 1", [], |row| row.get(0))
.expect("row must be committed and visible to a reader");
assert_eq!(v, "hello");
}
#[test]
fn spawn_fails_on_in_memory_pool() {
let cfg = PoolConfig {
path: None,
..PoolConfig::default()
};
let pool = ConnectionPool::new(cfg).unwrap();
let result = spawn(&pool, 8);
assert!(
result.is_err(),
"in-memory pools must reject spawn, not panic"
);
}
#[tokio::test]
async fn full_channel_applies_backpressure_not_immediate_error() {
let (tx, _rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
let handle = WriterTaskHandle { tx };
let first = tokio::spawn({
let handle = handle.clone();
async move {
let _ = handle.send(|_conn| Ok::<(), StorageError>(())).await;
}
});
tokio::time::sleep(Duration::from_millis(20)).await;
let second = tokio::time::timeout(
Duration::from_millis(100),
handle.send(|_conn| Ok::<(), StorageError>(())),
)
.await;
assert!(
second.is_err(),
"a full channel must apply backpressure (send suspends) rather \
than erroring immediately — no try_send escape hatch per ADR-067"
);
first.abort();
}
#[tokio::test]
async fn send_with_timeout_maps_full_channel_to_write_queue_full() {
let (tx, _rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
let handle = WriterTaskHandle { tx };
let first = tokio::spawn({
let handle = handle.clone();
async move {
let _ = handle.send(|_conn| Ok::<(), StorageError>(())).await;
}
});
tokio::time::sleep(Duration::from_millis(20)).await;
let result = handle
.send_with_timeout(
|_conn| Ok::<(), StorageError>(()),
Duration::from_millis(50),
)
.await;
match result {
Err(StorageError::WriteQueueFull { timeout_ms }) => assert_eq!(timeout_ms, 50),
other => panic!("expected WriteQueueFull, got {other:?}"),
}
first.abort();
}
#[tokio::test]
#[serial(tx_registry)]
async fn send_with_timeout_returns_op_result_when_op_outlives_the_timeout() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_slow_op.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
let result = handle
.send_with_timeout(
|conn| {
std::thread::sleep(Duration::from_millis(150));
conn.execute("INSERT INTO t (id, v) VALUES (1, 'slow')", [])
.map_err(|e| StorageError::Pool {
operation: "test_insert".into(),
message: e.to_string(),
})
},
Duration::from_millis(20),
)
.await;
let affected = result.expect(
"an accepted request must return its real result even when the \
op takes longer than the enqueue timeout, not WriteQueueFull",
);
assert_eq!(affected, 1);
let reader = pool.reader().expect("reader");
let v: String = reader
.conn()
.query_row("SELECT v FROM t WHERE id = 1", [], |row| row.get(0))
.expect("the slow op's write must have committed");
assert_eq!(v, "slow");
}
#[tokio::test]
#[serial(tx_registry)]
async fn operation_failure_with_successful_rollback_preserves_error_and_writer_continues() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_operation_rollback.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task spawn");
let original_error = handle
.send(|conn| -> Result<(), StorageError> {
conn.execute("INSERT INTO t (id, v) VALUES (1, 'rolled-back')", [])
.map_err(|e| StorageError::Pool {
operation: "test_operation_error_insert".into(),
message: e.to_string(),
})?;
Err(StorageError::Internal(
"intentional operation failure".into(),
))
})
.await;
assert!(
matches!(
&original_error,
Err(StorageError::Internal(message))
if message == "intentional operation failure"
),
"a confirmed rollback must preserve the operation error, got {original_error:?}"
);
let affected = handle
.send(|conn| {
conn.execute("INSERT INTO t (id, v) VALUES (2, 'committed')", [])
.map_err(|e| StorageError::Pool {
operation: "test_operation_error_followup_insert".into(),
message: e.to_string(),
})
})
.await
.expect("the writer must continue after a confirmed rollback");
assert_eq!(affected, 1);
let reader = pool.reader().expect("reader");
let rolled_back: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 1", [], |row| row.get(0))
.unwrap();
let committed: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 2", [], |row| row.get(0))
.unwrap();
assert_eq!(rolled_back, 0);
assert_eq!(committed, 1);
}
#[tokio::test]
#[serial(tx_registry)]
async fn commit_failure_with_successful_rollback_preserves_error_and_writer_continues() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_commit_rollback.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task spawn");
let commit_error = handle
.send(|conn| -> Result<usize, StorageError> {
let affected = conn
.execute("INSERT INTO t (id, v) VALUES (1, 'rolled-back')", [])
.map_err(|e| StorageError::Pool {
operation: "test_commit_error_insert".into(),
message: e.to_string(),
})?;
conn.authorizer(Some(deny_commit))
.map_err(|e| StorageError::Pool {
operation: "test_install_authorizer".into(),
message: e.to_string(),
})?;
Ok(affected)
})
.await;
assert!(
matches!(
&commit_error,
Err(StorageError::Pool { operation, .. })
if operation == "writer_task_commit"
),
"a confirmed rollback must preserve the commit error, got {commit_error:?}"
);
assert!(
commit_error
.as_ref()
.expect_err("COMMIT must be denied")
.is_retryable(),
"the existing retryable commit-error contract must remain unchanged after a \
confirmed rollback"
);
let affected = handle
.send(|conn| {
conn.authorizer(None::<fn(AuthContext<'_>) -> Authorization>)
.map_err(|e| StorageError::Pool {
operation: "test_remove_authorizer".into(),
message: e.to_string(),
})?;
conn.execute("INSERT INTO t (id, v) VALUES (2, 'committed')", [])
.map_err(|e| StorageError::Pool {
operation: "test_commit_error_followup_insert".into(),
message: e.to_string(),
})
})
.await
.expect("the writer must continue after the failed COMMIT is rolled back");
assert_eq!(affected, 1);
let reader = pool.reader().expect("reader");
let rolled_back: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 1", [], |row| row.get(0))
.unwrap();
let committed: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 2", [], |row| row.get(0))
.unwrap();
assert_eq!(rolled_back, 0);
assert_eq!(committed, 1);
}
#[test]
fn top_level_request_returning_with_open_transaction_reports_side_effects_unknown() {
let conn = Connection::open_in_memory().expect("in-memory connection");
conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY)")
.unwrap();
let (reply_tx, mut reply_rx) = oneshot::channel();
let request = WriteRequest {
op: Box::new(|conn| -> Result<usize, StorageError> {
conn.execute_batch("BEGIN IMMEDIATE")
.map_err(|e| StorageError::Pool {
operation: "test_top_level_begin".into(),
message: e.to_string(),
})?;
conn.execute("INSERT INTO t (id) VALUES (1)", [])
.map_err(|e| StorageError::Pool {
operation: "test_top_level_insert".into(),
message: e.to_string(),
})
}),
reply: reply_tx,
top_level: true,
};
let terminal_state = sealed::Sealed::execute_and_reply_top_level_reporting_terminal(
Box::new(request),
&conn,
);
assert_eq!(
terminal_state,
Some(WriterTaskRequestState::SideEffectsUnknown)
);
let reply = reply_rx
.try_recv()
.expect("active request must receive a typed terminal reply");
assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
assert!(
!conn.is_autocommit(),
"the fixture must prove the post-request autocommit check observed an open transaction"
);
}
#[test]
fn commit_failure_with_failed_rollback_reports_side_effects_unknown() {
let conn = Connection::open_in_memory().expect("in-memory connection");
conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
.unwrap();
let (reply_tx, mut reply_rx) = oneshot::channel();
let request = WriteRequest {
op: Box::new(|conn| -> Result<usize, StorageError> {
let affected = conn
.execute("INSERT INTO t (id) VALUES (1)", [])
.map_err(|e| StorageError::Pool {
operation: "test_insert_before_commit_failure".into(),
message: e.to_string(),
})?;
conn.authorizer(Some(deny_commit_and_rollback))
.map_err(|e| StorageError::Pool {
operation: "test_install_authorizer".into(),
message: e.to_string(),
})?;
Ok(affected)
}),
reply: reply_tx,
top_level: false,
};
let terminal_state =
sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
assert_eq!(
terminal_state,
Some(WriterTaskRequestState::SideEffectsUnknown)
);
let reply = reply_rx
.try_recv()
.expect("active request must receive a typed terminal reply");
assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
assert!(
!conn.is_autocommit(),
"the denied COMMIT and ROLLBACK must leave the test connection poisoned"
);
}
#[tokio::test]
#[serial(tx_registry)]
async fn poisoned_connection_retires_before_queued_top_level_request() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_rollback_poison.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task spawn");
let (started_tx, started_rx) = oneshot::channel::<()>();
let (release_tx, release_rx) = std_mpsc::channel::<()>();
let active = tokio::spawn({
let handle = handle.clone();
async move {
handle
.send(move |conn| -> Result<usize, StorageError> {
let affected = conn
.execute("INSERT INTO t (id, v) VALUES (1, 'active')", [])
.map_err(|e| StorageError::Pool {
operation: "test_active_insert".into(),
message: e.to_string(),
})?;
conn.authorizer(Some(deny_commit_and_rollback))
.map_err(|e| StorageError::Pool {
operation: "test_install_authorizer".into(),
message: e.to_string(),
})?;
let _ = started_tx.send(());
release_rx.recv().expect("test must release active op");
Ok(affected)
})
.await
}
});
tokio::time::timeout(Duration::from_secs(5), started_rx)
.await
.expect("active request did not start")
.expect("active request dropped its start signal");
let queued_ran = Arc::new(AtomicBool::new(false));
let queued_ran_in_op = Arc::clone(&queued_ran);
let queued_top_level = handle
.enqueue_inner(
move |conn| {
queued_ran_in_op.store(true, Ordering::SeqCst);
conn.execute("INSERT INTO t (id, v) VALUES (2, 'queued')", [])
.map_err(|e| StorageError::Pool {
operation: "test_queued_top_level_insert".into(),
message: e.to_string(),
})
},
true,
)
.await
.expect("top-level request must queue behind active request");
release_tx.send(()).expect("release active op");
let active_result = tokio::time::timeout(Duration::from_secs(5), active)
.await
.expect("active caller hung after rollback failure")
.expect("active caller task join");
assert_writer_task_terminal_state(
active_result,
WriterTaskRequestState::SideEffectsUnknown,
);
let queued_result = tokio::time::timeout(Duration::from_secs(5), queued_top_level)
.await
.expect("queued top-level caller hung after terminal failure")
.expect("terminal drain must preserve queued typed reply");
assert_writer_task_terminal_state(queued_result, WriterTaskRequestState::NotStarted);
assert!(
!queued_ran.load(Ordering::SeqCst),
"a top-level request must never run on the poisoned connection"
);
let future_ran = Arc::new(AtomicBool::new(false));
let future_ran_in_op = Arc::clone(&future_ran);
let future_result = handle
.send_top_level(move |_conn| {
future_ran_in_op.store(true, Ordering::SeqCst);
Ok::<(), StorageError>(())
})
.await;
assert_writer_task_terminal_state(future_result, WriterTaskRequestState::NotStarted);
assert!(!future_ran.load(Ordering::SeqCst));
}
#[test]
fn operation_failure_with_failed_rollback_reports_side_effects_unknown() {
let conn = Connection::open_in_memory().expect("in-memory connection");
conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
.unwrap();
let (reply_tx, mut reply_rx) = oneshot::channel();
let request = WriteRequest {
op: Box::new(|conn| -> Result<(), StorageError> {
conn.authorizer(Some(deny_rollback))
.map_err(|e| StorageError::Pool {
operation: "test_install_authorizer".into(),
message: e.to_string(),
})?;
Err(StorageError::Internal(
"intentional operation failure before denied rollback".into(),
))
}),
reply: reply_tx,
top_level: false,
};
let terminal_state =
sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
assert_eq!(
terminal_state,
Some(WriterTaskRequestState::SideEffectsUnknown)
);
let reply = reply_rx
.try_recv()
.expect("active request must receive a typed terminal reply");
assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
assert!(
!conn.is_autocommit(),
"the denied ROLLBACK must leave the test connection poisoned"
);
}
#[test]
fn wrapped_panic_with_failed_rollback_reports_side_effects_unknown() {
let conn = Connection::open_in_memory().expect("in-memory connection");
conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
.unwrap();
let (reply_tx, mut reply_rx) = oneshot::channel();
let request = WriteRequest {
op: Box::new(|conn| -> Result<(), StorageError> {
conn.execute_batch("INSERT INTO t (id) VALUES (1); COMMIT")
.map_err(|e| StorageError::Pool {
operation: "test_force_rollback_failure".into(),
message: e.to_string(),
})?;
panic!("intentional panic after illicit commit");
}),
reply: reply_tx,
top_level: false,
};
let terminal_state =
sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
assert_eq!(
terminal_state,
Some(WriterTaskRequestState::SideEffectsUnknown)
);
let reply = reply_rx
.try_recv()
.expect("active request must receive a typed terminal reply");
assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
let count: i64 = conn
.query_row("SELECT COUNT(*) FROM t", [], |row| row.get(0))
.unwrap();
assert_eq!(
count, 1,
"the fixture's committed side effect proves why the state must be unknown"
);
}
#[tokio::test]
#[serial(tx_registry)]
async fn wrapped_panic_rolls_back_and_terminally_fails_queue() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_wrapped_panic.db");
let cfg = PoolConfig {
path: Some(path),
write_queue_enabled: true,
write_queue_capacity: 8,
..PoolConfig::default()
};
let pool = ConnectionPool::new(cfg).unwrap();
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = pool
.writer_task_handle()
.expect("writer task lookup")
.expect("file-backed queued pool must spawn its writer task");
assert_eq!(pool.writer_task_spawn_count(), 1);
let (started_tx, started_rx) = oneshot::channel::<()>();
let (release_tx, release_rx) = std_mpsc::channel::<()>();
let active = tokio::spawn({
let handle = handle.clone();
async move {
handle
.send(move |conn| -> Result<usize, StorageError> {
conn.execute("INSERT INTO t (id, v) VALUES (1, 'active')", [])
.map_err(|e| StorageError::Pool {
operation: "test_active_insert".into(),
message: e.to_string(),
})?;
let _ = started_tx.send(());
release_rx.recv().expect("test must release active op");
panic!("intentional wrapped writer request panic");
})
.await
}
});
tokio::time::timeout(Duration::from_secs(5), started_rx)
.await
.expect("active request did not start")
.expect("active request dropped its start signal");
let queued_one_ran = Arc::new(AtomicBool::new(false));
let queued_one_ran_in_op = Arc::clone(&queued_one_ran);
let queued_one = handle
.enqueue(move |conn| {
queued_one_ran_in_op.store(true, Ordering::SeqCst);
conn.execute("INSERT INTO t (id, v) VALUES (2, 'queued-one')", [])
.map_err(|e| StorageError::Pool {
operation: "test_queued_one_insert".into(),
message: e.to_string(),
})
})
.await
.expect("first queued request must be accepted");
let queued_two_ran = Arc::new(AtomicBool::new(false));
let queued_two_ran_in_op = Arc::clone(&queued_two_ran);
let queued_two = handle
.enqueue(move |_conn| {
queued_two_ran_in_op.store(true, Ordering::SeqCst);
Ok::<String, StorageError>("queued-two-ran".to_string())
})
.await
.expect("second queued request must be accepted");
assert_eq!(
handle.queue_depth(),
2,
"both heterogeneous requests must be buffered behind the active op"
);
release_tx.send(()).expect("release active op");
let active_result = tokio::time::timeout(Duration::from_secs(5), active)
.await
.expect("active caller hung after panic")
.expect("active caller task join");
assert_writer_task_terminal_state(
active_result,
WriterTaskRequestState::TransactionRolledBack,
);
let queued_one_result = tokio::time::timeout(Duration::from_secs(5), queued_one)
.await
.expect("first queued caller hung after terminal failure")
.expect("terminal drain must preserve first typed reply");
assert_writer_task_terminal_state(queued_one_result, WriterTaskRequestState::NotStarted);
let queued_two_result = tokio::time::timeout(Duration::from_secs(5), queued_two)
.await
.expect("second queued caller hung after terminal failure")
.expect("terminal drain must preserve second typed reply");
assert_writer_task_terminal_state(queued_two_result, WriterTaskRequestState::NotStarted);
assert!(!queued_one_ran.load(Ordering::SeqCst));
assert!(!queued_two_ran.load(Ordering::SeqCst));
let future_ran = Arc::new(AtomicBool::new(false));
let future_ran_in_op = Arc::clone(&future_ran);
let future_result = handle
.send(move |_conn| {
future_ran_in_op.store(true, Ordering::SeqCst);
Ok::<(), StorageError>(())
})
.await;
assert_writer_task_terminal_state(future_result, WriterTaskRequestState::NotStarted);
assert!(!future_ran.load(Ordering::SeqCst));
let cached_after_failure = pool
.writer_task_handle()
.expect("cached writer task lookup")
.expect("pool retains its terminal handle");
assert_eq!(
pool.writer_task_spawn_count(),
1,
"a terminal writer task must not be restarted behind callers' backs"
);
let cached_result = cached_after_failure
.send(|_conn| Ok::<(), StorageError>(()))
.await;
assert_writer_task_terminal_state(cached_result, WriterTaskRequestState::NotStarted);
let reader = pool.reader().expect("reader");
let count: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t", [], |row| row.get(0))
.unwrap();
assert_eq!(
count, 0,
"the active transaction must be rolled back and queued ops must never run"
);
}
#[tokio::test]
async fn top_level_panic_reports_unknown_and_fails_queue_without_running_it() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("writer_task_top_level_panic.db");
let pool = file_pool(&path);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
.unwrap();
}
let handle = spawn(&pool, 8).expect("writer task spawn");
let (started_tx, started_rx) = oneshot::channel::<()>();
let (release_tx, release_rx) = std_mpsc::channel::<()>();
let active = tokio::spawn({
let handle = handle.clone();
async move {
handle
.send_top_level(move |conn| -> Result<usize, StorageError> {
conn.execute("INSERT INTO t (id, v) VALUES (10, 'autocommitted')", [])
.map_err(|e| StorageError::Pool {
operation: "test_top_level_insert".into(),
message: e.to_string(),
})?;
let _ = started_tx.send(());
release_rx.recv().expect("test must release top-level op");
panic!("intentional top-level writer request panic");
})
.await
}
});
tokio::time::timeout(Duration::from_secs(5), started_rx)
.await
.expect("top-level request did not start")
.expect("top-level request dropped its start signal");
let queued_ran = Arc::new(AtomicBool::new(false));
let queued_ran_in_op = Arc::clone(&queued_ran);
let queued = handle
.enqueue(move |conn| {
queued_ran_in_op.store(true, Ordering::SeqCst);
conn.execute("INSERT INTO t (id, v) VALUES (11, 'queued')", [])
.map_err(|e| StorageError::Pool {
operation: "test_top_level_queued_insert".into(),
message: e.to_string(),
})
})
.await
.expect("queued request must be accepted");
assert_eq!(handle.queue_depth(), 1);
release_tx.send(()).expect("release top-level op");
let active_result = tokio::time::timeout(Duration::from_secs(5), active)
.await
.expect("top-level caller hung after panic")
.expect("top-level caller task join");
assert_writer_task_terminal_state(
active_result,
WriterTaskRequestState::SideEffectsUnknown,
);
let queued_result = tokio::time::timeout(Duration::from_secs(5), queued)
.await
.expect("queued caller hung after top-level panic")
.expect("terminal drain must preserve queued typed reply");
assert_writer_task_terminal_state(queued_result, WriterTaskRequestState::NotStarted);
assert!(!queued_ran.load(Ordering::SeqCst));
let reader = pool.reader().expect("reader");
let active_count: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 10", [], |row| row.get(0))
.unwrap();
let queued_count: i64 = reader
.conn()
.query_row("SELECT COUNT(*) FROM t WHERE id = 11", [], |row| row.get(0))
.unwrap();
assert_eq!(
active_count, 1,
"the completed top-level statement autocommits before the panic"
);
assert_eq!(queued_count, 0, "the queued request must never run");
}
#[tokio::test]
async fn closed_receiver_rejects_all_send_surfaces_as_not_started() {
let (tx, rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(4);
drop(rx);
let handle = WriterTaskHandle { tx };
let send_result = handle.send(|_conn| Ok::<(), StorageError>(())).await;
assert_writer_task_terminal_state(send_result, WriterTaskRequestState::NotStarted);
let timed_result = handle
.send_with_timeout(|_conn| Ok::<(), StorageError>(()), Duration::from_secs(1))
.await;
assert_writer_task_terminal_state(timed_result, WriterTaskRequestState::NotStarted);
let top_level_result = handle
.send_top_level(|_conn| Ok::<(), StorageError>(()))
.await;
assert_writer_task_terminal_state(top_level_result, WriterTaskRequestState::NotStarted);
}
#[tokio::test]
async fn accepted_request_lost_reply_is_side_effects_unknown() {
let (tx, mut rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
let handle = WriterTaskHandle { tx };
let request_ran = Arc::new(AtomicBool::new(false));
let request_ran_in_op = Arc::clone(&request_ran);
let dropper = tokio::spawn(async move {
let request = rx.recv().await.expect("request must be accepted");
drop(request);
});
let result = tokio::time::timeout(
Duration::from_secs(5),
handle.send(move |_conn| {
request_ran_in_op.store(true, Ordering::SeqCst);
Ok::<(), StorageError>(())
}),
)
.await
.expect("caller hung after accepted request was dropped");
dropper.await.expect("dropper task join");
assert_writer_task_terminal_state(result, WriterTaskRequestState::SideEffectsUnknown);
assert!(!request_ran.load(Ordering::SeqCst));
}
}