feat: store working directory for sessions (#1559)

This commit is contained in:
Salman Mohammed
2025-03-07 11:12:57 -05:00
committed by GitHub
parent 2d0cd8e245
commit 32f20cd690
20 changed files with 334 additions and 144 deletions
+19 -4
View File
@@ -8,14 +8,18 @@ use axum::{
};
use bytes::Bytes;
use futures::{stream::StreamExt, Stream};
use goose::message::{Message, MessageContent};
use goose::session;
use goose::{
agents::SessionConfig,
message::{Message, MessageContent},
};
use mcp_core::role::Role;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::{
convert::Infallible,
path::PathBuf,
pin::Pin,
task::{Context, Poll},
time::Duration,
@@ -29,6 +33,7 @@ use tokio_stream::wrappers::ReceiverStream;
struct ChatRequest {
messages: Vec<Message>,
session_id: Option<String>,
session_working_dir: String,
}
// Custom SSE response type for streaming messages
@@ -108,8 +113,8 @@ async fn handler(
let (tx, rx) = mpsc::channel(100);
let stream = ReceiverStream::new(rx);
// Get messages directly from the request
let messages = request.messages;
let session_working_dir = request.session_working_dir;
// Generate a new session ID if not provided in the request
let session_id = request
@@ -149,7 +154,10 @@ async fn handler(
let mut stream = match agent
.reply(
&messages,
Some(session::Identifier::Name(session_id.clone())),
Some(SessionConfig {
id: session::Identifier::Name(session_id.clone()),
working_dir: PathBuf::from(session_working_dir),
}),
)
.await
{
@@ -246,6 +254,7 @@ async fn handler(
struct AskRequest {
prompt: String,
session_id: Option<String>,
session_working_dir: String,
}
#[derive(Debug, Serialize)]
@@ -269,6 +278,8 @@ async fn ask_handler(
return Err(StatusCode::UNAUTHORIZED);
}
let session_working_dir = request.session_working_dir;
// Generate a new session ID if not provided in the request
let session_id = request
.session_id
@@ -289,7 +300,10 @@ async fn ask_handler(
let mut stream = match agent
.reply(
&messages,
Some(session::Identifier::Name(session_id.clone())),
Some(SessionConfig {
id: session::Identifier::Name(session_id.clone()),
working_dir: PathBuf::from(session_working_dir),
}),
)
.await
{
@@ -464,6 +478,7 @@ mod tests {
serde_json::to_string(&AskRequest {
prompt: "test prompt".to_string(),
session_id: Some("test-session".to_string()),
session_working_dir: "test-working-dir".to_string(),
})
.unwrap(),
))