use parking_lot::{RwLock, RwLockReadGuard, RwLockWriteGuard};
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicU64, AtomicU8, Ordering};
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use uqa_sql::semantics::parameters::catalog::{find_parameter, message_levels};
use uqa_sql::semantics::parameters::definition::ParameterKind;
use crate::SessionStateSnapshot;
pub(crate) struct SessionStateLock {
state: RwLock<SessionStateSnapshot>,
client_level: Arc<AtomicU8>,
reset_client_level: AtomicU8,
cancellation: OnceLock<uqa_core::CancellationToken>,
reset_lock_timeout_ms: AtomicU64,
}
impl SessionStateLock {
pub(crate) fn new(state: SessionStateSnapshot) -> Self {
Self {
state: RwLock::new(state),
client_level: Arc::new(AtomicU8::new(message_levels::NOTICE)),
reset_client_level: AtomicU8::new(message_levels::NOTICE),
cancellation: OnceLock::new(),
reset_lock_timeout_ms: AtomicU64::new(0),
}
}
pub(crate) fn attach_cancellation(&self, token: &uqa_core::CancellationToken) {
if self.cancellation.set(token.clone()).is_ok() {
drop(self.write());
}
}
pub(crate) fn read(&self) -> RwLockReadGuard<'_, SessionStateSnapshot> {
self.state.read()
}
pub(crate) fn try_read_for(
&self,
timeout: std::time::Duration,
) -> Option<RwLockReadGuard<'_, SessionStateSnapshot>> {
self.state.try_read_for(timeout)
}
#[cfg(test)]
pub(crate) fn try_write(&self) -> Option<SessionStateWriteGuard<'_>> {
self.state
.try_write()
.map(|guard| SessionStateWriteGuard { guard, lock: self })
}
#[cfg(test)]
pub(crate) fn is_locked(&self) -> bool {
self.state.is_locked()
}
pub(crate) fn write(&self) -> SessionStateWriteGuard<'_> {
SessionStateWriteGuard {
guard: self.state.write(),
lock: self,
}
}
pub(crate) fn client_level(&self) -> Arc<AtomicU8> {
Arc::clone(&self.client_level)
}
pub(crate) fn set_reset_client_level(&self, level: u8) {
self.reset_client_level.store(level, Ordering::Release);
drop(self.write());
}
pub(crate) fn set_reset_lock_timeout(&self, milliseconds: u64) {
self.reset_lock_timeout_ms
.store(milliseconds, Ordering::Release);
drop(self.write());
}
}
pub(crate) struct SessionStateWriteGuard<'a> {
guard: RwLockWriteGuard<'a, SessionStateSnapshot>,
lock: &'a SessionStateLock,
}
impl Deref for SessionStateWriteGuard<'_> {
type Target = SessionStateSnapshot;
fn deref(&self) -> &Self::Target {
&self.guard
}
}
impl DerefMut for SessionStateWriteGuard<'_> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.guard
}
}
impl Drop for SessionStateWriteGuard<'_> {
fn drop(&mut self) {
let level = self
.guard
.session_vars
.get("client_min_messages")
.and_then(|setting| message_level(setting))
.unwrap_or_else(|| self.lock.reset_client_level.load(Ordering::Acquire));
self.lock.client_level.store(level, Ordering::Release);
if let Some(token) = self.lock.cancellation.get() {
let milliseconds = self
.guard
.session_vars
.get("lock_timeout")
.and_then(|setting| setting.parse::<u64>().ok())
.unwrap_or_else(|| self.lock.reset_lock_timeout_ms.load(Ordering::Acquire));
token
.set_lock_timeout((milliseconds != 0).then(|| Duration::from_millis(milliseconds)));
}
}
}
pub(crate) fn message_level(setting: &str) -> Option<u8> {
let ParameterKind::Enum { options, .. } = find_parameter("client_min_messages")?.kind else {
return None;
};
options
.iter()
.find(|option| option.name == setting)
.map(|option| option.value)
}