feat: persist streaming turn metadata

This commit is contained in:
xiaoxixi 2026-07-17 14:37:10 +08:00
parent b7d9a41f94
commit 740d0f4f33
6 changed files with 374 additions and 116 deletions

View File

@ -3,6 +3,52 @@ use std::collections::HashMap;
use crate::providers::ToolCall; use crate::providers::ToolCall;
/// Provider-private state required to faithfully replay an assistant message.
///
/// This is durable conversation data, but it is never presentation data. UI and
/// channel projections must not serialize it to end users.
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ProviderReasoningState {
pub provider: String,
pub payload: serde_json::Value,
}
impl ProviderReasoningState {
/// Decode persisted provider state without making conversation history
/// unreadable when an old or damaged payload is encountered.
pub fn from_json_lossy(value: &str) -> Option<Self> {
serde_json::from_str(value).ok()
}
}
/// Describes whether a persisted message represents a complete model result.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CompletionStatus {
#[default]
Completed,
Cancelled,
Interrupted,
}
impl CompletionStatus {
pub fn as_str(self) -> &'static str {
match self {
Self::Completed => "completed",
Self::Cancelled => "cancelled",
Self::Interrupted => "interrupted",
}
}
pub fn from_storage(value: &str) -> Self {
match value {
"cancelled" => Self::Cancelled,
"interrupted" => Self::Interrupted,
_ => Self::Completed,
}
}
}
// ============================================================================ // ============================================================================
// ContentBlock - Multimodal content representation (OpenAI-style) // ContentBlock - Multimodal content representation (OpenAI-style)
// ============================================================================ // ============================================================================
@ -85,6 +131,15 @@ pub struct ChatMessage {
pub role: String, pub role: String,
pub content: String, pub content: String,
pub reasoning_content: Option<String>, pub reasoning_content: Option<String>,
/// Opaque state used only when replaying history to the same provider.
#[serde(skip_serializing_if = "Option::is_none")]
pub provider_state: Option<ProviderReasoningState>,
#[serde(skip_serializing_if = "Option::is_none")]
pub turn_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub iteration: Option<u32>,
#[serde(default)]
pub completion_status: CompletionStatus,
pub media_refs: Vec<MediaRef>, pub media_refs: Vec<MediaRef>,
pub timestamp: i64, pub timestamp: i64,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
@ -124,6 +179,10 @@ impl ChatMessage {
role: "user".to_string(), role: "user".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -139,6 +198,10 @@ impl ChatMessage {
role: "user".to_string(), role: "user".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs, media_refs,
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -154,6 +217,10 @@ impl ChatMessage {
role: "assistant".to_string(), role: "assistant".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -172,6 +239,10 @@ impl ChatMessage {
role: "assistant".to_string(), role: "assistant".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -187,6 +258,10 @@ impl ChatMessage {
role: "assistant".to_string(), role: "assistant".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -202,6 +277,10 @@ impl ChatMessage {
role: "system".to_string(), role: "system".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -230,6 +309,10 @@ impl ChatMessage {
role: "tool".to_string(), role: "tool".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs, media_refs,
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: Some(tool_call_id.into()), tool_call_id: Some(tool_call_id.into()),
@ -245,6 +328,10 @@ impl ChatMessage {
role: "user".to_string(), role: "user".to_string(),
content: content.into(), content: content.into(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: CompletionStatus::Completed,
media_refs: Vec::new(), media_refs: Vec::new(),
timestamp: current_timestamp(), timestamp: current_timestamp(),
tool_call_id: None, tool_call_id: None,
@ -255,6 +342,24 @@ impl ChatMessage {
} }
} }
#[cfg(test)]
mod conversation_message_tests {
use super::*;
#[test]
fn damaged_provider_state_is_ignored() {
assert!(ProviderReasoningState::from_json_lossy("not-json").is_none());
}
#[test]
fn unknown_completion_status_is_backward_compatible() {
assert_eq!(
CompletionStatus::from_storage("future-status"),
CompletionStatus::Completed
);
}
}
// ============================================================================ // ============================================================================
// InboundMessage - Message from Channel to Bus (user input) // InboundMessage - Message from Channel to Bus (user input)
// ============================================================================ // ============================================================================

View File

@ -3,8 +3,8 @@ pub mod message;
pub use dispatcher::OutboundDispatcher; pub use dispatcher::OutboundDispatcher;
pub use message::{ pub use message::{
ChatMessage, ContentBlock, ControlMessage, InboundMessage, MediaItem, MediaRef, MessageSource, ChatMessage, CompletionStatus, ContentBlock, ControlMessage, InboundMessage, MediaItem,
OutboundMessage, SourceKind, MediaRef, MessageSource, OutboundMessage, ProviderReasoningState, SourceKind,
}; };
use std::sync::Arc; use std::sync::Arc;

View File

@ -258,6 +258,12 @@ impl Session {
role: m.role, role: m.role,
content: m.content, content: m.content,
reasoning_content: m.reasoning_content, reasoning_content: m.reasoning_content,
provider_state: m.provider_state.and_then(|state| {
crate::bus::ProviderReasoningState::from_json_lossy(&state)
}),
turn_id: m.turn_id,
iteration: m.iteration.and_then(|value| u32::try_from(value).ok()),
completion_status: m.completion_status,
media_refs: m media_refs: m
.media_refs .media_refs
.map(|refs| serde_json::from_str(&refs).unwrap_or_default()) .map(|refs| serde_json::from_str(&refs).unwrap_or_default())
@ -293,6 +299,12 @@ impl Session {
role: m.role, role: m.role,
content: m.content, content: m.content,
reasoning_content: m.reasoning_content, reasoning_content: m.reasoning_content,
provider_state: m.provider_state.and_then(|state| {
crate::bus::ProviderReasoningState::from_json_lossy(&state)
}),
turn_id: m.turn_id,
iteration: m.iteration.and_then(|value| u32::try_from(value).ok()),
completion_status: m.completion_status,
media_refs: m media_refs: m
.media_refs .media_refs
.map(|refs| serde_json::from_str(&refs).unwrap_or_default()) .map(|refs| serde_json::from_str(&refs).unwrap_or_default())
@ -386,6 +398,13 @@ impl Session {
role: message.role.clone(), role: message.role.clone(),
content: message.content.clone(), content: message.content.clone(),
reasoning_content: message.reasoning_content.clone(), reasoning_content: message.reasoning_content.clone(),
provider_state: message
.provider_state
.as_ref()
.and_then(|state| serde_json::to_string(state).ok()),
turn_id: message.turn_id.clone(),
iteration: message.iteration.map(i64::from),
completion_status: message.completion_status,
media_refs: if message.media_refs.is_empty() { media_refs: if message.media_refs.is_empty() {
None None
} else { } else {

View File

@ -1,5 +1,7 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use crate::bus::CompletionStatus;
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageMeta { pub struct MessageMeta {
pub id: String, pub id: String,
@ -8,6 +10,10 @@ pub struct MessageMeta {
pub role: String, pub role: String,
pub content: String, pub content: String,
pub reasoning_content: Option<String>, pub reasoning_content: Option<String>,
pub provider_state: Option<String>,
pub turn_id: Option<String>,
pub iteration: Option<i64>,
pub completion_status: CompletionStatus,
pub media_refs: Option<String>, pub media_refs: Option<String>,
pub tool_call_id: Option<String>, pub tool_call_id: Option<String>,
pub tool_name: Option<String>, pub tool_name: Option<String>,

View File

@ -9,15 +9,21 @@ pub use background_task::BackgroundTask;
pub use error::StorageError; pub use error::StorageError;
pub use scheduler::{DeliveryPolicy, JobKind, JobRun, ScheduledJob}; pub use scheduler::{DeliveryPolicy, JobKind, JobRun, ScheduledJob};
use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteSynchronous}; use sqlx::sqlite::{
SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteRow, SqliteSynchronous,
};
use sqlx::{Pool, Row, Sqlite}; use sqlx::{Pool, Row, Sqlite};
use std::path::Path; use std::path::Path;
use tokio::time::{Duration, sleep}; use tokio::time::{Duration, sleep};
const SCHEMA_VERSION: i64 = 3; const SCHEMA_VERSION: i64 = 4;
const INSERT_MESSAGE_SQL: &str = r#" const INSERT_MESSAGE_SQL: &str = r#"
INSERT INTO messages (id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at) INSERT INTO messages (
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) id, session_id, seq, role, content, reasoning_content, provider_state,
turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#; "#;
fn insert_message_query<'a>( fn insert_message_query<'a>(
@ -31,6 +37,10 @@ fn insert_message_query<'a>(
.bind(&msg.role) .bind(&msg.role)
.bind(&msg.content) .bind(&msg.content)
.bind(&msg.reasoning_content) .bind(&msg.reasoning_content)
.bind(&msg.provider_state)
.bind(&msg.turn_id)
.bind(msg.iteration)
.bind(msg.completion_status.as_str())
.bind(&msg.media_refs) .bind(&msg.media_refs)
.bind(&msg.tool_call_id) .bind(&msg.tool_call_id)
.bind(&msg.tool_name) .bind(&msg.tool_name)
@ -39,6 +49,28 @@ fn insert_message_query<'a>(
.bind(msg.created_at) .bind(msg.created_at)
} }
fn message_meta_from_row(row: SqliteRow) -> crate::storage::message::MessageMeta {
let completion_status: String = row.get("completion_status");
crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
provider_state: row.get("provider_state"),
turn_id: row.get("turn_id"),
iteration: row.get("iteration"),
completion_status: crate::bus::CompletionStatus::from_storage(&completion_status),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
}
}
pub struct Storage { pub struct Storage {
pub(crate) pool: Pool<Sqlite>, pub(crate) pool: Pool<Sqlite>,
} }
@ -111,6 +143,10 @@ impl Storage {
tool_calls TEXT, tool_calls TEXT,
source TEXT, source TEXT,
reasoning_content TEXT, reasoning_content TEXT,
provider_state TEXT,
turn_id TEXT,
iteration INTEGER,
completion_status TEXT NOT NULL DEFAULT 'completed',
created_at INTEGER NOT NULL, created_at INTEGER NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE
) )
@ -351,6 +387,14 @@ impl Storage {
for (table, column, definition) in [ for (table, column, definition) in [
("messages", "source", "source TEXT"), ("messages", "source", "source TEXT"),
("messages", "reasoning_content", "reasoning_content TEXT"), ("messages", "reasoning_content", "reasoning_content TEXT"),
("messages", "provider_state", "provider_state TEXT"),
("messages", "turn_id", "turn_id TEXT"),
("messages", "iteration", "iteration INTEGER"),
(
"messages",
"completion_status",
"completion_status TEXT NOT NULL DEFAULT 'completed'",
),
("sessions", "archived_at", "archived_at INTEGER"), ("sessions", "archived_at", "archived_at INTEGER"),
( (
"sessions", "sessions",
@ -829,7 +873,9 @@ impl Storage {
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> { ) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query( let rows = sqlx::query(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ? AND seq >= ? WHERE session_id = ? AND seq >= ?
ORDER BY seq ASC ORDER BY seq ASC
@ -840,23 +886,7 @@ impl Storage {
.fetch_all(self.pool()) .fetch_all(self.pool())
.await?; .await?;
Ok(rows Ok(rows.into_iter().map(message_meta_from_row).collect())
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
} }
pub async fn get_message( pub async fn get_message(
@ -866,8 +896,9 @@ impl Storage {
) -> Result<Option<crate::storage::message::MessageMeta>, StorageError> { ) -> Result<Option<crate::storage::message::MessageMeta>, StorageError> {
let row = sqlx::query( let row = sqlx::query(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
tool_call_id, tool_name, tool_calls, source, created_at turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ? AND id = ? WHERE session_id = ? AND id = ?
"#, "#,
@ -877,20 +908,7 @@ impl Storage {
.fetch_optional(self.pool()) .fetch_optional(self.pool())
.await?; .await?;
Ok(row.map(|row| crate::storage::message::MessageMeta { Ok(row.map(message_meta_from_row))
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
}))
} }
pub async fn get_max_message_seq(&self, session_id: &str) -> Result<i64, StorageError> { pub async fn get_max_message_seq(&self, session_id: &str) -> Result<i64, StorageError> {
@ -912,11 +930,13 @@ impl Storage {
let limit = limit.clamp(1, 2_000); let limit = limit.clamp(1, 2_000);
let rows = sqlx::query( let rows = sqlx::query(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
tool_call_id, tool_name, tool_calls, source, created_at turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM ( FROM (
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
tool_call_id, tool_name, tool_calls, source, created_at turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ? WHERE session_id = ?
ORDER BY seq DESC ORDER BY seq DESC
@ -930,23 +950,7 @@ impl Storage {
.fetch_all(self.pool()) .fetch_all(self.pool())
.await?; .await?;
Ok(rows Ok(rows.into_iter().map(message_meta_from_row).collect())
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
} }
pub async fn load_messages_after_timestamp( pub async fn load_messages_after_timestamp(
@ -956,7 +960,9 @@ impl Storage {
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> { ) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query( let rows = sqlx::query(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ? AND created_at > ? WHERE session_id = ? AND created_at > ?
ORDER BY seq ASC ORDER BY seq ASC
@ -967,23 +973,7 @@ impl Storage {
.fetch_all(self.pool()) .fetch_all(self.pool())
.await?; .await?;
Ok(rows Ok(rows.into_iter().map(message_meta_from_row).collect())
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
} }
pub async fn query_sessions_range( pub async fn query_sessions_range(
@ -1040,7 +1030,9 @@ impl Storage {
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> { ) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query( let rows = sqlx::query(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ? WHERE session_id = ?
ORDER BY seq DESC ORDER BY seq DESC
@ -1052,23 +1044,7 @@ impl Storage {
.fetch_all(self.pool()) .fetch_all(self.pool())
.await?; .await?;
let mut messages: Vec<_> = rows let mut messages: Vec<_> = rows.into_iter().map(message_meta_from_row).collect();
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect();
messages.reverse(); messages.reverse();
Ok(messages) Ok(messages)
} }
@ -1095,7 +1071,9 @@ impl Storage {
); );
let select_sql = format!( let select_sql = format!(
r#" r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at SELECT id, session_id, seq, role, content, reasoning_content, provider_state,
turn_id, iteration, completion_status, media_refs, tool_call_id,
tool_name, tool_calls, source, created_at
FROM messages FROM messages
WHERE session_id = ?{} WHERE session_id = ?{}
ORDER BY seq ASC ORDER BY seq ASC
@ -1127,23 +1105,7 @@ impl Storage {
.fetch_all(self.pool()) .fetch_all(self.pool())
.await?; .await?;
let messages: Vec<_> = rows let messages: Vec<_> = rows.into_iter().map(message_meta_from_row).collect();
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect();
Ok((messages, total)) Ok((messages, total))
} }
@ -1616,7 +1578,17 @@ mod tests {
let storage = Storage::new(&db_path).await.unwrap(); let storage = Storage::new(&db_path).await.unwrap();
for (table, expected) in [ for (table, expected) in [
("messages", vec!["source", "reasoning_content"]), (
"messages",
vec![
"source",
"reasoning_content",
"provider_state",
"turn_id",
"iteration",
"completion_status",
],
),
( (
"sessions", "sessions",
vec![ vec![
@ -1660,6 +1632,81 @@ mod tests {
} }
} }
#[tokio::test]
async fn v3_migration_preserves_existing_reasoning_and_defaults_completion() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("v3.db");
let pool = SqlitePoolOptions::new()
.connect_with(
SqliteConnectOptions::new()
.filename(&db_path)
.create_if_missing(true),
)
.await
.unwrap();
sqlx::query(
r#"
CREATE TABLE sessions (
id TEXT PRIMARY KEY, channel TEXT NOT NULL, chat_id TEXT NOT NULL,
dialog_id TEXT NOT NULL, title TEXT NOT NULL DEFAULT 'new',
created_at INTEGER NOT NULL, last_active_at INTEGER NOT NULL,
message_count INTEGER DEFAULT 0, routing_info TEXT, archived_at INTEGER,
deleted_at INTEGER, last_consolidated_at INTEGER,
last_compressed_message_at INTEGER,
UNIQUE(channel, chat_id, dialog_id)
)
"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"
CREATE TABLE messages (
id TEXT PRIMARY KEY, session_id TEXT NOT NULL, seq INTEGER NOT NULL,
role TEXT NOT NULL, content TEXT NOT NULL, reasoning_content TEXT,
media_refs TEXT, tool_call_id TEXT, tool_name TEXT, tool_calls TEXT,
source TEXT, created_at INTEGER NOT NULL
)
"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
"INSERT INTO sessions (id, channel, chat_id, dialog_id, created_at, last_active_at) VALUES ('cli:c:d', 'cli', 'c', 'd', 1, 1)",
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
"INSERT INTO messages (id, session_id, seq, role, content, reasoning_content, created_at) VALUES ('m1', 'cli:c:d', 1, 'assistant', 'answer', 'existing reasoning', 1)",
)
.execute(&pool)
.await
.unwrap();
sqlx::query("PRAGMA user_version = 3")
.execute(&pool)
.await
.unwrap();
drop(pool);
let storage = Storage::new(&db_path).await.unwrap();
let messages = storage.load_messages("cli:c:d", 0).await.unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(
messages[0].reasoning_content.as_deref(),
Some("existing reasoning")
);
assert_eq!(messages[0].provider_state, None);
assert_eq!(messages[0].turn_id, None);
assert_eq!(messages[0].iteration, None);
assert_eq!(
messages[0].completion_status,
crate::bus::CompletionStatus::Completed
);
}
#[tokio::test] #[tokio::test]
async fn test_upsert_and_get_session() { async fn test_upsert_and_get_session() {
let (storage, _dir) = create_test_storage().await; let (storage, _dir) = create_test_storage().await;
@ -1784,6 +1831,10 @@ mod tests {
role: "user".to_string(), role: "user".to_string(),
content: "你好".to_string(), content: "你好".to_string(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: crate::bus::CompletionStatus::Completed,
media_refs: None, media_refs: None,
tool_call_id: None, tool_call_id: None,
tool_name: None, tool_name: None,
@ -1821,6 +1872,67 @@ mod tests {
assert_eq!(recent[1].seq, 5); assert_eq!(recent[1].seq, 5);
} }
#[tokio::test]
async fn streaming_message_metadata_round_trips() {
let (storage, _dir) = create_test_storage().await;
let session_meta = crate::storage::session::SessionMeta {
id: "cli_chat:stream:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "stream".to_string(),
dialog_id: "dialog1".to_string(),
title: "Stream metadata".to_string(),
created_at: 1000,
last_active_at: 1000,
message_count: 0,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&session_meta).await.unwrap();
let message = crate::storage::message::MessageMeta {
id: "assistant-1".to_string(),
session_id: session_meta.id.clone(),
seq: 1,
role: "assistant".to_string(),
content: "partial".to_string(),
reasoning_content: Some("visible reasoning".to_string()),
provider_state: Some(
serde_json::json!({"provider":"anthropic","payload":{"signature":"opaque"}})
.to_string(),
),
turn_id: Some("turn-1".to_string()),
iteration: Some(2),
completion_status: crate::bus::CompletionStatus::Interrupted,
media_refs: None,
tool_call_id: None,
tool_name: None,
tool_calls: None,
source: None,
created_at: 1001,
};
storage
.append_message(&session_meta.id, &message)
.await
.unwrap();
let loaded = storage.load_messages(&session_meta.id, 0).await.unwrap();
assert_eq!(loaded.len(), 1);
assert_eq!(
loaded[0].reasoning_content.as_deref(),
Some("visible reasoning")
);
assert_eq!(loaded[0].provider_state, message.provider_state);
assert_eq!(loaded[0].turn_id.as_deref(), Some("turn-1"));
assert_eq!(loaded[0].iteration, Some(2));
assert_eq!(
loaded[0].completion_status,
crate::bus::CompletionStatus::Interrupted
);
}
#[tokio::test] #[tokio::test]
async fn test_persist_message_batch_is_atomic() { async fn test_persist_message_batch_is_atomic() {
let (storage, _dir) = create_test_storage().await; let (storage, _dir) = create_test_storage().await;
@ -1848,6 +1960,10 @@ mod tests {
role: "assistant".to_string(), role: "assistant".to_string(),
content: "must roll back".to_string(), content: "must roll back".to_string(),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: crate::bus::CompletionStatus::Completed,
media_refs: None, media_refs: None,
tool_call_id: None, tool_call_id: None,
tool_name: None, tool_name: None,

View File

@ -358,6 +358,10 @@ mod tests {
}, },
content: format!("消息内容 {}", i), content: format!("消息内容 {}", i),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: crate::bus::CompletionStatus::Completed,
media_refs: None, media_refs: None,
tool_call_id: None, tool_call_id: None,
tool_name: None, tool_name: None,
@ -420,6 +424,10 @@ mod tests {
}, },
content: format!("消息内容 {}", i), content: format!("消息内容 {}", i),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: crate::bus::CompletionStatus::Completed,
media_refs: None, media_refs: None,
tool_call_id: None, tool_call_id: None,
tool_name: None, tool_name: None,
@ -476,6 +484,10 @@ mod tests {
role: "user".to_string(), role: "user".to_string(),
content: format!("消息内容 {}", i), content: format!("消息内容 {}", i),
reasoning_content: None, reasoning_content: None,
provider_state: None,
turn_id: None,
iteration: None,
completion_status: crate::bus::CompletionStatus::Completed,
media_refs: None, media_refs: None,
tool_call_id: None, tool_call_id: None,
tool_name: None, tool_name: None,