fix(extension-manager): forward custom headers through OAuth connect path (#9388)

Signed-off-by: Cameron Yick <cameron.yick@datadoghq.com>
This commit is contained in:
Cameron Yick
2026-06-01 21:16:42 -04:00
committed by GitHub
parent cd8f718bbb
commit 586bb15d40
+219 -1
View File
@@ -510,6 +510,7 @@ async fn connect_with_auth(
auth_manager: rmcp::transport::AuthorizationManager,
uri: &str,
timeout: Duration,
headers: &HashMap<String, String>,
provider: SharedProvider,
client_name: String,
capabilities: GooseMcpClientCapabilities,
@@ -517,6 +518,15 @@ async fn connect_with_auth(
) -> ExtensionResult<Box<dyn McpClientTrait>> {
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<String, String>,
name: &str,
socket: Option<&str>,
credential_store: Box<dyn CredentialStore>,
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"
);
}
}