use super::access_token::{clear_rejected, is_rejected, is_valid_access_token, set_access_token}; use super::openai_compatible_oauth::OpenAICompatibleOAuthProvider; use super::{ClientConfig, ProviderModels}; use crate::config::paths; use anyhow::{Context, Error, Result, anyhow, bail}; use base64::Engine; use base64::engine::general_purpose::URL_SAFE_NO_PAD; use chrono::Utc; use indexmap::IndexMap; use inquire::Text; use reqwest::{Client as ReqwestClient, RequestBuilder, StatusCode}; use serde::{Deserialize, Serialize}; use serde_json::Value; use sha2::{Digest, Sha256}; use std::collections::HashMap; use std::fs; use std::io::{BufRead, BufReader, Write}; use std::net::TcpListener; use std::path::PathBuf; use std::sync::{Arc, OnceLock}; use std::time::Duration; use tokio::sync; use url::Url; use uuid::Uuid; #[derive(Debug, Clone, Copy, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum TokenRequestFormat { Json, FormUrlEncoded, } #[derive(Debug, Clone, Copy, Serialize, Deserialize, Default)] #[serde(rename_all = "snake_case")] pub enum OAuthFlow { #[default] Pkce, ClientCredentials, DeviceCode, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct OAuthConfig { pub client_id: String, pub token_url: String, #[serde(default)] pub flow: OAuthFlow, pub client_secret: Option, pub authorize_url: Option, pub redirect_uri: Option, pub redirect_port: Option, pub device_authorization_url: Option, #[serde(default)] pub scopes: Vec, pub token_request_format: Option, #[serde(default)] pub extra_authorize_params: IndexMap, #[serde(default)] pub extra_token_headers: IndexMap, #[serde(default)] pub extra_request_headers: IndexMap, #[serde(default)] pub echo_pkce_in_token_exchange: bool, #[serde(default)] pub use_pkce_in_device_flow: bool, #[serde(default = "default_true")] pub include_state_in_token_exchange: bool, } fn default_true() -> bool { true } impl OAuthConfig { pub fn merge(mut self, override_cfg: OAuthConfig) -> Self { self.client_id = override_cfg.client_id; self.token_url = override_cfg.token_url; self.flow = override_cfg.flow; if override_cfg.client_secret.is_some() { self.client_secret = override_cfg.client_secret; } if override_cfg.authorize_url.is_some() { self.authorize_url = override_cfg.authorize_url; } if override_cfg.redirect_uri.is_some() { self.redirect_uri = override_cfg.redirect_uri; } if override_cfg.redirect_port.is_some() { self.redirect_port = override_cfg.redirect_port; } if override_cfg.device_authorization_url.is_some() { self.device_authorization_url = override_cfg.device_authorization_url; } if !override_cfg.scopes.is_empty() { self.scopes = override_cfg.scopes; } if override_cfg.token_request_format.is_some() { self.token_request_format = override_cfg.token_request_format; } if !override_cfg.extra_authorize_params.is_empty() { self.extra_authorize_params .extend(override_cfg.extra_authorize_params); } if !override_cfg.extra_token_headers.is_empty() { self.extra_token_headers .extend(override_cfg.extra_token_headers); } if !override_cfg.extra_request_headers.is_empty() { self.extra_request_headers .extend(override_cfg.extra_request_headers); } self.echo_pkce_in_token_exchange = override_cfg.echo_pkce_in_token_exchange; self.use_pkce_in_device_flow = override_cfg.use_pkce_in_device_flow; self.include_state_in_token_exchange = override_cfg.include_state_in_token_exchange; self } } pub trait OAuthProvider: Send + Sync { fn provider_name(&self) -> &str; fn client_id(&self) -> &str; fn authorize_url(&self) -> &str; fn token_url(&self) -> &str; fn redirect_uri(&self) -> &str; fn scopes(&self) -> String; fn client_secret(&self) -> Option<&str> { None } fn extra_authorize_params(&self) -> Vec<(&str, &str)> { vec![] } /// Extra form/body parameters appended to every token request routed /// through `build_token_request` (authorization-code exchange, refresh, /// client_credentials, and device-code polling). Used e.g. for the /// RFC 8707 `resource` indicator required by the MCP spec. /// NOTE: these are merged AFTER the caller's params and will overwrite /// a colliding key; do not return protocol parameter names /// (grant_type, client_id, code, refresh_token, ...). fn extra_token_params(&self) -> Vec<(&str, &str)> { vec![] } fn token_request_format(&self) -> TokenRequestFormat { TokenRequestFormat::Json } fn uses_localhost_redirect(&self) -> bool { false } fn extra_token_headers(&self) -> Vec<(&str, &str)> { vec![] } fn extra_request_headers(&self) -> Vec<(&str, &str)> { vec![] } fn fixed_redirect_uri(&self) -> Option { None } fn extract_account_id(&self, _response: &Value) -> Option { None } fn include_state_in_token_exchange(&self) -> bool { true } fn flow(&self) -> OAuthFlow { OAuthFlow::Pkce } fn echo_pkce_in_token_exchange(&self) -> bool { false } fn device_authorization_url(&self) -> Option<&str> { None } fn use_pkce_in_device_flow(&self) -> bool { false } } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct OAuthTokens { pub access_token: String, #[serde(default)] pub refresh_token: Option, pub expires_at: i64, #[serde(default)] pub account_id: Option, } const TOKEN_ENDPOINT_TIMEOUT: Duration = Duration::from_secs(30); pub async fn run_oauth_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> { match provider.flow() { OAuthFlow::Pkce => run_pkce_flow(provider, client_name).await, OAuthFlow::ClientCredentials => { run_client_credentials_flow(provider, client_name).await?; println!( "Successfully authenticated client '{}' with {} via OAuth (client_credentials). Tokens saved.", client_name, provider.provider_name() ); Ok(()) } OAuthFlow::DeviceCode => run_device_code_flow(provider, client_name).await, } } async fn run_pkce_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> { let random_bytes: [u8; 32] = rand::random::<[u8; 32]>(); let code_verifier = URL_SAFE_NO_PAD.encode(random_bytes); let mut hasher = Sha256::new(); hasher.update(code_verifier.as_bytes()); let code_challenge = URL_SAFE_NO_PAD.encode(hasher.finalize()); let state = Uuid::new_v4().to_string(); let (redirect_uri, use_callback_listener) = if let Some(fixed) = provider.fixed_redirect_uri() { (fixed, true) } else if provider.uses_localhost_redirect() { let listener = TcpListener::bind("127.0.0.1:0")?; let port = listener.local_addr()?.port(); let uri = format!("http://127.0.0.1:{port}/callback"); drop(listener); (uri, true) } else { (provider.redirect_uri().to_string(), false) }; let scopes = provider.scopes(); let encoded_scopes = urlencoding::encode(&scopes); let encoded_redirect = urlencoding::encode(&redirect_uri); let mut authorize_url = format!( "{}?client_id={}&response_type=code&scope={}&redirect_uri={}&code_challenge={}&code_challenge_method=S256&state={}", provider.authorize_url(), provider.client_id(), encoded_scopes, encoded_redirect, code_challenge, state ); for (key, value) in provider.extra_authorize_params() { authorize_url.push_str(&format!( "&{}={}", urlencoding::encode(key), urlencoding::encode(value) )); } println!( "\nOpen this URL to authenticate with {} (client '{}'):\n", provider.provider_name(), client_name ); println!(" {authorize_url}\n"); let _ = open::that(&authorize_url); let (code, returned_state) = if use_callback_listener { let (code, state) = listen_for_oauth_callback(&redirect_uri)?; (code, Some(state)) } else { let input = Text::new("Paste the authorization code or callback URL:").prompt()?; parse_paste_input(input.trim())? }; if let Some(returned) = returned_state.as_deref() { if returned != state { bail!( "OAuth state mismatch: expected '{state}', got '{returned}'. \ This may indicate a CSRF attack or a stale authorization attempt." ); } } else { eprintln!( "Warning: no state returned in the paste; skipping CSRF check. \ If your provider's callback page shows a URL or code#state string, paste that instead." ); } let client = ReqwestClient::new(); let mut token_params = vec![ ("grant_type", "authorization_code"), ("client_id", provider.client_id()), ("code", code.as_str()), ("code_verifier", code_verifier.as_str()), ("redirect_uri", redirect_uri.as_str()), ]; if provider.include_state_in_token_exchange() { token_params.push(("state", state.as_str())); } if provider.echo_pkce_in_token_exchange() { token_params.push(("code_challenge", code_challenge.as_str())); token_params.push(("code_challenge_method", "S256")); } let request = build_token_request(&client, provider, &token_params); let response: Value = request.send().await?.json().await?; let access_token = response["access_token"] .as_str() .ok_or_else(|| { anyhow!( "Missing access_token in response (keys: {})", token_response_keys(&response) ) })? .to_string(); let refresh_token = response["refresh_token"].as_str().map(|s| s.to_string()); let expires_in = response["expires_in"].as_i64().ok_or_else(|| { anyhow!( "Missing expires_in in response (keys: {})", token_response_keys(&response) ) })?; let expires_at = Utc::now().timestamp() + expires_in; let account_id = provider.extract_account_id(&response); let tokens = OAuthTokens { access_token, refresh_token, expires_at, account_id, }; save_oauth_tokens(client_name, &tokens)?; println!( "Successfully authenticated client '{}' with {} via OAuth. Tokens saved.", client_name, provider.provider_name() ); Ok(()) } async fn run_client_credentials_flow( provider: &dyn OAuthProvider, client_name: &str, ) -> Result<()> { let client = ReqwestClient::builder() .timeout(TOKEN_ENDPOINT_TIMEOUT) .build()?; let scopes = provider.scopes(); let mut params: Vec<(&str, &str)> = vec![ ("grant_type", "client_credentials"), ("client_id", provider.client_id()), ]; if !scopes.is_empty() { params.push(("scope", scopes.as_str())); } let request = build_token_request(&client, provider, ¶ms); let response: Value = request.send().await?.json().await?; let access_token = response["access_token"] .as_str() .ok_or_else(|| { anyhow!( "Missing access_token in client_credentials response (keys: {})", token_response_keys(&response) ) })? .to_string(); let expires_in = response["expires_in"].as_i64().ok_or_else(|| { anyhow!( "Missing expires_in in client_credentials response (keys: {})", token_response_keys(&response) ) })?; let expires_at = Utc::now().timestamp() + expires_in; let tokens = OAuthTokens { access_token, refresh_token: None, expires_at, account_id: provider.extract_account_id(&response), }; save_oauth_tokens(client_name, &tokens)?; Ok(()) } async fn run_device_code_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> { let device_auth_url = provider.device_authorization_url().ok_or_else(|| { anyhow!( "Provider '{}' is configured with flow: device_code but has no device_authorization_url. \ Set `oauth.device_authorization_url` in your config.", provider.provider_name() ) })?; let client = ReqwestClient::new(); let pkce = if provider.use_pkce_in_device_flow() { let random_bytes: [u8; 32] = rand::random::<[u8; 32]>(); let verifier = URL_SAFE_NO_PAD.encode(random_bytes); let mut hasher = Sha256::new(); hasher.update(verifier.as_bytes()); let challenge = URL_SAFE_NO_PAD.encode(hasher.finalize()); Some((verifier, challenge)) } else { None }; let scopes = provider.scopes(); let mut device_params: Vec<(&str, &str)> = vec![("client_id", provider.client_id())]; if !scopes.is_empty() { device_params.push(("scope", scopes.as_str())); } if let Some((_, ref challenge)) = pkce { device_params.push(("code_challenge", challenge.as_str())); device_params.push(("code_challenge_method", "S256")); } let form: HashMap<&str, &str> = device_params.iter().copied().collect(); let mut device_request = client .post(device_auth_url) .header("Accept", "application/json") .form(&form); for (key, value) in provider.extra_token_headers() { device_request = device_request.header(key, value); } let device_response: Value = device_request.send().await?.json().await?; let device_code = device_response["device_code"] .as_str() .ok_or_else(|| { anyhow!( "Missing device_code in device authorization response (keys: {})", token_response_keys(&device_response) ) })? .to_string(); let user_code = device_response["user_code"] .as_str() .ok_or_else(|| { anyhow!( "Missing user_code in device authorization response (keys: {})", token_response_keys(&device_response) ) })? .to_string(); let verification_uri = device_response["verification_uri"] .as_str() .ok_or_else(|| { anyhow!( "Missing verification_uri in device authorization response (keys: {})", token_response_keys(&device_response) ) })? .to_string(); let verification_uri_complete = device_response["verification_uri_complete"] .as_str() .map(|s| s.to_string()); let expires_in = device_response["expires_in"].as_i64().unwrap_or(1800); let mut interval = device_response["interval"].as_u64().unwrap_or(5); let deadline = Utc::now().timestamp() + expires_in; println!( "\nAuthenticate with {} (client '{}'):", provider.provider_name(), client_name ); println!(" 1. Open: {verification_uri}"); println!(" 2. Enter code: {user_code}\n"); if let Some(complete) = &verification_uri_complete { println!(" (Or open the pre-filled URL: {complete})\n"); } let url_to_open = verification_uri_complete .as_deref() .unwrap_or(&verification_uri); if std::env::var(crate::sandbox::SANDBOX_ENV_FLAG).is_ok() && let Ok(qr) = qrcode::QrCode::new(url_to_open) { let rendered = qr .render::() .quiet_zone(true) .build(); println!("{rendered}\n"); } let _ = open::that(url_to_open); println!("Waiting for authorization (polling every {interval}s)...\n"); loop { tokio::time::sleep(std::time::Duration::from_secs(interval)).await; if Utc::now().timestamp() >= deadline { bail!( "Device code expired before user approval. Run `coyote --authenticate {}` to try again.", client_name ); } let mut token_params: Vec<(&str, &str)> = vec![ ("grant_type", "urn:ietf:params:oauth:grant-type:device_code"), ("device_code", device_code.as_str()), ("client_id", provider.client_id()), ]; if let Some((verifier, _)) = pkce.as_ref() { token_params.push(("code_verifier", verifier.as_str())); } let token_response: Value = build_token_request(&client, provider, &token_params) .header("Accept", "application/json") .send() .await? .json() .await?; if let Some(access_token) = token_response["access_token"].as_str() { let refresh_token = token_response["refresh_token"].as_str().map(str::to_string); let expires_at = match token_response["expires_in"].as_i64() { Some(secs) => Utc::now().timestamp() + secs, None => i64::MAX, }; let account_id = provider.extract_account_id(&token_response); let tokens = OAuthTokens { access_token: access_token.to_string(), refresh_token, expires_at, account_id, }; save_oauth_tokens(client_name, &tokens)?; println!( "Successfully authenticated client '{}' with {} via OAuth (device_code). Tokens saved.", client_name, provider.provider_name() ); return Ok(()); } let error_code = token_response["error"].as_str().unwrap_or(""); match error_code { "authorization_pending" => continue, "slow_down" => { interval += 5; println!("Server requested slower polling; increasing interval to {interval}s."); continue; } "expired_token" => bail!( "Device code expired. Run `coyote --authenticate {}` to try again.", client_name ), "access_denied" => bail!("Authorization was denied by the user."), other => bail!( "Device code polling failed: {} — {}", other, token_response["error_description"] .as_str() .unwrap_or("no description") ), } } } pub fn load_oauth_tokens(client_name: &str) -> Option { let path = paths::token_file(client_name); let content = fs::read_to_string(path).ok()?; serde_json::from_str(&content).ok() } fn save_oauth_tokens(client_name: &str, tokens: &OAuthTokens) -> Result<()> { let path = paths::token_file(client_name); if let Some(parent) = path.parent() { fs::create_dir_all(parent)?; } let json = serde_json::to_string_pretty(tokens)?; // Write-then-rename so a crash mid-write never truncates the live token file. let mut tmp = path.clone().into_os_string(); tmp.push(".tmp"); let tmp = PathBuf::from(tmp); // Tokens are live credentials: create the file owner-only, not umask-default. let mut options = fs::OpenOptions::new(); options.write(true).create(true).truncate(true); #[cfg(unix)] { use std::os::unix::fs::OpenOptionsExt; options.mode(0o600); } options.open(&tmp)?.write_all(json.as_bytes())?; fs::rename(&tmp, &path)?; Ok(()) } pub(crate) fn token_response_keys(response: &Value) -> String { match response.as_object() { Some(map) => { let keys: Vec<&str> = map.keys().map(String::as_str).collect(); format!("[{}]", keys.join(", ")) } None => "".to_string(), } } fn parse_refresh_response( status: StatusCode, response: &Value, previous_refresh_token: Option<&str>, ) -> Result<(String, Option, i64)> { if let Some(error) = response["error"].as_str() { let description = response["error_description"] .as_str() .unwrap_or("no description"); if matches!(error, "invalid_grant" | "invalid_token") { bail!( "OAuth refresh token was rejected ({error}: {description}). Please re-authenticate." ); } bail!("Token refresh failed ({error}: {description})"); } if !status.is_success() { bail!("Token refresh failed with HTTP status {status}"); } let access_token = response["access_token"] .as_str() .ok_or_else(|| { anyhow!( "Missing access_token in refresh response (keys: {})", token_response_keys(response) ) })? .to_string(); let refresh_token = response["refresh_token"] .as_str() .or(previous_refresh_token) .map(str::to_string); let expires_in = response["expires_in"].as_i64().ok_or_else(|| { anyhow!( "Missing expires_in in refresh response (keys: {})", token_response_keys(response) ) })?; Ok((access_token, refresh_token, expires_in)) } pub async fn refresh_oauth_token( client: &ReqwestClient, provider: &dyn OAuthProvider, client_name: &str, tokens: &OAuthTokens, ) -> Result { let refresh_token_val = tokens.refresh_token.as_deref().ok_or_else(|| { anyhow!( "No refresh token available for '{}'. Please re-authenticate.", client_name ) })?; let request = build_token_request( client, provider, &[ ("grant_type", "refresh_token"), ("client_id", provider.client_id()), ("refresh_token", refresh_token_val), ], ); let (status, response) = tokio::time::timeout(TOKEN_ENDPOINT_TIMEOUT, async { let response = request.send().await?; let status = response.status(); let body: Value = response.json().await?; Ok::<_, Error>((status, body)) }) .await .map_err(|_| { anyhow!( "Token refresh for '{}' timed out after {}s", client_name, TOKEN_ENDPOINT_TIMEOUT.as_secs() ) })??; let (access_token, refresh_token, expires_in) = parse_refresh_response(status, &response, tokens.refresh_token.as_deref())?; let expires_at = Utc::now().timestamp() + expires_in; let account_id = provider .extract_account_id(&response) .or_else(|| tokens.account_id.clone()); let new_tokens = OAuthTokens { access_token, refresh_token, expires_at, account_id, }; save_oauth_tokens(client_name, &new_tokens)?; Ok(new_tokens) } /// Per-client lock so concurrent requests perform a single refresh. /// Returns a clone of the Arc so the parking_lot guard is dropped before the /// caller awaits on the tokio mutex. fn refresh_guard(client_name: &str) -> Arc> { static GUARDS: OnceLock>>>> = OnceLock::new(); GUARDS .get_or_init(Default::default) .lock() .entry(client_name.to_string()) .or_default() .clone() } pub async fn prepare_oauth_access_token( client: &ReqwestClient, provider: &dyn OAuthProvider, client_name: &str, ) -> Result { if is_valid_access_token(client_name) { return Ok(true); } let tokens = match load_oauth_tokens(client_name) { Some(t) => t, None => return Ok(false), }; let tokens = if Utc::now().timestamp() >= tokens.expires_at || is_rejected(client_name, &tokens.access_token) { let guard = refresh_guard(client_name); let _guard = guard.lock().await; // A concurrent caller may have refreshed while we waited for the // lock; a valid in-memory token means the winner already populated // the cache. if is_valid_access_token(client_name) { return Ok(true); } let tokens = match load_oauth_tokens(client_name) { Some(t) => t, None => return Ok(false), }; if Utc::now().timestamp() >= tokens.expires_at || is_rejected(client_name, &tokens.access_token) { match provider.flow() { OAuthFlow::Pkce | OAuthFlow::DeviceCode => { refresh_oauth_token(client, provider, client_name, &tokens).await? } OAuthFlow::ClientCredentials => { run_client_credentials_flow(provider, client_name).await?; load_oauth_tokens(client_name).ok_or_else(|| { anyhow!("Token file missing after client_credentials refresh") })? } } } else { tokens } } else { tokens }; set_access_token( client_name, tokens.access_token, tokens.expires_at, tokens.account_id, ); // Clear even when the refresh returned the same token (some IdPs reuse // JWTs within validity); otherwise every request re-hits the token endpoint. clear_rejected(client_name); Ok(true) } fn build_token_request( client: &ReqwestClient, provider: &(impl OAuthProvider + ?Sized), params: &[(&str, &str)], ) -> RequestBuilder { let all_params: Vec<(&str, &str)> = params .iter() .copied() .chain(provider.extra_token_params()) .collect(); let mut request = match provider.token_request_format() { TokenRequestFormat::Json => { let body: serde_json::Map = all_params .iter() .map(|(k, v)| (k.to_string(), Value::String(v.to_string()))) .collect(); if let Some(secret) = provider.client_secret() { let mut body = body; body.insert( "client_secret".to_string(), Value::String(secret.to_string()), ); client.post(provider.token_url()).json(&body) } else { client.post(provider.token_url()).json(&body) } } TokenRequestFormat::FormUrlEncoded => { let mut form: HashMap = all_params .iter() .map(|(k, v)| (k.to_string(), v.to_string())) .collect(); if let Some(secret) = provider.client_secret() { form.insert("client_secret".to_string(), secret.to_string()); } client.post(provider.token_url()).form(&form) } }; for (key, value) in provider.extra_token_headers() { request = request.header(key, value); } request } fn parse_paste_input(input: &str) -> Result<(String, Option)> { if input.is_empty() { bail!("Empty input; paste the code, code#state, or callback URL from your browser."); } if input.starts_with("http://") || input.starts_with("https://") { let parsed = Url::parse(input).with_context(|| format!("Failed to parse pasted URL: {input}"))?; let code = parsed .query_pairs() .find(|(k, _)| k == "code") .map(|(_, v)| v.to_string()) .ok_or_else(|| { anyhow!("Pasted URL is missing the ?code= parameter. Paste the URL you were redirected to after approving.") })?; let state = parsed .query_pairs() .find(|(k, _)| k == "state") .map(|(_, v)| v.to_string()); return Ok((code, state)); } if let Some((code, state)) = input.split_once('#') { return Ok((code.to_string(), Some(state.to_string()))); } Ok((input.to_string(), None)) } fn listen_for_oauth_callback(redirect_uri: &str) -> Result<(String, String)> { let url: Url = redirect_uri.parse()?; let host = url.host_str().unwrap_or("127.0.0.1"); let port = url .port() .ok_or_else(|| anyhow!("No port in redirect URI"))?; let path = url.path(); println!("Waiting for OAuth callback on {redirect_uri} ..."); println!( "(If the browser shows a 'paste this code' page, ignore it. Coyote captures the callback automatically.)\n" ); let listener = TcpListener::bind(format!("{host}:{port}"))?; loop { let (mut stream, _) = listener.accept()?; let mut reader = BufReader::new(&stream); let mut request_line = String::new(); if reader.read_line(&mut request_line).is_err() || request_line.trim().is_empty() { continue; } let Some(request_path) = request_line.split_whitespace().nth(1) else { continue; }; let Ok(parsed) = format!("http://{host}:{port}{request_path}").parse::() else { continue; }; if !parsed.path().starts_with(path) { let _ = stream.write_all( b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n", ); continue; } let code = parsed .query_pairs() .find(|(k, _)| k == "code") .map(|(_, v)| v.to_string()) .ok_or_else(|| { let error = parsed .query_pairs() .find(|(k, _)| k == "error") .map(|(_, v)| v.to_string()) .unwrap_or_else(|| "unknown".to_string()); anyhow!("OAuth callback returned error: {error}") })?; let returned_state = parsed .query_pairs() .find(|(k, _)| k == "state") .map(|(_, v)| v.to_string()) .ok_or_else(|| anyhow!("Missing state parameter in OAuth callback"))?; let response_body = "

