Skip to main content

gcloud_gax/
conn.rs

1use std::fmt::Debug;
2use std::future::Future;
3use std::pin::Pin;
4use std::sync::atomic::{AtomicUsize, Ordering};
5use std::sync::Arc;
6use std::time::Duration;
7
8use http::header::AUTHORIZATION;
9use http::{HeaderValue, Request};
10use tonic::body::Body;
11use tonic::transport::{Channel as TonicChannel, ClientTlsConfig, Endpoint};
12use tonic::{Code, Status};
13use tower::filter::{AsyncFilter, AsyncPredicate};
14use tower::{BoxError, ServiceBuilder};
15
16use token_source::{TokenSource, TokenSourceProvider};
17
18pub type Channel = AsyncFilter<TonicChannel, AsyncAuthInterceptor>;
19
20#[derive(Clone, Debug)]
21pub struct AsyncAuthInterceptor {
22    token_source: Option<Arc<dyn TokenSource>>,
23}
24
25impl AsyncAuthInterceptor {
26    fn new(token_source: Arc<dyn TokenSource>) -> Self {
27        Self {
28            token_source: Some(token_source),
29        }
30    }
31    fn empty() -> Self {
32        Self { token_source: None }
33    }
34}
35
36impl AsyncPredicate<Request<Body>> for AsyncAuthInterceptor {
37    type Future = Pin<Box<dyn Future<Output = Result<Self::Request, BoxError>> + Send>>;
38    type Request = Request<Body>;
39
40    fn check(&mut self, request: Request<Body>) -> Self::Future {
41        let ts = match &self.token_source {
42            Some(ts) => ts.clone(),
43            None => return Box::pin(async move { Ok(request) }),
44        };
45        Box::pin(async move {
46            let token = ts
47                .token()
48                .await
49                .map_err(|e| Status::new(Code::Unauthenticated, format!("token error: {e:?}")))?;
50            let token_header = HeaderValue::from_str(token.as_str())
51                .map_err(|e| Status::new(Code::Unauthenticated, format!("token error: {e:?}")))?;
52            let (mut parts, body) = request.into_parts();
53            parts.headers.insert(AUTHORIZATION, token_header);
54            Ok(Request::from_parts(parts, body))
55        })
56    }
57}
58
59#[derive(thiserror::Error, Debug)]
60pub enum Error {
61    #[error(transparent)]
62    Auth(#[from] Box<dyn std::error::Error + Send + Sync>),
63
64    #[error("tonic error : {0}")]
65    TonicTransport(#[from] tonic::transport::Error),
66
67    #[error("invalid emulator host: {0}")]
68    InvalidEmulatorHOST(String),
69}
70
71#[derive(Debug)]
72pub enum Environment {
73    Emulator(String),
74    GoogleCloud(Box<dyn TokenSourceProvider>),
75}
76
77#[derive(Debug)]
78struct AtomicRing<T>
79where
80    T: Clone + Debug,
81{
82    index: AtomicUsize,
83    values: Vec<T>,
84}
85
86impl<T> AtomicRing<T>
87where
88    T: Clone + Debug,
89{
90    fn next(&self) -> T {
91        let current = self.index.fetch_add(1, Ordering::SeqCst);
92        //clone() reuses http/2 connection
93        self.values[current % self.values.len()].clone()
94    }
95}
96
97#[derive(Debug, Clone, Default)]
98pub struct ConnectionOptions {
99    pub timeout: Option<Duration>,
100    pub connect_timeout: Option<Duration>,
101    pub http2_keep_alive_interval: Option<Duration>,
102    pub keep_alive_timeout: Option<Duration>,
103    pub keep_alive_while_idle: Option<bool>,
104}
105
106impl ConnectionOptions {
107    fn apply(&self, mut endpoint: Endpoint) -> Endpoint {
108        endpoint = match self.timeout {
109            Some(t) => endpoint.timeout(t),
110            None => endpoint,
111        };
112        endpoint = match self.connect_timeout {
113            Some(t) => endpoint.connect_timeout(t),
114            None => endpoint,
115        };
116        endpoint = match self.http2_keep_alive_interval {
117            Some(d) => endpoint.http2_keep_alive_interval(d),
118            None => endpoint,
119        };
120        endpoint = match self.keep_alive_timeout {
121            Some(d) => endpoint.keep_alive_timeout(d),
122            None => endpoint,
123        };
124        endpoint = match self.keep_alive_while_idle {
125            Some(v) => endpoint.keep_alive_while_idle(v),
126            None => endpoint,
127        };
128        endpoint
129    }
130}
131
132#[derive(Debug)]
133pub struct ConnectionManager {
134    inner: AtomicRing<Channel>,
135}
136
137impl<'a> ConnectionManager {
138    pub async fn new(
139        pool_size: usize,
140        domain_name: impl Into<String>,
141        audience: &str,
142        environment: &Environment,
143        conn_options: &'a ConnectionOptions,
144    ) -> Result<Self, Error> {
145        let conns = match environment {
146            Environment::GoogleCloud(ts_provider) => {
147                Self::create_connections(pool_size, domain_name, audience, ts_provider.as_ref(), conn_options).await?
148            }
149            Environment::Emulator(host) => Self::create_emulator_connections(host, conn_options).await?,
150        };
151        Ok(Self {
152            inner: AtomicRing {
153                index: AtomicUsize::new(0),
154                values: conns,
155            },
156        })
157    }
158
159    async fn create_connections(
160        pool_size: usize,
161        domain_name: impl Into<String>,
162        audience: &str,
163        ts_provider: &dyn TokenSourceProvider,
164        conn_options: &'a ConnectionOptions,
165    ) -> Result<Vec<Channel>, Error> {
166        let pool_size = Self::get_pool_size(pool_size);
167
168        let tls_config = ClientTlsConfig::new().with_webpki_roots().domain_name(domain_name);
169        let mut conns = Vec::with_capacity(pool_size);
170
171        let ts = ts_provider.token_source();
172
173        for _i_ in 0..pool_size {
174            let endpoint = TonicChannel::from_shared(audience.to_string().into_bytes())
175                .map_err(|e| Error::InvalidEmulatorHOST(e.to_string()))?
176                .tls_config(tls_config.clone())?;
177            let endpoint = conn_options.apply(endpoint);
178
179            let con = Self::connect(endpoint).await?;
180            // use GCP token per call
181            let auth_filter = AsyncAuthInterceptor::new(Arc::clone(&ts));
182            let auth_con = ServiceBuilder::new().filter_async(auth_filter).service(con);
183            conns.push(auth_con);
184        }
185        Ok(conns)
186    }
187
188    async fn create_emulator_connections(
189        host: &str,
190        conn_options: &'a ConnectionOptions,
191    ) -> Result<Vec<Channel>, Error> {
192        let mut conns = Vec::with_capacity(1);
193        let endpoint = TonicChannel::from_shared(format!("http://{host}").into_bytes())
194            .map_err(|_| Error::InvalidEmulatorHOST(host.to_string()))?;
195        let endpoint = conn_options.apply(endpoint);
196
197        let con = Self::connect(endpoint).await?;
198        let auth_filter = AsyncAuthInterceptor::empty();
199        let auth_con = ServiceBuilder::new().filter_async(auth_filter).service(con);
200        conns.push(auth_con);
201        Ok(conns)
202    }
203
204    async fn connect(endpoint: Endpoint) -> Result<TonicChannel, tonic::transport::Error> {
205        let channel = endpoint.connect().await?;
206        Ok(channel)
207    }
208
209    fn get_pool_size(pool_size: usize) -> usize {
210        pool_size.max(1)
211    }
212
213    pub fn num(&self) -> usize {
214        self.inner.values.len()
215    }
216
217    pub fn conn(&self) -> Channel {
218        self.inner.next()
219    }
220}
221
222#[cfg(test)]
223mod test {
224    use std::collections::HashSet;
225    use std::sync::atomic::{AtomicUsize, Ordering};
226
227    use crate::conn::{AtomicRing, ConnectionManager};
228
229    #[test]
230    fn test_atomic_ring() {
231        let cm = AtomicRing::<&str> {
232            index: AtomicUsize::new(usize::MAX - 1),
233            values: vec!["a", "b", "c", "d"],
234        };
235        let mut values = HashSet::new();
236        assert_eq!(usize::MAX - 1, cm.index.load(Ordering::SeqCst));
237        assert!(values.insert(cm.next()));
238        assert_eq!(usize::MAX, cm.index.load(Ordering::SeqCst));
239        assert!(values.insert(cm.next()));
240        assert_eq!(0, cm.index.load(Ordering::SeqCst));
241        assert!(values.insert(cm.next()));
242        assert_eq!(1, cm.index.load(Ordering::SeqCst));
243        assert!(values.insert(cm.next()));
244        assert_eq!(2, cm.index.load(Ordering::SeqCst));
245        assert!(!values.insert(cm.next()));
246        assert_eq!(3, cm.index.load(Ordering::SeqCst));
247    }
248
249    #[test]
250    fn test_get_pool_size() {
251        assert_eq!(1, ConnectionManager::get_pool_size(0));
252        assert_eq!(1, ConnectionManager::get_pool_size(1));
253        assert_eq!(2, ConnectionManager::get_pool_size(2));
254    }
255}