1use crate::dialect::Dialect;
12use crate::driver::{Driver, DriverConnection};
13use rustlavel_core::{Error, Result};
14use std::collections::VecDeque;
15use std::sync::Arc;
16use tokio::sync::{Mutex, Semaphore};
17
18struct Inner {
19 driver: Arc<dyn Driver>,
20 idle: Mutex<VecDeque<(u64, Box<dyn DriverConnection>)>>,
24 permits: Arc<Semaphore>,
26}
27
28#[derive(Clone)]
29pub struct Pool {
30 inner: Arc<Inner>,
31}
32
33impl Pool {
34 pub fn new(driver: Arc<dyn Driver>) -> Self {
37 let permits = Arc::new(Semaphore::new(driver.max_connections().max(1)));
38 Pool { inner: Arc::new(Inner { driver, idle: Mutex::new(VecDeque::new()), permits }) }
39 }
40
41 pub async fn verify(&self) -> Result<()> {
44 let mut connection = self.acquire().await?;
45 connection.simple_query("select 1").await?;
46 Ok(())
47 }
48
49 pub fn driver(&self) -> &Arc<dyn Driver> {
50 &self.inner.driver
51 }
52
53 pub fn dialect(&self) -> Arc<dyn Dialect> {
54 self.inner.driver.dialect()
55 }
56
57 pub async fn acquire(&self) -> Result<PooledConnection> {
59 let permit = Arc::clone(&self.inner.permits)
60 .acquire_owned()
61 .await
62 .map_err(|_| Error::msg("the database pool has been closed"))?;
63
64 let generation = self.inner.driver.generation();
65
66 loop {
72 let Some((opened_under, connection)) = self.inner.idle.lock().await.pop_front() else {
73 break;
74 };
75
76 if opened_under == generation {
77 return Ok(PooledConnection {
78 connection: Some(connection),
79 generation,
80 pool: Arc::clone(&self.inner),
81 _permit: permit,
82 });
83 }
84
85 connection.close().await;
86 }
87
88 let connection = self.inner.driver.connect().await?;
89 Ok(PooledConnection {
90 connection: Some(connection),
91 generation,
92 pool: Arc::clone(&self.inner),
93 _permit: permit,
94 })
95 }
96
97 pub async fn idle_count(&self) -> usize {
99 self.inner.idle.lock().await.len()
100 }
101
102 pub async fn open_count(&self) -> usize {
113 let borrowed = self
114 .inner
115 .driver
116 .max_connections()
117 .max(1)
118 .saturating_sub(self.inner.permits.available_permits());
119 self.inner.idle.lock().await.len() + borrowed
120 }
121
122 pub async fn close_idle(&self, limit: usize) -> usize {
129 let mut closed = 0;
130 while closed < limit {
131 let Some((_, connection)) = self.inner.idle.lock().await.pop_front() else { break };
132 connection.close().await;
133 closed += 1;
134 }
135 closed
136 }
137
138 pub async fn close(&self) {
140 let mut idle = self.inner.idle.lock().await;
141 while let Some((_, connection)) = idle.pop_front() {
142 connection.close().await;
143 }
144 }
145
146 pub async fn retire_superseded(&self) -> usize {
157 let generation = self.inner.driver.generation();
158 let mut idle = self.inner.idle.lock().await;
159
160 let mut keeping = VecDeque::with_capacity(idle.len());
161 let mut closed = 0;
162
163 while let Some((opened_under, connection)) = idle.pop_front() {
164 if opened_under == generation {
165 keeping.push_back((opened_under, connection));
166 } else {
167 connection.close().await;
168 closed += 1;
169 }
170 }
171
172 *idle = keeping;
173 closed
174 }
175}
176
177pub struct PooledConnection {
179 connection: Option<Box<dyn DriverConnection>>,
180 generation: u64,
182 pool: Arc<Inner>,
183 _permit: tokio::sync::OwnedSemaphorePermit,
185}
186
187impl std::ops::Deref for PooledConnection {
188 type Target = dyn DriverConnection;
189
190 fn deref(&self) -> &(dyn DriverConnection + 'static) {
191 self.connection.as_deref().expect("connection is present until drop")
192 }
193}
194
195impl std::ops::DerefMut for PooledConnection {
196 fn deref_mut(&mut self) -> &mut (dyn DriverConnection + 'static) {
197 self.connection.as_deref_mut().expect("connection is present until drop")
198 }
199}
200
201impl Drop for PooledConnection {
202 fn drop(&mut self) {
203 let Some(connection) = self.connection.take() else { return };
204
205 if connection.is_broken() || connection.in_transaction() {
208 tokio::spawn(async move { connection.close().await });
209 return;
210 }
211
212 let pool = Arc::clone(&self.pool);
216 let generation = self.generation;
217 tokio::spawn(async move {
218 if generation != pool.driver.generation() {
219 connection.close().await;
220 return;
221 }
222 pool.idle.lock().await.push_back((generation, connection));
223 });
224 }
225}
226
227#[cfg(test)]
228mod tests {
229 use super::*;
230 use crate::dialect::Postgres;
231 use crate::driver::BoxFuture;
232
233 struct Counting {
236 opened: Arc<std::sync::atomic::AtomicUsize>,
237 closed: Arc<std::sync::atomic::AtomicUsize>,
238 generation: Arc<std::sync::atomic::AtomicU64>,
239 }
240
241 struct Nothing(Arc<std::sync::atomic::AtomicUsize>);
242
243 impl DriverConnection for Nothing {
244 fn query<'a>(
245 &'a mut self,
246 _sql: &'a str,
247 _params: &'a [crate::value::Value],
248 ) -> BoxFuture<'a, Result<crate::driver::QueryResult>> {
249 Box::pin(async { Err(Error::msg("not a real connection")) })
250 }
251
252 fn simple_query<'a>(
253 &'a mut self,
254 _sql: &'a str,
255 ) -> BoxFuture<'a, Result<crate::driver::QueryResult>> {
256 Box::pin(async { Err(Error::msg("not a real connection")) })
257 }
258
259 fn is_broken(&self) -> bool {
260 false
261 }
262
263 fn in_transaction(&self) -> bool {
264 false
265 }
266
267 fn close(self: Box<Self>) -> BoxFuture<'static, ()> {
268 self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
269 Box::pin(async {})
270 }
271 }
272
273 impl Driver for Counting {
274 fn dialect(&self) -> Arc<dyn Dialect> {
275 Arc::new(Postgres)
276 }
277
278 fn connect(&self) -> BoxFuture<'_, Result<Box<dyn DriverConnection>>> {
279 self.opened.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
280 let closed = Arc::clone(&self.closed);
281 Box::pin(async move { Ok(Box::new(Nothing(closed)) as Box<dyn DriverConnection>) })
282 }
283
284 fn describe(&self) -> String {
285 "test://counting".into()
286 }
287
288 fn generation(&self) -> u64 {
289 self.generation.load(std::sync::atomic::Ordering::Acquire)
290 }
291 }
292
293 fn counting() -> (Pool, Arc<std::sync::atomic::AtomicUsize>, Arc<std::sync::atomic::AtomicUsize>, Arc<std::sync::atomic::AtomicU64>) {
294 let opened = Arc::new(std::sync::atomic::AtomicUsize::new(0));
295 let closed = Arc::new(std::sync::atomic::AtomicUsize::new(0));
296 let generation = Arc::new(std::sync::atomic::AtomicU64::new(1));
297 let driver = Counting {
298 opened: Arc::clone(&opened),
299 closed: Arc::clone(&closed),
300 generation: Arc::clone(&generation),
301 };
302 (Pool::new(Arc::new(driver)), opened, closed, generation)
303 }
304
305 async fn settle() {
308 for _ in 0..8 {
309 tokio::task::yield_now().await;
310 }
311 }
312
313 #[tokio::test]
314 async fn a_connection_comes_back_and_is_reused() {
315 let (pool, opened, _, _) = counting();
316
317 drop(pool.acquire().await.unwrap());
318 settle().await;
319 drop(pool.acquire().await.unwrap());
320 settle().await;
321
322 assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 1, "the second borrow reused it");
323 }
324
325 #[tokio::test]
326 async fn an_idle_connection_from_a_rotated_credential_is_never_handed_out() {
327 let (pool, opened, closed, generation) = counting();
331
332 drop(pool.acquire().await.unwrap());
333 settle().await;
334 assert_eq!(pool.idle_count().await, 1);
335
336 generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
337
338 drop(pool.acquire().await.unwrap());
339 settle().await;
340
341 assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 2, "a fresh connection");
342 assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 1, "the stale one was closed");
343 }
344
345 #[tokio::test]
346 async fn a_borrowed_connection_is_retired_when_it_comes_back() {
347 let (pool, _, closed, generation) = counting();
351
352 let borrowed = pool.acquire().await.unwrap();
353 generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
354 drop(borrowed);
355 settle().await;
356
357 assert_eq!(pool.idle_count().await, 0);
358 assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 1);
359 }
360
361 #[tokio::test]
362 async fn retiring_early_closes_the_stale_and_keeps_the_current() {
363 let (pool, _, closed, generation) = counting();
364
365 let first = pool.acquire().await.unwrap();
367 let second = pool.acquire().await.unwrap();
368 drop(first);
369 drop(second);
370 settle().await;
371 assert_eq!(pool.idle_count().await, 2);
372
373 generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
374
375 assert_eq!(pool.retire_superseded().await, 2);
376 assert_eq!(pool.idle_count().await, 0);
377 assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 2);
378
379 drop(pool.acquire().await.unwrap());
381 settle().await;
382 assert_eq!(pool.retire_superseded().await, 0);
383 assert_eq!(pool.idle_count().await, 1);
384 }
385
386 #[tokio::test]
387 async fn a_pool_with_static_credentials_never_retires_anything() {
388 let (pool, opened, closed, _) = counting();
391
392 for _ in 0..5 {
393 drop(pool.acquire().await.unwrap());
394 settle().await;
395 }
396
397 assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 1);
398 assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 0);
399 }
400
401 struct Unreachable;
404
405 impl Driver for Unreachable {
406 fn dialect(&self) -> Arc<dyn Dialect> {
407 Arc::new(Postgres)
408 }
409
410 fn connect(&self) -> BoxFuture<'_, Result<Box<dyn DriverConnection>>> {
411 Box::pin(async { Err(Error::msg("nothing is listening")) })
412 }
413
414 fn describe(&self) -> String {
415 "test://unreachable".into()
416 }
417
418 fn max_connections(&self) -> usize {
419 3
420 }
421 }
422
423 #[tokio::test]
424 async fn a_pool_opens_nothing_until_it_is_used() {
425 let pool = Pool::new(Arc::new(Unreachable));
426 assert_eq!(pool.idle_count().await, 0);
427 }
428
429 #[tokio::test]
430 async fn acquiring_reports_the_drivers_failure() {
431 let pool = Pool::new(Arc::new(Unreachable));
432
433 let error = match pool.acquire().await {
434 Err(error) => error.to_string(),
435 Ok(_) => panic!("this driver cannot connect"),
436 };
437 assert!(error.contains("nothing is listening"), "{error}");
438 }
439
440 #[tokio::test]
441 async fn the_pool_carries_its_drivers_dialect() {
442 let pool = Pool::new(Arc::new(Unreachable));
443 assert_eq!(pool.dialect().name(), "postgres");
444 }
445}