drizzle 0.2.1

A type-safe SQL query builder for Rust
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
//! Shared savepoint orchestration for transaction drivers.
//!
//! Every driver implements its own `Transaction` type, but the savepoint
//! protocol (`SAVEPOINT N` → run callback → `RELEASE` / `ROLLBACK TO`) is
//! identical across SQLite and Postgres. The helpers in this module own
//! that protocol so individual drivers only supply the `execute_raw`
//! closure and the callback to invoke between bookends.
//!
//! Sync drivers use [`std::panic::catch_unwind`] so that a panic inside
//! the callback issues `ROLLBACK TO SAVEPOINT` before re-raising the
//! panic. Async drivers cannot reasonably catch panics across `.await`
//! points and simply propagate the `Err`.
//!
//! Synchronous drivers track nesting depth in an [`AtomicU32`]. Async drivers
//! use [`AsyncSavepointState`], which assigns monotonic names, orders cleanup
//! in LIFO order when futures overlap, and poisons the transaction if a
//! savepoint future is cancelled.

use core::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::{
    collections::HashMap,
    future::poll_fn,
    sync::Mutex,
    task::{Poll, Waker},
};

use drizzle_core::error::{DrizzleError, Result};

/// Reports a failed callback together with the cleanup that failed after it.
pub(crate) fn cleanup_error(
    scope: &str,
    original: DrizzleError,
    action: &str,
    cleanup: DrizzleError,
) -> DrizzleError {
    DrizzleError::TransactionError(
        format!("{scope} callback failed: {original}; {action} failed: {cleanup}").into(),
    )
}

fn trace_panic_cleanup_error(scope: &str, name: &str, action: &str, err: &DrizzleError) {
    #[cfg(feature = "tracing")]
    tracing::error!(
        scope,
        name,
        action,
        error = %err,
        "transaction cleanup failed after panic"
    );

    #[cfg(not(feature = "tracing"))]
    let _ = (scope, name, action, err);
}

/// Whether a statement failed on the server inside a transaction.
///
/// On PostgreSQL a server error aborts the transaction: every later statement
/// fails, and `COMMIT` rolls back instead of committing. Drivers record such
/// errors here so that committing reports the rollback instead of success.
/// `ROLLBACK TO SAVEPOINT` recovers the transaction, so savepoints restore
/// the flag to its value at `SAVEPOINT`. Drivers for databases whose
/// transactions survive a failed statement never set it.
#[derive(Debug, Default)]
pub struct AbortState(AtomicBool);

/// A state that is never marked, for drivers without aborting semantics.
static NEVER_ABORTED: AbortState = AbortState::new();

impl AbortState {
    pub const fn new() -> Self {
        Self(AtomicBool::new(false))
    }

    /// Records a server error.
    pub fn mark(&self) {
        self.0.store(true, Ordering::Relaxed);
    }

    /// Whether a server error aborted the transaction.
    pub fn is_aborted(&self) -> bool {
        self.0.load(Ordering::Relaxed)
    }

    fn set(&self, aborted: bool) {
        self.0.store(aborted, Ordering::Relaxed);
    }
}

/// The error returned when a statement failed inside a savepoint whose body
/// still returned `Ok`: the savepoint was rolled back to recover the
/// transaction.
fn failed_statement_in_savepoint() -> DrizzleError {
    DrizzleError::TransactionError(
        "a statement inside the savepoint failed, so the savepoint was rolled back".into(),
    )
}

/// The error `commit` returns for a transaction a failed statement aborted:
/// it was rolled back.
pub fn aborted_transaction_error() -> DrizzleError {
    DrizzleError::TransactionError(
        "a statement in the transaction failed, so it was rolled back instead of committed".into(),
    )
}

#[derive(Default)]
struct AsyncSavepointInner {
    next_id: u64,
    stack: Vec<u64>,
    poisoned: bool,
    waiters: HashMap<u64, Waker>,
}

/// Shared ordering and cancellation state for async transaction savepoints.
#[derive(Default)]
pub struct AsyncSavepointState(Mutex<AsyncSavepointInner>, AbortState);

