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 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 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}