use std::sync::{Arc, Mutex};
use std::time::Instant;
use anyhow::{Result, anyhow};
use tokio::sync::{OwnedMutexGuard, OwnedSemaphorePermit};
use super::config::Engine;
use super::connection::{Connection, query_outcome};
use super::dedicated::Dedicated;
use super::health::ConnectionTrouble;
use super::query::QueryResult;
use super::query_log::{QueryOutcome, QuerySource};
use super::{runtime, script, statement};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum TxnState {
#[default]
Idle,
Open,
Failed,
}
impl TxnState {
pub fn is_open(self) -> bool {
self != TxnState::Idle
}
}
#[derive(Clone)]
pub struct PinnedConnection {
connection: Arc<Connection>,
shared: Arc<Shared>,
}
struct Shared {
held: Arc<tokio::sync::Mutex<Held>>,
tracker: Mutex<Tracker>,
backend: Mutex<Option<u64>>,
}
#[derive(Default)]
pub(crate) struct Held {
session: Option<(Dedicated, OwnedSemaphorePermit)>,
}
impl Drop for Held {
fn drop(&mut self) {
if let Some(session) = self.session.take() {
let_go(session);
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct Tracker {
state: TxnState,
manual: bool,
}
impl PinnedConnection {
pub fn new(connection: Arc<Connection>) -> Self {
Self {
connection,
shared: Arc::new(Shared {
held: Arc::new(tokio::sync::Mutex::new(Held::default())),
tracker: Mutex::new(Tracker::default()),
backend: Mutex::new(None),
}),
}
}
pub fn connection(&self) -> &Arc<Connection> {
&self.connection
}
pub fn state(&self) -> TxnState {
self.shared
.tracker
.lock()
.map(|tracker| tracker.state)
.unwrap_or_default()
}
pub fn is_busy(&self) -> bool {
self.shared.held.try_lock().is_err()
}
pub async fn run_query(&self, sql: &str) -> Result<QueryResult> {
self.connection.refuse_write(sql)?;
let started = Instant::now();
let result = self.fetch(sql).await;
let (elapsed, outcome) = match &result {
Ok(result) => (result.elapsed, query_outcome(result)),
Err(error) => (started.elapsed(), QueryOutcome::Error(format!("{error:#}"))),
};
self.connection
.log(sql, QuerySource::User, elapsed, outcome);
result
}
async fn fetch(&self, sql: &str) -> Result<QueryResult> {
let mut guard = self.lock().await?;
let result = guard.fetch(sql).await;
guard.note(sql, result.as_ref().err());
if result
.as_ref()
.is_err_and(|error| ConnectionTrouble::of(error).is_some())
{
guard.discard();
} else {
guard.settle().await;
}
result
}
pub(crate) async fn lock(&self) -> Result<PinGuard> {
let mut held = self
.shared
.held
.clone()
.try_lock_owned()
.map_err(|_| anyhow!("this tab is still running something"))?;
if held.session.is_none() {
let permit = self.connection.pinned_permit()?;
let mut session = self.connection.dedicated().await?;
let backend = session.backend_id().await.ok().flatten();
if let Ok(mut slot) = self.shared.backend.lock() {
*slot = backend;
}
if let Ok(mut tracker) = self.shared.tracker.lock() {
*tracker = Tracker::default();
}
held.session = Some((session, permit));
}
Ok(PinGuard {
held,
shared: self.shared.clone(),
engine: self.connection.config.engine,
})
}
pub fn cancel(&self) {
let backend = self.shared.backend.lock().ok().and_then(|slot| *slot);
let Some(backend) = backend else {
return;
};
let connection = self.connection.clone();
drop(runtime::spawn(async move {
connection.cancel_backend(backend).await.ok();
}));
}
}
pub(crate) struct PinGuard {
held: OwnedMutexGuard<Held>,
shared: Arc<Shared>,
engine: Engine,
}
impl PinGuard {
pub(crate) fn session(&mut self) -> Result<&mut Dedicated> {
self.held
.session
.as_mut()
.map(|(session, _)| session)
.ok_or_else(|| anyhow!("the tab's connection was closed"))
}
pub(crate) async fn fetch(&mut self, sql: &str) -> Result<QueryResult> {
let mut flight = Flight {
guard: self,
landed: false,
};
let result = flight.guard.session()?.fetch(sql).await;
flight.landed = true;
result
}
pub(crate) async fn execute(&mut self, sql: &str) -> Result<()> {
let mut flight = Flight {
guard: self,
landed: false,
};
let result = flight.guard.session()?.execute(sql).await;
flight.landed = true;
result
}
pub(crate) fn note(&mut self, sql: &str, error: Option<&anyhow::Error>) {
if self.engine != Engine::MySql {
return;
}
if let Ok(mut tracker) = self.shared.tracker.lock() {
*tracker = mysql_after(*tracker, sql, error.and_then(mysql_error_number));
}
}
pub(crate) async fn settle(&mut self) {
let status = match self.session() {
Ok(session) => session.transaction_status().await,
Err(_) => return,
};
match status {
Ok(Some(state)) => {
if let Ok(mut tracker) = self.shared.tracker.lock() {
tracker.state = state;
}
}
Ok(None) => {}
Err(_) => self.discard(),
}
}
pub(crate) fn discard(&mut self) {
if let Some(session) = self.held.session.take() {
let_go(session);
}
if let Ok(mut tracker) = self.shared.tracker.lock() {
*tracker = Tracker::default();
}
if let Ok(mut backend) = self.shared.backend.lock() {
*backend = None;
}
}
pub(crate) fn state(&self) -> TxnState {
self.shared
.tracker
.lock()
.map(|tracker| tracker.state)
.unwrap_or_default()
}
}
struct Flight<'a> {
guard: &'a mut PinGuard,
landed: bool,
}
impl Drop for Flight<'_> {
fn drop(&mut self) {
if !self.landed {
self.guard.discard();
}
}
}
fn let_go(session: (Dedicated, OwnedSemaphorePermit)) {
if tokio::runtime::Handle::try_current().is_ok() {
drop(session);
} else {
drop(runtime::spawn(async move { drop(session) }));
}
}
fn mysql_error_number(error: &anyhow::Error) -> Option<u16> {
error
.downcast_ref::<sqlx::Error>()
.and_then(|error| error.as_database_error())
.and_then(|error| error.try_downcast_ref::<sqlx::mysql::MySqlDatabaseError>())
.map(|error| error.number())
}
const MYSQL_DEADLOCK: u16 = 1213;
fn mysql_after(tracker: Tracker, sql: &str, error: Option<u16>) -> Tracker {
let words = statement::leading_words(sql, 6, Engine::MySql);
let word = |index: usize| words.get(index).map(String::as_str).unwrap_or("");
let ok = error.is_none();
let open = Tracker {
state: TxnState::Open,
..tracker
};
let idle = Tracker {
state: TxnState::Idle,
..tracker
};
if error == Some(MYSQL_DEADLOCK) {
return idle;
}
match word(0) {
"BEGIN" if !ok => tracker,
"BEGIN" => open,
"START" if word(1) == "TRANSACTION" => {
if ok {
open
} else {
tracker
}
}
"COMMIT" | "ROLLBACK" if words.iter().any(|word| word == "TO") => tracker,
"COMMIT" | "ROLLBACK" if !ok => tracker,
"COMMIT" | "ROLLBACK" if chains(&words) => open,
"COMMIT" | "ROLLBACK" => idle,
"SET" if words.iter().any(|word| word == "AUTOCOMMIT") => {
if !ok {
return tracker;
}
let at = words
.iter()
.position(|word| word == "AUTOCOMMIT")
.unwrap_or_default();
let value = match word(at + 1) {
"=" => word(at + 2),
value => value,
};
match value {
"1" | "ON" | "TRUE" => Tracker {
state: TxnState::Idle,
manual: false,
},
_ => Tracker {
manual: true,
..tracker
},
}
}
_ if script::implicitly_commits(Engine::MySql, &words) => {
idle
}
_ if ok && tracker.manual => open,
_ => tracker,
}
}
fn chains(words: &[String]) -> bool {
words
.iter()
.position(|word| word == "CHAIN")
.is_some_and(|at| at == 0 || words[at - 1] != "NO")
}
#[cfg(test)]
mod tests {
use super::*;
fn after(steps: &[(&str, Option<u16>)]) -> TxnState {
steps
.iter()
.fold(Tracker::default(), |tracker, (sql, error)| {
mysql_after(tracker, sql, *error)
})
.state
}
#[test]
fn mysql_transaction_control_opens_and_ends_one() {
assert_eq!(after(&[("begin", None)]), TxnState::Open);
assert_eq!(
after(&[("start transaction read only", None)]),
TxnState::Open
);
assert_eq!(after(&[("begin", None), ("commit", None)]), TxnState::Idle);
assert_eq!(
after(&[("begin", None), ("rollback", None)]),
TxnState::Idle
);
assert_eq!(
after(&[
("begin", None),
("savepoint a", None),
("rollback to a", None)
]),
TxnState::Open
);
assert_eq!(
after(&[("begin", None), ("commit and chain", None)]),
TxnState::Open
);
assert_eq!(
after(&[("begin", None), ("commit and no chain", None)]),
TxnState::Idle
);
assert_eq!(
after(&[
("begin", None),
("commit work and no chain no release", None)
]),
TxnState::Idle
);
assert_eq!(
after(&[("begin", None), ("commit work and chain no release", None)]),
TxnState::Open
);
assert_eq!(after(&[("begin", Some(1064))]), TxnState::Idle);
}
#[test]
fn mysql_implicit_commits_end_a_transaction() {
assert_eq!(
after(&[("begin", None), ("create table t (a int)", None)]),
TxnState::Idle
);
assert_eq!(
after(&[("begin", None), ("lock tables t write", None)]),
TxnState::Idle
);
assert_eq!(
after(&[("begin", None), ("alter table nope add b int", Some(1146))]),
TxnState::Idle
);
assert_eq!(
after(&[("begin", None), ("create temporary table t (a int)", None)]),
TxnState::Open
);
}
#[test]
fn mysql_errors_keep_a_transaction_open_except_a_deadlock() {
assert_eq!(
after(&[("begin", None), ("insert into nope values (1)", Some(1146))]),
TxnState::Open
);
assert_eq!(
after(&[
("begin", None),
("update t set a = 1", Some(MYSQL_DEADLOCK))
]),
TxnState::Idle
);
}
#[test]
fn mysql_autocommit_off_opens_a_transaction_on_every_statement() {
assert_eq!(after(&[("set autocommit = 0", None)]), TxnState::Idle);
assert_eq!(
after(&[("set @@autocommit = 0", None), ("select 1", None)]),
TxnState::Open
);
assert_eq!(
after(&[
("set autocommit = 0", None),
("insert into t values (1)", None)
]),
TxnState::Open
);
assert_eq!(
after(&[
("set autocommit = 0", None),
("insert into t values (1)", None),
("commit", None),
]),
TxnState::Idle
);
assert_eq!(
after(&[
("set autocommit = 0", None),
("insert into t values (1)", None),
("commit", None),
("select 1", None),
]),
TxnState::Open
);
assert_eq!(
after(&[
("set session autocommit = off", None),
("insert into t values (1)", None),
("set autocommit = 1", None),
("select 1", None),
]),
TxnState::Idle
);
}
#[test]
fn mysql_autocommit_spelled_as_a_boolean_or_quoted_is_followed() {
for off in ["set autocommit = false", "set autocommit = 'OFF'"] {
assert_eq!(
after(&[(off, None), ("insert into t values (1)", None)]),
TxnState::Open,
"{off}"
);
}
assert_eq!(
after(&[
("set autocommit = false", None),
("insert into t values (1)", None),
("set autocommit = true", None),
]),
TxnState::Idle
);
}
#[test]
fn plain_statements_leave_mysql_idle() {
assert_eq!(
after(&[("select 1", None), ("insert into t values (1)", None)]),
TxnState::Idle
);
}
}