use std::time::Duration;
use chrono::Utc;
use serde::Deserialize;
use serde_json::json;
use tokio::net::TcpListener;
use url::Url;
use crate::auth::browser_flow::{
CallbackParams, make_code_challenge, make_code_verifier, make_random, wait_for_callback,
};
use crate::error::ClientError;
use super::McpTokens;
const CALLBACK_TIMEOUT: Duration = Duration::from_secs(600);
pub const PROTOCOL_VERSION: &str = "2025-06-18";
#[derive(Debug, Clone)]
pub struct Discovery {
pub resource: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub registration_endpoint: Option<String>,
}
#[derive(Debug, Deserialize)]
struct ProtectedResourceMetadata {
resource: String,
#[serde(default)]
authorization_servers: Vec<String>,
}
#[derive(Debug, Deserialize)]
struct AuthorizationServerMetadata {
authorization_endpoint: String,
token_endpoint: String,
#[serde(default)]
registration_endpoint: Option<String>,
}
pub async fn discover(http: &reqwest::Client, mcp_url: &str) -> Result<Discovery, ClientError> {
let origin = super::normalize_origin(mcp_url)?;
let suffixed = format!("{origin}/.well-known/oauth-protected-resource/mcp");
let bare = format!("{origin}/.well-known/oauth-protected-resource");
let metadata: ProtectedResourceMetadata = match fetch_json(http, &suffixed).await {
Ok(m) => m,
Err(_) => fetch_json(http, &bare).await?,
};
let as_origin = metadata
.authorization_servers
.first()
.ok_or_else(|| ClientError::Graph {
status: 502,
message: "protected resource metadata has no authorization_servers".to_string(),
})?
.trim_end_matches('/');
let as_metadata: AuthorizationServerMetadata = fetch_json(
http,
&format!("{as_origin}/.well-known/oauth-authorization-server"),
)
.await?;
Ok(Discovery {
resource: metadata.resource,
authorization_endpoint: as_metadata.authorization_endpoint,
token_endpoint: as_metadata.token_endpoint,
registration_endpoint: as_metadata.registration_endpoint,
})
}
async fn fetch_json<T: serde::de::DeserializeOwned>(
http: &reqwest::Client,
url: &str,
) -> Result<T, ClientError> {
let resp = http.get(url).send().await?;
let status = resp.status();
let bytes = resp.bytes().await?;
if !status.is_success() {
return Err(ClientError::Graph {
status: status.as_u16(),
message: String::from_utf8_lossy(&bytes).into_owned(),
});
}
Ok(serde_json::from_slice(&bytes)?)
}
pub async fn register(
http: &reqwest::Client,
discovery: &Discovery,
redirect_uri: &str,
client_name: &str,
) -> Result<String, ClientError> {
let Some(endpoint) = discovery.registration_endpoint.as_deref() else {
return Err(ClientError::Graph {
status: 400,
message: "server does not support dynamic client registration".to_string(),
});
};
let body = json!({
"redirect_uris": [redirect_uri],
"client_name": client_name,
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
});
let resp = http.post(endpoint).json(&body).send().await?;
let status = resp.status();
let bytes = resp.bytes().await?;
if !status.is_success() {
return Err(ClientError::Graph {
status: status.as_u16(),
message: String::from_utf8_lossy(&bytes).into_owned(),
});
}
let value: serde_json::Value = serde_json::from_slice(&bytes)?;
value
.get("client_id")
.and_then(|v| v.as_str())
.map(str::to_string)
.ok_or_else(|| ClientError::Graph {
status: 502,
message: "registration response missing client_id".to_string(),
})
}
pub async fn sign_in<F: FnOnce(&str)>(
http: &reqwest::Client,
mcp_url: &str,
client_name: &str,
on_open: F,
) -> Result<McpTokens, ClientError> {
let discovery = discover(http, mcp_url).await?;
let listener = TcpListener::bind("127.0.0.1:0")
.await
.map_err(ClientError::Io)?;
let port = listener.local_addr().map_err(ClientError::Io)?.port();
let redirect_uri = format!("http://localhost:{port}");
let client_id = register(http, &discovery, &redirect_uri, client_name).await?;
let verifier = make_code_verifier();
let challenge = make_code_challenge(&verifier);
let state = make_random(32);
let authorize_url = build_authorize_url(
&discovery.authorization_endpoint,
&client_id,
&redirect_uri,
&challenge,
&state,
&discovery.resource,
)?;
on_open(&authorize_url);
let CallbackParams {
code,
state: returned_state,
} = wait_for_callback(listener, CALLBACK_TIMEOUT, "MCP sign-in").await?;
if returned_state != state {
return Err(ClientError::Graph {
status: 400,
message: "OAuth state mismatch: possible CSRF or stale request".to_string(),
});
}
exchange_code(
http,
&discovery,
&client_id,
&code,
&verifier,
&redirect_uri,
&discovery.resource,
)
.await
}
fn build_authorize_url(
authorization_endpoint: &str,
client_id: &str,
redirect_uri: &str,
challenge: &str,
state: &str,
resource: &str,
) -> Result<String, ClientError> {
let mut url = Url::parse(authorization_endpoint).map_err(|e| ClientError::Graph {
status: 502,
message: format!("invalid authorization_endpoint: {e}"),
})?;
url.query_pairs_mut()
.append_pair("response_type", "code")
.append_pair("client_id", client_id)
.append_pair("redirect_uri", redirect_uri)
.append_pair("state", state)
.append_pair("code_challenge", challenge)
.append_pair("code_challenge_method", "S256")
.append_pair("resource", resource);
Ok(url.into())
}
#[derive(Deserialize)]
struct TokenResponse {
access_token: String,
refresh_token: Option<String>,
expires_in: Option<i64>,
}
#[derive(Debug, Deserialize)]
struct ErrorResponse {
error: String,
#[serde(default)]
error_description: Option<String>,
}
async fn exchange_code(
http: &reqwest::Client,
discovery: &Discovery,
client_id: &str,
code: &str,
code_verifier: &str,
redirect_uri: &str,
resource: &str,
) -> Result<McpTokens, ClientError> {
let params = [
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", redirect_uri),
("client_id", client_id),
("code_verifier", code_verifier),
("resource", resource),
];
let resp = http
.post(&discovery.token_endpoint)
.form(¶ms)
.send()
.await?;
let status = resp.status();
let bytes = resp.bytes().await?;
if !status.is_success() {
return Err(ClientError::Graph {
status: status.as_u16(),
message: String::from_utf8_lossy(&bytes).into_owned(),
});
}
let tr: TokenResponse = serde_json::from_slice(&bytes)?;
Ok(McpTokens {
server: resource.to_string(),
access_token: tr.access_token,
refresh_token: tr.refresh_token.unwrap_or_default(),
expires_at: Utc::now() + chrono::Duration::seconds(tr.expires_in.unwrap_or(3600)),
client_id: client_id.to_string(),
})
}
pub async fn refresh(
http: &reqwest::Client,
discovery: &Discovery,
current: &McpTokens,
) -> Result<McpTokens, ClientError> {
let params = [
("grant_type", "refresh_token"),
("refresh_token", current.refresh_token.as_str()),
("client_id", current.client_id.as_str()),
("resource", current.server.as_str()),
];
let resp = http
.post(&discovery.token_endpoint)
.form(¶ms)
.send()
.await?;
let status = resp.status();
let bytes = resp.bytes().await?;
if !status.is_success() {
let err: ErrorResponse =
serde_json::from_slice(&bytes).map_err(|_| ClientError::Graph {
status: status.as_u16(),
message: String::from_utf8_lossy(&bytes).into_owned(),
})?;
if matches!(err.error.as_str(), "invalid_grant" | "invalid_client") {
return Err(ClientError::SessionExpired {
email: current.server.clone(),
});
}
return Err(ClientError::Graph {
status: status.as_u16(),
message: err.error_description.unwrap_or(err.error),
});
}
let tr: TokenResponse = serde_json::from_slice(&bytes)?;
Ok(McpTokens {
server: current.server.clone(),
access_token: tr.access_token,
refresh_token: tr
.refresh_token
.unwrap_or_else(|| current.refresh_token.clone()),
expires_at: Utc::now() + chrono::Duration::seconds(tr.expires_in.unwrap_or(3600)),
client_id: current.client_id.clone(),
})
}
pub async fn valid_access_token(
http: &reqwest::Client,
mcp_url: &str,
tokens: &mut McpTokens,
) -> Result<String, ClientError> {
if !tokens.needs_refresh() {
return Ok(tokens.access_token.clone());
}
let discovery = discover(http, mcp_url).await?;
let refreshed = refresh(http, &discovery, tokens).await?;
*tokens = refreshed;
Ok(tokens.access_token.clone())
}
#[cfg(test)]
mod tests {
use chrono::Duration as ChronoDuration;
use wiremock::matchers::{body_partial_json, body_string_contains, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
#[test]
fn callback_timeout_is_at_least_ten_minutes() {
assert!(CALLBACK_TIMEOUT >= Duration::from_secs(600));
}
fn discovery_for(server: &MockServer) -> Discovery {
Discovery {
resource: format!("{}/mcp", server.uri()),
authorization_endpoint: format!("{}/authorize", server.uri()),
token_endpoint: format!("{}/token", server.uri()),
registration_endpoint: Some(format!("{}/register", server.uri())),
}
}
fn sample_tokens(server: &MockServer) -> McpTokens {
McpTokens {
server: format!("{}/mcp", server.uri()),
access_token: "OLD_AT".into(),
refresh_token: "OLD_RT".into(),
expires_at: Utc::now() - ChronoDuration::seconds(60),
client_id: "CID".into(),
}
}
async fn mount_protected_resource(server: &MockServer, at_mcp_suffix: bool) {
let path_str = if at_mcp_suffix {
"/.well-known/oauth-protected-resource/mcp"
} else {
"/.well-known/oauth-protected-resource"
};
Mock::given(method("GET"))
.and(path(path_str))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"resource": format!("{}/mcp", server.uri()),
"authorization_servers": [server.uri()],
})))
.mount(server)
.await;
}
async fn mount_as_metadata(server: &MockServer, registration_endpoint: Option<&str>) {
let mut body = serde_json::json!({
"authorization_endpoint": format!("{}/authorize", server.uri()),
"token_endpoint": format!("{}/token", server.uri()),
});
if let Some(reg) = registration_endpoint {
body["registration_endpoint"] = serde_json::json!(reg);
}
Mock::given(method("GET"))
.and(path("/.well-known/oauth-authorization-server"))
.respond_with(ResponseTemplate::new(200).set_body_json(body))
.mount(server)
.await;
}
#[tokio::test]
async fn discover_uses_mcp_suffixed_metadata_first() {
let server = MockServer::start().await;
mount_protected_resource(&server, true).await;
mount_as_metadata(&server, Some(&format!("{}/register", server.uri()))).await;
let http = reqwest::Client::new();
let discovery = discover(&http, &format!("{}/mcp", server.uri()))
.await
.unwrap();
assert_eq!(discovery.resource, format!("{}/mcp", server.uri()));
assert_eq!(
discovery.authorization_endpoint,
format!("{}/authorize", server.uri())
);
assert_eq!(discovery.token_endpoint, format!("{}/token", server.uri()));
assert_eq!(
discovery.registration_endpoint,
Some(format!("{}/register", server.uri()))
);
}
#[tokio::test]
async fn discover_falls_back_to_bare_metadata_path() {
let server = MockServer::start().await;
mount_protected_resource(&server, false).await;
mount_as_metadata(&server, None).await;
let http = reqwest::Client::new();
let discovery = discover(&http, &format!("{}/mcp", server.uri()))
.await
.unwrap();
assert_eq!(discovery.registration_endpoint, None);
}
#[tokio::test]
async fn register_posts_expected_json_and_returns_client_id() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/register"))
.and(body_partial_json(serde_json::json!({
"redirect_uris": ["http://localhost:54321"],
"client_name": "pidge",
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
})))
.respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({
"client_id": "issued-client-jwt",
"client_id_issued_at": 0,
"client_secret_expires_at": 0,
"redirect_uris": ["http://localhost:54321"],
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
})))
.expect(1)
.mount(&server)
.await;
let http = reqwest::Client::new();
let discovery = discovery_for(&server);
let client_id = register(&http, &discovery, "http://localhost:54321", "pidge")
.await
.unwrap();
assert_eq!(client_id, "issued-client-jwt");
}
#[tokio::test]
async fn register_errors_when_registration_endpoint_missing() {
let server = MockServer::start().await;
let mut discovery = discovery_for(&server);
discovery.registration_endpoint = None;
let http = reqwest::Client::new();
let err = register(&http, &discovery, "http://localhost:1", "pidge")
.await
.unwrap_err();
match err {
ClientError::Graph { status, message } => {
assert_eq!(status, 400);
assert_eq!(
message,
"server does not support dynamic client registration"
);
}
other => panic!("expected ClientError::Graph, got {other:?}"),
}
}
#[tokio::test]
async fn refresh_on_expiry_calls_refresh_grant_once_and_returns_new_tokens() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/token"))
.and(body_string_contains("grant_type=refresh_token"))
.and(body_string_contains("refresh_token=OLD_RT"))
.and(body_string_contains("client_id=CID"))
.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 http = reqwest::Client::new();
let discovery = discovery_for(&server);
let old_tokens = sample_tokens(&server);
let new_tokens = refresh(&http, &discovery, &old_tokens).await.unwrap();
assert_eq!(new_tokens.access_token, "NEW_AT");
assert_eq!(new_tokens.refresh_token, "NEW_RT");
assert_eq!(new_tokens.server, old_tokens.server);
assert_eq!(new_tokens.client_id, old_tokens.client_id);
}
#[tokio::test]
async fn valid_access_token_refreshes_expired_tokens_via_discovery() {
let server = MockServer::start().await;
mount_protected_resource(&server, true).await;
mount_as_metadata(&server, Some(&format!("{}/register", server.uri()))).await;
Mock::given(method("POST"))
.and(path("/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 http = reqwest::Client::new();
let mut tokens = sample_tokens(&server);
let server_url = tokens.server.clone();
let access = valid_access_token(&http, &server_url, &mut tokens)
.await
.unwrap();
assert_eq!(access, "NEW_AT");
assert_eq!(tokens.access_token, "NEW_AT");
}
#[tokio::test]
async fn refresh_maps_invalid_grant_to_session_expired() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/token"))
.respond_with(ResponseTemplate::new(400).set_body_json(serde_json::json!({
"error": "invalid_grant",
"error_description": "refresh token revoked"
})))
.mount(&server)
.await;
let http = reqwest::Client::new();
let discovery = discovery_for(&server);
let tokens = sample_tokens(&server);
let err = refresh(&http, &discovery, &tokens).await.unwrap_err();
match err {
ClientError::SessionExpired { email } => assert_eq!(email, tokens.server),
other => panic!("expected SessionExpired, got {other:?}"),
}
}
#[tokio::test]
async fn refresh_maps_invalid_client_to_session_expired() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/token"))
.respond_with(ResponseTemplate::new(401).set_body_json(serde_json::json!({
"error": "invalid_client",
"error_description": "unknown client"
})))
.mount(&server)
.await;
let http = reqwest::Client::new();
let discovery = discovery_for(&server);
let tokens = sample_tokens(&server);
let err = refresh(&http, &discovery, &tokens).await.unwrap_err();
match err {
ClientError::SessionExpired { email } => assert_eq!(email, tokens.server),
other => panic!("expected SessionExpired, got {other:?}"),
}
}
#[tokio::test]
async fn fresh_token_skips_the_network_entirely() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/token"))
.respond_with(ResponseTemplate::new(500))
.expect(0)
.mount(&server)
.await;
let http = reqwest::Client::new();
let mut tokens = sample_tokens(&server);
tokens.expires_at = Utc::now() + ChronoDuration::seconds(3600);
let server_url = tokens.server.clone();
let access = valid_access_token(&http, &server_url, &mut tokens)
.await
.unwrap();
assert_eq!(access, "OLD_AT");
}
#[test]
fn build_authorize_url_contains_every_query_pair() {
let url = build_authorize_url(
"https://mcp.example.com/authorize",
"client-id-here",
"http://localhost:54321",
"challenge-here",
"state-here",
"https://mcp.example.com/mcp",
)
.unwrap();
assert!(url.starts_with("https://mcp.example.com/authorize?"));
assert!(url.contains("response_type=code"));
assert!(url.contains("client_id=client-id-here"));
assert!(url.contains("redirect_uri=http%3A%2F%2Flocalhost%3A54321"));
assert!(url.contains("state=state-here"));
assert!(url.contains("code_challenge=challenge-here"));
assert!(url.contains("code_challenge_method=S256"));
assert!(url.contains("resource=https%3A%2F%2Fmcp.example.com%2Fmcp"));
}
fn query_param(url: &str, key: &str) -> Option<String> {
Url::parse(url)
.ok()?
.query_pairs()
.find(|(k, _)| k == key)
.map(|(_, v)| v.into_owned())
}
#[tokio::test]
async fn sign_in_reports_state_mismatch_when_callback_state_is_wrong() {
let server = MockServer::start().await;
mount_protected_resource(&server, true).await;
mount_as_metadata(&server, Some(&format!("{}/register", server.uri()))).await;
Mock::given(method("POST"))
.and(path("/register"))
.respond_with(
ResponseTemplate::new(201).set_body_json(serde_json::json!({"client_id": "cid-1"})),
)
.mount(&server)
.await;
let http = reqwest::Client::new();
let result = sign_in(&http, &format!("{}/mcp", server.uri()), "pidge", |url| {
let redirect_uri =
query_param(url, "redirect_uri").expect("redirect_uri present in authorize URL");
let callback_url = format!("{redirect_uri}/?code=ignored-code&state=totally-wrong");
tokio::spawn(async move {
let _ = reqwest::Client::new().get(&callback_url).send().await;
});
})
.await;
match result {
Err(ClientError::Graph { status, message }) => {
assert_eq!(status, 400);
assert!(message.contains("state mismatch"), "{message}");
}
other => panic!("expected a state-mismatch Graph error, got {other:?}"),
}
}
}