pidge_client/auth/
backend.rs1use 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#[async_trait]
22pub trait TokenBackend: Send + Sync {
23 async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError>;
25
26 async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError>;
28
29 fn on_access_token(&self, _email: &str, _access_token: &str) {}
33}
34
35#[derive(Debug, Default, Clone, Copy)]
38pub struct LocalBackend;
39
40impl LocalBackend {
41 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 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 #[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}