Session manager (#4648)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-09-26 14:56:45 -04:00
committed by GitHub
parent d0aa14a1ea
commit dc292883b7
61 changed files with 2895 additions and 5572 deletions
+5 -31
View File
@@ -364,10 +364,9 @@ mod schedule_tool_tests {
use goose::agents::platform_tools::PLATFORM_MANAGE_SCHEDULE_TOOL_NAME;
use goose::scheduler::{ScheduledJob, SchedulerError};
use goose::scheduler_trait::SchedulerTrait;
use goose::session::storage::SessionMetadata;
use goose::session::Session;
use std::sync::Arc;
// Mock scheduler for testing
struct MockScheduler {
jobs: tokio::sync::Mutex<Vec<ScheduledJob>>,
}
@@ -419,7 +418,7 @@ mod schedule_tool_tests {
&self,
_sched_id: &str,
_limit: usize,
) -> Result<Vec<(String, SessionMetadata)>, SchedulerError> {
) -> Result<Vec<(String, Session)>, SchedulerError> {
Ok(vec![])
}
@@ -853,7 +852,7 @@ mod final_output_tool_tests {
mod retry_tests {
use super::*;
use async_trait::async_trait;
use goose::agents::types::{RetryConfig, SessionConfig, SuccessCheck};
use goose::agents::types::{RetryConfig, SuccessCheck};
use goose::conversation::message::Message;
use goose::conversation::Conversation;
use goose::model::ModelConfig;
@@ -939,21 +938,10 @@ mod retry_tests {
"Valid config should pass validation"
);
let session_config = SessionConfig {
id: goose::session::Identifier::Name("test-retry".to_string()),
working_dir: std::env::current_dir()?,
schedule_id: None,
execution_mode: None,
max_turns: None,
retry_config: Some(retry_config),
};
let conversation =
Conversation::new(vec![Message::user().with_text("Complete this task")]).unwrap();
let reply_stream = agent
.reply(conversation, Some(session_config), None)
.await?;
let reply_stream = agent.reply(conversation, None, None).await?;
tokio::pin!(reply_stream);
let mut responses = Vec::new();
@@ -1051,10 +1039,8 @@ mod max_turns_tests {
use goose::model::ModelConfig;
use goose::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
use goose::providers::errors::ProviderError;
use goose::session::storage::Identifier;
use mcp_core::tool::ToolCall;
use rmcp::model::Tool;
use std::path::PathBuf;
struct MockToolProvider {}
@@ -1116,21 +1102,9 @@ mod max_turns_tests {
let provider = Arc::new(MockToolProvider::new());
agent.update_provider(provider).await?;
// The mock provider will call a non-existent tool, which will fail and allow the loop to continue
// Create session config with max_turns = 1
let session_config = goose::agents::SessionConfig {
id: Identifier::Name("test_session".to_string()),
working_dir: PathBuf::from("/tmp"),
schedule_id: None,
execution_mode: None,
max_turns: Some(1),
retry_config: None,
};
let conversation = Conversation::new(vec![Message::user().with_text("Hello")]).unwrap();
let reply_stream = agent
.reply(conversation, Some(session_config), None)
.await?;
let reply_stream = agent.reply(conversation, None, None).await?;
tokio::pin!(reply_stream);
let mut responses = Vec::new();
-79
View File
@@ -757,85 +757,6 @@ async fn test_schedule_tool_sessions_action_empty() {
assert!(calls.contains(&"sessions".to_string()));
}
#[tokio::test]
async fn test_schedule_tool_session_content_action() {
let (agent, _) = ScheduleToolTestBuilder::new().build().await;
// Test with a non-existent session
let arguments = json!({
"action": "session_content",
"session_id": "non_existent_session"
});
let result = agent
.handle_schedule_management(arguments, "test_req".to_string())
.await;
assert!(result.is_err());
if let Err(err) = result {
assert!(err
.message
.contains("Session 'non_existent_session' not found"));
} else {
panic!("Expected ExecutionError");
}
}
#[tokio::test]
async fn test_schedule_tool_session_content_action_with_real_session() {
let (agent, _) = ScheduleToolTestBuilder::new().build().await;
// Create a temporary session file in the proper session directory
let session_dir = goose::session::storage::ensure_session_dir().unwrap();
let session_id = "test_session_real";
let session_path = session_dir.join(format!("{}.jsonl", session_id));
// Create test metadata and messages
let metadata = create_test_session_metadata(2, "/tmp");
let messages = goose::conversation::Conversation::new_unvalidated(vec![
goose::conversation::message::Message::user().with_text("Hello"),
goose::conversation::message::Message::assistant().with_text("Hi there!"),
]);
// Save the session file
goose::session::storage::save_messages_with_metadata(&session_path, &metadata, &messages)
.unwrap();
// Test the session_content action
let arguments = json!({
"action": "session_content",
"session_id": session_id
});
let result = agent
.handle_schedule_management(arguments, "test_req".to_string())
.await;
// Clean up the test session file
let _ = std::fs::remove_file(&session_path);
// Verify the result
assert!(result.is_ok());
if let Ok(content) = result {
assert_eq!(content.len(), 1);
if let Some(text_content) = content[0].as_text() {
assert!(text_content
.text
.contains("Session 'test_session_real' Content:"));
assert!(text_content.text.contains("Metadata:"));
assert!(text_content.text.contains("Messages:"));
assert!(text_content.text.contains("Hello"));
assert!(text_content.text.contains("Hi there!"));
assert!(text_content.text.contains("Test session"));
} else {
panic!("Expected text content");
}
} else {
panic!("Expected successful result");
}
}
#[tokio::test]
async fn test_schedule_tool_session_content_action_missing_session_id() {
let (agent, _) = ScheduleToolTestBuilder::new().build().await;
+12 -13
View File
@@ -12,7 +12,7 @@ use tokio::sync::Mutex;
use goose::agents::Agent;
use goose::scheduler::{ScheduledJob, SchedulerError};
use goose::scheduler_trait::SchedulerTrait;
use goose::session::storage::SessionMetadata;
use goose::session::Session;
#[derive(Debug, Clone)]
pub enum MockBehavior {
@@ -30,7 +30,7 @@ pub struct ConfigurableMockScheduler {
call_log: Arc<Mutex<Vec<String>>>,
behaviors: Arc<Mutex<HashMap<String, MockBehavior>>>,
#[allow(clippy::type_complexity)]
sessions_data: Arc<Mutex<HashMap<String, Vec<(String, SessionMetadata)>>>>,
sessions_data: Arc<Mutex<HashMap<String, Vec<(String, Session)>>>>,
}
#[allow(dead_code)]
@@ -184,7 +184,7 @@ impl SchedulerTrait for ConfigurableMockScheduler {
&self,
sched_id: &str,
limit: usize,
) -> Result<Vec<(String, SessionMetadata)>, SchedulerError> {
) -> Result<Vec<(String, Session)>, SchedulerError> {
self.log_call("sessions").await;
match self.get_behavior("sessions").await {
@@ -362,11 +362,7 @@ impl ScheduleToolTestBuilder {
self
}
pub async fn with_sessions_data(
self,
job_id: &str,
sessions: Vec<(String, SessionMetadata)>,
) -> Self {
pub async fn with_sessions_data(self, job_id: &str, sessions: Vec<(String, Session)>) -> Self {
{
let mut sessions_data = self.scheduler.sessions_data.lock().await;
sessions_data.insert(job_id.to_string(), sessions);
@@ -381,13 +377,14 @@ impl ScheduleToolTestBuilder {
}
}
// Helper function to create test session metadata
pub fn create_test_session_metadata(message_count: usize, working_dir: &str) -> SessionMetadata {
SessionMetadata {
message_count,
pub fn create_test_session_metadata(message_count: usize, working_dir: &str) -> Session {
Session {
id: "".to_string(),
working_dir: PathBuf::from(working_dir),
description: "Test session".to_string(),
created_at: "".to_string(),
schedule_id: Some("test_job".to_string()),
recipe: None,
total_tokens: Some(100),
input_tokens: Some(50),
output_tokens: Some(50),
@@ -395,6 +392,8 @@ pub fn create_test_session_metadata(message_count: usize, working_dir: &str) ->
accumulated_input_tokens: Some(50),
accumulated_output_tokens: Some(50),
extension_data: Default::default(),
recipe: None,
updated_at: "".to_string(),
conversation: None,
message_count,
}
}
@@ -1,496 +0,0 @@
use futures::StreamExt;
use goose::agents::types::SessionConfig;
use goose::agents::{Agent, AgentEvent};
use goose::conversation::message::Message;
use goose::conversation::Conversation;
use goose::model::ModelConfig;
use goose::providers::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
use goose::providers::errors::ProviderError;
use goose::session;
use goose::session::storage::SessionMetadata;
use rmcp::model::Tool;
use std::sync::Arc;
use tempfile::TempDir;
use uuid::Uuid;
// Mock provider implementation for testing
struct MockProvider {
model_config: ModelConfig,
}
impl MockProvider {
fn new() -> Self {
Self {
model_config: ModelConfig::new_or_fail("mock-model"),
}
}
}
#[async_trait::async_trait]
impl Provider for MockProvider {
fn metadata() -> ProviderMetadata
where
Self: Sized,
{
ProviderMetadata::new(
"mock",
"Mock Provider",
"A mock provider for testing",
"mock-model",
vec!["mock-model"],
"https://example.com",
vec![],
)
}
async fn complete(
&self,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
// Return a simple mock response
Ok((
Message::assistant().with_text("Mock response"),
ProviderUsage::new(
"mock-model".to_string(),
Usage::new(Some(10), Some(20), Some(30)),
),
))
}
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
// Return a simple mock response
Ok((
Message::assistant().with_text("Mock response"),
ProviderUsage::new(
"mock-model".to_string(),
Usage::new(Some(10), Some(20), Some(30)),
),
))
}
fn get_model_config(&self) -> ModelConfig {
self.model_config.clone()
}
async fn stream(
&self,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<goose::providers::base::MessageStream, ProviderError> {
// Return a simple mock stream
let message = Message::assistant().with_text("Mock stream response");
let usage = ProviderUsage::new(
"mock-model".to_string(),
Usage::new(Some(10), Some(20), Some(30)),
);
Ok(goose::providers::base::stream_from_single_message(
message, usage,
))
}
fn supports_streaming(&self) -> bool {
true
}
async fn generate_session_name(
&self,
_messages: &Conversation,
) -> Result<String, ProviderError> {
Ok("Mock session description".to_string())
}
}
async fn create_test_session_dir() -> TempDir {
TempDir::new().unwrap()
}
async fn create_test_agent_with_mock_provider() -> Agent {
let agent = Agent::new();
let mock_provider = Arc::new(MockProvider::new());
agent.update_provider(mock_provider).await.unwrap();
agent
}
#[tokio::test]
async fn test_todo_add_persists_to_session() {
let temp_dir = create_test_session_dir().await;
let session_id = session::Identifier::Name(format!("test_session_{}", uuid::Uuid::new_v4()));
let agent = create_test_agent_with_mock_provider().await;
// Create a conversation with a TODO add request
let messages =
vec![Message::user().with_text("Add these tasks to my todo list: Buy milk, Call dentist")];
let conversation = Conversation::new(messages).unwrap();
let session_config = SessionConfig {
id: session_id.clone(),
working_dir: temp_dir.path().to_path_buf(),
schedule_id: None,
max_turns: Some(10),
execution_mode: Some("auto".to_string()),
retry_config: None,
};
// Process the conversation
let mut stream = agent
.reply(conversation, Some(session_config.clone()), None)
.await
.unwrap();
// Collect all events
while let Some(event) = stream.next().await {
if let Ok(_event) = event {
// Process events
}
}
// Verify TODO was persisted to session
let session_path = goose::session::storage::get_path(session_id).unwrap();
let metadata = goose::session::storage::read_metadata(&session_path).unwrap();
// Since we're using a mock provider, we can't test the actual TODO content
// but we can verify the metadata structure is correct
assert!(
metadata.extension_data.extension_states.is_empty()
|| !metadata.extension_data.extension_states.is_empty()
);
}
#[tokio::test]
async fn test_todo_list_reads_from_session() {
let temp_dir = create_test_session_dir().await;
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let agent = create_test_agent_with_mock_provider().await;
// Pre-populate session with TODO content
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let mut metadata = SessionMetadata::default();
use goose::session::extension_data::{ExtensionState, TodoState};
let todo_state = TodoState::new("- Task 1\n- Task 2\n- Task 3".to_string());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// Create a conversation requesting TODO list
let messages = vec![Message::user().with_text("Show me my todo list")];
let conversation = Conversation::new(messages).unwrap();
let session_config = SessionConfig {
id: session_id.clone(),
working_dir: temp_dir.path().to_path_buf(),
schedule_id: None,
max_turns: Some(10),
execution_mode: Some("auto".to_string()),
retry_config: None,
};
// Process the conversation
let mut stream = agent
.reply(conversation, Some(session_config), None)
.await
.unwrap();
// Collect all events
while let Some(event) = stream.next().await {
if let Ok(AgentEvent::Message(msg)) = event {
let _text = msg.as_concat_text();
// With mock provider, we can't verify the actual content
}
}
// Verify the TODO content is still in session
let metadata_after = goose::session::storage::read_metadata(&session_path).unwrap();
let todo_state_after = TodoState::from_extension_data(&metadata_after.extension_data);
assert!(todo_state_after.is_some());
assert_eq!(
todo_state_after.unwrap().content,
"- Task 1\n- Task 2\n- Task 3".to_string()
);
}
#[tokio::test]
async fn test_todo_isolation_between_sessions() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session1_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let session2_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
// Add TODO to session1
let session1_path = goose::session::storage::get_path(session1_id.clone()).unwrap();
let mut metadata1 = SessionMetadata::default();
let todo_state1 = TodoState::new("Session 1 tasks".to_string());
todo_state1
.to_extension_data(&mut metadata1.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session1_path, &metadata1)
.await
.unwrap();
// Add different TODO to session2
let session2_path = goose::session::storage::get_path(session2_id.clone()).unwrap();
let mut metadata2 = SessionMetadata::default();
let todo_state2 = TodoState::new("Session 2 tasks".to_string());
todo_state2
.to_extension_data(&mut metadata2.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session2_path, &metadata2)
.await
.unwrap();
// Verify isolation
let metadata1_read = goose::session::storage::read_metadata(&session1_path).unwrap();
let metadata2_read = goose::session::storage::read_metadata(&session2_path).unwrap();
let todo1 = TodoState::from_extension_data(&metadata1_read.extension_data).unwrap();
let todo2 = TodoState::from_extension_data(&metadata2_read.extension_data).unwrap();
assert_eq!(todo1.content, "Session 1 tasks");
assert_eq!(todo2.content, "Session 2 tasks");
}
#[tokio::test]
async fn test_todo_clear_removes_from_session() {
use goose::session::extension_data::{ExtensionState, TodoState};
let temp_dir = create_test_session_dir().await;
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let agent = create_test_agent_with_mock_provider().await;
// Pre-populate session with TODO content
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let mut metadata = SessionMetadata::default();
let todo_state = TodoState::new("- Task to clear".to_string());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// Create a conversation to clear TODO
let messages = vec![Message::user().with_text("Clear my entire todo list")];
let conversation = Conversation::new(messages).unwrap();
let session_config = SessionConfig {
id: session_id.clone(),
working_dir: temp_dir.path().to_path_buf(),
schedule_id: None,
max_turns: Some(10),
execution_mode: Some("auto".to_string()),
retry_config: None,
};
// Process the conversation
let mut stream = agent
.reply(conversation, Some(session_config), None)
.await
.unwrap();
// Consume the stream
while (stream.next().await).is_some() {}
// With mock provider, the TODO won't actually be cleared via tool calls
// but we can verify the structure is correct
let metadata_after = goose::session::storage::read_metadata(&session_path).unwrap();
let todo_state_after = TodoState::from_extension_data(&metadata_after.extension_data);
assert!(todo_state_after.is_some()); // Will still have the original content with mock
}
#[tokio::test]
async fn test_todo_persistence_across_agent_instances() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
// First agent instance adds TODO
{
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let mut metadata = SessionMetadata::default();
let todo_state = TodoState::new("Persistent task".to_string());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
}
// Second agent instance reads TODO
{
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let metadata = goose::session::storage::read_metadata(&session_path).unwrap();
let todo_state = TodoState::from_extension_data(&metadata.extension_data).unwrap();
assert_eq!(todo_state.content, "Persistent task");
}
}
#[tokio::test]
async fn test_todo_max_chars_limit() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
// Set a small limit for testing
std::env::set_var("GOOSE_TODO_MAX_CHARS", "50");
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let mut metadata = SessionMetadata::default();
// Try to set content that exceeds the limit
let long_content = "x".repeat(100);
let todo_state = TodoState::new(long_content.clone());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
// This should succeed at the storage level (storage doesn't enforce limits)
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// But when the agent tries to write through the TODO tool, it should enforce the limit
// This would be tested through the agent's dispatch_todo_tool_with_session method
// Clean up
std::env::remove_var("GOOSE_TODO_MAX_CHARS");
}
#[tokio::test]
async fn test_todo_with_special_characters() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let mut metadata = SessionMetadata::default();
// Test with various special characters
let special_content = r#"
- Task with "quotes"
- Task with 'single quotes'
- Task with emoji 🎉
- Task with unicode: 你好
- Task with newline
continuation
- Task with tab separation
"#;
let todo_state = TodoState::new(special_content.to_string());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// Read back and verify
let metadata_read = goose::session::storage::read_metadata(&session_path).unwrap();
let todo_state_read = TodoState::from_extension_data(&metadata_read.extension_data).unwrap();
assert_eq!(todo_state_read.content, special_content);
}
#[tokio::test]
async fn test_todo_concurrent_access() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
// Spawn multiple concurrent TODO operations
let mut handles = vec![];
for i in 0..5 {
let session_id_clone = session_id.clone();
let handle = tokio::spawn(async move {
let session_path = goose::session::storage::get_path(session_id_clone).unwrap();
let mut metadata = goose::session::storage::read_metadata(&session_path)
.unwrap_or_else(|_| SessionMetadata::default());
let current_content = TodoState::from_extension_data(&metadata.extension_data)
.map(|t| t.content)
.unwrap_or_default();
let new_todo = TodoState::new(format!("{}\n- Task {}", current_content, i));
new_todo
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata).await
});
handles.push(handle);
}
// Wait for all operations to complete
for handle in handles {
handle.await.unwrap().unwrap();
}
// Verify final state contains at least one task
let session_path = goose::session::storage::get_path(session_id).unwrap();
let metadata = goose::session::storage::read_metadata(&session_path).unwrap();
let todo_state = TodoState::from_extension_data(&metadata.extension_data).unwrap();
// Should contain at least one task (concurrent writes may overwrite)
assert!(todo_state.content.contains("Task"));
}
#[tokio::test]
async fn test_todo_empty_session_returns_empty() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
let metadata = goose::session::storage::read_metadata(&session_path)
.unwrap_or_else(|_| SessionMetadata::default());
let todo_state = TodoState::from_extension_data(&metadata.extension_data);
assert!(todo_state.is_none() || todo_state.unwrap().content.is_empty());
}
#[tokio::test]
async fn test_todo_update_preserves_other_metadata() {
use goose::session::extension_data::{ExtensionState, TodoState};
let session_id = session::Identifier::Name(format!("test_session_{}", Uuid::new_v4()));
let session_path = goose::session::storage::get_path(session_id.clone()).unwrap();
// Set initial metadata with various fields
let mut metadata = SessionMetadata::default();
#[allow(clippy::field_reassign_with_default)]
{
metadata.message_count = 5;
metadata.description = "Test session".to_string();
metadata.total_tokens = Some(1000);
}
let todo_state = TodoState::new("Initial TODO".to_string());
todo_state
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// Update only TODO content
let todo_state_updated = TodoState::new("Updated TODO".to_string());
todo_state_updated
.to_extension_data(&mut metadata.extension_data)
.unwrap();
goose::session::storage::update_metadata(&session_path, &metadata)
.await
.unwrap();
// Verify other fields are preserved
let metadata_read = goose::session::storage::read_metadata(&session_path).unwrap();
assert_eq!(metadata_read.message_count, 5);
assert_eq!(metadata_read.description, "Test session");
assert_eq!(metadata_read.total_tokens, Some(1000));
let todo_state_read = TodoState::from_extension_data(&metadata_read.extension_data).unwrap();
assert_eq!(todo_state_read.content, "Updated TODO");
}