feat(session_db): add durable agent session database (#3259)

This commit is contained in:
Steven Enamakel
2026-06-02 20:07:56 -07:00
committed by GitHub
parent 928f516d86
commit 468ca7bb6a
8 changed files with 1893 additions and 0 deletions
+4
View File
@@ -280,6 +280,8 @@ fn build_registered_controllers() -> Vec<RegisteredController> {
controllers.extend(crate::openhuman::whatsapp_data::all_whatsapp_data_registered_controllers());
// Mobile device pairing and management
controllers.extend(crate::openhuman::devices::all_devices_registered_controllers());
// Durable agent session database — queryable index over transcripts, lineage, tool calls
controllers.extend(crate::openhuman::session_db::all_session_db_registered_controllers());
controllers
}
@@ -399,6 +401,8 @@ fn build_declared_controller_schemas() -> Vec<ControllerSchema> {
schemas.extend(crate::openhuman::whatsapp_data::all_whatsapp_data_controller_schemas());
// Mobile device pairing and management
schemas.extend(crate::openhuman::devices::all_devices_controller_schemas());
// Durable agent session database
schemas.extend(crate::openhuman::session_db::all_session_db_controller_schemas());
schemas
}
+1
View File
@@ -92,6 +92,7 @@ pub mod screen_intelligence;
pub mod search;
pub mod security;
pub mod service;
pub mod session_db;
pub mod skills;
pub mod socket;
pub mod startup;
+28
View File
@@ -0,0 +1,28 @@
//! Durable agent session database.
//!
//! SQLite-backed store (WAL + FTS5) for sessions, messages, tool calls,
//! cost metadata, and parent/child lineage. Complements the existing
//! `session_raw/*.jsonl` transcript files — those remain the source of
//! truth for KV-cache resume; this module provides queryable indexing,
//! cross-session search, and orchestration recovery.
//!
//! Database path: `{workspace}/session_db/sessions.db`.
mod ops;
mod schemas;
mod store;
pub mod types;
pub use ops::{
get_session, import_transcript, list_sessions, record_message, record_session_end,
record_session_start, record_tool_call, search_sessions,
};
pub use schemas::{
all_controller_schemas as all_session_db_controller_schemas,
all_registered_controllers as all_session_db_registered_controllers,
};
pub use store::with_connection;
pub use types::{
SessionMessage, SessionRecord, SessionSearchParams, SessionSearchResult, SessionStatus,
SessionToolCall,
};
+696
View File
@@ -0,0 +1,696 @@
use anyhow::{Context, Result};
use chrono::{DateTime, Utc};
use rusqlite::{params, Connection};
use crate::openhuman::config::Config;
use super::store::with_connection;
use super::types::{
SessionMessage, SessionRecord, SessionSearchParams, SessionSearchResult, SessionStatus,
SessionToolCall,
};
const MAX_TOOL_OUTPUT_BYTES: usize = 32 * 1024;
pub fn record_session_start(
config: &Config,
id: &str,
agent_definition_id: &str,
agent_definition_name: &str,
session_key: &str,
parent_session_id: Option<&str>,
thread_id: Option<&str>,
source_channel: Option<&str>,
model: Option<&str>,
transcript_path: Option<&str>,
) -> Result<SessionRecord> {
let now = Utc::now();
log::debug!(
"[session_db] record_session_start id={id} agent={agent_definition_id} \
parent={} thread={} channel={}",
parent_session_id.unwrap_or("-"),
thread_id.unwrap_or("-"),
source_channel.unwrap_or("-"),
);
with_connection(config, |conn| {
conn.execute(
"INSERT INTO sessions (
id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
transcript_path, started_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, 'running', ?8, ?9, ?10)",
params![
id,
agent_definition_id,
agent_definition_name,
session_key,
parent_session_id,
thread_id,
source_channel,
model,
transcript_path,
now.to_rfc3339(),
],
)
.context("failed to insert session")?;
index_fts_session(conn, id, agent_definition_name)?;
Ok(())
})?;
get_session(config, id)
}
pub fn record_session_end(
config: &Config,
id: &str,
status: SessionStatus,
turn_count: u32,
input_tokens: u64,
output_tokens: u64,
cached_input_tokens: u64,
cost_usd: f64,
) -> Result<SessionRecord> {
let now = Utc::now();
log::debug!(
"[session_db] record_session_end id={id} status={} turns={turn_count} \
tokens_in={input_tokens} tokens_out={output_tokens} cost=${cost_usd:.6}",
status.as_str(),
);
with_connection(config, |conn| {
conn.execute(
"UPDATE sessions SET
status = ?1, turn_count = ?2, input_tokens = ?3,
output_tokens = ?4, cached_input_tokens = ?5,
cost_usd = ?6, ended_at = ?7
WHERE id = ?8",
params![
status.as_str(),
turn_count,
input_tokens as i64,
output_tokens as i64,
cached_input_tokens as i64,
cost_usd,
now.to_rfc3339(),
id,
],
)
.context("failed to update session end")?;
Ok(())
})?;
get_session(config, id)
}
pub fn record_message(
config: &Config,
session_id: &str,
role: &str,
content: &str,
model: Option<&str>,
input_tokens: Option<u64>,
output_tokens: Option<u64>,
cost_usd: Option<f64>,
) -> Result<i64> {
let now = Utc::now();
log::trace!(
"[session_db] record_message session={session_id} role={role} len={}",
content.len()
);
with_connection(config, |conn| {
conn.execute(
"INSERT INTO session_messages (
session_id, role, content, model,
input_tokens, output_tokens, cost_usd, created_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
params![
session_id,
role,
content,
model,
input_tokens.map(|v| v as i64),
output_tokens.map(|v| v as i64),
cost_usd,
now.to_rfc3339(),
],
)
.context("failed to insert session message")?;
let msg_id = conn.last_insert_rowid();
index_fts_content(conn, session_id, content)?;
Ok(msg_id)
})
}
pub fn record_tool_call(
config: &Config,
session_id: &str,
message_id: Option<i64>,
tool_name: &str,
tool_input: Option<&str>,
tool_output: Option<&str>,
status: &str,
duration_ms: Option<i64>,
) -> Result<i64> {
let now = Utc::now();
log::trace!(
"[session_db] record_tool_call session={session_id} tool={tool_name} status={status}"
);
let bounded_output = tool_output.map(|o| {
if o.len() <= MAX_TOOL_OUTPUT_BYTES {
o.to_string()
} else {
let mut cutoff = MAX_TOOL_OUTPUT_BYTES;
while cutoff > 0 && !o.is_char_boundary(cutoff) {
cutoff -= 1;
}
let mut truncated = o[..cutoff].to_string();
truncated.push_str("\n...[truncated]");
truncated
}
});
with_connection(config, |conn| {
conn.execute(
"INSERT INTO session_tool_calls (
session_id, message_id, tool_name, tool_input,
tool_output, status, duration_ms, created_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
params![
session_id,
message_id,
tool_name,
tool_input,
bounded_output,
status,
duration_ms,
now.to_rfc3339(),
],
)
.context("failed to insert tool call")?;
index_fts_tool(conn, session_id, tool_name)?;
Ok(conn.last_insert_rowid())
})
}
pub fn get_session(config: &Config, id: &str) -> Result<SessionRecord> {
with_connection(config, |conn| {
let mut stmt = conn.prepare(
"SELECT id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
FROM sessions WHERE id = ?1",
)?;
let mut rows = stmt.query(params![id])?;
if let Some(row) = rows.next()? {
map_session_row(row).map_err(Into::into)
} else {
anyhow::bail!("session '{id}' not found")
}
})
}
pub fn list_sessions(
config: &Config,
limit: Option<u32>,
offset: Option<u32>,
status: Option<&str>,
parent_id: Option<&str>,
) -> Result<SessionSearchResult> {
log::debug!(
"[session_db] list_sessions limit={} offset={} status={} parent={}",
limit.unwrap_or(50),
offset.unwrap_or(0),
status.unwrap_or("-"),
parent_id.unwrap_or("-"),
);
with_connection(config, |conn| {
let mut where_clauses: Vec<String> = Vec::new();
let mut param_values: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
if let Some(s) = status {
param_values.push(Box::new(s.to_string()));
where_clauses.push(format!("status = ?{}", param_values.len()));
}
if let Some(p) = parent_id {
param_values.push(Box::new(p.to_string()));
where_clauses.push(format!("parent_session_id = ?{}", param_values.len()));
}
let where_sql = if where_clauses.is_empty() {
String::new()
} else {
format!("WHERE {}", where_clauses.join(" AND "))
};
let lim = limit.unwrap_or(50).min(500) as i64;
let off = offset.unwrap_or(0) as i64;
let count_sql = format!("SELECT COUNT(*) FROM sessions {where_sql}");
let total: u64 = {
let mut stmt = conn.prepare(&count_sql)?;
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|b| b.as_ref()).collect();
stmt.query_row(params_ref.as_slice(), |r| r.get::<_, i64>(0))? as u64
};
param_values.push(Box::new(lim));
let lim_idx = param_values.len();
param_values.push(Box::new(off));
let off_idx = param_values.len();
let query_sql = format!(
"SELECT id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
FROM sessions {where_sql}
ORDER BY started_at DESC
LIMIT ?{lim_idx} OFFSET ?{off_idx}",
);
let mut stmt = conn.prepare(&query_sql)?;
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|b| b.as_ref()).collect();
let rows = stmt.query_map(params_ref.as_slice(), map_session_row)?;
let mut sessions = Vec::new();
for row in rows {
sessions.push(row?);
}
Ok(SessionSearchResult { sessions, total })
})
}
pub fn search_sessions(
config: &Config,
params: &SessionSearchParams,
) -> Result<SessionSearchResult> {
log::debug!(
"[session_db] search_sessions query={} agent={} tool={} channel={} thread={}",
params.query.as_deref().unwrap_or("-"),
params.agent_id.as_deref().unwrap_or("-"),
params.tool_name.as_deref().unwrap_or("-"),
params.source_channel.as_deref().unwrap_or("-"),
params.thread_id.as_deref().unwrap_or("-"),
);
with_connection(config, |conn| search_sessions_inner(conn, params))
}
fn search_sessions_inner(
conn: &Connection,
params: &SessionSearchParams,
) -> Result<SessionSearchResult> {
let lim = params.limit.unwrap_or(50).min(500) as i64;
let off = params.offset.unwrap_or(0) as i64;
let mut where_clauses: Vec<String> = Vec::new();
let mut param_values: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
if let Some(ref q) = params.query {
if !q.trim().is_empty() {
param_values.push(Box::new(q.clone()));
where_clauses.push(format!(
"s.id IN (SELECT session_id FROM sessions_fts WHERE sessions_fts MATCH ?{})",
param_values.len()
));
}
}
if let Some(ref agent) = params.agent_id {
param_values.push(Box::new(agent.clone()));
where_clauses.push(format!("s.agent_definition_id = ?{}", param_values.len()));
}
if let Some(ref tool) = params.tool_name {
param_values.push(Box::new(tool.clone()));
where_clauses.push(format!(
"s.id IN (SELECT DISTINCT session_id FROM session_tool_calls WHERE tool_name = ?{})",
param_values.len()
));
}
if let Some(ref channel) = params.source_channel {
param_values.push(Box::new(channel.clone()));
where_clauses.push(format!("s.source_channel = ?{}", param_values.len()));
}
if let Some(ref parent) = params.parent_session_id {
param_values.push(Box::new(parent.clone()));
where_clauses.push(format!("s.parent_session_id = ?{}", param_values.len()));
}
if let Some(ref status) = params.status {
param_values.push(Box::new(status.clone()));
where_clauses.push(format!("s.status = ?{}", param_values.len()));
}
if let Some(ref tid) = params.thread_id {
param_values.push(Box::new(tid.clone()));
where_clauses.push(format!("s.thread_id = ?{}", param_values.len()));
}
let where_sql = if where_clauses.is_empty() {
String::new()
} else {
format!("WHERE {}", where_clauses.join(" AND "))
};
let count_sql = format!("SELECT COUNT(*) FROM sessions s {where_sql}");
let total: u64 = {
let mut stmt = conn.prepare(&count_sql)?;
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|b| b.as_ref()).collect();
stmt.query_row(params_ref.as_slice(), |r| r.get::<_, i64>(0))? as u64
};
param_values.push(Box::new(lim));
let lim_idx = param_values.len();
param_values.push(Box::new(off));
let off_idx = param_values.len();
let query = format!(
"SELECT s.id, s.agent_definition_id, s.agent_definition_name, s.session_key,
s.parent_session_id, s.thread_id, s.source_channel, s.status, s.model,
s.turn_count, s.input_tokens, s.output_tokens, s.cached_input_tokens,
s.cost_usd, s.transcript_path, s.started_at, s.ended_at
FROM sessions s {where_sql}
ORDER BY s.started_at DESC
LIMIT ?{lim_idx} OFFSET ?{off_idx}",
);
let mut stmt = conn.prepare(&query)?;
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|b| b.as_ref()).collect();
let rows = stmt.query_map(params_ref.as_slice(), map_session_row)?;
let mut sessions = Vec::new();
for row in rows {
sessions.push(row?);
}
Ok(SessionSearchResult { sessions, total })
}
pub fn import_transcript(
config: &Config,
transcript_path: &std::path::Path,
) -> Result<SessionRecord> {
use crate::openhuman::agent::harness::session::transcript::read_transcript;
let transcript = read_transcript(transcript_path)
.with_context(|| format!("failed to read transcript {}", transcript_path.display()))?;
let stem = transcript_path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("unknown");
let session_key = stem.to_string();
let id = uuid::Uuid::new_v4().to_string();
let meta = &transcript.meta;
let parent_session_id: Option<&str> = None;
let is_subagent = stem.contains("__");
let parent_prefix = if is_subagent {
stem.rsplit_once("__").map(|(p, _)| p)
} else {
None
};
log::debug!(
"[session_db] import_transcript path={} stem={stem} agent={} turns={} cost=${:.6}",
transcript_path.display(),
meta.agent_name,
meta.turn_count,
meta.charged_amount_usd,
);
with_connection(config, |conn| {
let parent_id = if let Some(prefix) = parent_prefix {
let mut stmt = conn.prepare(
"SELECT id FROM sessions WHERE session_key = ?1 ORDER BY started_at DESC LIMIT 1",
)?;
let mut rows = stmt.query(params![prefix])?;
rows.next()?.map(|r| r.get::<_, String>(0)).transpose()?
} else {
None
};
let started = chrono::DateTime::parse_from_rfc3339(&meta.created)
.unwrap_or_else(|_| Utc::now().into())
.with_timezone(&Utc);
let ended = chrono::DateTime::parse_from_rfc3339(&meta.updated)
.ok()
.map(|dt| dt.with_timezone(&Utc));
conn.execute(
"INSERT OR IGNORE INTO sessions (
id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)",
params![
id,
meta.agent_name,
meta.agent_name,
session_key,
parent_id.as_deref().or(parent_session_id),
meta.thread_id,
meta.dispatcher,
"completed",
Option::<String>::None,
meta.turn_count as i64,
meta.input_tokens as i64,
meta.output_tokens as i64,
meta.cached_input_tokens as i64,
meta.charged_amount_usd,
transcript_path.to_string_lossy().as_ref(),
started.to_rfc3339(),
ended.map(|dt| dt.to_rfc3339()),
],
)
.context("failed to insert imported session")?;
index_fts_session(conn, &id, &meta.agent_name)?;
for msg in &transcript.messages {
conn.execute(
"INSERT INTO session_messages (session_id, role, content, created_at)
VALUES (?1, ?2, ?3, ?4)",
params![id, msg.role, msg.content, started.to_rfc3339()],
)?;
index_fts_content(conn, &id, &msg.content)?;
}
Ok(())
})?;
get_session(config, &id)
}
pub fn list_messages(
config: &Config,
session_id: &str,
limit: Option<u32>,
) -> Result<Vec<SessionMessage>> {
with_connection(config, |conn| {
let lim = limit.unwrap_or(200).min(1000) as i64;
let mut stmt = conn.prepare(
"SELECT id, session_id, role, content, model,
input_tokens, output_tokens, cost_usd, created_at
FROM session_messages
WHERE session_id = ?1
ORDER BY id ASC
LIMIT ?2",
)?;
let rows = stmt.query_map(params![session_id, lim], |row| {
Ok(SessionMessage {
id: row.get(0)?,
session_id: row.get(1)?,
role: row.get(2)?,
content: row.get(3)?,
model: row.get(4)?,
input_tokens: row.get::<_, Option<i64>>(5)?.map(|v| v as u64),
output_tokens: row.get::<_, Option<i64>>(6)?.map(|v| v as u64),
cost_usd: row.get(7)?,
created_at: parse_rfc3339(&row.get::<_, String>(8)?)
.map_err(sql_conversion_error)?,
})
})?;
let mut messages = Vec::new();
for row in rows {
messages.push(row?);
}
Ok(messages)
})
}
pub fn list_tool_calls(
config: &Config,
session_id: &str,
limit: Option<u32>,
) -> Result<Vec<SessionToolCall>> {
with_connection(config, |conn| {
let lim = limit.unwrap_or(200).min(1000) as i64;
let mut stmt = conn.prepare(
"SELECT id, session_id, message_id, tool_name, tool_input,
tool_output, status, duration_ms, created_at
FROM session_tool_calls
WHERE session_id = ?1
ORDER BY id ASC
LIMIT ?2",
)?;
let rows = stmt.query_map(params![session_id, lim], |row| {
Ok(SessionToolCall {
id: row.get(0)?,
session_id: row.get(1)?,
message_id: row.get(2)?,
tool_name: row.get(3)?,
tool_input: row.get(4)?,
tool_output: row.get(5)?,
status: row.get(6)?,
duration_ms: row.get(7)?,
created_at: parse_rfc3339(&row.get::<_, String>(8)?)
.map_err(sql_conversion_error)?,
})
})?;
let mut tool_calls = Vec::new();
for row in rows {
tool_calls.push(row?);
}
Ok(tool_calls)
})
}
pub fn list_children(config: &Config, session_id: &str) -> Result<Vec<SessionRecord>> {
with_connection(config, |conn| {
let mut stmt = conn.prepare(
"SELECT id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
FROM sessions
WHERE parent_session_id = ?1
ORDER BY started_at ASC",
)?;
let rows = stmt.query_map(params![session_id], map_session_row)?;
let mut children = Vec::new();
for row in rows {
children.push(row?);
}
Ok(children)
})
}
pub fn mark_interrupted(config: &Config) -> Result<usize> {
log::debug!("[session_db] mark_interrupted — marking all running sessions as interrupted");
with_connection(config, |conn| {
let now = Utc::now();
let changed = conn.execute(
"UPDATE sessions SET status = 'interrupted', ended_at = ?1
WHERE status = 'running'",
params![now.to_rfc3339()],
)?;
if changed > 0 {
log::info!("[session_db] marked {changed} running session(s) as interrupted");
}
Ok(changed)
})
}
fn index_fts_session(conn: &Connection, session_id: &str, agent_name: &str) -> Result<()> {
conn.execute(
"INSERT INTO sessions_fts (session_id, agent_definition_name, content, tool_name)
VALUES (?1, ?2, '', '')",
params![session_id, agent_name],
)
.context("failed to index session in FTS")?;
Ok(())
}
fn index_fts_content(conn: &Connection, session_id: &str, content: &str) -> Result<()> {
let snippet = if content.len() > 2000 {
&content[..2000]
} else {
content
};
conn.execute(
"INSERT INTO sessions_fts (session_id, agent_definition_name, content, tool_name)
VALUES (?1, '', ?2, '')",
params![session_id, snippet],
)
.context("failed to index content in FTS")?;
Ok(())
}
fn index_fts_tool(conn: &Connection, session_id: &str, tool_name: &str) -> Result<()> {
conn.execute(
"INSERT INTO sessions_fts (session_id, agent_definition_name, content, tool_name)
VALUES (?1, '', '', ?2)",
params![session_id, tool_name],
)
.context("failed to index tool call in FTS")?;
Ok(())
}
fn map_session_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<SessionRecord> {
let started_at_raw: String = row.get(15)?;
let ended_at_raw: Option<String> = row.get(16)?;
Ok(SessionRecord {
id: row.get(0)?,
agent_definition_id: row.get(1)?,
agent_definition_name: row.get(2)?,
session_key: row.get(3)?,
parent_session_id: row.get(4)?,
thread_id: row.get(5)?,
source_channel: row.get(6)?,
status: SessionStatus::parse(&row.get::<_, String>(7)?),
model: row.get(8)?,
turn_count: row.get::<_, i64>(9)? as u32,
input_tokens: row.get::<_, i64>(10)? as u64,
output_tokens: row.get::<_, i64>(11)? as u64,
cached_input_tokens: row.get::<_, i64>(12)? as u64,
cost_usd: row.get(13)?,
transcript_path: row.get(14)?,
started_at: parse_rfc3339(&started_at_raw).map_err(sql_conversion_error)?,
ended_at: match ended_at_raw {
Some(raw) => Some(parse_rfc3339(&raw).map_err(sql_conversion_error)?),
None => None,
},
})
}
fn parse_rfc3339(raw: &str) -> Result<DateTime<Utc>> {
let parsed = DateTime::parse_from_rfc3339(raw)
.with_context(|| format!("invalid RFC3339 timestamp in session DB: {raw}"))?;
Ok(parsed.with_timezone(&Utc))
}
fn sql_conversion_error(err: anyhow::Error) -> rusqlite::Error {
rusqlite::Error::ToSqlConversionFailure(err.into())
}
#[cfg(test)]
#[path = "ops_tests.rs"]
mod tests;
+349
View File
@@ -0,0 +1,349 @@
use super::*;
use crate::openhuman::session_db::store::with_memory_connection;
use crate::openhuman::session_db::types::SessionSearchParams;
fn insert_test_session(conn: &Connection, id: &str, agent_id: &str, key: &str) {
let now = Utc::now();
conn.execute(
"INSERT INTO sessions (
id, agent_definition_id, agent_definition_name, session_key,
status, started_at
) VALUES (?1, ?2, ?3, ?4, 'running', ?5)",
params![id, agent_id, agent_id, key, now.to_rfc3339()],
)
.unwrap();
index_fts_session(conn, id, agent_id).unwrap();
}
fn insert_test_session_with_parent(
conn: &Connection,
id: &str,
agent_id: &str,
key: &str,
parent_id: &str,
) {
let now = Utc::now();
conn.execute(
"INSERT INTO sessions (
id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, status, started_at
) VALUES (?1, ?2, ?3, ?4, ?5, 'running', ?6)",
params![id, agent_id, agent_id, key, parent_id, now.to_rfc3339()],
)
.unwrap();
index_fts_session(conn, id, agent_id).unwrap();
}
#[test]
fn map_session_row_roundtrip() {
with_memory_connection(|conn| {
insert_test_session(conn, "sess-1", "orchestrator", "1700000000_orchestrator");
let mut stmt = conn.prepare(
"SELECT id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
FROM sessions WHERE id = 'sess-1'",
)?;
let session = stmt.query_row([], map_session_row)?;
assert_eq!(session.id, "sess-1");
assert_eq!(session.agent_definition_id, "orchestrator");
assert_eq!(session.session_key, "1700000000_orchestrator");
assert_eq!(session.status, SessionStatus::Running);
assert!(session.parent_session_id.is_none());
assert!(session.ended_at.is_none());
Ok(())
})
.unwrap();
}
#[test]
fn search_by_agent_id() {
with_memory_connection(|conn| {
insert_test_session(conn, "a1", "orchestrator", "key1");
insert_test_session(conn, "a2", "researcher", "key2");
insert_test_session(conn, "a3", "orchestrator", "key3");
let params = SessionSearchParams {
agent_id: Some("orchestrator".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 2);
assert_eq!(result.sessions.len(), 2);
assert!(result
.sessions
.iter()
.all(|s| s.agent_definition_id == "orchestrator"));
Ok(())
})
.unwrap();
}
#[test]
fn search_by_fts_query() {
with_memory_connection(|conn| {
insert_test_session(conn, "b1", "orchestrator", "key1");
insert_test_session(conn, "b2", "researcher", "key2");
conn.execute(
"INSERT INTO session_messages (session_id, role, content, created_at)
VALUES ('b1', 'user', 'Fix the login bug in authentication', ?1)",
params![Utc::now().to_rfc3339()],
)?;
index_fts_content(conn, "b1", "Fix the login bug in authentication")?;
conn.execute(
"INSERT INTO session_messages (session_id, role, content, created_at)
VALUES ('b2', 'user', 'Deploy the new feature to production', ?1)",
params![Utc::now().to_rfc3339()],
)?;
index_fts_content(conn, "b2", "Deploy the new feature to production")?;
let params = SessionSearchParams {
query: Some("login".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 1);
assert_eq!(result.sessions[0].id, "b1");
Ok(())
})
.unwrap();
}
#[test]
fn search_by_tool_name() {
with_memory_connection(|conn| {
insert_test_session(conn, "c1", "orchestrator", "key1");
insert_test_session(conn, "c2", "researcher", "key2");
conn.execute(
"INSERT INTO session_tool_calls (session_id, tool_name, status, created_at)
VALUES ('c1', 'shell', 'ok', ?1)",
params![Utc::now().to_rfc3339()],
)?;
conn.execute(
"INSERT INTO session_tool_calls (session_id, tool_name, status, created_at)
VALUES ('c2', 'file_read', 'ok', ?1)",
params![Utc::now().to_rfc3339()],
)?;
let params = SessionSearchParams {
tool_name: Some("shell".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 1);
assert_eq!(result.sessions[0].id, "c1");
Ok(())
})
.unwrap();
}
#[test]
fn search_by_parent_session() {
with_memory_connection(|conn| {
insert_test_session(conn, "parent-1", "orchestrator", "key1");
insert_test_session_with_parent(conn, "child-1", "researcher", "key2", "parent-1");
insert_test_session_with_parent(conn, "child-2", "coder", "key3", "parent-1");
insert_test_session(conn, "unrelated", "other", "key4");
let params = SessionSearchParams {
parent_session_id: Some("parent-1".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 2);
Ok(())
})
.unwrap();
}
#[test]
fn search_pagination() {
with_memory_connection(|conn| {
for i in 0..10 {
insert_test_session(conn, &format!("p{i}"), "agent", &format!("key{i}"));
}
let params = SessionSearchParams {
limit: Some(3),
offset: Some(0),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 10);
assert_eq!(result.sessions.len(), 3);
let params2 = SessionSearchParams {
limit: Some(3),
offset: Some(3),
..Default::default()
};
let result2 = search_sessions_inner(conn, &params2)?;
assert_eq!(result2.total, 10);
assert_eq!(result2.sessions.len(), 3);
assert_ne!(result.sessions[0].id, result2.sessions[0].id);
Ok(())
})
.unwrap();
}
#[test]
fn search_empty_results() {
with_memory_connection(|conn| {
let params = SessionSearchParams {
agent_id: Some("nonexistent".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 0);
assert!(result.sessions.is_empty());
Ok(())
})
.unwrap();
}
#[test]
fn tool_output_truncation() {
with_memory_connection(|conn| {
let session_id = "trunc-sess";
insert_test_session(conn, session_id, "agent", "key");
let large_output = "x".repeat(MAX_TOOL_OUTPUT_BYTES + 1000);
let bounded = if large_output.len() <= MAX_TOOL_OUTPUT_BYTES {
large_output.clone()
} else {
let mut cutoff = MAX_TOOL_OUTPUT_BYTES;
while cutoff > 0 && !large_output.is_char_boundary(cutoff) {
cutoff -= 1;
}
let mut truncated = large_output[..cutoff].to_string();
truncated.push_str("\n...[truncated]");
truncated
};
conn.execute(
"INSERT INTO session_tool_calls (session_id, tool_name, tool_output, status, created_at)
VALUES (?1, 'test', ?2, 'ok', ?3)",
params![session_id, bounded, Utc::now().to_rfc3339()],
)?;
let stored: String = conn.query_row(
"SELECT tool_output FROM session_tool_calls WHERE session_id = ?1",
params![session_id],
|r| r.get(0),
)?;
assert!(stored.len() <= MAX_TOOL_OUTPUT_BYTES + 20);
assert!(stored.ends_with("[truncated]"));
Ok(())
})
.unwrap();
}
#[test]
fn mark_interrupted_updates_running() {
with_memory_connection(|conn| {
insert_test_session(conn, "run1", "agent", "key1");
insert_test_session(conn, "run2", "agent", "key2");
conn.execute(
"UPDATE sessions SET status = 'completed' WHERE id = 'run2'",
[],
)?;
let now = Utc::now();
let changed = conn.execute(
"UPDATE sessions SET status = 'interrupted', ended_at = ?1
WHERE status = 'running'",
params![now.to_rfc3339()],
)?;
assert_eq!(changed, 1);
let status: String =
conn.query_row("SELECT status FROM sessions WHERE id = 'run1'", [], |r| {
r.get(0)
})?;
assert_eq!(status, "interrupted");
let status2: String =
conn.query_row("SELECT status FROM sessions WHERE id = 'run2'", [], |r| {
r.get(0)
})?;
assert_eq!(status2, "completed");
Ok(())
})
.unwrap();
}
#[test]
fn session_end_updates_cost_fields() {
with_memory_connection(|conn| {
insert_test_session(conn, "cost-sess", "agent", "key");
let now = Utc::now();
conn.execute(
"UPDATE sessions SET
status = 'completed', turn_count = 5, input_tokens = 10000,
output_tokens = 2000, cached_input_tokens = 8000,
cost_usd = 0.0345, ended_at = ?1
WHERE id = 'cost-sess'",
params![now.to_rfc3339()],
)?;
let mut stmt = conn.prepare(
"SELECT id, agent_definition_id, agent_definition_name, session_key,
parent_session_id, thread_id, source_channel, status, model,
turn_count, input_tokens, output_tokens, cached_input_tokens,
cost_usd, transcript_path, started_at, ended_at
FROM sessions WHERE id = 'cost-sess'",
)?;
let session = stmt.query_row([], map_session_row)?;
assert_eq!(session.status, SessionStatus::Completed);
assert_eq!(session.turn_count, 5);
assert_eq!(session.input_tokens, 10000);
assert_eq!(session.output_tokens, 2000);
assert_eq!(session.cached_input_tokens, 8000);
assert!((session.cost_usd - 0.0345).abs() < f64::EPSILON);
assert!(session.ended_at.is_some());
Ok(())
})
.unwrap();
}
#[test]
fn combined_filters() {
with_memory_connection(|conn| {
insert_test_session(conn, "cf1", "orchestrator", "key1");
insert_test_session(conn, "cf2", "orchestrator", "key2");
insert_test_session(conn, "cf3", "researcher", "key3");
conn.execute(
"UPDATE sessions SET status = 'completed' WHERE id = 'cf1'",
[],
)?;
let params = SessionSearchParams {
agent_id: Some("orchestrator".to_string()),
status: Some("completed".to_string()),
..Default::default()
};
let result = search_sessions_inner(conn, &params)?;
assert_eq!(result.total, 1);
assert_eq!(result.sessions[0].id, "cf1");
Ok(())
})
.unwrap();
}
+492
View File
@@ -0,0 +1,492 @@
//! Controller schemas and JSON-RPC dispatchers for the session database.
use serde_json::{Map, Value};
use crate::core::all::{ControllerFuture, RegisteredController};
use crate::core::{ControllerSchema, FieldSchema, TypeSchema};
use crate::openhuman::config::rpc as config_rpc;
use crate::rpc::RpcOutcome;
use super::types::SessionSearchParams;
pub fn all_controller_schemas() -> Vec<ControllerSchema> {
vec![
schema_for("session_db_list"),
schema_for("session_db_get"),
schema_for("session_db_search"),
schema_for("session_db_get_messages"),
schema_for("session_db_get_tool_calls"),
schema_for("session_db_get_children"),
schema_for("session_db_import_transcript"),
]
}
pub fn all_registered_controllers() -> Vec<RegisteredController> {
vec![
RegisteredController {
schema: schema_for("session_db_list"),
handler: handle_session_db_list,
},
RegisteredController {
schema: schema_for("session_db_get"),
handler: handle_session_db_get,
},
RegisteredController {
schema: schema_for("session_db_search"),
handler: handle_session_db_search,
},
RegisteredController {
schema: schema_for("session_db_get_messages"),
handler: handle_session_db_get_messages,
},
RegisteredController {
schema: schema_for("session_db_get_tool_calls"),
handler: handle_session_db_get_tool_calls,
},
RegisteredController {
schema: schema_for("session_db_get_children"),
handler: handle_session_db_get_children,
},
RegisteredController {
schema: schema_for("session_db_import_transcript"),
handler: handle_session_db_import_transcript,
},
]
}
fn schema_for(function: &str) -> ControllerSchema {
match function {
"session_db_list" => ControllerSchema {
namespace: "session_db",
function: "list",
description: "List agent sessions with optional filters (status, parent) \
and pagination.",
inputs: vec![
optional_u64("limit", "Max sessions to return (default 50, max 500)."),
optional_u64("offset", "Pagination offset."),
optional_str(
"status",
"Filter by status (running, completed, failed, interrupted).",
),
optional_str("parentSessionId", "Filter by parent session ID."),
],
outputs: vec![json_output(
"result",
"SessionSearchResult with sessions array and total count.",
)],
},
"session_db_get" => ControllerSchema {
namespace: "session_db",
function: "get",
description: "Get a single session by ID.",
inputs: vec![required_str("id", "Session ID.")],
outputs: vec![json_output("session", "Full SessionRecord.")],
},
"session_db_search" => ControllerSchema {
namespace: "session_db",
function: "search",
description: "Search sessions by full-text query, agent ID, tool name, \
source channel, thread ID, parent, and/or status.",
inputs: vec![
optional_str("query", "Full-text search query."),
optional_str("agentId", "Filter by agent definition ID."),
optional_str("toolName", "Filter to sessions that used this tool."),
optional_str("sourceChannel", "Filter by source channel."),
optional_str("threadId", "Filter by thread ID."),
optional_str("parentSessionId", "Filter by parent session ID."),
optional_str("status", "Filter by status."),
optional_u64("limit", "Max results (default 50, max 500)."),
optional_u64("offset", "Pagination offset."),
],
outputs: vec![json_output(
"result",
"SessionSearchResult with sessions array and total count.",
)],
},
"session_db_get_messages" => ControllerSchema {
namespace: "session_db",
function: "get_messages",
description: "Get messages for a session.",
inputs: vec![
required_str("sessionId", "Session ID."),
optional_u64("limit", "Max messages (default 200, max 1000)."),
],
outputs: vec![json_output("messages", "Array of SessionMessage objects.")],
},
"session_db_get_tool_calls" => ControllerSchema {
namespace: "session_db",
function: "get_tool_calls",
description: "Get tool calls for a session.",
inputs: vec![
required_str("sessionId", "Session ID."),
optional_u64("limit", "Max tool calls (default 200, max 1000)."),
],
outputs: vec![json_output(
"toolCalls",
"Array of SessionToolCall objects.",
)],
},
"session_db_get_children" => ControllerSchema {
namespace: "session_db",
function: "get_children",
description: "Get child (sub-agent) sessions for a parent session.",
inputs: vec![required_str("sessionId", "Parent session ID.")],
outputs: vec![json_output(
"children",
"Array of child SessionRecord objects.",
)],
},
"session_db_import_transcript" => ControllerSchema {
namespace: "session_db",
function: "import_transcript",
description: "Import a JSONL transcript file into the session database.",
inputs: vec![required_str(
"path",
"Absolute path to the JSONL transcript file.",
)],
outputs: vec![json_output("session", "Imported SessionRecord.")],
},
_ => ControllerSchema {
namespace: "session_db",
function: "unknown",
description: "Unknown session_db controller.",
inputs: vec![],
outputs: vec![FieldSchema {
name: "error",
ty: TypeSchema::String,
comment: "Lookup error details.",
required: true,
}],
},
}
}
fn new_correlation_id() -> String {
uuid::Uuid::new_v4().simple().to_string()[..8].to_string()
}
fn handle_session_db_list(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] list.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] list.config_failed err={err}");
})?;
let limit = params
.get("limit")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let offset = params
.get("offset")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let status = params
.get("status")
.and_then(|v| v.as_str())
.map(String::from);
let parent_id = params
.get("parentSessionId")
.and_then(|v| v.as_str())
.map(String::from);
let result = super::ops::list_sessions(
&config,
limit,
offset,
status.as_deref(),
parent_id.as_deref(),
)
.map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] list.error err={s}");
s
})?;
let json = to_json(result);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] list.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_get(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get.config_failed err={err}");
})?;
let id = params
.get("id")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required param: id".to_string())?;
let session = super::ops::get_session(&config, id).map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get.error id={id} err={s}");
s
})?;
let json = to_json(session);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_search(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] search.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] search.config_failed err={err}");
})?;
let search_params: SessionSearchParams = if params.is_empty() {
SessionSearchParams::default()
} else {
serde_json::from_value(Value::Object(params)).map_err(|e| {
let s = format!("invalid search params: {e}");
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] search.bad_params err={s}");
s
})?
};
let result = super::ops::search_sessions(&config, &search_params).map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] search.error err={s}");
s
})?;
let json = to_json(result);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] search.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_get_messages(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_messages.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_messages.config_failed err={err}");
})?;
let session_id = params
.get("sessionId")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required param: sessionId".to_string())?;
let limit = params
.get("limit")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let messages = super::ops::list_messages(&config, session_id, limit).map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_messages.error err={s}");
s
})?;
let json = to_json(messages);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_messages.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_get_tool_calls(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_tool_calls.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_tool_calls.config_failed err={err}");
})?;
let session_id = params
.get("sessionId")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required param: sessionId".to_string())?;
let limit = params
.get("limit")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let tool_calls = super::ops::list_tool_calls(&config, session_id, limit).map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_tool_calls.error err={s}");
s
})?;
let json = to_json(tool_calls);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_tool_calls.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_get_children(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_children.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_children.config_failed err={err}");
})?;
let session_id = params
.get("sessionId")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required param: sessionId".to_string())?;
let children = super::ops::list_children(&config, session_id).map_err(|e| {
let s = e.to_string();
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_children.error err={s}");
s
})?;
let json = to_json(children);
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] get_children.exit ok={}", json.is_ok());
json
})
}
fn handle_session_db_import_transcript(params: Map<String, Value>) -> ControllerFuture {
Box::pin(async move {
let cid = new_correlation_id();
log::debug!(target: "session_db_rpc", "[session_db_rpc][{cid}] import_transcript.entry");
let config = config_rpc::load_config_with_timeout().await.inspect_err(|err| {
log::warn!(target: "session_db_rpc", "[session_db_rpc][{cid}] import_transcript.config_failed err={err}");
})?;
let path = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| "missing required param: path".to_string())?;
let session =
super::ops::import_transcript(&config, std::path::Path::new(path)).map_err(|e| {
let s = e.to_string();
log::warn!(
target: "session_db_rpc",
"[session_db_rpc][{cid}] import_transcript.error path={path} err={s}"
);
s
})?;
let json = to_json(session);
log::debug!(
target: "session_db_rpc",
"[session_db_rpc][{cid}] import_transcript.exit ok={}",
json.is_ok()
);
json
})
}
fn to_json<T: serde::Serialize>(value: T) -> Result<Value, String> {
RpcOutcome::new(value, vec![]).into_cli_compatible_json()
}
fn required_str(name: &'static str, comment: &'static str) -> FieldSchema {
FieldSchema {
name,
ty: TypeSchema::String,
comment,
required: true,
}
}
fn optional_str(name: &'static str, comment: &'static str) -> FieldSchema {
FieldSchema {
name,
ty: TypeSchema::Option(Box::new(TypeSchema::String)),
comment,
required: false,
}
}
fn optional_u64(name: &'static str, comment: &'static str) -> FieldSchema {
FieldSchema {
name,
ty: TypeSchema::Option(Box::new(TypeSchema::U64)),
comment,
required: false,
}
}
fn json_output(name: &'static str, comment: &'static str) -> FieldSchema {
FieldSchema {
name,
ty: TypeSchema::Json,
comment,
required: true,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn all_controller_schemas_lists_seven_functions() {
let schemas = all_controller_schemas();
assert_eq!(schemas.len(), 7);
for schema in &schemas {
assert_eq!(schema.namespace, "session_db");
}
}
#[test]
fn all_registered_controllers_match_schemas() {
let registered = all_registered_controllers();
let schemas = all_controller_schemas();
assert_eq!(registered.len(), schemas.len());
let schema_fns: Vec<&str> = schemas.iter().map(|s| s.function).collect();
for rc in &registered {
assert!(
schema_fns.contains(&rc.schema.function),
"registered controller '{}' not in schema list",
rc.schema.function
);
}
}
#[test]
fn schema_for_list_has_optional_inputs() {
let s = schema_for("session_db_list");
assert_eq!(s.function, "list");
assert!(s.inputs.iter().all(|i| !i.required));
}
#[test]
fn schema_for_get_requires_id() {
let s = schema_for("session_db_get");
assert_eq!(s.function, "get");
assert_eq!(s.inputs.len(), 1);
assert!(s.inputs[0].required);
assert_eq!(s.inputs[0].name, "id");
}
#[test]
fn schema_for_search_has_query_and_filters() {
let s = schema_for("session_db_search");
assert_eq!(s.function, "search");
let names: Vec<&str> = s.inputs.iter().map(|i| i.name).collect();
assert!(names.contains(&"query"));
assert!(names.contains(&"agentId"));
assert!(names.contains(&"toolName"));
assert!(names.contains(&"sourceChannel"));
assert!(names.contains(&"threadId"));
}
#[test]
fn schema_for_unknown_returns_error_shape() {
let s = schema_for("session_db_nonexistent");
assert_eq!(s.function, "unknown");
}
#[test]
fn new_correlation_id_is_eight_hex_chars() {
let cid = new_correlation_id();
assert_eq!(cid.len(), 8);
assert!(cid.chars().all(|c| c.is_ascii_hexdigit()));
}
}
+175
View File
@@ -0,0 +1,175 @@
use crate::openhuman::config::Config;
use anyhow::{Context, Result};
use rusqlite::Connection;
pub fn with_connection<T>(config: &Config, f: impl FnOnce(&Connection) -> Result<T>) -> Result<T> {
let db_path = config.workspace_dir.join("session_db").join("sessions.db");
if let Some(parent) = db_path.parent() {
std::fs::create_dir_all(parent).with_context(|| {
format!(
"failed to create session_db directory: {}",
parent.display()
)
})?;
}
let conn = Connection::open(&db_path)
.with_context(|| format!("failed to open session DB: {}", db_path.display()))?;
init_schema(&conn)?;
f(&conn)
}
#[cfg(test)]
pub fn with_memory_connection<T>(f: impl FnOnce(&Connection) -> Result<T>) -> Result<T> {
let conn = Connection::open_in_memory().context("failed to open in-memory session DB")?;
init_schema(&conn)?;
f(&conn)
}
fn init_schema(conn: &Connection) -> Result<()> {
conn.execute_batch(
"PRAGMA journal_mode = WAL;
PRAGMA foreign_keys = ON;
CREATE TABLE IF NOT EXISTS sessions (
id TEXT PRIMARY KEY,
agent_definition_id TEXT NOT NULL,
agent_definition_name TEXT NOT NULL,
session_key TEXT NOT NULL,
parent_session_id TEXT,
thread_id TEXT,
source_channel TEXT,
status TEXT NOT NULL DEFAULT 'running',
model TEXT,
turn_count INTEGER NOT NULL DEFAULT 0,
input_tokens INTEGER NOT NULL DEFAULT 0,
output_tokens INTEGER NOT NULL DEFAULT 0,
cached_input_tokens INTEGER NOT NULL DEFAULT 0,
cost_usd REAL NOT NULL DEFAULT 0.0,
transcript_path TEXT,
started_at TEXT NOT NULL,
ended_at TEXT,
FOREIGN KEY (parent_session_id) REFERENCES sessions(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_sessions_agent ON sessions(agent_definition_id);
CREATE INDEX IF NOT EXISTS idx_sessions_status ON sessions(status);
CREATE INDEX IF NOT EXISTS idx_sessions_started ON sessions(started_at);
CREATE INDEX IF NOT EXISTS idx_sessions_parent ON sessions(parent_session_id);
CREATE INDEX IF NOT EXISTS idx_sessions_thread ON sessions(thread_id);
CREATE INDEX IF NOT EXISTS idx_sessions_channel ON sessions(source_channel);
CREATE INDEX IF NOT EXISTS idx_sessions_key ON sessions(session_key);
CREATE TABLE IF NOT EXISTS session_messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id TEXT NOT NULL,
role TEXT NOT NULL,
content TEXT NOT NULL,
model TEXT,
input_tokens INTEGER,
output_tokens INTEGER,
cost_usd REAL,
created_at TEXT NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_messages_session ON session_messages(session_id);
CREATE TABLE IF NOT EXISTS session_tool_calls (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id TEXT NOT NULL,
message_id INTEGER,
tool_name TEXT NOT NULL,
tool_input TEXT,
tool_output TEXT,
status TEXT NOT NULL DEFAULT 'pending',
duration_ms INTEGER,
created_at TEXT NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE,
FOREIGN KEY (message_id) REFERENCES session_messages(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_tool_calls_session ON session_tool_calls(session_id);
CREATE INDEX IF NOT EXISTS idx_tool_calls_name ON session_tool_calls(tool_name);",
)
.context("failed to initialize session_db schema")?;
init_fts(conn)?;
Ok(())
}
fn init_fts(conn: &Connection) -> Result<()> {
let has_fts: bool = conn
.prepare("SELECT 1 FROM sqlite_master WHERE type='table' AND name='sessions_fts'")?
.exists([])?;
if !has_fts {
conn.execute_batch(
"CREATE VIRTUAL TABLE sessions_fts USING fts5(
session_id,
agent_definition_name,
content,
tool_name
);",
)
.context("failed to create sessions_fts virtual table")?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn schema_initializes_without_error() {
with_memory_connection(|conn| {
let count: i64 = conn.query_row("SELECT COUNT(*) FROM sessions", [], |r| r.get(0))?;
assert_eq!(count, 0);
Ok(())
})
.unwrap();
}
#[test]
fn schema_is_idempotent() {
let conn = Connection::open_in_memory().unwrap();
init_schema(&conn).unwrap();
init_schema(&conn).unwrap();
let count: i64 = conn
.query_row("SELECT COUNT(*) FROM sessions", [], |r| r.get(0))
.unwrap();
assert_eq!(count, 0);
}
#[test]
fn wal_mode_is_set() {
with_memory_connection(|conn| {
let mode: String = conn.query_row("PRAGMA journal_mode", [], |r| r.get(0))?;
// In-memory DBs may report "memory" instead of "wal"
assert!(mode == "wal" || mode == "memory");
Ok(())
})
.unwrap();
}
#[test]
fn fts_table_exists_after_init() {
with_memory_connection(|conn| {
let exists: bool = conn
.prepare("SELECT 1 FROM sqlite_master WHERE type='table' AND name='sessions_fts'")?
.exists([])?;
assert!(exists);
Ok(())
})
.unwrap();
}
#[test]
fn foreign_keys_are_enabled() {
with_memory_connection(|conn| {
let fk: i64 = conn.query_row("PRAGMA foreign_keys", [], |r| r.get(0))?;
assert_eq!(fk, 1);
Ok(())
})
.unwrap();
}
}
+148
View File
@@ -0,0 +1,148 @@
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SessionStatus {
Running,
Completed,
Failed,
Interrupted,
}
impl SessionStatus {
pub fn as_str(&self) -> &'static str {
match self {
Self::Running => "running",
Self::Completed => "completed",
Self::Failed => "failed",
Self::Interrupted => "interrupted",
}
}
pub fn parse(s: &str) -> Self {
match s {
"completed" => Self::Completed,
"failed" => Self::Failed,
"interrupted" => Self::Interrupted,
_ => Self::Running,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionRecord {
pub id: String,
pub agent_definition_id: String,
pub agent_definition_name: String,
pub session_key: String,
pub parent_session_id: Option<String>,
pub thread_id: Option<String>,
pub source_channel: Option<String>,
pub status: SessionStatus,
pub model: Option<String>,
pub turn_count: u32,
pub input_tokens: u64,
pub output_tokens: u64,
pub cached_input_tokens: u64,
pub cost_usd: f64,
pub transcript_path: Option<String>,
pub started_at: DateTime<Utc>,
pub ended_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionMessage {
pub id: i64,
pub session_id: String,
pub role: String,
pub content: String,
pub model: Option<String>,
pub input_tokens: Option<u64>,
pub output_tokens: Option<u64>,
pub cost_usd: Option<f64>,
pub created_at: DateTime<Utc>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionToolCall {
pub id: i64,
pub session_id: String,
pub message_id: Option<i64>,
pub tool_name: String,
pub tool_input: Option<String>,
pub tool_output: Option<String>,
pub status: String,
pub duration_ms: Option<i64>,
pub created_at: DateTime<Utc>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionSearchParams {
#[serde(default)]
pub query: Option<String>,
#[serde(default)]
pub agent_id: Option<String>,
#[serde(default)]
pub tool_name: Option<String>,
#[serde(default)]
pub source_channel: Option<String>,
#[serde(default)]
pub parent_session_id: Option<String>,
#[serde(default)]
pub status: Option<String>,
#[serde(default)]
pub thread_id: Option<String>,
#[serde(default)]
pub limit: Option<u32>,
#[serde(default)]
pub offset: Option<u32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionSearchResult {
pub sessions: Vec<SessionRecord>,
pub total: u64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn session_status_roundtrip() {
for status in [
SessionStatus::Running,
SessionStatus::Completed,
SessionStatus::Failed,
SessionStatus::Interrupted,
] {
assert_eq!(SessionStatus::parse(status.as_str()), status);
}
}
#[test]
fn session_status_parse_unknown_defaults_to_running() {
assert_eq!(SessionStatus::parse("bogus"), SessionStatus::Running);
assert_eq!(SessionStatus::parse(""), SessionStatus::Running);
}
#[test]
fn session_status_serde_roundtrip() {
let status = SessionStatus::Completed;
let json = serde_json::to_string(&status).unwrap();
assert_eq!(json, "\"completed\"");
let parsed: SessionStatus = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, status);
}
#[test]
fn session_search_params_defaults() {
let params = SessionSearchParams::default();
assert!(params.query.is_none());
assert!(params.agent_id.is_none());
assert!(params.limit.is_none());
assert!(params.offset.is_none());
}
}