use std::collections::HashMap;
use std::ffi::{c_int, c_uint, c_void, CStr};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, OnceLock, Weak};
use parking_lot::Mutex;
use rusqlite::{ffi, Connection};
use crate::SqliteError;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StartedStatement {
pub sql: String,
pub readonly: bool,
}
#[derive(Default)]
struct Records {
statements: Vec<StartedStatement>,
lost: bool,
}
struct Probe {
limit: usize,
records: Mutex<Records>,
}
static NEXT_HUB: AtomicUsize = AtomicUsize::new(1);
static ACTIVE_OBSERVERS: AtomicUsize = AtomicUsize::new(0);
static HUBS: OnceLock<Mutex<HashMap<usize, Weak<StatementObserverHub>>>> = OnceLock::new();
fn hubs() -> &'static Mutex<HashMap<usize, Weak<StatementObserverHub>>> {
HUBS.get_or_init(Mutex::default)
}
pub(crate) struct StatementObserverHub {
id: usize,
active: Mutex<Option<Arc<Probe>>>,
}
pub struct StatementStartObservation {
hub: Arc<StatementObserverHub>,
probe: Arc<Probe>,
}
impl StatementObserverHub {
pub(crate) fn new() -> Result<Arc<Self>, SqliteError> {
let id = NEXT_HUB
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |id| id.checked_add(1))
.map_err(|_| {
SqliteError::InvalidData("statement observer identity exhausted".into())
})?;
let hub = Arc::new(Self {
id,
active: Mutex::new(None),
});
hubs().lock().insert(id, Arc::downgrade(&hub));
Ok(hub)
}
pub(crate) fn observe(
self: &Arc<Self>,
limit: usize,
) -> Result<StatementStartObservation, SqliteError> {
if limit == 0 {
return Err(SqliteError::InvalidData(
"statement observation requires a nonzero record limit".into(),
));
}
let mut active = self.active.lock();
if active.is_some() {
return Err(SqliteError::InvalidData(
"a statement observation is already active for this pool".into(),
));
}
let probe = Arc::new(Probe {
limit,
records: Mutex::new(Records::default()),
});
*active = Some(Arc::clone(&probe));
ACTIVE_OBSERVERS.fetch_add(1, Ordering::Release);
Ok(StatementStartObservation {
hub: Arc::clone(self),
probe,
})
}
}
impl Drop for StatementObserverHub {
fn drop(&mut self) {
hubs().lock().remove(&self.id);
}
}
impl StatementStartObservation {
pub fn started_statements(&self) -> Result<Vec<StartedStatement>, SqliteError> {
let records = self.probe.records.lock();
if records.lost {
return Err(SqliteError::InvalidData(
"statement observation lost records; its count is incomplete".into(),
));
}
Ok(records.statements.clone())
}
}
impl Drop for StatementStartObservation {
fn drop(&mut self) {
let mut active = self.hub.active.lock();
if active
.as_ref()
.is_some_and(|probe| Arc::ptr_eq(probe, &self.probe))
{
*active = None;
ACTIVE_OBSERVERS.fetch_sub(1, Ordering::Release);
}
}
}
pub(crate) fn install(
conn: &Connection,
hub: &Arc<StatementObserverHub>,
) -> Result<(), SqliteError> {
let result = unsafe {
ffi::sqlite3_trace_v2(
conn.handle(),
ffi::SQLITE_TRACE_STMT as c_uint,
Some(trace),
hub.id as *mut c_void,
)
};
if result != ffi::SQLITE_OK {
return Err(rusqlite::Error::SqliteFailure(ffi::Error::new(result), None).into());
}
Ok(())
}
unsafe extern "C" fn trace(
event: c_uint,
context: *mut c_void,
statement: *mut c_void,
text: *mut c_void,
) -> c_int {
if event != ffi::SQLITE_TRACE_STMT as c_uint {
return 0;
}
if ACTIVE_OBSERVERS.load(Ordering::Acquire) == 0 {
return 0;
}
let hub = {
#[cfg(test)]
trace_counters::lock_acquired();
let registry = hubs().lock();
registry.get(&(context as usize)).and_then(Weak::upgrade)
};
let Some(hub) = hub else { return 0 };
#[cfg(test)]
trace_counters::lock_acquired();
let active = hub.active.lock();
let Some(probe) = active.as_ref() else {
return 0;
};
#[cfg(test)]
trace_counters::lock_acquired();
let mut records = probe.records.lock();
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let stmt = statement.cast::<ffi::sqlite3_stmt>();
let sql = unsafe { ffi::sqlite3_sql(stmt) };
if sql.is_null() || text.is_null() {
records.lost = true;
return;
}
let original = unsafe { CStr::from_ptr(sql) };
let reported = unsafe { CStr::from_ptr(text.cast()) };
if original.to_bytes() != reported.to_bytes() {
return;
}
if records.statements.len() == probe.limit {
records.lost = true;
return;
}
records.statements.push(StartedStatement {
sql: original.to_string_lossy().into_owned(),
readonly: unsafe { ffi::sqlite3_stmt_readonly(stmt) != 0 },
});
#[cfg(test)]
trace_counters::record_added();
}));
if result.is_err() {
records.lost = true;
}
0
}
#[cfg(test)]
mod trace_counters {
use std::cell::Cell;
std::thread_local! {
static COUNTS: Cell<(usize, usize)> = const { Cell::new((0, 0)) };
}
pub(super) fn lock_acquired() {
let _ = COUNTS.try_with(|counts| {
let (locks, records) = counts.get();
counts.set((locks.saturating_add(1), records));
});
}
pub(super) fn record_added() {
let _ = COUNTS.try_with(|counts| {
let (locks, records) = counts.get();
counts.set((locks, records.saturating_add(1)));
});
}
pub(super) fn reset() {
COUNTS.with(|counts| counts.set((0, 0)));
}
pub(super) fn snapshot() -> (usize, usize) {
COUNTS.with(Cell::get)
}
}
#[cfg(test)]
mod tests;