impl std::fmt::Debug for AsyncSavepointState {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        let state = self.0.lock().unwrap_or_else(|error| error.into_inner());
        f.debug_struct("AsyncSavepointState")
            .field("active", &state.stack.len())
            .field("poisoned", &state.poisoned)
            .finish()
    }
}

impl AsyncSavepointState {
    pub fn new() -> Self {
        Self(
            Mutex::new(AsyncSavepointInner {
                next_id: 0,
                stack: Vec::new(),
                poisoned: false,
                waiters: HashMap::new(),
            }),
            AbortState::new(),
        )
    }

    /// The transaction's record of server errors (see [`AbortState`]).
    pub fn aborted(&self) -> &AbortState {
        &self.1
    }

    /// Reject use after a savepoint future was cancelled or cleanup failed.
    pub fn ensure_usable(&self) -> Result<()> {
        let state = self.0.lock().unwrap_or_else(|error| error.into_inner());
        if state.poisoned {
            Err(DrizzleError::TransactionError(
                "transaction is unusable after cancelled or failed savepoint cleanup".into(),
            ))
        } else {
            Ok(())
        }
    }

    fn begin(&self) -> Result<u64> {
        let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
        if state.poisoned {
            return Err(DrizzleError::TransactionError(
                "transaction is unusable after cancelled or failed savepoint cleanup".into(),
            ));
        }
        let id = state.next_id;
        state.next_id = state.next_id.wrapping_add(1);
        state.stack.push(id);
        Ok(id)
    }

    async fn wait_until_top(&self, id: u64) -> Result<()> {
        poll_fn(|context| {
            let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
            if state.poisoned {
                return Poll::Ready(Err(DrizzleError::TransactionError(
                    "transaction is unusable after cancelled or failed savepoint cleanup".into(),
                )));
            }
            if state.stack.last() == Some(&id) {
                state.waiters.remove(&id);
                Poll::Ready(Ok(()))
            } else if state.stack.contains(&id) {
                state.waiters.insert(id, context.waker().clone());
                Poll::Pending
            } else {
                Poll::Ready(Err(DrizzleError::TransactionError(
                    "savepoint ordering state was lost".into(),
                )))
            }
        })
        .await
    }

    fn finish(&self, id: u64) -> Result<()> {
        let (result, waiters) = {
            let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
            let result = if state.stack.pop() == Some(id) {
                Ok(())
            } else {
                state.poisoned = true;
                Err(DrizzleError::TransactionError(
                    "savepoint cleanup completed out of order".into(),
                ))
            };
            let waiters = state
                .waiters
                .drain()
                .map(|(_, waker)| waker)
                .collect::<Vec<_>>();
            (result, waiters)
        };
        for waker in waiters {
            waker.wake();
        }
        result
    }

    fn poison(&self) {
        let waiters = {
            let mut state = self.0.lock().unwrap_or_else(|error| error.into_inner());
            state.poisoned = true;
            state
                .waiters
                .drain()
                .map(|(_, waker)| waker)
                .collect::<Vec<_>>()
        };
        for waker in waiters {
            waker.wake();
        }
    }
}

struct AsyncSavepointGuard<'a> {
    state: &'a AsyncSavepointState,
    armed: bool,
}

impl AsyncSavepointGuard<'_> {
    fn disarm(&mut self) {
        self.armed = false;
    }
}

impl Drop for AsyncSavepointGuard<'_> {
    fn drop(&mut self) {
        if self.armed {
            self.state.poison();
        }
    }
}

/// Run a synchronous transaction around `body`.
///
/// The driver owns how a transaction is begun and supplies commit/rollback
/// closures for its transaction type. This helper centralizes the shared
/// callback protocol: commit on `Ok`, rollback on `Err`, and rollback before
/// resuming a panic.
pub fn sync_transaction<Tx, R>(
    transaction: Tx,
    trace_name: &'static str,
    trace_commit: impl Fn(),
    trace_rollback: impl Fn(),
    body: impl FnOnce(&Tx) -> Result<R>,
    commit: impl FnOnce(Tx) -> Result<()>,
    rollback: impl FnOnce(Tx) -> Result<()>,
) -> Result<R> {
    let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| body(&transaction)));

    match outcome {
        Ok(Ok(value)) => {
            trace_commit();
            commit(transaction)?;
            Ok(value)
        }
        Ok(Err(e)) => {
            trace_rollback();
            match rollback(transaction) {
                Ok(()) => Err(e),
                Err(rollback_err) => Err(cleanup_error("transaction", e, "rollback", rollback_err)),
            }
        }
        Err(panic_payload) => {
            trace_rollback();
            if let Err(err) = rollback(transaction) {
                trace_panic_cleanup_error("transaction", trace_name, "rollback", &err);
            }
            std::panic::resume_unwind(panic_payload);
        }
    }
}

