Finish long-running agent parity
This commit is contained in:
162
src/database.rs
162
src/database.rs
@@ -119,6 +119,8 @@ pub struct StoredMessage {
|
||||
pub reasoning_complete: bool,
|
||||
pub content: String,
|
||||
pub system: bool,
|
||||
pub compaction: bool,
|
||||
pub compaction_tail_start: Option<i32>,
|
||||
}
|
||||
|
||||
#[derive(Insertable)]
|
||||
@@ -131,15 +133,8 @@ struct NewMessage<'a> {
|
||||
reasoning_complete: bool,
|
||||
content: &'a str,
|
||||
system: bool,
|
||||
}
|
||||
|
||||
pub struct MessageDraft {
|
||||
pub user: bool,
|
||||
pub tool: bool,
|
||||
pub reasoning: Option<String>,
|
||||
pub reasoning_complete: bool,
|
||||
pub content: String,
|
||||
pub system: bool,
|
||||
compaction: bool,
|
||||
compaction_tail_start: Option<i32>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -338,6 +333,8 @@ impl Database {
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: true,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
@@ -352,6 +349,8 @@ impl Database {
|
||||
reasoning_complete: true,
|
||||
content: prompt,
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
@@ -364,6 +363,8 @@ impl Database {
|
||||
reasoning_complete: !reasoning,
|
||||
content: "",
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
@@ -377,13 +378,13 @@ impl Database {
|
||||
&mut self,
|
||||
session_id: i32,
|
||||
result: &str,
|
||||
queued_users: &[String],
|
||||
reminder: Option<&str>,
|
||||
queued_user: Option<&str>,
|
||||
system_messages: &[String],
|
||||
reasoning: bool,
|
||||
) -> Result<Vec<StoredMessage>, String> {
|
||||
self.connection
|
||||
.transaction(|connection| {
|
||||
let mut stored = Vec::with_capacity(queued_users.len() + 3);
|
||||
let mut stored = Vec::with_capacity(system_messages.len() + 3);
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
@@ -394,11 +395,13 @@ impl Database {
|
||||
reasoning_complete: true,
|
||||
content: result,
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
for content in queued_users {
|
||||
if let Some(content) = queued_user {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
@@ -409,12 +412,14 @@ impl Database {
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
}
|
||||
if let Some(content) = reminder {
|
||||
for content in system_messages {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
@@ -425,6 +430,8 @@ impl Database {
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: true,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
@@ -439,6 +446,8 @@ impl Database {
|
||||
reasoning_complete: !reasoning,
|
||||
content: "",
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
@@ -466,11 +475,12 @@ impl Database {
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
pub fn replace_with_compacted_transcript(
|
||||
pub fn record_compaction(
|
||||
&mut self,
|
||||
session_id: i32,
|
||||
summary: &str,
|
||||
tail: &[MessageDraft],
|
||||
tail_start: Option<i32>,
|
||||
running_jobs: Option<&str>,
|
||||
context_used: u32,
|
||||
context_limit: u32,
|
||||
) -> Result<Vec<StoredMessage>, String> {
|
||||
@@ -480,8 +490,6 @@ impl Database {
|
||||
.map_err(|_| "Context limit is too large to save".to_owned())?;
|
||||
self.connection
|
||||
.transaction(|connection| {
|
||||
diesel::delete(messages::table.filter(messages::session_id.eq(session_id)))
|
||||
.execute(connection)?;
|
||||
diesel::update(sessions::table.find(session_id))
|
||||
.set((
|
||||
sessions::compacted_summary.eq(Some(summary)),
|
||||
@@ -490,18 +498,36 @@ impl Database {
|
||||
sessions::last_tokens_per_second.eq(None::<f32>),
|
||||
))
|
||||
.execute(connection)?;
|
||||
let mut stored = Vec::with_capacity(tail.len());
|
||||
for message in tail {
|
||||
let mut stored = Vec::with_capacity(2);
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: false,
|
||||
tool: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: summary,
|
||||
system: true,
|
||||
compaction: true,
|
||||
compaction_tail_start: tail_start,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
if let Some(content) = running_jobs {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: message.user,
|
||||
tool: message.tool,
|
||||
reasoning: message.reasoning.as_deref(),
|
||||
reasoning_complete: message.reasoning_complete,
|
||||
content: &message.content,
|
||||
system: message.system,
|
||||
user: false,
|
||||
tool: true,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: false,
|
||||
compaction: false,
|
||||
compaction_tail_start: None,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
@@ -586,7 +612,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_messages_survive_reopen_and_follow_session_deletion() {
|
||||
fn chat_and_compaction_history_survive_reopen_and_session_deletion() {
|
||||
let id = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
@@ -606,8 +632,8 @@ mod tests {
|
||||
.continue_tool_turn(
|
||||
session.id,
|
||||
"Tool result",
|
||||
&["Queued correction".into()],
|
||||
Some("Tool reminder"),
|
||||
Some("Queued correction"),
|
||||
&["Tool reminder".into()],
|
||||
false,
|
||||
)
|
||||
.unwrap();
|
||||
@@ -636,36 +662,86 @@ mod tests {
|
||||
assert_eq!(messages[5].content, "Tool reminder");
|
||||
assert!(!messages[6].user);
|
||||
assert!(!messages[6].tool);
|
||||
let compacted = reopened
|
||||
.replace_with_compacted_transcript(
|
||||
let first = reopened
|
||||
.record_compaction(
|
||||
session.id,
|
||||
"Keep the active task.",
|
||||
&[MessageDraft {
|
||||
user: true,
|
||||
tool: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: "Recent question".into(),
|
||||
system: false,
|
||||
}],
|
||||
"First durable state.",
|
||||
Some(messages[4].id),
|
||||
Some("bash job=1 status=running"),
|
||||
321,
|
||||
32_768,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(compacted.len(), 1);
|
||||
assert_eq!(first.len(), 2);
|
||||
assert!(first[0].compaction);
|
||||
assert_eq!(first[0].compaction_tail_start, Some(messages[4].id));
|
||||
drop(reopened);
|
||||
let mut reopened = Database::open(&path).unwrap();
|
||||
assert_eq!(
|
||||
reopened.load_projects().unwrap()[0].sessions[0]
|
||||
.compacted_summary
|
||||
.as_deref(),
|
||||
Some("Keep the active task.")
|
||||
Some("First durable state.")
|
||||
);
|
||||
assert_eq!(
|
||||
reopened.load_projects().unwrap()[0].sessions[0].context_used,
|
||||
321
|
||||
);
|
||||
assert_eq!(reopened.load_messages(session.id).unwrap().len(), 1);
|
||||
let messages = reopened.load_messages(session.id).unwrap();
|
||||
assert_eq!(messages.len(), 9);
|
||||
assert_eq!(messages[1].content, "Question");
|
||||
assert!(messages[7].compaction);
|
||||
assert_eq!(messages[7].content, "First durable state.");
|
||||
assert_eq!(messages[8].content, "bash job=1 status=running");
|
||||
reopened
|
||||
.continue_tool_turn(session.id, "Reloaded tool result", None, &[], false)
|
||||
.unwrap();
|
||||
let continued = reopened.load_messages(session.id).unwrap();
|
||||
let second_tail = continued[9].id;
|
||||
reopened
|
||||
.record_compaction(
|
||||
session.id,
|
||||
"Second durable state.",
|
||||
Some(second_tail),
|
||||
None,
|
||||
222,
|
||||
32_768,
|
||||
)
|
||||
.unwrap();
|
||||
reopened
|
||||
.continue_tool_turn(session.id, "Final tool result", None, &[], false)
|
||||
.unwrap();
|
||||
let continued = reopened.load_messages(session.id).unwrap();
|
||||
let third_tail = continued[12].id;
|
||||
reopened
|
||||
.record_compaction(
|
||||
session.id,
|
||||
"Third durable state.",
|
||||
Some(third_tail),
|
||||
None,
|
||||
111,
|
||||
32_768,
|
||||
)
|
||||
.unwrap();
|
||||
reopened
|
||||
.start_chat_turn(session.id, "After third compaction", &[], false)
|
||||
.unwrap();
|
||||
drop(reopened);
|
||||
|
||||
let mut reopened = Database::open(&path).unwrap();
|
||||
let history = reopened.load_messages(session.id).unwrap();
|
||||
assert_eq!(history.len(), 17);
|
||||
assert_eq!(history[1].content, "Question");
|
||||
let markers = history
|
||||
.iter()
|
||||
.filter(|message| message.compaction)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(markers.len(), 3);
|
||||
assert_eq!(markers[0].content, "First durable state.");
|
||||
assert_eq!(markers[1].content, "Second durable state.");
|
||||
assert_eq!(markers[2].content, "Third durable state.");
|
||||
assert_eq!(markers[2].compaction_tail_start, Some(third_tail));
|
||||
assert_eq!(history[15].content, "After third compaction");
|
||||
reopened.delete_session(session.id).unwrap();
|
||||
assert!(reopened.load_messages(session.id).unwrap().is_empty());
|
||||
drop(reopened);
|
||||
|
||||
Reference in New Issue
Block a user