//! Shared OAuth infrastructure: built-in credentials, callback server, landing pages. //! //! Every OAuth flow in the codebase (WASM tool auth, MCP server auth, NEAR AI login) //! uses the same callback port, landing page, and listener logic from this module. //! //! # Built-in Credentials //! //! Many CLI tools (gcloud, rclone, gdrive) ship with default OAuth credentials //! so users don't need to register their own OAuth app. Google explicitly //! documents that client_secret for "Desktop App" / "Installed App" types //! is NOT actually secret. //! //! Default credentials are hardcoded below. They can be overridden at: //! //! - **Compile time**: Set IRONCLAW_GOOGLE_CLIENT_ID / IRONCLAW_GOOGLE_CLIENT_SECRET //! env vars before building to replace the hardcoded defaults. //! - **Runtime**: Users can set GOOGLE_OAUTH_CLIENT_ID / GOOGLE_OAUTH_CLIENT_SECRET //! env vars, which take priority over built-in defaults. use std::time::Duration; use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; use tokio::net::TcpListener; // ── Built-in credentials ──────────────────────────────────────────────── pub struct OAuthCredentials { pub client_id: &'static str, pub client_secret: &'static str, } /// Google OAuth "Desktop App" credentials, shared across all Google tools. /// Compile-time env vars override the hardcoded defaults below. const GOOGLE_CLIENT_ID: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_ID") { Some(v) => v, None => "564604149681-efo25d43rs85v0tibdepsmdv5dsrhhr0.apps.googleusercontent.com", }; const GOOGLE_CLIENT_SECRET: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_SECRET") { Some(v) => v, None => "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", }; /// Returns built-in OAuth credentials for a provider, keyed by secret_name. /// /// The secret_name comes from the tool's capabilities.json `auth.secret_name` field. /// Returns `None` if no built-in credentials are configured for that provider. pub fn builtin_credentials(secret_name: &str) -> Option { match secret_name { "google_oauth_token" => Some(OAuthCredentials { client_id: GOOGLE_CLIENT_ID, client_secret: GOOGLE_CLIENT_SECRET, }), _ => None, } } // ── Shared callback server ────────────────────────────────────────────── /// Fixed port for all OAuth callbacks. /// /// Every redirect URI registered with providers must use this port: /// `http://localhost:9876/callback` (or `/auth/callback` for NEAR AI). pub const OAUTH_CALLBACK_PORT: u16 = 9876; /// Returns the OAuth callback base URL. /// /// Checks `IRONCLAW_OAUTH_CALLBACK_URL` env var first (useful for remote/VPS /// deployments where `127.0.0.1` is unreachable from the user's browser), /// then falls back to `http://127.0.0.1:{OAUTH_CALLBACK_PORT}`. pub fn callback_url() -> String { std::env::var("IRONCLAW_OAUTH_CALLBACK_URL") .ok() .filter(|v| !v.is_empty()) .unwrap_or_else(|| format!("http://127.0.0.1:{}", OAUTH_CALLBACK_PORT)) } /// Error from the OAuth callback listener. #[derive(Debug, thiserror::Error)] pub enum OAuthCallbackError { #[error("Port {0} is in use (another auth flow running?): {1}")] PortInUse(u16, String), #[error("Authorization denied by user")] Denied, #[error("Timed out waiting for authorization")] Timeout, #[error("IO error: {0}")] Io(String), } /// Bind the OAuth callback listener on the fixed port. /// /// Binds to IPv4 `127.0.0.1` first because callback URLs use `127.0.0.1` /// explicitly (e.g., NEAR AI redirects to `http://127.0.0.1:9876/auth/callback`). /// Falls back to IPv6 `[::1]` only if IPv4 binding fails for a reason other /// than `AddrInUse`. If the port is already occupied, fails immediately. pub async fn bind_callback_listener() -> Result { let ipv4_addr = format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT); match TcpListener::bind(&ipv4_addr).await { Ok(listener) => return Ok(listener), Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => { return Err(OAuthCallbackError::PortInUse( OAUTH_CALLBACK_PORT, e.to_string(), )); } Err(_) => { // IPv4 not available, fall back to IPv6 } } TcpListener::bind(format!("[::1]:{}", OAUTH_CALLBACK_PORT)) .await .map_err(|e| { if e.kind() == std::io::ErrorKind::AddrInUse { OAuthCallbackError::PortInUse(OAUTH_CALLBACK_PORT, e.to_string()) } else { OAuthCallbackError::Io(e.to_string()) } }) } /// Wait for an OAuth callback and extract a query parameter value. /// /// Listens for a GET request matching `path_prefix` (e.g., "/callback" or "/auth/callback"), /// extracts the value of `param_name` (e.g., "code" or "token"), and shows a branded /// landing page using `display_name` (e.g., "Google", "Notion", "NEAR AI"). /// /// Times out after 5 minutes. pub async fn wait_for_callback( listener: TcpListener, path_prefix: &str, param_name: &str, display_name: &str, ) -> Result { let path_prefix = path_prefix.to_string(); let param_name = param_name.to_string(); let display_name = display_name.to_string(); tokio::time::timeout(Duration::from_secs(300), async move { loop { let (mut socket, _) = listener .accept() .await .map_err(|e| OAuthCallbackError::Io(e.to_string()))?; let mut reader = BufReader::new(&mut socket); let mut request_line = String::new(); reader .read_line(&mut request_line) .await .map_err(|e| OAuthCallbackError::Io(e.to_string()))?; if let Some(path) = request_line.split_whitespace().nth(1) && path.starts_with(&path_prefix) && let Some(query) = path.split('?').nth(1) { // Check for error first if query.contains("error=") { let html = landing_html(&display_name, false); let response = format!( "HTTP/1.1 400 Bad Request\r\n\ Content-Type: text/html; charset=utf-8\r\n\ Connection: close\r\n\ \r\n\ {}", html ); let _ = socket.write_all(response.as_bytes()).await; return Err(OAuthCallbackError::Denied); } // Look for the target parameter for param in query.split('&') { let parts: Vec<&str> = param.splitn(2, '=').collect(); if parts.len() == 2 && parts[0] == param_name { let value = urlencoding::decode(parts[1]) .unwrap_or_else(|_| parts[1].into()) .into_owned(); let html = landing_html(&display_name, true); let response = format!( "HTTP/1.1 200 OK\r\n\ Content-Type: text/html; charset=utf-8\r\n\ Connection: close\r\n\ \r\n\ {}", html ); let _ = socket.write_all(response.as_bytes()).await; let _ = socket.shutdown().await; return Ok(value); } } } // Not the callback we're looking for let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n"; let _ = socket.write_all(response.as_bytes()).await; } }) .await .map_err(|_| OAuthCallbackError::Timeout)? } /// Escape a string for safe interpolation into HTML content. fn html_escape(s: &str) -> String { let mut out = String::with_capacity(s.len()); for c in s.chars() { match c { '&' => out.push_str("&"), '<' => out.push_str("<"), '>' => out.push_str(">"), '"' => out.push_str("""), '\'' => out.push_str("'"), _ => out.push(c), } } out } /// HTML landing page shown in the browser after an OAuth redirect. pub fn landing_html(provider_name: &str, success: bool) -> String { let safe_name = html_escape(provider_name); let (icon, heading, subtitle, accent) = if success { ( r##"
"##, format!("{} Connected", safe_name), "You can close this window and return to your terminal.", "#22c55e", ) } else { ( r##"
"##, "Authorization Failed".to_string(), "The request was denied. You can close this window and try again.", "#ef4444", ) }; format!( r#" IronClaw - {heading}
{icon}

