65 lines
2.6 KiB
Rust
65 lines
2.6 KiB
Rust
use crate::{
|
|
domain::sql::{CopyDirection, SqlLexer, SqlSubmission, StatementRange, TransactionState},
|
|
error::{AppError, SafeError},
|
|
};
|
|
use tokio::sync::mpsc;
|
|
use tokio_util::sync::CancellationToken;
|
|
|
|
pub const SQL_EVENT_CHANNEL_CAPACITY: usize = crate::domain::sql::SQL_EVENT_CHANNEL_CAPACITY;
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub enum ExecutorEvent {
|
|
StatementStarted { index: usize },
|
|
StatementCompleted { index: usize },
|
|
CopyRequired(CopyDirection),
|
|
Failed { index: usize },
|
|
Cancelled,
|
|
}
|
|
pub struct SqlExecutor;
|
|
impl SqlExecutor {
|
|
pub fn statements(submission: &SqlSubmission) -> Vec<StatementRange> {
|
|
SqlLexer::split(submission.as_str())
|
|
}
|
|
pub fn copy_mode(submission: &SqlSubmission, range: &StatementRange) -> Option<CopyDirection> {
|
|
SqlLexer::classify_copy(&submission.as_str()[range.range.clone()])
|
|
}
|
|
pub fn event_channel() -> (mpsc::Sender<ExecutorEvent>, mpsc::Receiver<ExecutorEvent>) {
|
|
mpsc::channel(SQL_EVENT_CHANNEL_CAPACITY)
|
|
}
|
|
pub async fn emit_plan(
|
|
submission: &SqlSubmission,
|
|
sender: mpsc::Sender<ExecutorEvent>,
|
|
cancellation: CancellationToken,
|
|
) -> Result<TransactionState, AppError> {
|
|
let ranges = Self::statements(submission);
|
|
let mut state = TransactionState::Idle;
|
|
for (index, range) in ranges.iter().enumerate() {
|
|
tokio::select! { _ = cancellation.cancelled() => { let _ = sender.send(ExecutorEvent::Cancelled).await; return Err(AppError(SafeError::input("Operation cancelled"))); }, result = sender.send(ExecutorEvent::StatementStarted { index }) => result.map_err(|_| AppError(SafeError::database()))? }
|
|
let statement = submission.as_str()[range.range.clone()].trim();
|
|
if let Some(copy) = SqlLexer::classify_copy(statement) {
|
|
sender
|
|
.send(ExecutorEvent::CopyRequired(copy))
|
|
.await
|
|
.map_err(|_| AppError(SafeError::database()))?;
|
|
} else {
|
|
state = transaction_state_for(statement, state);
|
|
sender
|
|
.send(ExecutorEvent::StatementCompleted { index })
|
|
.await
|
|
.map_err(|_| AppError(SafeError::database()))?;
|
|
}
|
|
}
|
|
Ok(state)
|
|
}
|
|
}
|
|
fn transaction_state_for(statement: &str, previous: TransactionState) -> TransactionState {
|
|
let upper = statement.trim().trim_end_matches(';').to_ascii_uppercase();
|
|
if upper == "BEGIN" || upper == "START TRANSACTION" {
|
|
TransactionState::Active
|
|
} else if upper == "COMMIT" || upper == "ROLLBACK" {
|
|
TransactionState::Idle
|
|
} else {
|
|
previous
|
|
}
|
|
}
|