Sessions required (#5548)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-11-03 21:04:44 -05:00
committed by GitHub
parent 86c3e42e43
commit 5b93ee587f
31 changed files with 956 additions and 2678 deletions
+141 -45
View File
@@ -18,7 +18,47 @@ use tokio::sync::OnceCell;
use tracing::{info, warn};
use utoipa::ToSchema;
const CURRENT_SCHEMA_VERSION: i32 = 4;
const CURRENT_SCHEMA_VERSION: i32 = 5;
#[derive(Debug, Clone, Copy, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum SessionType {
User,
Scheduled,
SubAgent,
Hidden,
}
impl Default for SessionType {
fn default() -> Self {
Self::User
}
}
impl std::fmt::Display for SessionType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SessionType::User => write!(f, "user"),
SessionType::SubAgent => write!(f, "sub_agent"),
SessionType::Hidden => write!(f, "hidden"),
SessionType::Scheduled => write!(f, "scheduled"),
}
}
}
impl std::str::FromStr for SessionType {
type Err = anyhow::Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"user" => Ok(SessionType::User),
"sub_agent" => Ok(SessionType::SubAgent),
"hidden" => Ok(SessionType::Hidden),
"scheduled" => Ok(SessionType::Scheduled),
_ => Err(anyhow::anyhow!("Invalid session type: {}", s)),
}
}
}
static SESSION_STORAGE: OnceCell<Arc<SessionStorage>> = OnceCell::const_new();
@@ -27,11 +67,12 @@ pub struct Session {
pub id: String,
#[schema(value_type = String)]
pub working_dir: PathBuf,
// Allow importing session exports from before 'description' was renamed to 'name'
#[serde(alias = "description")]
pub name: String,
#[serde(default)]
pub user_set_name: bool,
#[serde(default)]
pub session_type: SessionType,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
pub extension_data: ExtensionData,
@@ -52,6 +93,7 @@ pub struct SessionUpdateBuilder {
session_id: String,
name: Option<String>,
user_set_name: Option<bool>,
session_type: Option<SessionType>,
working_dir: Option<PathBuf>,
extension_data: Option<ExtensionData>,
total_tokens: Option<Option<i32>>,
@@ -78,6 +120,7 @@ impl SessionUpdateBuilder {
session_id,
name: None,
user_set_name: None,
session_type: None,
working_dir: None,
extension_data: None,
total_tokens: None,
@@ -110,6 +153,11 @@ impl SessionUpdateBuilder {
self
}
pub fn session_type(mut self, session_type: SessionType) -> Self {
self.session_type = Some(session_type);
self
}
pub fn working_dir(mut self, working_dir: PathBuf) -> Self {
self.working_dir = Some(working_dir);
self
@@ -183,10 +231,14 @@ impl SessionManager {
.map(Arc::clone)
}
pub async fn create_session(working_dir: PathBuf, name: String) -> Result<Session> {
pub async fn create_session(
working_dir: PathBuf,
name: String,
session_type: SessionType,
) -> Result<Session> {
Self::instance()
.await?
.create_session(working_dir, name)
.create_session(working_dir, name, session_type)
.await
}
@@ -306,6 +358,7 @@ impl Default for Session {
working_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")),
name: String::new(),
user_set_name: false,
session_type: SessionType::default(),
created_at: Default::default(),
updated_at: Default::default(),
extension_data: ExtensionData::default(),
@@ -353,11 +406,17 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session {
let user_set_name = row.try_get("user_set_name").unwrap_or(false);
let session_type_str: String = row
.try_get("session_type")
.unwrap_or_else(|_| "user".to_string());
let session_type = session_type_str.parse().unwrap_or_default();
Ok(Session {
id: row.try_get("id")?,
working_dir: PathBuf::from(row.try_get::<String, _>("working_dir")?),
name,
user_set_name,
session_type,
created_at: row.try_get("created_at")?,
updated_at: row.try_get("updated_at")?,
extension_data: serde_json::from_str(&row.try_get::<String, _>("extension_data")?)
@@ -446,6 +505,7 @@ impl SessionStorage {
name TEXT NOT NULL DEFAULT '',
description TEXT NOT NULL DEFAULT '',
user_set_name BOOLEAN DEFAULT FALSE,
session_type TEXT NOT NULL DEFAULT 'user',
working_dir TEXT NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
@@ -491,6 +551,9 @@ impl SessionStorage {
sqlx::query("CREATE INDEX idx_sessions_updated ON sessions(updated_at DESC)")
.execute(&pool)
.await?;
sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)")
.execute(&pool)
.await?;
Ok(Self { pool })
}
@@ -553,31 +616,32 @@ impl SessionStorage {
sqlx::query(
r#"
INSERT INTO sessions (
id, name, user_set_name, working_dir, created_at, updated_at, extension_data,
id, name, user_set_name, session_type, working_dir, created_at, updated_at, extension_data,
total_tokens, input_tokens, output_tokens,
accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens,
schedule_id, recipe_json, user_recipe_values_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&session.id)
.bind(&session.name)
.bind(session.user_set_name)
.bind(session.working_dir.to_string_lossy().as_ref())
.bind(session.created_at)
.bind(session.updated_at)
.bind(serde_json::to_string(&session.extension_data)?)
.bind(session.total_tokens)
.bind(session.input_tokens)
.bind(session.output_tokens)
.bind(session.accumulated_total_tokens)
.bind(session.accumulated_input_tokens)
.bind(session.accumulated_output_tokens)
.bind(&session.schedule_id)
.bind(recipe_json)
.bind(user_recipe_values_json)
.execute(&self.pool)
.await?;
.bind(&session.id)
.bind(&session.name)
.bind(session.user_set_name)
.bind(session.session_type.to_string())
.bind(session.working_dir.to_string_lossy().as_ref())
.bind(session.created_at)
.bind(session.updated_at)
.bind(serde_json::to_string(&session.extension_data)?)
.bind(session.total_tokens)
.bind(session.input_tokens)
.bind(session.output_tokens)
.bind(session.accumulated_total_tokens)
.bind(session.accumulated_input_tokens)
.bind(session.accumulated_output_tokens)
.bind(&session.schedule_id)
.bind(recipe_json)
.bind(user_recipe_values_json)
.execute(&self.pool)
.await?;
if let Some(conversation) = &session.conversation {
self.replace_conversation(&session.id, conversation).await?;
@@ -687,6 +751,19 @@ impl SessionStorage {
.execute(&self.pool)
.await?;
}
5 => {
sqlx::query(
r#"
ALTER TABLE sessions ADD COLUMN session_type TEXT NOT NULL DEFAULT 'user'
"#,
)
.execute(&self.pool)
.await?;
sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)")
.execute(&self.pool)
.await?;
}
_ => {
anyhow::bail!("Unknown migration version: {}", version);
}
@@ -695,11 +772,16 @@ impl SessionStorage {
Ok(())
}
async fn create_session(&self, working_dir: PathBuf, name: String) -> Result<Session> {
async fn create_session(
&self,
working_dir: PathBuf,
name: String,
session_type: SessionType,
) -> Result<Session> {
let today = chrono::Utc::now().format("%Y%m%d").to_string();
Ok(sqlx::query_as(
r#"
INSERT INTO sessions (id, name, user_set_name, working_dir, extension_data)
INSERT INTO sessions (id, name, user_set_name, session_type, working_dir, extension_data)
VALUES (
? || '_' || CAST(COALESCE((
SELECT MAX(CAST(SUBSTR(id, 10) AS INTEGER))
@@ -709,23 +791,25 @@ impl SessionStorage {
?,
FALSE,
?,
?,
'{}'
)
RETURNING *
"#,
)
.bind(&today)
.bind(&today)
.bind(&name)
.bind(working_dir.to_string_lossy().as_ref())
.fetch_one(&self.pool)
.await?)
.bind(&today)
.bind(&today)
.bind(&name)
.bind(session_type.to_string())
.bind(working_dir.to_string_lossy().as_ref())
.fetch_one(&self.pool)
.await?)
}
async fn get_session(&self, id: &str, include_messages: bool) -> Result<Session> {
let mut session = sqlx::query_as::<_, Session>(
r#"
SELECT id, working_dir, name, description, user_set_name, created_at, updated_at, extension_data,
SELECT id, working_dir, name, description, user_set_name, session_type, created_at, updated_at, extension_data,
total_tokens, input_tokens, output_tokens,
accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens,
schedule_id, recipe_json, user_recipe_values_json
@@ -733,10 +817,10 @@ impl SessionStorage {
WHERE id = ?
"#,
)
.bind(id)
.fetch_optional(&self.pool)
.await?
.ok_or_else(|| anyhow::anyhow!("Session not found"))?;
.bind(id)
.fetch_optional(&self.pool)
.await?
.ok_or_else(|| anyhow::anyhow!("Session not found"))?;
if include_messages {
let conv = self.get_conversation(&session.id).await?;
@@ -773,6 +857,7 @@ impl SessionStorage {
add_update!(builder.name, "name");
add_update!(builder.user_set_name, "user_set_name");
add_update!(builder.session_type, "session_type");
add_update!(builder.working_dir, "working_dir");
add_update!(builder.extension_data, "extension_data");
add_update!(builder.total_tokens, "total_tokens");
@@ -803,6 +888,9 @@ impl SessionStorage {
if let Some(user_set_name) = builder.user_set_name {
q = q.bind(user_set_name);
}
if let Some(session_type) = builder.session_type {
q = q.bind(session_type.to_string());
}
if let Some(wd) = builder.working_dir {
q = q.bind(wd.to_string_lossy().to_string());
}
@@ -872,7 +960,6 @@ impl SessionStorage {
let mut message = Message::new(role, created_timestamp, content);
message.metadata = metadata;
// TODO(Douwe): make id required
message = message.with_id(format!("msg_{}_{}", session_id, idx));
messages.push(message);
}
@@ -942,20 +1029,21 @@ impl SessionStorage {
async fn list_sessions(&self) -> Result<Vec<Session>> {
sqlx::query_as::<_, Session>(
r#"
SELECT s.id, s.working_dir, s.name, s.description, s.user_set_name, s.created_at, s.updated_at, s.extension_data,
SELECT s.id, s.working_dir, s.name, s.description, s.user_set_name, s.session_type, s.created_at, s.updated_at, s.extension_data,
s.total_tokens, s.input_tokens, s.output_tokens,
s.accumulated_total_tokens, s.accumulated_input_tokens, s.accumulated_output_tokens,
s.schedule_id, s.recipe_json, s.user_recipe_values_json,
COUNT(m.id) as message_count
FROM sessions s
INNER JOIN messages m ON s.id = m.session_id
WHERE s.session_type = 'user' OR s.session_type = 'scheduled'
GROUP BY s.id
ORDER BY s.updated_at DESC
"#,
)
.fetch_all(&self.pool)
.await
.map_err(Into::into)
.fetch_all(&self.pool)
.await
.map_err(Into::into)
}
async fn delete_session(&self, session_id: &str) -> Result<()> {
@@ -1008,7 +1096,11 @@ impl SessionStorage {
let import: Session = serde_json::from_str(json)?;
let session = self
.create_session(import.working_dir.clone(), import.name.clone())
.create_session(
import.working_dir.clone(),
import.name.clone(),
import.session_type,
)
.await?;
let mut builder = SessionUpdateBuilder::new(session.id.clone())
@@ -1084,7 +1176,7 @@ mod tests {
let description = format!("Test session {}", i);
let session = session_storage
.create_session(working_dir.clone(), description)
.create_session(working_dir.clone(), description, SessionType::User)
.await
.unwrap();
@@ -1176,7 +1268,11 @@ mod tests {
let storage = Arc::new(SessionStorage::create(&db_path).await.unwrap());
let original = storage
.create_session(PathBuf::from("/tmp/test"), DESCRIPTION.to_string())
.create_session(
PathBuf::from("/tmp/test"),
DESCRIPTION.to_string(),
SessionType::User,
)
.await
.unwrap();