{heading}

{subtitle}

IronClaw
"#, heading = heading, icon = icon, subtitle = subtitle, accent = accent, ) } #[cfg(test)] mod tests { use std::sync::Mutex; use crate::cli::oauth_defaults::{builtin_credentials, callback_url, landing_html}; /// Serializes env-mutating tests to prevent parallel races. static ENV_MUTEX: Mutex<()> = Mutex::new(()); #[test] fn test_callback_url_default() { let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); // Clear the env var to test default behavior let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL"); } let url = callback_url(); assert_eq!(url, "http://127.0.0.1:9876"); // Restore unsafe { if let Some(val) = original { std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val); } } } #[test] fn test_callback_url_env_override() { let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { std::env::set_var( "IRONCLAW_OAUTH_CALLBACK_URL", "https://myserver.example.com:9876", ); } let url = callback_url(); assert_eq!(url, "https://myserver.example.com:9876"); // Restore 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 test_unknown_provider_returns_none() { assert!(builtin_credentials("unknown_token").is_none()); } #[test] fn test_google_returns_based_on_compile_env() { let creds = builtin_credentials("google_oauth_token"); assert!(creds.is_some()); let creds = creds.unwrap(); assert!(!creds.client_id.is_empty()); assert!(!creds.client_secret.is_empty()); } #[test] fn test_landing_html_success_contains_key_elements() { let html = landing_html("Google", true); assert!(html.contains("Google Connected")); assert!(html.contains("charset")); assert!(html.contains("IronClaw")); assert!(html.contains("#22c55e")); // green accent assert!(!html.contains("Failed")); } #[test] fn test_landing_html_escapes_provider_name() { let html = landing_html("", true); assert!(!html.contains("