use anyhow::Result;
use pidge_client::ClientError;
use pidge_client::mcp::{
McpRpc, McpTokenStore, McpTokens, StoredServer, ToolResult, normalize_origin,
valid_access_token,
};
use pidge_core::TokenStorage;
pub trait McpCalls {
async fn call_tool(&mut self, name: &str, arguments: serde_json::Value) -> Result<ToolResult>;
}
impl McpCalls for McpRpc {
async fn call_tool(&mut self, name: &str, arguments: serde_json::Value) -> Result<ToolResult> {
let result = McpRpc::call_tool(self, name, arguments).await?;
check_tool_result(name, result)
}
}
fn check_tool_result(name: &str, result: ToolResult) -> Result<ToolResult> {
if result.is_error {
Err(anyhow::anyhow!("{name} reported an error: {}", result.text))
} else {
Ok(result)
}
}
pub(crate) fn candidate_backends(preferred: TokenStorage) -> [TokenStorage; 2] {
match preferred {
TokenStorage::Keychain => [TokenStorage::Keychain, TokenStorage::File],
TokenStorage::File => [TokenStorage::File, TokenStorage::Keychain],
}
}
pub(crate) fn backend_name(store: TokenStorage) -> &'static str {
match store {
TokenStorage::Keychain => "keychain",
TokenStorage::File => "file",
}
}
pub(crate) fn find_first_hit<T, E: std::fmt::Display>(
candidates: &[TokenStorage],
mut load: impl FnMut(TokenStorage) -> Result<Option<T>, E>,
) -> Result<Option<(T, TokenStorage)>, E> {
for (index, &backend) in candidates.iter().enumerate() {
match load(backend) {
Ok(Some(value)) => return Ok(Some((value, backend))),
Ok(None) => continue,
Err(e) if index == 0 => return Err(e),
Err(e) => {
eprintln!(
"warning: could not check the {} backend for a stored session: {e}; trying the next one",
backend_name(backend)
);
continue;
}
}
}
Ok(None)
}
pub(crate) fn try_each_backend<T, E: std::fmt::Display>(
candidates: &[TokenStorage],
mut f: impl FnMut(TokenStorage) -> Result<T, E>,
) -> Result<Vec<(TokenStorage, T)>, E> {
let mut out = Vec::new();
for (index, &backend) in candidates.iter().enumerate() {
match f(backend) {
Ok(value) => out.push((backend, value)),
Err(e) if index == 0 => return Err(e),
Err(e) => {
eprintln!(
"warning: could not reach the {} backend: {e}; skipping it",
backend_name(backend)
);
}
}
}
Ok(out)
}
pub(crate) fn find_stored_tokens(
url: &str,
preferred: TokenStorage,
) -> Result<Option<(McpTokens, TokenStorage)>> {
Ok(find_first_hit(&candidate_backends(preferred), |backend| {
McpTokenStore::load(url, backend)
})?)
}
pub(crate) fn preferred_backend_for(url: &str, servers: &[StoredServer]) -> TokenStorage {
let Ok(origin) = normalize_origin(url) else {
return TokenStorage::Keychain;
};
servers
.iter()
.find(|s| s.server == origin)
.map(|s| s.storage)
.unwrap_or(TokenStorage::Keychain)
}
pub(crate) enum SessionLookup {
Found(McpTokens, TokenStorage),
Expired,
Absent,
}
pub(crate) async fn lookup_session(
http: &reqwest::Client,
url: &str,
preferred: TokenStorage,
) -> Result<SessionLookup> {
let Some((mut tokens, backend)) = find_stored_tokens(url, preferred)? else {
return Ok(SessionLookup::Absent);
};
match refresh_if_needed(http, &mut tokens, backend).await {
Ok(()) => Ok(SessionLookup::Found(tokens, backend)),
Err(ClientError::SessionExpired { .. }) => Ok(SessionLookup::Expired),
Err(e) => Err(e.into()),
}
}
pub(crate) async fn refresh_if_needed(
http: &reqwest::Client,
tokens: &mut McpTokens,
store: TokenStorage,
) -> Result<(), ClientError> {
refresh_if_needed_with(http, tokens, |tokens| McpTokenStore::save(tokens, store)).await
}
async fn refresh_if_needed_with(
http: &reqwest::Client,
tokens: &mut McpTokens,
save: impl FnOnce(&McpTokens) -> Result<(), ClientError>,
) -> Result<(), ClientError> {
let server = tokens.server.clone();
let previous_access_token = tokens.access_token.clone();
let access_token = valid_access_token(http, &server, tokens).await?;
if access_token != previous_access_token {
save(tokens)?;
}
Ok(())
}
pub(crate) struct RefreshingRpc {
http: reqwest::Client,
inner: McpRpc,
tokens: McpTokens,
store: TokenStorage,
}
impl RefreshingRpc {
pub(crate) fn new(
http: reqwest::Client,
inner: McpRpc,
tokens: McpTokens,
store: TokenStorage,
) -> Self {
Self {
http,
inner,
tokens,
store,
}
}
pub(crate) async fn initialize(&mut self) -> Result<()> {
self.inner
.initialize()
.await
.map_err(|e| remap_401_to_session_expired(&self.tokens.server, e))
}
}
pub(crate) fn remap_401_to_session_expired(server: &str, err: ClientError) -> anyhow::Error {
match err {
ClientError::Graph { status: 401, .. } => ClientError::SessionExpired {
email: server.to_string(),
}
.into(),
other => other.into(),
}
}
impl McpCalls for RefreshingRpc {
async fn call_tool(&mut self, name: &str, arguments: serde_json::Value) -> Result<ToolResult> {
let server = self.tokens.server.clone();
refresh_if_needed(&self.http, &mut self.tokens, self.store).await?;
self.inner
.set_access_token(self.tokens.access_token.clone());
match <McpRpc as McpCalls>::call_tool(&mut self.inner, name, arguments).await {
Ok(result) => Ok(result),
Err(e) if is_unauthorized(&e) => {
Err(ClientError::SessionExpired { email: server }.into())
}
Err(e) => Err(e),
}
}
}
fn is_unauthorized(err: &anyhow::Error) -> bool {
matches!(
err.downcast_ref::<ClientError>(),
Some(ClientError::Graph { status: 401, .. })
)
}
pub(crate) fn is_session_expired(err: &anyhow::Error) -> bool {
err.chain().any(|cause| {
matches!(
cause.downcast_ref::<ClientError>(),
Some(ClientError::SessionExpired { .. } | ClientError::McpSessionExpired { .. })
)
})
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
fn tokens_expiring_in(server: &str, seconds: i64) -> McpTokens {
McpTokens {
server: server.to_string(),
access_token: "OLD_AT".into(),
refresh_token: "OLD_RT".into(),
expires_at: chrono::Utc::now() + chrono::Duration::seconds(seconds),
client_id: "CID".into(),
}
}
#[tokio::test]
async fn refresh_if_needed_does_not_save_an_unchanged_token() {
let mut tokens = tokens_expiring_in("http://127.0.0.1:9/mcp", 3600);
let mut saved = false;
refresh_if_needed_with(&reqwest::Client::new(), &mut tokens, |_| {
saved = true;
Ok(())
})
.await
.unwrap();
assert!(!saved, "an unchanged access token must not be written back");
assert_eq!(tokens.access_token, "OLD_AT");
}
#[tokio::test]
async fn refresh_if_needed_saves_a_refreshed_token() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/.well-known/oauth-protected-resource/mcp"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"resource": format!("{}/mcp", server.uri()),
"authorization_servers": [server.uri()],
})))
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/.well-known/oauth-authorization-server"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"authorization_endpoint": format!("{}/authorize", server.uri()),
"token_endpoint": format!("{}/token", server.uri()),
})))
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/token"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"access_token": "NEW_AT",
"refresh_token": "NEW_RT",
"expires_in": 3600
})))
.expect(1)
.mount(&server)
.await;
let mut tokens = tokens_expiring_in(&format!("{}/mcp", server.uri()), -60);
let mut saved: Option<McpTokens> = None;
refresh_if_needed_with(&reqwest::Client::new(), &mut tokens, |t| {
saved = Some(t.clone());
Ok(())
})
.await
.unwrap();
let saved = saved.expect("a refreshed token must be persisted");
assert_eq!(saved.access_token, "NEW_AT");
assert_eq!(saved.refresh_token, "NEW_RT");
assert_eq!(tokens.access_token, "NEW_AT");
}
#[test]
fn check_tool_result_bails_on_is_error() {
let result = ToolResult {
text: "boom".into(),
is_error: true,
};
let err = check_tool_result("accounts_list", result).unwrap_err();
assert!(err.to_string().contains("boom"), "{err}");
}
#[test]
fn check_tool_result_passes_through_ok_results() {
let result = ToolResult {
text: "fine".into(),
is_error: false,
};
assert_eq!(
check_tool_result("accounts_list", result).unwrap().text,
"fine"
);
}
#[test]
fn candidate_backends_tries_the_preferred_backend_first() {
assert_eq!(
candidate_backends(TokenStorage::Keychain),
[TokenStorage::Keychain, TokenStorage::File]
);
assert_eq!(
candidate_backends(TokenStorage::File),
[TokenStorage::File, TokenStorage::Keychain]
);
}
#[test]
fn find_first_hit_returns_the_first_hit_and_does_not_try_the_rest() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let mut tried = Vec::new();
let result = find_first_hit(&candidates, |backend| {
tried.push(backend);
Ok::<_, anyhow::Error>(if backend == TokenStorage::Keychain {
Some(42)
} else {
None
})
})
.unwrap();
assert_eq!(result, Some((42, TokenStorage::Keychain)));
assert_eq!(tried, vec![TokenStorage::Keychain]);
}
#[test]
fn find_first_hit_finds_a_hit_in_the_fallback_backend() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let result = find_first_hit(&candidates, |backend| {
Ok::<_, anyhow::Error>(match backend {
TokenStorage::Keychain => None,
TokenStorage::File => Some("found"),
})
})
.unwrap();
assert_eq!(result, Some(("found", TokenStorage::File)));
}
#[test]
fn find_first_hit_propagates_an_error_from_the_preferred_backend() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let mut tried = Vec::new();
let err = find_first_hit(&candidates, |backend| {
tried.push(backend);
Err::<Option<i32>, _>(anyhow::anyhow!("keychain unavailable"))
})
.unwrap_err();
assert!(err.to_string().contains("keychain unavailable"));
assert_eq!(
tried,
vec![TokenStorage::Keychain],
"must not try the fallback after a preferred-backend error"
);
}
#[test]
fn find_first_hit_treats_a_fallback_backend_error_as_a_miss() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let result = find_first_hit(&candidates, |backend| match backend {
TokenStorage::Keychain => Ok::<Option<i32>, anyhow::Error>(None),
TokenStorage::File => Err(anyhow::anyhow!("no secret service running")),
})
.unwrap();
assert_eq!(
result, None,
"a fallback-backend error must be reported as an overall miss, not fail the lookup"
);
}
#[test]
fn try_each_backend_collects_every_candidates_result() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let result = try_each_backend(&candidates, |backend| {
Ok::<_, anyhow::Error>(backend == TokenStorage::File)
})
.unwrap();
assert_eq!(
result,
vec![(TokenStorage::Keychain, false), (TokenStorage::File, true)]
);
}
#[test]
fn try_each_backend_propagates_an_error_from_the_preferred_backend() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let err = try_each_backend(&candidates, |_| {
Err::<bool, _>(anyhow::anyhow!("keychain unavailable"))
})
.unwrap_err();
assert!(err.to_string().contains("keychain unavailable"));
}
#[test]
fn try_each_backend_skips_a_failing_fallback_backend_instead_of_failing() {
let candidates = [TokenStorage::Keychain, TokenStorage::File];
let result = try_each_backend(&candidates, |backend| match backend {
TokenStorage::Keychain => Ok::<bool, anyhow::Error>(true),
TokenStorage::File => Err(anyhow::anyhow!("no secret service running")),
})
.unwrap();
assert_eq!(result, vec![(TokenStorage::Keychain, true)]);
}
#[test]
fn preferred_backend_for_uses_the_indexed_backend_when_present() {
let servers = vec![StoredServer {
server: "https://mcp.example.com".to_string(),
storage: TokenStorage::File,
}];
assert_eq!(
preferred_backend_for("https://mcp.example.com", &servers),
TokenStorage::File
);
}
#[test]
fn preferred_backend_for_defaults_to_keychain_when_not_indexed() {
assert_eq!(
preferred_backend_for("https://mcp.example.com", &[]),
TokenStorage::Keychain
);
}
#[test]
fn preferred_backend_for_matches_on_normalized_origin_not_the_exact_url() {
let servers = vec![StoredServer {
server: "https://mcp.example.com".to_string(),
storage: TokenStorage::File,
}];
assert_eq!(
preferred_backend_for("https://mcp.example.com/mcp", &servers),
TokenStorage::File
);
}
#[test]
fn remap_401_to_session_expired_converts_a_bare_401() {
let err = remap_401_to_session_expired(
"https://mcp.example.com",
ClientError::Graph {
status: 401,
message: "nope".into(),
},
);
assert!(matches!(
err.downcast_ref::<ClientError>(),
Some(ClientError::SessionExpired { .. })
));
}
#[test]
fn remap_401_to_session_expired_leaves_other_statuses_alone() {
let err = remap_401_to_session_expired(
"https://mcp.example.com",
ClientError::Graph {
status: 500,
message: "boom".into(),
},
);
assert!(matches!(
err.downcast_ref::<ClientError>(),
Some(ClientError::Graph { status: 500, .. })
));
}
#[test]
fn is_session_expired_recognizes_both_error_flavors() {
let a = anyhow::Error::from(ClientError::SessionExpired {
email: "a@b.se".into(),
});
let b = anyhow::Error::from(ClientError::McpSessionExpired {
server: "https://mcp.example.com".into(),
});
let c = anyhow::anyhow!("something unrelated");
assert!(is_session_expired(&a));
assert!(is_session_expired(&b));
assert!(!is_session_expired(&c));
}
}