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:
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user