Authentication successful!

You can close this tab and return to your terminal.

"; let response = format!( "HTTP/1.1 200 OK\r\nContent-Type: text/html\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", response_body.len(), response_body ); stream.write_all(response.as_bytes())?; return Ok((code, returned_state)); } } pub fn get_oauth_provider(provider_type: &str) -> Option> { match provider_type { "claude" => Some(Box::new(super::claude_oauth::ClaudeOAuthProvider)), "gemini" => Some(Box::new(super::gemini_oauth::GeminiOAuthProvider)), "openai" => Some(Box::new(super::openai_oauth::OpenAIOAuthProvider)), _ => None, } } pub fn get_oauth_provider_for_client( client_config: &ClientConfig, all_provider_models: &[ProviderModels], ) -> Option> { let (client_name, provider_type, auth) = client_config_info(client_config); if auth != Some("oauth") { return None; } match client_config { ClientConfig::OpenAICompatibleConfig(c) => { let base = all_provider_models .iter() .find(|p| p.provider == client_name) .and_then(|p| p.oauth.clone()); let user_oauth = c.oauth.clone().map(|b| *b); let merged = match (base, user_oauth) { (None, None) => return None, (Some(b), None) => b, (None, Some(u)) => u, (Some(b), Some(u)) => b.merge(u), }; Some(Box::new(OpenAICompatibleOAuthProvider { config: merged, client_name: client_name.to_string(), })) } _ => get_oauth_provider(provider_type), } } pub fn list_oauth_capable_clients(clients: &[ClientConfig]) -> Vec { clients .iter() .filter_map(|client_config| { let (name, _, auth) = client_config_info(client_config); if auth == Some("oauth") { Some(name.to_string()) } else { None } }) .collect() } pub(crate) fn client_config_info( client_config: &ClientConfig, ) -> (&str, &'static str, Option<&str>) { match client_config { ClientConfig::ClaudeConfig(c) => ( c.name.as_deref().unwrap_or("claude"), "claude", c.auth.as_deref(), ), ClientConfig::OpenAIConfig(c) => ( c.name.as_deref().unwrap_or("openai"), "openai", c.auth.as_deref(), ), ClientConfig::OpenAICompatibleConfig(c) => ( c.name.as_deref().unwrap_or("openai-compatible"), "openai-compatible", c.auth.as_deref(), ), ClientConfig::GeminiConfig(c) => ( c.name.as_deref().unwrap_or("gemini"), "gemini", c.auth.as_deref(), ), ClientConfig::CohereConfig(c) => (c.name.as_deref().unwrap_or("cohere"), "cohere", None), ClientConfig::AzureOpenAIConfig(c) => ( c.name.as_deref().unwrap_or("azure-openai"), "azure-openai", None, ), ClientConfig::VertexAIConfig(c) => { (c.name.as_deref().unwrap_or("vertexai"), "vertexai", None) } ClientConfig::BedrockConfig(c) => (c.name.as_deref().unwrap_or("bedrock"), "bedrock", None), ClientConfig::Unknown => ("unknown", "unknown", None), } } #[cfg(test)] mod tests { use std::ffi::OsString; use std::path::PathBuf; use std::str; use std::time::UNIX_EPOCH; use super::*; use crate::client::access_token::{distrust_access_token, get_access_token}; use crate::client::openai_compatible::OpenAICompatibleConfig; use crate::client::{ModelData, ProviderModels}; use crate::utils::get_env_name; use serial_test::serial; use std::{env, time::SystemTime}; fn with_temp_cache(f: F) { struct Restore { key: String, prev: Option, root: PathBuf, } impl Drop for Restore { fn drop(&mut self) { unsafe { match self.prev.take() { Some(v) => env::set_var(&self.key, v), None => env::remove_var(&self.key), } } let _ = fs::remove_dir_all(&self.root); } } let unique = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos(); let root = env::temp_dir().join(format!("coyote-client-oauth-test-{unique}")); fs::create_dir_all(&root).unwrap(); let env_key = get_env_name("cache_dir"); let prev = env::var_os(&env_key); unsafe { env::set_var(&env_key, &root); } let _restore = Restore { key: env_key, prev, root, }; f(); } fn base_config() -> OAuthConfig { OAuthConfig { client_id: "base-id".into(), token_url: "https://base.example/token".into(), flow: OAuthFlow::Pkce, client_secret: Some("base-secret".into()), authorize_url: Some("https://base.example/authorize".into()), redirect_uri: None, redirect_port: Some(1234), device_authorization_url: None, scopes: vec!["a".into(), "b".into()], token_request_format: Some(TokenRequestFormat::FormUrlEncoded), extra_authorize_params: IndexMap::from([("plan".into(), "base".into())]), extra_token_headers: IndexMap::new(), extra_request_headers: IndexMap::new(), echo_pkce_in_token_exchange: false, use_pkce_in_device_flow: false, include_state_in_token_exchange: true, } } fn empty_user_override(client_id: &str, token_url: &str) -> OAuthConfig { OAuthConfig { client_id: client_id.into(), token_url: token_url.into(), flow: OAuthFlow::Pkce, client_secret: None, authorize_url: None, redirect_uri: None, redirect_port: None, device_authorization_url: None, scopes: vec![], token_request_format: None, extra_authorize_params: IndexMap::new(), extra_token_headers: IndexMap::new(), extra_request_headers: IndexMap::new(), echo_pkce_in_token_exchange: false, use_pkce_in_device_flow: false, include_state_in_token_exchange: true, } } #[test] fn oauth_config_merge_user_wins_per_field() { let base = base_config(); let mut user = empty_user_override("user-id", "https://user.example/token"); user.client_secret = Some("user-secret".into()); user.scopes = vec!["c".into()]; user.extra_authorize_params = IndexMap::from([("plan".into(), "user".into())]); let merged = base.merge(user); assert_eq!(merged.client_id, "user-id"); assert_eq!(merged.token_url, "https://user.example/token"); assert_eq!(merged.client_secret.as_deref(), Some("user-secret")); assert_eq!( merged.authorize_url.as_deref(), Some("https://base.example/authorize") ); assert_eq!(merged.redirect_port, Some(1234)); assert_eq!(merged.scopes, vec!["c"]); assert_eq!( merged .extra_authorize_params .get("plan") .map(String::as_str), Some("user") ); } #[test] fn oauth_config_merge_empty_user_keeps_base_optionals() { let base = base_config(); let user = empty_user_override("user-id", "https://user.example/token"); let merged = base.merge(user); assert_eq!(merged.client_id, "user-id"); assert_eq!(merged.token_url, "https://user.example/token"); assert_eq!(merged.client_secret.as_deref(), Some("base-secret")); assert_eq!( merged.authorize_url.as_deref(), Some("https://base.example/authorize") ); assert_eq!(merged.redirect_port, Some(1234)); assert_eq!(merged.scopes, vec!["a", "b"]); assert!(matches!( merged.token_request_format, Some(TokenRequestFormat::FormUrlEncoded) )); assert_eq!( merged .extra_authorize_params .get("plan") .map(String::as_str), Some("base") ); } #[test] fn oauth_config_serde_roundtrip_from_yaml() { let yaml = r#" client_id: xai-client token_url: https://auth.x.ai/oauth2/token authorize_url: https://auth.x.ai/oauth2/authorize scopes: - openid - profile - api:access redirect_port: 56121 flow: pkce token_request_format: form_url_encoded extra_authorize_params: plan: generic referrer: coyote echo_pkce_in_token_exchange: true "#; let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); assert_eq!(cfg.client_id, "xai-client"); assert_eq!(cfg.token_url, "https://auth.x.ai/oauth2/token"); assert_eq!(cfg.scopes.len(), 3); assert_eq!(cfg.redirect_port, Some(56121)); assert!(matches!(cfg.flow, OAuthFlow::Pkce)); assert!(matches!( cfg.token_request_format, Some(TokenRequestFormat::FormUrlEncoded) )); assert!(cfg.echo_pkce_in_token_exchange); assert!(cfg.include_state_in_token_exchange); assert_eq!( cfg.extra_authorize_params.get("plan").map(String::as_str), Some("generic") ); } #[test] fn oauth_flow_defaults_to_pkce_when_missing() { let yaml = "client_id: x\ntoken_url: y"; let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); assert!(matches!(cfg.flow, OAuthFlow::Pkce)); } #[test] fn oauth_flow_client_credentials_parses() { let yaml = "client_id: x\ntoken_url: y\nflow: client_credentials"; let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); assert!(matches!(cfg.flow, OAuthFlow::ClientCredentials)); } fn make_provider_models(provider: &str, oauth: Option) -> ProviderModels { ProviderModels { provider: provider.into(), oauth, models: vec![ModelData::new("some-model")], } } fn make_openai_compat_client( name: &str, auth: Option<&str>, oauth: Option, ) -> ClientConfig { ClientConfig::OpenAICompatibleConfig(OpenAICompatibleConfig { name: Some(name.into()), api_base: Some("https://api.example/v1".into()), api_key: None, auth: auth.map(str::to_string), oauth: oauth.map(Box::new), models: vec![], patch: None, extra: None, }) } #[test] fn get_oauth_provider_for_client_merges_defaults_with_user_override() { let base = base_config(); let mut user = empty_user_override("user-id", "https://user.example/token"); user.echo_pkce_in_token_exchange = true; let models = vec![make_provider_models("acme", Some(base))]; let cc = make_openai_compat_client("acme", Some("oauth"), Some(user)); let provider = get_oauth_provider_for_client(&cc, &models).unwrap(); assert_eq!(provider.client_id(), "user-id"); assert_eq!(provider.token_url(), "https://user.example/token"); assert!(provider.echo_pkce_in_token_exchange()); assert_eq!( provider.fixed_redirect_uri().as_deref(), Some("http://127.0.0.1:1234/callback") ); } #[test] fn get_oauth_provider_for_client_uses_bundled_defaults_only() { let base = base_config(); let models = vec![make_provider_models("bundled-only", Some(base))]; let cc = make_openai_compat_client("bundled-only", Some("oauth"), None); let provider = get_oauth_provider_for_client(&cc, &models).unwrap(); assert_eq!(provider.client_id(), "base-id"); assert_eq!(provider.token_url(), "https://base.example/token"); } #[test] fn get_oauth_provider_for_client_uses_inline_only() { let user = OAuthConfig { client_id: "inline-id".into(), token_url: "https://inline.example/token".into(), ..empty_user_override("inline-id", "https://inline.example/token") }; let cc = make_openai_compat_client("inline-only", Some("oauth"), Some(user)); let provider = get_oauth_provider_for_client(&cc, &[]).unwrap(); assert_eq!(provider.client_id(), "inline-id"); } #[test] fn get_oauth_provider_for_client_returns_none_when_no_config_anywhere() { let cc = make_openai_compat_client("nothing", Some("oauth"), None); assert!(get_oauth_provider_for_client(&cc, &[]).is_none()); } #[test] fn get_oauth_provider_for_client_returns_none_when_auth_not_oauth() { let base = base_config(); let models = vec![make_provider_models("api-key-client", Some(base))]; let cc = make_openai_compat_client("api-key-client", None, None); assert!(get_oauth_provider_for_client(&cc, &models).is_none()); } #[test] fn openai_compatible_provider_joins_scopes_with_spaces() { let mut cfg = base_config(); cfg.scopes = vec!["one".into(), "two".into(), "three".into()]; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert_eq!(provider.scopes(), "one two three"); } #[test] fn openai_compatible_provider_prefers_redirect_uri_over_port() { let mut cfg = base_config(); cfg.redirect_uri = Some("http://127.0.0.1:9000/cb".into()); cfg.redirect_port = Some(9999); let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert_eq!( provider.fixed_redirect_uri().as_deref(), Some("http://127.0.0.1:9000/cb") ); } #[test] fn openai_compatible_provider_ephemeral_when_no_redirect() { let mut cfg = base_config(); cfg.redirect_uri = None; cfg.redirect_port = None; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert!(provider.uses_localhost_redirect()); assert!(provider.fixed_redirect_uri().is_none()); } #[test] fn default_extra_token_params_is_empty() { let provider = OpenAICompatibleOAuthProvider { config: base_config(), client_name: "test".into(), }; assert!(provider.extra_token_params().is_empty()); } #[test] fn oauth_flow_device_code_parses() { let yaml = "client_id: x\ntoken_url: y\nflow: device_code"; let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); assert!(matches!(cfg.flow, OAuthFlow::DeviceCode)); } #[test] fn oauth_config_merge_preserves_device_authorization_url_when_user_omits() { let mut base = base_config(); base.device_authorization_url = Some("https://base.example/device".into()); let user = empty_user_override("user-id", "https://user.example/token"); let merged = base.merge(user); assert_eq!( merged.device_authorization_url.as_deref(), Some("https://base.example/device") ); } #[test] fn oauth_config_merge_user_device_authorization_url_wins() { let mut base = base_config(); base.device_authorization_url = Some("https://base.example/device".into()); let mut user = empty_user_override("user-id", "https://user.example/token"); user.device_authorization_url = Some("https://user.example/device".into()); let merged = base.merge(user); assert_eq!( merged.device_authorization_url.as_deref(), Some("https://user.example/device") ); } #[test] fn oauth_config_merge_user_pkce_in_device_flow_wins() { let base = base_config(); let mut user = empty_user_override("user-id", "https://user.example/token"); user.use_pkce_in_device_flow = true; let merged = base.merge(user); assert!(merged.use_pkce_in_device_flow); } #[test] fn openai_compatible_provider_exposes_device_authorization_url() { let mut cfg = base_config(); cfg.device_authorization_url = Some("https://example/device".into()); let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert_eq!( provider.device_authorization_url(), Some("https://example/device") ); } #[test] fn openai_compatible_provider_device_authorization_url_none_when_unset() { let provider = OpenAICompatibleOAuthProvider { config: base_config(), client_name: "test".into(), }; assert!(provider.device_authorization_url().is_none()); } #[test] fn openai_compatible_provider_use_pkce_in_device_flow_defaults_false() { let provider = OpenAICompatibleOAuthProvider { config: base_config(), client_name: "test".into(), }; assert!(!provider.use_pkce_in_device_flow()); } #[test] fn openai_compatible_provider_use_pkce_in_device_flow_returns_true_when_set() { let mut cfg = base_config(); cfg.use_pkce_in_device_flow = true; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert!(provider.use_pkce_in_device_flow()); } #[test] fn parse_paste_input_full_callback_url() { let (code, state) = parse_paste_input("https://provider.example/oauth/callback?code=abc123&state=xyz") .unwrap(); assert_eq!(code, "abc123"); assert_eq!(state.as_deref(), Some("xyz")); } #[test] fn parse_paste_input_url_without_state() { let (code, state) = parse_paste_input("https://provider.example/oauth/callback?code=abc123").unwrap(); assert_eq!(code, "abc123"); assert!(state.is_none()); } #[test] fn parse_paste_input_code_state_fragment() { let (code, state) = parse_paste_input("abc123#xyz").unwrap(); assert_eq!(code, "abc123"); assert_eq!(state.as_deref(), Some("xyz")); } #[test] fn parse_paste_input_bare_code() { let (code, state) = parse_paste_input("abc123").unwrap(); assert_eq!(code, "abc123"); assert!(state.is_none()); } #[test] fn parse_paste_input_url_missing_code_fails() { let err = parse_paste_input("https://provider.example/oauth/callback?state=xyz").unwrap_err(); assert!(err.to_string().contains("code"), "unexpected error: {err}"); } #[test] fn parse_paste_input_empty_fails() { let err = parse_paste_input("").unwrap_err(); assert!(err.to_string().contains("Empty"), "unexpected error: {err}"); } #[test] fn openai_compatible_provider_fixed_redirect_uri_none_for_public_url() { let mut cfg = base_config(); cfg.redirect_uri = Some("https://provider.example.com/callback".into()); cfg.redirect_port = None; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert!(provider.fixed_redirect_uri().is_none()); } #[test] fn openai_compatible_provider_fixed_redirect_uri_some_for_localhost() { let mut cfg = base_config(); cfg.redirect_uri = Some("http://127.0.0.1:9999/cb".into()); cfg.redirect_port = None; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert_eq!( provider.fixed_redirect_uri().as_deref(), Some("http://127.0.0.1:9999/cb") ); } #[test] fn openai_compatible_provider_fixed_redirect_uri_some_for_localhost_hostname() { let mut cfg = base_config(); cfg.redirect_uri = Some("http://localhost:9999/cb".into()); cfg.redirect_port = None; let provider = OpenAICompatibleOAuthProvider { config: cfg, client_name: "test".into(), }; assert_eq!( provider.fixed_redirect_uri().as_deref(), Some("http://localhost:9999/cb") ); } #[test] fn oauth_config_serde_roundtrip_device_code_yaml() { let yaml = r#" client_id: my-client token_url: https://auth.example/oauth/token device_authorization_url: https://auth.example/oauth/device_authorization flow: device_code token_request_format: form_url_encoded use_pkce_in_device_flow: true scopes: - read - write "#; let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); assert_eq!(cfg.client_id, "my-client"); assert_eq!(cfg.token_url, "https://auth.example/oauth/token"); assert_eq!( cfg.device_authorization_url.as_deref(), Some("https://auth.example/oauth/device_authorization") ); assert!(matches!(cfg.flow, OAuthFlow::DeviceCode)); assert!(matches!( cfg.token_request_format, Some(TokenRequestFormat::FormUrlEncoded) )); assert!(cfg.use_pkce_in_device_flow); assert_eq!(cfg.scopes, vec!["read", "write"]); } struct ResourceStubProvider; impl OAuthProvider for ResourceStubProvider { fn provider_name(&self) -> &str { "stub" } fn client_id(&self) -> &str { "stub-client" } fn authorize_url(&self) -> &str { "https://as.example/authorize" } fn token_url(&self) -> &str { "https://as.example/token" } fn redirect_uri(&self) -> &str { "" } fn scopes(&self) -> String { String::new() } fn token_request_format(&self) -> TokenRequestFormat { TokenRequestFormat::FormUrlEncoded } fn extra_token_params(&self) -> Vec<(&str, &str)> { vec![("resource", "https://rs.example/mcp")] } } #[test] fn build_token_request_appends_extra_token_params_to_form_body() { let provider = ResourceStubProvider; let request = build_token_request( &ReqwestClient::new(), &provider, &[("grant_type", "authorization_code")], ) .build() .unwrap(); let body = str::from_utf8(request.body().unwrap().as_bytes().unwrap()).unwrap(); assert!( body.contains("resource=https%3A%2F%2Frs.example%2Fmcp"), "body missing resource param: {body}" ); assert!( body.contains("grant_type=authorization_code"), "body missing grant_type param: {body}" ); } #[test] #[serial] fn save_oauth_tokens_roundtrips_and_leaves_no_tmp_file() { with_temp_cache(|| { let tokens = OAuthTokens { access_token: "at-123".into(), refresh_token: Some("rt-456".into()), expires_at: 1234567890, account_id: Some("acct-789".into()), }; save_oauth_tokens("atomic-test", &tokens).unwrap(); let loaded = load_oauth_tokens("atomic-test").unwrap(); assert_eq!(loaded.access_token, "at-123"); assert_eq!(loaded.refresh_token.as_deref(), Some("rt-456")); assert_eq!(loaded.expires_at, 1234567890); assert_eq!(loaded.account_id.as_deref(), Some("acct-789")); let dir = paths::oauth_tokens_dir(); let leftover_tmp = fs::read_dir(&dir) .unwrap() .any(|e| e.unwrap().file_name().to_string_lossy().ends_with(".tmp")); assert!(!leftover_tmp, "temp file left behind in {dir:?}"); #[cfg(unix)] { use std::os::unix::fs::PermissionsExt; let mode = fs::metadata(paths::token_file("atomic-test")) .unwrap() .permissions() .mode(); assert_eq!(mode & 0o777, 0o600, "token file mode was {mode:o}"); } }); } #[test] #[serial] fn prepare_rejected_valid_file_token_attempts_refresh_branch() { with_temp_cache(|| { let client_name = "prepare-rejected-branch-test"; let expires_at = Utc::now().timestamp() + 3600; save_oauth_tokens( client_name, &OAuthTokens { access_token: "rejected-at".into(), refresh_token: None, expires_at, account_id: None, }, ) .unwrap(); set_access_token(client_name, "rejected-at".into(), expires_at, None); assert!(distrust_access_token(client_name, "rejected-at")); let err = tokio::runtime::Builder::new_current_thread() .enable_all() .build() .unwrap() .block_on(prepare_oauth_access_token( &ReqwestClient::new(), &ResourceStubProvider, client_name, )) .unwrap_err() .to_string(); // The timestamp-valid but rejected file token must not be trusted; // the refresh branch is taken and bails on the missing refresh token. assert!(err.contains("No refresh token"), "unexpected error: {err}"); }); } #[test] #[serial] fn prepare_trusts_differing_unmarked_valid_file_token() { with_temp_cache(|| { let client_name = "prepare-differing-token-test"; let expires_at = Utc::now().timestamp() + 3600; set_access_token(client_name, "rejected-at".into(), expires_at, None); assert!(distrust_access_token(client_name, "rejected-at")); save_oauth_tokens( client_name, &OAuthTokens { access_token: "fresh-at".into(), refresh_token: None, expires_at, account_id: None, }, ) .unwrap(); let ready = tokio::runtime::Builder::new_current_thread() .enable_all() .build() .unwrap() .block_on(prepare_oauth_access_token( &ReqwestClient::new(), &ResourceStubProvider, client_name, )) .unwrap(); assert!(ready); assert_eq!(get_access_token(client_name).unwrap(), "fresh-at"); assert!( !is_rejected(client_name, "rejected-at"), "marker not cleared after successful prepare" ); }); } #[test] fn parse_refresh_response_invalid_grant_redacts_and_prompts_reauth() { let response = serde_json::json!({ "error": "invalid_grant", "error_description": "refresh token revoked", "refresh_token": "planted-secret-token", }); let err = parse_refresh_response(StatusCode::BAD_REQUEST, &response, Some("old-rt")) .unwrap_err() .to_string(); assert!(err.contains("re-authenticate"), "unexpected error: {err}"); assert!( !err.contains("planted-secret-token"), "error leaked token material: {err}" ); } #[test] fn token_response_keys_lists_keys_without_values() { let response = serde_json::json!({ "access_token": "secret-at", "token_type": "SecretBearer", }); let keys = token_response_keys(&response); assert!(keys.contains("access_token"), "missing key name: {keys}"); assert!(keys.contains("token_type"), "missing key name: {keys}"); assert!(!keys.contains("secret-at"), "leaked value: {keys}"); assert!(!keys.contains("SecretBearer"), "leaked value: {keys}"); } #[test] fn parse_refresh_response_rotates_refresh_token_when_present() { let response = serde_json::json!({ "access_token": "new-at", "refresh_token": "new-rt", "expires_in": 3600, }); let (access_token, refresh_token, expires_in) = parse_refresh_response(StatusCode::OK, &response, Some("old-rt")).unwrap(); assert_eq!(access_token, "new-at"); assert_eq!(refresh_token.as_deref(), Some("new-rt")); assert_eq!(expires_in, 3600); } #[test] fn parse_refresh_response_keeps_old_refresh_token_when_absent() { let response = serde_json::json!({ "access_token": "new-at", "expires_in": 3600, }); let (_, refresh_token, _) = parse_refresh_response(StatusCode::OK, &response, Some("old-rt")).unwrap(); assert_eq!(refresh_token.as_deref(), Some("old-rt")); } }