318 lines
9.3 KiB
Rust
318 lines
9.3 KiB
Rust
use crate::error::Result;
|
|
use rusqlite::{params, Connection};
|
|
use serde::Serialize;
|
|
|
|
#[derive(Debug, Clone, Serialize)]
|
|
pub struct Conversation {
|
|
pub id: String,
|
|
pub title: String,
|
|
pub model_id: Option<String>,
|
|
pub system_prompt: Option<String>,
|
|
pub created_at: String,
|
|
pub updated_at: String,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize)]
|
|
pub struct Message {
|
|
pub id: String,
|
|
pub conversation_id: String,
|
|
pub role: String,
|
|
pub content: String,
|
|
pub tokens_in: Option<i64>,
|
|
pub tokens_out: Option<i64>,
|
|
pub elapsed_ms: Option<i64>,
|
|
pub first_token_ms: Option<i64>,
|
|
pub images: Vec<String>,
|
|
pub created_at: String,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize)]
|
|
pub struct MessageVersion {
|
|
pub id: String,
|
|
pub message_id: String,
|
|
pub content: String,
|
|
pub tokens_out: Option<i64>,
|
|
pub elapsed_ms: Option<i64>,
|
|
pub first_token_ms: Option<i64>,
|
|
pub seq: i64,
|
|
pub created_at: String,
|
|
}
|
|
|
|
pub fn create_conversation(
|
|
db: &Connection,
|
|
id: &str,
|
|
title: &str,
|
|
model_id: Option<&str>,
|
|
system_prompt: Option<&str>,
|
|
) -> Result<()> {
|
|
db.execute(
|
|
"INSERT INTO conversations (id, title, model_id, system_prompt) VALUES (?1, ?2, ?3, ?4)",
|
|
params![id, title, model_id, system_prompt],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn list_conversations(db: &Connection) -> Result<Vec<Conversation>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, title, model_id, system_prompt, created_at, updated_at
|
|
FROM conversations ORDER BY updated_at DESC",
|
|
)?;
|
|
let rows = stmt.query_map([], row_to_conversation)?;
|
|
let mut out = Vec::new();
|
|
for row in rows {
|
|
out.push(row?);
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
pub fn get_conversation(db: &Connection, id: &str) -> Result<Option<Conversation>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, title, model_id, system_prompt, created_at, updated_at
|
|
FROM conversations WHERE id = ?1",
|
|
)?;
|
|
let mut rows = stmt.query_map(params![id], row_to_conversation)?;
|
|
match rows.next() {
|
|
Some(row) => Ok(Some(row?)),
|
|
None => Ok(None),
|
|
}
|
|
}
|
|
|
|
pub fn touch_conversation(db: &Connection, id: &str) -> Result<()> {
|
|
db.execute(
|
|
"UPDATE conversations SET updated_at = datetime('now') WHERE id = ?1",
|
|
params![id],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn update_conversation_model(db: &Connection, id: &str, model_id: &str) -> Result<()> {
|
|
db.execute(
|
|
"UPDATE conversations SET model_id = ?1 WHERE id = ?2",
|
|
params![model_id, id],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn delete_conversation(db: &Connection, id: &str) -> Result<()> {
|
|
db.execute("DELETE FROM conversations WHERE id = ?1", params![id])?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn add_message(
|
|
db: &Connection,
|
|
conversation_id: &str,
|
|
role: &str,
|
|
content: &str,
|
|
tokens_in: Option<i64>,
|
|
tokens_out: Option<i64>,
|
|
elapsed_ms: Option<i64>,
|
|
first_token_ms: Option<i64>,
|
|
images: &[String],
|
|
) -> Result<String> {
|
|
let id = uuid::Uuid::new_v4().to_string();
|
|
let images_json = serde_json::to_string(images).unwrap_or_else(|_| "[]".to_string());
|
|
db.execute(
|
|
"INSERT INTO messages (id, conversation_id, role, content, tokens_in, tokens_out, elapsed_ms, first_token_ms, images_json)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
|
|
params![
|
|
id,
|
|
conversation_id,
|
|
role,
|
|
content,
|
|
tokens_in,
|
|
tokens_out,
|
|
elapsed_ms,
|
|
first_token_ms,
|
|
images_json
|
|
],
|
|
)?;
|
|
Ok(id)
|
|
}
|
|
|
|
pub fn list_messages(db: &Connection, conversation_id: &str) -> Result<Vec<Message>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, conversation_id, role, content, tokens_in, tokens_out, elapsed_ms, first_token_ms, images_json, created_at
|
|
FROM messages WHERE conversation_id = ?1 ORDER BY rowid ASC",
|
|
)?;
|
|
let rows = stmt.query_map(params![conversation_id], row_to_message)?;
|
|
let mut out = Vec::new();
|
|
for row in rows {
|
|
out.push(row?);
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
pub fn get_message(db: &Connection, id: &str) -> Result<Option<Message>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, conversation_id, role, content, tokens_in, tokens_out, elapsed_ms, first_token_ms, images_json, created_at
|
|
FROM messages WHERE id = ?1",
|
|
)?;
|
|
let mut rows = stmt.query_map(params![id], row_to_message)?;
|
|
match rows.next() {
|
|
Some(row) => Ok(Some(row?)),
|
|
None => Ok(None),
|
|
}
|
|
}
|
|
|
|
pub fn get_message_with_rowid(db: &Connection, id: &str) -> Result<Option<(i64, Message)>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, conversation_id, role, content, tokens_in, tokens_out, elapsed_ms, first_token_ms, images_json, created_at, rowid
|
|
FROM messages WHERE id = ?1",
|
|
)?;
|
|
let mut rows = stmt.query_map(params![id], |row| {
|
|
let rowid: i64 = row.get(10)?;
|
|
let msg = row_to_message(row)?;
|
|
Ok((rowid, msg))
|
|
})?;
|
|
match rows.next() {
|
|
Some(row) => Ok(Some(row?)),
|
|
None => Ok(None),
|
|
}
|
|
}
|
|
|
|
pub fn update_message_content(db: &Connection, id: &str, content: &str) -> Result<()> {
|
|
db.execute(
|
|
"UPDATE messages SET content = ?1 WHERE id = ?2",
|
|
params![content, id],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn update_message_stats(
|
|
db: &Connection,
|
|
id: &str,
|
|
tokens_out: Option<i64>,
|
|
elapsed_ms: Option<i64>,
|
|
first_token_ms: Option<i64>,
|
|
) -> Result<()> {
|
|
db.execute(
|
|
"UPDATE messages SET tokens_out = ?1, elapsed_ms = ?2, first_token_ms = ?3 WHERE id = ?4",
|
|
params![tokens_out, elapsed_ms, first_token_ms, id],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
/// 删除某条消息(含)之后的所有消息。
|
|
pub fn delete_messages_after(db: &Connection, conversation_id: &str, rowid: i64) -> Result<()> {
|
|
db.execute(
|
|
"DELETE FROM messages WHERE conversation_id = ?1 AND rowid >= ?2",
|
|
params![conversation_id, rowid],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
/// 删除某条消息之后(不含该消息)的所有消息。
|
|
pub fn delete_messages_strictly_after(
|
|
db: &Connection,
|
|
conversation_id: &str,
|
|
rowid: i64,
|
|
) -> Result<()> {
|
|
db.execute(
|
|
"DELETE FROM messages WHERE conversation_id = ?1 AND rowid > ?2",
|
|
params![conversation_id, rowid],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
pub fn save_message_version(
|
|
db: &Connection,
|
|
message_id: &str,
|
|
content: &str,
|
|
tokens_out: Option<i64>,
|
|
elapsed_ms: Option<i64>,
|
|
first_token_ms: Option<i64>,
|
|
) -> Result<String> {
|
|
let id = uuid::Uuid::new_v4().to_string();
|
|
let seq: i64 = db.query_row(
|
|
"SELECT COUNT(*) FROM message_versions WHERE message_id = ?1",
|
|
params![message_id],
|
|
|row| row.get::<_, i64>(0),
|
|
)? + 1;
|
|
db.execute(
|
|
"INSERT INTO message_versions (id, message_id, content, tokens_out, elapsed_ms, first_token_ms, seq)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
|
|
params![id, message_id, content, tokens_out, elapsed_ms, first_token_ms, seq],
|
|
)?;
|
|
Ok(id)
|
|
}
|
|
|
|
pub fn list_message_versions(db: &Connection, message_id: &str) -> Result<Vec<MessageVersion>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, message_id, content, tokens_out, elapsed_ms, first_token_ms, seq, created_at
|
|
FROM message_versions WHERE message_id = ?1 ORDER BY seq ASC",
|
|
)?;
|
|
let rows = stmt.query_map(params![message_id], row_to_version)?;
|
|
let mut out = Vec::new();
|
|
for row in rows {
|
|
out.push(row?);
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
pub fn get_message_version(db: &Connection, version_id: &str) -> Result<Option<MessageVersion>> {
|
|
let mut stmt = db.prepare(
|
|
"SELECT id, message_id, content, tokens_out, elapsed_ms, first_token_ms, seq, created_at
|
|
FROM message_versions WHERE id = ?1",
|
|
)?;
|
|
let mut rows = stmt.query_map(params![version_id], row_to_version)?;
|
|
match rows.next() {
|
|
Some(row) => Ok(Some(row?)),
|
|
None => Ok(None),
|
|
}
|
|
}
|
|
|
|
pub fn apply_message_version(
|
|
db: &Connection,
|
|
message_id: &str,
|
|
content: &str,
|
|
tokens_out: Option<i64>,
|
|
elapsed_ms: Option<i64>,
|
|
first_token_ms: Option<i64>,
|
|
) -> Result<()> {
|
|
db.execute(
|
|
"UPDATE messages SET content = ?1, tokens_out = ?2, elapsed_ms = ?3, first_token_ms = ?4 WHERE id = ?5",
|
|
params![content, tokens_out, elapsed_ms, first_token_ms, message_id],
|
|
)?;
|
|
Ok(())
|
|
}
|
|
|
|
fn row_to_conversation(row: &rusqlite::Row<'_>) -> rusqlite::Result<Conversation> {
|
|
Ok(Conversation {
|
|
id: row.get(0)?,
|
|
title: row.get(1)?,
|
|
model_id: row.get(2)?,
|
|
system_prompt: row.get(3)?,
|
|
created_at: row.get(4)?,
|
|
updated_at: row.get(5)?,
|
|
})
|
|
}
|
|
|
|
fn row_to_message(row: &rusqlite::Row<'_>) -> rusqlite::Result<Message> {
|
|
let images_json: String = row.get(8)?;
|
|
Ok(Message {
|
|
id: row.get(0)?,
|
|
conversation_id: row.get(1)?,
|
|
role: row.get(2)?,
|
|
content: row.get(3)?,
|
|
tokens_in: row.get(4)?,
|
|
tokens_out: row.get(5)?,
|
|
elapsed_ms: row.get(6)?,
|
|
first_token_ms: row.get(7)?,
|
|
images: serde_json::from_str(&images_json).unwrap_or_default(),
|
|
created_at: row.get(9)?,
|
|
})
|
|
}
|
|
|
|
fn row_to_version(row: &rusqlite::Row<'_>) -> rusqlite::Result<MessageVersion> {
|
|
Ok(MessageVersion {
|
|
id: row.get(0)?,
|
|
message_id: row.get(1)?,
|
|
content: row.get(2)?,
|
|
tokens_out: row.get(3)?,
|
|
elapsed_ms: row.get(4)?,
|
|
first_token_ms: row.get(5)?,
|
|
seq: row.get(6)?,
|
|
created_at: row.get(7)?,
|
|
})
|
|
}
|