fix(goose): propagate session_id across providers and MCP (#6584)

Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
Adrian Cole
2026-01-22 09:28:56 +09:00
committed by GitHub
parent f3bae7ea7a
commit 67de49abbb
61 changed files with 1457 additions and 616 deletions
+52 -18
View File
@@ -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());
}
+8 -4
View File
@@ -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);
+70 -33
View File
@@ -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,
+344 -135
View File
@@ -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()
);
}
}
+13 -10
View File
@@ -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();
+5 -4
View File
@@ -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);
+4 -7
View File
@@ -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),
+3 -3
View File
@@ -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)),