1use std::collections::{BTreeMap, BTreeSet};
2use std::sync::Arc;
3
4use tokio::sync::{Mutex, OnceCell};
5use tokio_util::sync::CancellationToken;
6
7use crate::{ConnectionFactory, FleetResult, HostId, HostRecord, PoolKey, TopologySnapshot};
8
9type ConnectionCell<C> = Arc<OnceCell<Arc<C>>>;
10type ConnectionCells<C> = BTreeMap<PoolKey, ConnectionCell<C>>;
11
12pub struct ConnectionPool<F>
14where
15 F: ConnectionFactory,
16{
17 factory: Arc<F>,
18 cells: Mutex<ConnectionCells<F::Connection>>,
19}
20
21impl<F> ConnectionPool<F>
22where
23 F: ConnectionFactory,
24{
25 #[must_use]
27 pub fn new(factory: Arc<F>) -> Self {
28 Self {
29 factory,
30 cells: Mutex::new(BTreeMap::new()),
31 }
32 }
33
34 pub async fn get_or_connect(
39 &self,
40 host: &HostRecord,
41 cancellation: &CancellationToken,
42 ) -> FleetResult<Arc<F::Connection>> {
43 let key = host.pool_key();
44 let cell = {
45 let mut cells = self.cells.lock().await;
46 Arc::clone(
47 cells
48 .entry(key.clone())
49 .or_insert_with(|| Arc::new(OnceCell::new())),
50 )
51 };
52
53 let result = cell
54 .get_or_try_init(|| async {
55 let connection = self.factory.connect(host, cancellation).await?;
56 Ok::<Arc<F::Connection>, crate::FleetError>(Arc::new(connection))
57 })
58 .await;
59
60 match result {
61 Ok(connection) => Ok(Arc::clone(connection)),
62 Err(error) => {
63 let mut cells = self.cells.lock().await;
64 if cells
65 .get(&key)
66 .is_some_and(|current| Arc::ptr_eq(current, &cell) && current.get().is_none())
67 {
68 cells.remove(&key);
69 }
70 Err(error)
71 }
72 }
73 }
74
75 pub async fn invalidate_host(&self, host: &HostId) -> FleetResult<usize> {
77 let removed = {
78 let mut cells = self.cells.lock().await;
79 let keys = cells
80 .keys()
81 .filter(|key| key.host() == host)
82 .cloned()
83 .collect::<Vec<_>>();
84 keys.into_iter()
85 .filter_map(|key| cells.remove(&key))
86 .collect::<Vec<_>>()
87 };
88 self.close_cells(removed).await
89 }
90
91 pub async fn retain_snapshot(&self, snapshot: &TopologySnapshot) -> FleetResult<usize> {
93 let current = snapshot
94 .hosts()
95 .map(HostRecord::pool_key)
96 .collect::<BTreeSet<_>>();
97 let removed = {
98 let mut cells = self.cells.lock().await;
99 let keys = cells
100 .keys()
101 .filter(|key| !current.contains(*key))
102 .cloned()
103 .collect::<Vec<_>>();
104 keys.into_iter()
105 .filter_map(|key| cells.remove(&key))
106 .collect::<Vec<_>>()
107 };
108 self.close_cells(removed).await
109 }
110
111 pub async fn shutdown(&self) -> FleetResult<usize> {
113 let removed = {
114 let mut cells = self.cells.lock().await;
115 std::mem::take(&mut *cells)
116 .into_values()
117 .collect::<Vec<_>>()
118 };
119 self.close_cells(removed).await
120 }
121
122 pub async fn len(&self) -> usize {
124 self.cells.lock().await.len()
125 }
126
127 pub async fn is_empty(&self) -> bool {
129 self.cells.lock().await.is_empty()
130 }
131
132 async fn close_cells(&self, cells: Vec<ConnectionCell<F::Connection>>) -> FleetResult<usize> {
133 let mut closed = 0;
134 for cell in cells {
135 if let Some(connection) = cell.get() {
136 self.factory.close(connection.as_ref()).await?;
137 closed += 1;
138 }
139 }
140 Ok(closed)
141 }
142}
143
144#[cfg(test)]
145#[path = "pool_tests.rs"]
146mod tests;