use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use parking_lot::Mutex;
use std::time::Duration;
use crate::any_driver::AnyBackend;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum TransactionIsolation {
ReadUncommitted,
#[default]
ReadCommitted,
RepeatableRead,
Serializable,
}
impl TransactionIsolation {
pub fn name(&self) -> &'static str {
match self {
TransactionIsolation::ReadUncommitted => "READ UNCOMMITTED",
TransactionIsolation::ReadCommitted => "READ COMMITTED",
TransactionIsolation::RepeatableRead => "REPEATABLE READ",
TransactionIsolation::Serializable => "SERIALIZABLE",
}
}
pub fn description(&self) -> &'static str {
match self {
TransactionIsolation::ReadUncommitted => "读未提交",
TransactionIsolation::ReadCommitted => "读已提交",
TransactionIsolation::RepeatableRead => "可重复读",
TransactionIsolation::Serializable => "串行化",
}
}
pub fn strictness(&self) -> u8 {
match self {
TransactionIsolation::ReadUncommitted => 0,
TransactionIsolation::ReadCommitted => 1,
TransactionIsolation::RepeatableRead => 2,
TransactionIsolation::Serializable => 3,
}
}
pub fn set_session_sql(&self, backend: AnyBackend) -> String {
match backend {
AnyBackend::MySql => {
format!("SET SESSION TRANSACTION ISOLATION LEVEL {}", self.name())
}
AnyBackend::Postgres => {
format!(
"SET SESSION CHARACTERISTICS AS TRANSACTION ISOLATION LEVEL {}",
self.name()
)
}
AnyBackend::Sqlite => {
String::new()
}
}
}
pub fn set_transaction_sql(&self, backend: AnyBackend) -> String {
match backend {
AnyBackend::MySql => format!("SET TRANSACTION ISOLATION LEVEL {}", self.name()),
AnyBackend::Postgres => format!("SET TRANSACTION ISOLATION LEVEL {}", self.name()),
AnyBackend::Sqlite => String::new(),
}
}
pub fn query_sql(&self, backend: AnyBackend) -> String {
match backend {
AnyBackend::MySql => "SELECT @@transaction_isolation".to_string(),
AnyBackend::Postgres => "SHOW transaction_isolation".to_string(),
AnyBackend::Sqlite => String::new(),
}
}
#[allow(clippy::should_implement_trait)]
pub fn from_str(s: &str) -> Option<Self> {
let upper = s.to_uppercase().replace('_', " ");
match upper.as_str() {
"READ UNCOMMITTED" => Some(TransactionIsolation::ReadUncommitted),
"READ COMMITTED" => Some(TransactionIsolation::ReadCommitted),
"REPEATABLE READ" => Some(TransactionIsolation::RepeatableRead),
"SERIALIZABLE" => Some(TransactionIsolation::Serializable),
_ => None,
}
}
}
impl std::fmt::Display for TransactionIsolation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.name())
}
}
#[derive(Debug, Clone)]
pub struct EnhancedPoolConfig {
pub max_connections: u32,
pub min_idle: Option<u32>,
pub acquire_timeout: Duration,
pub idle_timeout: Option<Duration>,
pub max_lifetime: Option<Duration>,
pub test_on_acquire: bool,
pub test_query: String,
pub pool_name: Option<String>,
}
impl Default for EnhancedPoolConfig {
fn default() -> Self {
Self {
max_connections: 10,
min_idle: None,
acquire_timeout: Duration::from_secs(30),
idle_timeout: Some(Duration::from_secs(600)),
max_lifetime: Some(Duration::from_secs(1800)),
test_on_acquire: false,
test_query: "SELECT 1".to_string(),
pool_name: None,
}
}
}
impl EnhancedPoolConfig {
pub fn builder() -> EnhancedPoolConfigBuilder {
EnhancedPoolConfigBuilder::default()
}
pub fn validate(&self) -> Result<(), String> {
if self.max_connections == 0 {
return Err("max_connections 不能为 0".to_string());
}
if let Some(min) = self.min_idle {
if min > self.max_connections {
return Err(format!(
"min_idle ({}) 不能大于 max_connections ({})",
min, self.max_connections
));
}
}
if self.acquire_timeout.is_zero() {
return Err("acquire_timeout 不能为 0".to_string());
}
if self.test_query.is_empty() {
return Err("test_query 不能为空".to_string());
}
Ok(())
}
pub fn summary(&self) -> String {
format!(
"PoolConfig{{max={}, min_idle={:?}, timeout={}ms, test_on_acquire={}, name={:?}}}",
self.max_connections,
self.min_idle,
self.acquire_timeout.as_millis(),
self.test_on_acquire,
self.pool_name
)
}
}
#[derive(Debug, Clone, Default)]
pub struct EnhancedPoolConfigBuilder {
config: EnhancedPoolConfig,
}
impl EnhancedPoolConfigBuilder {
pub fn max_connections(mut self, n: u32) -> Self {
self.config.max_connections = n;
self
}
pub fn min_idle(mut self, n: u32) -> Self {
self.config.min_idle = Some(n);
self
}
pub fn acquire_timeout_secs(mut self, secs: u64) -> Self {
self.config.acquire_timeout = Duration::from_secs(secs);
self
}
pub fn acquire_timeout_millis(mut self, millis: u64) -> Self {
self.config.acquire_timeout = Duration::from_millis(millis);
self
}
pub fn idle_timeout_secs(mut self, secs: u64) -> Self {
self.config.idle_timeout = Some(Duration::from_secs(secs));
self
}
pub fn max_lifetime_secs(mut self, secs: u64) -> Self {
self.config.max_lifetime = Some(Duration::from_secs(secs));
self
}
pub fn test_on_acquire(mut self) -> Self {
self.config.test_on_acquire = true;
self
}
pub fn test_query(mut self, sql: &str) -> Self {
self.config.test_query = sql.to_string();
self
}
pub fn name(mut self, name: &str) -> Self {
self.config.pool_name = Some(name.to_string());
self
}
pub fn build(self) -> Result<EnhancedPoolConfig, String> {
self.config.validate()?;
Ok(self.config)
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
struct CacheEntry {
statement_id: String,
created_seq: u64,
last_access_seq: u64,
hit_count: u64,
}
#[derive(Debug, Clone, Default)]
pub struct CacheStats {
pub hits: u64,
pub misses: u64,
pub evictions: u64,
pub size: usize,
pub capacity: usize,
}
impl CacheStats {
pub fn hit_rate(&self) -> f64 {
let total = self.hits + self.misses;
if total == 0 {
return 0.0;
}
self.hits as f64 / total as f64
}
pub fn summary(&self) -> String {
format!(
"CacheStats{{hits={}, misses={}, evictions={}, size={}, capacity={}, hit_rate={:.2}%, capacity_utilization={:.2}%}}",
self.hits,
self.misses,
self.evictions,
self.size,
self.capacity,
self.hit_rate() * 100.0,
self.capacity_utilization() * 100.0,
)
}
pub fn capacity_utilization(&self) -> f64 {
if self.capacity == 0 {
return 0.0;
}
self.size as f64 / self.capacity as f64
}
pub fn total_accesses(&self) -> u64 {
self.hits + self.misses
}
}
pub struct PreparedStatementCache {
entries: Mutex<HashMap<u64, CacheEntry>>,
capacity: usize,
stats: Mutex<CacheStats>,
access_seq: AtomicU64,
}
impl PreparedStatementCache {
pub fn new(capacity: usize) -> Self {
let capacity = capacity.max(1);
Self {
entries: Mutex::new(HashMap::with_capacity(capacity)),
capacity,
stats: Mutex::new(CacheStats {
hits: 0,
misses: 0,
evictions: 0,
size: 0,
capacity,
}),
access_seq: AtomicU64::new(0),
}
}
fn hash_sql(sql: &str) -> u64 {
const FNV_OFFSET: u64 = 0xcbf29ce484222325;
const FNV_PRIME: u64 = 0x100000001b3;
let mut hash = FNV_OFFSET;
for byte in sql.as_bytes() {
hash ^= *byte as u64;
hash = hash.wrapping_mul(FNV_PRIME);
}
hash
}
fn next_seq(&self) -> u64 {
self.access_seq.fetch_add(1, Ordering::Relaxed)
}
pub fn get(&self, sql: &str) -> Option<String> {
let hash = Self::hash_sql(sql);
let seq = self.next_seq();
let mut entries = self.entries.lock();
let mut stats = self.stats.lock();
if let Some(entry) = entries.get_mut(&hash) {
entry.last_access_seq = seq;
entry.hit_count += 1;
stats.hits += 1;
Some(entry.statement_id.clone())
} else {
stats.misses += 1;
None
}
}
pub fn put(&self, sql: &str, statement_id: &str) {
let hash = Self::hash_sql(sql);
let seq = self.next_seq();
let mut entries = self.entries.lock();
let mut stats = self.stats.lock();
if let Some(entry) = entries.get_mut(&hash) {
entry.statement_id = statement_id.to_string();
entry.last_access_seq = seq;
return;
}
if entries.len() >= self.capacity {
if let Some(&evict_hash) = entries
.iter()
.min_by_key(|(_, entry)| entry.last_access_seq)
.map(|(k, _)| k)
{
entries.remove(&evict_hash);
stats.evictions += 1;
}
}
entries.insert(
hash,
CacheEntry {
statement_id: statement_id.to_string(),
created_seq: seq,
last_access_seq: seq,
hit_count: 0,
},
);
stats.size = entries.len();
}
pub fn remove(&self, sql: &str) -> bool {
let hash = Self::hash_sql(sql);
let mut entries = self.entries.lock();
let mut stats = self.stats.lock();
let removed = entries.remove(&hash).is_some();
if removed {
stats.size = entries.len();
}
removed
}
pub fn clear(&self) {
let mut entries = self.entries.lock();
let mut stats = self.stats.lock();
entries.clear();
stats.size = 0;
}
pub fn stats(&self) -> CacheStats {
let stats = self.stats.lock();
stats.clone()
}
pub fn capacity(&self) -> usize {
self.capacity
}
pub fn len(&self) -> usize {
self.entries.lock().len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn reset_stats(&self) {
let mut stats = self.stats.lock();
stats.hits = 0;
stats.misses = 0;
stats.evictions = 0;
}
}
impl Default for PreparedStatementCache {
fn default() -> Self {
Self::new(256)
}
}
impl std::fmt::Debug for PreparedStatementCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let stats = self.stats();
f.debug_struct("PreparedStatementCache")
.field("capacity", &self.capacity)
.field("size", &stats.size)
.field("hits", &stats.hits)
.field("misses", &stats.misses)
.field("evictions", &stats.evictions)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_isolation_level_names() {
assert_eq!(
TransactionIsolation::ReadUncommitted.name(),
"READ UNCOMMITTED"
);
assert_eq!(TransactionIsolation::ReadCommitted.name(), "READ COMMITTED");
assert_eq!(
TransactionIsolation::RepeatableRead.name(),
"REPEATABLE READ"
);
assert_eq!(TransactionIsolation::Serializable.name(), "SERIALIZABLE");
}
#[test]
fn test_isolation_level_descriptions() {
assert_eq!(
TransactionIsolation::ReadUncommitted.description(),
"读未提交"
);
assert_eq!(
TransactionIsolation::ReadCommitted.description(),
"读已提交"
);
assert_eq!(
TransactionIsolation::RepeatableRead.description(),
"可重复读"
);
assert_eq!(TransactionIsolation::Serializable.description(), "串行化");
}
#[test]
fn test_isolation_level_strictness_order() {
assert!(
TransactionIsolation::ReadUncommitted.strictness()
< TransactionIsolation::ReadCommitted.strictness()
);
assert!(
TransactionIsolation::ReadCommitted.strictness()
< TransactionIsolation::RepeatableRead.strictness()
);
assert!(
TransactionIsolation::RepeatableRead.strictness()
< TransactionIsolation::Serializable.strictness()
);
}
#[test]
fn test_isolation_level_from_str() {
assert_eq!(
TransactionIsolation::from_str("READ COMMITTED"),
Some(TransactionIsolation::ReadCommitted)
);
assert_eq!(
TransactionIsolation::from_str("read committed"),
Some(TransactionIsolation::ReadCommitted)
);
assert_eq!(
TransactionIsolation::from_str("READ_COMMITTED"),
Some(TransactionIsolation::ReadCommitted)
);
assert_eq!(
TransactionIsolation::from_str("SERIALIZABLE"),
Some(TransactionIsolation::Serializable)
);
assert_eq!(TransactionIsolation::from_str("UNKNOWN"), None);
}
#[test]
fn test_isolation_level_set_session_sql_mysql() {
let sql = TransactionIsolation::ReadCommitted.set_session_sql(AnyBackend::MySql);
assert!(sql.contains("SET SESSION TRANSACTION ISOLATION LEVEL"));
assert!(sql.contains("READ COMMITTED"));
}
#[test]
fn test_isolation_level_set_session_sql_postgres() {
let sql = TransactionIsolation::Serializable.set_session_sql(AnyBackend::Postgres);
assert!(sql.contains("SET SESSION CHARACTERISTICS AS TRANSACTION ISOLATION LEVEL"));
assert!(sql.contains("SERIALIZABLE"));
}
#[test]
fn test_isolation_level_set_session_sql_sqlite_empty() {
let sql = TransactionIsolation::ReadCommitted.set_session_sql(AnyBackend::Sqlite);
assert!(sql.is_empty(), "SQLite 不支持设置隔离级别");
}
#[test]
fn test_isolation_level_set_transaction_sql() {
let mysql_sql = TransactionIsolation::RepeatableRead.set_transaction_sql(AnyBackend::MySql);
assert!(mysql_sql.contains("SET TRANSACTION ISOLATION LEVEL"));
assert!(mysql_sql.contains("REPEATABLE READ"));
let pg_sql = TransactionIsolation::RepeatableRead.set_transaction_sql(AnyBackend::Postgres);
assert!(pg_sql.contains("SET TRANSACTION ISOLATION LEVEL"));
let sqlite_sql =
TransactionIsolation::RepeatableRead.set_transaction_sql(AnyBackend::Sqlite);
assert!(sqlite_sql.is_empty());
}
#[test]
fn test_isolation_level_query_sql() {
assert_eq!(
TransactionIsolation::ReadCommitted.query_sql(AnyBackend::MySql),
"SELECT @@transaction_isolation"
);
assert_eq!(
TransactionIsolation::ReadCommitted.query_sql(AnyBackend::Postgres),
"SHOW transaction_isolation"
);
assert_eq!(
TransactionIsolation::ReadCommitted.query_sql(AnyBackend::Sqlite),
""
);
}
#[test]
fn test_isolation_level_display() {
let level = TransactionIsolation::Serializable;
let s = format!("{}", level);
assert_eq!(s, "SERIALIZABLE");
}
#[test]
fn test_isolation_level_default() {
let level = TransactionIsolation::default();
assert_eq!(level, TransactionIsolation::ReadCommitted);
}
#[test]
fn test_isolation_level_equality() {
assert_eq!(
TransactionIsolation::ReadCommitted,
TransactionIsolation::ReadCommitted
);
assert_ne!(
TransactionIsolation::ReadCommitted,
TransactionIsolation::Serializable
);
}
#[test]
fn test_pool_config_default() {
let config = EnhancedPoolConfig::default();
assert_eq!(config.max_connections, 10);
assert!(config.min_idle.is_none());
assert_eq!(config.acquire_timeout, Duration::from_secs(30));
assert_eq!(config.test_query, "SELECT 1");
assert!(!config.test_on_acquire);
}
#[test]
fn test_pool_config_builder_basic() {
let config = EnhancedPoolConfig::builder()
.max_connections(20)
.min_idle(5)
.acquire_timeout_secs(60)
.build()
.unwrap();
assert_eq!(config.max_connections, 20);
assert_eq!(config.min_idle, Some(5));
assert_eq!(config.acquire_timeout, Duration::from_secs(60));
}
#[test]
fn test_pool_config_builder_test_on_acquire() {
let config = EnhancedPoolConfig::builder()
.test_on_acquire()
.test_query("SELECT 1 FROM dual")
.build()
.unwrap();
assert!(config.test_on_acquire);
assert_eq!(config.test_query, "SELECT 1 FROM dual");
}
#[test]
fn test_pool_config_builder_with_name() {
let config = EnhancedPoolConfig::builder()
.name("primary-pool")
.build()
.unwrap();
assert_eq!(config.pool_name, Some("primary-pool".to_string()));
}
#[test]
fn test_pool_config_validate_max_connections_zero() {
let config = EnhancedPoolConfig {
max_connections: 0,
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_pool_config_validate_min_idle_exceeds_max() {
let config = EnhancedPoolConfig {
max_connections: 5,
min_idle: Some(10),
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_pool_config_validate_timeout_zero() {
let config = EnhancedPoolConfig {
acquire_timeout: Duration::from_secs(0),
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_pool_config_validate_empty_test_query() {
let config = EnhancedPoolConfig {
test_query: "".to_string(),
..Default::default()
};
assert!(config.validate().is_err());
}
#[test]
fn test_pool_config_validate_valid() {
let config = EnhancedPoolConfig::default();
assert!(config.validate().is_ok());
}
#[test]
fn test_pool_config_summary() {
let config = EnhancedPoolConfig::builder()
.max_connections(15)
.name("test-pool")
.build()
.unwrap();
let summary = config.summary();
assert!(summary.contains("max=15"));
assert!(summary.contains("name=Some(\"test-pool\")"));
}
#[test]
fn test_pool_config_builder_millis_timeout() {
let config = EnhancedPoolConfig::builder()
.acquire_timeout_millis(500)
.build()
.unwrap();
assert_eq!(config.acquire_timeout, Duration::from_millis(500));
}
#[test]
fn test_pool_config_builder_idle_and_lifetime() {
let config = EnhancedPoolConfig::builder()
.idle_timeout_secs(300)
.max_lifetime_secs(900)
.build()
.unwrap();
assert_eq!(config.idle_timeout, Some(Duration::from_secs(300)));
assert_eq!(config.max_lifetime, Some(Duration::from_secs(900)));
}
#[test]
fn test_cache_basic_put_and_get() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT * FROM users WHERE id = ?", "stmt_1");
let result = cache.get("SELECT * FROM users WHERE id = ?");
assert_eq!(result, Some("stmt_1".to_string()));
}
#[test]
fn test_cache_miss() {
let cache = PreparedStatementCache::new(10);
let result = cache.get("SELECT * FROM nonexist");
assert!(result.is_none());
let stats = cache.stats();
assert_eq!(stats.misses, 1);
assert_eq!(stats.hits, 0);
}
#[test]
fn test_cache_hit_increments_counter() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.get("SELECT 1");
cache.get("SELECT 1");
cache.get("SELECT 1");
let stats = cache.stats();
assert_eq!(stats.hits, 3);
}
#[test]
fn test_cache_remove() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
assert!(cache.remove("SELECT 1"));
assert!(cache.get("SELECT 1").is_none());
}
#[test]
fn test_cache_remove_nonexistent() {
let cache = PreparedStatementCache::new(10);
assert!(!cache.remove("SELECT 1"));
}
#[test]
fn test_cache_clear() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.put("SELECT 2", "stmt_2");
cache.clear();
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn test_cache_lru_eviction() {
let cache = PreparedStatementCache::new(2);
cache.put("sql_1", "stmt_1");
cache.put("sql_2", "stmt_2");
cache.get("sql_1");
cache.put("sql_3", "stmt_3");
assert!(cache.get("sql_1").is_some(), "sql_1 应被保留(最近使用)");
assert!(cache.get("sql_2").is_none(), "sql_2 应被 LRU 驱逐");
assert!(cache.get("sql_3").is_some(), "sql_3 应存在");
let stats = cache.stats();
assert!(stats.evictions >= 1, "应至少有 1 次驱逐");
}
#[test]
fn test_cache_update_existing() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.put("SELECT 1", "stmt_2");
let result = cache.get("SELECT 1");
assert_eq!(result, Some("stmt_2".to_string()));
assert_eq!(cache.len(), 1);
}
#[test]
fn test_cache_stats_hit_rate() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.get("SELECT 1");
cache.get("SELECT 1");
cache.get("SELECT 1");
cache.get("SELECT 2");
cache.get("SELECT 3");
let stats = cache.stats();
assert_eq!(stats.hits, 3);
assert_eq!(stats.misses, 2);
assert_eq!(stats.total_accesses(), 5);
let expected_rate = 3.0 / 5.0;
assert!((stats.hit_rate() - expected_rate).abs() < 0.001);
}
#[test]
fn test_cache_stats_summary() {
let cache = PreparedStatementCache::new(100);
cache.put("SELECT 1", "stmt_1");
cache.get("SELECT 1");
let stats = cache.stats();
let summary = stats.summary();
assert!(summary.contains("hits=1"));
assert!(summary.contains("capacity=100"));
assert!(summary.contains("hit_rate="));
}
#[test]
fn test_cache_capacity_utilization() {
let cache = PreparedStatementCache::new(10);
cache.put("sql_1", "stmt_1");
cache.put("sql_2", "stmt_2");
let stats = cache.stats();
assert_eq!(stats.size, 2);
assert_eq!(stats.capacity, 10);
assert!((stats.capacity_utilization() - 0.2).abs() < 0.001);
}
#[test]
fn test_cache_reset_stats() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.get("SELECT 1");
cache.get("SELECT 2");
cache.reset_stats();
let stats = cache.stats();
assert_eq!(stats.hits, 0);
assert_eq!(stats.misses, 0);
assert_eq!(stats.evictions, 0);
assert_eq!(stats.size, 1);
}
#[test]
fn test_cache_default_capacity() {
let cache = PreparedStatementCache::default();
assert_eq!(cache.capacity(), 256);
}
#[test]
fn test_cache_min_capacity_1() {
let cache = PreparedStatementCache::new(0);
assert_eq!(cache.capacity(), 1, "容量为 0 时应自动设为 1");
}
#[test]
fn test_cache_debug_format() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
let debug_str = format!("{:?}", cache);
assert!(debug_str.contains("PreparedStatementCache"));
assert!(debug_str.contains("capacity: 10"));
assert!(debug_str.contains("size: 1"));
}
#[test]
fn test_cache_concurrent_access() {
use std::sync::Arc;
use std::thread;
let cache = Arc::new(PreparedStatementCache::new(100));
let mut handles = Vec::new();
for i in 0..4 {
let c = cache.clone();
handles.push(thread::spawn(move || {
for j in 0..10 {
let sql = format!("SELECT {}", i * 10 + j);
c.put(&sql, &format!("stmt_{}", i * 10 + j));
c.get(&sql);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(cache.len(), 40);
let stats = cache.stats();
assert!(stats.hits >= 40, "每个 put 后立即 get 应产生 40 次命中");
}
#[test]
fn test_cache_same_sql_different_whitespace_same_hash() {
let cache = PreparedStatementCache::new(10);
cache.put("SELECT 1", "stmt_1");
cache.put("SELECT 1", "stmt_2"); assert_eq!(cache.len(), 2, "不同空格的 SQL 应为不同条目");
}
}