use crate::{
config::McpOAuthConfig,
http_body::{DEFAULT_BOUNDED_BODY_MAX_BYTES, read_bounded_response_text},
mcp::{McpError, McpResult},
persistence::{CrossProcessFileLock, atomic_write_with_permissions, sync_parent_dir},
};
use base64::Engine;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::{
collections::BTreeMap,
fmt, fs, io,
io::{Read, Write},
net::{TcpListener, TcpStream},
path::{Path, PathBuf},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
mpsc::Receiver,
},
time::{Duration, Instant},
};
use chrono::Utc;
const TOKEN_FILE_MODE: u32 = 0o600;
const HTTP_CLIENT_NAME: &str = "magi-code";
const REFRESH_SKEW_SECONDS: i64 = 60;
pub(crate) const CALLBACK_PATH: &str = "/mcp/oauth/callback";
pub(crate) const LOGIN_WAIT_TIMEOUT: Duration = Duration::from_secs(300);
#[cfg(not(test))]
const CALLBACK_STREAM_TIMEOUT: Duration = Duration::from_secs(5);
#[cfg(test)]
const CALLBACK_STREAM_TIMEOUT: Duration = Duration::from_millis(100);
#[derive(Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub(crate) struct ProtectedResourceMetadata {
pub(crate) resource: String,
#[serde(default)]
pub(crate) authorization_servers: Vec<String>,
#[serde(default, flatten)]
pub(crate) extra: BTreeMap<String, Value>,
}
impl fmt::Debug for ProtectedResourceMetadata {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProtectedResourceMetadata")
.field("resource", &self.resource)
.field("authorization_servers", &self.authorization_servers)
.field("extra", &self.extra)
.finish()
}
}
#[derive(Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub(crate) struct AuthorizationServerMetadata {
pub(crate) issuer: String,
pub(crate) authorization_endpoint: String,
pub(crate) token_endpoint: String,
#[serde(default)]
pub(crate) registration_endpoint: Option<String>,
#[serde(default)]
pub(crate) scopes_supported: Option<Vec<String>>,
#[serde(default)]
pub(crate) response_types_supported: Option<Vec<String>>,
#[serde(default)]
pub(crate) grant_types_supported: Option<Vec<String>>,
#[serde(default)]
pub(crate) code_challenge_methods_supported: Option<Vec<String>>,
#[serde(default, flatten)]
pub(crate) extra: BTreeMap<String, Value>,
}
impl fmt::Debug for AuthorizationServerMetadata {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AuthorizationServerMetadata")
.field("issuer", &self.issuer)
.field("authorization_endpoint", &self.authorization_endpoint)
.field("token_endpoint", &self.token_endpoint)
.field("registration_endpoint", &self.registration_endpoint)
.field("scopes_supported", &self.scopes_supported)
.field("response_types_supported", &self.response_types_supported)
.field("grant_types_supported", &self.grant_types_supported)
.field(
"code_challenge_methods_supported",
&self.code_challenge_methods_supported,
)
.field("extra", &self.extra)
.finish()
}
}
#[derive(Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub(crate) struct RegistrationResponse {
pub(crate) client_id: String,
#[serde(default)]
pub(crate) client_secret: Option<String>,
#[serde(default, flatten)]
pub(crate) extra: BTreeMap<String, Value>,
}
impl fmt::Debug for RegistrationResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RegistrationResponse")
.field("client_id", &self.client_id)
.field(
"client_secret",
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
)
.field("extra", &self.extra)
.finish()
}
}
#[derive(Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub(crate) struct TokenResponse {
pub(crate) access_token: String,
#[serde(default)]
pub(crate) token_type: String,
#[serde(default)]
pub(crate) expires_in: Option<u64>,
#[serde(default)]
pub(crate) refresh_token: Option<String>,
#[serde(default)]
pub(crate) scope: Option<String>,
#[serde(default, flatten)]
pub(crate) extra: BTreeMap<String, Value>,
}
impl fmt::Debug for TokenResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TokenResponse")
.field("access_token", &"[REDACTED]")
.field("token_type", &self.token_type)
.field("expires_in", &self.expires_in)
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "[REDACTED]"),
)
.field("scope", &self.scope)
.field("extra", &self.extra)
.finish()
}
}
#[derive(Debug, Deserialize)]
struct OAuthErrorResponse {
error: Option<String>,
}
fn token_error_message(status: reqwest::StatusCode, body: &str, refresh: bool) -> String {
let parsed = serde_json::from_str::<OAuthErrorResponse>(body).ok();
let code = parsed
.as_ref()
.and_then(|error| error.error.as_deref())
.filter(|code| is_known_oauth_error_code(code))
.unwrap_or_default();
match (refresh, code) {
(false, "invalid_grant") => {
"authorization code invalid or expired; try mcp login again".to_string()
}
(false, "invalid_client") => "client_id rejected by server".to_string(),
(true, "invalid_grant") => {
"refresh token expired or revoked; run magi-code mcp login <server>".to_string()
}
(_, "") if parsed.is_some() => "OAuth token exchange failed".to_string(),
(_, "") => format!("MCP OAuth token endpoint returned HTTP {status}"),
(_, code) => format!("MCP OAuth token endpoint returned OAuth error {code}"),
}
}
fn is_known_oauth_error_code(code: &str) -> bool {
matches!(
code,
"invalid_request"
| "invalid_client"
| "invalid_grant"
| "unauthorized_client"
| "unsupported_grant_type"
| "invalid_scope"
| "server_error"
| "temporarily_unavailable"
)
}
#[derive(Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub(crate) struct StoredToken {
pub(crate) client_id: String,
pub(crate) access_token: String,
#[serde(default)]
pub(crate) refresh_token: Option<String>,
#[serde(default)]
pub(crate) expires_at: Option<i64>,
#[serde(default)]
pub(crate) granted_scopes: Vec<String>,
#[serde(default)]
pub(crate) client_secret: Option<String>,
#[serde(default)]
pub(crate) authorization_server: Option<String>,
#[serde(default)]
pub(crate) issuer: Option<String>,
#[serde(default)]
pub(crate) token_endpoint: Option<String>,
#[serde(default)]
pub(crate) resource: Option<String>,
pub(crate) server_url: String,
pub(crate) token_received_at: i64,
}
impl fmt::Debug for StoredToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StoredToken")
.field("client_id", &self.client_id)
.field("access_token", &"[REDACTED]")
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "[REDACTED]"),
)
.field("expires_at", &self.expires_at)
.field("granted_scopes", &self.granted_scopes)
.field(
"client_secret",
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
)
.field("authorization_server", &self.authorization_server)
.field("issuer", &self.issuer)
.field("token_endpoint", &self.token_endpoint)
.field("resource", &self.resource)
.field("server_url", &self.server_url)
.field("token_received_at", &self.token_received_at)
.finish()
}
}
#[derive(Clone)]
pub(crate) struct McpOAuthLoginOutcome {
pub(crate) authorization_url: String,
pub(crate) redirect_uri: String,
pub(crate) state: String,
pub(crate) verifier: String,
pub(crate) client_id: String,
pub(crate) client_secret: Option<String>,
pub(crate) metadata: AuthorizationServerMetadata,
pub(crate) resource: String,
pub(crate) authorization_server: String,
}
impl fmt::Debug for McpOAuthLoginOutcome {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("McpOAuthLoginOutcome")
.field("authorization_url", &sanitize_url(&self.authorization_url))
.field("redirect_uri", &self.redirect_uri)
.field("state", &"[REDACTED]")
.field("verifier", &"[REDACTED]")
.field("client_id", &self.client_id)
.field(
"client_secret",
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
)
.field("metadata", &self.metadata)
.field("resource", &self.resource)
.field("authorization_server", &self.authorization_server)
.finish()
}
}
#[derive(Clone)]
pub(crate) struct TokenProvider {
inner: Arc<Mutex<TokenProviderInner>>,
}
#[derive(Clone, Debug)]
struct TokenProviderInner {
mc_home: PathBuf,
server_name: String,
server_url: String,
oauth: McpOAuthConfig,
client: reqwest::blocking::Client,
}
impl TokenProvider {
pub(crate) fn new(
mc_home: PathBuf,
server_name: String,
server_url: String,
oauth: McpOAuthConfig,
client: reqwest::blocking::Client,
) -> Self {
Self {
inner: Arc::new(Mutex::new(TokenProviderInner {
mc_home,
server_name,
server_url,
oauth,
client,
})),
}
}
pub(crate) fn access_token(&self) -> McpResult<String> {
self.with_inner(false)
}
pub(crate) fn force_refresh_access_token(&self) -> McpResult<String> {
self.with_inner(true)
}
fn with_inner(&self, force_refresh: bool) -> McpResult<String> {
let inner = self
.inner
.lock()
.map_err(|_| McpError::Transport("MCP OAuth token provider lock poisoned".to_string()))?
.clone();
inner.access_token(force_refresh)
}
}
impl std::fmt::Debug for TokenProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TokenProvider").finish_non_exhaustive()
}
}
impl TokenProviderInner {
fn access_token(&self, force_refresh: bool) -> McpResult<String> {
loop {
let token = read_token_locked(&self.mc_home, &self.server_name)?.ok_or_else(|| {
McpError::Config(format!(
"not authenticated; run magi-code mcp login {}",
self.server_name
))
})?;
self.validate_stored_token_metadata(&token)?;
if !force_refresh && !token_needs_refresh(&token) {
if token.access_token.trim().is_empty() {
return Err(McpError::Config(format!(
"not authenticated; run magi-code mcp login {}",
self.server_name
)));
}
return Ok(token.access_token);
}
let refreshed = self.refresh_stored_token(&token)?;
if write_token_if_unchanged(&self.mc_home, &self.server_name, &token, &refreshed)? {
return Ok(refreshed.access_token);
}
}
}
fn validate_stored_token_metadata(&self, token: &StoredToken) -> McpResult<()> {
if !validate_token_url(token, &self.server_url) {
return Err(self.relogin_error("server URL changed since last login"));
}
if let Some(endpoint) = token
.token_endpoint
.as_deref()
.filter(|value| !value.is_empty())
{
validate_oauth_endpoint("stored token_endpoint", endpoint)?;
}
if token.resource.is_none()
&& token.authorization_server.is_none()
&& token.issuer.is_none()
{
return Ok(());
}
let protected = discover_protected_resource(&self.server_url, &self.client)?;
let current_resource = protected
.as_ref()
.map(|metadata| metadata.resource.as_str())
.unwrap_or(&self.server_url);
if token
.resource
.as_deref()
.is_some_and(|stored| stored != current_resource)
{
return Err(self.relogin_error("OAuth resource changed since last login"));
}
let auth_server = select_authorization_server(
self.oauth.authorization_server.as_deref(),
protected.as_ref(),
)?;
if token
.authorization_server
.as_deref()
.is_some_and(|stored| stored != auth_server)
{
return Err(self.relogin_error("OAuth authorization server changed since last login"));
}
if token.issuer.is_some() {
let metadata = discover_authorization_server(&auth_server, &self.client)?;
if token.issuer.as_deref() != Some(metadata.issuer.as_str()) {
return Err(self.relogin_error("OAuth issuer changed since last login"));
}
}
Ok(())
}
fn relogin_error(&self, reason: &str) -> McpError {
McpError::Config(format!(
"{reason}; run magi-code mcp login {}",
self.server_name
))
}
fn refresh_stored_token(&self, stored: &StoredToken) -> McpResult<StoredToken> {
let refresh = stored
.refresh_token
.as_deref()
.filter(|value| !value.is_empty())
.ok_or_else(|| {
McpError::Config(format!(
"token expired; run magi-code mcp login {}",
self.server_name
))
})?;
let token_endpoint = match stored.token_endpoint.as_deref() {
Some(endpoint) if !endpoint.is_empty() => endpoint.to_string(),
_ => discover_token_endpoint(&self.server_url, &self.oauth, &self.client)?,
};
validate_oauth_endpoint("token_endpoint", &token_endpoint)?;
let response = refresh_token(
&token_endpoint,
refresh,
&stored.client_id,
&self.server_url,
stored.client_secret.as_deref(),
&self.client,
)
.map_err(|_| {
McpError::Config(format!(
"refresh token expired or revoked; run magi-code mcp login {}",
self.server_name
))
})?;
let mut next = stored_token_from_response(
response,
&stored.client_id,
stored.client_secret.clone(),
&self.server_url,
stored.authorization_server.clone(),
stored.issuer.clone(),
Some(token_endpoint),
stored.resource.clone(),
)?;
if next.refresh_token.is_none() {
next.refresh_token = stored.refresh_token.clone();
}
Ok(next)
}
}
pub(crate) fn token_needs_refresh(token: &StoredToken) -> bool {
token
.expires_at
.is_some_and(|expires| expires <= Utc::now().timestamp() + REFRESH_SKEW_SECONDS)
}
pub(crate) fn auth_status(
mc_home: &Path,
server_name: &str,
server_url: &str,
) -> McpResult<&'static str> {
match read_token(mc_home, server_name)? {
None => Ok("not authenticated"),
Some(token) if !validate_token_url(&token, server_url) => Ok("invalid (needs login)"),
Some(token) if token_needs_refresh(&token) && token.refresh_token.is_some() => {
Ok("expired (refreshable)")
}
Some(token) if token_needs_refresh(&token) => Ok("expired (needs login)"),
Some(_) => Ok("authenticated"),
}
}
pub(crate) fn generate_code_verifier() -> String {
random_urlsafe(64)
}
pub(crate) fn code_challenge(verifier: &str) -> String {
let digest = Sha256::digest(verifier.as_bytes());
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest)
}
pub(crate) fn generate_state() -> String {
random_urlsafe(32)
}
pub(crate) fn authorization_url(
authorization_endpoint: &str,
client_id: &str,
redirect_uri: &str,
challenge: &str,
state: &str,
resource: &str,
scopes: &[String],
) -> String {
let mut pairs = vec![
("response_type", "code"),
("client_id", client_id),
("redirect_uri", redirect_uri),
("code_challenge", challenge),
("code_challenge_method", "S256"),
("state", state),
("resource", resource),
];
let scope = scopes.join(" ");
if !scope.is_empty() {
pairs.push(("scope", scope.as_str()));
}
let query = pairs
.into_iter()
.map(|(key, value)| format!("{}={}", pct(key), pct(value)))
.collect::<Vec<_>>()
.join("&");
format!("{authorization_endpoint}?{query}")
}
pub(crate) fn discover_protected_resource(
origin_url: &str,
client: &reqwest::blocking::Client,
) -> McpResult<Option<ProtectedResourceMetadata>> {
let mut first_error: Option<McpError> = None;
for metadata_url in protected_resource_metadata_urls(origin_url)? {
match fetch_protected_resource_metadata(client, metadata_url, origin_url) {
Ok(Some(metadata)) => return Ok(Some(metadata)),
Ok(None) => {}
Err(error) => {
if first_error.is_none() {
first_error = Some(error);
}
}
}
}
if let Some(metadata) = discover_protected_resource_from_www_authenticate(origin_url, client)? {
return Ok(Some(metadata));
}
if let Some(error) = first_error {
return Err(error);
}
Ok(None)
}
fn fetch_protected_resource_metadata(
client: &reqwest::blocking::Client,
metadata_url: reqwest::Url,
origin_url: &str,
) -> McpResult<Option<ProtectedResourceMetadata>> {
validate_oauth_endpoint("resource_metadata", metadata_url.as_str())?;
let response = client.get(metadata_url).send().map_err(|_| {
McpError::Transport(format!(
"network error during OAuth protected-resource discovery at {}",
sanitize_url(origin_url)
))
})?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(McpError::Transport(format!(
"OAuth protected-resource discovery failed at {} with HTTP {}",
sanitize_url(origin_url),
response.status()
)));
}
let text = read_oauth_success_text(response, "OAuth protected-resource metadata")?;
let metadata: ProtectedResourceMetadata = serde_json::from_str(&text).map_err(|_| {
McpError::Protocol {
code: -32700,
message: "OAuth protected-resource metadata is malformed JSON (expected resource and authorization_servers)".to_string(),
}
})?;
validate_protected_resource_metadata(&metadata)?;
Ok(Some(metadata))
}
fn discover_protected_resource_from_www_authenticate(
origin_url: &str,
client: &reqwest::blocking::Client,
) -> McpResult<Option<ProtectedResourceMetadata>> {
let response = match client.get(origin_url).send() {
Ok(response) => response,
Err(_) => return Ok(None),
};
if response.status() != reqwest::StatusCode::UNAUTHORIZED {
return Ok(None);
}
let Some(metadata_url) = response
.headers()
.get(reqwest::header::WWW_AUTHENTICATE)
.and_then(|value| value.to_str().ok())
.and_then(parse_www_authenticate_resource_metadata)
else {
return Ok(None);
};
let url = reqwest::Url::parse(&metadata_url).map_err(|_| {
McpError::Config("OAuth WWW-Authenticate resource_metadata URL is invalid".to_string())
})?;
fetch_protected_resource_metadata(client, url, origin_url)
}
pub(crate) fn discover_authorization_server(
auth_server_url: &str,
client: &reqwest::blocking::Client,
) -> McpResult<AuthorizationServerMetadata> {
let mut saw_404 = false;
validate_oauth_endpoint("authorization_server", auth_server_url)?;
for metadata_url in authorization_server_metadata_urls(auth_server_url)? {
let response = client.get(metadata_url).send().map_err(|_| {
McpError::Transport(format!(
"could not discover OAuth authorization server at {}: network error",
sanitize_url(auth_server_url)
))
})?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
saw_404 = true;
continue;
}
if !response.status().is_success() {
return Err(McpError::Transport(format!(
"could not discover OAuth authorization server at {}: HTTP {}",
sanitize_url(auth_server_url),
response.status()
)));
}
let text = read_oauth_success_text(response, "OAuth authorization-server metadata")?;
let metadata: AuthorizationServerMetadata =
serde_json::from_str(&text).map_err(|_| McpError::Protocol {
code: -32700,
message: format!(
"OAuth authorization-server metadata at {} is malformed JSON",
sanitize_url(auth_server_url)
),
})?;
validate_authorization_server_metadata(&metadata)?;
return Ok(metadata);
}
let _ = saw_404;
Err(McpError::Config(format!(
"could not discover OAuth authorization server at {}",
sanitize_url(auth_server_url)
)))
}
pub(crate) fn parse_www_authenticate_resource_metadata(header: &str) -> Option<String> {
header.split(',').find_map(|part| {
let (key, value) = part.trim().split_once('=')?;
let key = key.split_whitespace().last().unwrap_or(key).trim();
if !key.eq_ignore_ascii_case("resource_metadata") {
return None;
}
Some(value.trim().trim_matches('"').to_string()).filter(|value| !value.is_empty())
})
}
pub(crate) fn select_authorization_server(
configured: Option<&str>,
protected: Option<&ProtectedResourceMetadata>,
) -> McpResult<String> {
if let Some(configured) = configured {
validate_oauth_endpoint("authorization_server", configured)?;
if let Some(protected) = protected
&& !protected.authorization_servers.is_empty()
&& !protected
.authorization_servers
.iter()
.any(|server| server == configured)
{
return Err(McpError::Config(
"configured OAuth authorization_server is not listed in protected-resource metadata".to_string(),
));
}
return Ok(configured.to_string());
}
let server = protected
.and_then(|metadata| metadata.authorization_servers.first())
.ok_or_else(|| {
McpError::Config(
"MCP OAuth protected-resource metadata missing authorization_servers".to_string(),
)
})?;
validate_oauth_endpoint("authorization_server", server)?;
Ok(server.to_string())
}
pub(crate) fn register_client(
registration_endpoint: &str,
redirect_uris: Vec<String>,
client: &reqwest::blocking::Client,
) -> McpResult<RegistrationResponse> {
validate_oauth_endpoint("registration_endpoint", registration_endpoint)?;
let body = serde_json::json!({
"redirect_uris": redirect_uris,
"client_name": HTTP_CLIENT_NAME,
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"token_endpoint_auth_method": "none",
"code_challenge_method": "S256"
});
let response = client
.post(registration_endpoint)
.json(&body)
.send()
.map_err(|_| {
McpError::Transport("MCP OAuth dynamic client registration failed".to_string())
})?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Err(McpError::Config(
"server does not support dynamic client registration; configure client_id".to_string(),
));
}
if !response.status().is_success() {
return Err(McpError::Transport(format!(
"MCP OAuth dynamic client registration was rejected with HTTP {}",
response.status()
)));
}
let text = read_oauth_success_text(response, "MCP OAuth dynamic client registration")?;
let registered: RegistrationResponse =
serde_json::from_str(&text).map_err(|_| McpError::Protocol {
code: -32700,
message: "MCP OAuth dynamic client registration response is malformed JSON".to_string(),
})?;
if registered.client_id.trim().is_empty() {
return Err(McpError::Protocol {
code: -32602,
message: "MCP OAuth dynamic client registration response missing client_id".to_string(),
});
}
Ok(registered)
}
fn validate_token_server_name(server_name: &str) -> McpResult<()> {
if server_name.is_empty()
|| server_name.contains("__")
|| !server_name
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-')
{
return Err(McpError::Config(format!(
"MCP server name '{server_name}' is not a valid token store name"
)));
}
Ok(())
}
pub(crate) fn token_file_path(mc_home: &Path, server_name: &str) -> McpResult<PathBuf> {
validate_token_server_name(server_name)?;
Ok(mc_home
.join("mcp-tokens")
.join(format!("{server_name}.json")))
}
pub(crate) fn read_token(mc_home: &Path, server_name: &str) -> McpResult<Option<StoredToken>> {
let path = token_file_path(mc_home, server_name)?;
let _lock = CrossProcessFileLock::acquire(&path).map_err(lock_error)?;
read_token_unlocked(mc_home, server_name)
}
fn read_token_unlocked(mc_home: &Path, server_name: &str) -> McpResult<Option<StoredToken>> {
let path = token_file_path(mc_home, server_name)?;
match read_token_file_text(&path, server_name)? {
Some(text) => serde_json::from_str(&text).map(Some).map_err(|_| {
McpError::Config(format!(
"MCP OAuth token file for server '{server_name}' is corrupt"
))
}),
None => Ok(None),
}
}
#[cfg(windows)]
fn validate_token_file_path_before_open(path: &Path) -> McpResult<()> {
use std::os::windows::fs::MetadataExt;
const FILE_ATTRIBUTE_REPARSE_POINT: u32 = 0x400;
match fs::symlink_metadata(path) {
Ok(metadata) if metadata.file_attributes() & FILE_ATTRIBUTE_REPARSE_POINT != 0 => {
Err(McpError::Config(
"MCP OAuth token file must be a regular private file; reparse-point token files are not allowed".to_string(),
))
}
Ok(_) => Ok(()),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(McpError::Transport(format!(
"failed to stat MCP OAuth token file: {error}"
))),
}
}
#[cfg(not(windows))]
fn validate_token_file_path_before_open(_path: &Path) -> McpResult<()> {
Ok(())
}
#[cfg(unix)]
fn path_is_symlink(path: &Path) -> bool {
fs::symlink_metadata(path)
.map(|metadata| metadata.file_type().is_symlink())
.unwrap_or(false)
}
#[cfg(all(unix, target_os = "linux"))]
fn o_no_follow() -> i32 {
0x20000
}
#[cfg(all(unix, not(target_os = "linux")))]
fn o_no_follow() -> i32 {
0x100
}
#[cfg(unix)]
fn validate_open_token_file(file: &fs::File) -> McpResult<()> {
use std::os::unix::fs::PermissionsExt;
let metadata = file.metadata().map_err(|error| {
McpError::Transport(format!("failed to stat MCP OAuth token file: {error}"))
})?;
if !metadata.is_file() {
return Err(McpError::Config(
"MCP OAuth token file must be a regular private file".to_string(),
));
}
if metadata.permissions().mode() & 0o077 != 0 {
return Err(McpError::Config(
"MCP OAuth token file permissions must be private/owner-only (0600 or stricter)"
.to_string(),
));
}
Ok(())
}
#[cfg(not(unix))]
fn validate_open_token_file(_file: &fs::File) -> McpResult<()> {
Ok(())
}
fn read_token_file_text(path: &Path, server_name: &str) -> McpResult<Option<String>> {
validate_token_file_path_before_open(path)?;
let mut options = fs::OpenOptions::new();
options.read(true);
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt;
options.custom_flags(o_no_follow());
}
let mut file = match options.open(path) {
Ok(file) => file,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(None),
#[cfg(unix)]
Err(error) if path_is_symlink(path) => {
let _ = error;
return Err(McpError::Config(
"MCP OAuth token file must be a regular private file; symlinked token files are not allowed".to_string(),
));
}
Err(error) => {
return Err(McpError::Transport(format!(
"failed to read MCP OAuth token file for server '{server_name}': {error}"
)));
}
};
validate_open_token_file(&file)?;
let mut text = String::new();
file.read_to_string(&mut text).map_err(|error| {
McpError::Transport(format!(
"failed to read MCP OAuth token file for server '{server_name}': {error}"
))
})?;
Ok(Some(text))
}
pub(crate) fn write_token(mc_home: &Path, server_name: &str, token: &StoredToken) -> McpResult<()> {
let path = token_file_path(mc_home, server_name)?;
let _lock = CrossProcessFileLock::acquire(&path).map_err(lock_error)?;
write_token_unlocked(mc_home, server_name, token)
}
fn write_token_unlocked(mc_home: &Path, server_name: &str, token: &StoredToken) -> McpResult<()> {
let path = token_file_path(mc_home, server_name)?;
let bytes = serde_json::to_vec_pretty(token).map_err(|_| {
McpError::Config(format!(
"failed to serialize MCP OAuth token for server '{server_name}'"
))
})?;
atomic_write_with_permissions(&path, &bytes, Some(TOKEN_FILE_MODE)).map_err(|error| {
McpError::Transport(format!(
"failed to write MCP OAuth token file for server '{server_name}': {error}"
))
})
}
pub(crate) fn delete_token(mc_home: &Path, server_name: &str) -> McpResult<()> {
let path = token_file_path(mc_home, server_name)?;
let _lock = CrossProcessFileLock::acquire(&path).map_err(lock_error)?;
match fs::remove_file(&path) {
Ok(()) => sync_parent_after_token_delete(&path, server_name),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(McpError::Transport(format!(
"failed to delete MCP OAuth token file for server '{server_name}': {error}"
))),
}
}
fn sync_parent_after_token_delete(path: &Path, server_name: &str) -> McpResult<()> {
let parent = path.parent().ok_or_else(|| {
McpError::Transport(format!(
"failed to sync parent directory after deleting MCP OAuth token file for server '{server_name}': token path has no parent"
))
})?;
sync_parent_dir(parent).map_err(|error| {
McpError::Transport(format!(
"failed to sync parent directory after deleting MCP OAuth token file for server '{server_name}': {error}"
))
})
}
fn read_token_locked(mc_home: &Path, server_name: &str) -> McpResult<Option<StoredToken>> {
with_token_lock(mc_home, server_name, || {
read_token_unlocked(mc_home, server_name)
})
}
fn write_token_if_unchanged(
mc_home: &Path,
server_name: &str,
expected: &StoredToken,
next: &StoredToken,
) -> McpResult<bool> {
with_token_lock(mc_home, server_name, || {
let current = read_token_unlocked(mc_home, server_name)?;
if current.as_ref() != Some(expected) {
return Ok(false);
}
write_token_unlocked(mc_home, server_name, next)?;
Ok(true)
})
}
fn with_token_lock<T>(
mc_home: &Path,
server_name: &str,
f: impl FnOnce() -> McpResult<T>,
) -> McpResult<T> {
let path = token_file_path(mc_home, server_name)?;
let _lock = CrossProcessFileLock::acquire(&path).map_err(lock_error)?;
f()
}
pub(crate) fn validate_token_url(token: &StoredToken, current_url: &str) -> bool {
token.server_url == current_url
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn exchange_code(
token_endpoint: &str,
code: &str,
redirect_uri: &str,
client_id: &str,
code_verifier: &str,
resource: &str,
client_secret: Option<&str>,
client: &reqwest::blocking::Client,
) -> McpResult<TokenResponse> {
let mut form = vec![
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", redirect_uri),
("client_id", client_id),
("code_verifier", code_verifier),
("resource", resource),
];
if let Some(secret) = client_secret.filter(|value| !value.is_empty()) {
form.push(("client_secret", secret));
}
post_token_form(token_endpoint, &form, client)
}
pub(crate) fn refresh_token(
token_endpoint: &str,
refresh_token_value: &str,
client_id: &str,
resource: &str,
client_secret: Option<&str>,
client: &reqwest::blocking::Client,
) -> McpResult<TokenResponse> {
let mut form = vec![
("grant_type", "refresh_token"),
("refresh_token", refresh_token_value),
("client_id", client_id),
("resource", resource),
];
if let Some(secret) = client_secret.filter(|value| !value.is_empty()) {
form.push(("client_secret", secret));
}
post_token_form(token_endpoint, &form, client)
}
fn post_token_form(
token_endpoint: &str,
form: &[(&str, &str)],
client: &reqwest::blocking::Client,
) -> McpResult<TokenResponse> {
validate_oauth_endpoint("token_endpoint", token_endpoint)?;
let refresh = form
.iter()
.any(|(key, value)| *key == "grant_type" && *value == "refresh_token");
let response = client.post(token_endpoint).form(form).send().map_err(|_| {
McpError::Transport(format!(
"network error during OAuth token request to {}",
sanitize_url(token_endpoint)
))
})?;
let status = response.status();
if !status.is_success() {
let body = bounded_response_text(response, 512);
return Err(McpError::Transport(token_error_message(
status, &body, refresh,
)));
}
let text = read_oauth_success_text(response, "MCP OAuth token")?;
let token: TokenResponse = serde_json::from_str(&text).map_err(|_| McpError::Protocol {
code: -32700,
message: "MCP OAuth token response is malformed JSON".to_string(),
})?;
validate_token_response(&token)?;
Ok(token)
}
fn bounded_response_text(mut response: reqwest::blocking::Response, max: usize) -> String {
let mut bytes = Vec::new();
let _ = response
.by_ref()
.take(max as u64 + 1)
.read_to_end(&mut bytes);
bytes.truncate(max);
String::from_utf8_lossy(&bytes).into_owned()
}
fn read_oauth_success_text(
response: reqwest::blocking::Response,
response_label: &str,
) -> McpResult<String> {
read_bounded_response_text(response, DEFAULT_BOUNDED_BODY_MAX_BYTES).map_err(|error| {
let error = error.to_string();
if error.contains("response exceeded") {
McpError::Transport(format!(
"{response_label} response exceeded {DEFAULT_BOUNDED_BODY_MAX_BYTES} bytes"
))
} else {
McpError::Transport(format!("{response_label} response read failed: {error}"))
}
})
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn stored_token_from_response(
response: TokenResponse,
client_id: &str,
client_secret: Option<String>,
server_url: &str,
authorization_server: Option<String>,
issuer: Option<String>,
token_endpoint: Option<String>,
resource: Option<String>,
) -> McpResult<StoredToken> {
validate_token_response(&response)?;
let now = Utc::now().timestamp();
let scopes = response
.scope
.as_deref()
.unwrap_or("")
.split_whitespace()
.map(ToString::to_string)
.collect::<Vec<_>>();
let expires_at = response
.expires_in
.map(|seconds| {
let seconds = i64::try_from(seconds).map_err(|_| McpError::Protocol {
code: -32602,
message: "MCP OAuth token response expires_in is too large".to_string(),
})?;
now.checked_add(seconds).ok_or_else(|| McpError::Protocol {
code: -32602,
message: "MCP OAuth token response expires_in is too large".to_string(),
})
})
.transpose()?;
Ok(StoredToken {
client_id: client_id.to_string(),
access_token: response.access_token,
refresh_token: response.refresh_token.filter(|value| !value.is_empty()),
expires_at,
granted_scopes: scopes,
client_secret,
authorization_server,
issuer,
token_endpoint,
resource,
server_url: server_url.to_string(),
token_received_at: now,
})
}
fn validate_token_response(response: &TokenResponse) -> McpResult<()> {
if response.access_token.trim().is_empty() {
return Err(McpError::Protocol {
code: -32602,
message: "MCP OAuth token response missing access_token".to_string(),
});
}
if !response.token_type.is_empty() && !response.token_type.eq_ignore_ascii_case("bearer") {
return Err(McpError::Protocol {
code: -32602,
message: "MCP OAuth token response token_type is not bearer".to_string(),
});
}
Ok(())
}
pub(crate) fn prepare_login(
server_url: &str,
oauth: &McpOAuthConfig,
redirect_uri: &str,
client: &reqwest::blocking::Client,
) -> McpResult<McpOAuthLoginOutcome> {
let protected = discover_protected_resource(server_url, client)?;
let auth_server =
select_authorization_server(oauth.authorization_server.as_deref(), protected.as_ref())?;
let metadata = discover_authorization_server(&auth_server, client)?;
let registration = if oauth.client_id.is_none() {
match metadata.registration_endpoint.as_deref() {
Some(endpoint) => Some(register_client(
endpoint,
vec![redirect_uri.to_string()],
client,
)?),
None => {
return Err(McpError::Config(
"server does not support dynamic client registration; configure client_id"
.to_string(),
));
}
}
} else {
None
};
let client_id = oauth
.client_id
.clone()
.or_else(|| registration.as_ref().map(|r| r.client_id.clone()))
.ok_or_else(|| McpError::Config("MCP OAuth client_id missing".to_string()))?;
let client_secret = registration.and_then(|r| r.client_secret);
let verifier = generate_code_verifier();
let challenge = code_challenge(&verifier);
let state = generate_state();
let resource = protected
.as_ref()
.map(|p| p.resource.clone())
.unwrap_or_else(|| server_url.to_string());
let authorization_url = authorization_url(
&metadata.authorization_endpoint,
&client_id,
redirect_uri,
&challenge,
&state,
&resource,
&oauth.scopes,
);
Ok(McpOAuthLoginOutcome {
authorization_url,
redirect_uri: redirect_uri.to_string(),
state,
verifier,
client_id,
client_secret,
metadata,
resource,
authorization_server: auth_server,
})
}
fn discover_token_endpoint(
server_url: &str,
oauth: &McpOAuthConfig,
client: &reqwest::blocking::Client,
) -> McpResult<String> {
let protected = discover_protected_resource(server_url, client)?;
let auth_server =
select_authorization_server(oauth.authorization_server.as_deref(), protected.as_ref())?;
Ok(discover_authorization_server(&auth_server, client)?.token_endpoint)
}
pub(crate) fn bind_callback_listener() -> McpResult<(TcpListener, String)> {
let listener = TcpListener::bind("127.0.0.1:0").map_err(|_| {
McpError::Transport("could not bind MCP OAuth callback on 127.0.0.1".to_string())
})?;
listener
.set_nonblocking(true)
.map_err(McpError::transport)?;
let addr = listener.local_addr().map_err(McpError::transport)?;
Ok((listener, format!("http://{addr}{CALLBACK_PATH}")))
}
pub(crate) fn capture_loopback_or_manual_code(
listener: Option<TcpListener>,
expected_state: &str,
timeout: Duration,
cancel: &AtomicBool,
manual_rx: Option<&Receiver<String>>,
) -> McpResult<String> {
let deadline = Instant::now() + timeout;
loop {
if cancel.load(Ordering::SeqCst) {
return Err(McpError::Transport(
"MCP OAuth login cancelled; token unchanged".to_string(),
));
}
if let Some(rx) = manual_rx {
match rx.try_recv() {
Ok(input) => return parse_manual_fallback_input(&input, expected_state),
Err(std::sync::mpsc::TryRecvError::Empty)
| Err(std::sync::mpsc::TryRecvError::Disconnected) => {}
}
}
if let Some(listener) = &listener {
match listener.accept() {
Ok((mut stream, _)) => return handle_callback_stream(&mut stream, expected_state),
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {}
Err(error) => {
return Err(McpError::Transport(format!(
"MCP OAuth callback failed: {error}"
)));
}
}
}
if Instant::now() >= deadline {
return Err(McpError::Transport(
"OAuth login timed out (no callback received within 5 minutes)".to_string(),
));
}
std::thread::sleep(Duration::from_millis(50));
}
}
fn handle_callback_stream(stream: &mut TcpStream, expected_state: &str) -> McpResult<String> {
stream
.set_read_timeout(Some(CALLBACK_STREAM_TIMEOUT))
.map_err(McpError::transport)?;
stream
.set_write_timeout(Some(CALLBACK_STREAM_TIMEOUT))
.map_err(McpError::transport)?;
let mut buf = [0_u8; 4096];
let n = stream
.read(&mut buf)
.map_err(|_| McpError::Transport("MCP OAuth callback read failed".to_string()))?;
let request = String::from_utf8_lossy(&buf[..n]);
let first = request.lines().next().unwrap_or_default();
let result = parse_callback_request_line(first, expected_state);
let (status, body) = if result.is_ok() {
(
"200 OK",
"MCP OAuth login complete. You can close this tab.",
)
} else {
(
"400 Bad Request",
"MCP OAuth login failed. Return to your terminal.",
)
};
let response = format!(
"HTTP/1.1 {status}\r\ncontent-type: text/html\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
let _ = stream.write_all(response.as_bytes());
result
}
pub(crate) fn parse_callback_request_line(line: &str, expected_state: &str) -> McpResult<String> {
let Some(target) = line
.strip_prefix("GET ")
.and_then(|rest| rest.split_whitespace().next())
else {
return Err(McpError::Transport(
"MCP OAuth callback was malformed".to_string(),
));
};
parse_redirect_target(target, expected_state)
}
pub(crate) fn parse_manual_fallback_input(input: &str, expected_state: &str) -> McpResult<String> {
let trimmed = input.trim();
if trimmed.is_empty() {
return Err(McpError::Transport(
"manual MCP OAuth fallback was empty".to_string(),
));
}
if trimmed.contains('?') || trimmed.starts_with("http://") || trimmed.starts_with("https://") {
let target = trimmed
.split_once("://")
.and_then(|(_, rest)| rest.find('/').map(|idx| &rest[idx..]))
.unwrap_or(trimmed);
return parse_redirect_target(target, expected_state);
}
if trimmed.contains(char::is_whitespace) || trimmed.contains('&') || trimmed.contains('=') {
return Err(McpError::Transport(
"manual MCP OAuth fallback was malformed".to_string(),
));
}
Ok(trimmed.to_string())
}
fn parse_redirect_target(target: &str, expected_state: &str) -> McpResult<String> {
let (path, query) = target.split_once('?').unwrap_or((target, ""));
if path != CALLBACK_PATH {
return Err(McpError::Transport(
"MCP OAuth callback used an unexpected path".to_string(),
));
}
let params = parse_query(query)?;
if params.iter().any(|(key, _)| key == "error") {
return Err(McpError::Transport(
"MCP OAuth provider rejected login".to_string(),
));
}
let state = params
.iter()
.find(|(key, _)| key == "state")
.map(|(_, value)| value.as_str())
.unwrap_or_default();
if state != expected_state {
return Err(McpError::Transport(
"OAuth state mismatch — possible CSRF attack or stale login attempt".to_string(),
));
}
params
.into_iter()
.find(|(key, value)| key == "code" && !value.is_empty())
.map(|(_, value)| value)
.ok_or_else(|| {
McpError::Transport("OAuth callback received without authorization code".to_string())
})
}
fn parse_query(query: &str) -> McpResult<Vec<(String, String)>> {
query
.split('&')
.filter(|part| !part.is_empty())
.map(|part| {
let (k, v) = part.split_once('=').unwrap_or((part, ""));
Ok((decode_pct(k)?, decode_pct(v)?))
})
.collect()
}
fn decode_pct(input: &str) -> McpResult<String> {
let mut out = Vec::new();
let bytes = input.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' {
if i + 2 >= bytes.len() {
return Err(McpError::Transport(
"MCP OAuth callback query contains malformed percent escape".to_string(),
));
}
let high = hex_digit(bytes[i + 1]).ok_or_else(|| {
McpError::Transport(
"MCP OAuth callback query contains malformed percent escape".to_string(),
)
})?;
let low = hex_digit(bytes[i + 2]).ok_or_else(|| {
McpError::Transport(
"MCP OAuth callback query contains malformed percent escape".to_string(),
)
})?;
out.push((high << 4) | low);
i += 3;
} else {
out.push(if bytes[i] == b'+' { b' ' } else { bytes[i] });
i += 1;
}
}
String::from_utf8(out).map_err(|_| {
McpError::Transport("MCP OAuth callback query contains invalid UTF-8".to_string())
})
}
fn hex_digit(byte: u8) -> Option<u8> {
match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
b'A'..=b'F' => Some(byte - b'A' + 10),
_ => None,
}
}
fn lock_error(error: anyhow::Error) -> McpError {
McpError::Transport(format!("failed to lock MCP OAuth token file: {error}"))
}
fn protected_resource_metadata_urls(origin_url: &str) -> McpResult<Vec<reqwest::Url>> {
let parsed = reqwest::Url::parse(origin_url)
.map_err(|_| McpError::Config("MCP OAuth server URL is invalid".to_string()))?;
let mut urls = Vec::new();
let mut root = parsed.clone();
root.set_path("/.well-known/oauth-protected-resource");
root.set_query(None);
root.set_fragment(None);
urls.push(root);
let endpoint_dir = endpoint_directory_path(parsed.path());
let endpoint_path = format!("{endpoint_dir}.well-known/oauth-protected-resource");
let mut endpoint = parsed;
endpoint.set_path(&endpoint_path);
endpoint.set_query(None);
endpoint.set_fragment(None);
if !urls.iter().any(|url| url == &endpoint) {
urls.push(endpoint);
}
Ok(urls)
}
fn authorization_server_metadata_urls(auth_server_url: &str) -> McpResult<Vec<reqwest::Url>> {
let parsed = reqwest::Url::parse(auth_server_url).map_err(|_| {
McpError::Config("MCP OAuth authorization server URL is invalid".to_string())
})?;
let mut urls = Vec::new();
let mut root = parsed.clone();
root.set_path("/.well-known/oauth-authorization-server");
root.set_query(None);
root.set_fragment(None);
urls.push(root);
let endpoint_dir = endpoint_directory_path(parsed.path());
let mut path_relative = parsed.clone();
path_relative.set_path(&format!(
"{endpoint_dir}.well-known/oauth-authorization-server"
));
path_relative.set_query(None);
path_relative.set_fragment(None);
if !urls.iter().any(|url| url == &path_relative) {
urls.push(path_relative);
}
let mut oidc = parsed;
oidc.set_path(&format!("{endpoint_dir}.well-known/openid-configuration"));
oidc.set_query(None);
oidc.set_fragment(None);
if !urls.iter().any(|url| url == &oidc) {
urls.push(oidc);
}
Ok(urls)
}
fn endpoint_directory_path(path: &str) -> String {
let trimmed = path.trim_end_matches('/');
match trimmed.rsplit_once('/') {
Some(("", _)) | None => "/".to_string(),
Some((parent, _)) => format!("{parent}/"),
}
}
fn validate_protected_resource_metadata(metadata: &ProtectedResourceMetadata) -> McpResult<()> {
if metadata.resource.trim().is_empty() {
return Err(McpError::Protocol {
code: -32602,
message: "MCP OAuth protected-resource metadata missing resource".to_string(),
});
}
Ok(())
}
fn validate_authorization_server_metadata(metadata: &AuthorizationServerMetadata) -> McpResult<()> {
if metadata.issuer.trim().is_empty() {
return Err(metadata_missing("issuer"));
}
if metadata.authorization_endpoint.trim().is_empty() {
return Err(metadata_missing("authorization_endpoint"));
}
if metadata.token_endpoint.trim().is_empty() {
return Err(metadata_missing("token_endpoint"));
}
validate_oauth_endpoint("authorization_endpoint", &metadata.authorization_endpoint)?;
validate_oauth_endpoint("token_endpoint", &metadata.token_endpoint)?;
if let Some(endpoint) = metadata.registration_endpoint.as_deref() {
validate_oauth_endpoint("registration_endpoint", endpoint)?;
}
if metadata
.code_challenge_methods_supported
.as_ref()
.is_some_and(|methods| !methods.iter().any(|method| method == "S256"))
{
return Err(McpError::Config(
"MCP OAuth authorization server does not support PKCE S256".to_string(),
));
}
Ok(())
}
fn validate_oauth_endpoint(field: &str, endpoint: &str) -> McpResult<()> {
crate::config::validate_mcp_http_url_field("oauth", field, endpoint)
.map_err(|error| McpError::Config(format!("MCP OAuth {field} is not permitted: {error}")))
}
fn metadata_missing(field: &str) -> McpError {
McpError::Protocol {
code: -32602,
message: format!("MCP OAuth authorization-server metadata missing {field}"),
}
}
fn random_urlsafe(bytes: usize) -> String {
let mut out = Vec::new();
while out.len() < bytes {
out.extend_from_slice(uuid::Uuid::new_v4().as_bytes());
}
out.truncate(bytes);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(out)
}
pub(crate) fn sanitize_url(url: &str) -> String {
match reqwest::Url::parse(url) {
Ok(mut parsed) => {
let _ = parsed.set_username("");
let _ = parsed.set_password(None);
parsed.set_query(None);
parsed.set_fragment(None);
let host = parsed.host_str().unwrap_or("<unknown>");
let port = parsed
.port()
.map(|port| format!(":{port}"))
.unwrap_or_default();
format!("{}://{}{}{}", parsed.scheme(), host, port, parsed.path())
}
Err(_) => "<invalid-url>".to_string(),
}
}
fn pct(input: &str) -> String {
let mut out = String::new();
for b in input.bytes() {
match b {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
out.push(char::from(b));
}
_ => out.push_str(&format!("%{b:02X}")),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use std::{
io::{Read, Write},
net::{TcpListener, TcpStream},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
mpsc,
},
thread,
time::Duration,
};
fn stored_token() -> StoredToken {
StoredToken {
client_id: "client-1".to_string(),
access_token: "fake-access-token".to_string(),
refresh_token: Some("fake-refresh-token".to_string()),
expires_at: Some(1_900_000_000),
granted_scopes: vec!["search".to_string()],
client_secret: None,
authorization_server: None,
issuer: None,
token_endpoint: None,
resource: Some("https://mcp.example.test/mcp".to_string()),
server_url: "https://mcp.example.test/mcp".to_string(),
token_received_at: Utc::now().timestamp(),
}
}
#[test]
fn oauth_debug_redacts_secret_fields() {
let token = TokenResponse {
access_token: "fake-access-token".to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(3600),
refresh_token: Some("fake-refresh-token".to_string()),
scope: Some("search".to_string()),
extra: BTreeMap::new(),
};
let registration = RegistrationResponse {
client_id: "client-1".to_string(),
client_secret: Some("fake-client-secret".to_string()),
extra: BTreeMap::new(),
};
let stored = stored_token();
let login = McpOAuthLoginOutcome {
authorization_url: "https://auth.example.test/authorize?code_challenge=fake-verifier&state=fake-state&client_secret=fake-client-secret".to_string(),
redirect_uri: "http://127.0.0.1:1234/mcp/oauth/callback".to_string(),
state: "fake-state".to_string(),
verifier: "fake-verifier".to_string(),
client_id: "client-1".to_string(),
client_secret: Some("fake-client-secret".to_string()),
metadata: AuthorizationServerMetadata {
issuer: "https://auth.example.test".to_string(),
authorization_endpoint: "https://auth.example.test/authorize".to_string(),
token_endpoint: "https://auth.example.test/token".to_string(),
registration_endpoint: None,
scopes_supported: None,
response_types_supported: None,
grant_types_supported: None,
code_challenge_methods_supported: None,
extra: BTreeMap::new(),
},
resource: "https://mcp.example.test/mcp".to_string(),
authorization_server: "https://auth.example.test".to_string(),
};
for debug in [
format!("{token:?}"),
format!("{registration:?}"),
format!("{stored:?}"),
format!("{login:?}"),
] {
assert!(debug.contains("[REDACTED]"), "{debug}");
assert!(!debug.contains("fake-access-token"), "{debug}");
assert!(!debug.contains("fake-refresh-token"), "{debug}");
assert!(!debug.contains("fake-client-secret"), "{debug}");
assert!(!debug.contains("fake-verifier"), "{debug}");
assert!(!debug.contains("fake-state"), "{debug}");
assert!(!debug.contains("code_challenge="), "{debug}");
}
}
#[test]
fn token_error_message_whitelists_oauth_error_codes() {
let malicious = token_error_message(
reqwest::StatusCode::BAD_REQUEST,
r#"{"error":"fake_access_token_secret_12345"}"#,
false,
);
assert_eq!(malicious, "OAuth token exchange failed");
assert!(!malicious.contains("fake_access_token_secret_12345"));
let known = token_error_message(
reqwest::StatusCode::BAD_REQUEST,
r#"{"error":"invalid_scope"}"#,
false,
);
assert!(known.contains("invalid_scope"), "{known}");
}
#[test]
fn pkce_verifier_charset_challenge_determinism_and_state_uniqueness() {
let verifier = generate_code_verifier();
assert!((43..=128).contains(&verifier.len()), "{verifier}");
assert!(verifier.bytes().all(|byte| matches!(
byte,
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~'
)));
assert_eq!(
code_challenge("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
);
assert_ne!(generate_state(), generate_state());
}
#[test]
fn authorization_url_percent_encodes_parameters() {
let url = authorization_url(
"https://auth.example.test/authorize",
"client id",
"http://127.0.0.1:1234/mcp/oauth/callback",
"challenge",
"state",
"https://mcp.example.test/mcp?a=b",
&["search".to_string(), "offline_access".to_string()],
);
assert!(url.contains("client_id=client%20id"), "{url}");
assert!(
url.contains("redirect_uri=http%3A%2F%2F127.0.0.1%3A1234%2Fmcp%2Foauth%2Fcallback"),
"{url}"
);
assert!(
url.contains("resource=https%3A%2F%2Fmcp.example.test%2Fmcp%3Fa%3Db"),
"{url}"
);
assert!(url.contains("scope=search%20offline_access"), "{url}");
}
#[test]
fn metadata_deserialization_tolerates_unknown_fields() {
let protected: ProtectedResourceMetadata = serde_json::from_str(
r#"{"resource":"https://mcp.example.test/mcp","authorization_servers":["https://auth.example.test"],"future":true}"#,
)
.unwrap();
assert_eq!(protected.extra["future"], true);
let auth: AuthorizationServerMetadata = serde_json::from_str(
r#"{"issuer":"https://auth.example.test","authorization_endpoint":"https://auth.example.test/authorize","token_endpoint":"https://auth.example.test/token","registration_endpoint":"https://auth.example.test/register","code_challenge_methods_supported":["S256"],"future":"kept"}"#,
)
.unwrap();
assert_eq!(auth.extra["future"], "kept");
assert!(format!("{auth:?}").contains("authorization_endpoint"));
}
#[test]
fn select_authorization_server_prefers_configured_then_metadata() {
let protected = ProtectedResourceMetadata {
resource: "https://mcp.example.test/mcp".to_string(),
authorization_servers: vec!["https://auth.example.test".to_string()],
extra: BTreeMap::new(),
};
assert_eq!(
select_authorization_server(Some("https://auth.example.test"), Some(&protected))
.unwrap(),
"https://auth.example.test"
);
assert!(
select_authorization_server(Some("https://override.example.test"), Some(&protected))
.is_err()
);
assert_eq!(
select_authorization_server(None, Some(&protected)).unwrap(),
"https://auth.example.test"
);
assert!(select_authorization_server(None, None).is_err());
}
#[test]
fn www_authenticate_resource_metadata_parser_extracts_url() {
assert_eq!(
parse_www_authenticate_resource_metadata(
r#"Bearer realm="mcp", resource_metadata="https://mcp.example.test/.well-known/oauth-protected-resource""#,
)
.as_deref(),
Some("https://mcp.example.test/.well-known/oauth-protected-resource")
);
}
#[test]
fn token_storage_round_trip_permissions_delete_and_url_validation() {
let temp = tempfile::TempDir::new().unwrap();
let token = stored_token();
write_token(temp.path(), "remote", &token).unwrap();
let path = token_file_path(temp.path(), "remote").unwrap();
assert!(path.exists());
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mode = fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600);
}
let read = read_token(temp.path(), "remote").unwrap().unwrap();
assert_eq!(read, token);
assert!(validate_token_url(&read, "https://mcp.example.test/mcp"));
assert!(!validate_token_url(&read, "https://mcp.example.test/other"));
delete_token(temp.path(), "remote").unwrap();
assert!(!path.exists());
delete_token(temp.path(), "remote").unwrap();
}
#[test]
fn token_storage_nonexistent_and_corrupt_file() {
let temp = tempfile::TempDir::new().unwrap();
assert!(read_token(temp.path(), "missing").unwrap().is_none());
let path = token_file_path(temp.path(), "bad").unwrap();
fs::create_dir_all(path.parent().unwrap()).unwrap();
fs::write(&path, "not json").unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mut permissions = fs::metadata(&path).unwrap().permissions();
permissions.set_mode(0o600);
fs::set_permissions(&path, permissions).unwrap();
}
let error = read_token(temp.path(), "bad").unwrap_err().to_string();
assert!(error.contains("corrupt"), "{error}");
}
#[test]
fn token_storage_rejects_path_traversal_server_names() {
let temp = tempfile::TempDir::new().unwrap();
let token = stored_token();
for malicious in ["../auth", "..", ".", "a/b", "a\\b", "a:b", "a b"] {
let err = write_token(temp.path(), malicious, &token)
.unwrap_err()
.to_string();
assert!(
err.contains("not a valid token store name"),
"{malicious}: {err}"
);
assert!(read_token(temp.path(), malicious).is_err());
assert!(delete_token(temp.path(), malicious).is_err());
}
assert!(!temp.path().join("auth.json").exists());
write_token(temp.path(), "my-server_1", &token).unwrap();
assert!(read_token(temp.path(), "my-server_1").unwrap().is_some());
}
#[test]
fn token_read_rejects_symlinked_token_file() {
let temp = tempfile::TempDir::new().unwrap();
let token = stored_token();
write_token(temp.path(), "remote", &token).unwrap();
let real_path = token_file_path(temp.path(), "remote").unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::symlink;
let link_path = token_file_path(temp.path(), "linked").unwrap();
fs::create_dir_all(link_path.parent().unwrap()).unwrap();
symlink(&real_path, &link_path).unwrap();
let err = read_token(temp.path(), "linked").unwrap_err().to_string();
assert!(
err.contains("symlink") || err.contains("regular private file"),
"{err}"
);
assert!(!err.contains("fake-access-token"), "{err}");
}
}
#[test]
fn discovery_and_registration_use_loopback_http() {
let protected_body = r#"{"resource":"http://127.0.0.1/mcp","authorization_servers":["http://127.0.0.1/auth"],"future":true}"#;
let protected_url =
serve_once("/.well-known/oauth-protected-resource", protected_body, 200);
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let protected = discover_protected_resource(&format!("{protected_url}/mcp"), &client)
.unwrap()
.unwrap();
assert_eq!(protected.resource, "http://127.0.0.1/mcp");
let auth_body = r#"{"issuer":"http://127.0.0.1/auth","authorization_endpoint":"https://auth.example.test/authorize","token_endpoint":"https://auth.example.test/token","registration_endpoint":"http://127.0.0.1/register","code_challenge_methods_supported":["S256"]}"#;
let auth_url = serve_once("/.well-known/oauth-authorization-server", auth_body, 200);
let auth = discover_authorization_server(&auth_url, &client).unwrap();
assert_eq!(auth.token_endpoint, "https://auth.example.test/token");
let (registration_url, body_rx) = serve_once_with_body(
"/register",
r#"{"client_id":"registered-client","client_secret":"fake-client-secret"}"#,
201,
);
let registration = register_client(
&format!("{registration_url}/register"),
vec!["http://127.0.0.1:1234/mcp/oauth/callback".to_string()],
&client,
)
.unwrap();
assert_eq!(registration.client_id, "registered-client");
let body = body_rx.recv_timeout(Duration::from_secs(2)).unwrap();
assert!(
body.contains("\"token_endpoint_auth_method\":\"none\""),
"{body}"
);
assert!(!format!("{registration:?}").contains("fake-client-secret"));
}
#[test]
fn token_exchange_redirect_does_not_send_credentials_to_second_origin() {
let target = TcpListener::bind("127.0.0.1:0").unwrap();
target.set_nonblocking(true).unwrap();
let target_url = format!("http://{}/token", target.local_addr().unwrap());
let origin = TcpListener::bind("127.0.0.1:0").unwrap();
origin.set_nonblocking(true).unwrap();
let origin_url = format!("http://{}/token", origin.local_addr().unwrap());
let origin_thread = thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(2);
loop {
match origin.accept() {
Ok((mut stream, _)) => {
let mut buffer = [0_u8; 4096];
let n = stream.read(&mut buffer).unwrap();
let request = String::from_utf8_lossy(&buffer[..n]);
assert!(request.contains("code=authorization-code"), "{request}");
assert!(request.contains("client_secret=client-secret"), "{request}");
loopback_write_response_with_headers(
&mut stream,
"302 Found",
"text/plain",
"redirect",
&[&format!("Location: {target_url}")],
);
break;
}
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
if Instant::now() >= deadline {
panic!("redirect origin received no request");
}
thread::sleep(Duration::from_millis(10));
}
Err(error) => panic!("redirect origin accept failed: {error}"),
}
}
});
let client = reqwest::blocking::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.timeout(Duration::from_secs(1))
.build()
.unwrap();
let error = exchange_code(
&origin_url,
"authorization-code",
"http://127.0.0.1:1234/mcp/oauth/callback",
"client-1",
"verifier",
"http://127.0.0.1/mcp",
Some("client-secret"),
&client,
)
.unwrap_err()
.to_string();
assert!(error.contains("HTTP 302"), "{error}");
origin_thread.join().unwrap();
thread::sleep(Duration::from_millis(100));
assert!(matches!(target.accept(), Err(error) if error.kind() == io::ErrorKind::WouldBlock));
}
#[test]
fn token_exchange_refresh_and_callback_validation_are_redacted() {
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let (token_url, body_rx) = serve_once_with_body(
"/token",
r#"{"access_token":"fake-access-token","token_type":"Bearer","expires_in":3600,"refresh_token":"fake-refresh-token","scope":"search"}"#,
200,
);
let token = exchange_code(
&format!("{token_url}/token"),
"fake-code",
"http://127.0.0.1:1234/mcp/oauth/callback",
"client-1",
"verifier",
"http://127.0.0.1/mcp",
None,
&client,
)
.unwrap();
assert_eq!(token.access_token, "fake-access-token");
assert!(
body_rx
.recv_timeout(Duration::from_secs(2))
.unwrap()
.contains("grant_type=authorization_code")
);
let (refresh_url, refresh_rx) = serve_once_with_body(
"/token",
r#"{"access_token":"new-access-token","token_type":"Bearer","expires_in":3600}"#,
200,
);
let refreshed = refresh_token(
&format!("{refresh_url}/token"),
"fake-refresh-token",
"client-1",
"http://127.0.0.1/mcp",
None,
&client,
)
.unwrap();
assert_eq!(refreshed.access_token, "new-access-token");
assert!(
refresh_rx
.recv_timeout(Duration::from_secs(2))
.unwrap()
.contains("grant_type=refresh_token")
);
assert_eq!(
parse_callback_request_line(
"GET /mcp/oauth/callback?code=callback-code&state=expected HTTP/1.1",
"expected",
)
.unwrap(),
"callback-code"
);
let error = parse_callback_request_line(
"GET /mcp/oauth/callback?code=callback-code&state=wrong HTTP/1.1",
"expected",
)
.unwrap_err()
.to_string();
assert!(error.contains("state mismatch"), "{error}");
assert!(!error.contains("callback-code"), "{error}");
}
#[test]
fn token_exchange_rejects_oversized_success_body() {
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let body = "x".repeat(DEFAULT_BOUNDED_BODY_MAX_BYTES as usize + 1);
let body: &'static str = Box::leak(body.into_boxed_str());
let (token_url, _body_rx) = serve_once_with_body("/token", body, 200);
let error = exchange_code(
&format!("{token_url}/token"),
"fake-code",
"http://127.0.0.1:1234/mcp/oauth/callback",
"client-1",
"verifier",
"http://127.0.0.1/mcp",
None,
&client,
)
.unwrap_err()
.to_string();
assert!(
error.contains("MCP OAuth token response exceeded"),
"{error}"
);
}
#[test]
fn discovery_handles_404_as_unavailable() {
let url = serve_once("/.well-known/oauth-protected-resource", "not found", 404);
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
assert!(
discover_protected_resource(&format!("{url}/mcp"), &client)
.unwrap()
.is_none()
);
}
#[test]
fn discovery_uses_endpoint_path_protected_resource_metadata() {
let server = DiscoveryServer::start(vec![
("/.well-known/oauth-protected-resource", 404, "not found"),
(
"/api/.well-known/oauth-protected-resource",
200,
r#"{"resource":"resource-from-path","authorization_servers":["http://auth.example.test"]}"#,
),
]);
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let metadata = discover_protected_resource(&format!("{}/api/mcp", server.base), &client)
.unwrap()
.unwrap();
assert_eq!(metadata.resource, "resource-from-path");
}
#[test]
fn discovery_uses_oidc_authorization_server_fallback() {
let server = DiscoveryServer::start(vec![
("/.well-known/oauth-authorization-server", 404, "not found"),
(
"/.well-known/openid-configuration",
200,
r#"{"issuer":"http://issuer.example.test","authorization_endpoint":"https://issuer.example.test/authorize","token_endpoint":"https://issuer.example.test/token","code_challenge_methods_supported":["S256"]}"#,
),
]);
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let metadata = discover_authorization_server(&server.base, &client).unwrap();
assert_eq!(metadata.issuer, "http://issuer.example.test");
}
#[test]
fn authorization_server_metadata_rejects_non_https_endpoints() {
let metadata = AuthorizationServerMetadata {
issuer: "https://auth.example.test".to_string(),
authorization_endpoint: "http://auth.example.test/authorize".to_string(),
token_endpoint: "https://auth.example.test/token".to_string(),
registration_endpoint: None,
scopes_supported: None,
response_types_supported: None,
grant_types_supported: None,
code_challenge_methods_supported: Some(vec!["S256".to_string()]),
extra: BTreeMap::new(),
};
let error = validate_authorization_server_metadata(&metadata)
.unwrap_err()
.to_string();
assert!(
error.contains("authorization_endpoint must use https"),
"{error}"
);
}
#[test]
fn stored_token_from_response_sets_checked_expiry() {
let token = stored_token_from_response(
TokenResponse {
access_token: "access".to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(3600),
refresh_token: None,
scope: None,
extra: BTreeMap::new(),
},
"client-1",
None,
"https://mcp.example.test/mcp",
None,
None,
None,
None,
)
.unwrap();
assert!(token.expires_at.unwrap() >= token.token_received_at + 3600);
}
#[test]
fn stored_token_from_response_rejects_expiry_overflow() {
match stored_token_from_response(
TokenResponse {
access_token: "access".to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(u64::MAX),
refresh_token: None,
scope: None,
extra: BTreeMap::new(),
},
"client-1",
None,
"https://mcp.example.test/mcp",
None,
None,
None,
None,
) {
Err(McpError::Protocol { code, message }) => {
assert_eq!(code, -32602);
assert_eq!(message, "MCP OAuth token response expires_in is too large");
}
other => panic!("expected protocol error for expires_in overflow, got {other:?}"),
}
}
#[test]
fn concurrent_force_refreshes_keep_token_file_updates_consistent() {
let temp = tempfile::TempDir::new().unwrap();
let token_server = CountingTokenServer::start();
let mut token = stored_token();
token.expires_at = Some(Utc::now().timestamp() - 60);
token.token_endpoint = Some(format!("{}/token", token_server.base));
token.resource = None;
write_token(temp.path(), "remote", &token).unwrap();
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let oauth = McpOAuthConfig {
client_id: Some("client-1".to_string()),
scopes: Vec::new(),
authorization_server: None,
};
let provider_a = TokenProvider::new(
temp.path().to_path_buf(),
"remote".to_string(),
token.server_url.clone(),
oauth.clone(),
client.clone(),
);
let provider_b = TokenProvider::new(
temp.path().to_path_buf(),
"remote".to_string(),
token.server_url.clone(),
oauth,
client,
);
let a = thread::spawn(move || provider_a.force_refresh_access_token().unwrap());
let b = thread::spawn(move || provider_b.force_refresh_access_token().unwrap());
let first = a.join().unwrap();
let second = b.join().unwrap();
assert!(first.starts_with("new-access-token-"), "{first}");
assert!(second.starts_with("new-access-token-"), "{second}");
let refresh_count = *token_server.count.lock().unwrap();
assert!((2..=3).contains(&refresh_count), "{refresh_count}");
assert!(
read_token(temp.path(), "remote")
.unwrap()
.unwrap()
.access_token
.starts_with("new-access-token-")
);
}
#[test]
fn refresh_releases_token_file_lock_during_network_request() {
let temp = tempfile::TempDir::new().unwrap();
let (request_started_tx, request_started_rx) = mpsc::channel();
let (release_response_tx, release_response_rx) = mpsc::channel();
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let server = thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buffer = [0_u8; 4096];
let n = stream.read(&mut buffer).unwrap();
let request = String::from_utf8_lossy(&buffer[..n]);
assert!(request.contains("grant_type=refresh_token"), "{request}");
request_started_tx.send(()).unwrap();
release_response_rx
.recv_timeout(Duration::from_secs(2))
.unwrap();
loopback_write_response(
&mut stream,
"200 OK",
"application/json",
r#"{"access_token":"new-access-token","token_type":"Bearer","expires_in":3600}"#,
);
});
let mut token = stored_token();
token.expires_at = Some(Utc::now().timestamp() - 60);
token.token_endpoint = Some(format!("{base}/token"));
token.resource = None;
write_token(temp.path(), "remote", &token).unwrap();
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let provider = TokenProvider::new(
temp.path().to_path_buf(),
"remote".to_string(),
token.server_url.clone(),
McpOAuthConfig {
client_id: Some("client-1".to_string()),
scopes: Vec::new(),
authorization_server: None,
},
client,
);
let refresh = thread::spawn(move || provider.access_token().unwrap());
request_started_rx
.recv_timeout(Duration::from_secs(2))
.unwrap();
let token_path = token_file_path(temp.path(), "remote").unwrap();
let (lock_acquired_tx, lock_acquired_rx) = mpsc::channel();
let lock = thread::spawn(move || {
let lock = CrossProcessFileLock::acquire(&token_path).unwrap();
lock_acquired_tx.send(()).unwrap();
lock
});
lock_acquired_rx
.recv_timeout(Duration::from_secs(1))
.expect("token file lock should be available while refresh request waits");
let held_lock = lock.join().unwrap();
drop(held_lock);
release_response_tx.send(()).unwrap();
assert_eq!(refresh.join().unwrap(), "new-access-token");
server.join().unwrap();
}
struct DiscoveryServer {
base: String,
handle: Option<thread::JoinHandle<()>>,
}
impl DiscoveryServer {
fn start(routes: Vec<(&'static str, u16, &'static str)>) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let handle = thread::spawn(move || {
for _ in 0..routes.len() {
let (mut stream, _) = listener.accept().unwrap();
let mut buffer = [0_u8; 4096];
let n = stream.read(&mut buffer).unwrap();
let request = String::from_utf8_lossy(&buffer[..n]);
let path = request
.lines()
.next()
.and_then(|line| line.split_whitespace().nth(1))
.unwrap_or("");
let (status, body) = routes
.iter()
.find(|(route, _, _)| *route == path)
.map(|(_, status, body)| (*status, *body))
.unwrap_or((404, "not found"));
let reason = if status == 200 { "OK" } else { "Not Found" };
let response = format!(
"HTTP/1.1 {status} {reason}
Content-Type: application/json
Content-Length: {}
Connection: close
{body}",
body.len()
);
stream.write_all(response.as_bytes()).unwrap();
}
});
Self {
base,
handle: Some(handle),
}
}
}
impl Drop for DiscoveryServer {
fn drop(&mut self) {
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
struct CountingTokenServer {
base: String,
count: Arc<Mutex<usize>>,
stop: Arc<AtomicBool>,
handle: Option<thread::JoinHandle<()>>,
}
impl CountingTokenServer {
fn start() -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let count = Arc::new(Mutex::new(0));
let thread_count = Arc::clone(&count);
let stop = Arc::new(AtomicBool::new(false));
let thread_stop = Arc::clone(&stop);
let handle = thread::spawn(move || {
while !thread_stop.load(Ordering::SeqCst) {
match listener.accept() {
Ok((mut stream, _)) => {
stream.set_nonblocking(false).unwrap();
let mut buffer = [0_u8; 4096];
let n = stream.read(&mut buffer).unwrap();
let request = String::from_utf8_lossy(&buffer[..n]);
assert!(request.contains("grant_type=refresh_token"), "{request}");
let mut count = thread_count.lock().unwrap();
*count += 1;
let body = format!(
r#"{{"access_token":"new-access-token-{}","token_type":"Bearer","expires_in":3600}}"#,
*count
);
loopback_write_response(
&mut stream,
"200 OK",
"application/json",
&body,
);
}
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(10));
}
Err(_) => break,
}
}
});
Self {
base,
count,
stop,
handle: Some(handle),
}
}
}
impl Drop for CountingTokenServer {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
if let Ok(url) = reqwest::Url::parse(&self.base)
&& let Ok(mut addrs) = url.socket_addrs(|| None)
&& let Some(addr) = addrs.pop()
{
let _ = std::net::TcpStream::connect(addr);
}
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
#[derive(Clone, Copy)]
enum LoopbackMode {
Normal,
BadState,
TokenRevoked,
}
struct LoopbackOAuthMcpServer {
url: String,
stop: Arc<AtomicBool>,
handle: Option<thread::JoinHandle<()>>,
}
impl LoopbackOAuthMcpServer {
fn start(mode: LoopbackMode) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let addr = listener.local_addr().unwrap();
let base = format!("http://{addr}");
let stop = Arc::new(AtomicBool::new(false));
let thread_stop = Arc::clone(&stop);
let thread_base = base.clone();
let handle = thread::spawn(move || {
while !thread_stop.load(Ordering::SeqCst) {
match listener.accept() {
Ok((mut stream, _)) => {
handle_loopback_request(&mut stream, &thread_base, mode);
}
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(10));
}
Err(_) => break,
}
}
});
Self {
url: format!("{base}/mcp"),
stop,
handle: Some(handle),
}
}
}
impl Drop for LoopbackOAuthMcpServer {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
if let Ok(parsed) = reqwest::Url::parse(&self.url)
&& let Ok(mut addrs) = parsed.socket_addrs(|| None)
&& let Some(addr) = addrs.pop()
{
let _ = TcpStream::connect(addr);
}
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
fn loopback_write_response(
stream: &mut TcpStream,
status: &str,
content_type: &str,
body: &str,
) {
loopback_write_response_with_headers(stream, status, content_type, body, &[]);
}
fn loopback_write_response_with_headers(
stream: &mut TcpStream,
status: &str,
content_type: &str,
body: &str,
extra: &[&str],
) {
let mut response = format!(
"HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n",
body.len()
);
for header in extra {
response.push_str(header);
response.push_str("\r\n");
}
response.push_str("\r\n");
response.push_str(body);
stream.write_all(response.as_bytes()).unwrap();
}
fn handle_loopback_request(stream: &mut TcpStream, base: &str, mode: LoopbackMode) {
let mut buf = [0_u8; 8192];
let n = stream.read(&mut buf).unwrap_or(0);
if n == 0 {
return;
}
let request = String::from_utf8_lossy(&buf[..n]);
let head = request.split("\r\n\r\n").next().unwrap_or("");
let body = request.split("\r\n\r\n").nth(1).unwrap_or("");
let first = head.lines().next().unwrap_or("");
let path = first.split_whitespace().nth(1).unwrap_or("");
match path.split('?').next().unwrap_or(path) {
"/.well-known/oauth-protected-resource" => loopback_write_response(
stream,
"200 OK",
"application/json",
&format!(r#"{{"resource":"{base}/mcp","authorization_servers":["{base}"]}}"#),
),
"/.well-known/oauth-authorization-server" => loopback_write_response(
stream,
"200 OK",
"application/json",
&format!(
r#"{{"issuer":"{base}","authorization_endpoint":"{base}/authorize","token_endpoint":"{base}/token","registration_endpoint":"{base}/register","scopes_supported":["read","tools"],"code_challenge_methods_supported":["S256"]}}"#
),
),
"/register" => loopback_write_response(
stream,
"201 Created",
"application/json",
r#"{"client_id":"loopback-client","client_secret":"fake_client_secret_abc"}"#,
),
"/authorize" => {
let params = path
.split_once('?')
.map(|(_, q)| parse_query(q).unwrap())
.unwrap_or_default();
let redirect_uri = params
.iter()
.find(|(k, _)| k == "redirect_uri")
.unwrap()
.1
.clone();
let state = params.iter().find(|(k, _)| k == "state").unwrap().1.clone();
let returned_state = if matches!(mode, LoopbackMode::BadState) {
"wrong_state"
} else {
&state
};
let location = format!(
"{redirect_uri}?code=fake_authorization_code_123&state={returned_state}"
);
loopback_write_response_with_headers(
stream,
"302 Found",
"text/plain",
"",
&[&format!("Location: {location}")],
);
}
"/token" if body.contains("grant_type=authorization_code") => loopback_write_response(
stream,
"200 OK",
"application/json",
r#"{"access_token":"fake_access_token_12345","refresh_token":"fake_refresh_abc","token_type":"Bearer","expires_in":3600,"scope":"read tools"}"#,
),
"/token"
if body.contains("grant_type=refresh_token")
&& matches!(mode, LoopbackMode::TokenRevoked) =>
{
loopback_write_response(
stream,
"400 Bad Request",
"application/json",
r#"{"error":"invalid_grant","error_description":"fake_refresh_abc revoked"}"#,
)
}
"/token" if body.contains("grant_type=refresh_token") => loopback_write_response(
stream,
"200 OK",
"application/json",
r#"{"access_token":"fake_access_token_refreshed","token_type":"Bearer","expires_in":3600}"#,
),
"/mcp" => {
if first.starts_with("DELETE ") {
loopback_write_response(stream, "204 No Content", "text/plain", "");
return;
}
if first.starts_with("GET ") {
loopback_write_response(
stream,
"405 Method Not Allowed",
"text/plain",
"no stream",
);
return;
}
let lower_head = head.to_ascii_lowercase();
if !lower_head.contains("authorization: bearer fake_access_token_12345")
&& !lower_head.contains("authorization: bearer fake_access_token_refreshed")
{
loopback_write_response_with_headers(
stream,
"401 Unauthorized",
"text/plain",
"unauthorized",
&[&format!(
"WWW-Authenticate: Bearer resource_metadata=\"{base}/.well-known/oauth-protected-resource\""
)],
);
return;
}
let value: Value = serde_json::from_str(body).unwrap();
let id = value.get("id").cloned().unwrap_or(serde_json::json!(1));
match value.get("method").and_then(Value::as_str).unwrap_or("") {
"initialize" => loopback_write_response(
stream,
"200 OK",
"application/json",
&format!(
r#"{{"jsonrpc":"2.0","id":{id},"result":{{"protocolVersion":"2025-03-26","capabilities":{{"tools":{{}}}},"serverInfo":{{"name":"loopback-oauth","version":"1"}}}}}}"#
),
),
"notifications/initialized" => {
loopback_write_response(stream, "204 No Content", "text/plain", "")
}
"tools/list" => loopback_write_response(
stream,
"200 OK",
"application/json",
&format!(
r#"{{"jsonrpc":"2.0","id":{id},"result":{{"tools":[{{"name":"echo","description":"Echo","inputSchema":{{"type":"object"}}}}]}}}}"#
),
),
"tools/call" => loopback_write_response(
stream,
"200 OK",
"application/json",
&format!(
r#"{{"jsonrpc":"2.0","id":{id},"result":{{"content":[{{"type":"text","text":"ok"}}],"isError":false}}}}"#
),
),
_ => loopback_write_response(stream, "404 Not Found", "text/plain", "missing"),
}
}
_ => loopback_write_response(stream, "404 Not Found", "text/plain", "missing"),
}
}
#[test]
fn loopback_oauth_fixture_full_login_mcp_request_logout_and_secret_non_leak() {
let server = LoopbackOAuthMcpServer::start(LoopbackMode::Normal);
let temp = tempfile::TempDir::new().unwrap();
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.redirect(reqwest::redirect::Policy::none())
.build()
.unwrap();
let redirect_uri = "http://127.0.0.1:7777/mcp/oauth/callback";
let oauth = McpOAuthConfig {
client_id: None,
scopes: vec!["read".to_string(), "tools".to_string()],
authorization_server: None,
};
let login = prepare_login(&server.url, &oauth, redirect_uri, &client).unwrap();
let redirect = client
.get(&login.authorization_url)
.send()
.unwrap()
.headers()
.get(reqwest::header::LOCATION)
.unwrap()
.to_str()
.unwrap()
.to_string();
let code = parse_manual_fallback_input(&redirect, &login.state).unwrap();
let token_response = exchange_code(
&login.metadata.token_endpoint,
&code,
&login.redirect_uri,
&login.client_id,
&login.verifier,
&login.resource,
login.client_secret.as_deref(),
&client,
)
.unwrap();
let stored = stored_token_from_response(
token_response,
&login.client_id,
login.client_secret,
&server.url,
Some(login.authorization_server),
Some(login.metadata.issuer),
Some(login.metadata.token_endpoint),
Some(login.resource),
)
.unwrap();
write_token(temp.path(), "remote", &stored).unwrap();
let config = crate::config::McpServerConfig::Http(crate::config::McpHttpServerConfig {
url: server.url.clone(),
headers: BTreeMap::new(),
oauth: Some(oauth),
enabled: true,
timeout: Some(5),
});
let mut mcp =
crate::mcp::McpClient::connect_named(Some("remote"), &config, Some(temp.path()))
.unwrap();
mcp.send_request("initialize", Some(serde_json::json!({})))
.unwrap();
let tools = mcp.list_tools().unwrap();
assert_eq!(tools[0].name, "echo");
mcp.shutdown();
delete_token(temp.path(), "remote").unwrap();
assert!(read_token(temp.path(), "remote").unwrap().is_none());
let combined = format!("{stored:?}");
for secret in [
"fake_access_token_12345",
"fake_refresh_abc",
"fake_client_secret_abc",
"fake_authorization_code_123",
] {
assert!(!combined.contains(secret), "{combined}");
}
}
#[test]
fn loopback_oauth_fixture_bad_state_and_revoked_refresh_are_actionable_and_redacted() {
let bad = LoopbackOAuthMcpServer::start(LoopbackMode::BadState);
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(5))
.redirect(reqwest::redirect::Policy::none())
.build()
.unwrap();
let oauth = McpOAuthConfig {
client_id: None,
scopes: Vec::new(),
authorization_server: None,
};
let login = prepare_login(
&bad.url,
&oauth,
"http://127.0.0.1:7777/mcp/oauth/callback",
&client,
)
.unwrap();
let redirect = client
.get(&login.authorization_url)
.send()
.unwrap()
.headers()
.get(reqwest::header::LOCATION)
.unwrap()
.to_str()
.unwrap()
.to_string();
let error = parse_manual_fallback_input(&redirect, &login.state)
.unwrap_err()
.to_string();
assert!(error.contains("state mismatch"), "{error}");
assert!(!error.contains("fake_authorization_code_123"), "{error}");
let revoked = LoopbackOAuthMcpServer::start(LoopbackMode::TokenRevoked);
let token_endpoint = format!("{}/token", revoked.url.trim_end_matches("/mcp"));
let err = refresh_token(
&token_endpoint,
"fake_refresh_abc",
"client",
&revoked.url,
None,
&client,
)
.unwrap_err()
.to_string();
assert!(err.contains("refresh token expired or revoked"), "{err}");
assert!(!err.contains("fake_refresh_abc"), "{err}");
}
fn serve_once(expected_path: &'static str, body: &'static str, status: u16) -> String {
serve_once_with_body(expected_path, body, status).0
}
fn serve_once_with_body(
expected_path: &'static str,
body: &'static str,
status: u16,
) -> (String, mpsc::Receiver<String>) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = mpsc::channel();
thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buffer = [0_u8; 4096];
let n = stream.read(&mut buffer).unwrap();
let request = String::from_utf8_lossy(&buffer[..n]);
let first_line = request.lines().next().unwrap_or_default();
assert!(first_line.contains(expected_path), "{first_line}");
if let Some((_, body)) = request.split_once("\r\n\r\n") {
let _ = tx.send(body.to_string());
}
let reason = if status == 200 || status == 201 {
"OK"
} else {
"Not Found"
};
let response = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
stream.write_all(response.as_bytes()).unwrap();
});
(format!("http://{addr}"), rx)
}
}