Skip to main content

a3s_orm/drivers/sqlite/
savepoint.rs

1use async_trait::async_trait;
2use tokio::sync::OwnedMutexGuard;
3
4use crate::{ExecuteResult, Executor, QueryResult};
5
6use super::{SqliteError, SqliteExecutor, SqliteRow};
7
8pub struct SqliteSavepoint {
9    executor: SqliteExecutor,
10    name: String,
11    guard: Option<OwnedMutexGuard<()>>,
12    completed: bool,
13}
14
15impl SqliteSavepoint {
16    pub(crate) async fn begin(
17        executor: SqliteExecutor,
18        guard: OwnedMutexGuard<()>,
19        id: u64,
20    ) -> Result<Self, SqliteError> {
21        let name = format!("a3s_sp_{id}");
22        executor
23            .execute_control(format!("SAVEPOINT \"{name}\""))
24            .await?;
25        Ok(Self {
26            executor,
27            name,
28            guard: Some(guard),
29            completed: false,
30        })
31    }
32
33    pub(crate) async fn release(mut self) -> Result<(), SqliteError> {
34        self.executor
35            .execute_control(format!("RELEASE SAVEPOINT \"{}\"", self.name))
36            .await?;
37        self.completed = true;
38        Ok(())
39    }
40
41    pub(crate) async fn rollback(mut self) -> Result<(), SqliteError> {
42        self.executor
43            .execute_control(cleanup_sql(&self.name))
44            .await?;
45        self.completed = true;
46        Ok(())
47    }
48}
49
50#[async_trait]
51impl Executor for SqliteSavepoint {
52    type Row = SqliteRow;
53    type Error = SqliteError;
54
55    async fn execute(&self, query: &crate::CompiledQuery) -> Result<ExecuteResult, Self::Error> {
56        self.executor.execute_unlocked(query).await
57    }
58
59    async fn fetch_all(
60        &self,
61        query: &crate::CompiledQuery,
62    ) -> Result<QueryResult<Self::Row>, Self::Error> {
63        self.executor.fetch_all_unlocked(query).await
64    }
65}
66
67impl Drop for SqliteSavepoint {
68    fn drop(&mut self) {
69        if self.completed {
70            return;
71        }
72        let Some(guard) = self.guard.take() else {
73            return;
74        };
75        let executor = self.executor.clone();
76        let sql = cleanup_sql(&self.name);
77        if let Ok(runtime) = tokio::runtime::Handle::try_current() {
78            runtime.spawn(async move {
79                let _guard = guard;
80                let _ = executor.execute_control(sql).await;
81            });
82        }
83    }
84}
85
86fn cleanup_sql(name: &str) -> String {
87    format!("ROLLBACK TO SAVEPOINT \"{name}\"; RELEASE SAVEPOINT \"{name}\"")
88}