From 586bb15d40f02fafe15aeca4149ab846d9606d92 Mon Sep 17 00:00:00 2001 From: Cameron Yick Date: Mon, 1 Jun 2026 21:16:42 -0400 Subject: [PATCH] fix(extension-manager): forward custom headers through OAuth connect path (#9388) Signed-off-by: Cameron Yick --- crates/goose/src/agents/extension_manager.rs | 220 ++++++++++++++++++- 1 file changed, 219 insertions(+), 1 deletion(-) diff --git a/crates/goose/src/agents/extension_manager.rs b/crates/goose/src/agents/extension_manager.rs index da7e2ce9..cf7f6764 100644 --- a/crates/goose/src/agents/extension_manager.rs +++ b/crates/goose/src/agents/extension_manager.rs @@ -510,6 +510,7 @@ async fn connect_with_auth( auth_manager: rmcp::transport::AuthorizationManager, uri: &str, timeout: Duration, + headers: &HashMap, provider: SharedProvider, client_name: String, capabilities: GooseMcpClientCapabilities, @@ -517,6 +518,15 @@ async fn connect_with_auth( ) -> ExtensionResult> { let mut auth_headers = HeaderMap::new(); auth_headers.insert(reqwest::header::USER_AGENT, GOOSE_USER_AGENT); + for (key, value) in headers { + auth_headers.insert( + HeaderName::try_from(key) + .map_err(|_| ExtensionError::ConfigError(format!("invalid header: {}", key)))?, + value.parse().map_err(|_| { + ExtensionError::ConfigError(format!("invalid header value: {}", key)) + })?, + ); + } #[allow(unused_mut)] let mut auth_client_builder = reqwest::Client::builder().default_headers(auth_headers); #[cfg(target_os = "linux")] @@ -551,6 +561,7 @@ async fn create_streamable_http_client( headers: &HashMap, name: &str, socket: Option<&str>, + credential_store: Box, provider: SharedProvider, client_name: String, capabilities: GooseMcpClientCapabilities, @@ -611,7 +622,6 @@ async fn create_streamable_http_client( // If we have stored OAuth credentials, try refreshing and connecting directly. // This avoids the unnecessary 401 → browser re-auth cycle on every new session. - let credential_store = GooseCredentialStore::new(name.to_string()); if credential_store.load().await.is_ok_and(|c| c.is_some()) { match oauth_flow(&uri.to_string(), &name.to_string()).await { Ok(auth_manager) => { @@ -619,6 +629,7 @@ async fn create_streamable_http_client( auth_manager, uri, timeout_duration, + headers, provider, client_name, capabilities, @@ -652,6 +663,7 @@ async fn create_streamable_http_client( auth_manager, uri, timeout_duration, + headers, provider, client_name, capabilities, @@ -854,6 +866,7 @@ impl ExtensionManager { &resolved_headers, name, resolved_socket.as_deref(), + Box::new(GooseCredentialStore::new(name.to_string())), self.provider.clone(), self.client_name.clone(), self.mcp_client_capabilities(), @@ -2711,4 +2724,209 @@ mod tests { ); assert!(should_attempt_oauth_fallback(&Err(err))); } + + #[tokio::test] + async fn test_invalid_header_name_returns_config_error() { + let mut headers = HashMap::new(); + headers.insert("bad header name".to_string(), "value".to_string()); + + let temp_dir = tempdir().unwrap(); + let provider: SharedProvider = Arc::new(Mutex::new(None)); + let capabilities = GooseMcpClientCapabilities { + mcpui: false, + host_info: None, + }; + + let result = create_streamable_http_client( + "http://localhost:1", + None, + &headers, + "test-ext", + None, + Box::new(rmcp::transport::auth::InMemoryCredentialStore::new()), + provider, + "goose-test".to_string(), + capabilities, + temp_dir.path(), + ) + .await; + + let Err(ExtensionError::ConfigError(msg)) = result else { + panic!("expected ConfigError, got a different result"); + }; + assert!( + msg.contains("invalid header"), + "unexpected error message: {msg}" + ); + } + + #[tokio::test] + async fn test_invalid_header_value_returns_config_error() { + let mut headers = HashMap::new(); + headers.insert("x-valid-name".to_string(), "bad\r\nvalue".to_string()); + + let temp_dir = tempdir().unwrap(); + let provider: SharedProvider = Arc::new(Mutex::new(None)); + let capabilities = GooseMcpClientCapabilities { + mcpui: false, + host_info: None, + }; + + let result = create_streamable_http_client( + "http://localhost:1", + None, + &headers, + "test-ext", + None, + Box::new(rmcp::transport::auth::InMemoryCredentialStore::new()), + provider, + "goose-test".to_string(), + capabilities, + temp_dir.path(), + ) + .await; + + let Err(ExtensionError::ConfigError(msg)) = result else { + panic!("expected ConfigError, got a different result"); + }; + assert!( + msg.contains("invalid header value"), + "unexpected error message: {msg}" + ); + } + + #[tokio::test] + async fn test_custom_headers_forwarded_to_http_extension() { + use wiremock::matchers::any; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let mock_server = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(200)) + .mount(&mock_server) + .await; + + let mut headers = HashMap::new(); + headers.insert("x-api-key".to_string(), "test-secret-123".to_string()); + + let temp_dir = tempdir().unwrap(); + let provider: SharedProvider = Arc::new(Mutex::new(None)); + let capabilities = GooseMcpClientCapabilities { + mcpui: false, + host_info: None, + }; + + // The MCP handshake will fail against the stub server. We only care that + // the outgoing HTTP request carried the custom header. + let _ = create_streamable_http_client( + &mock_server.uri(), + None, + &headers, + "test-ext", + None, + Box::new(rmcp::transport::auth::InMemoryCredentialStore::new()), + provider, + "goose-test".to_string(), + capabilities, + temp_dir.path(), + ) + .await; + + let received = mock_server.received_requests().await.unwrap(); + assert!( + !received.is_empty(), + "expected at least one HTTP request to reach the mock server" + ); + let header_found = received.iter().any(|req| { + req.headers + .get("x-api-key") + .map(|v| v == "test-secret-123") + .unwrap_or(false) + }); + assert!( + header_found, + "custom header x-api-key was not forwarded to the extension server" + ); + } + + /// Directly exercises `connect_with_auth`, which is the code path fixed by + /// the PR (custom headers were dropped when the OAuth connection path was + /// taken). Uses a pre-seeded `InMemoryCredentialStore` with a fake, + /// non-expiring token so `get_access_token()` returns immediately without + /// touching any OAuth endpoints or the system keychain. + #[tokio::test] + async fn test_custom_headers_forwarded_oauth_path() { + use rmcp::transport::auth::{ + InMemoryCredentialStore, OAuthTokenResponse, StoredCredentials, + }; + use wiremock::matchers::any; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let mock_server = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(200)) + .mount(&mock_server) + .await; + + let mut headers = HashMap::new(); + headers.insert("x-api-key".to_string(), "test-secret-oauth".to_string()); + + // Build a fake, non-expiring token. token_received_at=None skips the + // expiry check, so get_access_token() returns without any network call. + let token_response: OAuthTokenResponse = serde_json::from_value(serde_json::json!({ + "access_token": "fake-test-token", + "token_type": "bearer", + })) + .expect("valid fake token JSON"); + let creds = StoredCredentials::new( + "test-client".to_string(), + Some(token_response), + vec![], + None, + ); + let store = InMemoryCredentialStore::new(); + store.save(creds).await.unwrap(); + + let mut auth_manager = rmcp::transport::AuthorizationManager::new(mock_server.uri()) + .await + .expect("AuthorizationManager::new should not make network calls"); + auth_manager.set_credential_store(store); + + let temp_dir = tempdir().unwrap(); + let provider: SharedProvider = Arc::new(Mutex::new(None)); + let capabilities = GooseMcpClientCapabilities { + mcpui: false, + host_info: None, + }; + + // connect_with_auth will fail (mock server isn't an MCP server) but we + // only care that the outgoing request carried the custom header. + let _ = connect_with_auth( + auth_manager, + &mock_server.uri(), + Duration::from_secs(5), + &headers, + provider, + "goose-test".to_string(), + capabilities, + temp_dir.path(), + ) + .await; + + let received = mock_server.received_requests().await.unwrap(); + assert!( + !received.is_empty(), + "expected at least one HTTP request to reach the mock server" + ); + let header_found = received.iter().any(|req| { + req.headers + .get("x-api-key") + .map(|v| v == "test-secret-oauth") + .unwrap_or(false) + }); + assert!( + header_found, + "custom header x-api-key was not forwarded through the OAuth connection path" + ); + } }