use async_trait::async_trait;
use crossbeam_queue::ArrayQueue;
use futures::StreamExt;
#[cfg(feature = "circuit-breaker")]
use parking_lot::Mutex as PlMutex;
#[cfg(feature = "rate-limit")]
use parking_lot::RwLock as PlRwLock;
use std::future::Future;
use std::ops::{Deref, DerefMut};
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Notify;
#[cfg(feature = "circuit-breaker")]
use crate::circuit_breaker::{CircuitBreaker, CircuitState, DefaultCircuitBreaker};
use crate::error::PoolError;
#[cfg(feature = "rate-limit")]
use crate::rate_limiter::RateLimiter;
pub type QueryRows = Vec<std::collections::HashMap<String, crate::value::Value>>;
pub type QueryStreamItem =
Result<std::collections::HashMap<String, crate::value::Value>, crate::DbError>;
pub trait Connection: Send + Sync {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>>;
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>>;
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>>;
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>>;
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>>;
fn is_connected(&self) -> bool;
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>>;
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>>;
fn execute_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [crate::value::Value],
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
let _ = (sql, params);
Box::pin(async move {
Err(crate::DbError::Internal(
"execute_with_params not implemented for this adapter".to_string(),
))
})
}
fn query_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [crate::value::Value],
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>> {
let _ = (sql, params);
Box::pin(async move {
Err(crate::DbError::Internal(
"query_with_params not implemented for this adapter".to_string(),
))
})
}
fn query_values<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<crate::value::QueryValues, crate::DbError>> + Send + 'a>>
{
let _ = sql;
Box::pin(async move {
Err(crate::DbError::Internal(
"query_values not implemented for this adapter".to_string(),
))
})
}
fn query_values_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [crate::value::Value],
) -> Pin<Box<dyn Future<Output = Result<crate::value::QueryValues, crate::DbError>> + Send + 'a>>
{
let _ = (sql, params);
Box::pin(async move {
Err(crate::DbError::Internal(
"query_values_with_params not implemented for this adapter".to_string(),
))
})
}
fn query_stream<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn futures::Stream<Item = QueryStreamItem> + Send + 'a>> {
let sql_owned = sql.to_string();
let stream = futures::stream::once(async move { self.query(&sql_owned).await })
.map(|result| {
let items: Vec<QueryStreamItem> = match result {
Ok(rows) => rows.into_iter().map(Ok).collect(),
Err(e) => vec![Err(e)],
};
futures::stream::iter(items)
})
.flatten();
Box::pin(stream)
}
fn query_stream_cursor<'a>(
&'a mut self,
sql: &'a str,
_batch_size: usize,
) -> Pin<Box<dyn futures::Stream<Item = QueryStreamItem> + Send + 'a>> {
self.query_stream(sql)
}
fn execute_batch<'a>(
&'a mut self,
sqls: &'a [String],
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move {
let mut total = 0u64;
for sql in sqls {
total += self.execute(sql).await?;
}
Ok(total)
})
}
fn execute_batch_params<'a>(
&'a mut self,
sql: &'a str,
params_batch: &'a [Vec<crate::value::Value>],
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move {
let mut total = 0u64;
for params in params_batch {
total += self.execute_with_params(sql, params).await?;
}
Ok(total)
})
}
}
pub struct PooledConnection {
conn: Box<dyn Connection>,
created_at: Instant,
last_used_at: Instant,
pool: Option<Pool>,
}
impl PooledConnection {
fn new(conn: Box<dyn Connection>, pool: Pool) -> Self {
let now = Instant::now();
Self {
conn,
created_at: now,
last_used_at: now,
pool: Some(pool),
}
}
fn is_expired(&self, max_lifetime: Duration) -> bool {
self.created_at.elapsed() >= max_lifetime
}
fn is_idle_too_long(&self, idle_timeout: Duration) -> bool {
self.last_used_at.elapsed() >= idle_timeout
}
pub fn created_at(&self) -> Instant {
self.created_at
}
pub fn into_inner(mut self) -> Box<dyn Connection> {
self.pool = None; std::mem::replace(&mut self.conn, Box::new(ClosedConnection))
}
}
impl Drop for PooledConnection {
fn drop(&mut self) {
if let Some(pool) = self.pool.take() {
let conn = std::mem::replace(&mut self.conn, Box::new(ClosedConnection));
let pooled = PooledConnection {
conn,
created_at: self.created_at,
last_used_at: self.last_used_at,
pool: None,
};
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
pool.release(pooled).await;
});
} else {
drop(pooled);
pool.total_count.fetch_sub(1, Ordering::SeqCst);
}
}
}
}
struct ClosedConnection;
impl Connection for ClosedConnection {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async {
Err(crate::DbError::ConnectionError(
"connection already returned to pool".to_string(),
))
})
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>> {
Box::pin(async {
Err(crate::DbError::ConnectionError(
"connection already returned to pool".to_string(),
))
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async {
Err(crate::DbError::ConnectionError(
"connection already returned to pool".to_string(),
))
})
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
fn is_connected(&self) -> bool {
false
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async { false })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
}
impl Deref for PooledConnection {
type Target = dyn Connection;
fn deref(&self) -> &Self::Target {
self.conn.as_ref()
}
}
impl DerefMut for PooledConnection {
fn deref_mut(&mut self) -> &mut Self::Target {
self.conn.as_mut()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TlsVersion {
#[default]
Tls12,
Tls13,
}
#[derive(Debug, Clone, Default)]
pub struct TlsConfig {
pub enabled: bool,
pub ca_cert_path: Option<String>,
pub client_cert_path: Option<String>,
pub client_key_path: Option<String>,
pub min_version: TlsVersion,
}
#[derive(Debug, Clone)]
pub enum PoolEvent {
ConnectionCreated,
ConnectionClosed,
ConnectionAcquired,
ConnectionReleased,
AcquireTimeout,
}
pub type PoolEventCallback = Arc<dyn Fn(PoolEvent) + Send + Sync>;
pub struct PoolConfig {
pub max_size: u32,
pub min_idle: u32,
pub acquire_timeout: Duration,
pub idle_timeout: Duration,
pub max_lifetime: Duration,
pub connection_timeout: Duration,
pub tls: Option<TlsConfig>,
pub query_timeout: Option<Duration>,
pub max_rows: Option<usize>,
pub memory_limit: Option<usize>,
pub on_event: Option<PoolEventCallback>,
pub test_before_acquire: bool,
pub prewarm: bool,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
max_size: 100,
min_idle: 0,
acquire_timeout: Duration::from_secs(30),
idle_timeout: Duration::from_secs(600),
max_lifetime: Duration::from_secs(1800),
connection_timeout: Duration::from_secs(10),
tls: None,
query_timeout: Some(Duration::from_secs(30)),
max_rows: None,
memory_limit: None,
on_event: None,
test_before_acquire: false,
prewarm: false,
}
}
}
impl Clone for PoolConfig {
fn clone(&self) -> Self {
Self {
max_size: self.max_size,
min_idle: self.min_idle,
acquire_timeout: self.acquire_timeout,
idle_timeout: self.idle_timeout,
max_lifetime: self.max_lifetime,
connection_timeout: self.connection_timeout,
tls: self.tls.clone(),
query_timeout: self.query_timeout,
max_rows: self.max_rows,
memory_limit: self.memory_limit,
on_event: self.on_event.clone(),
test_before_acquire: self.test_before_acquire,
prewarm: self.prewarm,
}
}
}
impl PoolConfig {
pub fn validate(&self) -> Result<(), PoolError> {
if self.max_size == 0 {
return Err(PoolError::InvalidConfig("max_size cannot be 0".to_string()));
}
if self.min_idle > self.max_size {
return Err(PoolError::InvalidConfig(
"min_idle cannot exceed max_size".to_string(),
));
}
const MAX_DURATION_SECS: u64 = u32::MAX as u64; for (name, dur) in [
("acquire_timeout", self.acquire_timeout),
("idle_timeout", self.idle_timeout),
("max_lifetime", self.max_lifetime),
("connection_timeout", self.connection_timeout),
] {
if dur.as_secs() > MAX_DURATION_SECS {
return Err(PoolError::InvalidConfig(format!(
"{name} ({:?}) exceeds maximum allowed duration ({} seconds)",
dur, MAX_DURATION_SECS
)));
}
}
Ok(())
}
#[must_use]
pub fn with_prewarm(mut self, prewarm: bool) -> Self {
self.prewarm = prewarm;
self
}
}
pub struct PoolStatus {
pub idle: u32,
pub active: u32,
pub max: u32,
pub min: u32,
pub waiters: u32,
}
impl std::fmt::Debug for PoolStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PoolStatus")
.field("idle", &self.idle)
.field("active", &self.active)
.field("max", &self.max)
.field("min", &self.min)
.field("waiters", &self.waiters)
.finish()
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct PoolMetrics {
pub acquire_count: u64,
pub acquire_failed_count: u64,
pub acquire_wait_time: Duration,
pub release_count: u64,
pub connection_created_count: u64,
pub connection_closed_count: u64,
}
impl PoolMetrics {
#[must_use]
pub fn average_acquire_wait_time(&self) -> Duration {
if self.acquire_count == 0 {
Duration::ZERO
} else {
self.acquire_wait_time / self.acquire_count as u32
}
}
}
pub struct PoolConfigBuilder {
config: PoolConfig,
}
impl PoolConfigBuilder {
pub fn new() -> Self {
Self {
config: PoolConfig::default(),
}
}
pub fn max_size(mut self, size: u32) -> Self {
self.config.max_size = size;
self
}
pub fn min_idle(mut self, count: u32) -> Self {
self.config.min_idle = count;
self
}
pub fn acquire_timeout(mut self, timeout_secs: u64) -> Self {
self.config.acquire_timeout = Duration::from_secs(timeout_secs);
self
}
pub fn idle_timeout(mut self, timeout_secs: u64) -> Self {
self.config.idle_timeout = Duration::from_secs(timeout_secs);
self
}
pub fn max_lifetime(mut self, lifetime_secs: u64) -> Self {
self.config.max_lifetime = Duration::from_secs(lifetime_secs);
self
}
pub fn tls(mut self, tls: TlsConfig) -> Self {
self.config.tls = Some(tls);
self
}
pub fn query_timeout(mut self, timeout: Duration) -> Self {
self.config.query_timeout = Some(timeout);
self
}
pub fn max_rows(mut self, max_rows: usize) -> Self {
self.config.max_rows = Some(max_rows);
self
}
pub fn memory_limit(mut self, memory_limit: usize) -> Self {
self.config.memory_limit = Some(memory_limit);
self
}
pub fn on_event(mut self, callback: PoolEventCallback) -> Self {
self.config.on_event = Some(callback);
self
}
pub fn test_before_acquire(mut self, enabled: bool) -> Self {
self.config.test_before_acquire = enabled;
self
}
pub fn prewarm(mut self, enabled: bool) -> Self {
self.config.prewarm = enabled;
self
}
pub fn build(self) -> Result<PoolConfig, PoolError> {
self.config.validate()?;
Ok(self.config)
}
}
impl Default for PoolConfigBuilder {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
pub trait ConnectionFactory: Send + Sync {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError>;
}
pub struct Pool {
config: PoolConfig,
factory: Arc<dyn ConnectionFactory>,
idle: Arc<ArrayQueue<PooledConnection>>,
total_count: Arc<AtomicU32>,
closed: Arc<AtomicBool>,
notify: Arc<Notify>,
waiters_count: Arc<AtomicU32>,
dynamic_max_size: Arc<AtomicU32>,
#[cfg(feature = "circuit-breaker")]
circuit_breaker: Arc<PlMutex<DefaultCircuitBreaker>>,
#[cfg(feature = "rate-limit")]
rate_limiter: Arc<PlRwLock<Option<Arc<dyn RateLimiter>>>>,
#[cfg(feature = "rate-limit")]
rate_limit_key: String,
acquire_count: Arc<AtomicU64>,
acquire_failed_count: Arc<AtomicU64>,
acquire_wait_time_ns: Arc<AtomicU64>,
release_count: Arc<AtomicU64>,
connection_created_count: Arc<AtomicU64>,
connection_closed_count: Arc<AtomicU64>,
}
impl Clone for Pool {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
factory: self.factory.clone(),
idle: self.idle.clone(),
total_count: self.total_count.clone(),
closed: self.closed.clone(),
notify: Arc::clone(&self.notify),
waiters_count: self.waiters_count.clone(),
dynamic_max_size: self.dynamic_max_size.clone(),
#[cfg(feature = "circuit-breaker")]
circuit_breaker: Arc::clone(&self.circuit_breaker),
#[cfg(feature = "rate-limit")]
rate_limiter: Arc::clone(&self.rate_limiter),
#[cfg(feature = "rate-limit")]
rate_limit_key: self.rate_limit_key.clone(),
acquire_count: self.acquire_count.clone(),
acquire_failed_count: self.acquire_failed_count.clone(),
acquire_wait_time_ns: self.acquire_wait_time_ns.clone(),
release_count: self.release_count.clone(),
connection_created_count: self.connection_created_count.clone(),
connection_closed_count: self.connection_closed_count.clone(),
}
}
}
impl Pool {
pub fn new(config: PoolConfig, factory: Arc<dyn ConnectionFactory>) -> Result<Self, PoolError> {
config.validate()?;
let max_size = config.max_size as usize;
let dynamic_max = config.max_size;
Ok(Self {
config,
factory,
idle: Arc::new(ArrayQueue::new(max_size)),
total_count: Arc::new(AtomicU32::new(0)),
closed: Arc::new(AtomicBool::new(false)),
notify: Arc::new(Notify::new()),
waiters_count: Arc::new(AtomicU32::new(0)),
dynamic_max_size: Arc::new(AtomicU32::new(dynamic_max)),
#[cfg(feature = "circuit-breaker")]
circuit_breaker: Arc::new(PlMutex::new(DefaultCircuitBreaker::new(
5,
std::time::Duration::from_secs(30),
))),
#[cfg(feature = "rate-limit")]
rate_limiter: Arc::new(PlRwLock::new(None)),
#[cfg(feature = "rate-limit")]
rate_limit_key: "pool".to_string(),
acquire_count: Arc::new(AtomicU64::new(0)),
acquire_failed_count: Arc::new(AtomicU64::new(0)),
acquire_wait_time_ns: Arc::new(AtomicU64::new(0)),
release_count: Arc::new(AtomicU64::new(0)),
connection_created_count: Arc::new(AtomicU64::new(0)),
connection_closed_count: Arc::new(AtomicU64::new(0)),
})
}
pub async fn new_async(
config: PoolConfig,
factory: Arc<dyn ConnectionFactory>,
) -> Result<Self, PoolError> {
let pool = Self::new(config, factory)?;
if pool.config.prewarm {
pool.prewarm().await;
}
Ok(pool)
}
pub async fn prewarm(&self) {
if !self.config.prewarm {
return;
}
let min_idle = self.config.min_idle as usize;
let mut warmed = 0;
for i in 0..min_idle {
if self.closed.load(Ordering::Acquire) {
break;
}
let current_max = self.dynamic_max_size.load(Ordering::Acquire);
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break;
}
let created = loop {
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break None;
}
match self.total_count.compare_exchange(
current,
current + 1,
Ordering::SeqCst,
Ordering::Acquire,
) {
Ok(_) => break Some(()),
Err(_) => continue,
}
};
if created.is_some() {
match tokio::time::timeout(self.config.connection_timeout, self.factory.create())
.await
{
Ok(Ok(conn)) => {
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_success();
}
self.emit_event(PoolEvent::ConnectionCreated);
let pooled = PooledConnection::new(conn, self.clone());
if self.idle.push(pooled).is_err() {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
tracing::warn!(
target: "sz_orm::pool::prewarm",
"prewarm connection {} failed: idle queue full",
i
);
} else {
warmed += 1;
self.notify.notify_one();
}
}
Ok(Err(e)) => {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
tracing::warn!(
target: "sz_orm::pool::prewarm",
"prewarm connection {} failed: {}",
i,
e
);
}
Err(_) => {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
tracing::warn!(
target: "sz_orm::pool::prewarm",
"prewarm connection {} timeout",
i
);
}
}
}
}
if warmed > 0 {
tracing::info!(
target: "sz_orm::pool::prewarm",
"pool prewarm completed: {}/{} connections established",
warmed,
min_idle
);
}
}
#[cfg(feature = "auto-prewarm")]
pub async fn progressive_prewarm(
&self,
batch_size: u32,
interval: std::time::Duration,
total_timeout: std::time::Duration,
progress: &crate::prewarm::PrewarmProgress,
) {
use std::time::Instant;
let min_idle = self.config.min_idle;
if min_idle == 0 || !self.config.prewarm {
progress.mark_completed();
return;
}
let start = Instant::now();
let batch = batch_size.max(1);
let mut warmed_total: u32 = 0;
while warmed_total < min_idle {
if start.elapsed() >= total_timeout {
tracing::warn!(
target: "sz_orm::pool::prewarm",
"progressive prewarm timeout: {}/{} connections established",
warmed_total,
min_idle
);
break;
}
if self.closed.load(Ordering::Acquire) {
break;
}
let remaining = min_idle - warmed_total;
let this_batch = batch.min(remaining);
for _ in 0..this_batch {
let current_max = self.dynamic_max_size.load(Ordering::Acquire);
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break;
}
let created = loop {
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break None;
}
match self.total_count.compare_exchange(
current,
current + 1,
Ordering::SeqCst,
Ordering::Acquire,
) {
Ok(_) => break Some(()),
Err(_) => continue,
}
};
if created.is_some() {
match tokio::time::timeout(
self.config.connection_timeout,
self.factory.create(),
)
.await
{
Ok(Ok(conn)) => {
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_success();
}
self.emit_event(PoolEvent::ConnectionCreated);
let pooled = PooledConnection::new(conn, self.clone());
if self.idle.push(pooled).is_err() {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
progress.record_failure();
} else {
progress.record_success();
warmed_total += 1;
self.notify.notify_one();
}
}
Ok(Err(_)) => {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
progress.record_failure();
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
}
Err(_) => {
let _ = self.total_count.fetch_sub(1, Ordering::SeqCst);
progress.record_failure();
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
}
}
}
}
if warmed_total < min_idle && interval > std::time::Duration::ZERO {
tokio::time::sleep(interval).await;
}
}
progress.set_elapsed(start.elapsed());
progress.mark_completed();
tracing::info!(
target: "sz_orm::pool::prewarm",
"progressive prewarm completed: {} warmed, {} failed, elapsed {:?}",
progress.snapshot().warmed,
progress.snapshot().failed,
start.elapsed()
);
}
pub fn config(&self) -> &PoolConfig {
&self.config
}
#[cfg(feature = "circuit-breaker")]
pub fn configure_circuit_breaker(
&self,
failure_threshold: usize,
reset_timeout: std::time::Duration,
) {
let new_cb = DefaultCircuitBreaker::new(failure_threshold, reset_timeout);
let mut guard = self.circuit_breaker.lock();
*guard = new_cb;
}
#[cfg(feature = "circuit-breaker")]
pub fn reset_circuit_breaker(&self) -> bool {
let mut guard = self.circuit_breaker.lock();
guard.reset()
}
#[cfg(feature = "circuit-breaker")]
pub fn circuit_state(&self) -> CircuitState {
let guard = self.circuit_breaker.lock();
guard.state()
}
#[cfg(feature = "rate-limit")]
pub fn set_rate_limiter(&self, limiter: Option<Arc<dyn RateLimiter>>) {
let mut guard = self.rate_limiter.write();
*guard = limiter;
}
#[cfg(feature = "rate-limit")]
pub fn with_rate_limit_key(mut self, key: impl Into<String>) -> Self {
self.rate_limit_key = key.into();
self
}
fn emit_event(&self, event: PoolEvent) {
if matches!(event, PoolEvent::ConnectionCreated) {
self.connection_created_count
.fetch_add(1, Ordering::Relaxed);
}
if let Some(ref callback) = self.config.on_event {
callback(event);
}
}
async fn close_connection(&self, pooled: PooledConnection) {
let mut pooled = pooled;
let _ = pooled.conn.close().await;
self.connection_closed_count.fetch_add(1, Ordering::Relaxed);
}
#[tracing::instrument(skip(self), fields(max_size = self.config.max_size, acquire_timeout = ?self.config.acquire_timeout))]
pub async fn acquire(&self) -> Result<PooledConnection, PoolError> {
if self.closed.load(Ordering::Acquire) {
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::Closed);
}
#[cfg(feature = "circuit-breaker")]
{
let mut guard = self.circuit_breaker.lock();
if !guard.can_execute() {
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::CircuitOpen);
}
}
#[cfg(feature = "rate-limit")]
{
let guard = self.rate_limiter.read();
if let Some(ref limiter) = *guard {
match limiter.try_acquire(&self.rate_limit_key) {
Ok(result) if !result.allowed => {
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::RateLimited {
remaining: result.remaining,
reset_at: result.reset_at,
});
}
Ok(_) => {} Err(_) => {
}
}
}
}
let deadline = Instant::now() + self.config.acquire_timeout;
let mut backoff = Duration::from_millis(1);
const MAX_BACKOFF: Duration = Duration::from_millis(100);
loop {
let mut to_close: Vec<PooledConnection> = Vec::new();
let acquired: Option<PooledConnection> = {
let mut found: Option<PooledConnection> = None;
while let Some(pooled) = self.idle.pop() {
if pooled.is_expired(self.config.max_lifetime) {
to_close.push(pooled);
continue;
}
if pooled.is_idle_too_long(self.config.idle_timeout) {
to_close.push(pooled);
continue;
}
if !pooled.conn.is_connected() {
to_close.push(pooled);
continue;
}
found = Some(pooled);
break;
}
found
};
for pooled in to_close {
self.close_connection(pooled).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
}
if let Some(mut pooled) = acquired {
if self.config.test_before_acquire {
let ping_timeout = self.config.connection_timeout / 2;
let alive = match tokio::time::timeout(ping_timeout, pooled.conn.ping()).await {
Ok(true) => true,
Ok(false) => false,
Err(_) => false, };
if !alive {
self.close_connection(pooled).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
continue;
}
}
pooled.pool = Some(self.clone());
self.acquire_count.fetch_add(1, Ordering::Relaxed);
return Ok(pooled);
}
let current_max = self.dynamic_max_size.load(Ordering::Acquire);
let created = loop {
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break None; }
match self.total_count.compare_exchange(
current,
current + 1,
Ordering::SeqCst,
Ordering::Acquire,
) {
Ok(_) => break Some(()), Err(_) => continue, }
};
if created.is_some() {
match tokio::time::timeout(self.config.connection_timeout, self.factory.create())
.await
{
Ok(Ok(conn)) => {
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_success();
}
self.emit_event(PoolEvent::ConnectionCreated);
self.emit_event(PoolEvent::ConnectionAcquired);
self.acquire_count.fetch_add(1, Ordering::Relaxed);
return Ok(PooledConnection::new(conn, self.clone()));
}
Ok(Err(e)) => {
self.total_count.fetch_sub(1, Ordering::SeqCst);
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::ConnectionFailed(e.to_string()));
}
Err(_) => {
self.total_count.fetch_sub(1, Ordering::SeqCst);
#[cfg(feature = "circuit-breaker")]
{
self.circuit_breaker.lock().record_failure();
}
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::Timeout);
}
}
}
let now = Instant::now();
if now >= deadline {
self.emit_event(PoolEvent::AcquireTimeout);
self.acquire_failed_count.fetch_add(1, Ordering::Relaxed);
return Err(PoolError::Timeout);
}
self.waiters_count.fetch_add(1, Ordering::SeqCst);
let wait = std::cmp::min(backoff, deadline - now);
match tokio::time::timeout(wait, self.notify.notified()).await {
Ok(()) => {
backoff = Duration::from_millis(1);
}
Err(_) => {
backoff = std::cmp::min(backoff * 2, MAX_BACKOFF);
}
}
self.waiters_count.fetch_sub(1, Ordering::SeqCst);
self.acquire_wait_time_ns
.fetch_add(wait.as_nanos() as u64, Ordering::Relaxed);
}
}
#[tracing::instrument(skip(self, pooled))]
pub async fn release(&self, mut pooled: PooledConnection) {
pooled.pool = None;
self.release_count.fetch_add(1, Ordering::Relaxed);
if self.closed.load(Ordering::Acquire) {
self.close_connection(pooled).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
self.emit_event(PoolEvent::ConnectionClosed);
return;
}
if !pooled.conn.is_connected() {
self.close_connection(pooled).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
self.emit_event(PoolEvent::ConnectionClosed);
return;
}
pooled.last_used_at = Instant::now();
if let Err(rejected) = self.idle.push(pooled) {
self.close_connection(rejected).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
self.emit_event(PoolEvent::ConnectionClosed);
} else {
self.emit_event(PoolEvent::ConnectionReleased);
}
self.notify.notify_one();
}
pub async fn status(&self) -> PoolStatus {
let idle_count = self.idle.len() as u32;
let active = self.total_count.load(Ordering::Acquire);
let waiters = self.waiters_count.load(Ordering::Acquire);
PoolStatus {
idle: idle_count,
active,
max: self.dynamic_max_size.load(Ordering::Acquire),
min: self.config.min_idle,
waiters,
}
}
pub fn pool_metrics(&self) -> PoolMetrics {
PoolMetrics {
acquire_count: self.acquire_count.load(Ordering::Acquire),
acquire_failed_count: self.acquire_failed_count.load(Ordering::Acquire),
acquire_wait_time: Duration::from_nanos(
self.acquire_wait_time_ns.load(Ordering::Acquire),
),
release_count: self.release_count.load(Ordering::Acquire),
connection_created_count: self.connection_created_count.load(Ordering::Acquire),
connection_closed_count: self.connection_closed_count.load(Ordering::Acquire),
}
}
#[tracing::instrument(skip(self))]
pub async fn reap_idle(&self) {
let mut all: Vec<PooledConnection> = Vec::new();
while let Some(pooled) = self.idle.pop() {
all.push(pooled);
}
let mut to_close = Vec::new();
for pooled in all {
if pooled.is_idle_too_long(self.config.idle_timeout)
|| pooled.is_expired(self.config.max_lifetime)
{
to_close.push(pooled);
} else {
if let Err(rejected) = self.idle.push(pooled) {
self.close_connection(rejected).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
}
}
}
for pooled in to_close {
self.close_connection(pooled).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
}
}
pub async fn close_all(&self) {
self.closed.store(true, Ordering::Release);
let mut to_close: Vec<PooledConnection> = Vec::new();
while let Some(pooled) = self.idle.pop() {
to_close.push(pooled);
}
let closed_count: u32 = to_close.len() as u32;
for pooled in to_close {
self.close_connection(pooled).await;
}
self.total_count.fetch_sub(closed_count, Ordering::SeqCst);
}
pub async fn health_check(&self) -> u32 {
let mut to_check: Vec<PooledConnection> = Vec::new();
while let Some(pooled) = self.idle.pop() {
to_check.push(pooled);
}
let mut removed: u32 = 0;
let mut alive: Vec<PooledConnection> = Vec::with_capacity(to_check.len());
for mut pooled in to_check.drain(..) {
if !pooled.conn.is_connected() {
self.close_connection(pooled).await;
removed += 1;
continue;
}
let ping_timeout = self.config.connection_timeout / 2;
match tokio::time::timeout(ping_timeout, pooled.conn.ping()).await {
Ok(true) => alive.push(pooled),
Ok(false) => {
self.close_connection(pooled).await;
removed += 1;
}
Err(_) => {
self.close_connection(pooled).await;
removed += 1;
}
}
}
let alive_count: u32 = alive.len() as u32;
for pooled in alive {
if let Err(rejected) = self.idle.push(pooled) {
self.close_connection(rejected).await;
removed += 1;
}
}
if removed > 0 {
self.total_count.fetch_sub(removed, Ordering::SeqCst);
}
if alive_count > 0 {
self.notify.notify_one();
}
removed
}
pub async fn shutdown(&self) {
self.closed.store(true, Ordering::SeqCst);
self.notify.notify_waiters();
self.close_all().await;
let deadline = Instant::now() + Duration::from_secs(30);
while self.total_count.load(Ordering::SeqCst) > 0 {
if Instant::now() >= deadline {
break;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
pub fn resize(&self, new_max: usize) {
self.set_max_size(new_max as u32);
}
pub fn set_max_size(&self, new_max: u32) {
self.dynamic_max_size.store(new_max, Ordering::SeqCst);
}
pub fn max_size(&self) -> u32 {
self.dynamic_max_size.load(Ordering::Acquire)
}
pub async fn warmup(&self, min_idle: usize) -> Result<(), PoolError> {
for _ in 0..min_idle {
let current_max = self.dynamic_max_size.load(Ordering::Acquire);
let current = self.total_count.load(Ordering::Acquire);
if current >= current_max {
break;
}
match self.total_count.compare_exchange(
current,
current + 1,
Ordering::SeqCst,
Ordering::Acquire,
) {
Ok(_) => {}
Err(_) => continue, }
match self.factory.create().await {
Ok(conn) => {
let now = Instant::now();
let pooled = PooledConnection {
conn,
created_at: now,
last_used_at: now,
pool: None,
};
if let Err(rejected) = self.idle.push(pooled) {
self.close_connection(rejected).await;
self.total_count.fetch_sub(1, Ordering::SeqCst);
}
self.emit_event(PoolEvent::ConnectionCreated);
}
Err(_) => {
self.total_count.fetch_sub(1, Ordering::SeqCst);
break;
}
}
}
Ok(())
}
pub async fn query_with_timeout(&self, sql: &str) -> Result<QueryRows, crate::DbError> {
let timeout = self.config.query_timeout.unwrap_or(Duration::from_secs(30));
let mut conn = self.acquire().await.map_err(crate::DbError::PoolError)?;
tokio::time::timeout(timeout, conn.query(sql))
.await
.map_err(|_| crate::DbError::QueryError(format!("Query timeout after {:?}", timeout)))?
}
}
#[cfg(test)]
mod tests {
use super::*;
struct MockConnection {
connected: bool,
}
impl MockConnection {
fn new() -> Self {
Self { connected: true }
}
}
impl Connection for MockConnection {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(1) })
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, crate::value::Value>>,
crate::DbError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async move { Ok(vec![]) })
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
self.connected
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move {
self.connected = false;
Ok(())
})
}
}
struct MockConnectionFactory;
#[async_trait]
impl ConnectionFactory for MockConnectionFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
Ok(Box::new(MockConnection::new()))
}
}
#[tokio::test]
async fn test_pool_config_builder() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(50).min_idle(10).build()?;
assert_eq!(config.max_size, 50);
assert_eq!(config.min_idle, 10);
Ok(())
}
#[test]
fn test_pool_status_display() {
let status = PoolStatus {
idle: 5,
active: 10,
max: 100,
min: 5,
waiters: 0,
};
let display = format!("{:?}", status);
assert!(display.contains("idle"));
assert!(display.contains("active"));
}
#[test]
fn test_default_pool_config() {
let config = PoolConfig::default();
assert_eq!(config.max_size, 100);
assert_eq!(config.min_idle, 0);
assert_eq!(config.acquire_timeout.as_secs(), 30);
assert_eq!(config.idle_timeout.as_secs(), 600);
assert_eq!(config.max_lifetime.as_secs(), 1800);
}
#[tokio::test]
async fn test_pool_config_clone() {
let config = PoolConfig::default();
let cloned = config.clone();
assert_eq!(cloned.max_size, config.max_size);
assert_eq!(cloned.min_idle, config.min_idle);
}
#[test]
fn test_pool_config_builder_default() -> Result<(), Box<dyn std::error::Error>> {
let builder = PoolConfigBuilder::new();
let config = builder.build()?;
assert_eq!(config.max_size, 100);
Ok(())
}
#[test]
fn test_pool_config_validate() {
let result = PoolConfigBuilder::new().max_size(0).build();
assert!(result.is_err());
let result = PoolConfigBuilder::new().max_size(10).min_idle(20).build();
assert!(result.is_err());
}
#[test]
fn test_pool_config_validate_duration_upper_bound() {
use std::time::Duration;
let config = PoolConfig {
max_size: 10,
min_idle: 1,
acquire_timeout: Duration::from_secs(u64::MAX),
idle_timeout: Duration::from_secs(1),
max_lifetime: Duration::from_secs(1),
connection_timeout: Duration::from_secs(5),
tls: None,
query_timeout: None,
max_rows: None,
memory_limit: None,
on_event: None,
test_before_acquire: false,
prewarm: false,
};
assert!(config.validate().is_err());
let config = PoolConfig {
max_size: 10,
min_idle: 1,
acquire_timeout: Duration::from_secs(u32::MAX as u64),
idle_timeout: Duration::from_secs(1),
max_lifetime: Duration::from_secs(1),
connection_timeout: Duration::from_secs(5),
tls: None,
query_timeout: None,
max_rows: None,
memory_limit: None,
on_event: None,
test_before_acquire: false,
prewarm: false,
};
assert!(config.validate().is_ok());
let config = PoolConfig {
max_size: 10,
min_idle: 1,
acquire_timeout: Duration::from_secs(u32::MAX as u64 + 1),
idle_timeout: Duration::from_secs(1),
max_lifetime: Duration::from_secs(1),
connection_timeout: Duration::from_secs(5),
tls: None,
query_timeout: None,
max_rows: None,
memory_limit: None,
on_event: None,
test_before_acquire: false,
prewarm: false,
};
assert!(config.validate().is_err());
}
#[test]
fn test_pool_config_test_before_acquire_default() {
let config = PoolConfig::default();
assert!(!config.test_before_acquire);
}
#[test]
fn test_pool_config_builder_test_before_acquire() {
let config = PoolConfigBuilder::new()
.test_before_acquire(true)
.build()
.unwrap();
assert!(config.test_before_acquire);
}
#[tokio::test]
async fn test_pool_acquire_and_release() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(5).min_idle(1).build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let conn = pool.acquire().await?;
let status = pool.status().await;
assert_eq!(status.active, 1);
assert_eq!(status.idle, 0);
pool.release(conn).await;
let status = pool.status().await;
assert_eq!(status.idle, 1);
let _conn2 = pool.acquire().await?;
let status = pool.status().await;
assert_eq!(status.idle, 0);
Ok(())
}
#[tokio::test]
async fn test_pool_status() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(10).min_idle(2).build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let status = pool.status().await;
assert_eq!(status.max, 10);
assert_eq!(status.min, 2);
assert_eq!(status.active, 0);
Ok(())
}
#[tokio::test]
async fn test_pool_close_all() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(5).build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let conn1 = pool.acquire().await?;
let conn2 = pool.acquire().await?;
pool.release(conn1).await;
pool.release(conn2).await;
pool.close_all().await;
let status = pool.status().await;
assert_eq!(status.idle, 0);
assert_eq!(status.active, 0);
Ok(())
}
#[tokio::test]
async fn test_pool_reap_idle() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new()
.max_size(5)
.idle_timeout(0) .build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let conn = pool.acquire().await?;
pool.release(conn).await;
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
pool.reap_idle().await;
let status = pool.status().await;
assert_eq!(status.idle, 0);
Ok(())
}
#[tokio::test]
async fn test_h7_acquire_timeout_default_30s() {
let config = PoolConfig::default();
assert_eq!(
config.acquire_timeout,
Duration::from_secs(30),
"H-7: acquire_timeout 默认应为 30s"
);
}
#[tokio::test]
async fn test_h7_acquire_timeout_configurable() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new()
.max_size(1)
.acquire_timeout(5) .build()?;
assert_eq!(config.acquire_timeout, Duration::from_secs(5));
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let _conn1 = pool.acquire().await?;
let fast_config = PoolConfigBuilder::new()
.max_size(1)
.acquire_timeout(0) .build()?;
let fast_pool = Pool::new(fast_config, Arc::new(MockConnectionFactory))?;
let _fast_conn = fast_pool.acquire().await?; let result = fast_pool.acquire().await;
assert!(
matches!(result, Err(PoolError::Timeout)),
"H-7: 应返回 Timeout"
);
Ok(())
}
#[tokio::test]
async fn test_m7_health_check_removes_nothing_when_all_healthy(
) -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(5).build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let conn1 = pool.acquire().await?;
let conn2 = pool.acquire().await?;
let conn3 = pool.acquire().await?;
pool.release(conn1).await;
pool.release(conn2).await;
pool.release(conn3).await;
let removed = pool.health_check().await;
assert_eq!(removed, 0, "Healthy connections should not be removed");
let status = pool.status().await;
assert_eq!(status.idle, 3);
assert_eq!(status.active, 3);
Ok(())
}
#[tokio::test]
async fn test_m7_health_check_returns_zero_for_empty_pool(
) -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(5).build()?;
let factory = Arc::new(MockConnectionFactory);
let pool = Pool::new(config, factory)?;
let removed = pool.health_check().await;
assert_eq!(removed, 0);
Ok(())
}
struct CountingFactory {
count: AtomicU32,
}
impl CountingFactory {
fn new() -> Self {
Self {
count: AtomicU32::new(0),
}
}
fn created_count(&self) -> u32 {
self.count.load(Ordering::SeqCst)
}
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
#[tokio::test]
async fn test_production_bug_max_lifetime_never_expires(
) -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfig {
max_size: 5,
min_idle: 0,
acquire_timeout: Duration::from_secs(30),
idle_timeout: Duration::from_secs(600),
max_lifetime: Duration::from_millis(100), connection_timeout: Duration::from_secs(10),
tls: None,
query_timeout: None,
max_rows: None,
memory_limit: None,
on_event: None,
test_before_acquire: false,
prewarm: false,
};
let factory = Arc::new(CountingFactory::new());
let pool = Pool::new(config, factory.clone())?;
let conn = pool.acquire().await?;
assert_eq!(factory.created_count(), 1, "应创建 1 个连接");
pool.release(conn).await;
tokio::time::sleep(Duration::from_millis(150)).await;
let conn2 = pool.acquire().await?;
assert_eq!(
factory.created_count(),
2,
"超过 max_lifetime 后应创建新连接(旧连接应被回收)"
);
pool.release(conn2).await;
Ok(())
}
#[tokio::test]
async fn test_drop_auto_release_connection() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(2).build()?;
let factory = Arc::new(CountingFactory::new());
let pool = Pool::new(config, factory.clone())?;
{
let _conn = pool.acquire().await?;
assert_eq!(factory.created_count(), 1, "应创建 1 个连接");
let status = pool.status().await;
assert_eq!(status.active, 1, "active 应为 1");
assert_eq!(status.idle, 0, "idle 应为 0");
}
tokio::time::sleep(Duration::from_millis(50)).await;
let status = pool.status().await;
assert_eq!(status.idle, 1, "Drop 后连接应自动归还,idle 应为 1");
assert_eq!(status.active, 1, "total_count 应为 1");
assert_eq!(factory.created_count(), 1, "应复用归还的连接,不创建新连接");
Ok(())
}
#[tokio::test]
async fn test_drop_auto_release_then_reuse() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(1).build()?;
let factory = Arc::new(CountingFactory::new());
let pool = Pool::new(config, factory.clone())?;
{
let _conn = pool.acquire().await?;
}
tokio::time::sleep(Duration::from_millis(50)).await;
let conn = pool.acquire().await?;
assert_eq!(factory.created_count(), 1, "应复用归还的连接,不创建新连接");
pool.release(conn).await;
Ok(())
}
#[tokio::test]
async fn test_into_inner_does_not_return_to_pool() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(2).build()?;
let factory = Arc::new(CountingFactory::new());
let pool = Pool::new(config, factory.clone())?;
let conn = pool.acquire().await?;
assert_eq!(factory.created_count(), 1);
let _raw_conn = conn.into_inner();
tokio::time::sleep(Duration::from_millis(50)).await;
let status = pool.status().await;
assert_eq!(status.idle, 0, "into_inner 后连接不应归还");
assert_eq!(status.active, 1, "total_count 仍为 1(连接被外部持有)");
Ok(())
}
#[tokio::test]
async fn test_explicit_release_no_double_return() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(2).build()?;
let factory = Arc::new(CountingFactory::new());
let pool = Pool::new(config, factory.clone())?;
let conn = pool.acquire().await?;
pool.release(conn).await;
let status = pool.status().await;
assert_eq!(status.idle, 1, "release 后 idle 应为 1");
let conn = pool.acquire().await?;
pool.release(conn).await;
let status = pool.status().await;
assert_eq!(status.idle, 1, "再次 release 后 idle 仍应为 1(不重复)");
assert_eq!(status.active, 1, "total_count 应为 1");
Ok(())
}
struct CursorMockConn {
rows: QueryRows,
call_count: usize,
}
impl CursorMockConn {
fn new(rows: QueryRows) -> Self {
Self {
rows,
call_count: 0,
}
}
}
impl Connection for CursorMockConn {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(1) })
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>> {
Box::pin(async move {
self.call_count += 1;
Ok(self.rows.clone())
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
}
struct CursorOverrideMockConn {
rows: Vec<crate::value::Value>,
yielded: usize,
}
impl CursorOverrideMockConn {
fn new(rows: Vec<crate::value::Value>) -> Self {
Self { rows, yielded: 0 }
}
}
impl Connection for CursorOverrideMockConn {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(1) })
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>> {
Box::pin(async move {
Ok(self
.rows
.iter()
.map(|v| {
let mut m = std::collections::HashMap::new();
m.insert("v".to_string(), v.clone());
m
})
.collect())
})
}
fn query_stream<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn futures::Stream<Item = QueryStreamItem> + Send + 'a>> {
Box::pin(futures::stream::iter(
self.rows
.iter()
.enumerate()
.map(|(i, v)| {
self.yielded = i + 1;
let mut m = std::collections::HashMap::new();
m.insert("v".to_string(), v.clone());
Ok(m)
})
.collect::<Vec<_>>(),
))
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
}
#[tokio::test]
async fn test_query_stream_default_impl_yields_all_rows() {
use futures::StreamExt;
let rows: QueryRows = vec![
std::collections::HashMap::from([
("id".to_string(), crate::value::Value::I64(1)),
(
"name".to_string(),
crate::value::Value::String("alice".to_string()),
),
]),
std::collections::HashMap::from([
("id".to_string(), crate::value::Value::I64(2)),
(
"name".to_string(),
crate::value::Value::String("bob".to_string()),
),
]),
std::collections::HashMap::from([
("id".to_string(), crate::value::Value::I64(3)),
(
"name".to_string(),
crate::value::Value::String("carol".to_string()),
),
]),
];
let mut conn = CursorMockConn::new(rows);
let mut stream = conn.query_stream("SELECT id, name FROM users");
let mut received: Vec<QueryStreamItem> = Vec::new();
while let Some(item) = stream.next().await {
received.push(item);
}
assert_eq!(received.len(), 3, "应收到 3 行");
assert!(received.iter().all(|r| r.is_ok()), "所有项应为 Ok");
drop(stream);
assert_eq!(conn.call_count, 1, "默认实现应调用 query() 一次");
}
#[tokio::test]
async fn test_query_stream_default_empty_result() {
use futures::StreamExt;
let mut conn = CursorMockConn::new(Vec::new());
let mut stream = conn.query_stream("SELECT * FROM empty_table");
let mut count = 0;
while let Some(_item) = stream.next().await {
count += 1;
}
assert_eq!(count, 0, "空结果集应产生 0 项");
}
#[tokio::test]
async fn test_query_stream_default_error_propagation() {
use futures::StreamExt;
struct ErrorMockConn;
impl Connection for ErrorMockConn {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>>
{
Box::pin(async move { Ok(1) })
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryRows, crate::DbError>> + Send + 'a>>
{
Box::pin(async move { Err(crate::DbError::Internal("query failed".to_string())) })
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
}
let mut conn = ErrorMockConn;
let mut stream = conn.query_stream("SELECT * FROM bad_table");
let item = stream.next().await;
assert!(item.is_some(), "应产生一项");
assert!(item.unwrap().is_err(), "该项应为 Err");
}
#[tokio::test]
async fn test_query_stream_override_yields_rows_one_by_one() {
use futures::StreamExt;
let rows = vec![
crate::value::Value::I64(10),
crate::value::Value::I64(20),
crate::value::Value::I64(30),
crate::value::Value::I64(40),
crate::value::Value::I64(50),
];
let mut conn = CursorOverrideMockConn::new(rows);
let values: Vec<i64> = {
let mut stream = conn.query_stream("SELECT v FROM seq");
let mut vals: Vec<i64> = Vec::new();
while let Some(Ok(row)) = stream.next().await {
if let crate::value::Value::I64(v) = row.get("v").unwrap() {
vals.push(*v);
}
}
vals
};
assert_eq!(values, vec![10, 20, 30, 40, 50], "应按顺序收到全部 5 行");
assert_eq!(conn.yielded, 5, "应逐行 yield 5 次(真游标覆盖)");
}
#[tokio::test]
async fn test_query_stream_override_early_drop() {
use futures::StreamExt;
let rows = vec![
crate::value::Value::I64(1),
crate::value::Value::I64(2),
crate::value::Value::I64(3),
];
let mut conn = CursorOverrideMockConn::new(rows);
{
let mut stream = conn.query_stream("SELECT v FROM seq");
let first = stream.next().await;
assert!(first.is_some(), "第一项应存在");
drop(stream);
}
assert!(conn.is_connected(), "提前 drop 流后连接仍应可用");
}
#[tokio::test]
async fn test_pool_prewarm() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(5)
.prewarm(true)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new(config, factory)?;
let status_before = pool.status().await;
assert_eq!(status_before.idle, 0, "预热前 idle 应为 0");
pool.prewarm().await;
let status_after = pool.status().await;
assert!(
status_after.idle >= 5,
"预热后 idle 应 >= 5,实际: {}",
status_after.idle
);
assert_eq!(
create_count.load(Ordering::SeqCst),
5,
"工厂应被调用 5 次(min_idle)"
);
Ok(())
}
#[tokio::test]
async fn test_pool_prewarm_failure_non_blocking() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicBool;
struct FailingFactory {
failed: Arc<AtomicBool>,
}
#[async_trait]
impl ConnectionFactory for FailingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.failed.store(true, Ordering::SeqCst);
Err(crate::DbError::Internal(
"simulated connection failure".to_string(),
))
}
}
let failed = Arc::new(AtomicBool::new(false));
let mut config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(3)
.prewarm(true)
.build()?;
config.connection_timeout = std::time::Duration::from_secs(1);
let factory = Arc::new(FailingFactory {
failed: failed.clone(),
});
let pool = Pool::new(config, factory)?;
pool.prewarm().await;
assert!(failed.load(Ordering::SeqCst), "工厂应被调用且失败");
let status = pool.status().await;
assert_eq!(status.max, 10, "池配置应正常");
Ok(())
}
#[tokio::test]
async fn test_pool_prewarm_disabled() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(5)
.prewarm(false) .build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new(config, factory)?;
pool.prewarm().await;
assert_eq!(
create_count.load(Ordering::SeqCst),
0,
"prewarm=false 时工厂不应被调用"
);
let status = pool.status().await;
assert_eq!(status.idle, 0, "idle 应为 0");
Ok(())
}
#[tokio::test]
async fn test_pool_new_async_with_prewarm() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(5)
.prewarm(true)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new_async(config, factory).await?;
let status = pool.status().await;
assert!(
status.idle >= 5,
"new_async prewarm=true 后 idle 应 >= 5,实际: {}",
status.idle
);
assert_eq!(create_count.load(Ordering::SeqCst), 5, "工厂应被调用 5 次");
Ok(())
}
#[tokio::test]
async fn test_pool_new_async_without_prewarm() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(5)
.prewarm(false)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new_async(config, factory).await?;
let status = pool.status().await;
assert_eq!(status.idle, 0, "prewarm=false 时 idle 应为 0");
assert_eq!(create_count.load(Ordering::SeqCst), 0, "工厂不应被调用");
Ok(())
}
#[tokio::test]
async fn test_pool_new_async_failure_non_blocking() -> Result<(), Box<dyn std::error::Error>> {
struct FailingFactory;
#[async_trait]
impl ConnectionFactory for FailingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
Err(crate::DbError::Internal("simulated failure".to_string()))
}
}
let mut config = PoolConfigBuilder::new()
.max_size(10)
.min_idle(3)
.prewarm(true)
.build()?;
config.connection_timeout = std::time::Duration::from_secs(1);
let pool = Pool::new_async(config, Arc::new(FailingFactory)).await?;
let status = pool.status().await;
assert_eq!(status.max, 10, "池配置应正常");
Ok(())
}
#[cfg(feature = "auto-prewarm")]
#[tokio::test]
async fn test_pool_progressive_prewarm() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(20)
.min_idle(6)
.prewarm(true)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new(config, factory)?;
let progress = crate::prewarm::PrewarmProgress::new(6);
pool.progressive_prewarm(
2,
std::time::Duration::from_millis(5),
std::time::Duration::from_secs(10),
&progress,
)
.await;
let snap = progress.snapshot();
assert!(
snap.warmed >= 6,
"progressive_prewarm 后 warmed 应 >= 6,实际: {}",
snap.warmed
);
assert!(snap.is_completed, "应标记完成");
assert_eq!(create_count.load(Ordering::SeqCst), 6, "工厂应被调用 6 次");
let status = pool.status().await;
assert!(status.idle >= 6, "池中 idle 应 >= 6");
Ok(())
}
#[cfg(feature = "auto-prewarm")]
#[tokio::test]
async fn test_pool_progressive_prewarm_timeout_zero() -> Result<(), Box<dyn std::error::Error>>
{
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(20)
.min_idle(10)
.prewarm(true)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new(config, factory)?;
let progress = crate::prewarm::PrewarmProgress::new(10);
pool.progressive_prewarm(
2,
std::time::Duration::from_millis(5),
std::time::Duration::ZERO,
&progress,
)
.await;
let snap = progress.snapshot();
assert!(snap.is_completed, "应标记完成");
assert!(
snap.warmed <= 2,
"total_timeout=0 时最多建一批(batch_size=2),实际: {}",
snap.warmed
);
Ok(())
}
#[cfg(feature = "auto-prewarm")]
#[tokio::test]
async fn test_pool_progressive_prewarm_disabled() -> Result<(), Box<dyn std::error::Error>> {
use std::sync::atomic::AtomicU32;
let create_count = Arc::new(AtomicU32::new(0));
let create_count_clone = create_count.clone();
struct CountingFactory {
count: Arc<AtomicU32>,
}
#[async_trait]
impl ConnectionFactory for CountingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
self.count.fetch_add(1, Ordering::SeqCst);
Ok(Box::new(MockConnection::new()))
}
}
let config = PoolConfigBuilder::new()
.max_size(20)
.min_idle(10)
.prewarm(false)
.build()?;
let factory = Arc::new(CountingFactory {
count: create_count_clone,
});
let pool = Pool::new(config, factory)?;
let progress = crate::prewarm::PrewarmProgress::new(10);
pool.progressive_prewarm(
2,
std::time::Duration::from_millis(5),
std::time::Duration::from_secs(10),
&progress,
)
.await;
let snap = progress.snapshot();
assert!(snap.is_completed, "应标记完成");
assert_eq!(snap.warmed, 0, "prewarm=false 时不应建连");
assert_eq!(create_count.load(Ordering::SeqCst), 0, "工厂不应被调用");
Ok(())
}
#[cfg(feature = "auto-prewarm")]
#[tokio::test]
async fn test_pool_progressive_prewarm_failure_non_blocking(
) -> Result<(), Box<dyn std::error::Error>> {
struct FailingFactory;
#[async_trait]
impl ConnectionFactory for FailingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
Err(crate::DbError::Internal("simulated failure".to_string()))
}
}
let mut config = PoolConfigBuilder::new()
.max_size(20)
.min_idle(5)
.prewarm(true)
.build()?;
config.connection_timeout = std::time::Duration::from_secs(1);
let pool = Pool::new(config, Arc::new(FailingFactory))?;
let progress = crate::prewarm::PrewarmProgress::new(5);
pool.progressive_prewarm(
2,
std::time::Duration::from_millis(5),
std::time::Duration::from_secs(5),
&progress,
)
.await;
let snap = progress.snapshot();
assert!(snap.is_completed, "应标记完成");
assert_eq!(snap.warmed, 0, "全部失败时 warmed=0");
assert!(snap.failed > 0, "应有失败记录");
Ok(())
}
#[tokio::test]
async fn test_pool_metrics_acquire_release() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(10).build()?;
let pool = Pool::new(config, Arc::new(MockConnectionFactory))?;
let metrics = pool.pool_metrics();
assert_eq!(metrics.acquire_count, 0);
assert_eq!(metrics.release_count, 0);
assert_eq!(metrics.connection_created_count, 0);
let conn = pool.acquire().await?;
let metrics = pool.pool_metrics();
assert_eq!(metrics.acquire_count, 1);
assert_eq!(metrics.connection_created_count, 1);
assert_eq!(metrics.acquire_failed_count, 0);
pool.release(conn).await;
let metrics = pool.pool_metrics();
assert_eq!(metrics.release_count, 1);
assert_eq!(metrics.connection_closed_count, 0);
Ok(())
}
#[tokio::test]
async fn test_pool_metrics_acquire_failed() -> Result<(), Box<dyn std::error::Error>> {
struct FailingFactory;
#[async_trait]
impl ConnectionFactory for FailingFactory {
async fn create(&self) -> Result<Box<dyn Connection>, crate::DbError> {
Err(crate::DbError::Internal("simulated failure".to_string()))
}
}
let config = PoolConfigBuilder::new().max_size(10).build()?;
let pool = Pool::new(config, Arc::new(FailingFactory))?;
let result = pool.acquire().await;
assert!(result.is_err());
let metrics = pool.pool_metrics();
assert_eq!(metrics.acquire_failed_count, 1);
assert_eq!(metrics.acquire_count, 0);
Ok(())
}
#[tokio::test]
async fn test_pool_metrics_connection_closed() -> Result<(), Box<dyn std::error::Error>> {
let config = PoolConfigBuilder::new().max_size(10).build()?;
let pool = Pool::new(config, Arc::new(MockConnectionFactory))?;
let conn = pool.acquire().await?;
pool.release(conn).await;
let status = pool.status().await;
assert_eq!(status.idle, 1);
pool.close_all().await;
let metrics = pool.pool_metrics();
assert_eq!(metrics.connection_closed_count, 1);
assert_eq!(metrics.connection_created_count, 1);
Ok(())
}
#[test]
fn test_pool_metrics_average_wait_time() {
let metrics = PoolMetrics {
acquire_count: 4,
acquire_failed_count: 1,
acquire_wait_time: Duration::from_millis(200),
release_count: 4,
connection_created_count: 2,
connection_closed_count: 0,
};
assert_eq!(
metrics.average_acquire_wait_time(),
Duration::from_millis(50)
);
let empty = PoolMetrics::default();
assert_eq!(empty.average_acquire_wait_time(), Duration::ZERO);
}
}