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
14pub(crate) const DEFAULT_POOL_CAPACITY: usize = 50;
16
17#[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: u64,
31}
32
33#[derive(Debug)]
35struct Pooled {
36 generation: u64,
37 item: PoolItem,
38}
39
40#[derive(Debug, Clone, PartialEq, Eq, Hash)]
43struct PoolKey {
44 prefix: PoolPrefix,
45 uri: String,
46 connection: Option<String>,
47}
48
49impl ConnectionPool {
50 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 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 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 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 pub fn put(&self, uri: &Uri, settings: Option<&AppliedRequestSettings>, value: PoolItem) {
145 self.put_with_generation(uri, settings, value);
146 }
147
148 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 pub fn generation(&self) -> u64 {
172 self.inner.lock().expect("must lock").generation
173 }
174
175 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 PoolKey {
210 prefix: class.into(),
211 uri: uri.to_string(),
212 connection: settings.map(AppliedRequestSettings::connection_key),
213 }
214 }
215}
216
217#[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#[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 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 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 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 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 let result = pool.get(PoolPrefix::Core, &uri, None);
500 assert!(result.is_none());
501
502 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 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 let result = pool_clone.get(PoolPrefix::Platform, &uri, None);
617 assert!(result.is_some());
618 }
619}