1use rusqlite::Connection;
18use std::panic::{catch_unwind, AssertUnwindSafe};
19use tokio::sync::{mpsc, oneshot};
20
21use khive_storage::error::{StorageError, WriterTaskRequestState};
22
23use crate::error::SqliteError;
24use crate::pool::ConnectionPool;
25
26type WriteOp<R> = Box<dyn FnOnce(&Connection) -> Result<R, StorageError> + Send>;
36
37pub struct WriteRequest<R: Send + 'static> {
52 op: WriteOp<R>,
53 reply: oneshot::Sender<Result<R, StorageError>>,
54 top_level: bool,
55}
56
57mod sealed {
58 pub trait Sealed {
63 fn execute_and_reply_reporting_terminal(
64 self: Box<Self>,
65 conn: &rusqlite::Connection,
66 ) -> Option<khive_storage::error::WriterTaskRequestState>;
67
68 fn execute_and_reply_top_level_reporting_terminal(
69 self: Box<Self>,
70 conn: &rusqlite::Connection,
71 ) -> Option<khive_storage::error::WriterTaskRequestState>;
72 }
73}
74
75pub trait AnyWriteRequest: sealed::Sealed + Send {
81 fn execute_and_reply(self: Box<Self>, conn: &Connection);
94
95 fn execute_and_reply_top_level(self: Box<Self>, conn: &Connection);
105
106 fn reply_error(self: Box<Self>, err: StorageError);
118
119 fn is_top_level(&self) -> bool;
123}
124
125#[derive(Debug, Clone, Copy, PartialEq, Eq)]
126enum RollbackDisposition {
127 RolledBack,
128 SideEffectsUnknown,
129}
130
131fn rollback_after_failure(conn: &Connection, failure_context: &'static str) -> RollbackDisposition {
134 match conn.execute_batch("ROLLBACK") {
135 Ok(()) if conn.is_autocommit() => RollbackDisposition::RolledBack,
136 Ok(()) => {
137 tracing::error!(
138 failure_context,
139 "writer task: ROLLBACK returned success but the connection is still in a \
140 transaction; request side effects are unknown"
141 );
142 RollbackDisposition::SideEffectsUnknown
143 }
144 Err(rollback_error) => {
145 tracing::error!(
146 error = %rollback_error,
147 failure_context,
148 "writer task: rollback after request failure failed; request side effects are \
149 unknown"
150 );
151 RollbackDisposition::SideEffectsUnknown
152 }
153 }
154}
155
156impl<R: Send + 'static> sealed::Sealed for WriteRequest<R> {
157 fn execute_and_reply_reporting_terminal(
158 self: Box<Self>,
159 conn: &Connection,
160 ) -> Option<WriterTaskRequestState> {
161 let WriteRequest { op, reply, .. } = *self;
165 match catch_unwind(AssertUnwindSafe(|| op(conn))) {
166 Ok(Ok(value)) => match conn.execute_batch("COMMIT") {
167 Ok(()) if conn.is_autocommit() => {
168 let _ = reply.send(Ok(value));
171 None
172 }
173 Ok(()) => {
174 tracing::error!(
175 "writer task: COMMIT returned success but the connection is still in a \
176 transaction; request side effects are unknown"
177 );
178 let request_state = WriterTaskRequestState::SideEffectsUnknown;
179 let _ = reply.send(Err(writer_task_terminated(request_state)));
180 Some(request_state)
181 }
182 Err(commit_error) => match rollback_after_failure(conn, "commit failure") {
183 RollbackDisposition::RolledBack => {
184 let _ = reply.send(Err(StorageError::Pool {
185 operation: "writer_task_commit".into(),
186 message: commit_error.to_string(),
187 }));
188 None
189 }
190 RollbackDisposition::SideEffectsUnknown => {
191 let request_state = WriterTaskRequestState::SideEffectsUnknown;
192 let _ = reply.send(Err(writer_task_terminated(request_state)));
193 Some(request_state)
194 }
195 },
196 },
197 Ok(Err(operation_error)) => {
198 match rollback_after_failure(conn, "request operation failure") {
199 RollbackDisposition::RolledBack => {
200 let _ = reply.send(Err(operation_error));
201 None
202 }
203 RollbackDisposition::SideEffectsUnknown => {
204 let request_state = WriterTaskRequestState::SideEffectsUnknown;
205 let _ = reply.send(Err(writer_task_terminated(request_state)));
206 Some(request_state)
207 }
208 }
209 }
210 Err(_panic_payload) => {
211 let request_state = match rollback_after_failure(conn, "request panic") {
215 RollbackDisposition::RolledBack => {
216 WriterTaskRequestState::TransactionRolledBack
217 }
218 RollbackDisposition::SideEffectsUnknown => {
219 WriterTaskRequestState::SideEffectsUnknown
220 }
221 };
222 let _ = reply.send(Err(writer_task_terminated(request_state)));
223 Some(request_state)
224 }
225 }
226 }
227
228 fn execute_and_reply_top_level_reporting_terminal(
229 self: Box<Self>,
230 conn: &Connection,
231 ) -> Option<WriterTaskRequestState> {
232 let WriteRequest { op, reply, .. } = *self;
233 match catch_unwind(AssertUnwindSafe(|| op(conn))) {
234 Ok(outcome) if conn.is_autocommit() => {
235 let _ = reply.send(outcome);
238 None
239 }
240 Ok(_outcome) => {
241 tracing::error!(
242 "writer task: top-level request returned with an open transaction; request \
243 side effects are unknown"
244 );
245 let request_state = WriterTaskRequestState::SideEffectsUnknown;
246 let _ = reply.send(Err(writer_task_terminated(request_state)));
247 Some(request_state)
248 }
249 Err(_panic_payload) => {
250 let request_state = WriterTaskRequestState::SideEffectsUnknown;
254 let _ = reply.send(Err(writer_task_terminated(request_state)));
255 Some(request_state)
256 }
257 }
258 }
259}
260
261impl<R: Send + 'static> AnyWriteRequest for WriteRequest<R> {
262 fn execute_and_reply(self: Box<Self>, conn: &Connection) {
263 let _ = sealed::Sealed::execute_and_reply_reporting_terminal(self, conn);
264 }
265
266 fn execute_and_reply_top_level(self: Box<Self>, conn: &Connection) {
267 let _ = sealed::Sealed::execute_and_reply_top_level_reporting_terminal(self, conn);
268 }
269
270 fn reply_error(self: Box<Self>, err: StorageError) {
271 let _ = self.reply.send(Err(err));
274 }
275
276 fn is_top_level(&self) -> bool {
277 self.top_level
278 }
279}
280
281fn writer_task_terminated(request_state: WriterTaskRequestState) -> StorageError {
282 StorageError::WriterTaskTerminated { request_state }
283}
284
285#[derive(Clone, Debug)]
289pub struct WriterTaskHandle {
290 tx: mpsc::Sender<Box<dyn AnyWriteRequest + Send>>,
291}
292
293impl WriterTaskHandle {
294 async fn enqueue<R, F>(
308 &self,
309 op: F,
310 ) -> Result<oneshot::Receiver<Result<R, StorageError>>, StorageError>
311 where
312 R: Send + 'static,
313 F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
314 {
315 self.enqueue_inner(op, false).await
316 }
317
318 async fn enqueue_inner<R, F>(
322 &self,
323 op: F,
324 top_level: bool,
325 ) -> Result<oneshot::Receiver<Result<R, StorageError>>, StorageError>
326 where
327 R: Send + 'static,
328 F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
329 {
330 let (reply_tx, reply_rx) = oneshot::channel();
331 let request = WriteRequest {
332 op: Box::new(op),
333 reply: reply_tx,
334 top_level,
335 };
336
337 self.tx
338 .send(Box::new(request))
339 .await
340 .map_err(|_| writer_task_terminated(WriterTaskRequestState::NotStarted))?;
341
342 Ok(reply_rx)
343 }
344
345 pub async fn send<R, F>(&self, op: F) -> Result<R, StorageError>
352 where
353 R: Send + 'static,
354 F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
355 {
356 let reply_rx = self.enqueue(op).await?;
357 reply_rx
358 .await
359 .map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
360 }
361
362 pub async fn send_with_timeout<R, F>(
377 &self,
378 op: F,
379 timeout: std::time::Duration,
380 ) -> Result<R, StorageError>
381 where
382 R: Send + 'static,
383 F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
384 {
385 let reply_rx = match tokio::time::timeout(timeout, self.enqueue(op)).await {
386 Ok(Ok(reply_rx)) => reply_rx,
387 Ok(Err(e)) => return Err(e),
388 Err(_elapsed) => {
389 return Err(StorageError::WriteQueueFull {
390 timeout_ms: timeout.as_millis() as u64,
391 })
392 }
393 };
394
395 reply_rx
396 .await
397 .map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
398 }
399
400 pub async fn send_top_level<R, F>(&self, op: F) -> Result<R, StorageError>
412 where
413 R: Send + 'static,
414 F: FnOnce(&Connection) -> Result<R, StorageError> + Send + 'static,
415 {
416 let reply_rx = self.enqueue_inner(op, true).await?;
417 reply_rx
418 .await
419 .map_err(|_| writer_task_terminated(WriterTaskRequestState::SideEffectsUnknown))?
420 }
421
422 pub fn queue_depth(&self) -> usize {
431 self.tx.max_capacity() - self.tx.capacity()
432 }
433
434 pub fn capacity(&self) -> usize {
437 self.tx.max_capacity()
438 }
439}
440
441pub fn spawn(pool: &ConnectionPool, capacity: usize) -> Result<WriterTaskHandle, SqliteError> {
462 let conn = pool.open_standalone_writer()?;
463 let origin = pool.origin();
464 let (tx, rx) = mpsc::channel(capacity.max(1));
465 tokio::spawn(run_writer_task(conn, rx, origin));
466 Ok(WriterTaskHandle { tx })
467}
468
469async fn close_and_fail_queued_requests(rx: &mut mpsc::Receiver<Box<dyn AnyWriteRequest + Send>>) {
477 rx.close();
478 while let Some(request) = rx.recv().await {
479 request.reply_error(writer_task_terminated(WriterTaskRequestState::NotStarted));
480 }
481}
482
483async fn run_writer_task(
497 mut conn: Connection,
498 mut rx: mpsc::Receiver<Box<dyn AnyWriteRequest + Send>>,
499 origin: khive_storage::tx_registry::TxOrigin,
500) {
501 while let Some(request) = rx.recv().await {
502 let origin = origin.clone();
503 let outcome = tokio::task::spawn_blocking(move || {
504 if !conn.is_autocommit() {
509 tracing::error!(
510 "writer task: connection is not in autocommit mode before request dispatch; \
511 retiring the poisoned writer without running the request"
512 );
513 let request_state = WriterTaskRequestState::NotStarted;
514 request.reply_error(writer_task_terminated(request_state));
515 return (conn, Some(request_state));
516 }
517
518 let terminal_state = if request.is_top_level() {
519 sealed::Sealed::execute_and_reply_top_level_reporting_terminal(request, &conn)
527 } else {
528 let _tx_handle = khive_storage::tx_registry::register_scoped(
529 Some("writer_task_tx".to_string()),
530 origin,
531 );
532 match conn.execute_batch("BEGIN IMMEDIATE") {
533 Ok(()) => sealed::Sealed::execute_and_reply_reporting_terminal(request, &conn),
534 Err(e) => {
535 tracing::warn!(
541 error = %e,
542 "writer task: BEGIN IMMEDIATE failed; replying an \
543 error without running the request's operation"
544 );
545 request.reply_error(StorageError::Pool {
546 operation: "writer_task_begin".into(),
547 message: e.to_string(),
548 });
549 None
550 }
551 }
552 };
553 (conn, terminal_state)
554 })
555 .await;
556
557 match outcome {
558 Ok((returned_conn, None)) => conn = returned_conn,
559 Ok((_returned_conn, Some(request_state))) => {
560 tracing::error!(
561 request_state = %request_state,
562 "writer task reached a terminal request or connection state; closing and \
563 failing the queue without restarting"
564 );
565 close_and_fail_queued_requests(&mut rx).await;
566 return;
567 }
568 Err(join_err) => {
569 tracing::error!(
570 error = %join_err,
571 "writer task blocking closure failed outside the request \
572 panic boundary; closing and failing the queue without restarting"
573 );
574 close_and_fail_queued_requests(&mut rx).await;
575 return;
576 }
577 }
578 }
579}
580
581#[cfg(test)]
582mod tests {
583 use super::*;
584 use crate::pool::PoolConfig;
585 use rusqlite::hooks::{AuthAction, AuthContext, Authorization, TransactionOperation};
586 use serial_test::serial;
587 use std::sync::atomic::{AtomicBool, Ordering};
588 use std::sync::mpsc as std_mpsc;
589 use std::sync::Arc;
590 use std::time::Duration;
591
592 fn file_pool(path: &std::path::Path) -> ConnectionPool {
593 let cfg = PoolConfig {
594 path: Some(path.to_path_buf()),
595 ..PoolConfig::default()
596 };
597 ConnectionPool::new(cfg).expect("pool open")
598 }
599
600 fn deny_commit_and_rollback(ctx: AuthContext<'_>) -> Authorization {
601 match ctx.action {
602 AuthAction::Transaction {
605 operation: TransactionOperation::Unknown | TransactionOperation::Rollback,
606 } => Authorization::Deny,
607 _ => Authorization::Allow,
608 }
609 }
610
611 fn deny_commit(ctx: AuthContext<'_>) -> Authorization {
612 match ctx.action {
613 AuthAction::Transaction {
614 operation: TransactionOperation::Unknown,
615 } => Authorization::Deny,
616 _ => Authorization::Allow,
617 }
618 }
619
620 fn deny_rollback(ctx: AuthContext<'_>) -> Authorization {
621 match ctx.action {
622 AuthAction::Transaction {
623 operation: TransactionOperation::Rollback,
624 } => Authorization::Deny,
625 _ => Authorization::Allow,
626 }
627 }
628
629 fn assert_writer_task_terminal_state<T: std::fmt::Debug>(
630 result: Result<T, StorageError>,
631 expected: WriterTaskRequestState,
632 ) {
633 match result {
634 Err(StorageError::WriterTaskTerminated { request_state }) => {
635 assert_eq!(request_state, expected)
636 }
637 other => panic!("expected WriterTaskTerminated({expected:?}), got {other:?}"),
638 }
639 }
640
641 #[tokio::test]
648 #[serial(tx_registry)]
649 async fn begin_immediate_failure_replies_error_without_running_op() {
650 let dir = tempfile::tempdir().unwrap();
656 let path = dir.path().join("writer_task_begin_failure.db");
657 let cfg = PoolConfig {
658 path: Some(path.clone()),
659 busy_timeout: Duration::from_millis(150),
660 ..PoolConfig::default()
661 };
662 let pool = ConnectionPool::new(cfg).unwrap();
663 {
664 let writer = pool.try_writer().unwrap();
665 writer
666 .conn()
667 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
668 .unwrap();
669 }
670
671 let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
672
673 let lock_holder = pool.try_writer().unwrap();
674 lock_holder.conn().execute_batch("BEGIN IMMEDIATE").unwrap();
675
676 let op_ran = Arc::new(AtomicBool::new(false));
677 let op_ran_clone = Arc::clone(&op_ran);
678 let result = handle
679 .send(move |conn| {
680 op_ran_clone.store(true, Ordering::SeqCst);
681 conn.execute("INSERT INTO t (id, v) VALUES (99, 'should-not-land')", [])
682 .map_err(|e| StorageError::Pool {
683 operation: "test_insert".into(),
684 message: e.to_string(),
685 })
686 })
687 .await;
688
689 assert!(
690 matches!(
691 &result,
692 Err(StorageError::Pool { operation, .. }) if operation == "writer_task_begin"
693 ),
694 "expected a writer_task_begin Pool error on BEGIN IMMEDIATE \
695 failure, got {result:?}"
696 );
697 assert!(
698 !op_ran.load(Ordering::SeqCst),
699 "the request's operation closure must never run when BEGIN \
700 IMMEDIATE fails — running it would land a partial write in \
701 autocommit mode for a request the caller is told failed"
702 );
703
704 lock_holder.conn().execute_batch("ROLLBACK").unwrap();
707 drop(lock_holder);
708
709 let reader = pool.reader().expect("reader");
710 let count: i64 = reader
711 .conn()
712 .query_row("SELECT COUNT(*) FROM t WHERE id = 99", [], |row| row.get(0))
713 .unwrap();
714 assert_eq!(
715 count, 0,
716 "no row must have landed from the request whose BEGIN IMMEDIATE failed"
717 );
718 }
719
720 #[tokio::test]
724 #[serial(tx_registry)]
725 async fn writer_task_executes_op_and_commits() {
726 let dir = tempfile::tempdir().unwrap();
727 let path = dir.path().join("writer_task_commit.db");
728 let pool = file_pool(&path);
729 {
730 let writer = pool.try_writer().unwrap();
731 writer
732 .conn()
733 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
734 .unwrap();
735 }
736
737 let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
738
739 let affected = handle
740 .send(|conn| {
741 conn.execute("INSERT INTO t (id, v) VALUES (1, 'hello')", [])
742 .map_err(|e| StorageError::Pool {
743 operation: "test_insert".into(),
744 message: e.to_string(),
745 })
746 })
747 .await
748 .expect("op should succeed");
749 assert_eq!(affected, 1);
750
751 let reader = pool.reader().expect("reader");
755 let v: String = reader
756 .conn()
757 .query_row("SELECT v FROM t WHERE id = 1", [], |row| row.get(0))
758 .expect("row must be committed and visible to a reader");
759 assert_eq!(v, "hello");
760 }
761
762 #[test]
763 fn spawn_fails_on_in_memory_pool() {
764 let cfg = PoolConfig {
770 path: None,
771 ..PoolConfig::default()
772 };
773 let pool = ConnectionPool::new(cfg).unwrap();
774 let result = spawn(&pool, 8);
775 assert!(
776 result.is_err(),
777 "in-memory pools must reject spawn, not panic"
778 );
779 }
780
781 #[tokio::test]
782 async fn full_channel_applies_backpressure_not_immediate_error() {
783 let (tx, _rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
788 let handle = WriterTaskHandle { tx };
789
790 let first = tokio::spawn({
793 let handle = handle.clone();
794 async move {
795 let _ = handle.send(|_conn| Ok::<(), StorageError>(())).await;
796 }
797 });
798
799 tokio::time::sleep(Duration::from_millis(20)).await;
801
802 let second = tokio::time::timeout(
805 Duration::from_millis(100),
806 handle.send(|_conn| Ok::<(), StorageError>(())),
807 )
808 .await;
809
810 assert!(
811 second.is_err(),
812 "a full channel must apply backpressure (send suspends) rather \
813 than erroring immediately — no try_send escape hatch per ADR-067"
814 );
815
816 first.abort();
817 }
818
819 #[tokio::test]
820 async fn send_with_timeout_maps_full_channel_to_write_queue_full() {
821 let (tx, _rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
822 let handle = WriterTaskHandle { tx };
823
824 let first = tokio::spawn({
825 let handle = handle.clone();
826 async move {
827 let _ = handle.send(|_conn| Ok::<(), StorageError>(())).await;
828 }
829 });
830 tokio::time::sleep(Duration::from_millis(20)).await;
831
832 let result = handle
833 .send_with_timeout(
834 |_conn| Ok::<(), StorageError>(()),
835 Duration::from_millis(50),
836 )
837 .await;
838
839 match result {
840 Err(StorageError::WriteQueueFull { timeout_ms }) => assert_eq!(timeout_ms, 50),
841 other => panic!("expected WriteQueueFull, got {other:?}"),
842 }
843
844 first.abort();
845 }
846
847 #[tokio::test]
853 #[serial(tx_registry)]
854 async fn send_with_timeout_returns_op_result_when_op_outlives_the_timeout() {
855 let dir = tempfile::tempdir().unwrap();
862 let path = dir.path().join("writer_task_slow_op.db");
863 let pool = file_pool(&path);
864 {
865 let writer = pool.try_writer().unwrap();
866 writer
867 .conn()
868 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
869 .unwrap();
870 }
871
872 let handle = spawn(&pool, 8).expect("writer task should spawn on a file-backed pool");
873
874 let result = handle
875 .send_with_timeout(
876 |conn| {
877 std::thread::sleep(Duration::from_millis(150));
880 conn.execute("INSERT INTO t (id, v) VALUES (1, 'slow')", [])
881 .map_err(|e| StorageError::Pool {
882 operation: "test_insert".into(),
883 message: e.to_string(),
884 })
885 },
886 Duration::from_millis(20),
887 )
888 .await;
889
890 let affected = result.expect(
891 "an accepted request must return its real result even when the \
892 op takes longer than the enqueue timeout, not WriteQueueFull",
893 );
894 assert_eq!(affected, 1);
895
896 let reader = pool.reader().expect("reader");
899 let v: String = reader
900 .conn()
901 .query_row("SELECT v FROM t WHERE id = 1", [], |row| row.get(0))
902 .expect("the slow op's write must have committed");
903 assert_eq!(v, "slow");
904 }
905
906 #[tokio::test]
907 #[serial(tx_registry)]
908 async fn operation_failure_with_successful_rollback_preserves_error_and_writer_continues() {
909 let dir = tempfile::tempdir().unwrap();
910 let path = dir.path().join("writer_task_operation_rollback.db");
911 let pool = file_pool(&path);
912 {
913 let writer = pool.try_writer().unwrap();
914 writer
915 .conn()
916 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
917 .unwrap();
918 }
919 let handle = spawn(&pool, 8).expect("writer task spawn");
920
921 let original_error = handle
922 .send(|conn| -> Result<(), StorageError> {
923 conn.execute("INSERT INTO t (id, v) VALUES (1, 'rolled-back')", [])
924 .map_err(|e| StorageError::Pool {
925 operation: "test_operation_error_insert".into(),
926 message: e.to_string(),
927 })?;
928 Err(StorageError::Internal(
929 "intentional operation failure".into(),
930 ))
931 })
932 .await;
933 assert!(
934 matches!(
935 &original_error,
936 Err(StorageError::Internal(message))
937 if message == "intentional operation failure"
938 ),
939 "a confirmed rollback must preserve the operation error, got {original_error:?}"
940 );
941
942 let affected = handle
943 .send(|conn| {
944 conn.execute("INSERT INTO t (id, v) VALUES (2, 'committed')", [])
945 .map_err(|e| StorageError::Pool {
946 operation: "test_operation_error_followup_insert".into(),
947 message: e.to_string(),
948 })
949 })
950 .await
951 .expect("the writer must continue after a confirmed rollback");
952 assert_eq!(affected, 1);
953
954 let reader = pool.reader().expect("reader");
955 let rolled_back: i64 = reader
956 .conn()
957 .query_row("SELECT COUNT(*) FROM t WHERE id = 1", [], |row| row.get(0))
958 .unwrap();
959 let committed: i64 = reader
960 .conn()
961 .query_row("SELECT COUNT(*) FROM t WHERE id = 2", [], |row| row.get(0))
962 .unwrap();
963 assert_eq!(rolled_back, 0);
964 assert_eq!(committed, 1);
965 }
966
967 #[tokio::test]
968 #[serial(tx_registry)]
969 async fn commit_failure_with_successful_rollback_preserves_error_and_writer_continues() {
970 let dir = tempfile::tempdir().unwrap();
971 let path = dir.path().join("writer_task_commit_rollback.db");
972 let pool = file_pool(&path);
973 {
974 let writer = pool.try_writer().unwrap();
975 writer
976 .conn()
977 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
978 .unwrap();
979 }
980 let handle = spawn(&pool, 8).expect("writer task spawn");
981
982 let commit_error = handle
983 .send(|conn| -> Result<usize, StorageError> {
984 let affected = conn
985 .execute("INSERT INTO t (id, v) VALUES (1, 'rolled-back')", [])
986 .map_err(|e| StorageError::Pool {
987 operation: "test_commit_error_insert".into(),
988 message: e.to_string(),
989 })?;
990 conn.authorizer(Some(deny_commit))
991 .map_err(|e| StorageError::Pool {
992 operation: "test_install_authorizer".into(),
993 message: e.to_string(),
994 })?;
995 Ok(affected)
996 })
997 .await;
998 assert!(
999 matches!(
1000 &commit_error,
1001 Err(StorageError::Pool { operation, .. })
1002 if operation == "writer_task_commit"
1003 ),
1004 "a confirmed rollback must preserve the commit error, got {commit_error:?}"
1005 );
1006 assert!(
1007 commit_error
1008 .as_ref()
1009 .expect_err("COMMIT must be denied")
1010 .is_retryable(),
1011 "the existing retryable commit-error contract must remain unchanged after a \
1012 confirmed rollback"
1013 );
1014
1015 let affected = handle
1016 .send(|conn| {
1017 conn.authorizer(None::<fn(AuthContext<'_>) -> Authorization>)
1018 .map_err(|e| StorageError::Pool {
1019 operation: "test_remove_authorizer".into(),
1020 message: e.to_string(),
1021 })?;
1022 conn.execute("INSERT INTO t (id, v) VALUES (2, 'committed')", [])
1023 .map_err(|e| StorageError::Pool {
1024 operation: "test_commit_error_followup_insert".into(),
1025 message: e.to_string(),
1026 })
1027 })
1028 .await
1029 .expect("the writer must continue after the failed COMMIT is rolled back");
1030 assert_eq!(affected, 1);
1031
1032 let reader = pool.reader().expect("reader");
1033 let rolled_back: i64 = reader
1034 .conn()
1035 .query_row("SELECT COUNT(*) FROM t WHERE id = 1", [], |row| row.get(0))
1036 .unwrap();
1037 let committed: i64 = reader
1038 .conn()
1039 .query_row("SELECT COUNT(*) FROM t WHERE id = 2", [], |row| row.get(0))
1040 .unwrap();
1041 assert_eq!(rolled_back, 0);
1042 assert_eq!(committed, 1);
1043 }
1044
1045 #[test]
1046 fn top_level_request_returning_with_open_transaction_reports_side_effects_unknown() {
1047 let conn = Connection::open_in_memory().expect("in-memory connection");
1048 conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY)")
1049 .unwrap();
1050 let (reply_tx, mut reply_rx) = oneshot::channel();
1051 let request = WriteRequest {
1052 op: Box::new(|conn| -> Result<usize, StorageError> {
1053 conn.execute_batch("BEGIN IMMEDIATE")
1054 .map_err(|e| StorageError::Pool {
1055 operation: "test_top_level_begin".into(),
1056 message: e.to_string(),
1057 })?;
1058 conn.execute("INSERT INTO t (id) VALUES (1)", [])
1059 .map_err(|e| StorageError::Pool {
1060 operation: "test_top_level_insert".into(),
1061 message: e.to_string(),
1062 })
1063 }),
1064 reply: reply_tx,
1065 top_level: true,
1066 };
1067
1068 let terminal_state = sealed::Sealed::execute_and_reply_top_level_reporting_terminal(
1069 Box::new(request),
1070 &conn,
1071 );
1072 assert_eq!(
1073 terminal_state,
1074 Some(WriterTaskRequestState::SideEffectsUnknown)
1075 );
1076 let reply = reply_rx
1077 .try_recv()
1078 .expect("active request must receive a typed terminal reply");
1079 assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
1080 assert!(
1081 !conn.is_autocommit(),
1082 "the fixture must prove the post-request autocommit check observed an open transaction"
1083 );
1084 }
1085
1086 #[test]
1087 fn commit_failure_with_failed_rollback_reports_side_effects_unknown() {
1088 let conn = Connection::open_in_memory().expect("in-memory connection");
1089 conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
1090 .unwrap();
1091 let (reply_tx, mut reply_rx) = oneshot::channel();
1092 let request = WriteRequest {
1093 op: Box::new(|conn| -> Result<usize, StorageError> {
1094 let affected = conn
1095 .execute("INSERT INTO t (id) VALUES (1)", [])
1096 .map_err(|e| StorageError::Pool {
1097 operation: "test_insert_before_commit_failure".into(),
1098 message: e.to_string(),
1099 })?;
1100 conn.authorizer(Some(deny_commit_and_rollback))
1101 .map_err(|e| StorageError::Pool {
1102 operation: "test_install_authorizer".into(),
1103 message: e.to_string(),
1104 })?;
1105 Ok(affected)
1106 }),
1107 reply: reply_tx,
1108 top_level: false,
1109 };
1110
1111 let terminal_state =
1112 sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
1113 assert_eq!(
1114 terminal_state,
1115 Some(WriterTaskRequestState::SideEffectsUnknown)
1116 );
1117 let reply = reply_rx
1118 .try_recv()
1119 .expect("active request must receive a typed terminal reply");
1120 assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
1121 assert!(
1122 !conn.is_autocommit(),
1123 "the denied COMMIT and ROLLBACK must leave the test connection poisoned"
1124 );
1125 }
1126
1127 #[tokio::test]
1128 #[serial(tx_registry)]
1129 async fn poisoned_connection_retires_before_queued_top_level_request() {
1130 let dir = tempfile::tempdir().unwrap();
1131 let path = dir.path().join("writer_task_rollback_poison.db");
1132 let pool = file_pool(&path);
1133 {
1134 let writer = pool.try_writer().unwrap();
1135 writer
1136 .conn()
1137 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
1138 .unwrap();
1139 }
1140 let handle = spawn(&pool, 8).expect("writer task spawn");
1141 let (started_tx, started_rx) = oneshot::channel::<()>();
1142 let (release_tx, release_rx) = std_mpsc::channel::<()>();
1143
1144 let active = tokio::spawn({
1145 let handle = handle.clone();
1146 async move {
1147 handle
1148 .send(move |conn| -> Result<usize, StorageError> {
1149 let affected = conn
1150 .execute("INSERT INTO t (id, v) VALUES (1, 'active')", [])
1151 .map_err(|e| StorageError::Pool {
1152 operation: "test_active_insert".into(),
1153 message: e.to_string(),
1154 })?;
1155 conn.authorizer(Some(deny_commit_and_rollback))
1156 .map_err(|e| StorageError::Pool {
1157 operation: "test_install_authorizer".into(),
1158 message: e.to_string(),
1159 })?;
1160 let _ = started_tx.send(());
1161 release_rx.recv().expect("test must release active op");
1162 Ok(affected)
1163 })
1164 .await
1165 }
1166 });
1167
1168 tokio::time::timeout(Duration::from_secs(5), started_rx)
1169 .await
1170 .expect("active request did not start")
1171 .expect("active request dropped its start signal");
1172
1173 let queued_ran = Arc::new(AtomicBool::new(false));
1174 let queued_ran_in_op = Arc::clone(&queued_ran);
1175 let queued_top_level = handle
1176 .enqueue_inner(
1177 move |conn| {
1178 queued_ran_in_op.store(true, Ordering::SeqCst);
1179 conn.execute("INSERT INTO t (id, v) VALUES (2, 'queued')", [])
1180 .map_err(|e| StorageError::Pool {
1181 operation: "test_queued_top_level_insert".into(),
1182 message: e.to_string(),
1183 })
1184 },
1185 true,
1186 )
1187 .await
1188 .expect("top-level request must queue behind active request");
1189 release_tx.send(()).expect("release active op");
1190
1191 let active_result = tokio::time::timeout(Duration::from_secs(5), active)
1192 .await
1193 .expect("active caller hung after rollback failure")
1194 .expect("active caller task join");
1195 assert_writer_task_terminal_state(
1196 active_result,
1197 WriterTaskRequestState::SideEffectsUnknown,
1198 );
1199
1200 let queued_result = tokio::time::timeout(Duration::from_secs(5), queued_top_level)
1201 .await
1202 .expect("queued top-level caller hung after terminal failure")
1203 .expect("terminal drain must preserve queued typed reply");
1204 assert_writer_task_terminal_state(queued_result, WriterTaskRequestState::NotStarted);
1205 assert!(
1206 !queued_ran.load(Ordering::SeqCst),
1207 "a top-level request must never run on the poisoned connection"
1208 );
1209
1210 let future_ran = Arc::new(AtomicBool::new(false));
1211 let future_ran_in_op = Arc::clone(&future_ran);
1212 let future_result = handle
1213 .send_top_level(move |_conn| {
1214 future_ran_in_op.store(true, Ordering::SeqCst);
1215 Ok::<(), StorageError>(())
1216 })
1217 .await;
1218 assert_writer_task_terminal_state(future_result, WriterTaskRequestState::NotStarted);
1219 assert!(!future_ran.load(Ordering::SeqCst));
1220 }
1221
1222 #[test]
1223 fn operation_failure_with_failed_rollback_reports_side_effects_unknown() {
1224 let conn = Connection::open_in_memory().expect("in-memory connection");
1225 conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
1226 .unwrap();
1227 let (reply_tx, mut reply_rx) = oneshot::channel();
1228 let request = WriteRequest {
1229 op: Box::new(|conn| -> Result<(), StorageError> {
1230 conn.authorizer(Some(deny_rollback))
1231 .map_err(|e| StorageError::Pool {
1232 operation: "test_install_authorizer".into(),
1233 message: e.to_string(),
1234 })?;
1235 Err(StorageError::Internal(
1236 "intentional operation failure before denied rollback".into(),
1237 ))
1238 }),
1239 reply: reply_tx,
1240 top_level: false,
1241 };
1242
1243 let terminal_state =
1244 sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
1245 assert_eq!(
1246 terminal_state,
1247 Some(WriterTaskRequestState::SideEffectsUnknown)
1248 );
1249 let reply = reply_rx
1250 .try_recv()
1251 .expect("active request must receive a typed terminal reply");
1252 assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
1253 assert!(
1254 !conn.is_autocommit(),
1255 "the denied ROLLBACK must leave the test connection poisoned"
1256 );
1257 }
1258
1259 #[test]
1260 fn wrapped_panic_with_failed_rollback_reports_side_effects_unknown() {
1261 let conn = Connection::open_in_memory().expect("in-memory connection");
1262 conn.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY); BEGIN IMMEDIATE")
1263 .unwrap();
1264 let (reply_tx, mut reply_rx) = oneshot::channel();
1265 let request = WriteRequest {
1266 op: Box::new(|conn| -> Result<(), StorageError> {
1267 conn.execute_batch("INSERT INTO t (id) VALUES (1); COMMIT")
1272 .map_err(|e| StorageError::Pool {
1273 operation: "test_force_rollback_failure".into(),
1274 message: e.to_string(),
1275 })?;
1276 panic!("intentional panic after illicit commit");
1277 }),
1278 reply: reply_tx,
1279 top_level: false,
1280 };
1281
1282 let terminal_state =
1283 sealed::Sealed::execute_and_reply_reporting_terminal(Box::new(request), &conn);
1284 assert_eq!(
1285 terminal_state,
1286 Some(WriterTaskRequestState::SideEffectsUnknown)
1287 );
1288 let reply = reply_rx
1289 .try_recv()
1290 .expect("active request must receive a typed terminal reply");
1291 assert_writer_task_terminal_state(reply, WriterTaskRequestState::SideEffectsUnknown);
1292
1293 let count: i64 = conn
1294 .query_row("SELECT COUNT(*) FROM t", [], |row| row.get(0))
1295 .unwrap();
1296 assert_eq!(
1297 count, 1,
1298 "the fixture's committed side effect proves why the state must be unknown"
1299 );
1300 }
1301
1302 #[tokio::test]
1306 #[serial(tx_registry)]
1307 async fn wrapped_panic_rolls_back_and_terminally_fails_queue() {
1308 let dir = tempfile::tempdir().unwrap();
1309 let path = dir.path().join("writer_task_wrapped_panic.db");
1310 let cfg = PoolConfig {
1311 path: Some(path),
1312 write_queue_enabled: true,
1313 write_queue_capacity: 8,
1314 ..PoolConfig::default()
1315 };
1316 let pool = ConnectionPool::new(cfg).unwrap();
1317 {
1318 let writer = pool.try_writer().unwrap();
1319 writer
1320 .conn()
1321 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
1322 .unwrap();
1323 }
1324
1325 let handle = pool
1328 .writer_task_handle()
1329 .expect("writer task lookup")
1330 .expect("file-backed queued pool must spawn its writer task");
1331 assert_eq!(pool.writer_task_spawn_count(), 1);
1332
1333 let (started_tx, started_rx) = oneshot::channel::<()>();
1334 let (release_tx, release_rx) = std_mpsc::channel::<()>();
1335 let active = tokio::spawn({
1336 let handle = handle.clone();
1337 async move {
1338 handle
1339 .send(move |conn| -> Result<usize, StorageError> {
1340 conn.execute("INSERT INTO t (id, v) VALUES (1, 'active')", [])
1341 .map_err(|e| StorageError::Pool {
1342 operation: "test_active_insert".into(),
1343 message: e.to_string(),
1344 })?;
1345 let _ = started_tx.send(());
1346 release_rx.recv().expect("test must release active op");
1347 panic!("intentional wrapped writer request panic");
1348 })
1349 .await
1350 }
1351 });
1352
1353 tokio::time::timeout(Duration::from_secs(5), started_rx)
1354 .await
1355 .expect("active request did not start")
1356 .expect("active request dropped its start signal");
1357
1358 let queued_one_ran = Arc::new(AtomicBool::new(false));
1359 let queued_one_ran_in_op = Arc::clone(&queued_one_ran);
1360 let queued_one = handle
1361 .enqueue(move |conn| {
1362 queued_one_ran_in_op.store(true, Ordering::SeqCst);
1363 conn.execute("INSERT INTO t (id, v) VALUES (2, 'queued-one')", [])
1364 .map_err(|e| StorageError::Pool {
1365 operation: "test_queued_one_insert".into(),
1366 message: e.to_string(),
1367 })
1368 })
1369 .await
1370 .expect("first queued request must be accepted");
1371
1372 let queued_two_ran = Arc::new(AtomicBool::new(false));
1373 let queued_two_ran_in_op = Arc::clone(&queued_two_ran);
1374 let queued_two = handle
1375 .enqueue(move |_conn| {
1376 queued_two_ran_in_op.store(true, Ordering::SeqCst);
1377 Ok::<String, StorageError>("queued-two-ran".to_string())
1378 })
1379 .await
1380 .expect("second queued request must be accepted");
1381
1382 assert_eq!(
1383 handle.queue_depth(),
1384 2,
1385 "both heterogeneous requests must be buffered behind the active op"
1386 );
1387 release_tx.send(()).expect("release active op");
1388
1389 let active_result = tokio::time::timeout(Duration::from_secs(5), active)
1390 .await
1391 .expect("active caller hung after panic")
1392 .expect("active caller task join");
1393 assert_writer_task_terminal_state(
1394 active_result,
1395 WriterTaskRequestState::TransactionRolledBack,
1396 );
1397
1398 let queued_one_result = tokio::time::timeout(Duration::from_secs(5), queued_one)
1399 .await
1400 .expect("first queued caller hung after terminal failure")
1401 .expect("terminal drain must preserve first typed reply");
1402 assert_writer_task_terminal_state(queued_one_result, WriterTaskRequestState::NotStarted);
1403
1404 let queued_two_result = tokio::time::timeout(Duration::from_secs(5), queued_two)
1405 .await
1406 .expect("second queued caller hung after terminal failure")
1407 .expect("terminal drain must preserve second typed reply");
1408 assert_writer_task_terminal_state(queued_two_result, WriterTaskRequestState::NotStarted);
1409 assert!(!queued_one_ran.load(Ordering::SeqCst));
1410 assert!(!queued_two_ran.load(Ordering::SeqCst));
1411
1412 let future_ran = Arc::new(AtomicBool::new(false));
1413 let future_ran_in_op = Arc::clone(&future_ran);
1414 let future_result = handle
1415 .send(move |_conn| {
1416 future_ran_in_op.store(true, Ordering::SeqCst);
1417 Ok::<(), StorageError>(())
1418 })
1419 .await;
1420 assert_writer_task_terminal_state(future_result, WriterTaskRequestState::NotStarted);
1421 assert!(!future_ran.load(Ordering::SeqCst));
1422
1423 let cached_after_failure = pool
1424 .writer_task_handle()
1425 .expect("cached writer task lookup")
1426 .expect("pool retains its terminal handle");
1427 assert_eq!(
1428 pool.writer_task_spawn_count(),
1429 1,
1430 "a terminal writer task must not be restarted behind callers' backs"
1431 );
1432 let cached_result = cached_after_failure
1433 .send(|_conn| Ok::<(), StorageError>(()))
1434 .await;
1435 assert_writer_task_terminal_state(cached_result, WriterTaskRequestState::NotStarted);
1436
1437 let reader = pool.reader().expect("reader");
1438 let count: i64 = reader
1439 .conn()
1440 .query_row("SELECT COUNT(*) FROM t", [], |row| row.get(0))
1441 .unwrap();
1442 assert_eq!(
1443 count, 0,
1444 "the active transaction must be rolled back and queued ops must never run"
1445 );
1446 }
1447
1448 #[tokio::test]
1449 async fn top_level_panic_reports_unknown_and_fails_queue_without_running_it() {
1450 let dir = tempfile::tempdir().unwrap();
1451 let path = dir.path().join("writer_task_top_level_panic.db");
1452 let pool = file_pool(&path);
1453 {
1454 let writer = pool.try_writer().unwrap();
1455 writer
1456 .conn()
1457 .execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)")
1458 .unwrap();
1459 }
1460 let handle = spawn(&pool, 8).expect("writer task spawn");
1461
1462 let (started_tx, started_rx) = oneshot::channel::<()>();
1463 let (release_tx, release_rx) = std_mpsc::channel::<()>();
1464 let active = tokio::spawn({
1465 let handle = handle.clone();
1466 async move {
1467 handle
1468 .send_top_level(move |conn| -> Result<usize, StorageError> {
1469 conn.execute("INSERT INTO t (id, v) VALUES (10, 'autocommitted')", [])
1470 .map_err(|e| StorageError::Pool {
1471 operation: "test_top_level_insert".into(),
1472 message: e.to_string(),
1473 })?;
1474 let _ = started_tx.send(());
1475 release_rx.recv().expect("test must release top-level op");
1476 panic!("intentional top-level writer request panic");
1477 })
1478 .await
1479 }
1480 });
1481
1482 tokio::time::timeout(Duration::from_secs(5), started_rx)
1483 .await
1484 .expect("top-level request did not start")
1485 .expect("top-level request dropped its start signal");
1486
1487 let queued_ran = Arc::new(AtomicBool::new(false));
1488 let queued_ran_in_op = Arc::clone(&queued_ran);
1489 let queued = handle
1490 .enqueue(move |conn| {
1491 queued_ran_in_op.store(true, Ordering::SeqCst);
1492 conn.execute("INSERT INTO t (id, v) VALUES (11, 'queued')", [])
1493 .map_err(|e| StorageError::Pool {
1494 operation: "test_top_level_queued_insert".into(),
1495 message: e.to_string(),
1496 })
1497 })
1498 .await
1499 .expect("queued request must be accepted");
1500 assert_eq!(handle.queue_depth(), 1);
1501 release_tx.send(()).expect("release top-level op");
1502
1503 let active_result = tokio::time::timeout(Duration::from_secs(5), active)
1504 .await
1505 .expect("top-level caller hung after panic")
1506 .expect("top-level caller task join");
1507 assert_writer_task_terminal_state(
1508 active_result,
1509 WriterTaskRequestState::SideEffectsUnknown,
1510 );
1511
1512 let queued_result = tokio::time::timeout(Duration::from_secs(5), queued)
1513 .await
1514 .expect("queued caller hung after top-level panic")
1515 .expect("terminal drain must preserve queued typed reply");
1516 assert_writer_task_terminal_state(queued_result, WriterTaskRequestState::NotStarted);
1517 assert!(!queued_ran.load(Ordering::SeqCst));
1518
1519 let reader = pool.reader().expect("reader");
1520 let active_count: i64 = reader
1521 .conn()
1522 .query_row("SELECT COUNT(*) FROM t WHERE id = 10", [], |row| row.get(0))
1523 .unwrap();
1524 let queued_count: i64 = reader
1525 .conn()
1526 .query_row("SELECT COUNT(*) FROM t WHERE id = 11", [], |row| row.get(0))
1527 .unwrap();
1528 assert_eq!(
1529 active_count, 1,
1530 "the completed top-level statement autocommits before the panic"
1531 );
1532 assert_eq!(queued_count, 0, "the queued request must never run");
1533 }
1534
1535 #[tokio::test]
1536 async fn closed_receiver_rejects_all_send_surfaces_as_not_started() {
1537 let (tx, rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(4);
1541 drop(rx);
1542
1543 let handle = WriterTaskHandle { tx };
1544 let send_result = handle.send(|_conn| Ok::<(), StorageError>(())).await;
1545 assert_writer_task_terminal_state(send_result, WriterTaskRequestState::NotStarted);
1546
1547 let timed_result = handle
1548 .send_with_timeout(|_conn| Ok::<(), StorageError>(()), Duration::from_secs(1))
1549 .await;
1550 assert_writer_task_terminal_state(timed_result, WriterTaskRequestState::NotStarted);
1551
1552 let top_level_result = handle
1553 .send_top_level(|_conn| Ok::<(), StorageError>(()))
1554 .await;
1555 assert_writer_task_terminal_state(top_level_result, WriterTaskRequestState::NotStarted);
1556 }
1557
1558 #[tokio::test]
1559 async fn accepted_request_lost_reply_is_side_effects_unknown() {
1560 let (tx, mut rx) = mpsc::channel::<Box<dyn AnyWriteRequest + Send>>(1);
1564 let handle = WriterTaskHandle { tx };
1565 let request_ran = Arc::new(AtomicBool::new(false));
1566 let request_ran_in_op = Arc::clone(&request_ran);
1567
1568 let dropper = tokio::spawn(async move {
1569 let request = rx.recv().await.expect("request must be accepted");
1570 drop(request);
1571 });
1572 let result = tokio::time::timeout(
1573 Duration::from_secs(5),
1574 handle.send(move |_conn| {
1575 request_ran_in_op.store(true, Ordering::SeqCst);
1576 Ok::<(), StorageError>(())
1577 }),
1578 )
1579 .await
1580 .expect("caller hung after accepted request was dropped");
1581 dropper.await.expect("dropper task join");
1582
1583 assert_writer_task_terminal_state(result, WriterTaskRequestState::SideEffectsUnknown);
1584 assert!(!request_ran.load(Ordering::SeqCst));
1585 }
1586}