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);
|
||||
|
||||
Reference in New Issue
Block a user