rust_grpc_lib/pool.rs
1//! Process-wide gRPC channel pool.
2//!
3//! Channels are keyed by endpoint string and created lazily on first use.
4//! Subsequent calls with the same endpoint string reuse the cached channel.
5//! The pool is backed by a `RwLock<HashMap<String, Channel>>` and is safe to
6//! use from multiple threads.
7//!
8//! The sole public function is [`get_channel`]. Consumers do not call it
9//! directly — the `from_endpoint_with_provider` and `from_endpoint`
10//! constructors generated by `#[derive(GrpcClient)]` /
11//! `#[derive(GrpcNoAuthClient)]` call it internally.
12//!
13//! Keepalive timing is controlled by two environment variables read at channel
14//! creation time:
15//!
16//! | Variable | Default |
17//! |---|---|
18//! | `RUST_GRPC_LIB_KEEP_ALIVE_INTERVAL_SECS` | `30` |
19//! | `RUST_GRPC_LIB_KEEP_ALIVE_TIMEOUT_SECS` | `10` |
20
21use std::{
22 collections::HashMap,
23 sync::{LazyLock, RwLock},
24 time::Duration,
25};
26
27use rust_env_var_lib::env_var;
28use tonic::transport::{Channel, Endpoint, Error};
29
30#[cfg(test)]
31mod tests;
32
33type ChannelMap = RwLock<HashMap<String, Channel>>;
34
35const KEEP_ALIVE_INTERVAL_VAR: &str = "RUST_GRPC_LIB_KEEP_ALIVE_INTERVAL_SECS";
36const KEEP_ALIVE_INTERVAL_DEFAULT: u64 = 30;
37const KEEP_ALIVE_TIMEOUT_VAR: &str = "RUST_GRPC_LIB_KEEP_ALIVE_TIMEOUT_SECS";
38const KEEP_ALIVE_TIMEOUT_DEFAULT: u64 = 10;
39
40static POOL: LazyLock<ChannelMap> = LazyLock::new(RwLock::default);
41
42/// Return a [`Channel`] for `endpoint`, reusing a cached one if it exists.
43///
44/// The channel is connected lazily: no network activity occurs until the first
45/// RPC is made. Keepalive settings are applied at creation time from the
46/// environment variables documented in the module-level doc.
47///
48/// # Errors
49///
50/// Returns [`tonic::transport::Error`] if `endpoint` is not a valid URI.
51///
52/// # Panics
53///
54/// Panics if called outside a Tokio runtime context (required by tonic's
55/// transport layer).
56pub fn get_channel(endpoint: &str) -> Result<Channel, Error> {
57 get_or_create_channel(endpoint, &POOL)
58}
59
60/// Look up an existing channel for `endpoint` in the pool, or create and
61/// insert one if none exists.
62///
63/// Uses a double-checked lock: a read lock is taken first to avoid write
64/// contention on the common path where the channel already exists.
65///
66/// ### Poison recovery
67/// If the thread should panic (somehow) while the lock is held, the pool will
68/// become "poisoned". The next read or write lock will return an error wrapping
69/// the lock handle.
70///
71/// As the only mutations to the pool are insertions, the state of the pool should
72/// always be valid. It is therefore safe to simply extract the lock handle from the
73/// error and proceed.
74fn get_or_create_channel(endpoint: &str, pool: &ChannelMap) -> Result<Channel, Error> {
75 if let Some(channel) = pool
76 .read()
77 .unwrap_or_else(|e| e.into_inner())
78 .get(endpoint)
79 .cloned()
80 {
81 return Ok(channel);
82 }
83
84 let mut lock = pool.write().unwrap_or_else(|e| e.into_inner());
85 // Re-check after acquiring the write lock: another thread may have
86 // inserted the channel between our read and write lock acquisitions.
87 if let Some(channel) = lock.get(endpoint).cloned() {
88 return Ok(channel);
89 }
90
91 let channel = Endpoint::new(endpoint.to_string())?
92 .http2_keep_alive_interval(duration_from_env(
93 KEEP_ALIVE_INTERVAL_VAR,
94 KEEP_ALIVE_INTERVAL_DEFAULT,
95 ))
96 .keep_alive_timeout(duration_from_env(
97 KEEP_ALIVE_TIMEOUT_VAR,
98 KEEP_ALIVE_TIMEOUT_DEFAULT,
99 ))
100 .keep_alive_while_idle(true)
101 .connect_lazy();
102
103 lock.insert(endpoint.to_string(), channel.clone());
104 Ok(channel)
105}
106
107fn duration_from_env(var_name: &str, default_secs: u64) -> Duration {
108 let seconds = env_var::get(var_name).or(default_secs);
109 Duration::from_secs(seconds)
110}