Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
220 changes: 219 additions & 1 deletion crates/goose/src/agents/extension_manager.rs
Original file line number Diff line number Diff line change
Expand Up @@ -510,13 +510,23 @@ 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,
roots_dir: &std::path::Path,
) -> 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")]
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -611,14 +622,14 @@ 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) => {
return connect_with_auth(
auth_manager,
uri,
timeout_duration,
headers,
provider,
client_name,
capabilities,
Expand Down Expand Up @@ -652,6 +663,7 @@ async fn create_streamable_http_client(
auth_manager,
uri,
timeout_duration,
headers,
provider,
client_name,
capabilities,
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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)
Comment thread
hydrosquall marked this conversation as resolved.
.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"
);
}
}