Skip to main content

pidge_client/auth/
backend.rs

1//! Pluggable token persistence for [`super::AuthClient`].
2//!
3//! The CLI resolves each account's backend (OS keychain or plaintext file)
4//! from `config.yaml`; that behaviour lives in [`LocalBackend`] and stays the
5//! default. Hosted consumers (e.g. the remote MCP server) implement
6//! [`TokenBackend`] themselves so tokens can live in a per-user secret store
7//! instead; `AuthClient` and `GraphClient` don't care which.
8
9use async_trait::async_trait;
10use pidge_core::TokenStorage;
11
12use crate::auth::jwt;
13use crate::auth::token_store::TokenStore;
14use crate::auth::tokens::TokenSet;
15use crate::error::ClientError;
16
17/// Where an account's [`TokenSet`] is loaded from and saved to.
18///
19/// Implementations must be safe to share across tasks; `AuthClient` holds one
20/// behind an `Arc`.
21#[async_trait]
22pub trait TokenBackend: Send + Sync {
23    /// Load the stored tokens for `email`. `Ok(None)` means "no session".
24    async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError>;
25
26    /// Persist `tokens` for `email`, replacing whatever was stored before.
27    async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError>;
28
29    /// Hook invoked with every access token handed out by
30    /// [`super::AuthClient::get_valid_token`]. The default is a no-op; the
31    /// CLI backend uses it to backfill account metadata.
32    fn on_access_token(&self, _email: &str, _access_token: &str) {}
33}
34
35/// The CLI's backend: consults `config.yaml` for the account's
36/// [`TokenStorage`] and dispatches to the keychain or file store.
37#[derive(Debug, Default, Clone, Copy)]
38pub struct LocalBackend;
39
40impl LocalBackend {
41    /// Resolve the token storage backend for an email by consulting `config.yaml`.
42    /// Falls back to [`TokenStorage::Keychain`] (the default) if the config can't
43    /// be read or the account isn't listed yet.
44    fn storage_for(email: &str) -> TokenStorage {
45        pidge_core::Config::load()
46            .ok()
47            .and_then(|c| c.find(email).map(|a| a.storage))
48            .unwrap_or_default()
49    }
50}
51
52#[async_trait]
53impl TokenBackend for LocalBackend {
54    async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError> {
55        TokenStore::load(email, Self::storage_for(email))
56    }
57
58    async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError> {
59        TokenStore::save(email, tokens, Self::storage_for(email))
60    }
61
62    /// Opportunistic backfill: accounts added before pidge requested the
63    /// `openid` scope have an empty tenant_id in config. Microsoft Graph
64    /// access tokens are JWTs that carry the `tid` claim, so we can fix
65    /// this once per such account on the next Graph call without any
66    /// user action. Silent on any failure: cosmetic, not a correctness
67    /// requirement.
68    fn on_access_token(&self, email: &str, access_token: &str) {
69        let Ok(mut config) = pidge_core::Config::load() else {
70            return;
71        };
72        let Some(existing) = config.find(email).cloned() else {
73            return;
74        };
75        if !existing.tenant_id.is_empty() {
76            return;
77        }
78        let Some(tid) = jwt::extract_tenant_id(access_token) else {
79            return;
80        };
81        if tid.is_empty() {
82            return;
83        }
84        let mut updated = existing;
85        updated.tenant_id = tid;
86        config.add_account(updated);
87        let _ = config.save();
88    }
89}
90
91#[cfg(test)]
92mod tests {
93    use std::collections::HashMap;
94    use std::sync::{Arc, Mutex};
95
96    use chrono::{Duration, Utc};
97    use wiremock::matchers::{method, path};
98    use wiremock::{Mock, MockServer, ResponseTemplate};
99
100    use super::*;
101    use crate::auth::AuthClient;
102
103    /// The kind of backend a hosted consumer would write: tokens live in a
104    /// map instead of the keychain, and the refresh path must round-trip
105    /// through it.
106    #[derive(Default)]
107    struct MemoryBackend {
108        tokens: Mutex<HashMap<String, TokenSet>>,
109        seen_access_tokens: Mutex<Vec<String>>,
110    }
111
112    #[async_trait]
113    impl TokenBackend for MemoryBackend {
114        async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError> {
115            Ok(self.tokens.lock().unwrap().get(email).cloned())
116        }
117
118        async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError> {
119            self.tokens
120                .lock()
121                .unwrap()
122                .insert(email.to_string(), tokens.clone());
123            Ok(())
124        }
125
126        fn on_access_token(&self, _email: &str, access_token: &str) {
127            self.seen_access_tokens
128                .lock()
129                .unwrap()
130                .push(access_token.to_string());
131        }
132    }
133
134    #[tokio::test]
135    async fn custom_backend_receives_refreshed_tokens() {
136        let server = MockServer::start().await;
137        Mock::given(method("POST"))
138            .and(path("/oauth2/v2.0/token"))
139            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
140                "access_token": "NEW_AT",
141                "refresh_token": "NEW_RT",
142                "expires_in": 3600
143            })))
144            .expect(1)
145            .mount(&server)
146            .await;
147
148        let backend = Arc::new(MemoryBackend::default());
149        backend
150            .save(
151                "jane@example.com",
152                &TokenSet {
153                    access_token: "OLD_AT".into(),
154                    refresh_token: "OLD_RT".into(),
155                    expires_at: Utc::now() - Duration::seconds(60),
156                },
157            )
158            .await
159            .unwrap();
160
161        let auth = AuthClient::for_test("cid", server.uri()).with_backend(backend.clone());
162        let token = auth.get_valid_token("jane@example.com").await.unwrap();
163
164        assert_eq!(token, "NEW_AT");
165        let stored = backend.load("jane@example.com").await.unwrap().unwrap();
166        assert_eq!(stored.refresh_token, "NEW_RT");
167        assert_eq!(*backend.seen_access_tokens.lock().unwrap(), vec!["NEW_AT"]);
168    }
169
170    #[tokio::test]
171    async fn fresh_token_is_returned_without_touching_microsoft() {
172        let server = MockServer::start().await;
173        Mock::given(method("POST"))
174            .and(path("/oauth2/v2.0/token"))
175            .respond_with(ResponseTemplate::new(500))
176            .expect(0)
177            .mount(&server)
178            .await;
179
180        let backend = Arc::new(MemoryBackend::default());
181        backend
182            .save(
183                "jane@example.com",
184                &TokenSet {
185                    access_token: "FRESH".into(),
186                    refresh_token: "RT".into(),
187                    expires_at: Utc::now() + Duration::seconds(3600),
188                },
189            )
190            .await
191            .unwrap();
192
193        let auth = AuthClient::for_test("cid", server.uri()).with_backend(backend);
194        assert_eq!(
195            auth.get_valid_token("jane@example.com").await.unwrap(),
196            "FRESH"
197        );
198    }
199
200    #[tokio::test]
201    async fn missing_session_maps_to_session_expired() {
202        let server = MockServer::start().await;
203        let auth = AuthClient::for_test("cid", server.uri())
204            .with_backend(Arc::new(MemoryBackend::default()));
205        let err = auth
206            .get_valid_token("nobody@example.com")
207            .await
208            .unwrap_err();
209        assert!(matches!(err, ClientError::SessionExpired { .. }));
210    }
211}