Skip to main content

rs_dapi_client/
connection_pool.rs

1use std::{
2    fmt::Display,
3    sync::{Arc, Mutex},
4};
5
6use lru::LruCache;
7
8use crate::{
9    request_settings::AppliedRequestSettings,
10    transport::{CoreGrpcClient, PlatformGrpcClient},
11    Uri,
12};
13
14/// Default capacity of the [ConnectionPool].
15pub(crate) const DEFAULT_POOL_CAPACITY: usize = 50;
16
17/// ConnectionPool represents pool of connections to DAPI nodes.
18///
19/// It can be cloned and shared between threads.
20/// Cloning the pool will create a new reference to the same pool.
21#[derive(Debug, Clone)]
22pub struct ConnectionPool {
23    inner: Arc<Mutex<PoolState>>,
24}
25
26#[derive(Debug)]
27struct PoolState {
28    connections: LruCache<PoolKey, Pooled>,
29    /// Generation of the connection pooled most recently.
30    generation: u64,
31}
32
33/// A pooled connection and the generation it was pooled at.
34#[derive(Debug)]
35struct Pooled {
36    generation: u64,
37    item: PoolItem,
38}
39
40/// Identity of a pooled connection: the client type, the node, and the
41/// connection-affecting settings (`None` when none were given).
42#[derive(Debug, Clone, PartialEq, Eq, Hash)]
43struct PoolKey {
44    prefix: PoolPrefix,
45    uri: String,
46    connection: Option<String>,
47}
48
49impl ConnectionPool {
50    /// Create a new pool with a given capacity.
51    /// The pool will evict the least recently used item when the capacity is reached.
52    ///
53    /// # Panics
54    ///
55    /// Panics if the capacity is zero.
56    pub fn new(capacity: usize) -> Self {
57        Self {
58            inner: Arc::new(Mutex::new(PoolState {
59                connections: LruCache::new(capacity.try_into().expect("must be non-zero")),
60                generation: 0,
61            })),
62        }
63    }
64}
65
66impl Default for ConnectionPool {
67    fn default() -> Self {
68        Self::new(DEFAULT_POOL_CAPACITY)
69    }
70}
71
72impl ConnectionPool {
73    /// Get item from the pool for the given uri and settings.
74    ///
75    /// # Arguments
76    /// * `prefix` -  Prefix for the item in the pool. Used to distinguish between Core and Platform clients.
77    /// * `uri` - URI of the node.
78    /// * `settings` - Applied request settings.
79    pub fn get(
80        &self,
81        prefix: PoolPrefix,
82        uri: &Uri,
83        settings: Option<&AppliedRequestSettings>,
84    ) -> Option<PoolItem> {
85        let key = Self::key(prefix, uri, settings);
86        self.inner
87            .lock()
88            .expect("must lock")
89            .connections
90            .get(&key)
91            .map(|pooled| pooled.item.clone())
92    }
93
94    /// Get value from cache or create it using provided closure.
95    /// If value is already in the cache, it will be returned.
96    /// If value is not in the cache, it will be created by calling `create()` and stored in the cache.
97    ///
98    /// # Arguments
99    /// * `prefix` -  Prefix for the item in the pool. Used to distinguish between Core and Platform clients.
100    /// * `uri` - URI of the node.
101    /// * `settings` - Applied request settings.
102    pub fn get_or_create<E>(
103        &self,
104        prefix: PoolPrefix,
105        uri: &Uri,
106        settings: Option<&AppliedRequestSettings>,
107        create: impl FnOnce() -> Result<PoolItem, E>,
108    ) -> Result<PoolItem, E> {
109        self.get_or_create_with_generation(prefix, uri, settings, create)
110            .map(|(item, _)| item)
111    }
112
113    /// Like [ConnectionPool::get_or_create], and also returns the generation
114    /// the returned connection was pooled at (see [ConnectionPool::generation]).
115    ///
116    /// The generation is read under the same lock that finds or stores the
117    /// connection, so it is the returned connection's own even when other
118    /// threads replace it right afterwards.
119    pub fn get_or_create_with_generation<E>(
120        &self,
121        prefix: PoolPrefix,
122        uri: &Uri,
123        settings: Option<&AppliedRequestSettings>,
124        create: impl FnOnce() -> Result<PoolItem, E>,
125    ) -> Result<(PoolItem, u64), E> {
126        let key = Self::key(prefix, uri, settings);
127        let cached = self
128            .inner
129            .lock()
130            .expect("must lock")
131            .connections
132            .get(&key)
133            .map(|pooled| (pooled.item.clone(), pooled.generation));
134        if let Some(cached) = cached {
135            return Ok(cached);
136        }
137
138        let item = create()?;
139        let generation = self.put_with_generation(uri, settings, item.clone());
140        Ok((item, generation))
141    }
142
143    /// Put item into the pool for the given uri and settings.
144    pub fn put(&self, uri: &Uri, settings: Option<&AppliedRequestSettings>, value: PoolItem) {
145        self.put_with_generation(uri, settings, value);
146    }
147
148    /// Put item into the pool and return the generation it was pooled at.
149    fn put_with_generation(
150        &self,
151        uri: &Uri,
152        settings: Option<&AppliedRequestSettings>,
153        value: PoolItem,
154    ) -> u64 {
155        let key = Self::key(&value, uri, settings);
156        let mut state = self.inner.lock().expect("must lock");
157        state.generation += 1;
158        let generation = state.generation;
159        state.connections.put(
160            key,
161            Pooled {
162                generation,
163                item: value,
164            },
165        );
166        generation
167    }
168
169    /// Generation of the connection pooled most recently. Every connection
170    /// put into the pool afterwards gets a higher generation.
171    pub fn generation(&self) -> u64 {
172        self.inner.lock().expect("must lock").generation
173    }
174
175    /// Drop every connection to `uri` pooled at or before `generation`,
176    /// whatever its prefix and connection settings.
177    ///
178    /// A request that misses its deadline may have been sent over a half-open
179    /// connection: the network path died after the request left, and nothing
180    /// on the idle channel would ever notice. Keeping it pooled would stall
181    /// the next request sent to the same node, so the executor evicts it and
182    /// the next request dials a fresh connection.
183    ///
184    /// Connections pooled after `generation` stay. They were dialed after the
185    /// timed-out attempt took its connection, for example by a concurrent
186    /// request that already timed out on the same node and reconnected.
187    pub fn remove_uri(&self, uri: &Uri, generation: u64) {
188        let uri = uri.to_string();
189        let mut state = self.inner.lock().expect("must lock");
190        let stale: Vec<PoolKey> = state
191            .connections
192            .iter()
193            .filter(|(key, pooled)| key.uri == uri && pooled.generation <= generation)
194            .map(|(key, _)| key.clone())
195            .collect();
196        for key in stale {
197            state.connections.pop(&key);
198        }
199    }
200
201    fn key<C: Into<PoolPrefix>>(
202        class: C,
203        uri: &Uri,
204        settings: Option<&AppliedRequestSettings>,
205    ) -> PoolKey {
206        // Only connection-affecting settings participate in the key (see
207        // `AppliedRequestSettings::connection_key`), so requests differing only
208        // in per-request knobs (timeout, retries, banning) share a connection.
209        PoolKey {
210            prefix: class.into(),
211            uri: uri.to_string(),
212            connection: settings.map(AppliedRequestSettings::connection_key),
213        }
214    }
215}
216
217/// Item stored in the pool.
218///
219/// We use an enum as we need to represent two different types of clients.
220#[derive(Clone, Debug)]
221pub enum PoolItem {
222    Core(CoreGrpcClient),
223    Platform(PlatformGrpcClient),
224}
225
226impl From<PlatformGrpcClient> for PoolItem {
227    fn from(client: PlatformGrpcClient) -> Self {
228        Self::Platform(client)
229    }
230}
231impl From<CoreGrpcClient> for PoolItem {
232    fn from(client: CoreGrpcClient) -> Self {
233        Self::Core(client)
234    }
235}
236
237impl From<PoolItem> for PlatformGrpcClient {
238    fn from(client: PoolItem) -> Self {
239        match client {
240            PoolItem::Platform(client) => client,
241            _ => {
242                tracing::error!(
243                    ?client,
244                    "invalid connection fetched from pool: expected platform client"
245                );
246                panic!("ClientType is not Platform: {:?}", client)
247            }
248        }
249    }
250}
251
252impl From<PoolItem> for CoreGrpcClient {
253    fn from(client: PoolItem) -> Self {
254        match client {
255            PoolItem::Core(client) => client,
256            _ => {
257                tracing::error!(
258                    ?client,
259                    "invalid connection fetched from pool: expected core client"
260                );
261                panic!("ClientType is not Core: {:?}", client)
262            }
263        }
264    }
265}
266
267/// Prefix for the item in the pool. Used to distinguish between Core and Platform clients.
268#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
269pub enum PoolPrefix {
270    Core,
271    Platform,
272}
273impl Display for PoolPrefix {
274    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
275        match self {
276            PoolPrefix::Core => write!(f, "Core"),
277            PoolPrefix::Platform => write!(f, "Platform"),
278        }
279    }
280}
281impl From<&PoolItem> for PoolPrefix {
282    fn from(item: &PoolItem) -> Self {
283        match item {
284            PoolItem::Core(_) => PoolPrefix::Core,
285            PoolItem::Platform(_) => PoolPrefix::Platform,
286        }
287    }
288}
289
290#[cfg(test)]
291mod tests {
292    use super::*;
293    use crate::RequestSettings;
294    use dapi_grpc::tonic::transport::Channel;
295    use std::str::FromStr;
296    use std::time::Duration;
297
298    fn test_uri() -> Uri {
299        Uri::from_str("http://127.0.0.1:3000").unwrap()
300    }
301
302    fn make_platform_pool_item() -> PoolItem {
303        let channel = Channel::builder(test_uri()).connect_lazy();
304        PoolItem::Platform(PlatformGrpcClient::new(channel))
305    }
306
307    fn make_core_pool_item() -> PoolItem {
308        let channel = Channel::builder(test_uri()).connect_lazy();
309        PoolItem::Core(CoreGrpcClient::new(channel))
310    }
311
312    #[test]
313    fn test_connection_pool_new() {
314        let pool = ConnectionPool::new(10);
315        let result = pool.get(PoolPrefix::Platform, &test_uri(), None);
316        assert!(result.is_none());
317    }
318
319    #[test]
320    fn test_connection_pool_default() {
321        let pool = ConnectionPool::default();
322        let result = pool.get(PoolPrefix::Core, &test_uri(), None);
323        assert!(result.is_none());
324    }
325
326    #[tokio::test]
327    async fn test_connection_pool_put_and_get_platform() {
328        let pool = ConnectionPool::new(10);
329        let uri = test_uri();
330        let item = make_platform_pool_item();
331
332        pool.put(&uri, None, item);
333
334        let result = pool.get(PoolPrefix::Platform, &uri, None);
335        assert!(result.is_some());
336        assert!(matches!(result.unwrap(), PoolItem::Platform(_)));
337    }
338
339    #[tokio::test]
340    async fn test_connection_pool_put_and_get_core() {
341        let pool = ConnectionPool::new(10);
342        let uri = test_uri();
343        let item = make_core_pool_item();
344
345        pool.put(&uri, None, item);
346
347        let result = pool.get(PoolPrefix::Core, &uri, None);
348        assert!(result.is_some());
349        assert!(matches!(result.unwrap(), PoolItem::Core(_)));
350    }
351
352    #[tokio::test]
353    async fn test_connection_pool_get_or_create_creates_new() {
354        let pool = ConnectionPool::new(10);
355        let uri = test_uri();
356
357        let result: Result<PoolItem, String> =
358            pool.get_or_create(PoolPrefix::Platform, &uri, None, || {
359                Ok(make_platform_pool_item())
360            });
361
362        assert!(result.is_ok());
363
364        // Second call should return cached version
365        let mut create_called = false;
366        let result2: Result<PoolItem, String> =
367            pool.get_or_create(PoolPrefix::Platform, &uri, None, || {
368                create_called = true;
369                Ok(make_platform_pool_item())
370            });
371
372        assert!(result2.is_ok());
373        assert!(
374            !create_called,
375            "create should not be called for cached item"
376        );
377    }
378
379    #[test]
380    fn test_connection_pool_get_or_create_error_not_cached() {
381        let pool = ConnectionPool::new(10);
382        let uri = test_uri();
383
384        let result: Result<PoolItem, String> =
385            pool.get_or_create(PoolPrefix::Platform, &uri, None, || {
386                Err("creation failed".to_string())
387            });
388
389        assert!(result.is_err());
390
391        // Pool should still be empty after failed creation
392        let cached = pool.get(PoolPrefix::Platform, &uri, None);
393        assert!(cached.is_none());
394    }
395
396    #[test]
397    fn test_pool_prefix_display() {
398        assert_eq!(format!("{}", PoolPrefix::Core), "Core");
399        assert_eq!(format!("{}", PoolPrefix::Platform), "Platform");
400    }
401
402    #[tokio::test]
403    async fn test_pool_prefix_from_pool_item() {
404        let platform_item = make_platform_pool_item();
405        let prefix: PoolPrefix = (&platform_item).into();
406        assert!(matches!(prefix, PoolPrefix::Platform));
407
408        let core_item = make_core_pool_item();
409        let prefix: PoolPrefix = (&core_item).into();
410        assert!(matches!(prefix, PoolPrefix::Core));
411    }
412
413    #[tokio::test]
414    async fn test_pool_item_from_platform_client() {
415        let channel = Channel::builder(test_uri()).connect_lazy();
416        let client = PlatformGrpcClient::new(channel);
417        let item: PoolItem = client.into();
418        assert!(matches!(item, PoolItem::Platform(_)));
419    }
420
421    #[tokio::test]
422    async fn test_pool_item_from_core_client() {
423        let channel = Channel::builder(test_uri()).connect_lazy();
424        let client = CoreGrpcClient::new(channel);
425        let item: PoolItem = client.into();
426        assert!(matches!(item, PoolItem::Core(_)));
427    }
428
429    #[tokio::test]
430    async fn test_pool_item_into_platform_client() {
431        let item = make_platform_pool_item();
432        let _client: PlatformGrpcClient = item.into();
433    }
434
435    #[tokio::test]
436    async fn test_pool_item_into_core_client() {
437        let item = make_core_pool_item();
438        let _client: CoreGrpcClient = item.into();
439    }
440
441    #[tokio::test]
442    #[should_panic(expected = "ClientType is not Platform")]
443    async fn test_pool_item_core_into_platform_panics() {
444        let item = make_core_pool_item();
445        let _client: PlatformGrpcClient = item.into();
446    }
447
448    #[tokio::test]
449    #[should_panic(expected = "ClientType is not Core")]
450    async fn test_pool_item_platform_into_core_panics() {
451        let item = make_platform_pool_item();
452        let _client: CoreGrpcClient = item.into();
453    }
454
455    #[tokio::test]
456    async fn test_connection_pool_shares_client_across_per_request_settings() {
457        let pool = ConnectionPool::new(10);
458        let uri = test_uri();
459
460        // Settings differing only in per-request knobs (timeout, retries,
461        // banning) must map to the same pooled connection...
462        let stored = RequestSettings {
463            timeout: Some(Duration::from_secs(30)),
464            retries: Some(3),
465            ban_failed_address: Some(false),
466            ..RequestSettings::default()
467        }
468        .finalize();
469        pool.put(&uri, Some(&stored), make_platform_pool_item());
470
471        let default = RequestSettings::default().finalize();
472        assert!(
473            pool.get(PoolPrefix::Platform, &uri, Some(&default))
474                .is_some(),
475            "per-request settings must not split pooled connections"
476        );
477
478        // ...while connection-affecting settings still get their own entry.
479        let connect = RequestSettings {
480            connect_timeout: Some(Duration::from_secs(3)),
481            ..RequestSettings::default()
482        }
483        .finalize();
484        assert!(
485            pool.get(PoolPrefix::Platform, &uri, Some(&connect))
486                .is_none(),
487            "connection-affecting settings must key separate connections"
488        );
489    }
490
491    #[tokio::test]
492    async fn test_connection_pool_different_prefixes_different_keys() {
493        let pool = ConnectionPool::new(10);
494        let uri = test_uri();
495
496        pool.put(&uri, None, make_platform_pool_item());
497
498        // Core prefix should not find a Platform item
499        let result = pool.get(PoolPrefix::Core, &uri, None);
500        assert!(result.is_none());
501
502        // Platform prefix should find it
503        let result = pool.get(PoolPrefix::Platform, &uri, None);
504        assert!(result.is_some());
505    }
506
507    #[tokio::test]
508    async fn should_remove_every_pooled_connection_to_an_uri() {
509        let pool = ConnectionPool::new(10);
510        let uri = test_uri();
511        let longer_port = Uri::from_str("http://127.0.0.1:30001").unwrap();
512        let connect_timeout = RequestSettings {
513            connect_timeout: Some(Duration::from_secs(3)),
514            ..RequestSettings::default()
515        }
516        .finalize();
517
518        pool.put(&uri, None, make_platform_pool_item());
519        pool.put(&uri, None, make_core_pool_item());
520        pool.put(&uri, Some(&connect_timeout), make_platform_pool_item());
521        pool.put(&longer_port, None, make_platform_pool_item());
522
523        pool.remove_uri(&uri, pool.generation());
524
525        assert!(pool.get(PoolPrefix::Platform, &uri, None).is_none());
526        assert!(pool.get(PoolPrefix::Core, &uri, None).is_none());
527        assert!(pool
528            .get(PoolPrefix::Platform, &uri, Some(&connect_timeout))
529            .is_none());
530        assert!(
531            pool.get(PoolPrefix::Platform, &longer_port, None).is_some(),
532            "a URI that merely starts with the evicted one must stay pooled"
533        );
534    }
535
536    #[tokio::test]
537    async fn should_remove_only_the_exact_uri_when_another_extends_it_past_a_colon() {
538        let cases = [
539            ("http://node", "http://node:443"),
540            ("http://node/grpc", "http://node/grpc:8080"),
541        ];
542        for (evicted, kept) in cases {
543            let pool = ConnectionPool::new(10);
544            let evicted = Uri::from_str(evicted).unwrap();
545            let kept = Uri::from_str(kept).unwrap();
546            pool.put(&evicted, None, make_platform_pool_item());
547            pool.put(&kept, None, make_platform_pool_item());
548
549            pool.remove_uri(&evicted, pool.generation());
550
551            assert!(
552                pool.get(PoolPrefix::Platform, &evicted, None).is_none(),
553                "{evicted} must be evicted"
554            );
555            assert!(
556                pool.get(PoolPrefix::Platform, &kept, None).is_some(),
557                "{kept} must stay pooled when {evicted} is evicted"
558            );
559        }
560    }
561
562    #[tokio::test]
563    async fn should_keep_connections_pooled_after_the_given_generation() {
564        let pool = ConnectionPool::new(10);
565        let uri = test_uri();
566        pool.put(&uri, None, make_platform_pool_item());
567        let used_by_attempt = pool.generation();
568        // Replaced after the attempt took its connection.
569        pool.put(&uri, None, make_platform_pool_item());
570
571        pool.remove_uri(&uri, used_by_attempt);
572        assert!(
573            pool.get(PoolPrefix::Platform, &uri, None).is_some(),
574            "a connection pooled after the given generation must stay"
575        );
576
577        pool.remove_uri(&uri, pool.generation());
578        assert!(pool.get(PoolPrefix::Platform, &uri, None).is_none());
579    }
580
581    #[tokio::test]
582    async fn should_return_the_generation_of_the_connection_it_takes() {
583        let pool = ConnectionPool::new(10);
584        let uri = test_uri();
585
586        let (_, created) = pool
587            .get_or_create_with_generation(PoolPrefix::Platform, &uri, None, || {
588                Ok::<_, String>(make_platform_pool_item())
589            })
590            .unwrap();
591        assert_eq!(created, pool.generation());
592
593        let (_, taken) = pool
594            .get_or_create_with_generation(PoolPrefix::Platform, &uri, None, || {
595                Err("the pooled connection must be reused".to_string())
596            })
597            .unwrap();
598        pool.put(&uri, None, make_platform_pool_item());
599
600        assert_eq!(taken, created);
601        assert!(
602            taken < pool.generation(),
603            "a replacement pooled afterwards must get a higher generation"
604        );
605    }
606
607    #[tokio::test]
608    async fn test_connection_pool_clone_shares_data() {
609        let pool = ConnectionPool::new(10);
610        let pool_clone = pool.clone();
611        let uri = test_uri();
612
613        pool.put(&uri, None, make_platform_pool_item());
614
615        // Clone should see the same data
616        let result = pool_clone.get(PoolPrefix::Platform, &uri, None);
617        assert!(result.is_some());
618    }
619}