/// Run a savepoint block synchronously around `body`.
///
/// `execute_raw` issues a raw, parameterless statement against the active
/// transaction; the helper invokes it three times: `SAVEPOINT`,
/// `RELEASE SAVEPOINT`, and `ROLLBACK TO SAVEPOINT`.
///
/// If `body` panics, the helper attempts `ROLLBACK TO SAVEPOINT` and
/// `RELEASE SAVEPOINT` (errors ignored) and re-raises the panic via
/// [`std::panic::resume_unwind`].
pub fn sync_savepoint<R>(
    depth: &AtomicU32,
    execute_raw: impl FnMut(&str) -> Result<()>,
    body: impl FnOnce() -> Result<R>,
) -> Result<R> {
    sync_savepoint_tracking(depth, &NEVER_ABORTED, execute_raw, body)
}

/// [`sync_savepoint`] for a transaction whose server errors are recorded in
/// `aborted`.
///
/// If a statement failed inside the savepoint, the body's `Ok` cannot be
/// released (the transaction is aborted), so the savepoint is rolled back
/// and the call fails. Rolling back to the savepoint recovers the
/// transaction and restores `aborted`. A failed `RELEASE` is also rolled
/// back before it is reported, so the enclosing transaction stays usable.
pub fn sync_savepoint_tracking<R>(
    depth: &AtomicU32,
    aborted: &AbortState,
    mut execute_raw: impl FnMut(&str) -> Result<()>,
    body: impl FnOnce() -> Result<R>,
) -> Result<R> {
    let level = depth.load(Ordering::Relaxed);
    let sp = format!("drizzle_sp_{level}");
    depth.store(level + 1, Ordering::Relaxed);

    execute_raw(&format!("SAVEPOINT {sp}"))?;
    let aborted_before = aborted.is_aborted();

    let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(body));

    depth.store(level, Ordering::Relaxed);

    let outcome = match outcome {
        Ok(Ok(_)) if aborted.is_aborted() && !aborted_before => {
            Ok(Err(failed_statement_in_savepoint()))
        }
        other => other,
    };

    match outcome {
        Ok(Ok(value)) => {
            if let Err(release_err) = execute_raw(&format!("RELEASE SAVEPOINT {sp}")) {
                if let Err(rollback_err) = execute_raw(&format!("ROLLBACK TO SAVEPOINT {sp}")) {
                    return Err(cleanup_error(
                        "savepoint release",
                        release_err,
                        "rollback to savepoint",
                        rollback_err,
                    ));
                }
                aborted.set(aborted_before);
                return Err(release_err);
            }
            Ok(value)
        }
        Ok(Err(e)) => {
            if let Err(rollback_err) = execute_raw(&format!("ROLLBACK TO SAVEPOINT {sp}")) {
                return Err(cleanup_error(
                    "savepoint",
                    e,
                    "rollback to savepoint",
                    rollback_err,
                ));
            }
            aborted.set(aborted_before);
            if let Err(release_err) = execute_raw(&format!("RELEASE SAVEPOINT {sp}")) {
                return Err(cleanup_error(
                    "savepoint",
                    e,
                    "release savepoint after rollback",
                    release_err,
                ));
            }
            Err(e)
        }
        Err(panic_payload) => {
            if let Err(err) = execute_raw(&format!("ROLLBACK TO SAVEPOINT {sp}")) {
                trace_panic_cleanup_error("savepoint", &sp, "rollback to savepoint", &err);
            } else {
                aborted.set(aborted_before);
            }
            if let Err(err) = execute_raw(&format!("RELEASE SAVEPOINT {sp}")) {
                trace_panic_cleanup_error(
                    "savepoint",
                    &sp,
                    "release savepoint after rollback",
                    &err,
                );
            }
            std::panic::resume_unwind(panic_payload);
        }
    }
}

