feat: sessions api, view & resume prev sessions (#1453)

* Centralize session files to goose::session module
* Write session metadata and messages in jsonl
* Refactor CLI build_session to use goose::session functions
* Track session's token usage by adding optional session_id in agent.reply(...)
* NOTE: Only sessions saved through the updates goose::session functions will show up in GUI

Co-authored-by: Bradley Axen <baxen@squareup.com>
This commit is contained in:
Salman Mohammed
2025-03-03 11:49:15 -05:00
committed by GitHub
parent 68b8c5d19d
commit 9ae9045584
25 changed files with 1413 additions and 257 deletions
+11 -2
View File
@@ -4,10 +4,12 @@ use anyhow::Result;
use async_trait::async_trait;
use futures::stream::BoxStream;
use serde_json::Value;
use std::sync::Arc;
use super::extension::{ExtensionConfig, ExtensionResult};
use crate::message::Message;
use crate::providers::base::ProviderUsage;
use crate::providers::base::{Provider, ProviderUsage};
use crate::session;
use mcp_core::prompt::Prompt;
use mcp_core::protocol::GetPromptResult;
@@ -15,7 +17,11 @@ use mcp_core::protocol::GetPromptResult;
#[async_trait]
pub trait Agent: Send + Sync {
/// Create a stream that yields each message as it's generated by the agent
async fn reply(&self, messages: &[Message]) -> Result<BoxStream<'_, Result<Message>>>;
async fn reply(
&self,
messages: &[Message],
session_id: Option<session::Identifier>,
) -> Result<BoxStream<'_, Result<Message>>>;
/// Add a new MCP client to the agent
async fn add_extension(&mut self, config: ExtensionConfig) -> ExtensionResult<()>;
@@ -48,4 +54,7 @@ pub trait Agent: Send + Sync {
/// Get a prompt result with the given name and arguments
/// Returns the prompt text that would be used as user input
async fn get_prompt(&self, name: &str, arguments: Value) -> Result<GetPromptResult>;
/// Get a reference to the provider used by this agent
async fn provider(&self) -> Arc<Box<dyn Provider>>;
}
+4 -4
View File
@@ -30,7 +30,7 @@ pub struct Capabilities {
clients: HashMap<String, McpClientBox>,
instructions: HashMap<String, String>,
resource_capable_extensions: HashSet<String>,
provider: Box<dyn Provider>,
provider: Arc<Box<dyn Provider>>,
provider_usage: Mutex<Vec<ProviderUsage>>,
system_prompt_override: Option<String>,
system_prompt_extensions: Vec<String>,
@@ -90,7 +90,7 @@ impl Capabilities {
clients: HashMap::new(),
instructions: HashMap::new(),
resource_capable_extensions: HashSet::new(),
provider,
provider: Arc::new(provider),
provider_usage: Mutex::new(Vec::new()),
system_prompt_override: None,
system_prompt_extensions: Vec::new(),
@@ -202,8 +202,8 @@ impl Capabilities {
}
/// Get a reference to the provider
pub fn provider(&self) -> &dyn Provider {
&*self.provider
pub fn provider(&self) -> Arc<Box<dyn Provider>> {
Arc::clone(&self.provider)
}
/// Record provider usage
+21 -2
View File
@@ -3,6 +3,7 @@
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::Mutex;
use tracing::{debug, instrument};
@@ -12,8 +13,8 @@ use crate::agents::extension::{ExtensionConfig, ExtensionResult};
use crate::message::{Message, ToolRequest};
use crate::providers::base::Provider;
use crate::providers::base::ProviderUsage;
use crate::register_agent;
use crate::token_counter::TokenCounter;
use crate::{register_agent, session};
use anyhow::{anyhow, Result};
use indoc::indoc;
use mcp_core::prompt::Prompt;
@@ -73,6 +74,7 @@ impl Agent for ReferenceAgent {
async fn reply(
&self,
messages: &[Message],
session_id: Option<session::Identifier>,
) -> anyhow::Result<BoxStream<'_, anyhow::Result<Message>>> {
let mut messages = messages.to_vec();
let reply_span = tracing::Span::current();
@@ -143,7 +145,19 @@ impl Agent for ReferenceAgent {
&messages,
&tools,
).await?;
capabilities.record_usage(usage).await;
capabilities.record_usage(usage.clone()).await;
// record usage for the session in the session file
if let Some(session_id) = session_id.clone() {
// TODO: track session_id in langfuse tracing
let session_file = session::get_path(session_id);
let mut metadata = session::read_metadata(&session_file)?;
metadata.total_tokens = usage.usage.total_tokens;
// The message count is the number of messages in the session + 1 for the response
// The message count does not include the tool response till next iteration
metadata.message_count = messages.len() + 1;
session::update_metadata(&session_file, &metadata).await?;
}
// Yield the assistant's response
yield response.clone();
@@ -233,6 +247,11 @@ impl Agent for ReferenceAgent {
Err(anyhow!("Prompt '{}' not found", name))
}
async fn provider(&self) -> Arc<Box<dyn Provider>> {
let capabilities = self.capabilities.lock().await;
capabilities.provider()
}
}
register_agent!("reference", ReferenceAgent);
+22 -2
View File
@@ -3,6 +3,7 @@
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use tracing::{debug, error, instrument, warn};
@@ -18,6 +19,7 @@ use crate::providers::base::Provider;
use crate::providers::base::ProviderUsage;
use crate::providers::errors::ProviderError;
use crate::register_agent;
use crate::session;
use crate::token_counter::TokenCounter;
use crate::truncate::{truncate_messages, OldestFirstTruncation};
use anyhow::{anyhow, Result};
@@ -143,10 +145,11 @@ impl Agent for TruncateAgent {
}
}
#[instrument(skip(self, messages), fields(user_message))]
#[instrument(skip(self, messages, session_id), fields(user_message))]
async fn reply(
&self,
messages: &[Message],
session_id: Option<session::Identifier>,
) -> anyhow::Result<BoxStream<'_, anyhow::Result<Message>>> {
let mut messages = messages.to_vec();
let reply_span = tracing::Span::current();
@@ -223,7 +226,19 @@ impl Agent for TruncateAgent {
&tools,
).await {
Ok((response, usage)) => {
capabilities.record_usage(usage).await;
capabilities.record_usage(usage.clone()).await;
// record usage for the session in the session file
if let Some(session_id) = session_id.clone() {
// TODO: track session_id in langfuse tracing
let session_file = session::get_path(session_id);
let mut metadata = session::read_metadata(&session_file)?;
metadata.total_tokens = usage.usage.total_tokens;
// The message count is the number of messages in the session + 1 for the response
// The message count does not include the tool response till next iteration
metadata.message_count = messages.len() + 1;
session::update_metadata(&session_file, &metadata).await?;
}
// Reset truncation attempt
truncation_attempt = 0;
@@ -435,6 +450,11 @@ impl Agent for TruncateAgent {
Err(anyhow!("Prompt '{}' not found", name))
}
async fn provider(&self) -> Arc<Box<dyn Provider>> {
let capabilities = self.capabilities.lock().await;
capabilities.provider()
}
}
register_agent!("truncate", TruncateAgent);
+1
View File
@@ -4,6 +4,7 @@ pub mod message;
pub mod model;
pub mod prompt_template;
pub mod providers;
pub mod session;
pub mod token_counter;
pub mod tracing;
pub mod truncate;
+8
View File
@@ -0,0 +1,8 @@
pub mod storage;
// Re-export common session types and functions
pub use storage::{
ensure_session_dir, generate_description, generate_session_id, get_most_recent_session,
get_path, list_sessions, persist_messages, read_messages, read_metadata, update_metadata,
Identifier, SessionMetadata,
};
+411
View File
@@ -0,0 +1,411 @@
use crate::message::Message;
use crate::providers::base::Provider;
use anyhow::Result;
use chrono::Local;
use etcetera::{choose_app_strategy, AppStrategy, AppStrategyArgs};
use serde::{Deserialize, Serialize};
use std::fs::{self, File};
use std::io::{self, BufRead, Write};
use std::path::{Path, PathBuf};
use std::sync::Arc;
/// Metadata for a session, stored as the first line in the session file
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionMetadata {
/// A short description of the session, typically 3 words or less
pub description: String,
/// Number of messages in the session
pub message_count: usize,
/// The total number of tokens used in the session. Retrieved from the provider's last usage.
pub total_tokens: Option<i32>,
}
impl SessionMetadata {
pub fn new() -> Self {
Self {
description: String::new(),
message_count: 0,
total_tokens: None,
}
}
}
impl Default for SessionMetadata {
fn default() -> Self {
Self::new()
}
}
// The single app name used for all Goose applications
const APP_NAME: &str = "goose";
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Identifier {
Name(String),
Path(PathBuf),
}
pub fn get_path(id: Identifier) -> PathBuf {
match id {
Identifier::Name(name) => {
let session_dir = ensure_session_dir().expect("Failed to create session directory");
session_dir.join(format!("{}.jsonl", name))
}
Identifier::Path(path) => path,
}
}
/// Ensure the session directory exists and return its path
pub fn ensure_session_dir() -> Result<PathBuf> {
let app_strategy = AppStrategyArgs {
top_level_domain: "Block".to_string(),
author: "Block".to_string(),
app_name: APP_NAME.to_string(),
};
let data_dir = choose_app_strategy(app_strategy)
.expect("goose requires a home dir")
.data_dir()
.join("sessions");
if !data_dir.exists() {
fs::create_dir_all(&data_dir)?;
}
Ok(data_dir)
}
/// Get the path to the most recently modified session file
pub fn get_most_recent_session() -> Result<PathBuf> {
let session_dir = ensure_session_dir()?;
let mut entries = fs::read_dir(&session_dir)?
.filter_map(|entry| entry.ok())
.filter(|entry| entry.path().extension().is_some_and(|ext| ext == "jsonl"))
.collect::<Vec<_>>();
if entries.is_empty() {
return Err(anyhow::anyhow!("No session files found"));
}
// Sort by modification time, most recent first
entries.sort_by(|a, b| {
b.metadata()
.and_then(|m| m.modified())
.unwrap_or(std::time::SystemTime::UNIX_EPOCH)
.cmp(
&a.metadata()
.and_then(|m| m.modified())
.unwrap_or(std::time::SystemTime::UNIX_EPOCH),
)
});
Ok(entries[0].path())
}
/// List all available session files
pub fn list_sessions() -> Result<Vec<(String, PathBuf)>> {
let session_dir = ensure_session_dir()?;
let entries = fs::read_dir(&session_dir)?
.filter_map(|entry| {
let entry = entry.ok()?;
let path = entry.path();
if path.extension().is_some_and(|ext| ext == "jsonl") {
let name = path.file_stem()?.to_string_lossy().to_string();
Some((name, path))
} else {
None
}
})
.collect::<Vec<_>>();
Ok(entries)
}
/// Generate a session ID using timestamp format (yyyymmdd_hhmmss)
pub fn generate_session_id() -> String {
Local::now().format("%Y%m%d_%H%M%S").to_string()
}
/// Read messages from a session file
///
/// Creates the file if it doesn't exist, reads and deserializes all messages if it does.
/// The first line of the file is expected to be metadata, and the rest are messages.
pub fn read_messages(session_file: &Path) -> Result<Vec<Message>> {
let file = fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(session_file)?;
let reader = io::BufReader::new(file);
let mut lines = reader.lines();
let mut messages = Vec::new();
// Read the first line as metadata or create default if empty/missing
if let Some(line) = lines.next() {
let line = line?;
// Try to parse as metadata, but if it fails, treat it as a message
if let Ok(_metadata) = serde_json::from_str::<SessionMetadata>(&line) {
// Metadata successfully parsed, continue with the rest of the lines as messages
} else {
// This is not metadata, it's a message
messages.push(serde_json::from_str::<Message>(&line)?);
}
}
// Read the rest of the lines as messages
for line in lines {
messages.push(serde_json::from_str::<Message>(&line?)?);
}
Ok(messages)
}
/// Read session metadata from a session file
///
/// Returns default empty metadata if the file doesn't exist or has no metadata.
pub fn read_metadata(session_file: &Path) -> Result<SessionMetadata> {
if !session_file.exists() {
return Ok(SessionMetadata::new());
}
let file = fs::File::open(session_file)?;
let mut reader = io::BufReader::new(file);
let mut first_line = String::new();
// Read just the first line
if reader.read_line(&mut first_line)? > 0 {
// Try to parse as metadata
match serde_json::from_str::<SessionMetadata>(&first_line) {
Ok(metadata) => Ok(metadata),
Err(_) => {
// If the first line isn't metadata, return default
Ok(SessionMetadata::new())
}
}
} else {
// Empty file, return default
Ok(SessionMetadata::new())
}
}
/// Write messages to a session file with metadata
///
/// Overwrites the file with metadata as the first line, followed by all messages in JSONL format.
/// If a provider is supplied, it will automatically generate a description when appropriate.
pub async fn persist_messages(
session_file: &Path,
messages: &[Message],
provider: Option<Arc<Box<dyn Provider>>>,
) -> Result<()> {
// Read existing metadata
let mut metadata = read_metadata(session_file)?;
// Count user messages
let user_message_count = messages
.iter()
.filter(|m| m.role == mcp_core::role::Role::User)
.filter(|m| !m.as_concat_text().trim().is_empty())
.count();
// Check if we need to update the description (after 1st or 3rd user message)
if let Some(provider) = provider {
if user_message_count < 4 {
// Generate description
let mut description_prompt = "Based on the conversation so far, provide a concise header for this session in 4 words or less. This will be used for finding the session later in a UI with limited space - reply *ONLY* with the header. Avoid filler words such as help, summary, exchange, request etc that do not help distinguish different conversations.".to_string();
// get context from messages so far
let context: Vec<String> = messages.iter().map(|m| m.as_concat_text()).collect();
if !context.is_empty() {
description_prompt = format!(
"Here are the first few user messages:\n{}\n\n{}",
context.join("\n"),
description_prompt
);
}
// Generate the description
let message = Message::user().with_text(&description_prompt);
match provider
.complete(
"Reply with only a description in four words or less.",
&[message],
&[],
)
.await
{
Ok((response, _)) => {
metadata.description = response.as_concat_text();
}
Err(e) => {
tracing::error!("Failed to generate session description: {:?}", e);
}
}
}
}
// Write the file with metadata and messages
save_messages_with_metadata(session_file, &metadata, messages)
}
/// Write messages to a session file with the provided metadata
///
/// Overwrites the file with metadata as the first line, followed by all messages in JSONL format.
pub fn save_messages_with_metadata(
session_file: &Path,
metadata: &SessionMetadata,
messages: &[Message],
) -> Result<()> {
let file = File::create(session_file).expect("The path specified does not exist");
let mut writer = io::BufWriter::new(file);
// Write metadata as the first line
serde_json::to_writer(&mut writer, &metadata)?;
writeln!(writer)?;
// Write all messages
for message in messages {
serde_json::to_writer(&mut writer, &message)?;
writeln!(writer)?;
}
writer.flush()?;
Ok(())
}
/// Generate a description for the session using the provider
///
/// This function is called when appropriate to generate a short description
/// of the session based on the conversation history.
pub async fn generate_description(
session_file: &Path,
messages: &[Message],
provider: &dyn Provider,
) -> Result<()> {
// Create a special message asking for a 3-word description
let mut description_prompt = "Based on the conversation so far, provide a concise description of this session in 4 words or less. This will be used for finding the session later in a UI with limited space - reply *ONLY* with the description".to_string();
// get context from messages so far
let context: Vec<String> = messages
.iter()
.filter(|m| m.role == mcp_core::role::Role::User)
.take(3) // Use up to first 3 user messages for context
.map(|m| m.as_concat_text())
.collect();
if !context.is_empty() {
description_prompt = format!(
"Here are the first few user messages:\n{}\n\n{}",
context.join("\n"),
description_prompt
);
}
// Generate the description
let message = Message::user().with_text(&description_prompt);
let result = provider
.complete(
"Reply with only a description in four words or less",
&[message],
&[],
)
.await?;
let description = result.0.as_concat_text();
// Read current metadata
let mut metadata = read_metadata(session_file)?;
// Update description
metadata.description = description;
// Update the file with the new metadata and existing messages
update_metadata(session_file, &metadata).await?;
Ok(())
}
/// Update only the metadata in a session file, preserving all messages
pub async fn update_metadata(session_file: &Path, metadata: &SessionMetadata) -> Result<()> {
// Read all messages from the file
let messages = read_messages(session_file)?;
// Rewrite the file with the new metadata and existing messages
save_messages_with_metadata(session_file, metadata, &messages)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::MessageContent;
use tempfile::tempdir;
#[tokio::test]
async fn test_read_write_messages() -> Result<()> {
let dir = tempdir()?;
let file_path = dir.path().join("test.jsonl");
// Create some test messages
let messages = vec![
Message::user().with_text("Hello"),
Message::assistant().with_text("Hi there"),
];
// Write messages
persist_messages(&file_path, &messages, None).await?;
// Read them back
let read_messages = read_messages(&file_path)?;
// Compare
assert_eq!(messages.len(), read_messages.len());
for (orig, read) in messages.iter().zip(read_messages.iter()) {
assert_eq!(orig.role, read.role);
assert_eq!(orig.content.len(), read.content.len());
// Compare first text content
if let (Some(MessageContent::Text(orig_text)), Some(MessageContent::Text(read_text))) =
(orig.content.first(), read.content.first())
{
assert_eq!(orig_text.text, read_text.text);
} else {
panic!("Messages don't match expected structure");
}
}
Ok(())
}
#[test]
fn test_empty_file() -> Result<()> {
let dir = tempdir()?;
let file_path = dir.path().join("empty.jsonl");
// Reading an empty file should return empty vec
let messages = read_messages(&file_path)?;
assert!(messages.is_empty());
Ok(())
}
#[test]
fn test_generate_session_id() {
let id = generate_session_id();
// Check that it follows the timestamp format (yyyymmdd_hhmmss)
assert_eq!(id.len(), 15); // 8 chars for date + 1 for underscore + 6 for time
assert!(id.contains('_'));
// Split by underscore and check parts
let parts: Vec<&str> = id.split('_').collect();
assert_eq!(parts.len(), 2);
// Date part should be 8 digits
assert_eq!(parts[0].len(), 8);
// Time part should be 6 digits
assert_eq!(parts[1].len(), 6);
}
}