fix(goose): propagate session_id across providers and MCP (#6584)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -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],
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user