Make hosted OAuth and MCP auth generic (#1375)

* Make hosted OAuth and MCP auth generic

* Address PR feedback and lint issues

* Suppress built-in Google secret in hosted proxy flows

* Align hosted OAuth secret suppression with proxy config

* Harden hosted OAuth callback helpers

* Tighten hosted OAuth URL rewriting
This commit is contained in:
Henry Park
2026-03-19 15:50:54 -07:00
committed by GitHub
parent 65062f3cc0
commit c4ab382522
8 changed files with 785 additions and 232 deletions
+355 -109
View File
@@ -45,6 +45,56 @@ struct PendingAuth {
task_handle: Option<tokio::task::JoinHandle<()>>,
}
struct HostedOAuthFlowStart {
name: String,
kind: ExtensionKind,
auth_url: String,
expected_state: String,
flow: crate::cli::oauth_defaults::PendingOAuthFlow,
}
fn hosted_proxy_client_secret(
client_secret: &Option<String>,
builtin: Option<&crate::cli::oauth_defaults::OAuthCredentials>,
exchange_proxy_configured: bool,
) -> Option<String> {
if !exchange_proxy_configured {
return client_secret.clone();
}
let builtin_secret = builtin.map(|credentials| credentials.client_secret);
match (client_secret, builtin_secret) {
(Some(resolved), Some(baked_in)) if resolved == baked_in => None,
_ => client_secret.clone(),
}
}
fn normalize_oauth_callback_path(path: &str) -> String {
let trimmed_path = path.trim_end_matches('/');
if trimmed_path.is_empty() {
"/oauth/callback".to_string()
} else if trimmed_path.ends_with("/oauth/callback") {
trimmed_path.to_string()
} else {
format!("{trimmed_path}/oauth/callback")
}
}
fn normalize_hosted_callback_url(callback_url: &str) -> String {
if let Ok(mut parsed) = url::Url::parse(callback_url) {
let normalized_path = normalize_oauth_callback_path(parsed.path());
parsed.set_path(&normalized_path);
return parsed.to_string();
}
let normalized_callback_url = callback_url.trim_end_matches('/');
if normalized_callback_url.ends_with("/oauth/callback") {
normalized_callback_url.to_string()
} else {
format!("{normalized_callback_url}/oauth/callback")
}
}
/// Runtime infrastructure needed for hot-activating WASM channels.
///
/// Set after construction via [`ExtensionManager::set_channel_runtime`] once the
@@ -547,7 +597,9 @@ impl ExtensionManager {
async fn gateway_callback_redirect_uri(&self) -> Option<String> {
use crate::cli::oauth_defaults;
if oauth_defaults::use_gateway_callback() {
return Some(format!("{}/oauth/callback", oauth_defaults::callback_url()));
return Some(normalize_hosted_callback_url(
&oauth_defaults::callback_url(),
));
}
// Use gateway_base_url from enable_gateway_mode()
if let Some(ref base) = *self.gateway_base_url.read().await {
@@ -924,6 +976,98 @@ impl ExtensionManager {
&self.pending_oauth_flows
}
async fn clear_pending_extension_auth(&self, name: &str) {
{
let mut pending = self.pending_auth.write().await;
if let Some(old) = pending.remove(name)
&& let Some(handle) = old.task_handle
{
handle.abort();
}
}
let mut flows = self.pending_oauth_flows.write().await;
flows.retain(|_, flow| flow.extension_name != name);
}
fn rewrite_oauth_state_param(
auth_url: String,
expected_state: &str,
hosted_state: &str,
) -> String {
if hosted_state == expected_state {
return auth_url;
}
let Ok(mut parsed) = url::Url::parse(&auth_url) else {
return auth_url.replace(
&format!("state={}", urlencoding::encode(expected_state)),
&format!("state={}", urlencoding::encode(hosted_state)),
);
};
let mut replaced = false;
let pairs: Vec<(String, String)> = parsed
.query_pairs()
.map(|(key, value)| {
if key == "state" {
replaced = true;
(key.into_owned(), hosted_state.to_string())
} else {
(key.into_owned(), value.into_owned())
}
})
.collect();
{
let mut query_pairs = parsed.query_pairs_mut();
query_pairs.clear();
for (key, value) in pairs {
query_pairs.append_pair(&key, &value);
}
if !replaced {
query_pairs.append_pair("state", hosted_state);
}
}
parsed.to_string()
}
async fn start_gateway_oauth_flow(&self, request: HostedOAuthFlowStart) -> AuthResult {
use crate::cli::oauth_defaults;
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await;
let hosted_state = oauth_defaults::build_platform_state(&request.expected_state);
let auth_url = Self::rewrite_oauth_state_param(
request.auth_url,
&request.expected_state,
&hosted_state,
);
self.pending_oauth_flows
.write()
.await
.insert(request.expected_state, request.flow);
self.pending_auth.write().await.insert(
request.name.clone(),
PendingAuth {
_name: request.name.clone(),
_kind: request.kind,
created_at: std::time::Instant::now(),
task_handle: None,
},
);
AuthResult::awaiting_authorization(
request.name,
request.kind,
auth_url,
"gateway".to_string(),
)
}
/// Broadcast an extension status change to the web UI via SSE.
async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) {
if let Some(ref sender) = *self.sse_sender.read().await {
@@ -2383,6 +2527,7 @@ impl ExtensionManager {
use crate::cli::oauth_defaults;
let is_gateway = self.should_use_gateway_mode();
self.clear_pending_extension_auth(name).await;
// Build redirect URI: gateway uses the public callback URL,
// local mode binds a random port.
@@ -2440,19 +2585,8 @@ impl ExtensionManager {
let code_verifier = oauth_result.code_verifier;
if is_gateway {
// Gateway mode: store pending flow for the /oauth/callback handler.
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await;
// Platform routing: prepend instance name to state
let platform_state = oauth_defaults::build_platform_state(&expected_state);
let auth_url = if platform_state != expected_state {
oauth_result.url.replace(
&format!("state={}", urlencoding::encode(&expected_state)),
&format!("state={}", urlencoding::encode(&platform_state)),
)
} else {
oauth_result.url
};
let mut token_exchange_extra_params = HashMap::new();
token_exchange_extra_params.insert("resource".to_string(), resource.clone());
let flow = oauth_defaults::PendingOAuthFlow {
extension_name: name.to_string(),
@@ -2471,7 +2605,7 @@ impl ExtensionManager {
secrets: Arc::clone(&self.secrets),
sse_sender: self.sse_sender.read().await.clone(),
gateway_token: self.gateway_token.clone(),
resource: Some(resource),
token_exchange_extra_params,
client_id_secret_name: if server.oauth.is_none() {
Some(server.client_id_secret_name())
} else {
@@ -2480,27 +2614,15 @@ impl ExtensionManager {
created_at: std::time::Instant::now(),
};
self.pending_oauth_flows
.write()
.await
.insert(expected_state, flow);
self.pending_auth.write().await.insert(
name.to_string(),
PendingAuth {
_name: name.to_string(),
_kind: ExtensionKind::McpServer,
created_at: std::time::Instant::now(),
task_handle: None,
},
);
Ok(AuthResult::awaiting_authorization(
name,
ExtensionKind::McpServer,
auth_url,
"gateway".to_string(),
))
Ok(self
.start_gateway_oauth_flow(HostedOAuthFlowStart {
name: name.to_string(),
kind: ExtensionKind::McpServer,
auth_url: oauth_result.url,
expected_state,
flow,
})
.await)
} else {
// Local mode: return URL for manual opening
self.pending_auth.write().await.insert(
@@ -2901,9 +3023,10 @@ impl ExtensionManager {
Enter it in the Setup tab or set {} env var",
name, env_name
);
// Only mention the Google-specific build flag for Google providers
if auth.secret_name.to_lowercase().contains("google") {
msg.push_str(", or build with IRONCLAW_GOOGLE_CLIENT_ID");
if let Some(override_env) =
crate::cli::oauth_defaults::builtin_client_id_override_env(&auth.secret_name)
{
msg.push_str(&format!(", or build with {override_env}"));
}
msg.push('.');
msg
@@ -2919,20 +3042,7 @@ impl ExtensionManager {
)
.await;
// Cancel any existing pending auth for this tool (frees port 9876 in TCP mode)
{
let mut pending = self.pending_auth.write().await;
if let Some(old) = pending.remove(name)
&& let Some(handle) = old.task_handle
{
handle.abort();
}
}
// Also clean up any gateway-mode pending flows for this tool
{
let mut flows = self.pending_oauth_flows.write().await;
flows.retain(|_, flow| flow.extension_name != name);
}
self.clear_pending_extension_auth(name).await;
let redirect_uri = self
.gateway_callback_redirect_uri()
@@ -2963,30 +3073,24 @@ impl ExtensionManager {
.unwrap_or_else(|| name.to_string());
if self.should_use_gateway_mode() {
// Gateway mode: store pending flow state for the web gateway's
// `/oauth/callback` handler to complete the exchange. No TCP listener
// needed — the OAuth provider redirects to the gateway URL.
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await;
// Wrap the CSRF nonce with instance name for platform routing.
// Nginx at auth.DOMAIN parses `instance:nonce` to route the callback
// to the correct container. The flow is keyed by the raw nonce.
let platform_state = oauth_defaults::build_platform_state(&expected_state);
let auth_url = if platform_state != expected_state {
auth_url.replace(
&format!("state={}", urlencoding::encode(&expected_state)),
&format!("state={}", urlencoding::encode(&platform_state)),
)
} else {
auth_url
};
// When an exchange proxy is configured, omit the client_secret if it
// was resolved from built-in defaults (desktop app credentials). The
// proxy holds the correct web-app secret for platform-registered OAuth
// apps. Sending the desktop secret would cause a client_id/secret
// mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web
// app, not the desktop app.
let proxy_client_secret = hosted_proxy_client_secret(
&client_secret,
builtin.as_ref(),
oauth_defaults::exchange_proxy_url().is_some(),
);
let flow = oauth_defaults::PendingOAuthFlow {
extension_name: name.to_string(),
display_name: display_name.clone(),
token_url: oauth.token_url.clone(),
client_id: client_id.clone(),
client_secret: client_secret.clone(),
client_secret: proxy_client_secret,
redirect_uri: redirect_uri.clone(),
code_verifier,
access_token_field: oauth.access_token_field.clone(),
@@ -2998,35 +3102,20 @@ impl ExtensionManager {
secrets: Arc::clone(&self.secrets),
sse_sender: self.sse_sender.read().await.clone(),
gateway_token: self.gateway_token.clone(),
resource: None,
token_exchange_extra_params: std::collections::HashMap::new(),
client_id_secret_name: None,
created_at: std::time::Instant::now(),
};
// Key by raw nonce (without instance prefix) — the callback handler
// strips the prefix before lookup.
self.pending_oauth_flows
.write()
.await
.insert(expected_state, flow);
// Register pending auth without a task handle (gateway handles completion)
self.pending_auth.write().await.insert(
name.to_string(),
PendingAuth {
_name: name.to_string(),
_kind: ExtensionKind::WasmTool,
created_at: std::time::Instant::now(),
task_handle: None,
},
);
Ok(AuthResult::awaiting_authorization(
name,
ExtensionKind::WasmTool,
auth_url,
"gateway".to_string(),
))
Ok(self
.start_gateway_oauth_flow(HostedOAuthFlowStart {
name: name.to_string(),
kind: ExtensionKind::WasmTool,
auth_url,
expected_state,
flow,
})
.await)
} else {
// TCP listener mode: bind port 9876 and spawn a background task
// to wait for the callback. This is the original flow for local/desktop use.
@@ -5241,7 +5330,8 @@ mod tests {
use crate::extensions::manager::{
ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult,
TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates,
combine_install_errors, fallback_decision, infer_kind_from_url, send_telegram_text_message,
combine_install_errors, fallback_decision, hosted_proxy_client_secret, infer_kind_from_url,
normalize_hosted_callback_url, send_telegram_text_message,
telegram_message_matches_verification_code,
};
use crate::extensions::{
@@ -6510,7 +6600,7 @@ mod tests {
secrets: Arc::clone(&secrets),
sse_sender: None,
gateway_token: None,
resource: None,
token_exchange_extra_params: std::collections::HashMap::new(),
client_id_secret_name: None,
created_at: std::time::Instant::now(),
},
@@ -6534,7 +6624,7 @@ mod tests {
secrets,
sse_sender: None,
gateway_token: None,
resource: None,
token_exchange_extra_params: std::collections::HashMap::new(),
client_id_secret_name: None,
created_at: std::time::Instant::now(),
},
@@ -6701,9 +6791,6 @@ mod tests {
// The root cause was that `should_use_gateway_mode()` only checked the
// `IRONCLAW_OAUTH_CALLBACK_URL` env var, ignoring `self.tunnel_url`.
/// Serializes env-mutating tests to prevent parallel races.
static GATEWAY_ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(());
/// Build a minimal ExtensionManager with a custom tunnel_url.
fn make_manager_with_tunnel(tunnel_url: Option<String>) -> ExtensionManager {
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
@@ -6736,9 +6823,11 @@ mod tests {
#[test]
fn should_use_gateway_mode_true_for_tunnel_url() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let _guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
@@ -6758,7 +6847,9 @@ mod tests {
#[test]
fn should_use_gateway_mode_false_without_tunnel() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let _guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
@@ -6779,7 +6870,9 @@ mod tests {
#[test]
fn should_use_gateway_mode_false_for_loopback_tunnel() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let _guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
@@ -6807,9 +6900,11 @@ mod tests {
impl EnvGuard {
fn new() -> Self {
let guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
@@ -6822,7 +6917,7 @@ mod tests {
impl Drop for EnvGuard {
fn drop(&mut self) {
// SAFETY: Under GATEWAY_ENV_MUTEX (still held by _mutex), no concurrent env access.
// SAFETY: Under ENV_MUTEX (still held by _mutex), no concurrent env access.
unsafe {
if let Some(ref val) = self.original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
@@ -6863,6 +6958,90 @@ mod tests {
);
}
#[test]
fn gateway_callback_redirect_uri_does_not_duplicate_callback_path_from_env() {
let _guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::set_var(
"IRONCLAW_OAUTH_CALLBACK_URL",
"https://oauth.test.example/oauth/callback",
);
}
let mgr = make_manager_with_tunnel(None);
assert_eq!(
tokio_test::block_on(mgr.gateway_callback_redirect_uri()),
Some("https://oauth.test.example/oauth/callback".to_string()),
);
unsafe {
if let Some(val) = original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
} else {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
}
}
#[test]
fn gateway_callback_redirect_uri_trims_trailing_slash_from_env_callback() {
let _guard = crate::config::helpers::ENV_MUTEX
.lock()
.expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::set_var(
"IRONCLAW_OAUTH_CALLBACK_URL",
"https://oauth.test.example/oauth/callback/",
);
}
let mgr = make_manager_with_tunnel(None);
assert_eq!(
tokio_test::block_on(mgr.gateway_callback_redirect_uri()),
Some("https://oauth.test.example/oauth/callback".to_string()),
);
unsafe {
if let Some(val) = original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
} else {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
}
}
#[test]
fn normalize_hosted_callback_url_preserves_query_params() {
assert_eq!(
normalize_hosted_callback_url("https://oauth.test.example?source=hosted"),
"https://oauth.test.example/oauth/callback?source=hosted"
);
assert_eq!(
normalize_hosted_callback_url(
"https://oauth.test.example/oauth/callback?source=hosted"
),
"https://oauth.test.example/oauth/callback?source=hosted"
);
}
#[test]
fn rewrite_oauth_state_param_updates_only_state_query_param() {
let auth_url =
"https://auth.example.com/authorize?client_id=abc&state=old-state&hint=state%3Dkeep";
assert_eq!(
ExtensionManager::rewrite_oauth_state_param(
auth_url.to_string(),
"old-state",
"new-hosted-state",
),
"https://auth.example.com/authorize?client_id=abc&state=new-hosted-state&hint=state%3Dkeep"
);
}
#[tokio::test]
async fn gateway_mode_enabled_explicitly() {
let _env = EnvGuard::new();
@@ -7217,4 +7396,71 @@ mod tests {
panic!("URL missing token: {url}"); // safety: test assertion
}
}
// ── proxy_client_secret suppression ─────────────────────────────
#[test]
fn test_proxy_client_secret_suppressed_when_builtin_matches_with_exchange_proxy() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = hosted_proxy_client_secret(&secret, builtin_ref, true);
assert_eq!(
result, None,
"built-in desktop secret must be suppressed when the exchange proxy is configured"
);
}
#[test]
fn test_proxy_client_secret_kept_when_not_builtin_with_exchange_proxy() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let secret = Some("user-entered-custom-secret".to_string());
let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!(
result,
Some("user-entered-custom-secret".to_string()),
"non-builtin secret must be kept even when the exchange proxy is configured"
);
}
#[test]
fn test_proxy_client_secret_kept_without_exchange_proxy_even_for_builtin_secret() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = hosted_proxy_client_secret(&secret, builtin_ref, false);
assert_eq!(
result, secret,
"built-in secret must be kept when the callback will exchange directly"
);
}
#[test]
fn test_proxy_client_secret_none_stays_none() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let result = hosted_proxy_client_secret(&None, builtin.as_ref(), true);
assert_eq!(
result, None,
"None secret stays None even when the exchange proxy is configured"
);
}
#[test]
fn test_proxy_client_secret_no_builtin_provider() {
// MCP/non-Google providers have no builtin credentials
let builtin = crate::cli::oauth_defaults::builtin_credentials("mcp_notion_access_token");
assert!(builtin.is_none());
let secret = Some("dcr-secret".to_string());
let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!(
result,
Some("dcr-secret".to_string()),
"non-builtin provider secret must be kept"
);
}
}