Skip to main content

soma_fleet/
pool.rs

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
12/// Async connection pool keyed by stable host identity and exact topology revision.
13pub 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    /// Creates an empty pool using the supplied connection factory.
26    #[must_use]
27    pub fn new(factory: Arc<F>) -> Self {
28        Self {
29            factory,
30            cells: Mutex::new(BTreeMap::new()),
31        }
32    }
33
34    /// Returns an existing exact-revision connection or opens it once.
35    ///
36    /// Concurrent cold-cache callers for the same revision share one
37    /// initialization cell. The map lock is never held across an await.
38    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    /// Invalidates every cached revision for one host and closes each handle.
76    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    /// Evicts connections absent from the current topology snapshot.
92    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    /// Closes and removes every initialized connection.
112    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    /// Returns cached revision-key count, including in-flight initializations.
123    pub async fn len(&self) -> usize {
124        self.cells.lock().await.len()
125    }
126
127    /// Returns whether no revision keys are cached.
128    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;