Skip to main content

xmtp_api_backend/middleware/
auth.rs

1use 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)]
15/// Longer than the transport's first reconnect delay (100 ms), so a reopen
16/// during a lockout reliably meets it instead of racing it.
17const AUTH_LOCKOUT_COOLDOWN: std::time::Duration = std::time::Duration::from_millis(500);
18
19#[cfg(not(test))]
20use xmtp_common::time::now_secs;
21// Use a fixed clock to keep expiry tests stable.
22#[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    /// Held across the callback so only one runs at a time. It is separate from
70    /// `state` because `AuthHandle::set` takes `state`, and a callback that
71    /// pushes its credential through the handle would deadlock on a lock this
72    /// function held across the await. Tokio's mutex is not reentrant.
73    refresh: tokio::sync::Mutex<()>,
74}
75
76impl AuthInner {
77    /// Store only while the state lock is held. This operation cannot be cancelled.
78    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/// Add credentials and refetch after expiry or rejection.
119/// Shared handles share credentials, callback serialization, and lockout state.
120/// Without a callback, expired or rejected credentials remain in use until `set`.
121#[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    /// Keep the credential and its generation together across each network call.
147    /// Commit callback state only after the callback completes, including a probe.
148    async fn get_credential(&self) -> Result<(Arc<Credential>, u64), AuthError> {
149        let inner = &self.handle.inner;
150        // Take the refresh lock first so only one callback runs at a time. The
151        // state lock is taken and released around it, never held across the
152        // await, because `AuthHandle::set` needs it while a callback runs.
153        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            // `AuthHandle::set` may have run during the callback. Its credential
173            // is newer, so keep it and leave its state alone.
174            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                    // A callback failure that trips the lockout reports
195                    // `Exhausted`, like a rejection that trips it. The wire must
196                    // wait for the cool-down, not read this as permanent.
197                    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    /// Return whether a current rejection permits one immediate replay.
213    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            // A rejection that trips the lockout reports `Exhausted`, the same
237            // error every later call gets, so the cool-down has one error. A
238            // long-lived transport can then tell a timed lockout apart from a
239            // permanent rejection and wait instead of shutting down.
240            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    /// Rebuild all request parts. Replace the credential header instead of appending it.
258    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    /// CFG-062: this middleware exists only when a callback or a handle was
277    /// configured, so its presence in the stack is the credential source.
278    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        // CFG-045: the configuration read carries no credential and never
293        // invokes the app's callback. A client asks what the deployment
294        // requires before it can know whether it needs one.
295        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;