distrust_access_token compare-and-invalidates the in-memory entry only when the cached token equals the rejected one, so a concurrent refresh is never clobbered. is_valid_access_token and both expiry checks in prepare_oauth_access_token treat marked tokens as expired, forcing a refresh of provider-rejected tokens that are still locally unexpired. The marker is cleared after every completed refresh, including ones that return the same token.
1839 lines
59 KiB
Rust
1839 lines
59 KiB
Rust
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<String>,
|
|
pub authorize_url: Option<String>,
|
|
pub redirect_uri: Option<String>,
|
|
pub redirect_port: Option<u16>,
|
|
pub device_authorization_url: Option<String>,
|
|
#[serde(default)]
|
|
pub scopes: Vec<String>,
|
|
pub token_request_format: Option<TokenRequestFormat>,
|
|
#[serde(default)]
|
|
pub extra_authorize_params: IndexMap<String, String>,
|
|
#[serde(default)]
|
|
pub extra_token_headers: IndexMap<String, String>,
|
|
#[serde(default)]
|
|
pub extra_request_headers: IndexMap<String, String>,
|
|
#[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<String> {
|
|
None
|
|
}
|
|
|
|
fn extract_account_id(&self, _response: &Value) -> Option<String> {
|
|
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<String>,
|
|
pub expires_at: i64,
|
|
#[serde(default)]
|
|
pub account_id: Option<String>,
|
|
}
|
|
|
|
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::<qrcode::render::unicode::Dense1x2>()
|
|
.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<OAuthTokens> {
|
|
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 => "<non-object response>".to_string(),
|
|
}
|
|
}
|
|
|
|
fn parse_refresh_response(
|
|
status: StatusCode,
|
|
response: &Value,
|
|
previous_refresh_token: Option<&str>,
|
|
) -> Result<(String, Option<String>, 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<OAuthTokens> {
|
|
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<sync::Mutex<()>> {
|
|
static GUARDS: OnceLock<parking_lot::Mutex<HashMap<String, Arc<sync::Mutex<()>>>>> =
|
|
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<bool> {
|
|
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<String, Value> = 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<String, String> = 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<String>)> {
|
|
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::<Url>() 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 = "<html><body><h2>Authentication successful!</h2><p>You can close this tab and return to your terminal.</p></body></html>";
|
|
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<Box<dyn OAuthProvider>> {
|
|
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<Box<dyn OAuthProvider>> {
|
|
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<String> {
|
|
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: FnOnce()>(f: F) {
|
|
struct Restore {
|
|
key: String,
|
|
prev: Option<OsString>,
|
|
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<OAuthConfig>) -> ProviderModels {
|
|
ProviderModels {
|
|
provider: provider.into(),
|
|
oauth,
|
|
models: vec![ModelData::new("some-model")],
|
|
}
|
|
}
|
|
|
|
fn make_openai_compat_client(
|
|
name: &str,
|
|
auth: Option<&str>,
|
|
oauth: Option<OAuthConfig>,
|
|
) -> 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"));
|
|
}
|
|
}
|