use async_trait::async_trait;
use pidge_core::TokenStorage;
use crate::auth::jwt;
use crate::auth::token_store::TokenStore;
use crate::auth::tokens::TokenSet;
use crate::error::ClientError;
#[async_trait]
pub trait TokenBackend: Send + Sync {
async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError>;
async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError>;
fn on_access_token(&self, _email: &str, _access_token: &str) {}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct LocalBackend;
impl LocalBackend {
fn storage_for(email: &str) -> TokenStorage {
pidge_core::Config::load()
.ok()
.and_then(|c| c.find(email).map(|a| a.storage))
.unwrap_or_default()
}
}
#[async_trait]
impl TokenBackend for LocalBackend {
async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError> {
TokenStore::load(email, Self::storage_for(email))
}
async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError> {
TokenStore::save(email, tokens, Self::storage_for(email))
}
fn on_access_token(&self, email: &str, access_token: &str) {
let Ok(mut config) = pidge_core::Config::load() else {
return;
};
let Some(existing) = config.find(email).cloned() else {
return;
};
if !existing.tenant_id.is_empty() {
return;
}
let Some(tid) = jwt::extract_tenant_id(access_token) else {
return;
};
if tid.is_empty() {
return;
}
let mut updated = existing;
updated.tenant_id = tid;
config.add_account(updated);
let _ = config.save();
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use chrono::{Duration, Utc};
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
use crate::auth::AuthClient;
#[derive(Default)]
struct MemoryBackend {
tokens: Mutex<HashMap<String, TokenSet>>,
seen_access_tokens: Mutex<Vec<String>>,
}
#[async_trait]
impl TokenBackend for MemoryBackend {
async fn load(&self, email: &str) -> Result<Option<TokenSet>, ClientError> {
Ok(self.tokens.lock().unwrap().get(email).cloned())
}
async fn save(&self, email: &str, tokens: &TokenSet) -> Result<(), ClientError> {
self.tokens
.lock()
.unwrap()
.insert(email.to_string(), tokens.clone());
Ok(())
}
fn on_access_token(&self, _email: &str, access_token: &str) {
self.seen_access_tokens
.lock()
.unwrap()
.push(access_token.to_string());
}
}
#[tokio::test]
async fn custom_backend_receives_refreshed_tokens() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/oauth2/v2.0/token"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"access_token": "NEW_AT",
"refresh_token": "NEW_RT",
"expires_in": 3600
})))
.expect(1)
.mount(&server)
.await;
let backend = Arc::new(MemoryBackend::default());
backend
.save(
"jane@example.com",
&TokenSet {
access_token: "OLD_AT".into(),
refresh_token: "OLD_RT".into(),
expires_at: Utc::now() - Duration::seconds(60),
},
)
.await
.unwrap();
let auth = AuthClient::for_test("cid", server.uri()).with_backend(backend.clone());
let token = auth.get_valid_token("jane@example.com").await.unwrap();
assert_eq!(token, "NEW_AT");
let stored = backend.load("jane@example.com").await.unwrap().unwrap();
assert_eq!(stored.refresh_token, "NEW_RT");
assert_eq!(*backend.seen_access_tokens.lock().unwrap(), vec!["NEW_AT"]);
}
#[tokio::test]
async fn fresh_token_is_returned_without_touching_microsoft() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/oauth2/v2.0/token"))
.respond_with(ResponseTemplate::new(500))
.expect(0)
.mount(&server)
.await;
let backend = Arc::new(MemoryBackend::default());
backend
.save(
"jane@example.com",
&TokenSet {
access_token: "FRESH".into(),
refresh_token: "RT".into(),
expires_at: Utc::now() + Duration::seconds(3600),
},
)
.await
.unwrap();
let auth = AuthClient::for_test("cid", server.uri()).with_backend(backend);
assert_eq!(
auth.get_valid_token("jane@example.com").await.unwrap(),
"FRESH"
);
}
#[tokio::test]
async fn missing_session_maps_to_session_expired() {
let server = MockServer::start().await;
let auth = AuthClient::for_test("cid", server.uri())
.with_backend(Arc::new(MemoryBackend::default()));
let err = auth
.get_valid_token("nobody@example.com")
.await
.unwrap_err();
assert!(matches!(err, ClientError::SessionExpired { .. }));
}
}