2229 lines
76 KiB
Rust
2229 lines
76 KiB
Rust
use anyhow::Result;
|
|
use axum::http::{HeaderMap, HeaderName};
|
|
use chrono::{DateTime, Utc};
|
|
use futures::stream::{FuturesUnordered, StreamExt};
|
|
use futures::{future, FutureExt};
|
|
use once_cell::sync::Lazy;
|
|
use rmcp::service::{ClientInitializeError, ServiceError};
|
|
use rmcp::transport::streamable_http_client::{
|
|
AuthRequiredError, StreamableHttpClientTransportConfig, StreamableHttpError,
|
|
};
|
|
use rmcp::transport::{
|
|
ConfigureCommandExt, DynamicTransportError, StreamableHttpClientTransport, TokioChildProcess,
|
|
};
|
|
use std::collections::HashMap;
|
|
use std::path::PathBuf;
|
|
use std::process::Stdio;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
use tempfile::{tempdir, TempDir};
|
|
use tokio::io::AsyncReadExt;
|
|
use tokio::process::Command;
|
|
use tokio::sync::Mutex;
|
|
use tokio_stream::wrappers::ReceiverStream;
|
|
use tokio_util::sync::CancellationToken;
|
|
use tracing::{error, warn};
|
|
|
|
use super::container::Container;
|
|
use super::extension::{
|
|
ExtensionConfig, ExtensionError, ExtensionInfo, ExtensionResult, PlatformExtensionContext,
|
|
ToolInfo, PLATFORM_EXTENSIONS,
|
|
};
|
|
use super::tool_execution::ToolCallResult;
|
|
use super::types::SharedProvider;
|
|
use crate::agents::extension::{Envs, ProcessExit};
|
|
use crate::agents::extension_malware_check;
|
|
use crate::agents::mcp_client::{GooseMcpClientCapabilities, McpClient, McpClientTrait};
|
|
use crate::builtin_extension::get_builtin_extension;
|
|
use crate::config::extensions::name_to_key;
|
|
use crate::config::search_path::SearchPaths;
|
|
use crate::config::{get_all_extensions, Config};
|
|
use crate::oauth::oauth_flow;
|
|
use crate::prompt_template;
|
|
use crate::subprocess::configure_subprocess;
|
|
use rmcp::model::{
|
|
CallToolRequestParams, Content, ErrorCode, ErrorData, GetPromptResult, Prompt, Resource,
|
|
ResourceContents, ServerInfo, Tool,
|
|
};
|
|
use rmcp::transport::auth::AuthClient;
|
|
use schemars::_private::NoSerialize;
|
|
use serde_json::Value;
|
|
|
|
type McpClientBox = Arc<dyn McpClientTrait>;
|
|
|
|
static RE_ENV_BRACES: Lazy<regex::Regex> =
|
|
Lazy::new(|| regex::Regex::new(r"\$\{\s*([A-Za-z_][A-Za-z0-9_]*)\s*\}").expect("valid regex"));
|
|
|
|
static RE_ENV_SIMPLE: Lazy<regex::Regex> =
|
|
Lazy::new(|| regex::Regex::new(r"\$([A-Za-z_][A-Za-z0-9_]*)").expect("valid regex"));
|
|
|
|
struct Extension {
|
|
pub config: ExtensionConfig,
|
|
|
|
client: McpClientBox,
|
|
server_info: Option<ServerInfo>,
|
|
_temp_dir: Option<tempfile::TempDir>,
|
|
}
|
|
|
|
impl Extension {
|
|
fn new(
|
|
config: ExtensionConfig,
|
|
client: McpClientBox,
|
|
server_info: Option<ServerInfo>,
|
|
temp_dir: Option<tempfile::TempDir>,
|
|
) -> Self {
|
|
Self {
|
|
client,
|
|
config,
|
|
server_info,
|
|
_temp_dir: temp_dir,
|
|
}
|
|
}
|
|
|
|
fn supports_resources(&self) -> bool {
|
|
self.server_info
|
|
.as_ref()
|
|
.and_then(|info| info.capabilities.resources.as_ref())
|
|
.is_some()
|
|
}
|
|
|
|
fn get_instructions(&self) -> Option<String> {
|
|
self.server_info
|
|
.as_ref()
|
|
.and_then(|info| info.instructions.clone())
|
|
}
|
|
|
|
fn get_client(&self) -> McpClientBox {
|
|
self.client.clone()
|
|
}
|
|
}
|
|
|
|
pub struct ExtensionManagerCapabilities {
|
|
pub mcpui: bool,
|
|
}
|
|
|
|
/// Manages goose extensions / MCP clients and their interactions
|
|
pub struct ExtensionManager {
|
|
extensions: Mutex<HashMap<String, Extension>>,
|
|
context: PlatformExtensionContext,
|
|
provider: SharedProvider,
|
|
tools_cache: Mutex<Option<Arc<Vec<Tool>>>>,
|
|
tools_cache_version: AtomicU64,
|
|
client_name: String,
|
|
capabilities: ExtensionManagerCapabilities,
|
|
}
|
|
|
|
/// A flattened representation of a resource used by the agent to prepare inference
|
|
#[derive(Debug, Clone)]
|
|
pub struct ResourceItem {
|
|
pub extension_name: String, // The name of the extension that owns the resource
|
|
pub uri: String, // The URI of the resource
|
|
pub name: String, // The name of the resource
|
|
pub content: String, // The content of the resource
|
|
pub timestamp: DateTime<Utc>, // The timestamp of the resource
|
|
pub priority: f32, // The priority of the resource
|
|
pub token_count: Option<u32>, // The token count of the resource (filled in by the agent)
|
|
}
|
|
|
|
impl ResourceItem {
|
|
pub fn new(
|
|
extension_name: String,
|
|
uri: String,
|
|
name: String,
|
|
content: String,
|
|
timestamp: DateTime<Utc>,
|
|
priority: f32,
|
|
) -> Self {
|
|
Self {
|
|
extension_name,
|
|
uri,
|
|
name,
|
|
content,
|
|
timestamp,
|
|
priority,
|
|
token_count: None,
|
|
}
|
|
}
|
|
}
|
|
|
|
fn resolve_command(cmd: &str) -> PathBuf {
|
|
SearchPaths::builder()
|
|
.with_npm()
|
|
.resolve(cmd)
|
|
.unwrap_or_else(|_| {
|
|
// let the OS raise the error
|
|
PathBuf::from(cmd)
|
|
})
|
|
}
|
|
|
|
fn require_str_parameter<'a>(v: &'a serde_json::Value, name: &str) -> Result<&'a str, ErrorData> {
|
|
let v = v.get(name).ok_or_else(|| {
|
|
ErrorData::new(
|
|
ErrorCode::INVALID_PARAMS,
|
|
format!("The parameter {name} is required"),
|
|
None,
|
|
)
|
|
})?;
|
|
match v.as_str() {
|
|
Some(r) => Ok(r),
|
|
None => Err(ErrorData::new(
|
|
ErrorCode::INVALID_PARAMS,
|
|
format!("The parameter {name} must be a string"),
|
|
None,
|
|
)),
|
|
}
|
|
}
|
|
|
|
pub fn get_parameter_names(tool: &Tool) -> Vec<String> {
|
|
let mut names: Vec<String> = tool
|
|
.input_schema
|
|
.get("properties")
|
|
.and_then(|props| props.as_object())
|
|
.map(|props| props.keys().cloned().collect())
|
|
.unwrap_or_default();
|
|
names.sort();
|
|
names
|
|
}
|
|
|
|
const TOOL_EXTENSION_META_KEY: &str = "goose_extension";
|
|
|
|
pub fn get_tool_owner(tool: &Tool) -> Option<String> {
|
|
tool.meta
|
|
.as_ref()
|
|
.and_then(|m| m.0.get(TOOL_EXTENSION_META_KEY))
|
|
.and_then(|v| v.as_str())
|
|
.map(|s| s.to_string())
|
|
}
|
|
|
|
fn is_unprefixed_extension(config: &ExtensionConfig) -> bool {
|
|
if let ExtensionConfig::Platform { name, .. } = config {
|
|
PLATFORM_EXTENSIONS
|
|
.get(name_to_key(name).as_str())
|
|
.is_some_and(|def| def.unprefixed_tools)
|
|
} else {
|
|
false
|
|
}
|
|
}
|
|
|
|
/// Returns true if the named extension is a first-class platform extension
|
|
/// whose tools are exposed unprefixed and remain visible during code execution mode.
|
|
pub fn is_first_class_extension(name: &str) -> bool {
|
|
PLATFORM_EXTENSIONS
|
|
.get(name_to_key(name).as_str())
|
|
.is_some_and(|def| def.unprefixed_tools)
|
|
}
|
|
|
|
/// Result of resolving a tool call to its owning extension
|
|
struct ResolvedTool {
|
|
extension_name: String,
|
|
actual_tool_name: String,
|
|
client: McpClientBox,
|
|
}
|
|
|
|
async fn child_process_client(
|
|
mut command: Command,
|
|
timeout: &Option<u64>,
|
|
provider: SharedProvider,
|
|
working_dir: Option<&PathBuf>,
|
|
docker_container: Option<String>,
|
|
client_name: String,
|
|
capabilities: GooseMcpClientCapabilities,
|
|
) -> ExtensionResult<McpClient> {
|
|
configure_subprocess(&mut command);
|
|
|
|
if let Ok(path) = SearchPaths::builder().path() {
|
|
command.env("PATH", path);
|
|
}
|
|
|
|
// Use explicitly passed working_dir, falling back to GOOSE_WORKING_DIR env var
|
|
let effective_working_dir = working_dir
|
|
.map(|p| p.to_path_buf())
|
|
.or_else(|| std::env::var("GOOSE_WORKING_DIR").ok().map(PathBuf::from));
|
|
|
|
if let Some(ref dir) = effective_working_dir {
|
|
if dir.exists() && dir.is_dir() {
|
|
tracing::info!("Setting MCP process working directory: {:?}", dir);
|
|
command.current_dir(dir);
|
|
} else {
|
|
tracing::warn!(
|
|
"Working directory doesn't exist or isn't a directory: {:?}",
|
|
dir
|
|
);
|
|
}
|
|
}
|
|
|
|
let (transport, mut stderr) = TokioChildProcess::builder(command)
|
|
.stderr(Stdio::piped())
|
|
.spawn()?;
|
|
let mut stderr = stderr.take().ok_or_else(|| {
|
|
ExtensionError::SetupError("failed to attach child process stderr".to_owned())
|
|
})?;
|
|
|
|
let stderr_task = tokio::spawn(async move {
|
|
let mut all_stderr = Vec::new();
|
|
stderr.read_to_end(&mut all_stderr).await?;
|
|
Ok::<String, std::io::Error>(String::from_utf8_lossy(&all_stderr).into())
|
|
});
|
|
|
|
let client_result = McpClient::connect_with_container(
|
|
transport,
|
|
Duration::from_secs(timeout.unwrap_or(crate::config::DEFAULT_EXTENSION_TIMEOUT)),
|
|
provider,
|
|
docker_container,
|
|
client_name,
|
|
capabilities,
|
|
)
|
|
.await;
|
|
|
|
match client_result {
|
|
Ok(client) => Ok(client),
|
|
Err(error) => {
|
|
let error_task_out = stderr_task.await?;
|
|
Err::<McpClient, ExtensionError>(match error_task_out {
|
|
Ok(stderr_content) => ProcessExit::new(stderr_content, error).into(),
|
|
Err(e) => e.into(),
|
|
})
|
|
}
|
|
}
|
|
}
|
|
|
|
fn extract_auth_error(
|
|
res: &Result<McpClient, ClientInitializeError>,
|
|
) -> Option<&AuthRequiredError> {
|
|
match res {
|
|
Ok(_) => None,
|
|
Err(err) => match err {
|
|
ClientInitializeError::TransportError {
|
|
error: DynamicTransportError { error, .. },
|
|
..
|
|
} => error
|
|
.downcast_ref::<StreamableHttpError<reqwest::Error>>()
|
|
.and_then(|auth_error| match auth_error {
|
|
StreamableHttpError::AuthRequired(auth_required_error) => {
|
|
Some(auth_required_error)
|
|
}
|
|
_ => None,
|
|
}),
|
|
_ => None,
|
|
},
|
|
}
|
|
}
|
|
|
|
/// Merge environment variables from direct envs and keychain-stored env_keys
|
|
pub(crate) async fn merge_environments(
|
|
envs: &Envs,
|
|
env_keys: &[String],
|
|
ext_name: &str,
|
|
config: &Config,
|
|
) -> Result<HashMap<String, String>, ExtensionError> {
|
|
let mut all_envs = envs.get_env();
|
|
|
|
for key in env_keys {
|
|
if all_envs.contains_key(key) {
|
|
continue;
|
|
}
|
|
|
|
match config.get(key, true) {
|
|
Ok(value) => {
|
|
if value.is_null() {
|
|
warn!(
|
|
key = %key,
|
|
ext_name = %ext_name,
|
|
"Secret key not found in config (returned null)."
|
|
);
|
|
continue;
|
|
}
|
|
|
|
if let Some(str_val) = value.as_str() {
|
|
all_envs.insert(key.clone(), str_val.to_string());
|
|
} else {
|
|
warn!(
|
|
key = %key,
|
|
ext_name = %ext_name,
|
|
value_type = %value.get("type").and_then(|t| t.as_str()).unwrap_or("unknown"),
|
|
"Secret value is not a string; skipping."
|
|
);
|
|
}
|
|
}
|
|
Err(e) => {
|
|
error!(
|
|
key = %key,
|
|
ext_name = %ext_name,
|
|
error = %e,
|
|
"Failed to fetch secret from config."
|
|
);
|
|
return Err(ExtensionError::ConfigError(format!(
|
|
"Failed to fetch secret '{}' from config: {}",
|
|
key, e
|
|
)));
|
|
}
|
|
}
|
|
}
|
|
|
|
Ok(all_envs)
|
|
}
|
|
|
|
/// Substitute environment variables in a string. Supports both ${VAR} and $VAR syntax.
|
|
pub(crate) fn substitute_env_vars(value: &str, env_map: &HashMap<String, String>) -> String {
|
|
let mut result = value.to_string();
|
|
|
|
for cap in RE_ENV_BRACES.captures_iter(value) {
|
|
if let Some(var_name) = cap.get(1) {
|
|
if let Some(env_value) = env_map.get(var_name.as_str()) {
|
|
result = result.replace(&cap[0], env_value);
|
|
}
|
|
}
|
|
}
|
|
|
|
let snapshot = result.clone();
|
|
for cap in RE_ENV_SIMPLE.captures_iter(&snapshot) {
|
|
if let Some(var_name) = cap.get(1) {
|
|
if !value.contains(&format!("${{{}}}", var_name.as_str())) {
|
|
if let Some(env_value) = env_map.get(var_name.as_str()) {
|
|
result = result.replace(&cap[0], env_value);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
result
|
|
}
|
|
|
|
const GOOSE_USER_AGENT: reqwest::header::HeaderValue =
|
|
reqwest::header::HeaderValue::from_static(concat!("goose/", env!("CARGO_PKG_VERSION")));
|
|
|
|
async fn create_streamable_http_client(
|
|
uri: &str,
|
|
timeout: Option<u64>,
|
|
headers: &HashMap<String, String>,
|
|
name: &str,
|
|
provider: SharedProvider,
|
|
client_name: String,
|
|
capabilities: GooseMcpClientCapabilities,
|
|
) -> ExtensionResult<Box<dyn McpClientTrait>> {
|
|
let mut default_headers = HeaderMap::new();
|
|
|
|
default_headers.insert(reqwest::header::USER_AGENT, GOOSE_USER_AGENT);
|
|
|
|
for (key, value) in headers {
|
|
default_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))
|
|
})?,
|
|
);
|
|
}
|
|
|
|
let http_client = reqwest::Client::builder()
|
|
.default_headers(default_headers)
|
|
.build()
|
|
.map_err(|_| ExtensionError::ConfigError("could not construct http client".to_string()))?;
|
|
|
|
let transport = StreamableHttpClientTransport::with_client(
|
|
http_client,
|
|
StreamableHttpClientTransportConfig {
|
|
uri: uri.into(),
|
|
..Default::default()
|
|
},
|
|
);
|
|
|
|
let timeout_duration =
|
|
Duration::from_secs(timeout.unwrap_or(crate::config::DEFAULT_EXTENSION_TIMEOUT));
|
|
|
|
let client_res = McpClient::connect(
|
|
transport,
|
|
timeout_duration,
|
|
provider.clone(),
|
|
client_name.clone(),
|
|
capabilities.clone(),
|
|
)
|
|
.await;
|
|
|
|
if extract_auth_error(&client_res).is_some() {
|
|
let auth_manager = oauth_flow(&uri.to_string(), &name.to_string())
|
|
.await
|
|
.map_err(|_| ExtensionError::SetupError("auth error".to_string()))?;
|
|
let mut auth_headers = HeaderMap::new();
|
|
auth_headers.insert(reqwest::header::USER_AGENT, GOOSE_USER_AGENT);
|
|
let auth_http_client = reqwest::Client::builder()
|
|
.default_headers(auth_headers)
|
|
.build()
|
|
.map_err(|_| {
|
|
ExtensionError::ConfigError("could not construct http client".to_string())
|
|
})?;
|
|
let auth_client = AuthClient::new(auth_http_client, auth_manager);
|
|
let transport = StreamableHttpClientTransport::with_client(
|
|
auth_client,
|
|
StreamableHttpClientTransportConfig {
|
|
uri: uri.into(),
|
|
..Default::default()
|
|
},
|
|
);
|
|
Ok(Box::new(
|
|
McpClient::connect(
|
|
transport,
|
|
timeout_duration,
|
|
provider,
|
|
client_name,
|
|
capabilities,
|
|
)
|
|
.await?,
|
|
))
|
|
} else {
|
|
Ok(Box::new(client_res?))
|
|
}
|
|
}
|
|
|
|
impl ExtensionManager {
|
|
pub fn new(
|
|
provider: SharedProvider,
|
|
session_manager: Arc<crate::session::SessionManager>,
|
|
client_name: String,
|
|
capabilities: ExtensionManagerCapabilities,
|
|
) -> Self {
|
|
Self {
|
|
extensions: Mutex::new(HashMap::new()),
|
|
context: PlatformExtensionContext {
|
|
extension_manager: None,
|
|
session_manager,
|
|
session: None,
|
|
},
|
|
provider,
|
|
tools_cache: Mutex::new(None),
|
|
tools_cache_version: AtomicU64::new(0),
|
|
client_name,
|
|
capabilities,
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
pub fn new_without_provider(data_dir: std::path::PathBuf) -> Self {
|
|
let session_manager = Arc::new(crate::session::SessionManager::new(data_dir));
|
|
Self::new(
|
|
Arc::new(Mutex::new(None)),
|
|
session_manager,
|
|
"goose-cli".to_string(),
|
|
ExtensionManagerCapabilities { mcpui: false },
|
|
)
|
|
}
|
|
|
|
pub fn get_context(&self) -> &PlatformExtensionContext {
|
|
&self.context
|
|
}
|
|
|
|
pub fn get_provider(&self) -> &SharedProvider {
|
|
&self.provider
|
|
}
|
|
|
|
pub async fn supports_resources(&self) -> bool {
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.values()
|
|
.any(|ext| ext.supports_resources())
|
|
}
|
|
|
|
/// Add an extension with an optional working directory.
|
|
/// If working_dir is None, falls back to current_dir.
|
|
#[allow(clippy::too_many_lines)]
|
|
pub async fn add_extension(
|
|
self: &Arc<Self>,
|
|
config: ExtensionConfig,
|
|
working_dir: Option<PathBuf>,
|
|
container: Option<&Container>,
|
|
session_id: Option<&str>,
|
|
) -> ExtensionResult<()> {
|
|
let sanitized_name = config.key();
|
|
|
|
if self.extensions.lock().await.contains_key(&sanitized_name) {
|
|
return Ok(());
|
|
}
|
|
|
|
let mut temp_dir = None;
|
|
|
|
let client: Box<dyn McpClientTrait> = match &config {
|
|
ExtensionConfig::Sse { .. } => {
|
|
return Err(ExtensionError::ConfigError(
|
|
"SSE is unsupported, migrate to streamable_http".to_string(),
|
|
));
|
|
}
|
|
ExtensionConfig::StreamableHttp {
|
|
uri,
|
|
timeout,
|
|
headers,
|
|
name,
|
|
envs,
|
|
env_keys,
|
|
..
|
|
} => {
|
|
let config = Config::global();
|
|
let all_envs = merge_environments(envs, env_keys, &sanitized_name, config).await?;
|
|
let resolved_headers = headers
|
|
.iter()
|
|
.map(|(k, v)| (k.clone(), substitute_env_vars(v, &all_envs)))
|
|
.collect();
|
|
let capability = GooseMcpClientCapabilities {
|
|
mcpui: self.capabilities.mcpui,
|
|
};
|
|
|
|
create_streamable_http_client(
|
|
uri,
|
|
*timeout,
|
|
&resolved_headers,
|
|
name,
|
|
self.provider.clone(),
|
|
self.client_name.clone(),
|
|
capability,
|
|
)
|
|
.await?
|
|
}
|
|
ExtensionConfig::Builtin { name, timeout, .. } => {
|
|
let timeout_duration = Duration::from_secs(timeout.unwrap_or(300));
|
|
let normalized_name = name_to_key(name);
|
|
let extension_fn =
|
|
get_builtin_extension(normalized_name.as_str()).ok_or_else(|| {
|
|
ExtensionError::ConfigError(format!("Unknown builtin extension: {}", name))
|
|
})?;
|
|
|
|
if let Some(container) = container {
|
|
let container_id = container.id();
|
|
tracing::info!(
|
|
container = %container_id,
|
|
builtin = %name,
|
|
"Starting builtin extension inside Docker container"
|
|
);
|
|
let normalized_name = name_to_key(name);
|
|
let command = Command::new("docker").configure(|command| {
|
|
command
|
|
.arg("exec")
|
|
.arg("-i")
|
|
.arg(container_id)
|
|
.arg("goose")
|
|
.arg("mcp")
|
|
.arg(&normalized_name);
|
|
});
|
|
|
|
let effective_working_dir = working_dir
|
|
.clone()
|
|
.unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
|
|
|
|
let capabilities = GooseMcpClientCapabilities {
|
|
mcpui: self.capabilities.mcpui,
|
|
};
|
|
|
|
let client = child_process_client(
|
|
command,
|
|
timeout,
|
|
self.provider.clone(),
|
|
Some(&effective_working_dir),
|
|
Some(container_id.to_string()),
|
|
self.client_name.clone(),
|
|
capabilities,
|
|
)
|
|
.await?;
|
|
Box::new(client)
|
|
} else {
|
|
// Non-containerized builtin runs in-process via duplex channels.
|
|
// Working directory is passed per-request via call_tool metadata, not here.
|
|
let (server_read, client_write) = tokio::io::duplex(65536);
|
|
let (client_read, server_write) = tokio::io::duplex(65536);
|
|
extension_fn(server_read, server_write);
|
|
|
|
let capabilities = GooseMcpClientCapabilities {
|
|
mcpui: self.capabilities.mcpui,
|
|
};
|
|
|
|
Box::new(
|
|
McpClient::connect(
|
|
(client_read, client_write),
|
|
timeout_duration,
|
|
self.provider.clone(),
|
|
self.client_name.clone(),
|
|
capabilities,
|
|
)
|
|
.await?,
|
|
)
|
|
}
|
|
}
|
|
ExtensionConfig::Stdio {
|
|
cmd,
|
|
args,
|
|
envs,
|
|
env_keys,
|
|
timeout,
|
|
..
|
|
} => {
|
|
let config = Config::global();
|
|
let mut all_envs =
|
|
merge_environments(envs, env_keys, &sanitized_name, config).await?;
|
|
|
|
if let Some(sid) = session_id {
|
|
all_envs.insert("AGENT_SESSION_ID".to_string(), sid.to_string());
|
|
}
|
|
|
|
// Check for malicious packages before launching the process
|
|
extension_malware_check::deny_if_malicious_cmd_args(cmd, args).await?;
|
|
|
|
let command = if let Some(container) = container {
|
|
let container_id = container.id();
|
|
tracing::info!(
|
|
container = %container_id,
|
|
cmd = %cmd,
|
|
"Starting stdio extension inside Docker container"
|
|
);
|
|
Command::new("docker").configure(|command| {
|
|
command.arg("exec").arg("-i");
|
|
for (key, value) in &all_envs {
|
|
command.arg("-e").arg(format!("{}={}", key, value));
|
|
}
|
|
command.arg(container_id);
|
|
command.arg(cmd);
|
|
command.args(args);
|
|
})
|
|
} else {
|
|
let cmd = resolve_command(cmd);
|
|
Command::new(cmd).configure(|command| {
|
|
command.args(args).envs(all_envs);
|
|
})
|
|
};
|
|
|
|
let effective_working_dir = working_dir
|
|
.clone()
|
|
.unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
|
|
let capabilities = GooseMcpClientCapabilities {
|
|
mcpui: self.capabilities.mcpui,
|
|
};
|
|
let client = child_process_client(
|
|
command,
|
|
timeout,
|
|
self.provider.clone(),
|
|
Some(&effective_working_dir),
|
|
container.map(|c| c.id().to_string()),
|
|
self.client_name.clone(),
|
|
capabilities,
|
|
)
|
|
.await?;
|
|
Box::new(client)
|
|
}
|
|
ExtensionConfig::Platform { name, .. } => {
|
|
let normalized_key = name_to_key(name);
|
|
let def = PLATFORM_EXTENSIONS
|
|
.get(normalized_key.as_str())
|
|
.ok_or_else(|| {
|
|
ExtensionError::ConfigError(format!("Unknown platform extension: {}", name))
|
|
})?;
|
|
let mut context = self.context.clone();
|
|
context.extension_manager = Some(Arc::downgrade(self));
|
|
if let Some(id) = session_id {
|
|
if let Ok(session) = self.context.session_manager.get_session(id, false).await {
|
|
context.session = Some(Arc::new(session));
|
|
}
|
|
}
|
|
|
|
(def.client_factory)(context)
|
|
}
|
|
ExtensionConfig::InlinePython {
|
|
name,
|
|
code,
|
|
timeout,
|
|
dependencies,
|
|
..
|
|
} => {
|
|
let dir = tempdir()?;
|
|
let file_path = dir.path().join(format!("{}.py", name));
|
|
temp_dir = Some(dir);
|
|
std::fs::write(&file_path, code)?;
|
|
|
|
let command = Command::new("uvx").configure(|command| {
|
|
command.arg("--with").arg("mcp");
|
|
dependencies.iter().flatten().for_each(|dep| {
|
|
command.arg("--with").arg(dep);
|
|
});
|
|
command.arg("python").arg(file_path.to_str().unwrap());
|
|
});
|
|
|
|
// Compute working_dir for InlinePython (runs as child process via uvx)
|
|
let effective_working_dir = working_dir
|
|
.clone()
|
|
.unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
|
|
|
|
let capabilities = GooseMcpClientCapabilities {
|
|
mcpui: self.capabilities.mcpui,
|
|
};
|
|
|
|
let client = child_process_client(
|
|
command,
|
|
timeout,
|
|
self.provider.clone(),
|
|
Some(&effective_working_dir),
|
|
container.map(|c| c.id().to_string()),
|
|
self.client_name.clone(),
|
|
capabilities,
|
|
)
|
|
.await?;
|
|
|
|
Box::new(client)
|
|
}
|
|
ExtensionConfig::Frontend { .. } => {
|
|
return Err(ExtensionError::ConfigError(
|
|
"Invalid extension type: Frontend extensions cannot be added as server extensions".to_string()
|
|
));
|
|
}
|
|
};
|
|
|
|
let server_info = client.get_info().cloned();
|
|
|
|
let mut extensions = self.extensions.lock().await;
|
|
extensions.insert(
|
|
sanitized_name,
|
|
Extension::new(config, Arc::from(client), server_info, temp_dir),
|
|
);
|
|
drop(extensions);
|
|
self.invalidate_tools_cache_and_bump_version().await;
|
|
|
|
Ok(())
|
|
}
|
|
|
|
pub async fn add_client(
|
|
&self,
|
|
name: String,
|
|
config: ExtensionConfig,
|
|
client: McpClientBox,
|
|
info: Option<ServerInfo>,
|
|
temp_dir: Option<TempDir>,
|
|
) {
|
|
let normalized = name_to_key(&name);
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.insert(normalized, Extension::new(config, client, info, temp_dir));
|
|
self.invalidate_tools_cache_and_bump_version().await;
|
|
}
|
|
|
|
/// Get extensions info for building the system prompt
|
|
pub async fn get_extensions_info(&self, working_dir: &std::path::Path) -> Vec<ExtensionInfo> {
|
|
let working_dir_str = working_dir.to_string_lossy();
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.iter()
|
|
.map(|(name, ext)| {
|
|
let instructions = ext.get_instructions().unwrap_or_default();
|
|
let instructions = instructions.replace("{{WORKING_DIR}}", &working_dir_str);
|
|
ExtensionInfo::new(name, &instructions, ext.supports_resources())
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
/// Get aggregated usage statistics
|
|
pub async fn remove_extension(&self, name: &str) -> ExtensionResult<()> {
|
|
let sanitized_name = name_to_key(name);
|
|
self.extensions.lock().await.remove(&sanitized_name);
|
|
self.invalidate_tools_cache_and_bump_version().await;
|
|
Ok(())
|
|
}
|
|
|
|
pub async fn get_extension_and_tool_counts(&self, session_id: &str) -> (usize, usize) {
|
|
let enabled_extensions_count = self.extensions.lock().await.len();
|
|
|
|
let total_tools = self
|
|
.get_prefixed_tools(session_id, None)
|
|
.await
|
|
.map(|tools| tools.len())
|
|
.unwrap_or(0);
|
|
|
|
(enabled_extensions_count, total_tools)
|
|
}
|
|
|
|
pub async fn list_extensions(&self) -> ExtensionResult<Vec<String>> {
|
|
Ok(self.extensions.lock().await.keys().cloned().collect())
|
|
}
|
|
|
|
pub async fn is_extension_enabled(&self, name: &str) -> bool {
|
|
let normalized = name_to_key(name);
|
|
self.extensions.lock().await.contains_key(&normalized)
|
|
}
|
|
|
|
pub async fn get_extension_configs(&self) -> Vec<ExtensionConfig> {
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.values()
|
|
.map(|ext| ext.config.clone())
|
|
.collect()
|
|
}
|
|
|
|
/// Get all tools from all clients with proper prefixing
|
|
pub async fn get_prefixed_tools(
|
|
&self,
|
|
session_id: &str,
|
|
extension_name: Option<String>,
|
|
) -> ExtensionResult<Vec<Tool>> {
|
|
let all_tools = self.get_all_tools_cached(session_id).await?;
|
|
Ok(self.filter_tools(&all_tools, extension_name.as_deref(), None))
|
|
}
|
|
|
|
pub async fn get_prefixed_tools_excluding(
|
|
&self,
|
|
session_id: &str,
|
|
exclude: &str,
|
|
) -> ExtensionResult<Vec<Tool>> {
|
|
let all_tools = self.get_all_tools_cached(session_id).await?;
|
|
Ok(self.filter_tools(&all_tools, None, Some(exclude)))
|
|
}
|
|
|
|
fn filter_tools(
|
|
&self,
|
|
tools: &[Tool],
|
|
extension_name: Option<&str>,
|
|
exclude: Option<&str>,
|
|
) -> Vec<Tool> {
|
|
let extension_name_normalized = extension_name.map(name_to_key);
|
|
let exclude_normalized = exclude.map(name_to_key);
|
|
|
|
tools
|
|
.iter()
|
|
.filter(|tool| {
|
|
let tool_owner = get_tool_owner(tool)
|
|
.map(|s| name_to_key(&s))
|
|
.unwrap_or_else(|| tool.name.split("__").next().unwrap_or("").to_string());
|
|
|
|
if let Some(ref excluded) = exclude_normalized {
|
|
if tool_owner == *excluded {
|
|
return false;
|
|
}
|
|
}
|
|
|
|
if let Some(ref name_filter) = extension_name_normalized {
|
|
tool_owner == *name_filter
|
|
} else {
|
|
true
|
|
}
|
|
})
|
|
.cloned()
|
|
.collect()
|
|
}
|
|
|
|
async fn get_all_tools_cached(&self, session_id: &str) -> ExtensionResult<Arc<Vec<Tool>>> {
|
|
{
|
|
let cache = self.tools_cache.lock().await;
|
|
if let Some(ref tools) = *cache {
|
|
return Ok(Arc::clone(tools));
|
|
}
|
|
}
|
|
|
|
let version_before = self.tools_cache_version.load(Ordering::SeqCst);
|
|
let tools = Arc::new(self.fetch_all_tools(session_id).await?);
|
|
|
|
{
|
|
let mut cache = self.tools_cache.lock().await;
|
|
let version_after = self.tools_cache_version.load(Ordering::SeqCst);
|
|
if version_after == version_before && cache.is_none() {
|
|
*cache = Some(Arc::clone(&tools));
|
|
}
|
|
}
|
|
|
|
Ok(tools)
|
|
}
|
|
|
|
async fn invalidate_tools_cache_and_bump_version(&self) {
|
|
self.tools_cache_version.fetch_add(1, Ordering::SeqCst);
|
|
*self.tools_cache.lock().await = None;
|
|
}
|
|
|
|
async fn fetch_all_tools(&self, session_id: &str) -> ExtensionResult<Vec<Tool>> {
|
|
let clients: Vec<_> = self
|
|
.extensions
|
|
.lock()
|
|
.await
|
|
.iter()
|
|
.map(|(name, ext)| (name.clone(), ext.config.clone(), ext.get_client()))
|
|
.collect();
|
|
|
|
let cancel_token = CancellationToken::default();
|
|
let client_futures = clients.into_iter().map(|(name, config, client)| {
|
|
let cancel_token = cancel_token.clone();
|
|
let ext_name = name.clone();
|
|
async move {
|
|
let mut tools = Vec::new();
|
|
let mut client_tools = match client
|
|
.list_tools(session_id, None, cancel_token.clone())
|
|
.await
|
|
{
|
|
Ok(t) => t,
|
|
Err(e) => {
|
|
warn!(extension = %ext_name, error = %e, "Failed to list tools");
|
|
return (name, vec![]);
|
|
}
|
|
};
|
|
|
|
let expose_unprefixed = is_unprefixed_extension(&config);
|
|
|
|
loop {
|
|
for tool in client_tools.tools {
|
|
if config.is_tool_available(&tool.name) {
|
|
let public_name = if expose_unprefixed {
|
|
tool.name.to_string()
|
|
} else {
|
|
format!("{}__{}", name, tool.name)
|
|
};
|
|
|
|
let mut meta_map = tool
|
|
.meta
|
|
.as_ref()
|
|
.map(|m| m.0.clone())
|
|
.unwrap_or_default();
|
|
meta_map.insert(
|
|
TOOL_EXTENSION_META_KEY.to_string(),
|
|
serde_json::Value::String(name.clone()),
|
|
);
|
|
|
|
tools.push(Tool {
|
|
name: public_name.into(),
|
|
description: tool.description,
|
|
input_schema: tool.input_schema,
|
|
annotations: tool.annotations,
|
|
output_schema: tool.output_schema,
|
|
execution: tool.execution,
|
|
icons: tool.icons,
|
|
title: tool.title,
|
|
meta: Some(rmcp::model::Meta(meta_map)),
|
|
});
|
|
}
|
|
}
|
|
|
|
if client_tools.next_cursor.is_none() {
|
|
break;
|
|
}
|
|
|
|
client_tools = match client
|
|
.list_tools(session_id, client_tools.next_cursor, cancel_token.clone())
|
|
.await
|
|
{
|
|
Ok(t) => t,
|
|
Err(e) => {
|
|
warn!(extension = %ext_name, error = %e, "Failed to list tools (pagination)");
|
|
break;
|
|
}
|
|
};
|
|
}
|
|
|
|
(name, tools)
|
|
}
|
|
});
|
|
|
|
let results = future::join_all(client_futures).await;
|
|
|
|
let mut seen_names: std::collections::HashSet<String> = std::collections::HashSet::new();
|
|
let mut tools = Vec::new();
|
|
for (ext_name, client_tools) in results {
|
|
for tool in client_tools {
|
|
let tool_name = tool.name.to_string();
|
|
if seen_names.contains(&tool_name) {
|
|
warn!(
|
|
tool = %tool_name,
|
|
extension = %ext_name,
|
|
"Duplicate tool name - skipping"
|
|
);
|
|
continue;
|
|
}
|
|
seen_names.insert(tool_name);
|
|
tools.push(tool);
|
|
}
|
|
}
|
|
|
|
Ok(tools)
|
|
}
|
|
|
|
/// Get the extension prompt including client instructions
|
|
pub async fn get_planning_prompt(&self, tools_info: Vec<ToolInfo>) -> String {
|
|
let mut context: HashMap<&str, Value> = HashMap::new();
|
|
context.insert("tools", serde_json::to_value(tools_info).unwrap());
|
|
|
|
prompt_template::render_template("plan.md", &context).expect("Prompt should render")
|
|
}
|
|
|
|
// Function that gets executed for read_resource tool
|
|
pub async fn read_resource_tool(
|
|
&self,
|
|
session_id: &str,
|
|
params: Value,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<Vec<Content>, ErrorData> {
|
|
let uri = require_str_parameter(¶ms, "uri")?;
|
|
|
|
let extension_name = params.get("extension_name").and_then(|v| v.as_str());
|
|
|
|
// If extension name is provided, we can just look it up
|
|
if let Some(ext_name) = extension_name {
|
|
let read_result = self
|
|
.read_resource(session_id, uri, ext_name, cancellation_token.clone())
|
|
.await?;
|
|
|
|
let mut result = Vec::new();
|
|
for content in read_result.contents {
|
|
if let ResourceContents::TextResourceContents { text, .. } = content {
|
|
let content_str = format!("{}\n\n{}", uri, text);
|
|
result.push(Content::text(content_str));
|
|
}
|
|
}
|
|
return Ok(result);
|
|
}
|
|
|
|
// If extension name is not provided, we need to search for the resource across all extensions
|
|
// Loop through each extension and try to read the resource, don't raise an error if the resource is not found
|
|
// TODO: do we want to find if a provided uri is in multiple extensions?
|
|
// currently it will return the first match and skip any others
|
|
let extension_names: Vec<String> = self
|
|
.extensions
|
|
.lock()
|
|
.await
|
|
.iter()
|
|
.filter(|(_name, ext)| ext.supports_resources())
|
|
.map(|(name, _)| name.clone())
|
|
.collect();
|
|
|
|
for extension_name in extension_names {
|
|
let read_result = self
|
|
.read_resource(session_id, uri, &extension_name, cancellation_token.clone())
|
|
.await;
|
|
match read_result {
|
|
Ok(read_result) => {
|
|
let mut result = Vec::new();
|
|
for content in read_result.contents {
|
|
if let ResourceContents::TextResourceContents { text, .. } = content {
|
|
let content_str = format!("{}\n\n{}", uri, text);
|
|
result.push(Content::text(content_str));
|
|
}
|
|
}
|
|
return Ok(result);
|
|
}
|
|
Err(_) => continue,
|
|
}
|
|
}
|
|
|
|
// None of the extensions had the resource so we raise an error
|
|
let available_extensions = self
|
|
.extensions
|
|
.lock()
|
|
.await
|
|
.keys()
|
|
.map(|s| s.as_str())
|
|
.collect::<Vec<&str>>()
|
|
.join(", ");
|
|
let error_msg = format!(
|
|
"Resource with uri '{}' not found. Here are the available extensions: {}",
|
|
uri, available_extensions
|
|
);
|
|
|
|
Err(ErrorData::new(
|
|
ErrorCode::RESOURCE_NOT_FOUND,
|
|
error_msg,
|
|
None,
|
|
))
|
|
}
|
|
|
|
pub async fn read_resource(
|
|
&self,
|
|
session_id: &str,
|
|
uri: &str,
|
|
extension_name: &str,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<rmcp::model::ReadResourceResult, ErrorData> {
|
|
let available_extensions = self
|
|
.extensions
|
|
.lock()
|
|
.await
|
|
.keys()
|
|
.map(|s| s.as_str())
|
|
.collect::<Vec<&str>>()
|
|
.join(", ");
|
|
let error_msg = format!(
|
|
"Extension '{}' not found. Here are the available extensions: {}",
|
|
extension_name, available_extensions
|
|
);
|
|
|
|
let client = self
|
|
.get_server_client(extension_name)
|
|
.await
|
|
.ok_or(ErrorData::new(ErrorCode::INVALID_PARAMS, error_msg, None))?;
|
|
|
|
client
|
|
.read_resource(session_id, uri, cancellation_token)
|
|
.await
|
|
.map_err(|_| {
|
|
ErrorData::new(
|
|
ErrorCode::INTERNAL_ERROR,
|
|
format!("Could not read resource with uri: {}", uri),
|
|
None,
|
|
)
|
|
})
|
|
}
|
|
|
|
pub async fn get_ui_resources(
|
|
&self,
|
|
session_id: &str,
|
|
) -> Result<Vec<(String, Resource)>, ErrorData> {
|
|
let mut ui_resources = Vec::new();
|
|
|
|
let extensions_to_check: Vec<(String, McpClientBox)> = {
|
|
let extensions = self.extensions.lock().await;
|
|
extensions
|
|
.iter()
|
|
.map(|(name, ext)| (name.clone(), ext.get_client()))
|
|
.collect()
|
|
};
|
|
|
|
for (extension_name, client) in extensions_to_check {
|
|
match client
|
|
.list_resources(session_id, None, CancellationToken::default())
|
|
.await
|
|
{
|
|
Ok(list_response) => {
|
|
for resource in list_response.resources {
|
|
if resource.uri.starts_with("ui://") {
|
|
ui_resources.push((extension_name.clone(), resource));
|
|
}
|
|
}
|
|
}
|
|
Err(e) => {
|
|
warn!("Failed to list resources for {}: {:?}", extension_name, e);
|
|
}
|
|
}
|
|
}
|
|
|
|
Ok(ui_resources)
|
|
}
|
|
|
|
async fn list_resources_from_extension(
|
|
&self,
|
|
session_id: &str,
|
|
extension_name: &str,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<Vec<Content>, ErrorData> {
|
|
let client = self
|
|
.get_server_client(extension_name)
|
|
.await
|
|
.ok_or_else(|| {
|
|
ErrorData::new(
|
|
ErrorCode::INVALID_PARAMS,
|
|
format!("Extension {} is not valid", extension_name),
|
|
None,
|
|
)
|
|
})?;
|
|
|
|
client
|
|
.list_resources(session_id, None, cancellation_token)
|
|
.await
|
|
.map_err(|e| {
|
|
ErrorData::new(
|
|
ErrorCode::INTERNAL_ERROR,
|
|
format!("Unable to list resources for {}, {:?}", extension_name, e),
|
|
None,
|
|
)
|
|
})
|
|
.map(|lr| {
|
|
let resource_list = lr
|
|
.resources
|
|
.into_iter()
|
|
.map(|r| format!("{} - {}, uri: ({})", extension_name, r.name, r.uri))
|
|
.collect::<Vec<String>>()
|
|
.join("\n");
|
|
|
|
vec![Content::text(resource_list)]
|
|
})
|
|
}
|
|
|
|
pub async fn list_resources(
|
|
&self,
|
|
session_id: &str,
|
|
params: Value,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<Vec<Content>, ErrorData> {
|
|
let extension = params.get("extension").and_then(|v| v.as_str());
|
|
|
|
match extension {
|
|
Some(extension_name) => {
|
|
// Handle single extension case
|
|
self.list_resources_from_extension(session_id, extension_name, cancellation_token)
|
|
.await
|
|
}
|
|
None => {
|
|
// Handle all extensions case using FuturesUnordered
|
|
let mut futures = FuturesUnordered::new();
|
|
|
|
// Create futures for each resource_capable_extension
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.iter()
|
|
.filter(|(_name, ext)| ext.supports_resources())
|
|
.map(|(name, _ext)| name.clone())
|
|
.for_each(|name| {
|
|
let token = cancellation_token.clone();
|
|
futures.push(async move {
|
|
self.list_resources_from_extension(session_id, name.as_str(), token)
|
|
.await
|
|
});
|
|
});
|
|
|
|
let mut all_resources = Vec::new();
|
|
let mut errors = Vec::new();
|
|
|
|
// Process results as they complete
|
|
while let Some(result) = futures.next().await {
|
|
match result {
|
|
Ok(content) => {
|
|
all_resources.extend(content);
|
|
}
|
|
Err(tool_error) => {
|
|
errors.push(tool_error);
|
|
}
|
|
}
|
|
}
|
|
|
|
if !errors.is_empty() {
|
|
tracing::error!(
|
|
errors = ?errors
|
|
.into_iter()
|
|
.map(|e| format!("{:?}", e))
|
|
.collect::<Vec<_>>(),
|
|
"errors from listing resources"
|
|
);
|
|
}
|
|
|
|
Ok(all_resources)
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn resolve_tool(
|
|
&self,
|
|
session_id: &str,
|
|
tool_name: &str,
|
|
) -> Result<ResolvedTool, ErrorData> {
|
|
if let Some((prefix, actual)) = tool_name.split_once("__") {
|
|
let owner = name_to_key(prefix);
|
|
if let Some(client) = self.get_server_client(&owner).await {
|
|
return Ok(ResolvedTool {
|
|
extension_name: owner,
|
|
actual_tool_name: actual.to_string(),
|
|
client,
|
|
});
|
|
}
|
|
}
|
|
|
|
let tools = self.get_all_tools_cached(session_id).await.map_err(|e| {
|
|
ErrorData::new(
|
|
ErrorCode::INTERNAL_ERROR,
|
|
format!("Failed to get tools: {}", e),
|
|
None,
|
|
)
|
|
})?;
|
|
|
|
if let Some(tool) = tools.iter().find(|t| *t.name == *tool_name) {
|
|
let owner = get_tool_owner(tool).ok_or_else(|| {
|
|
ErrorData::new(
|
|
ErrorCode::RESOURCE_NOT_FOUND,
|
|
format!("Tool '{}' has no owner", tool_name),
|
|
None,
|
|
)
|
|
})?;
|
|
|
|
let actual_tool_name = tool_name
|
|
.strip_prefix(&format!("{owner}__"))
|
|
.unwrap_or(tool_name)
|
|
.to_string();
|
|
|
|
let client = self.get_server_client(&owner).await.ok_or_else(|| {
|
|
ErrorData::new(
|
|
ErrorCode::RESOURCE_NOT_FOUND,
|
|
format!("Extension '{}' not found for tool '{}'", owner, tool_name),
|
|
None,
|
|
)
|
|
})?;
|
|
|
|
return Ok(ResolvedTool {
|
|
extension_name: owner,
|
|
actual_tool_name,
|
|
client,
|
|
});
|
|
}
|
|
|
|
Err(ErrorData::new(
|
|
ErrorCode::RESOURCE_NOT_FOUND,
|
|
format!("Tool '{}' not found", tool_name),
|
|
None,
|
|
))
|
|
}
|
|
|
|
pub async fn dispatch_tool_call(
|
|
&self,
|
|
session_id: &str,
|
|
tool_call: CallToolRequestParams,
|
|
working_dir: Option<&std::path::Path>,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<ToolCallResult> {
|
|
let tool_name_str = tool_call.name.to_string();
|
|
let resolved = self.resolve_tool(session_id, &tool_name_str).await?;
|
|
|
|
if let Some(extension) = self.extensions.lock().await.get(&resolved.extension_name) {
|
|
if !extension
|
|
.config
|
|
.is_tool_available(&resolved.actual_tool_name)
|
|
{
|
|
return Err(ErrorData::new(
|
|
ErrorCode::RESOURCE_NOT_FOUND,
|
|
format!(
|
|
"Tool '{}' is not available for extension '{}'",
|
|
resolved.actual_tool_name, resolved.extension_name
|
|
),
|
|
None,
|
|
)
|
|
.into());
|
|
}
|
|
}
|
|
|
|
let arguments = tool_call.arguments.clone();
|
|
let client = resolved.client.clone();
|
|
let notifications_receiver = client.subscribe().await;
|
|
let session_id = session_id.to_string();
|
|
let actual_tool_name = resolved.actual_tool_name;
|
|
let working_dir_str = working_dir.map(|p| p.to_string_lossy().to_string());
|
|
|
|
let fut = async move {
|
|
tracing::debug!(
|
|
"dispatch_tool_call: calling client.call_tool tool={} session_id={} working_dir={:?}",
|
|
actual_tool_name,
|
|
session_id,
|
|
working_dir_str
|
|
);
|
|
client
|
|
.call_tool(
|
|
&session_id,
|
|
&actual_tool_name,
|
|
arguments,
|
|
working_dir_str.as_deref(),
|
|
cancellation_token,
|
|
)
|
|
.await
|
|
.map_err(|e| match e {
|
|
ServiceError::McpError(error_data) => error_data,
|
|
_ => {
|
|
ErrorData::new(ErrorCode::INTERNAL_ERROR, e.to_string(), e.maybe_to_value())
|
|
}
|
|
})
|
|
};
|
|
|
|
Ok(ToolCallResult {
|
|
result: Box::new(fut.boxed()),
|
|
notification_stream: Some(Box::new(ReceiverStream::new(notifications_receiver))),
|
|
})
|
|
}
|
|
|
|
pub async fn list_prompts_from_extension(
|
|
&self,
|
|
session_id: &str,
|
|
extension_name: &str,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<Vec<Prompt>, ErrorData> {
|
|
let client = self
|
|
.get_server_client(extension_name)
|
|
.await
|
|
.ok_or_else(|| {
|
|
ErrorData::new(
|
|
ErrorCode::INVALID_PARAMS,
|
|
format!("Extension {} is not valid", extension_name),
|
|
None,
|
|
)
|
|
})?;
|
|
|
|
client
|
|
.list_prompts(session_id, None, cancellation_token)
|
|
.await
|
|
.map_err(|e| {
|
|
ErrorData::new(
|
|
ErrorCode::INTERNAL_ERROR,
|
|
format!("Unable to list prompts for {}, {:?}", extension_name, e),
|
|
None,
|
|
)
|
|
})
|
|
.map(|lp| lp.prompts)
|
|
}
|
|
|
|
pub async fn list_prompts(
|
|
&self,
|
|
session_id: &str,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<HashMap<String, Vec<Prompt>>, ErrorData> {
|
|
let mut futures = FuturesUnordered::new();
|
|
|
|
let names: Vec<_> = self.extensions.lock().await.keys().cloned().collect();
|
|
for extension_name in names {
|
|
let token = cancellation_token.clone();
|
|
futures.push(async move {
|
|
(
|
|
extension_name.clone(),
|
|
self.list_prompts_from_extension(session_id, extension_name.as_str(), token)
|
|
.await,
|
|
)
|
|
});
|
|
}
|
|
|
|
let mut all_prompts = HashMap::new();
|
|
let mut errors = Vec::new();
|
|
|
|
// Process results as they complete
|
|
while let Some(result) = futures.next().await {
|
|
let (name, prompts) = result;
|
|
match prompts {
|
|
Ok(content) => {
|
|
all_prompts.insert(name.to_string(), content);
|
|
}
|
|
Err(tool_error) => {
|
|
errors.push(tool_error);
|
|
}
|
|
}
|
|
}
|
|
|
|
if !errors.is_empty() {
|
|
tracing::debug!(
|
|
errors = ?errors
|
|
.into_iter()
|
|
.map(|e| format!("{:?}", e))
|
|
.collect::<Vec<_>>(),
|
|
"errors from listing prompts"
|
|
);
|
|
}
|
|
|
|
Ok(all_prompts)
|
|
}
|
|
|
|
pub async fn get_prompt(
|
|
&self,
|
|
session_id: &str,
|
|
extension_name: &str,
|
|
name: &str,
|
|
arguments: Value,
|
|
cancellation_token: CancellationToken,
|
|
) -> Result<GetPromptResult> {
|
|
let client = self
|
|
.get_server_client(extension_name)
|
|
.await
|
|
.ok_or_else(|| anyhow::anyhow!("Extension {} not found", extension_name))?;
|
|
|
|
client
|
|
.get_prompt(session_id, name, arguments, cancellation_token)
|
|
.await
|
|
.map_err(|e| anyhow::anyhow!("Failed to get prompt: {}", e))
|
|
}
|
|
|
|
pub async fn search_available_extensions(&self) -> Result<Vec<Content>, ErrorData> {
|
|
let mut output_parts = vec![];
|
|
|
|
// First get disabled extensions from current config
|
|
let mut disabled_extensions: Vec<String> = vec![];
|
|
for extension in get_all_extensions() {
|
|
if !extension.enabled {
|
|
let config = extension.config.clone();
|
|
let description = match &config {
|
|
ExtensionConfig::Builtin {
|
|
description,
|
|
display_name,
|
|
..
|
|
} => {
|
|
if description.is_empty() {
|
|
display_name.as_deref().unwrap_or("Built-in extension")
|
|
} else {
|
|
description
|
|
}
|
|
}
|
|
ExtensionConfig::Sse { .. } => "SSE extension (unsupported)",
|
|
ExtensionConfig::Platform { description, .. }
|
|
| ExtensionConfig::StreamableHttp { description, .. }
|
|
| ExtensionConfig::Stdio { description, .. }
|
|
| ExtensionConfig::Frontend { description, .. }
|
|
| ExtensionConfig::InlinePython { description, .. } => description,
|
|
};
|
|
disabled_extensions.push(format!("- {} - {}", config.name(), description));
|
|
}
|
|
}
|
|
|
|
// Get currently enabled extensions that can be disabled
|
|
let enabled_extensions: Vec<String> =
|
|
self.extensions.lock().await.keys().cloned().collect();
|
|
|
|
// Build output string
|
|
if !disabled_extensions.is_empty() {
|
|
output_parts.push(format!(
|
|
"Extensions available to enable:\n{}\n",
|
|
disabled_extensions.join("\n")
|
|
));
|
|
} else {
|
|
output_parts.push("No extensions available to enable.\n".to_string());
|
|
}
|
|
|
|
if !enabled_extensions.is_empty() {
|
|
output_parts.push(format!(
|
|
"\n\nExtensions available to disable:\n{}\n",
|
|
enabled_extensions
|
|
.iter()
|
|
.map(|name| format!("- {}", name))
|
|
.collect::<Vec<_>>()
|
|
.join("\n")
|
|
));
|
|
} else {
|
|
output_parts.push("No extensions that can be disabled.\n".to_string());
|
|
}
|
|
|
|
Ok(vec![Content::text(output_parts.join("\n"))])
|
|
}
|
|
|
|
async fn get_server_client(&self, name: impl Into<String>) -> Option<McpClientBox> {
|
|
let normalized = name_to_key(&name.into());
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.get(&normalized)
|
|
.map(|ext| ext.get_client())
|
|
}
|
|
|
|
pub async fn collect_moim(
|
|
&self,
|
|
session_id: &str,
|
|
working_dir: &std::path::Path,
|
|
) -> Option<String> {
|
|
// Skip MOIM for models with small context windows to avoid consuming limited context
|
|
const MIN_CONTEXT_FOR_MOIM: usize = 32_000;
|
|
if let Ok(provider_guard) = self.provider.try_lock() {
|
|
if let Some(provider) = provider_guard.as_ref() {
|
|
if provider.get_model_config().context_limit() < MIN_CONTEXT_FOR_MOIM {
|
|
return None;
|
|
}
|
|
}
|
|
}
|
|
|
|
// Use minute-level granularity to prevent conversation changes every second
|
|
let timestamp = chrono::Local::now().format("%Y-%m-%d %H:%M:00").to_string();
|
|
let mut content = format!(
|
|
"<info-msg>\nIt is currently {}\nWorking directory: {}\n",
|
|
timestamp,
|
|
working_dir.display()
|
|
);
|
|
|
|
if let Ok(session) = self
|
|
.context
|
|
.session_manager
|
|
.get_session(session_id, false)
|
|
.await
|
|
{
|
|
if let (Some(total), Some(config)) =
|
|
(session.total_tokens, session.model_config.as_ref())
|
|
{
|
|
let limit = config.context_limit();
|
|
if total > 0 && limit > 0 {
|
|
let pct = (total as f64 / limit as f64 * 100.0).round() as u32;
|
|
content.push_str(&format!(
|
|
"Context: ~{}k/{}k tokens used ({}%)\n",
|
|
total / 1000,
|
|
limit / 1000,
|
|
pct
|
|
));
|
|
}
|
|
}
|
|
}
|
|
|
|
let platform_clients: Vec<(String, McpClientBox)> = {
|
|
let extensions = self.extensions.lock().await;
|
|
extensions
|
|
.iter()
|
|
.filter_map(|(name, extension)| {
|
|
if let ExtensionConfig::Platform { .. } = &extension.config {
|
|
Some((name.clone(), extension.get_client()))
|
|
} else {
|
|
None
|
|
}
|
|
})
|
|
.collect()
|
|
};
|
|
|
|
for (name, client) in platform_clients {
|
|
if let Some(moim_content) = client.get_moim(session_id).await {
|
|
tracing::debug!("MOIM content from {}: {} chars", name, moim_content.len());
|
|
content.push('\n');
|
|
content.push_str(&moim_content);
|
|
}
|
|
}
|
|
|
|
content.push_str("\n</info-msg>");
|
|
|
|
Some(content)
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use rmcp::model::CallToolResult;
|
|
use rmcp::model::{InitializeResult, JsonObject};
|
|
use rmcp::{object, ServiceError as Error};
|
|
|
|
use rmcp::model::ListPromptsResult;
|
|
use rmcp::model::ListResourcesResult;
|
|
use rmcp::model::ListToolsResult;
|
|
use rmcp::model::ReadResourceResult;
|
|
use rmcp::model::ServerNotification;
|
|
|
|
use tokio::sync::mpsc;
|
|
|
|
impl ExtensionManager {
|
|
async fn add_mock_extension(&self, name: String, client: McpClientBox) {
|
|
self.add_mock_extension_with_tools(name, client, vec![])
|
|
.await;
|
|
}
|
|
|
|
async fn add_mock_extension_with_tools(
|
|
&self,
|
|
name: String,
|
|
client: McpClientBox,
|
|
available_tools: Vec<String>,
|
|
) {
|
|
let sanitized_name = name_to_key(&name);
|
|
let config = ExtensionConfig::Builtin {
|
|
name: name.clone(),
|
|
display_name: Some(name.clone()),
|
|
description: "built-in".to_string(),
|
|
timeout: None,
|
|
bundled: None,
|
|
available_tools,
|
|
};
|
|
let extension = Extension::new(config, client, None, None);
|
|
self.extensions
|
|
.lock()
|
|
.await
|
|
.insert(sanitized_name, extension);
|
|
self.invalidate_tools_cache_and_bump_version().await;
|
|
}
|
|
}
|
|
|
|
struct MockClient {}
|
|
|
|
#[async_trait::async_trait]
|
|
impl McpClientTrait for MockClient {
|
|
fn get_info(&self) -> Option<&InitializeResult> {
|
|
None
|
|
}
|
|
|
|
async fn list_resources(
|
|
&self,
|
|
_session_id: &str,
|
|
_next_cursor: Option<String>,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<ListResourcesResult, Error> {
|
|
Err(Error::TransportClosed)
|
|
}
|
|
|
|
async fn read_resource(
|
|
&self,
|
|
_session_id: &str,
|
|
_uri: &str,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<ReadResourceResult, Error> {
|
|
Err(Error::TransportClosed)
|
|
}
|
|
|
|
async fn list_tools(
|
|
&self,
|
|
_session_id: &str,
|
|
_next_cursor: Option<String>,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<ListToolsResult, Error> {
|
|
use serde_json::json;
|
|
use std::sync::Arc;
|
|
Ok(ListToolsResult {
|
|
tools: vec![
|
|
Tool::new(
|
|
"tool".to_string(),
|
|
"A basic tool".to_string(),
|
|
Arc::new(json!({}).as_object().unwrap().clone()),
|
|
),
|
|
Tool::new(
|
|
"available_tool".to_string(),
|
|
"An available tool".to_string(),
|
|
Arc::new(json!({}).as_object().unwrap().clone()),
|
|
),
|
|
Tool::new(
|
|
"hidden_tool".to_string(),
|
|
"hidden tool".to_string(),
|
|
Arc::new(json!({}).as_object().unwrap().clone()),
|
|
),
|
|
],
|
|
next_cursor: None,
|
|
meta: None,
|
|
})
|
|
}
|
|
|
|
async fn call_tool(
|
|
&self,
|
|
_session_id: &str,
|
|
name: &str,
|
|
_arguments: Option<JsonObject>,
|
|
_working_dir: Option<&str>,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<CallToolResult, Error> {
|
|
match name {
|
|
"tool" | "test__tool" | "available_tool" | "hidden_tool" => Ok(CallToolResult {
|
|
content: vec![],
|
|
is_error: None,
|
|
structured_content: None,
|
|
meta: None,
|
|
}),
|
|
_ => Err(Error::TransportClosed),
|
|
}
|
|
}
|
|
|
|
async fn list_prompts(
|
|
&self,
|
|
_session_id: &str,
|
|
_next_cursor: Option<String>,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<ListPromptsResult, Error> {
|
|
Err(Error::TransportClosed)
|
|
}
|
|
|
|
async fn get_prompt(
|
|
&self,
|
|
_session_id: &str,
|
|
_name: &str,
|
|
_arguments: Value,
|
|
_cancellation_token: CancellationToken,
|
|
) -> Result<GetPromptResult, Error> {
|
|
Err(Error::TransportClosed)
|
|
}
|
|
|
|
async fn subscribe(&self) -> mpsc::Receiver<ServerNotification> {
|
|
mpsc::channel(1).1
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_dispatch_tool_call() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
// Add some mock clients using the helper method
|
|
extension_manager
|
|
.add_mock_extension("test_client".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
extension_manager
|
|
.add_mock_extension("__cli__ent__".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
extension_manager
|
|
.add_mock_extension("client 🚀".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "test_client__tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
assert!(result.is_ok());
|
|
|
|
let tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "test_client__available_tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
assert!(result.is_ok());
|
|
|
|
let tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "__cli__ent____tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
assert!(result.is_ok());
|
|
|
|
let tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "client___tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
assert!(result.is_ok());
|
|
|
|
let invalid_tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "client___tools".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
invalid_tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
if let Err(err) = result {
|
|
let tool_err = err.downcast_ref::<ErrorData>().expect("Expected ErrorData");
|
|
assert_eq!(tool_err.code, ErrorCode::RESOURCE_NOT_FOUND);
|
|
} else {
|
|
panic!("Expected ErrorData with ErrorCode::RESOURCE_NOT_FOUND");
|
|
}
|
|
|
|
let invalid_tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "_client__tools".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
invalid_tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
if let Err(err) = result {
|
|
let tool_err = err.downcast_ref::<ErrorData>().expect("Expected ErrorData");
|
|
assert_eq!(tool_err.code, ErrorCode::RESOURCE_NOT_FOUND);
|
|
} else {
|
|
panic!("Expected ErrorData with ErrorCode::RESOURCE_NOT_FOUND");
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_tool_availability_filtering() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
// Only "available_tool" should be available to the LLM
|
|
let available_tools = vec!["available_tool".to_string()];
|
|
|
|
extension_manager
|
|
.add_mock_extension_with_tools(
|
|
"test_extension".to_string(),
|
|
Arc::new(MockClient {}),
|
|
available_tools,
|
|
)
|
|
.await;
|
|
|
|
let tools = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let tool_names: Vec<String> = tools.iter().map(|t| t.name.to_string()).collect();
|
|
assert!(!tool_names.iter().any(|name| name == "test_extension__tool")); // Default unavailable
|
|
assert!(tool_names
|
|
.iter()
|
|
.any(|name| name == "test_extension__available_tool"));
|
|
assert!(!tool_names
|
|
.iter()
|
|
.any(|name| name == "test_extension__hidden_tool"));
|
|
assert!(tool_names.len() == 1);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_tool_availability_defaults_to_available() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
extension_manager
|
|
.add_mock_extension_with_tools(
|
|
"test_extension".to_string(),
|
|
Arc::new(MockClient {}),
|
|
vec![], // Empty available_tools means all tools are available by default
|
|
)
|
|
.await;
|
|
|
|
let tools = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let tool_names: Vec<String> = tools.iter().map(|t| t.name.to_string()).collect();
|
|
assert!(tool_names.iter().any(|name| name == "test_extension__tool"));
|
|
assert!(tool_names
|
|
.iter()
|
|
.any(|name| name == "test_extension__available_tool"));
|
|
assert!(tool_names
|
|
.iter()
|
|
.any(|name| name == "test_extension__hidden_tool"));
|
|
assert!(tool_names.len() == 3);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_dispatch_unavailable_tool_returns_error() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
let available_tools = vec!["available_tool".to_string()];
|
|
|
|
extension_manager
|
|
.add_mock_extension_with_tools(
|
|
"test_extension".to_string(),
|
|
Arc::new(MockClient {}),
|
|
available_tools,
|
|
)
|
|
.await;
|
|
|
|
let unavailable_tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "test_extension__tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
unavailable_tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
|
|
if let Err(err) = result {
|
|
let tool_err = err.downcast_ref::<ErrorData>().expect("Expected ErrorData");
|
|
assert_eq!(tool_err.code, ErrorCode::RESOURCE_NOT_FOUND);
|
|
} else {
|
|
panic!("Expected ErrorData with ErrorCode::RESOURCE_NOT_FOUND");
|
|
}
|
|
|
|
// Try to call an available tool - should succeed
|
|
let available_tool_call = CallToolRequestParams {
|
|
meta: None,
|
|
task: None,
|
|
name: "test_extension__available_tool".to_string().into(),
|
|
arguments: Some(object!({})),
|
|
};
|
|
|
|
let result = extension_manager
|
|
.dispatch_tool_call(
|
|
"test-session-id",
|
|
available_tool_call,
|
|
None,
|
|
CancellationToken::default(),
|
|
)
|
|
.await;
|
|
|
|
assert!(result.is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_streamable_http_header_env_substitution() {
|
|
let mut env_map = HashMap::new();
|
|
env_map.insert("AUTH_TOKEN".to_string(), "secret123".to_string());
|
|
env_map.insert("API_KEY".to_string(), "key456".to_string());
|
|
|
|
// Test ${VAR} syntax
|
|
let result = substitute_env_vars("Bearer ${ AUTH_TOKEN }", &env_map);
|
|
assert_eq!(result, "Bearer secret123");
|
|
|
|
// Test ${VAR} syntax without spaces
|
|
let result = substitute_env_vars("Bearer ${AUTH_TOKEN}", &env_map);
|
|
assert_eq!(result, "Bearer secret123");
|
|
|
|
// Test $VAR syntax
|
|
let result = substitute_env_vars("Bearer $AUTH_TOKEN", &env_map);
|
|
assert_eq!(result, "Bearer secret123");
|
|
|
|
// Test multiple substitutions
|
|
let result = substitute_env_vars("Key: $API_KEY, Token: ${AUTH_TOKEN}", &env_map);
|
|
assert_eq!(result, "Key: key456, Token: secret123");
|
|
|
|
// Test no substitution when variable doesn't exist
|
|
let result = substitute_env_vars("Bearer ${UNKNOWN_VAR}", &env_map);
|
|
assert_eq!(result, "Bearer ${UNKNOWN_VAR}");
|
|
|
|
// Test mixed content
|
|
let result = substitute_env_vars(
|
|
"Authorization: Bearer ${AUTH_TOKEN} and API ${API_KEY}",
|
|
&env_map,
|
|
);
|
|
assert_eq!(result, "Authorization: Bearer secret123 and API key456");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_collect_moim_uses_minute_granularity() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let em = ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
let working_dir = std::path::Path::new("/tmp");
|
|
|
|
if let Some(moim) = em.collect_moim("test-session-id", working_dir).await {
|
|
// Timestamp should end with :00 (seconds fixed to 00)
|
|
assert!(
|
|
moim.contains(":00\n"),
|
|
"Timestamp should use minute granularity"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_tools_cache_invalidated_on_add_extension() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
extension_manager
|
|
.add_mock_extension("ext_a".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tools_after_first = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools_after_first
|
|
.iter()
|
|
.map(|t| t.name.to_string())
|
|
.collect();
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(!tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
|
|
extension_manager
|
|
.add_mock_extension("ext_b".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tools_after_second = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools_after_second
|
|
.iter()
|
|
.map(|t| t.name.to_string())
|
|
.collect();
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_tools_cache_invalidated_on_remove_extension() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
extension_manager
|
|
.add_mock_extension("ext_a".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
extension_manager
|
|
.add_mock_extension("ext_b".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tools_before = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools_before.iter().map(|t| t.name.to_string()).collect();
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
|
|
extension_manager.remove_extension("ext_b").await.unwrap();
|
|
|
|
let tools_after = extension_manager
|
|
.get_prefixed_tools("test-session-id", None)
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools_after.iter().map(|t| t.name.to_string()).collect();
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(!tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_get_prefixed_tools_excluding() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
extension_manager
|
|
.add_mock_extension("ext_a".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
extension_manager
|
|
.add_mock_extension("ext_b".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tools = extension_manager
|
|
.get_prefixed_tools_excluding("test-session-id", "ext_a")
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools.iter().map(|t| t.name.to_string()).collect();
|
|
|
|
assert!(!tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_get_prefixed_tools_by_extension_name() {
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let extension_manager =
|
|
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
|
|
|
extension_manager
|
|
.add_mock_extension("ext_a".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
extension_manager
|
|
.add_mock_extension("ext_b".to_string(), Arc::new(MockClient {}))
|
|
.await;
|
|
|
|
let tools = extension_manager
|
|
.get_prefixed_tools("test-session-id", Some("ext_a".to_string()))
|
|
.await
|
|
.unwrap();
|
|
let tool_names: Vec<String> = tools.iter().map(|t| t.name.to_string()).collect();
|
|
|
|
assert!(tool_names.iter().any(|n| n.starts_with("ext_a__")));
|
|
assert!(!tool_names.iter().any(|n| n.starts_with("ext_b__")));
|
|
}
|
|
}
|