feat: MCP support for agentic CLI providers (#6972)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -58,6 +58,7 @@ indicatif = "0.18.1"
|
||||
tokio-util = { workspace = true, features = ["compat", "rt"] }
|
||||
anstream = "0.6.18"
|
||||
open = "5.3.2"
|
||||
url = { workspace = true }
|
||||
urlencoding = { workspace = true }
|
||||
clap_complete = "4.5.62"
|
||||
|
||||
@@ -70,4 +71,5 @@ disable-update = []
|
||||
|
||||
[dev-dependencies]
|
||||
tempfile = { workspace = true }
|
||||
test-case = { workspace = true }
|
||||
tokio = { workspace = true }
|
||||
|
||||
@@ -328,7 +328,7 @@ async fn handle_oauth_configuration(provider_name: &str, key_name: &str) -> anyh
|
||||
|
||||
// Create a temporary provider instance to handle OAuth
|
||||
let temp_model = ModelConfig::new("temp")?;
|
||||
match create(provider_name, temp_model).await {
|
||||
match create(provider_name, temp_model, Vec::new()).await {
|
||||
Ok(provider) => match provider.configure_oauth().await {
|
||||
Ok(_) => {
|
||||
let _ = cliclack::log::success("OAuth authentication completed successfully!");
|
||||
@@ -683,7 +683,7 @@ pub async fn configure_provider_dialog() -> anyhow::Result<bool> {
|
||||
spin.start("Attempting to fetch supported models...");
|
||||
let models_res = {
|
||||
let temp_model_config = ModelConfig::new(&provider_meta.default_model)?;
|
||||
let temp_provider = create(provider_name, temp_model_config).await?;
|
||||
let temp_provider = create(provider_name, temp_model_config, Vec::new()).await?;
|
||||
retry_operation(&RetryConfig::default(), || async {
|
||||
temp_provider.fetch_recommended_models().await
|
||||
})
|
||||
@@ -1445,7 +1445,6 @@ pub async fn configure_tool_permissions_dialog() -> anyhow::Result<()> {
|
||||
let model_config = ModelConfig::new(&model)?;
|
||||
|
||||
let agent = Agent::new();
|
||||
let new_provider = create(&provider_name, model_config).await?;
|
||||
|
||||
let session = agent
|
||||
.config
|
||||
@@ -1457,8 +1456,8 @@ pub async fn configure_tool_permissions_dialog() -> anyhow::Result<()> {
|
||||
)
|
||||
.await?;
|
||||
|
||||
agent.update_provider(new_provider, &session.id).await?;
|
||||
if let Some(config) = get_extension_by_name(&selected_extension_name) {
|
||||
let extension_config = get_extension_by_name(&selected_extension_name);
|
||||
if let Some(config) = extension_config.as_ref() {
|
||||
agent
|
||||
.add_extension(config.clone(), &session.id)
|
||||
.await
|
||||
@@ -1478,6 +1477,10 @@ pub async fn configure_tool_permissions_dialog() -> anyhow::Result<()> {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let extensions = extension_config.into_iter().collect::<Vec<_>>();
|
||||
let new_provider = create(&provider_name, model_config, extensions).await?;
|
||||
agent.update_provider(new_provider, &session.id).await?;
|
||||
|
||||
let permission_manager = PermissionManager::instance();
|
||||
let selected_tools = agent
|
||||
.list_tools(&session.id, Some(selected_extension_name.clone()))
|
||||
@@ -1667,7 +1670,7 @@ pub async fn handle_openrouter_auth() -> anyhow::Result<()> {
|
||||
}
|
||||
};
|
||||
|
||||
match create("openrouter", model_config).await {
|
||||
match create("openrouter", model_config, Vec::new()).await {
|
||||
Ok(provider) => {
|
||||
let model_config = provider.get_model_config();
|
||||
let test_result = provider
|
||||
@@ -1747,7 +1750,7 @@ pub async fn handle_tetrate_auth() -> anyhow::Result<()> {
|
||||
}
|
||||
};
|
||||
|
||||
match create("tetrate", model_config).await {
|
||||
match create("tetrate", model_config, Vec::new()).await {
|
||||
Ok(provider) => {
|
||||
let test_result = provider.fetch_supported_models().await;
|
||||
|
||||
|
||||
@@ -181,16 +181,16 @@ async fn create_agent(provider_name: &str, model: &str) -> Result<Agent> {
|
||||
)
|
||||
.await?;
|
||||
|
||||
let provider = goose::providers::create(provider_name, model_config).await?;
|
||||
agent.update_provider(provider, &init_session.id).await?;
|
||||
|
||||
let enabled_configs = goose::config::get_enabled_extensions();
|
||||
for config in enabled_configs {
|
||||
for config in &enabled_configs {
|
||||
if let Err(e) = agent.add_extension(config.clone(), &init_session.id).await {
|
||||
eprintln!("Warning: Failed to load extension {}: {}", config.name(), e);
|
||||
}
|
||||
}
|
||||
|
||||
let provider = goose::providers::create(provider_name, model_config, enabled_configs).await?;
|
||||
agent.update_provider(provider, &init_session.id).await?;
|
||||
|
||||
Ok(agent)
|
||||
}
|
||||
|
||||
|
||||
@@ -187,7 +187,12 @@ where
|
||||
|
||||
let original_env = setup_environment(config)?;
|
||||
|
||||
let inner_provider = create(&factory_name, ModelConfig::new(config.model_name)?).await?;
|
||||
let inner_provider = create(
|
||||
&factory_name,
|
||||
ModelConfig::new(config.model_name)?,
|
||||
Vec::new(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
let test_provider = Arc::new(TestProvider::new_recording(inner_provider, &file_path));
|
||||
(
|
||||
|
||||
@@ -3,14 +3,13 @@ use crate::cli::StreamableHttpOptions;
|
||||
use super::output;
|
||||
use super::CliSession;
|
||||
use console::style;
|
||||
use goose::agents::{Agent, Container};
|
||||
use goose::config::get_enabled_extensions;
|
||||
use goose::agents::{Agent, Container, ExtensionError};
|
||||
use goose::config::resolve_extensions_for_new_session;
|
||||
use goose::config::{get_all_extensions, Config, ExtensionConfig};
|
||||
use goose::providers::create;
|
||||
use goose::recipe::Recipe;
|
||||
use goose::session::session_manager::SessionType;
|
||||
use goose::session::{EnabledExtensionsState, ExtensionState};
|
||||
use goose::session::EnabledExtensionsState;
|
||||
use rustyline::EditMode;
|
||||
use std::collections::BTreeSet;
|
||||
use std::process;
|
||||
@@ -490,27 +489,19 @@ async fn handle_resumed_session_workdir(agent: &Agent, session_id: &str, interac
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_and_load_extensions(
|
||||
agent: Agent,
|
||||
async fn collect_extension_configs(
|
||||
agent: &Agent,
|
||||
session_config: &SessionBuilderConfig,
|
||||
recipe: Option<&Recipe>,
|
||||
session_id: &str,
|
||||
provider_for_debug: Arc<dyn goose::providers::base::Provider>,
|
||||
) -> Arc<Agent> {
|
||||
for warning in goose::config::get_warnings() {
|
||||
eprintln!("{}", style(format!("Warning: {}", warning)).yellow());
|
||||
}
|
||||
|
||||
) -> Result<Vec<ExtensionConfig>, ExtensionError> {
|
||||
let configured_extensions: Vec<ExtensionConfig> = if session_config.resume {
|
||||
agent
|
||||
.config
|
||||
.session_manager
|
||||
.get_session(session_id, false)
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|s| EnabledExtensionsState::from_extension_data(&s.extension_data))
|
||||
.map(|state| state.extensions)
|
||||
.unwrap_or_else(get_enabled_extensions)
|
||||
EnabledExtensionsState::for_session(
|
||||
&agent.config.session_manager,
|
||||
session_id,
|
||||
Config::global(),
|
||||
)
|
||||
.await
|
||||
} else if session_config.no_profile {
|
||||
Vec::new()
|
||||
} else {
|
||||
@@ -523,17 +514,33 @@ async fn resolve_and_load_extensions(
|
||||
&session_config.builtins,
|
||||
);
|
||||
|
||||
let mut extensions_to_load: Vec<(String, ExtensionConfig)> = configured_extensions
|
||||
.iter()
|
||||
.map(|cfg| (cfg.name(), cfg.clone()))
|
||||
let mut all: Vec<ExtensionConfig> = configured_extensions;
|
||||
all.extend(cli_flag_extensions.into_iter().map(|(_, cfg)| cfg));
|
||||
|
||||
Ok(all)
|
||||
}
|
||||
|
||||
async fn resolve_and_load_extensions(
|
||||
agent: Agent,
|
||||
extensions: Vec<ExtensionConfig>,
|
||||
provider_for_debug: Arc<dyn goose::providers::base::Provider>,
|
||||
interactive: bool,
|
||||
session_id: &str,
|
||||
) -> Arc<Agent> {
|
||||
for warning in goose::config::get_warnings() {
|
||||
eprintln!("{}", style(format!("Warning: {}", warning)).yellow());
|
||||
}
|
||||
|
||||
let extensions_to_load: Vec<(String, ExtensionConfig)> = extensions
|
||||
.into_iter()
|
||||
.map(|cfg| (cfg.name(), cfg))
|
||||
.collect();
|
||||
extensions_to_load.extend(cli_flag_extensions);
|
||||
|
||||
load_extensions(
|
||||
agent,
|
||||
extensions_to_load,
|
||||
provider_for_debug,
|
||||
session_config.interactive,
|
||||
interactive,
|
||||
session_id,
|
||||
)
|
||||
.await
|
||||
@@ -598,7 +605,28 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession {
|
||||
.apply_recipe_components(recipe.and_then(|r| r.response.clone()), true)
|
||||
.await;
|
||||
|
||||
let new_provider = match create(&resolved.provider_name, resolved.model_config).await {
|
||||
let session_id = resolve_session_id(&session_config, &session_manager).await;
|
||||
|
||||
if session_config.resume {
|
||||
handle_resumed_session_workdir(&agent, &session_id, session_config.interactive).await;
|
||||
}
|
||||
|
||||
let extensions_for_provider =
|
||||
match collect_extension_configs(&agent, &session_config, recipe, &session_id).await {
|
||||
Ok(exts) => exts,
|
||||
Err(e) => {
|
||||
output::render_error(&format!("Failed to collect extensions: {}", e));
|
||||
process::exit(1);
|
||||
}
|
||||
};
|
||||
|
||||
let new_provider = match create(
|
||||
&resolved.provider_name,
|
||||
resolved.model_config,
|
||||
extensions_for_provider.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(provider) => provider,
|
||||
Err(e) => {
|
||||
output::render_error(&format!(
|
||||
@@ -624,8 +652,6 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession {
|
||||
tracing::info!("🤖 Using model: {}", resolved.model_name);
|
||||
}
|
||||
|
||||
let session_id = resolve_session_id(&session_config, &session_manager).await;
|
||||
|
||||
agent
|
||||
.update_provider(new_provider, &session_id)
|
||||
.await
|
||||
@@ -645,17 +671,13 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession {
|
||||
}
|
||||
}
|
||||
|
||||
if session_config.resume {
|
||||
handle_resumed_session_workdir(&agent, &session_id, session_config.interactive).await;
|
||||
}
|
||||
|
||||
// Extensions are loaded after session creation because we may change directory when resuming
|
||||
let agent_ptr = resolve_and_load_extensions(
|
||||
agent,
|
||||
&session_config,
|
||||
recipe,
|
||||
&session_id,
|
||||
extensions_for_provider,
|
||||
Arc::clone(&provider_for_display),
|
||||
session_config.interactive,
|
||||
&session_id,
|
||||
)
|
||||
.await;
|
||||
|
||||
|
||||
@@ -286,9 +286,14 @@ impl CliSession {
|
||||
}
|
||||
|
||||
let cmd = parts.remove(0).to_string();
|
||||
let name = std::path::Path::new(&cmd)
|
||||
.file_name()
|
||||
.and_then(|f| f.to_str())
|
||||
.unwrap_or("unnamed")
|
||||
.to_string();
|
||||
|
||||
Ok(ExtensionConfig::Stdio {
|
||||
name: String::new(),
|
||||
name,
|
||||
cmd,
|
||||
args: parts.iter().map(|s| s.to_string()).collect(),
|
||||
envs: Envs::new(envs),
|
||||
@@ -301,8 +306,29 @@ impl CliSession {
|
||||
}
|
||||
|
||||
pub fn parse_streamable_http_extension(extension_url: &str, timeout: u64) -> ExtensionConfig {
|
||||
let name = url::Url::parse(extension_url)
|
||||
.ok()
|
||||
.map(|u| {
|
||||
let mut s = String::new();
|
||||
if let Some(host) = u.host_str() {
|
||||
s.push_str(host);
|
||||
}
|
||||
if let Some(port) = u.port() {
|
||||
s.push('_');
|
||||
s.push_str(&port.to_string());
|
||||
}
|
||||
let path = u.path().trim_matches('/');
|
||||
if !path.is_empty() {
|
||||
s.push('_');
|
||||
s.push_str(path);
|
||||
}
|
||||
s
|
||||
})
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| "unnamed".to_string());
|
||||
|
||||
ExtensionConfig::StreamableHttp {
|
||||
name: String::new(),
|
||||
name,
|
||||
uri: extension_url.to_string(),
|
||||
envs: Envs::new(HashMap::new()),
|
||||
env_keys: Vec::new(),
|
||||
@@ -1841,7 +1867,8 @@ async fn get_reasoner() -> Result<Arc<dyn Provider>, anyhow::Error> {
|
||||
|
||||
let model_config =
|
||||
ModelConfig::new_with_context_env(model, Some("GOOSE_PLANNER_CONTEXT_LIMIT"))?;
|
||||
let reasoner = create(&provider, model_config).await?;
|
||||
let extensions = goose::config::extensions::get_enabled_extensions_with_config(config);
|
||||
let reasoner = create(&provider, model_config, extensions).await?;
|
||||
|
||||
Ok(reasoner)
|
||||
}
|
||||
@@ -1862,7 +1889,10 @@ fn format_elapsed_time(duration: std::time::Duration) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use goose::agents::extension::Envs;
|
||||
use goose::config::ExtensionConfig;
|
||||
use std::time::Duration;
|
||||
use test_case::test_case;
|
||||
|
||||
#[test]
|
||||
fn test_format_elapsed_time_under_60_seconds() {
|
||||
@@ -1929,4 +1959,95 @@ mod tests {
|
||||
let duration = Duration::from_millis(60500);
|
||||
assert_eq!(format_elapsed_time(duration), "1m 00s");
|
||||
}
|
||||
|
||||
#[test_case(
|
||||
"/usr/bin/my-server",
|
||||
ExtensionConfig::Stdio {
|
||||
name: "my-server".into(),
|
||||
cmd: "/usr/bin/my-server".into(),
|
||||
args: vec![],
|
||||
envs: Envs::default(),
|
||||
env_keys: vec![],
|
||||
description: goose::config::DEFAULT_EXTENSION_DESCRIPTION.to_string(),
|
||||
timeout: Some(goose::config::DEFAULT_EXTENSION_TIMEOUT),
|
||||
bundled: None,
|
||||
available_tools: vec![],
|
||||
}
|
||||
; "name_from_cmd_basename"
|
||||
)]
|
||||
#[test_case(
|
||||
"MY_SECRET=s3cret npx -y @modelcontextprotocol/server-everything",
|
||||
ExtensionConfig::Stdio {
|
||||
name: "npx".into(),
|
||||
cmd: "npx".into(),
|
||||
args: vec!["-y".into(), "@modelcontextprotocol/server-everything".into()],
|
||||
envs: Envs::new([("MY_SECRET".into(), "s3cret".into())].into()),
|
||||
env_keys: vec![],
|
||||
description: goose::config::DEFAULT_EXTENSION_DESCRIPTION.to_string(),
|
||||
timeout: Some(goose::config::DEFAULT_EXTENSION_TIMEOUT),
|
||||
bundled: None,
|
||||
available_tools: vec![],
|
||||
}
|
||||
; "env_prefix_name_from_cmd"
|
||||
)]
|
||||
fn test_parse_stdio_extension(input: &str, expected: ExtensionConfig) {
|
||||
assert_eq!(CliSession::parse_stdio_extension(input).unwrap(), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_stdio_extension_no_command() {
|
||||
assert!(CliSession::parse_stdio_extension("").is_err());
|
||||
}
|
||||
|
||||
#[test_case(
|
||||
"https://mcp.kiwi.com", 300,
|
||||
ExtensionConfig::StreamableHttp {
|
||||
name: "mcp.kiwi.com".into(),
|
||||
uri: "https://mcp.kiwi.com".into(),
|
||||
envs: Envs::default(),
|
||||
env_keys: vec![],
|
||||
headers: HashMap::new(),
|
||||
description: goose::config::DEFAULT_EXTENSION_DESCRIPTION.to_string(),
|
||||
timeout: Some(300),
|
||||
bundled: None,
|
||||
available_tools: vec![],
|
||||
}
|
||||
; "name_from_host"
|
||||
)]
|
||||
#[test_case(
|
||||
"http://localhost:8080/api", 300,
|
||||
ExtensionConfig::StreamableHttp {
|
||||
name: "localhost_8080_api".into(),
|
||||
uri: "http://localhost:8080/api".into(),
|
||||
envs: Envs::default(),
|
||||
env_keys: vec![],
|
||||
headers: HashMap::new(),
|
||||
description: goose::config::DEFAULT_EXTENSION_DESCRIPTION.to_string(),
|
||||
timeout: Some(300),
|
||||
bundled: None,
|
||||
available_tools: vec![],
|
||||
}
|
||||
; "port_and_path"
|
||||
)]
|
||||
#[test_case(
|
||||
"http://localhost:9090/other", 300,
|
||||
ExtensionConfig::StreamableHttp {
|
||||
name: "localhost_9090_other".into(),
|
||||
uri: "http://localhost:9090/other".into(),
|
||||
envs: Envs::default(),
|
||||
env_keys: vec![],
|
||||
headers: HashMap::new(),
|
||||
description: goose::config::DEFAULT_EXTENSION_DESCRIPTION.to_string(),
|
||||
timeout: Some(300),
|
||||
bundled: None,
|
||||
available_tools: vec![],
|
||||
}
|
||||
; "different_port_and_path"
|
||||
)]
|
||||
fn test_parse_streamable_http_extension(url: &str, timeout: u64, expected: ExtensionConfig) {
|
||||
assert_eq!(
|
||||
CliSession::parse_streamable_http_extension(url, timeout),
|
||||
expected
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user