diff --git a/crates/goose/src/providers/oauth.rs b/crates/goose/src/providers/oauth.rs index 1c4bac231..1c207230a 100644 --- a/crates/goose/src/providers/oauth.rs +++ b/crates/goose/src/providers/oauth.rs @@ -6,15 +6,68 @@ use axum::{extract::Query, response::Html, routing::get, Router}; use base64::Engine; use chrono::{DateTime, Utc}; use once_cell::sync::Lazy; +use reqwest::{redirect::Policy, Client}; use serde::{Deserialize, Serialize}; use serde_json::Value; use sha2::Digest; use std::{collections::HashMap, fs, net::SocketAddr, path::PathBuf, sync::Arc}; use tokio::sync::{oneshot, Mutex as TokioMutex}; -use url::Url; +use url::{Host, Url}; static OAUTH_MUTEX: Lazy> = Lazy::new(|| TokioMutex::new(())); +#[derive(Clone, Copy)] +enum EndpointTransport { + Https, + LoopbackHttp, +} + +fn endpoint_transport(url: &Url) -> Result { + if url.scheme() == "https" { + return Ok(EndpointTransport::Https); + } + + let is_loopback = match url.host() { + Some(Host::Domain(domain)) => domain.eq_ignore_ascii_case("localhost"), + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + None => false, + }; + + if url.scheme() == "http" && is_loopback { + Ok(EndpointTransport::LoopbackHttp) + } else { + anyhow::bail!("OAuth endpoint must use HTTPS unless it targets loopback") + } +} + +fn secure_client_for_endpoint(endpoint: &Url) -> Result { + let builder = Client::builder(); + match endpoint_transport(endpoint)? { + EndpointTransport::Https => Ok(builder.https_only(true).build()?), + EndpointTransport::LoopbackHttp => Ok(builder + .no_proxy() + .redirect(Policy::custom(|attempt| { + if attempt.previous().len() >= 10 { + return attempt.error("too many redirects"); + } + + if endpoint_transport(attempt.url()).is_ok() { + attempt.follow() + } else { + attempt.error("redirect violates the OAuth endpoint transport policy") + } + })) + .build()?), + } +} + +fn secure_endpoint(raw_endpoint: &str) -> Result<(Url, Client)> { + let endpoint = Url::parse(raw_endpoint)?; + let client = secure_client_for_endpoint(&endpoint)?; + Ok((endpoint, client)) +} + #[derive(Debug, Clone)] struct OidcEndpoints { authorization_endpoint: String, @@ -96,12 +149,12 @@ impl TokenCache { } async fn get_workspace_endpoints(host: &str) -> Result { - let base_url = Url::parse(host).expect("Invalid host URL"); + let base_url = Url::parse(host)?; let oidc_url = base_url .join("oidc/.well-known/oauth-authorization-server") - .expect("Invalid OIDC URL"); + .map_err(|error| anyhow::anyhow!("Invalid OIDC URL: {error}"))?; - let client = reqwest::Client::new(); + let client = secure_client_for_endpoint(&oidc_url)?; let resp = client.get(oidc_url.clone()).send().await?; if !resp.status().is_success() { @@ -124,6 +177,7 @@ async fn get_workspace_endpoints(host: &str) -> Result { .and_then(|v| v.as_str()) .ok_or_else(|| anyhow::anyhow!("token_endpoint not found in OIDC configuration"))? .to_string(); + secure_endpoint(&token_endpoint)?; Ok(OidcEndpoints { authorization_endpoint, @@ -249,9 +303,9 @@ impl OAuthFlow { ("client_id", &self.client_id), ]; - let client = reqwest::Client::new(); + let (token_endpoint, client) = secure_endpoint(&self.endpoints.token_endpoint)?; let resp = client - .post(&self.endpoints.token_endpoint) + .post(token_endpoint) .header("Content-Type", "application/x-www-form-urlencoded") .form(¶ms) .send() @@ -278,9 +332,9 @@ impl OAuthFlow { tracing::debug!("Refreshing token using refresh_token"); - let client = reqwest::Client::new(); + let (token_endpoint, client) = secure_endpoint(&self.endpoints.token_endpoint)?; let resp = client - .post(&self.endpoints.token_endpoint) + .post(token_endpoint) .header("Content-Type", "application/x-www-form-urlencoded") .form(¶ms) .send() @@ -615,4 +669,179 @@ mod tests { Ok(()) } + + #[tokio::test] + async fn refresh_token_rejects_plaintext_remote_endpoint() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "attacker-token", + "expires_in": 3600 + }))) + .mount(&server) + .await; + + let flow = OAuthFlow::new( + OidcEndpoints { + authorization_endpoint: "https://workspace.example/authorize".to_string(), + token_endpoint: format!("{}/token", server.uri()).replace("127.0.0.1", "0.0.0.0"), + }, + "databricks-cli".to_string(), + "http://localhost".to_string(), + vec!["all-apis".to_string(), "offline_access".to_string()], + ); + + let result = flow.refresh_token("long-lived-refresh-secret").await; + + assert!( + result.is_err(), + "plaintext remote token endpoint was accepted" + ); + assert!(server.received_requests().await.unwrap().is_empty()); + } + + #[tokio::test] + async fn code_exchange_rejects_plaintext_remote_endpoint() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "attacker-token", + "refresh_token": "attacker-refresh", + "expires_in": 3600 + }))) + .mount(&server) + .await; + + let flow = OAuthFlow::new( + OidcEndpoints { + authorization_endpoint: "https://workspace.example/authorize".to_string(), + token_endpoint: format!("{}/token", server.uri()).replace("127.0.0.1", "0.0.0.0"), + }, + "databricks-cli".to_string(), + "http://localhost".to_string(), + vec!["all-apis".to_string(), "offline_access".to_string()], + ); + + let result = flow + .exchange_code_for_token_with_redirect("authorization-code", "http://localhost") + .await; + + assert!( + result.is_err(), + "plaintext remote token endpoint was accepted" + ); + assert!(server.received_requests().await.unwrap().is_empty()); + } + + #[tokio::test] + async fn refresh_token_accepts_loopback_endpoint() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "loopback-token", + "expires_in": 3600 + }))) + .mount(&server) + .await; + + let flow = OAuthFlow::new( + OidcEndpoints { + authorization_endpoint: "https://workspace.example/authorize".to_string(), + token_endpoint: format!("{}/token", server.uri()), + }, + "databricks-cli".to_string(), + "http://localhost".to_string(), + vec!["all-apis".to_string(), "offline_access".to_string()], + ); + + let token = flow.refresh_token("loopback-refresh-secret").await.unwrap(); + + assert_eq!(token.access_token, "loopback-token"); + assert_eq!(server.received_requests().await.unwrap().len(), 1); + } + + #[tokio::test] + async fn code_exchange_accepts_loopback_endpoint() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "loopback-token", + "refresh_token": "loopback-refresh-token", + "expires_in": 3600 + }))) + .mount(&server) + .await; + + let flow = OAuthFlow::new( + OidcEndpoints { + authorization_endpoint: "https://workspace.example/authorize".to_string(), + token_endpoint: format!("{}/token", server.uri()), + }, + "databricks-cli".to_string(), + "http://localhost".to_string(), + vec!["all-apis".to_string(), "offline_access".to_string()], + ); + + let token = flow + .exchange_code_for_token_with_redirect("authorization-code", "http://localhost") + .await + .unwrap(); + + assert_eq!(token.access_token, "loopback-token"); + assert_eq!(server.received_requests().await.unwrap().len(), 1); + } + + #[tokio::test] + async fn discovery_rejects_plaintext_remote_token_endpoint() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/oidc/.well-known/oauth-authorization-server")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "authorization_endpoint": "https://workspace.example/authorize", + "token_endpoint": format!("{}/token", server.uri()).replace("127.0.0.1", "0.0.0.0") + }))) + .mount(&server) + .await; + + let result = get_workspace_endpoints(&server.uri()).await; + + assert!( + result.is_err(), + "plaintext remote token endpoint was accepted" + ); + } + + #[tokio::test] + async fn loopback_token_endpoint_rejects_remote_plaintext_redirects() { + for status in [307, 308] { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/token")) + .respond_with(ResponseTemplate::new(status).insert_header( + "Location", + format!("{}/redirected", server.uri()).replace("127.0.0.1", "0.0.0.0"), + )) + .mount(&server) + .await; + + let flow = OAuthFlow::new( + OidcEndpoints { + authorization_endpoint: "https://workspace.example/authorize".to_string(), + token_endpoint: format!("{}/token", server.uri()), + }, + "databricks-cli".to_string(), + "http://localhost".to_string(), + vec!["all-apis".to_string(), "offline_access".to_string()], + ); + + let result = flow.refresh_token("loopback-refresh-secret").await; + + assert!(result.is_err(), "remote plaintext redirect was followed"); + assert_eq!(server.received_requests().await.unwrap().len(), 1); + } + } }