Files
xianren_studio/crates/core/src/sessions.rs
T

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)?,
})
}