fix(goose): propagate session_id across providers and MCP (#6584)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -23,8 +23,7 @@ use crate::agents::subagent_task_config::TaskConfig;
|
||||
use crate::agents::subagent_tool::{
|
||||
create_subagent_tool, handle_subagent_tool, SUBAGENT_TOOL_NAME,
|
||||
};
|
||||
use crate::agents::types::SessionConfig;
|
||||
use crate::agents::types::{FrontendTool, SharedProvider, ToolResultReceiver};
|
||||
use crate::agents::types::{FrontendTool, SessionConfig, SharedProvider, ToolResultReceiver};
|
||||
use crate::config::permission::PermissionManager;
|
||||
use crate::config::{get_enabled_extensions, Config, GooseMode};
|
||||
use crate::context_mgmt::{
|
||||
@@ -785,7 +784,7 @@ impl Agent {
|
||||
pub async fn list_tools(&self, session_id: &str, extension_name: Option<String>) -> Vec<Tool> {
|
||||
let mut prefixed_tools = self
|
||||
.extension_manager
|
||||
.get_prefixed_tools(extension_name.clone())
|
||||
.get_prefixed_tools(session_id, extension_name.clone())
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
|
||||
@@ -999,7 +998,14 @@ impl Agent {
|
||||
)
|
||||
);
|
||||
|
||||
match compact_messages(self.provider().await?.as_ref(), &conversation_to_compact, false).await {
|
||||
match compact_messages(
|
||||
self.provider().await?.as_ref(),
|
||||
&session_config.id,
|
||||
&conversation_to_compact,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((compacted_conversation, summarization_usage)) => {
|
||||
session_manager.replace_conversation(&session_config.id, &compacted_conversation).await?;
|
||||
self.update_session_metrics(&session_config, &summarization_usage, true).await?;
|
||||
@@ -1041,7 +1047,7 @@ impl Agent {
|
||||
cancel_token: Option<CancellationToken>,
|
||||
) -> Result<BoxStream<'_, Result<AgentEvent>>> {
|
||||
let context = self
|
||||
.prepare_reply_context(&session_config.id, conversation, &session.working_dir)
|
||||
.prepare_reply_context(&session.id, conversation, session.working_dir.as_path())
|
||||
.await?;
|
||||
let ReplyContext {
|
||||
mut conversation,
|
||||
@@ -1108,6 +1114,7 @@ impl Agent {
|
||||
|
||||
let mut stream = Self::stream_response_from_provider(
|
||||
self.provider().await?,
|
||||
&session_config.id,
|
||||
&system_prompt,
|
||||
conversation_with_moim.messages(),
|
||||
&tools,
|
||||
@@ -1388,7 +1395,14 @@ impl Agent {
|
||||
)
|
||||
);
|
||||
|
||||
match compact_messages(self.provider().await?.as_ref(), &conversation, false).await {
|
||||
match compact_messages(
|
||||
self.provider().await?.as_ref(),
|
||||
&session_config.id,
|
||||
&conversation,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((compacted_conversation, usage)) => {
|
||||
session_manager.replace_conversation(&session_config.id, &compacted_conversation).await?;
|
||||
self.update_session_metrics(&session_config, &usage, true).await?;
|
||||
@@ -1533,18 +1547,23 @@ impl Agent {
|
||||
prompt_manager.set_system_prompt_override(template);
|
||||
}
|
||||
|
||||
pub async fn list_extension_prompts(&self) -> HashMap<String, Vec<Prompt>> {
|
||||
pub async fn list_extension_prompts(&self, session_id: &str) -> HashMap<String, Vec<Prompt>> {
|
||||
self.extension_manager
|
||||
.list_prompts(CancellationToken::default())
|
||||
.list_prompts(session_id, CancellationToken::default())
|
||||
.await
|
||||
.expect("Failed to list prompts")
|
||||
}
|
||||
|
||||
pub async fn get_prompt(&self, name: &str, arguments: Value) -> Result<GetPromptResult> {
|
||||
pub async fn get_prompt(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Value,
|
||||
) -> Result<GetPromptResult> {
|
||||
// First find which extension has this prompt
|
||||
let prompts = self
|
||||
.extension_manager
|
||||
.list_prompts(CancellationToken::default())
|
||||
.list_prompts(session_id, CancellationToken::default())
|
||||
.await
|
||||
.map_err(|e| anyhow!("Failed to list prompts: {}", e))?;
|
||||
|
||||
@@ -1555,7 +1574,13 @@ impl Agent {
|
||||
{
|
||||
return self
|
||||
.extension_manager
|
||||
.get_prompt(extension, name, arguments, CancellationToken::default())
|
||||
.get_prompt(
|
||||
session_id,
|
||||
extension,
|
||||
name,
|
||||
arguments,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow!("Failed to get prompt: {}", e));
|
||||
}
|
||||
@@ -1563,8 +1588,11 @@ impl Agent {
|
||||
Err(anyhow!("Prompt '{}' not found", name))
|
||||
}
|
||||
|
||||
pub async fn get_plan_prompt(&self) -> Result<String> {
|
||||
let tools = self.extension_manager.get_prefixed_tools(None).await?;
|
||||
pub async fn get_plan_prompt(&self, session_id: &str) -> Result<String> {
|
||||
let tools = self
|
||||
.extension_manager
|
||||
.get_prefixed_tools(session_id, None)
|
||||
.await?;
|
||||
let tools_info = tools
|
||||
.into_iter()
|
||||
.map(|tool| {
|
||||
@@ -1591,13 +1619,19 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn create_recipe(&self, mut messages: Conversation) -> Result<Recipe> {
|
||||
pub async fn create_recipe(
|
||||
&self,
|
||||
session_id: &str,
|
||||
mut messages: Conversation,
|
||||
) -> Result<Recipe> {
|
||||
tracing::info!("Starting recipe creation with {} messages", messages.len());
|
||||
|
||||
let extensions_info = self.extension_manager.get_extensions_info().await;
|
||||
tracing::debug!("Retrieved {} extensions info", extensions_info.len());
|
||||
let (extension_count, tool_count) =
|
||||
self.extension_manager.get_extension_and_tool_counts().await;
|
||||
let (extension_count, tool_count) = self
|
||||
.extension_manager
|
||||
.get_extension_and_tool_counts(session_id)
|
||||
.await;
|
||||
|
||||
// Get model name from provider
|
||||
let provider = self.provider().await.map_err(|e| {
|
||||
@@ -1619,7 +1653,7 @@ impl Agent {
|
||||
let recipe_prompt = prompt_manager.get_recipe_prompt().await;
|
||||
let tools = self
|
||||
.extension_manager
|
||||
.get_prefixed_tools(None)
|
||||
.get_prefixed_tools(session_id, None)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
tracing::error!("Failed to get tools for recipe creation: {}", e);
|
||||
@@ -1651,7 +1685,7 @@ impl Agent {
|
||||
tracing::error!("{}", error);
|
||||
error
|
||||
})?
|
||||
.complete(&system_prompt, messages.messages(), &tools)
|
||||
.complete(session_id, &system_prompt, messages.messages(), &tools)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
tracing::error!("Provider completion failed during recipe creation: {}", e);
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::agents::extension::PlatformExtensionContext;
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait};
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use indoc::indoc;
|
||||
@@ -281,6 +281,7 @@ impl ChatRecallClient {
|
||||
impl McpClientTrait for ChatRecallClient {
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -293,12 +294,11 @@ impl McpClientTrait for ChatRecallClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let session_id = &meta.session_id;
|
||||
let content = match name {
|
||||
"chatrecall" => self.handle_chatrecall(session_id, arguments).await,
|
||||
_ => Err(format!("Unknown tool: {}", name)),
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use crate::agents::extension::PlatformExtensionContext;
|
||||
use crate::agents::extension_manager::get_parameter_names;
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait};
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use boa_engine::builtins::promise::PromiseState;
|
||||
@@ -458,7 +458,7 @@ impl CodeExecutionClient {
|
||||
Ok(Self { info, context })
|
||||
}
|
||||
|
||||
async fn get_tool_infos(&self) -> Vec<ToolInfo> {
|
||||
async fn get_tool_infos(&self, session_id: &str) -> Vec<ToolInfo> {
|
||||
let Some(manager) = self
|
||||
.context
|
||||
.extension_manager
|
||||
@@ -468,7 +468,10 @@ impl CodeExecutionClient {
|
||||
return Vec::new();
|
||||
};
|
||||
|
||||
match manager.get_prefixed_tools_excluding(EXTENSION_NAME).await {
|
||||
match manager
|
||||
.get_prefixed_tools_excluding(session_id, EXTENSION_NAME)
|
||||
.await
|
||||
{
|
||||
Ok(tools) if !tools.is_empty() => {
|
||||
tools.iter().filter_map(ToolInfo::from_mcp_tool).collect()
|
||||
}
|
||||
@@ -488,7 +491,7 @@ impl CodeExecutionClient {
|
||||
.ok_or("Missing required parameter: code")?
|
||||
.to_string();
|
||||
|
||||
let tools = self.get_tool_infos().await;
|
||||
let tools = self.get_tool_infos(session_id).await;
|
||||
let (call_tx, call_rx) = mpsc::unbounded_channel();
|
||||
let tool_handler = tokio::spawn(Self::run_tool_handler(
|
||||
session_id.to_string(),
|
||||
@@ -506,6 +509,7 @@ impl CodeExecutionClient {
|
||||
|
||||
async fn handle_read_module(
|
||||
&self,
|
||||
session_id: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
) -> Result<Vec<Content>, String> {
|
||||
let path = arguments
|
||||
@@ -514,7 +518,7 @@ impl CodeExecutionClient {
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or("Missing required parameter: module_path")?;
|
||||
|
||||
let tools = self.get_tool_infos().await;
|
||||
let tools = self.get_tool_infos(session_id).await;
|
||||
let parts: Vec<&str> = path.trim_start_matches('/').split('/').collect();
|
||||
|
||||
match parts.as_slice() {
|
||||
@@ -549,6 +553,7 @@ impl CodeExecutionClient {
|
||||
|
||||
async fn handle_search_modules(
|
||||
&self,
|
||||
session_id: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
) -> Result<Vec<Content>, String> {
|
||||
let terms = arguments
|
||||
@@ -580,7 +585,7 @@ impl CodeExecutionClient {
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
|
||||
let tools = self.get_tool_infos().await;
|
||||
let tools = self.get_tool_infos(session_id).await;
|
||||
Self::handle_search(&tools, &terms_vec, use_regex)
|
||||
}
|
||||
|
||||
@@ -707,6 +712,7 @@ impl McpClientTrait for CodeExecutionClient {
|
||||
#[allow(clippy::too_many_lines)]
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -831,15 +837,15 @@ impl McpClientTrait for CodeExecutionClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let content = match name {
|
||||
"execute_code" => self.handle_execute_code(&meta.session_id, arguments).await,
|
||||
"read_module" => self.handle_read_module(arguments).await,
|
||||
"search_modules" => self.handle_search_modules(arguments).await,
|
||||
"execute_code" => self.handle_execute_code(session_id, arguments).await,
|
||||
"read_module" => self.handle_read_module(session_id, arguments).await,
|
||||
"search_modules" => self.handle_search_modules(session_id, arguments).await,
|
||||
_ => Err(format!("Unknown tool: {name}")),
|
||||
};
|
||||
|
||||
@@ -855,8 +861,8 @@ impl McpClientTrait for CodeExecutionClient {
|
||||
Some(&self.info)
|
||||
}
|
||||
|
||||
async fn get_moim(&self, _session_id: &str) -> Option<String> {
|
||||
let tools = self.get_tool_infos().await;
|
||||
async fn get_moim(&self, session_id: &str) -> Option<String> {
|
||||
let tools = self.get_tool_infos(session_id).await;
|
||||
if tools.is_empty() {
|
||||
return None;
|
||||
}
|
||||
@@ -909,9 +915,9 @@ mod tests {
|
||||
|
||||
let result = client
|
||||
.call_tool(
|
||||
"test-session-id",
|
||||
"execute_code",
|
||||
Some(args),
|
||||
McpMeta::new("test-session-id"),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
@@ -946,9 +952,9 @@ mod tests {
|
||||
|
||||
let result = client
|
||||
.call_tool(
|
||||
"test-session-id",
|
||||
"execute_code",
|
||||
Some(args),
|
||||
McpMeta::new("test-session-id"),
|
||||
CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
@@ -984,7 +990,9 @@ mod tests {
|
||||
Value::String("nonexistent".to_string()),
|
||||
);
|
||||
|
||||
let result = client.handle_read_module(Some(args)).await;
|
||||
let result = client
|
||||
.handle_read_module("test-session-id", Some(args))
|
||||
.await;
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
|
||||
@@ -88,6 +88,7 @@ impl Agent {
|
||||
|
||||
let (compacted_conversation, _usage) = compact_messages(
|
||||
self.provider().await?.as_ref(),
|
||||
session_id,
|
||||
&conversation,
|
||||
true, // is_manual_compact
|
||||
)
|
||||
@@ -128,11 +129,11 @@ impl Agent {
|
||||
async fn handle_prompts_command(
|
||||
&self,
|
||||
params: &[&str],
|
||||
_session_id: &str,
|
||||
session_id: &str,
|
||||
) -> Result<Option<Message>> {
|
||||
let extension_filter = params.first().map(|s| s.to_string());
|
||||
|
||||
let prompts = self.list_extension_prompts().await;
|
||||
let prompts = self.list_extension_prompts(session_id).await;
|
||||
|
||||
if let Some(filter) = &extension_filter {
|
||||
if !prompts.contains_key(filter) {
|
||||
@@ -182,7 +183,7 @@ impl Agent {
|
||||
let is_info = params.get(1).map(|s| *s == "--info").unwrap_or(false);
|
||||
|
||||
if is_info {
|
||||
let prompts = self.list_extension_prompts().await;
|
||||
let prompts = self.list_extension_prompts(session_id).await;
|
||||
let mut prompt_info = None;
|
||||
|
||||
for (extension, prompt_list) in prompts {
|
||||
@@ -225,7 +226,10 @@ impl Agent {
|
||||
let arguments_value = serde_json::to_value(arguments)
|
||||
.map_err(|e| anyhow!("Failed to serialize arguments: {}", e))?;
|
||||
|
||||
match self.get_prompt(&prompt_name, arguments_value).await {
|
||||
match self
|
||||
.get_prompt(session_id, &prompt_name, arguments_value)
|
||||
.await
|
||||
{
|
||||
Ok(prompt_result) => {
|
||||
for (i, prompt_message) in prompt_result.messages.into_iter().enumerate() {
|
||||
let msg = Message::from(prompt_message);
|
||||
|
||||
@@ -33,7 +33,7 @@ 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::{McpClient, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{McpClient, McpClientTrait};
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{get_all_extensions, Config};
|
||||
use crate::oauth::oauth_flow;
|
||||
@@ -685,11 +685,11 @@ impl ExtensionManager {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_extension_and_tool_counts(&self) -> (usize, usize) {
|
||||
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(None)
|
||||
.get_prefixed_tools(session_id, None)
|
||||
.await
|
||||
.map(|tools| tools.len())
|
||||
.unwrap_or(0);
|
||||
@@ -717,14 +717,19 @@ impl ExtensionManager {
|
||||
/// 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().await?;
|
||||
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, exclude: &str) -> ExtensionResult<Vec<Tool>> {
|
||||
let all_tools = self.get_all_tools_cached().await?;
|
||||
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)))
|
||||
}
|
||||
|
||||
@@ -755,7 +760,7 @@ impl ExtensionManager {
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn get_all_tools_cached(&self) -> ExtensionResult<Arc<Vec<Tool>>> {
|
||||
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 {
|
||||
@@ -764,7 +769,7 @@ impl ExtensionManager {
|
||||
}
|
||||
|
||||
let version_before = self.tools_cache_version.load(Ordering::SeqCst);
|
||||
let tools = Arc::new(self.fetch_all_tools().await?);
|
||||
let tools = Arc::new(self.fetch_all_tools(session_id).await?);
|
||||
|
||||
{
|
||||
let mut cache = self.tools_cache.lock().await;
|
||||
@@ -782,7 +787,7 @@ impl ExtensionManager {
|
||||
*self.tools_cache.lock().await = None;
|
||||
}
|
||||
|
||||
async fn fetch_all_tools(&self) -> ExtensionResult<Vec<Tool>> {
|
||||
async fn fetch_all_tools(&self, session_id: &str) -> ExtensionResult<Vec<Tool>> {
|
||||
let clients: Vec<_> = self
|
||||
.extensions
|
||||
.lock()
|
||||
@@ -799,7 +804,7 @@ impl ExtensionManager {
|
||||
let mut tools = Vec::new();
|
||||
let client_guard = client.lock().await;
|
||||
let mut client_tools = match client_guard
|
||||
.list_tools(None, cancel_token.clone())
|
||||
.list_tools(session_id, None, cancel_token.clone())
|
||||
.await
|
||||
{
|
||||
Ok(t) => t,
|
||||
@@ -830,7 +835,7 @@ impl ExtensionManager {
|
||||
}
|
||||
|
||||
client_tools = match client_guard
|
||||
.list_tools(client_tools.next_cursor, cancel_token.clone())
|
||||
.list_tools(session_id, client_tools.next_cursor, cancel_token.clone())
|
||||
.await
|
||||
{
|
||||
Ok(t) => t,
|
||||
@@ -876,6 +881,7 @@ impl ExtensionManager {
|
||||
// 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> {
|
||||
@@ -886,7 +892,7 @@ impl ExtensionManager {
|
||||
// If extension name is provided, we can just look it up
|
||||
if let Some(ext_name) = extension_name {
|
||||
let read_result = self
|
||||
.read_resource(uri, ext_name, cancellation_token.clone())
|
||||
.read_resource(session_id, uri, ext_name, cancellation_token.clone())
|
||||
.await?;
|
||||
|
||||
let mut result = Vec::new();
|
||||
@@ -909,7 +915,7 @@ impl ExtensionManager {
|
||||
|
||||
for extension_name in extension_names {
|
||||
let read_result = self
|
||||
.read_resource(uri, &extension_name, cancellation_token.clone())
|
||||
.read_resource(session_id, uri, &extension_name, cancellation_token.clone())
|
||||
.await;
|
||||
match read_result {
|
||||
Ok(read_result) => {
|
||||
@@ -949,6 +955,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn read_resource(
|
||||
&self,
|
||||
session_id: &str,
|
||||
uri: &str,
|
||||
extension_name: &str,
|
||||
cancellation_token: CancellationToken,
|
||||
@@ -973,7 +980,7 @@ impl ExtensionManager {
|
||||
|
||||
let client_guard = client.lock().await;
|
||||
client_guard
|
||||
.read_resource(uri, cancellation_token)
|
||||
.read_resource(session_id, uri, cancellation_token)
|
||||
.await
|
||||
.map_err(|_| {
|
||||
ErrorData::new(
|
||||
@@ -984,7 +991,10 @@ impl ExtensionManager {
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn get_ui_resources(&self) -> Result<Vec<(String, Resource)>, ErrorData> {
|
||||
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)> = {
|
||||
@@ -999,7 +1009,7 @@ impl ExtensionManager {
|
||||
let client_guard = client.lock().await;
|
||||
|
||||
match client_guard
|
||||
.list_resources(None, CancellationToken::default())
|
||||
.list_resources(session_id, None, CancellationToken::default())
|
||||
.await
|
||||
{
|
||||
Ok(list_response) => {
|
||||
@@ -1020,6 +1030,7 @@ impl ExtensionManager {
|
||||
|
||||
async fn list_resources_from_extension(
|
||||
&self,
|
||||
session_id: &str,
|
||||
extension_name: &str,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<Vec<Content>, ErrorData> {
|
||||
@@ -1036,7 +1047,7 @@ impl ExtensionManager {
|
||||
|
||||
let client_guard = client.lock().await;
|
||||
client_guard
|
||||
.list_resources(None, cancellation_token)
|
||||
.list_resources(session_id, None, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ErrorData::new(
|
||||
@@ -1059,6 +1070,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn list_resources(
|
||||
&self,
|
||||
session_id: &str,
|
||||
params: Value,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<Vec<Content>, ErrorData> {
|
||||
@@ -1067,7 +1079,7 @@ impl ExtensionManager {
|
||||
match extension {
|
||||
Some(extension_name) => {
|
||||
// Handle single extension case
|
||||
self.list_resources_from_extension(extension_name, cancellation_token)
|
||||
self.list_resources_from_extension(session_id, extension_name, cancellation_token)
|
||||
.await
|
||||
}
|
||||
None => {
|
||||
@@ -1084,7 +1096,7 @@ impl ExtensionManager {
|
||||
.for_each(|name| {
|
||||
let token = cancellation_token.clone();
|
||||
futures.push(async move {
|
||||
self.list_resources_from_extension(&name.clone(), token)
|
||||
self.list_resources_from_extension(session_id, name.as_str(), token)
|
||||
.await
|
||||
});
|
||||
});
|
||||
@@ -1190,9 +1202,8 @@ impl ExtensionManager {
|
||||
session_id
|
||||
);
|
||||
let client_guard = client.lock().await;
|
||||
let meta = McpMeta::new(&session_id);
|
||||
client_guard
|
||||
.call_tool(&tool_name, arguments, meta, cancellation_token)
|
||||
.call_tool(&session_id, &tool_name, arguments, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| match e {
|
||||
ServiceError::McpError(error_data) => error_data,
|
||||
@@ -1210,6 +1221,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn list_prompts_from_extension(
|
||||
&self,
|
||||
session_id: &str,
|
||||
extension_name: &str,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<Vec<Prompt>, ErrorData> {
|
||||
@@ -1226,7 +1238,7 @@ impl ExtensionManager {
|
||||
|
||||
let client_guard = client.lock().await;
|
||||
client_guard
|
||||
.list_prompts(None, cancellation_token)
|
||||
.list_prompts(session_id, None, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ErrorData::new(
|
||||
@@ -1240,6 +1252,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn list_prompts(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<HashMap<String, Vec<Prompt>>, ErrorData> {
|
||||
let mut futures = FuturesUnordered::new();
|
||||
@@ -1250,7 +1263,7 @@ impl ExtensionManager {
|
||||
futures.push(async move {
|
||||
(
|
||||
extension_name.clone(),
|
||||
self.list_prompts_from_extension(extension_name.as_str(), token)
|
||||
self.list_prompts_from_extension(session_id, extension_name.as_str(), token)
|
||||
.await,
|
||||
)
|
||||
});
|
||||
@@ -1287,6 +1300,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn get_prompt(
|
||||
&self,
|
||||
session_id: &str,
|
||||
extension_name: &str,
|
||||
name: &str,
|
||||
arguments: Value,
|
||||
@@ -1299,7 +1313,7 @@ impl ExtensionManager {
|
||||
|
||||
let client_guard = client.lock().await;
|
||||
client_guard
|
||||
.get_prompt(name, arguments, cancellation_token)
|
||||
.get_prompt(session_id, name, arguments, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to get prompt: {}", e))
|
||||
}
|
||||
@@ -1470,6 +1484,7 @@ mod tests {
|
||||
|
||||
async fn list_resources(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
@@ -1478,6 +1493,7 @@ mod tests {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_uri: &str,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
@@ -1486,6 +1502,7 @@ mod tests {
|
||||
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -1516,9 +1533,9 @@ mod tests {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
name: &str,
|
||||
_arguments: Option<JsonObject>,
|
||||
_meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
match name {
|
||||
@@ -1534,6 +1551,7 @@ mod tests {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
@@ -1542,6 +1560,7 @@ mod tests {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_name: &str,
|
||||
_arguments: Value,
|
||||
_cancellation_token: CancellationToken,
|
||||
@@ -1767,7 +1786,10 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let tools = extension_manager.get_prefixed_tools(None).await.unwrap();
|
||||
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
|
||||
@@ -1794,7 +1816,10 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let tools = extension_manager.get_prefixed_tools(None).await.unwrap();
|
||||
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"));
|
||||
@@ -1956,7 +1981,10 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let tools_after_first = extension_manager.get_prefixed_tools(None).await.unwrap();
|
||||
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())
|
||||
@@ -1971,7 +1999,10 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let tools_after_second = extension_manager.get_prefixed_tools(None).await.unwrap();
|
||||
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())
|
||||
@@ -1999,14 +2030,20 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let tools_before = extension_manager.get_prefixed_tools(None).await.unwrap();
|
||||
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(None).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__")));
|
||||
@@ -2032,7 +2069,7 @@ mod tests {
|
||||
.await;
|
||||
|
||||
let tools = extension_manager
|
||||
.get_prefixed_tools_excluding("ext_a")
|
||||
.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();
|
||||
@@ -2061,7 +2098,7 @@ mod tests {
|
||||
.await;
|
||||
|
||||
let tools = extension_manager
|
||||
.get_prefixed_tools(Some("ext_a".to_string()))
|
||||
.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();
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::agents::extension::PlatformExtensionContext;
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait};
|
||||
use crate::config::get_extension_by_name;
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
@@ -224,6 +224,7 @@ impl ExtensionManagerClient {
|
||||
|
||||
async fn handle_list_resources(
|
||||
&self,
|
||||
session_id: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
) -> Result<Vec<Content>, ExtensionManagerToolError> {
|
||||
if let Some(weak_ref) = &self.context.extension_manager {
|
||||
@@ -233,7 +234,11 @@ impl ExtensionManagerClient {
|
||||
.unwrap_or(serde_json::Value::Object(serde_json::Map::new()));
|
||||
|
||||
match extension_manager
|
||||
.list_resources(params, tokio_util::sync::CancellationToken::default())
|
||||
.list_resources(
|
||||
session_id,
|
||||
params,
|
||||
tokio_util::sync::CancellationToken::default(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(content) => Ok(content),
|
||||
@@ -251,6 +256,7 @@ impl ExtensionManagerClient {
|
||||
|
||||
async fn handle_read_resource(
|
||||
&self,
|
||||
session_id: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
) -> Result<Vec<Content>, ExtensionManagerToolError> {
|
||||
if let Some(weak_ref) = &self.context.extension_manager {
|
||||
@@ -260,7 +266,11 @@ impl ExtensionManagerClient {
|
||||
.unwrap_or(serde_json::Value::Object(serde_json::Map::new()));
|
||||
|
||||
match extension_manager
|
||||
.read_resource_tool(params, tokio_util::sync::CancellationToken::default())
|
||||
.read_resource_tool(
|
||||
session_id,
|
||||
params,
|
||||
tokio_util::sync::CancellationToken::default(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(content) => Ok(content),
|
||||
@@ -390,6 +400,7 @@ impl ExtensionManagerClient {
|
||||
impl McpClientTrait for ExtensionManagerClient {
|
||||
async fn list_resources(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
@@ -398,6 +409,7 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_uri: &str,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
@@ -407,6 +419,7 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -419,9 +432,9 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let result = match name {
|
||||
@@ -429,8 +442,8 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
self.handle_search_available_extensions().await
|
||||
}
|
||||
MANAGE_EXTENSIONS_TOOL_NAME => self.handle_manage_extensions(arguments).await,
|
||||
LIST_RESOURCES_TOOL_NAME => self.handle_list_resources(arguments).await,
|
||||
READ_RESOURCE_TOOL_NAME => self.handle_read_resource(arguments).await,
|
||||
LIST_RESOURCES_TOOL_NAME => self.handle_list_resources(session_id, arguments).await,
|
||||
READ_RESOURCE_TOOL_NAME => self.handle_read_resource(session_id, arguments).await,
|
||||
_ => Err(ExtensionManagerToolError::UnknownTool {
|
||||
tool_name: name.to_string(),
|
||||
}),
|
||||
@@ -455,6 +468,7 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
@@ -463,6 +477,7 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_name: &str,
|
||||
_arguments: Value,
|
||||
_cancellation_token: CancellationToken,
|
||||
|
||||
@@ -26,47 +26,35 @@ use rmcp::{
|
||||
ClientHandler, ErrorData, Peer, RoleClient, ServiceError, ServiceExt,
|
||||
};
|
||||
use serde_json::Value;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, OnceLock},
|
||||
time::Duration,
|
||||
};
|
||||
use tokio::sync::{
|
||||
mpsc::{self, Sender},
|
||||
Mutex,
|
||||
};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use uuid::Uuid;
|
||||
|
||||
pub type BoxError = Box<dyn std::error::Error + Sync + Send>;
|
||||
|
||||
pub type Error = rmcp::ServiceError;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct McpMeta {
|
||||
pub session_id: String,
|
||||
}
|
||||
|
||||
impl McpMeta {
|
||||
pub fn new(session_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
session_id: session_id.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn inject_into_extensions(&self, extensions: Extensions) -> Extensions {
|
||||
inject_session_id_into_extensions(extensions, &self.session_id)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait McpClientTrait: Send + Sync {
|
||||
async fn list_tools(
|
||||
&self,
|
||||
session_id: &str,
|
||||
next_cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error>;
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error>;
|
||||
|
||||
@@ -74,6 +62,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn list_resources(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
@@ -82,6 +71,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_uri: &str,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
@@ -90,6 +80,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
@@ -98,6 +89,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_name: &str,
|
||||
_arguments: Value,
|
||||
_cancel_token: CancellationToken,
|
||||
@@ -117,6 +109,10 @@ pub trait McpClientTrait: Send + Sync {
|
||||
pub struct GooseClient {
|
||||
notification_handlers: Arc<Mutex<Vec<Sender<ServerNotification>>>>,
|
||||
provider: SharedProvider,
|
||||
// Single-slot because calls are serialized per MCP client; see send_request_with_session.
|
||||
current_session_id: Arc<Mutex<Option<String>>>,
|
||||
// Connection-scoped fallback for server-initiated sampling.
|
||||
client_session_id: OnceLock<String>,
|
||||
}
|
||||
|
||||
impl GooseClient {
|
||||
@@ -127,8 +123,52 @@ impl GooseClient {
|
||||
GooseClient {
|
||||
notification_handlers: handlers,
|
||||
provider,
|
||||
current_session_id: Arc::new(Mutex::new(None)),
|
||||
client_session_id: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn set_current_session_id(&self, session_id: &str) {
|
||||
let mut slot = self.current_session_id.lock().await;
|
||||
*slot = Some(session_id.to_string());
|
||||
}
|
||||
|
||||
async fn clear_current_session_id(&self) {
|
||||
let mut slot = self.current_session_id.lock().await;
|
||||
*slot = None;
|
||||
}
|
||||
|
||||
async fn current_session_id(&self) -> Option<String> {
|
||||
let slot = self.current_session_id.lock().await;
|
||||
slot.clone()
|
||||
}
|
||||
|
||||
async fn resolve_session_id(&self, extensions: &Extensions) -> String {
|
||||
// Prefer explicit MCP metadata, then the active request scope.
|
||||
if let Some(session_id) = Self::session_id_from_extensions(extensions) {
|
||||
return session_id;
|
||||
}
|
||||
if let Some(session_id) = self.current_session_id().await {
|
||||
return session_id;
|
||||
}
|
||||
// Fallback for server-initiated sampling not tied to a request session.
|
||||
self.client_session_id()
|
||||
}
|
||||
|
||||
fn client_session_id(&self) -> String {
|
||||
self.client_session_id
|
||||
.get_or_init(|| Uuid::new_v4().to_string())
|
||||
.clone()
|
||||
}
|
||||
|
||||
fn session_id_from_extensions(extensions: &Extensions) -> Option<String> {
|
||||
let meta = extensions.get::<Meta>()?;
|
||||
meta.0
|
||||
.iter()
|
||||
.find(|(key, _)| key.eq_ignore_ascii_case(SESSION_ID_HEADER))
|
||||
.and_then(|(_, value)| value.as_str())
|
||||
.map(|value| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientHandler for GooseClient {
|
||||
@@ -175,7 +215,7 @@ impl ClientHandler for GooseClient {
|
||||
async fn create_message(
|
||||
&self,
|
||||
params: CreateMessageRequestParam,
|
||||
_context: RequestContext<RoleClient>,
|
||||
context: RequestContext<RoleClient>,
|
||||
) -> Result<CreateMessageResult, ErrorData> {
|
||||
let provider = self
|
||||
.provider
|
||||
@@ -189,6 +229,9 @@ impl ClientHandler for GooseClient {
|
||||
))?
|
||||
.clone();
|
||||
|
||||
// Prefer explicit MCP metadata, then the active request scope.
|
||||
let session_id = self.resolve_session_id(&context.extensions).await;
|
||||
|
||||
let provider_ready_messages: Vec<crate::conversation::message::Message> = params
|
||||
.messages
|
||||
.iter()
|
||||
@@ -211,7 +254,7 @@ impl ClientHandler for GooseClient {
|
||||
.unwrap_or("You are a general-purpose AI agent called goose");
|
||||
|
||||
let (response, usage) = provider
|
||||
.complete(system_prompt, &provider_ready_messages, &[])
|
||||
.complete(&session_id, system_prompt, &provider_ready_messages, &[])
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ErrorData::new(
|
||||
@@ -336,19 +379,38 @@ impl McpClient {
|
||||
})
|
||||
}
|
||||
|
||||
async fn send_request(
|
||||
async fn send_request_with_session(
|
||||
&self,
|
||||
session_id: &str,
|
||||
request: ClientRequest,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ServerResult, Error> {
|
||||
let handle = self
|
||||
.client
|
||||
.lock()
|
||||
.await
|
||||
.send_cancellable_request(request, PeerRequestOptions::no_options())
|
||||
.await?;
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
// ExtensionManager serializes calls per MCP connection, so one current_session_id slot
|
||||
// is sufficient for mapping callbacks to the active request session.
|
||||
let handle = {
|
||||
let client = self.client.lock().await;
|
||||
client.service().set_current_session_id(session_id).await;
|
||||
client
|
||||
.send_cancellable_request(request, PeerRequestOptions::no_options())
|
||||
.await
|
||||
};
|
||||
|
||||
await_response(handle, self.timeout, &cancel_token).await
|
||||
let handle = match handle {
|
||||
Ok(handle) => handle,
|
||||
Err(err) => {
|
||||
let client = self.client.lock().await;
|
||||
client.service().clear_current_session_id().await;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let result = await_response(handle, self.timeout, &cancel_token).await;
|
||||
|
||||
let client = self.client.lock().await;
|
||||
client.service().clear_current_session_id().await;
|
||||
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
@@ -399,15 +461,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_resources(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListResourcesRequest(ListResourcesRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -421,17 +485,19 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn read_resource(
|
||||
&self,
|
||||
session_id: &str,
|
||||
uri: &str,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ReadResourceRequest(ReadResourceRequest {
|
||||
params: ReadResourceRequestParam {
|
||||
uri: uri.to_string(),
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -445,15 +511,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_tools(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListToolsRequest(ListToolsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -467,27 +535,26 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: name.to_string().into(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: meta.inject_into_extensions(Default::default()),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
.await?;
|
||||
let request = ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: name.to_string().into(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: Default::default(),
|
||||
});
|
||||
|
||||
match res {
|
||||
let result = self
|
||||
.send_request_with_session(session_id, request, cancel_token)
|
||||
.await;
|
||||
|
||||
match result? {
|
||||
ServerResult::CallToolResult(result) => Ok(result),
|
||||
_ => Err(ServiceError::UnexpectedResponse),
|
||||
}
|
||||
@@ -495,15 +562,17 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn list_prompts(
|
||||
&self,
|
||||
session_id: &str,
|
||||
cursor: Option<String>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::ListPromptsRequest(ListPromptsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor }),
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -517,6 +586,7 @@ impl McpClientTrait for McpClient {
|
||||
|
||||
async fn get_prompt(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Value,
|
||||
cancel_token: CancellationToken,
|
||||
@@ -526,14 +596,15 @@ impl McpClientTrait for McpClient {
|
||||
_ => None,
|
||||
};
|
||||
let res = self
|
||||
.send_request(
|
||||
.send_request_with_session(
|
||||
session_id,
|
||||
ClientRequest::GetPromptRequest(GetPromptRequest {
|
||||
params: GetPromptRequestParam {
|
||||
name: name.to_string(),
|
||||
arguments,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions: inject_current_session_id_into_extensions(Default::default()),
|
||||
extensions: Default::default(),
|
||||
}),
|
||||
cancel_token,
|
||||
)
|
||||
@@ -571,103 +642,241 @@ fn inject_session_id_into_extensions(mut extensions: Extensions, session_id: &st
|
||||
extensions
|
||||
}
|
||||
|
||||
/// Injects session ID from task-local context into Extensions._meta.
|
||||
fn inject_current_session_id_into_extensions(extensions: Extensions) -> Extensions {
|
||||
if let Some(session_id) = crate::session_context::current_session_id() {
|
||||
inject_session_id_into_extensions(extensions, &session_id)
|
||||
} else {
|
||||
extensions
|
||||
fn inject_session_id_into_request(request: ClientRequest, session_id: &str) -> ClientRequest {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListResourcesRequest(req)
|
||||
}
|
||||
ClientRequest::ReadResourceRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ReadResourceRequest(req)
|
||||
}
|
||||
ClientRequest::ListToolsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListToolsRequest(req)
|
||||
}
|
||||
ClientRequest::CallToolRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::CallToolRequest(req)
|
||||
}
|
||||
ClientRequest::ListPromptsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::ListPromptsRequest(req)
|
||||
}
|
||||
ClientRequest::GetPromptRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
ClientRequest::GetPromptRequest(req)
|
||||
}
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rmcp::model::Meta;
|
||||
use test_case::test_case;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_id_in_mcp_meta() {
|
||||
fn new_client() -> GooseClient {
|
||||
GooseClient::new(Arc::new(Mutex::new(Vec::new())), Arc::new(Mutex::new(None)))
|
||||
}
|
||||
|
||||
fn request_extensions(request: &ClientRequest) -> Option<&Extensions> {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ReadResourceRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ListToolsRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::CallToolRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::ListPromptsRequest(req) => Some(&req.extensions),
|
||||
ClientRequest::GetPromptRequest(req) => Some(&req.extensions),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn list_resources_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListResourcesRequest(ListResourcesRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn read_resource_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ReadResourceRequest(ReadResourceRequest {
|
||||
params: ReadResourceRequestParam {
|
||||
uri: "test://resource".to_string(),
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_tools_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListToolsRequest(ListToolsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn call_tool_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::CallToolRequest(CallToolRequest {
|
||||
params: CallToolRequestParam {
|
||||
task: None,
|
||||
name: "tool".to_string().into(),
|
||||
arguments: None,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_prompts_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::ListPromptsRequest(ListPromptsRequest {
|
||||
params: Some(PaginatedRequestParam { cursor: None }),
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
fn get_prompt_request(extensions: Extensions) -> ClientRequest {
|
||||
ClientRequest::GetPromptRequest(GetPromptRequest {
|
||||
params: GetPromptRequestParam {
|
||||
name: "prompt".to_string(),
|
||||
arguments: None,
|
||||
},
|
||||
method: Default::default(),
|
||||
extensions,
|
||||
})
|
||||
}
|
||||
|
||||
#[test_case(
|
||||
Some("ext-session"),
|
||||
Some("current-session"),
|
||||
"ext-session";
|
||||
"extensions win"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
Some("current-session"),
|
||||
"current-session";
|
||||
"current when no extensions"
|
||||
)]
|
||||
#[test_case(
|
||||
None,
|
||||
None,
|
||||
"client-session";
|
||||
"client fallback when no session"
|
||||
)]
|
||||
fn test_resolve_session_id(
|
||||
ext_session: Option<&str>,
|
||||
current_session: Option<&str>,
|
||||
expected: &str,
|
||||
) {
|
||||
let runtime = tokio::runtime::Runtime::new().unwrap();
|
||||
runtime.block_on(async {
|
||||
let client = new_client();
|
||||
// Make the fallback deterministic so the expected value can live in the test_case row.
|
||||
client
|
||||
.client_session_id
|
||||
.get_or_init(|| "client-session".to_string());
|
||||
if let Some(session_id) = current_session {
|
||||
let mut slot = client.current_session_id.lock().await;
|
||||
*slot = Some(session_id.to_string());
|
||||
}
|
||||
|
||||
let mut extensions = Extensions::new();
|
||||
if let Some(session_id) = ext_session {
|
||||
extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
}
|
||||
|
||||
let resolved = client.resolve_session_id(&extensions).await;
|
||||
|
||||
assert_eq!(resolved, expected);
|
||||
});
|
||||
}
|
||||
|
||||
#[test_case(list_resources_request; "list_resources")]
|
||||
#[test_case(read_resource_request; "read_resource")]
|
||||
#[test_case(list_tools_request; "list_tools")]
|
||||
#[test_case(call_tool_request; "call_tool")]
|
||||
#[test_case(list_prompts_request; "list_prompts")]
|
||||
#[test_case(get_prompt_request; "get_prompt")]
|
||||
fn test_request_injects_session(request_builder: fn(Extensions) -> ClientRequest) {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-id";
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
serde_json::from_value::<Meta>(json!({
|
||||
"Goose-Session-Id": "old-session-id",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let request = request_builder(extensions);
|
||||
let request = inject_session_id_into_request(request, session_id);
|
||||
let extensions = request_extensions(&request).expect("request should have extensions");
|
||||
let meta = extensions
|
||||
.get::<Meta>()
|
||||
.expect("extensions should contain meta");
|
||||
|
||||
assert_eq!(
|
||||
meta.0.get(SESSION_ID_HEADER),
|
||||
Some(&Value::String(session_id.to_string()))
|
||||
);
|
||||
assert_eq!(
|
||||
meta.0.get("other-key"),
|
||||
Some(&Value::String("preserve-me".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_session_id_in_mcp_meta() {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-789";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let extensions = inject_current_session_id_into_extensions(Default::default());
|
||||
let meta = extensions.get::<Meta>().unwrap();
|
||||
let extensions = inject_session_id_into_extensions(Default::default(), session_id);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
})
|
||||
.await;
|
||||
assert_eq!(
|
||||
&mcp_meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_no_session_id_in_mcp_when_absent() {
|
||||
let extensions = inject_current_session_id_into_extensions(Default::default());
|
||||
let meta = extensions.get::<Meta>();
|
||||
|
||||
assert!(meta.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_all_mcp_operations_include_session() {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "consistent-session-id";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let ext1 = inject_current_session_id_into_extensions(Default::default());
|
||||
let ext2 = inject_current_session_id_into_extensions(Default::default());
|
||||
let ext3 = inject_current_session_id_into_extensions(Default::default());
|
||||
|
||||
for ext in [&ext1, &ext2, &ext3] {
|
||||
assert_eq!(
|
||||
&ext.get::<Meta>().unwrap().0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_id_case_insensitive_replacement() {
|
||||
use rmcp::model::{Extensions, Meta};
|
||||
#[test]
|
||||
fn test_session_id_case_insensitive_replacement() {
|
||||
use rmcp::model::Extensions;
|
||||
use serde_json::{from_value, json};
|
||||
|
||||
let session_id = "new-session-id";
|
||||
crate::session_context::with_session_id(Some(session_id.to_string()), async {
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
from_value::<Meta>(json!({
|
||||
"GOOSE-SESSION-ID": "old-session-1",
|
||||
"Goose-Session-Id": "old-session-2",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
from_value::<Meta>(json!({
|
||||
"GOOSE-SESSION-ID": "old-session-1",
|
||||
"Goose-Session-Id": "old-session-2",
|
||||
"other-key": "preserve-me"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let extensions = inject_current_session_id_into_extensions(extensions);
|
||||
let meta = extensions.get::<Meta>().unwrap();
|
||||
let extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
&meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id,
|
||||
"other-key": "preserve-me"
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
})
|
||||
.await;
|
||||
assert_eq!(
|
||||
&mcp_meta.0,
|
||||
json!({
|
||||
SESSION_ID_HEADER: session_id,
|
||||
"other-key": "preserve-me"
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,8 @@ use serde_json::{json, Value};
|
||||
use tracing::debug;
|
||||
|
||||
use super::super::agents::Agent;
|
||||
use crate::agents::code_execution_extension::EXTENSION_NAME as CODE_EXECUTION_EXTENSION;
|
||||
use crate::agents::subagent_tool::SUBAGENT_TOOL_NAME;
|
||||
use crate::conversation::message::{Message, MessageContent, ToolRequest};
|
||||
use crate::conversation::Conversation;
|
||||
use crate::providers::base::{stream_from_single_message, MessageStream, Provider, ProviderUsage};
|
||||
@@ -15,11 +17,6 @@ use crate::providers::toolshim::{
|
||||
augment_message_with_tool_calls, convert_tool_messages_to_text,
|
||||
modify_system_prompt_for_tool_json, OllamaInterpreter,
|
||||
};
|
||||
|
||||
use crate::agents::code_execution_extension::EXTENSION_NAME as CODE_EXECUTION_EXTENSION;
|
||||
use crate::agents::subagent_tool::SUBAGENT_TOOL_NAME;
|
||||
#[cfg(test)]
|
||||
use crate::session::SessionType;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
fn coerce_value(s: &str, schema: &Value) -> Value {
|
||||
@@ -139,8 +136,10 @@ impl Agent {
|
||||
|
||||
// Prepare system prompt
|
||||
let extensions_info = self.extension_manager.get_extensions_info().await;
|
||||
let (extension_count, tool_count) =
|
||||
self.extension_manager.get_extension_and_tool_counts().await;
|
||||
let (extension_count, tool_count) = self
|
||||
.extension_manager
|
||||
.get_extension_and_tool_counts(session_id)
|
||||
.await;
|
||||
|
||||
// Get model name from provider
|
||||
let provider = self.provider().await?;
|
||||
@@ -175,6 +174,7 @@ impl Agent {
|
||||
/// Handles toolshim transformations if needed
|
||||
pub(crate) async fn stream_response_from_provider(
|
||||
provider: Arc<dyn Provider>,
|
||||
session_id: &str,
|
||||
system_prompt: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
@@ -201,6 +201,7 @@ impl Agent {
|
||||
debug!("WAITING_LLM_STREAM_START");
|
||||
let result = provider
|
||||
.stream(
|
||||
session_id,
|
||||
system_prompt.as_str(),
|
||||
messages_for_provider.messages(),
|
||||
&tools,
|
||||
@@ -212,6 +213,7 @@ impl Agent {
|
||||
debug!("WAITING_LLM_START");
|
||||
let complete_result = provider
|
||||
.complete(
|
||||
session_id,
|
||||
system_prompt.as_str(),
|
||||
messages_for_provider.messages(),
|
||||
&tools,
|
||||
@@ -408,6 +410,7 @@ mod tests {
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::{Provider, ProviderUsage, Usage};
|
||||
use crate::providers::errors::ProviderError;
|
||||
use crate::session::session_manager::SessionType;
|
||||
use async_trait::async_trait;
|
||||
use rmcp::object;
|
||||
|
||||
@@ -432,6 +435,7 @@ mod tests {
|
||||
|
||||
async fn complete_with_model(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_model_config: &ModelConfig,
|
||||
_system: &str,
|
||||
_messages: &[Message],
|
||||
@@ -452,7 +456,7 @@ mod tests {
|
||||
.config
|
||||
.session_manager
|
||||
.create_session(
|
||||
std::path::PathBuf::default(),
|
||||
std::env::current_dir().unwrap(),
|
||||
"test-prepare-tools".to_string(),
|
||||
SessionType::Hidden,
|
||||
)
|
||||
@@ -488,9 +492,8 @@ mod tests {
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let working_dir = std::env::current_dir()?;
|
||||
let (tools, _toolshim_tools, _system_prompt) = agent
|
||||
.prepare_tools_and_prompt(&session.id, &working_dir)
|
||||
.prepare_tools_and_prompt(&session.id, session.working_dir.as_path())
|
||||
.await?;
|
||||
|
||||
let names: Vec<String> = tools.iter().map(|t| t.name.clone().into_owned()).collect();
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::agents::extension::PlatformExtensionContext;
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait};
|
||||
use crate::config::paths::Paths;
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
@@ -263,6 +263,7 @@ impl SkillsClient {
|
||||
impl McpClientTrait for SkillsClient {
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -280,9 +281,9 @@ impl McpClientTrait for SkillsClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let content = match name {
|
||||
@@ -601,7 +602,7 @@ Content from dir3
|
||||
};
|
||||
|
||||
let result = client
|
||||
.list_tools(None, CancellationToken::new())
|
||||
.list_tools("test-session-id", None, CancellationToken::new())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(result.tools.len(), 0);
|
||||
@@ -656,7 +657,7 @@ Content
|
||||
};
|
||||
|
||||
let result = client
|
||||
.list_tools(None, CancellationToken::new())
|
||||
.list_tools("test-session-id", None, CancellationToken::new())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(result.tools.len(), 1);
|
||||
|
||||
@@ -182,13 +182,10 @@ fn get_agent_messages(
|
||||
retry_config: recipe.retry,
|
||||
};
|
||||
|
||||
let mut stream = crate::session_context::with_session_id(Some(session_id.clone()), async {
|
||||
agent
|
||||
.reply(user_message, session_config, cancellation_token)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
.map_err(|e| anyhow!("Failed to get reply from agent: {}", e))?;
|
||||
let mut stream = agent
|
||||
.reply(user_message, session_config, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| anyhow!("Failed to get reply from agent: {}", e))?;
|
||||
while let Some(message_result) = stream.next().await {
|
||||
match message_result {
|
||||
Ok(AgentEvent::Message(msg)) => conversation.push(msg),
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::agents::extension::PlatformExtensionContext;
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait, McpMeta};
|
||||
use crate::agents::mcp_client::{Error, McpClientTrait};
|
||||
use crate::session::extension_data;
|
||||
use crate::session::extension_data::ExtensionState;
|
||||
use anyhow::Result;
|
||||
@@ -158,6 +158,7 @@ impl TodoClient {
|
||||
impl McpClientTrait for TodoClient {
|
||||
async fn list_tools(
|
||||
&self,
|
||||
_session_id: &str,
|
||||
_next_cursor: Option<String>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
@@ -170,12 +171,11 @@ impl McpClientTrait for TodoClient {
|
||||
|
||||
async fn call_tool(
|
||||
&self,
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let session_id = &meta.session_id;
|
||||
let content = match name {
|
||||
"todo_write" => self.handle_write_todo(session_id, arguments).await,
|
||||
_ => Err(format!("Unknown tool: {}", name)),
|
||||
|
||||
Reference in New Issue
Block a user