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,
|
auth_manager: rmcp::transport::AuthorizationManager,
|
||||||
uri: &str,
|
uri: &str,
|
||||||
timeout: Duration,
|
timeout: Duration,
|
||||||
|
headers: &HashMap<String, String>,
|
||||||
provider: SharedProvider,
|
provider: SharedProvider,
|
||||||
client_name: String,
|
client_name: String,
|
||||||
capabilities: GooseMcpClientCapabilities,
|
capabilities: GooseMcpClientCapabilities,
|
||||||
@@ -517,6 +518,15 @@ async fn connect_with_auth(
|
|||||||
) -> ExtensionResult<Box<dyn McpClientTrait>> {
|
) -> ExtensionResult<Box<dyn McpClientTrait>> {
|
||||||
let mut auth_headers = HeaderMap::new();
|
let mut auth_headers = HeaderMap::new();
|
||||||
auth_headers.insert(reqwest::header::USER_AGENT, GOOSE_USER_AGENT);
|
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)]
|
#[allow(unused_mut)]
|
||||||
let mut auth_client_builder = reqwest::Client::builder().default_headers(auth_headers);
|
let mut auth_client_builder = reqwest::Client::builder().default_headers(auth_headers);
|
||||||
#[cfg(target_os = "linux")]
|
#[cfg(target_os = "linux")]
|
||||||
@@ -551,6 +561,7 @@ async fn create_streamable_http_client(
|
|||||||
headers: &HashMap<String, String>,
|
headers: &HashMap<String, String>,
|
||||||
name: &str,
|
name: &str,
|
||||||
socket: Option<&str>,
|
socket: Option<&str>,
|
||||||
|
credential_store: Box<dyn CredentialStore>,
|
||||||
provider: SharedProvider,
|
provider: SharedProvider,
|
||||||
client_name: String,
|
client_name: String,
|
||||||
capabilities: GooseMcpClientCapabilities,
|
capabilities: GooseMcpClientCapabilities,
|
||||||
@@ -611,7 +622,6 @@ async fn create_streamable_http_client(
|
|||||||
|
|
||||||
// If we have stored OAuth credentials, try refreshing and connecting directly.
|
// If we have stored OAuth credentials, try refreshing and connecting directly.
|
||||||
// This avoids the unnecessary 401 → browser re-auth cycle on every new session.
|
// 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()) {
|
if credential_store.load().await.is_ok_and(|c| c.is_some()) {
|
||||||
match oauth_flow(&uri.to_string(), &name.to_string()).await {
|
match oauth_flow(&uri.to_string(), &name.to_string()).await {
|
||||||
Ok(auth_manager) => {
|
Ok(auth_manager) => {
|
||||||
@@ -619,6 +629,7 @@ async fn create_streamable_http_client(
|
|||||||
auth_manager,
|
auth_manager,
|
||||||
uri,
|
uri,
|
||||||
timeout_duration,
|
timeout_duration,
|
||||||
|
headers,
|
||||||
provider,
|
provider,
|
||||||
client_name,
|
client_name,
|
||||||
capabilities,
|
capabilities,
|
||||||
@@ -652,6 +663,7 @@ async fn create_streamable_http_client(
|
|||||||
auth_manager,
|
auth_manager,
|
||||||
uri,
|
uri,
|
||||||
timeout_duration,
|
timeout_duration,
|
||||||
|
headers,
|
||||||
provider,
|
provider,
|
||||||
client_name,
|
client_name,
|
||||||
capabilities,
|
capabilities,
|
||||||
@@ -854,6 +866,7 @@ impl ExtensionManager {
|
|||||||
&resolved_headers,
|
&resolved_headers,
|
||||||
name,
|
name,
|
||||||
resolved_socket.as_deref(),
|
resolved_socket.as_deref(),
|
||||||
|
Box::new(GooseCredentialStore::new(name.to_string())),
|
||||||
self.provider.clone(),
|
self.provider.clone(),
|
||||||
self.client_name.clone(),
|
self.client_name.clone(),
|
||||||
self.mcp_client_capabilities(),
|
self.mcp_client_capabilities(),
|
||||||
@@ -2711,4 +2724,209 @@ mod tests {
|
|||||||
);
|
);
|
||||||
assert!(should_attempt_oauth_fallback(&Err(err)));
|
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