Support long-running agent turns
This commit is contained in:
134
src/database.rs
134
src/database.rs
@@ -118,6 +118,7 @@ pub struct StoredMessage {
|
||||
pub reasoning: Option<String>,
|
||||
pub reasoning_complete: bool,
|
||||
pub content: String,
|
||||
pub system: bool,
|
||||
}
|
||||
|
||||
#[derive(Insertable)]
|
||||
@@ -129,6 +130,7 @@ struct NewMessage<'a> {
|
||||
reasoning: Option<&'a str>,
|
||||
reasoning_complete: bool,
|
||||
content: &'a str,
|
||||
system: bool,
|
||||
}
|
||||
|
||||
pub struct MessageDraft {
|
||||
@@ -137,6 +139,7 @@ pub struct MessageDraft {
|
||||
pub reasoning: Option<String>,
|
||||
pub reasoning_complete: bool,
|
||||
pub content: String,
|
||||
pub system: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -318,10 +321,28 @@ impl Database {
|
||||
&mut self,
|
||||
session_id: i32,
|
||||
prompt: &str,
|
||||
system_messages: &[String],
|
||||
reasoning: bool,
|
||||
) -> Result<(StoredMessage, StoredMessage), String> {
|
||||
) -> Result<Vec<StoredMessage>, String> {
|
||||
self.connection
|
||||
.transaction(|connection| {
|
||||
let mut stored = Vec::with_capacity(system_messages.len() + 2);
|
||||
for content in system_messages {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: false,
|
||||
tool: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: true,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
}
|
||||
let user = diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
@@ -330,6 +351,7 @@ impl Database {
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: prompt,
|
||||
system: false,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
@@ -341,10 +363,12 @@ impl Database {
|
||||
reasoning: reasoning.then_some(""),
|
||||
reasoning_complete: !reasoning,
|
||||
content: "",
|
||||
system: false,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
Ok((user, assistant))
|
||||
stored.extend([user, assistant]);
|
||||
Ok(stored)
|
||||
})
|
||||
.map_err(|error: diesel::result::Error| error.to_string())
|
||||
}
|
||||
@@ -353,21 +377,59 @@ impl Database {
|
||||
&mut self,
|
||||
session_id: i32,
|
||||
result: &str,
|
||||
queued_users: &[String],
|
||||
reminder: Option<&str>,
|
||||
reasoning: bool,
|
||||
) -> Result<(StoredMessage, StoredMessage), String> {
|
||||
) -> Result<Vec<StoredMessage>, String> {
|
||||
self.connection
|
||||
.transaction(|connection| {
|
||||
let tool = diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: false,
|
||||
tool: true,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: result,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
let mut stored = Vec::with_capacity(queued_users.len() + 3);
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: false,
|
||||
tool: true,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: result,
|
||||
system: false,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
for content in queued_users {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: true,
|
||||
tool: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: false,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
}
|
||||
if let Some(content) = reminder {
|
||||
stored.push(
|
||||
diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
user: false,
|
||||
tool: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content,
|
||||
system: true,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
);
|
||||
}
|
||||
let assistant = diesel::insert_into(messages::table)
|
||||
.values(NewMessage {
|
||||
session_id,
|
||||
@@ -376,10 +438,12 @@ impl Database {
|
||||
reasoning: reasoning.then_some(""),
|
||||
reasoning_complete: !reasoning,
|
||||
content: "",
|
||||
system: false,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?;
|
||||
Ok((tool, assistant))
|
||||
stored.push(assistant);
|
||||
Ok(stored)
|
||||
})
|
||||
.map_err(|error: diesel::result::Error| error.to_string())
|
||||
}
|
||||
@@ -437,6 +501,7 @@ impl Database {
|
||||
reasoning: message.reasoning.as_deref(),
|
||||
reasoning_complete: message.reasoning_complete,
|
||||
content: &message.content,
|
||||
system: message.system,
|
||||
})
|
||||
.returning(StoredMessage::as_returning())
|
||||
.get_result(connection)?,
|
||||
@@ -530,14 +595,21 @@ mod tests {
|
||||
let mut database = Database::open(&path).unwrap();
|
||||
let project = database.create_project("DS4", "/tmp/ds4-chat").unwrap();
|
||||
let session = database.create_session(project.id, "Chat").unwrap();
|
||||
let (_, assistant) = database
|
||||
.start_chat_turn(session.id, "Question", true)
|
||||
let mut opening = database
|
||||
.start_chat_turn(session.id, "Question", &["Date context".into()], true)
|
||||
.unwrap();
|
||||
let assistant = opening.pop().unwrap();
|
||||
database
|
||||
.update_message(assistant.id, Some("Reasoning"), true, "Answer")
|
||||
.unwrap();
|
||||
database
|
||||
.continue_tool_turn(session.id, "Tool result", false)
|
||||
.continue_tool_turn(
|
||||
session.id,
|
||||
"Tool result",
|
||||
&["Queued correction".into()],
|
||||
Some("Tool reminder"),
|
||||
false,
|
||||
)
|
||||
.unwrap();
|
||||
database
|
||||
.update_session_context(session.id, 1_234, 65_536, Some(12.5))
|
||||
@@ -550,17 +622,20 @@ mod tests {
|
||||
assert_eq!(projects[0].sessions[0].context_limit, 65_536);
|
||||
assert_eq!(projects[0].sessions[0].last_tokens_per_second, Some(12.5));
|
||||
let messages = reopened.load_messages(session.id).unwrap();
|
||||
assert_eq!(messages.len(), 4);
|
||||
assert!(messages[0].user);
|
||||
assert!(!messages[0].tool);
|
||||
assert_eq!(messages[0].content, "Question");
|
||||
assert_eq!(messages[1].reasoning.as_deref(), Some("Reasoning"));
|
||||
assert!(messages[1].reasoning_complete);
|
||||
assert_eq!(messages[1].content, "Answer");
|
||||
assert!(messages[2].tool);
|
||||
assert_eq!(messages[2].content, "Tool result");
|
||||
assert!(!messages[3].user);
|
||||
assert!(!messages[3].tool);
|
||||
assert_eq!(messages.len(), 7);
|
||||
assert!(messages[0].system);
|
||||
assert_eq!(messages[1].content, "Question");
|
||||
assert_eq!(messages[2].reasoning.as_deref(), Some("Reasoning"));
|
||||
assert!(messages[2].reasoning_complete);
|
||||
assert_eq!(messages[2].content, "Answer");
|
||||
assert!(messages[3].tool);
|
||||
assert_eq!(messages[3].content, "Tool result");
|
||||
assert!(messages[4].user);
|
||||
assert_eq!(messages[4].content, "Queued correction");
|
||||
assert!(messages[5].system);
|
||||
assert_eq!(messages[5].content, "Tool reminder");
|
||||
assert!(!messages[6].user);
|
||||
assert!(!messages[6].tool);
|
||||
let compacted = reopened
|
||||
.replace_with_compacted_transcript(
|
||||
session.id,
|
||||
@@ -571,6 +646,7 @@ mod tests {
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: "Recent question".into(),
|
||||
system: false,
|
||||
}],
|
||||
321,
|
||||
32_768,
|
||||
|
||||
Reference in New Issue
Block a user