use async_trait::async_trait;
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use turnframe_core::ids::{AccountId, ConversationId, TurnId};
use turnframe_core::replay::TurnPhase;
use turnframe_core::response::AssistantTurn;
use turnframe_core::turn::TurnInput;
use crate::error::StoreError;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ConversationRecord {
pub id: ConversationId,
pub account_id: AccountId,
pub created_at: DateTime<Utc>,
#[serde(default)]
pub metadata: serde_json::Value,
}
impl ConversationRecord {
#[must_use]
pub fn new(id: ConversationId, account_id: AccountId, created_at: DateTime<Utc>) -> Self {
Self {
id,
account_id,
created_at,
metadata: serde_json::Value::Null,
}
}
#[must_use]
pub fn with_metadata(mut self, metadata: serde_json::Value) -> Self {
self.metadata = metadata;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct StoredUserTurn {
pub input: TurnInput,
pub received_at: DateTime<Utc>,
}
impl StoredUserTurn {
#[must_use]
pub fn new(input: TurnInput, received_at: DateTime<Utc>) -> Self {
Self { input, received_at }
}
#[must_use]
pub fn account_id(&self) -> &AccountId {
&self.input.actor.account_id
}
#[must_use]
pub fn turn_id(&self) -> TurnId {
self.input.turn_id
}
#[must_use]
pub fn conversation_id(&self) -> ConversationId {
self.input.conversation_id
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct StoredTurn {
pub user: StoredUserTurn,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub assistant: Option<AssistantTurn>,
pub phase: TurnPhase,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TurnPhaseMarker {
pub account_id: AccountId,
pub conversation_id: ConversationId,
pub turn_id: TurnId,
pub phase: TurnPhase,
pub updated_at: DateTime<Utc>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum RecoveryScope {
Account(AccountId),
AllAccounts,
}
impl RecoveryScope {
#[must_use]
pub fn includes(&self, account: &AccountId) -> bool {
match self {
Self::Account(scoped) => scoped == account,
Self::AllAccounts => true,
}
}
}
#[async_trait]
pub trait ConversationReader: Send + Sync {
async fn load_conversation(
&self,
account: &AccountId,
id: &ConversationId,
) -> Result<ConversationRecord, StoreError>;
async fn load_recent_turns(
&self,
account: &AccountId,
conversation: &ConversationId,
limit: usize,
) -> Result<Vec<StoredTurn>, StoreError>;
async fn load_turn(
&self,
account: &AccountId,
turn_id: &TurnId,
) -> Result<StoredTurn, StoreError>;
async fn turn_phase(
&self,
account: &AccountId,
turn_id: &TurnId,
) -> Result<TurnPhaseMarker, StoreError>;
async fn list_unfinished_turns(
&self,
scope: RecoveryScope,
limit: usize,
) -> Result<Vec<TurnPhaseMarker>, StoreError>;
}
#[async_trait]
pub trait ConversationWriter: Send + Sync {
async fn create_conversation(&self, record: ConversationRecord) -> Result<(), StoreError>;
async fn append_user_turn(&self, turn: StoredUserTurn) -> Result<(), StoreError>;
async fn append_assistant_turn(
&self,
account: &AccountId,
turn: AssistantTurn,
) -> Result<(), StoreError>;
async fn set_turn_phase(
&self,
account: &AccountId,
turn_id: &TurnId,
phase: TurnPhase,
) -> Result<TurnPhaseMarker, StoreError>;
}
pub trait ConversationStore: ConversationReader + ConversationWriter {}
impl<T: ConversationReader + ConversationWriter + ?Sized> ConversationStore for T {}
#[cfg(test)]
mod tests {
use super::*;
use turnframe_core::locale::Locale;
use turnframe_core::turn::ActorContext;
#[test]
fn stored_user_turn_accessors() {
let input = TurnInput {
turn_id: TurnId::nil(),
conversation_id: ConversationId::nil(),
actor: ActorContext::new("acct", "u"),
text: Some("hi".into()),
interaction_response: None,
attachments: Vec::new(),
origin: None,
locale: Locale::from("en"),
effort: None,
};
let turn = StoredUserTurn::new(input, DateTime::<Utc>::UNIX_EPOCH);
assert_eq!(turn.account_id(), &AccountId::from("acct"));
assert_eq!(turn.turn_id(), TurnId::nil());
assert_eq!(turn.conversation_id(), ConversationId::nil());
}
#[test]
fn recovery_scope_membership() {
let a = AccountId::from("a");
assert!(RecoveryScope::AllAccounts.includes(&a));
assert!(RecoveryScope::Account(a.clone()).includes(&a));
assert!(!RecoveryScope::Account(AccountId::from("b")).includes(&a));
}
#[test]
fn conversation_record_round_trips() {
let record = ConversationRecord::new(
ConversationId::nil(),
AccountId::from("a"),
DateTime::<Utc>::UNIX_EPOCH,
)
.with_metadata(serde_json::json!({"title": "t"}));
let json = serde_json::to_string(&record).unwrap();
assert_eq!(
serde_json::from_str::<ConversationRecord>(&json).unwrap(),
record
);
}
}