xmtp_api_backend/middleware/
auth.rs1use crate::endpoints::backend::GET_CONFIGURATION_PATH;
2use arc_swap::ArcSwap;
3use prost::bytes::Bytes;
4use std::sync::Arc;
5use tokio::sync::OnceCell;
6use xmtp_common::{BoxDynError, MaybeSend, MaybeSync, time::Instant};
7#[cfg(not(test))]
8use xmtp_configuration::AUTH_LOCKOUT_COOLDOWN;
9use xmtp_configuration::MAX_CONSECUTIVE_AUTH_FAILURES;
10use xmtp_proto::api::{
11 ApiClientError, AuthError, BytesStream, Client, IsConnectedCheck, grpc_status,
12};
13
14#[cfg(test)]
15const AUTH_LOCKOUT_COOLDOWN: std::time::Duration = std::time::Duration::from_millis(500);
18
19#[cfg(not(test))]
20use xmtp_common::time::now_secs;
21#[cfg(test)]
23fn now_secs() -> i64 {
24 1_000_000
25}
26
27#[derive(Clone, Debug, PartialEq, Eq)]
28pub struct Credential {
29 name: http::header::HeaderName,
30 value: http::header::HeaderValue,
31 expires_at_seconds: i64,
32}
33
34impl Credential {
35 pub fn new(
36 name: Option<http::header::HeaderName>,
37 value: http::header::HeaderValue,
38 expires_at_seconds: i64,
39 ) -> Self {
40 Self {
41 name: name.unwrap_or(http::header::AUTHORIZATION),
42 value,
43 expires_at_seconds,
44 }
45 }
46}
47
48#[derive(Clone, Copy, Default)]
49struct AuthState {
50 generation: u64,
51 stale: bool,
52 failures: u32,
53 locked_until: Option<Instant>,
54}
55
56impl AuthState {
57 fn fail(&mut self) {
58 self.failures = (self.failures + 1).min(MAX_CONSECUTIVE_AUTH_FAILURES);
59 if self.failures == MAX_CONSECUTIVE_AUTH_FAILURES && self.locked_until.is_none() {
60 self.locked_until = Some(Instant::now() + AUTH_LOCKOUT_COOLDOWN);
61 }
62 }
63}
64
65#[derive(Default)]
66struct AuthInner {
67 current: OnceCell<ArcSwap<Credential>>,
68 state: tokio::sync::Mutex<AuthState>,
69 refresh: tokio::sync::Mutex<()>,
74}
75
76impl AuthInner {
77 fn store(&self, credential: Credential) {
79 if let Some(current) = self.current.get() {
80 current.store(Arc::new(credential));
81 } else {
82 self.current
83 .set(ArcSwap::from_pointee(credential))
84 .unwrap_or_else(|_| unreachable!("state lock protects initialization"));
85 }
86 }
87}
88
89#[derive(Default, Clone)]
90pub struct AuthHandle {
91 inner: Arc<AuthInner>,
92}
93
94impl AuthHandle {
95 pub fn new() -> Self {
96 Self::default()
97 }
98
99 pub async fn set(&self, credential: Credential) {
100 let mut state = self.inner.state.lock().await;
101 self.inner.store(credential);
102 *state = AuthState {
103 generation: state.generation + 1,
104 ..AuthState::default()
105 };
106 }
107
108 pub fn id(&self) -> usize {
109 Arc::as_ptr(&self.inner) as usize
110 }
111}
112
113#[xmtp_common::async_trait]
114pub trait AuthCallback: MaybeSend + MaybeSync {
115 async fn on_auth_required(&self) -> Result<Credential, BoxDynError>;
116}
117
118#[derive(Clone)]
122pub struct AuthMiddleware<C> {
123 inner: C,
124 handle: AuthHandle,
125 callback: Option<Arc<dyn AuthCallback>>,
126}
127
128impl<C> AuthMiddleware<C> {
129 #[track_caller]
130 pub fn new(
131 inner: C,
132 callback: Option<Arc<dyn AuthCallback>>,
133 handle: Option<AuthHandle>,
134 ) -> Self {
135 assert!(
136 callback.is_some() || handle.is_some(),
137 "Either a callback or a handle must be provided"
138 );
139 Self {
140 inner,
141 handle: handle.unwrap_or_default(),
142 callback,
143 }
144 }
145
146 async fn get_credential(&self) -> Result<(Arc<Credential>, u64), AuthError> {
149 let inner = &self.handle.inner;
150 let _refresh = inner.refresh.lock().await;
154 let mut state = *inner.state.lock().await;
155 if let Some(until) = state.locked_until {
156 if until > Instant::now() {
157 return Err(AuthError::Exhausted);
158 }
159 state.locked_until = None;
160 state.failures = MAX_CONSECUTIVE_AUTH_FAILURES - 1;
161 state.stale = true;
162 }
163 let mut credential = inner.current.get().map(|current| current.load_full());
164 let needs_refresh = state.stale
165 || credential
166 .as_ref()
167 .is_none_or(|credential| credential.expires_at_seconds <= now_secs());
168 if needs_refresh && let Some(callback) = &self.callback {
169 let generation_before = state.generation;
170 let result = callback.on_auth_required().await;
171 let mut guard = inner.state.lock().await;
172 if guard.generation != generation_before {
175 return Ok((
176 inner
177 .current
178 .get()
179 .map(|current| current.load_full())
180 .ok_or(AuthError::MissingCredential)?,
181 guard.generation,
182 ));
183 }
184 match result {
185 Ok(new) => {
186 inner.store(new);
187 credential = inner.current.get().map(|current| current.load_full());
188 state.generation += 1;
189 state.stale = false;
190 }
191 Err(_) => {
192 state.fail();
193 *guard = state;
194 return Err(if state.locked_until.is_some() {
198 AuthError::Exhausted
199 } else {
200 AuthError::CallbackFailed { retryable: true }
201 });
202 }
203 }
204 *guard = state;
205 } else {
206 *inner.state.lock().await = state;
207 }
208 let credential = credential.ok_or(AuthError::MissingCredential)?;
209 Ok((credential, state.generation))
210 }
211
212 async fn finish<T>(
214 &self,
215 generation: u64,
216 result: Result<T, ApiClientError>,
217 ) -> (Result<T, ApiClientError>, bool) {
218 let rejected = result
219 .as_ref()
220 .err()
221 .and_then(|error| grpc_status(error))
222 .is_some_and(|status| status.code() == tonic::Code::Unauthenticated);
223 let mut state = self.handle.inner.state.lock().await;
224 let current = generation == state.generation;
225 if current {
226 if result.is_ok() {
227 state.failures = 0;
228 } else if rejected {
229 state.stale = true;
230 if self.callback.is_some() {
231 state.fail();
232 }
233 }
234 }
235 if rejected {
236 if self.callback.is_some()
241 && state
242 .locked_until
243 .is_some_and(|until| until > Instant::now())
244 {
245 return (Err(AuthError::Exhausted.into()), false);
246 }
247 let retryable = self.callback.is_some();
248 (
249 Err(AuthError::CredentialRejected { retryable }.into()),
250 current && retryable,
251 )
252 } else {
253 (result, false)
254 }
255 }
256
257 fn request_builder(
259 parts: &http::request::Parts,
260 credential: &Credential,
261 ) -> http::request::Builder {
262 let mut request = http::Request::builder()
263 .method(parts.method.clone())
264 .uri(parts.uri.clone())
265 .version(parts.version);
266 let headers = request.headers_mut().expect("validated request parts");
267 *headers = parts.headers.clone();
268 headers.insert(credential.name.clone(), credential.value.clone());
269 *request.extensions_mut().expect("validated request parts") = parts.extensions.clone();
270 request
271 }
272}
273
274#[xmtp_common::async_trait]
275impl<C: Client> Client for AuthMiddleware<C> {
276 fn has_credential_source(&self) -> bool {
279 true
280 }
281
282 fn host(&self) -> &str {
283 self.inner.host()
284 }
285
286 async fn request(
287 &self,
288 request: http::request::Builder,
289 path: http::uri::PathAndQuery,
290 body: Bytes,
291 ) -> Result<http::Response<Bytes>, ApiClientError> {
292 if path.path() == GET_CONFIGURATION_PATH {
296 return self.inner.request(request, path, body).await;
297 }
298 let (parts, ()) = request.body(())?.into_parts();
299 let (credential, generation) = self.get_credential().await?;
300 let result = self
301 .inner
302 .request(
303 Self::request_builder(&parts, &credential),
304 path.clone(),
305 body.clone(),
306 )
307 .await;
308 let (result, replay) = self.finish(generation, result).await;
309 if !replay {
310 return result;
311 }
312 let (credential, generation) = self.get_credential().await?;
313 let result = self
314 .inner
315 .request(Self::request_builder(&parts, &credential), path, body)
316 .await;
317 self.finish(generation, result).await.0
318 }
319
320 async fn stream(
321 &self,
322 request: http::request::Builder,
323 path: http::uri::PathAndQuery,
324 body: Bytes,
325 ) -> Result<http::Response<BytesStream>, ApiClientError> {
326 let (parts, ()) = request.body(())?.into_parts();
327 let (credential, generation) = self.get_credential().await?;
328 let result = self
329 .inner
330 .stream(
331 Self::request_builder(&parts, &credential),
332 path.clone(),
333 body.clone(),
334 )
335 .await;
336 let (result, replay) = self.finish(generation, result).await;
337 if !replay {
338 return result;
339 }
340 let (credential, generation) = self.get_credential().await?;
341 let result = self
342 .inner
343 .stream(Self::request_builder(&parts, &credential), path, body)
344 .await;
345 self.finish(generation, result).await.0
346 }
347
348 async fn bidi_stream(
349 &self,
350 request: http::request::Builder,
351 path: http::uri::PathAndQuery,
352 body: xmtp_common::BoxDynStream<'static, Bytes>,
353 ) -> Result<http::Response<BytesStream>, ApiClientError> {
354 let (parts, ()) = request.body(())?.into_parts();
355 let (credential, generation) = self.get_credential().await?;
356 let result = self
357 .inner
358 .bidi_stream(Self::request_builder(&parts, &credential), path, body)
359 .await;
360 self.finish(generation, result).await.0
361 }
362}
363
364#[xmtp_common::async_trait]
365impl<C: IsConnectedCheck> IsConnectedCheck for AuthMiddleware<C> {
366 async fn is_connected(&self) -> bool {
367 self.inner.is_connected().await
368 }
369}
370
371#[cfg(test)]
372mod tests;
373
374#[cfg(test)]
375mod refetch_tests;