/// Run a savepoint block asynchronously around `body`.
///
/// Overlapping futures may run their bodies concurrently, but cleanup waits
/// for LIFO order. Dropping this future poisons `state`; transaction drivers
/// must call [`AsyncSavepointState::ensure_usable`] before subsequent work so
/// the outer transaction is forced to roll back.
pub async fn async_savepoint<R, Exec, ExecFut, BodyFut>(
    state: &AsyncSavepointState,
    mut execute_raw: Exec,
    body: BodyFut,
) -> Result<R>
where
    Exec: FnMut(String) -> ExecFut,
    ExecFut: core::future::Future<Output = Result<()>>,
    BodyFut: core::future::Future<Output = Result<R>>,
{
    let id = state.begin()?;
    let sp = format!("drizzle_sp_{id}");
    let mut guard = AsyncSavepointGuard { state, armed: true };
    let aborted = state.aborted();

    execute_raw(format!("SAVEPOINT {sp}")).await?;
    let aborted_before = aborted.is_aborted();

    let outcome = body.await;
    state.wait_until_top(id).await?;

    // A statement that failed inside the savepoint aborted the transaction
    // even if the body returned `Ok`; only rolling back to the savepoint
    // recovers it.
    let outcome = match outcome {
        Ok(_) if aborted.is_aborted() && !aborted_before => Err(failed_statement_in_savepoint()),
        other => other,
    };

    match outcome {
        Ok(value) => {
            if let Err(release_err) = execute_raw(format!("RELEASE SAVEPOINT {sp}")).await {
                if let Err(rollback_err) = execute_raw(format!("ROLLBACK TO SAVEPOINT {sp}")).await
                {
                    return Err(cleanup_error(
                        "savepoint release",
                        release_err,
                        "rollback to savepoint",
                        rollback_err,
                    ));
                }
                aborted.set(aborted_before);
                state.finish(id)?;
                guard.disarm();
                return Err(release_err);
            }
            state.finish(id)?;
            guard.disarm();
            Ok(value)
        }
        Err(e) => {
            if let Err(rollback_err) = execute_raw(format!("ROLLBACK TO SAVEPOINT {sp}")).await {
                return Err(cleanup_error(
                    "savepoint",
                    e,
                    "rollback to savepoint",
                    rollback_err,
                ));
            }
            aborted.set(aborted_before);
            if let Err(release_err) = execute_raw(format!("RELEASE SAVEPOINT {sp}")).await {
                return Err(cleanup_error(
                    "savepoint",
                    e,
                    "release savepoint after rollback",
                    release_err,
                ));
            }
            state.finish(id)?;
            guard.disarm();
            Err(e)
        }
    }
}

#[cfg(test)]
mod tests {
    use super::{AsyncSavepointState, async_savepoint};
    use core::{future::Future, pin::Pin, task::Context};

    #[test]
    fn overlapping_async_savepoints_wait_for_lifo_cleanup() {
        let state = AsyncSavepointState::new();
        let outer = state.begin().expect("outer savepoint");
        let inner = state.begin().expect("inner savepoint");

        let mut outer_turn = Box::pin(state.wait_until_top(outer));
        let mut context = Context::from_waker(std::task::Waker::noop());
        assert!(Pin::new(&mut outer_turn).poll(&mut context).is_pending());

        state.finish(inner).expect("finish inner");
        assert!(Pin::new(&mut outer_turn).poll(&mut context).is_ready());
        state.finish(outer).expect("finish outer");
        state.ensure_usable().expect("state remains usable");
    }

    #[test]
    fn dropping_async_savepoint_future_poisons_state() {
        let state = AsyncSavepointState::new();
        let mut future = Box::pin(async_savepoint(
            &state,
            |_| std::future::ready(Ok(())),
            std::future::pending::<drizzle_core::error::Result<()>>(),
        ));
        let mut context = Context::from_waker(std::task::Waker::noop());
        assert!(Pin::new(&mut future).poll(&mut context).is_pending());

        drop(future);
        assert!(state.ensure_usable().is_err());
        assert!(state.begin().is_err());
    }
}