Skip to main content

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}