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
+9 -5
View File
@@ -355,6 +355,7 @@ mod tests {
impl Provider for MockToolProvider {
async fn complete(
&self,
_session_id: &str,
_system_prompt: &str,
_messages: &[Message],
_tools: &[Tool],
@@ -376,12 +377,14 @@ mod tests {
async fn complete_with_model(
&self,
session_id: &str,
_model_config: &ModelConfig,
system_prompt: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
self.complete(system_prompt, messages, tools).await
self.complete(session_id, system_prompt, messages, tools)
.await
}
fn get_model_config(&self) -> ModelConfig {
@@ -496,7 +499,7 @@ mod tests {
use goose::config::GooseMode;
use goose::session::SessionManager;
async fn setup_agent_with_extension_manager() -> Agent {
async fn setup_agent_with_extension_manager() -> (Agent, String) {
// Add the TODO extension to the config so it can be discovered by search_available_extensions
// Set it as disabled initially so tests can enable it
let todo_extension_entry = ExtensionEntry {
@@ -515,6 +518,7 @@ mod tests {
// Create agent with session_id from the start
let temp_dir = tempfile::tempdir().unwrap();
let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf()));
let session_id = "test-session-id".to_string();
let config = AgentConfig::new(
session_manager,
PermissionManager::instance(),
@@ -536,13 +540,13 @@ mod tests {
.add_extension(ext_config)
.await
.expect("Failed to add extension manager");
agent
(agent, session_id)
}
#[tokio::test]
async fn test_extension_manager_tools_available() {
let agent = setup_agent_with_extension_manager().await;
let tools = agent.list_tools("test-session-id", None).await;
let (agent, session_id) = setup_agent_with_extension_manager().await;
let tools = agent.list_tools(&session_id, None).await;
// Note: Tool names are prefixed with the normalized extension name "extensionmanager"
// not the display name "Extension Manager"
@@ -59,6 +59,7 @@ impl Provider for MockProvider {
async fn complete_with_model(
&self,
_session_id: &str,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
+16 -2
View File
@@ -100,7 +100,12 @@ impl ProviderTester {
let (response, _) = self
.provider
.complete("You are a helpful assistant.", &[message], &[])
.complete(
"test-session-id",
"You are a helpful assistant.",
&[message],
&[],
)
.await?;
assert_eq!(
@@ -138,6 +143,7 @@ impl ProviderTester {
let (response1, _) = self
.provider
.complete(
"test-session-id",
"You are a helpful weather assistant.",
std::slice::from_ref(&message),
std::slice::from_ref(&weather_tool),
@@ -186,6 +192,7 @@ impl ProviderTester {
let (response2, _) = self
.provider
.complete(
"test-session-id",
"You are a helpful weather assistant.",
&[message, response1, weather],
&[weather_tool],
@@ -228,7 +235,12 @@ impl ProviderTester {
let result = self
.provider
.complete("You are a helpful assistant.", &messages, &[])
.complete(
"test-session-id",
"You are a helpful assistant.",
&messages,
&[],
)
.await;
println!("=== {}::context_length_exceeded_error ===", self.name);
@@ -286,6 +298,7 @@ impl ProviderTester {
let result = self
.provider
.complete(
"test-session-id",
"You are a helpful assistant. Describe what you see in the image briefly.",
&[message_with_image],
&[],
@@ -338,6 +351,7 @@ impl ProviderTester {
let result2 = self
.provider
.complete(
"test-session-id",
"You are a helpful assistant.",
&[user_message, tool_request, tool_response],
&[screenshot_tool],
@@ -3,7 +3,6 @@ use goose::model::ModelConfig;
use goose::providers::api_client::{ApiClient, AuthMethod};
use goose::providers::base::Provider;
use goose::providers::openai::OpenAiProvider;
use goose::session_context;
use goose::session_context::SESSION_ID_HEADER;
use serde_json::json;
use std::sync::Arc;
@@ -81,30 +80,19 @@ async fn setup_mock_server() -> (MockServer, HeaderCapture, Box<dyn Provider>) {
(mock_server, capture, provider)
}
async fn make_request(provider: &dyn Provider, session_id: Option<&str>) {
async fn make_request(provider: &dyn Provider, session_id: &str) {
let message = Message::user().with_text("test message");
let request_fn = async {
provider
.complete("You are a helpful assistant.", &[message], &[])
.await
.unwrap()
};
match session_id {
Some(id) => {
session_context::with_session_id(Some(id.to_string()), request_fn).await;
}
None => {
request_fn.await;
}
}
let _ = provider
.complete(session_id, "You are a helpful assistant.", &[message], &[])
.await
.unwrap();
}
#[tokio::test]
async fn test_session_id_propagation_to_llm() {
let (_, capture, provider) = setup_mock_server().await;
make_request(provider.as_ref(), Some("integration-test-session-123")).await;
make_request(provider.as_ref(), "integration-test-session-123").await;
assert_eq!(
capture.get_captured(),
@@ -113,26 +101,29 @@ async fn test_session_id_propagation_to_llm() {
}
#[tokio::test]
async fn test_no_session_id_when_absent() {
async fn test_session_id_always_present() {
let (_, capture, provider) = setup_mock_server().await;
make_request(provider.as_ref(), None).await;
make_request(provider.as_ref(), "test-session-id").await;
assert_eq!(capture.get_captured(), vec![None]);
assert_eq!(
capture.get_captured(),
vec![Some("test-session-id".to_string())]
);
}
#[tokio::test]
async fn test_session_id_matches_across_calls() {
let (_, capture, provider) = setup_mock_server().await;
let test_session_id = "consistent-session-456";
make_request(provider.as_ref(), Some(test_session_id)).await;
make_request(provider.as_ref(), Some(test_session_id)).await;
make_request(provider.as_ref(), Some(test_session_id)).await;
let session_id = "consistent-session-456";
make_request(provider.as_ref(), session_id).await;
make_request(provider.as_ref(), session_id).await;
make_request(provider.as_ref(), session_id).await;
assert_eq!(
capture.get_captured(),
vec![Some(test_session_id.to_string()); 3]
vec![Some(session_id.to_string()); 3]
);
}
@@ -142,8 +133,8 @@ async fn test_different_sessions_have_different_ids() {
let session_id_1 = "session-one";
let session_id_2 = "session-two";
make_request(provider.as_ref(), Some(session_id_1)).await;
make_request(provider.as_ref(), Some(session_id_2)).await;
make_request(provider.as_ref(), session_id_1).await;
make_request(provider.as_ref(), session_id_2).await;
assert_eq!(
capture.get_captured(),
+27 -4
View File
@@ -29,6 +29,7 @@ mod tetrate_streaming_tests {
let mut stream = provider
.stream(
"test-session-id",
"You are a helpful assistant that counts numbers.",
&messages,
&[],
@@ -100,6 +101,7 @@ mod tetrate_streaming_tests {
let mut stream = provider
.stream(
"test-session-id",
"You are a helpful assistant with access to weather information.",
&messages,
&[weather_tool],
@@ -146,7 +148,12 @@ mod tetrate_streaming_tests {
let messages = vec![Message::user().with_text("")];
let mut stream = provider
.stream("You are a helpful assistant.", &messages, &[])
.stream(
"test-session-id",
"You are a helpful assistant.",
&messages,
&[],
)
.await?;
let mut chunk_count = 0;
@@ -177,6 +184,7 @@ mod tetrate_streaming_tests {
let mut stream = provider
.stream(
"test-session-id",
"You are a helpful assistant that writes detailed essays.",
&messages,
&[],
@@ -235,7 +243,12 @@ mod tetrate_streaming_tests {
let messages = vec![Message::user().with_text("Hello")];
let result = provider
.stream("You are a helpful assistant.", &messages, &[])
.stream(
"test-session-id",
"You are a helpful assistant.",
&messages,
&[],
)
.await;
// We expect this to fail with an authentication error
@@ -258,11 +271,21 @@ mod tetrate_streaming_tests {
let messages2 = vec![Message::user().with_text("Say 'Stream 2'")];
let stream1 = provider
.stream("You are a helpful assistant.", &messages1, &[])
.stream(
"test-session-id",
"You are a helpful assistant.",
&messages1,
&[],
)
.await?;
let stream2 = provider
.stream("You are a helpful assistant.", &messages2, &[])
.stream(
"test-session-id",
"You are a helpful assistant.",
&messages2,
&[],
)
.await?;
// Process both streams concurrently