1use async_trait::async_trait;
16use chrono::{DateTime, Utc};
17use sqlx::PgConnection;
18use turnframe_core::ids::{AccountId, ConversationId, TurnId};
19use turnframe_core::replay::TurnPhase;
20use turnframe_core::response::AssistantTurn;
21use turnframe_store::conversation::{
22 ConversationReader, ConversationRecord, ConversationWriter, RecoveryScope, StoredTurn,
23 StoredUserTurn, TurnPhaseMarker,
24};
25use turnframe_store::error::{StoreError, identity_mismatch};
26use uuid::Uuid;
27
28use crate::codec::{column, from_json, from_label, label, limit_to_sql, now, to_json};
29use crate::error::store_error;
30use crate::store::{PgStores, commit};
31
32const TERMINAL_PHASES: &str = "('delivered', 'failed')";
34
35pub(crate) async fn create_conversation(
37 conn: &mut PgConnection,
38 record: ConversationRecord,
39) -> Result<(), StoreError> {
40 sqlx::query(
41 "INSERT INTO tf_conversation (account_id, conversation_id, created_at, metadata)
42 VALUES ($1, $2, $3, $4)",
43 )
44 .bind(record.account_id.as_str())
45 .bind(record.id.as_uuid())
46 .bind(record.created_at)
47 .bind(to_json(&record.metadata)?)
48 .execute(conn)
49 .await
50 .map_err(|error| store_error(&error))?;
51 Ok(())
52}
53
54pub(crate) async fn load_conversation(
56 conn: &mut PgConnection,
57 account: &AccountId,
58 id: &ConversationId,
59) -> Result<ConversationRecord, StoreError> {
60 let row = sqlx::query(
61 "SELECT created_at, metadata FROM tf_conversation
62 WHERE account_id = $1 AND conversation_id = $2",
63 )
64 .bind(account.as_str())
65 .bind(id.as_uuid())
66 .fetch_optional(conn)
67 .await
68 .map_err(|error| store_error(&error))?
69 .ok_or(StoreError::NotFound)?;
70 Ok(ConversationRecord {
71 id: *id,
72 account_id: account.clone(),
73 created_at: column(&row, "created_at")?,
74 metadata: column(&row, "metadata")?,
75 })
76}
77
78pub(crate) async fn append_user_turn(
84 conn: &mut PgConnection,
85 turn: StoredUserTurn,
86 at: DateTime<Utc>,
87) -> Result<(), StoreError> {
88 let account = turn.account_id().clone();
89 let turn_id = turn.turn_id();
90 let conversation = turn.conversation_id();
91 let inserted = sqlx::query(
92 "INSERT INTO tf_turn (account_id, turn_id, conversation_id, received_at, user_turn)
93 SELECT $1, $2, $3, $4, $5
94 WHERE EXISTS (
95 SELECT 1 FROM tf_conversation
96 WHERE account_id = $1 AND conversation_id = $3
97 )",
98 )
99 .bind(account.as_str())
100 .bind(turn_id.as_uuid())
101 .bind(conversation.as_uuid())
102 .bind(turn.received_at)
103 .bind(to_json(&turn.input)?)
104 .execute(&mut *conn)
105 .await
106 .map_err(|error| store_error(&error))?;
107 if inserted.rows_affected() == 0 {
108 return Err(StoreError::NotFound);
109 }
110 sqlx::query(
111 "INSERT INTO tf_turn_phase (account_id, turn_id, conversation_id, phase, updated_at)
112 VALUES ($1, $2, $3, $4, $5)",
113 )
114 .bind(account.as_str())
115 .bind(turn_id.as_uuid())
116 .bind(conversation.as_uuid())
117 .bind(label(&TurnPhase::Received)?)
118 .bind(at)
119 .execute(conn)
120 .await
121 .map_err(|error| store_error(&error))?;
122 Ok(())
123}
124
125pub(crate) async fn append_assistant_turn(
127 conn: &mut PgConnection,
128 account: &AccountId,
129 turn: AssistantTurn,
130) -> Result<(), StoreError> {
131 let attached = sqlx::query(
132 "UPDATE tf_turn SET assistant_turn = $4
133 WHERE account_id = $1 AND turn_id = $2 AND conversation_id = $3
134 AND assistant_turn IS NULL",
135 )
136 .bind(account.as_str())
137 .bind(turn.turn_id.as_uuid())
138 .bind(turn.conversation_id.as_uuid())
139 .bind(to_json(&turn)?)
140 .execute(&mut *conn)
141 .await
142 .map_err(|error| store_error(&error))?;
143 if attached.rows_affected() == 1 {
144 return Ok(());
145 }
146 let row = sqlx::query(
149 "SELECT conversation_id, assistant_turn IS NOT NULL AS answered FROM tf_turn
150 WHERE account_id = $1 AND turn_id = $2",
151 )
152 .bind(account.as_str())
153 .bind(turn.turn_id.as_uuid())
154 .fetch_optional(conn)
155 .await
156 .map_err(|error| store_error(&error))?
157 .ok_or(StoreError::NotFound)?;
158 let stored: Uuid = column(&row, "conversation_id")?;
159 if stored != *turn.conversation_id.as_uuid() {
160 return Err(identity_mismatch());
161 }
162 Err(StoreError::Conflict)
163}
164
165pub(crate) async fn load_recent_turns(
167 conn: &mut PgConnection,
168 account: &AccountId,
169 conversation: &ConversationId,
170 limit: usize,
171) -> Result<Vec<StoredTurn>, StoreError> {
172 load_conversation(&mut *conn, account, conversation).await?;
175 let rows = sqlx::query(
176 "SELECT t.received_at, t.user_turn, t.assistant_turn, p.phase
177 FROM tf_turn t
178 JOIN tf_turn_phase p ON p.account_id = t.account_id AND p.turn_id = t.turn_id
179 WHERE t.account_id = $1 AND t.conversation_id = $2
180 ORDER BY t.received_at DESC, t.turn_id DESC
181 LIMIT $3",
182 )
183 .bind(account.as_str())
184 .bind(conversation.as_uuid())
185 .bind(limit_to_sql(limit))
186 .fetch_all(conn)
187 .await
188 .map_err(|error| store_error(&error))?;
189 let mut turns = rows
190 .iter()
191 .map(decode_turn)
192 .collect::<Result<Vec<StoredTurn>, StoreError>>()?;
193 turns.reverse();
194 Ok(turns)
195}
196
197pub(crate) async fn load_turn(
199 conn: &mut PgConnection,
200 account: &AccountId,
201 turn_id: &TurnId,
202) -> Result<StoredTurn, StoreError> {
203 let row = sqlx::query(
204 "SELECT t.received_at, t.user_turn, t.assistant_turn, p.phase
205 FROM tf_turn t
206 JOIN tf_turn_phase p ON p.account_id = t.account_id AND p.turn_id = t.turn_id
207 WHERE t.account_id = $1 AND t.turn_id = $2",
208 )
209 .bind(account.as_str())
210 .bind(turn_id.as_uuid())
211 .fetch_optional(conn)
212 .await
213 .map_err(|error| store_error(&error))?
214 .ok_or(StoreError::NotFound)?;
215 decode_turn(&row)
216}
217
218pub(crate) async fn turn_phase(
220 conn: &mut PgConnection,
221 account: &AccountId,
222 turn_id: &TurnId,
223) -> Result<TurnPhaseMarker, StoreError> {
224 let row = sqlx::query(
225 "SELECT conversation_id, phase, updated_at FROM tf_turn_phase
226 WHERE account_id = $1 AND turn_id = $2",
227 )
228 .bind(account.as_str())
229 .bind(turn_id.as_uuid())
230 .fetch_optional(conn)
231 .await
232 .map_err(|error| store_error(&error))?
233 .ok_or(StoreError::NotFound)?;
234 decode_marker(&row, account, *turn_id)
235}
236
237pub(crate) async fn set_turn_phase(
244 conn: &mut PgConnection,
245 account: &AccountId,
246 turn_id: &TurnId,
247 phase: TurnPhase,
248 at: DateTime<Utc>,
249) -> Result<TurnPhaseMarker, StoreError> {
250 let statement = format!(
251 "UPDATE tf_turn_phase SET phase = $3, updated_at = $4
252 WHERE account_id = $1 AND turn_id = $2
253 AND (phase = $3 OR phase NOT IN {TERMINAL_PHASES})
254 RETURNING conversation_id, phase, updated_at"
255 );
256 let updated = sqlx::query(&statement)
257 .bind(account.as_str())
258 .bind(turn_id.as_uuid())
259 .bind(label(&phase)?)
260 .bind(at)
261 .fetch_optional(&mut *conn)
262 .await
263 .map_err(|error| store_error(&error))?;
264 match updated {
265 Some(row) => decode_marker(&row, account, *turn_id),
266 None => Err(refusal(conn, account, turn_id).await?),
267 }
268}
269
270async fn refusal(
273 conn: &mut PgConnection,
274 account: &AccountId,
275 turn_id: &TurnId,
276) -> Result<StoreError, StoreError> {
277 let exists = sqlx::query("SELECT 1 FROM tf_turn_phase WHERE account_id = $1 AND turn_id = $2")
278 .bind(account.as_str())
279 .bind(turn_id.as_uuid())
280 .fetch_optional(conn)
281 .await
282 .map_err(|error| store_error(&error))?
283 .is_some();
284 Ok(if exists {
285 StoreError::Conflict
286 } else {
287 StoreError::NotFound
288 })
289}
290
291pub(crate) async fn list_unfinished_turns(
293 conn: &mut PgConnection,
294 scope: RecoveryScope,
295 limit: usize,
296) -> Result<Vec<TurnPhaseMarker>, StoreError> {
297 let scoped = match &scope {
298 RecoveryScope::Account(account) => Some(account.as_str()),
299 RecoveryScope::AllAccounts => None,
300 };
301 let statement = format!(
302 "SELECT p.account_id, p.conversation_id, p.turn_id, p.phase, p.updated_at
303 FROM tf_turn_phase p
304 JOIN tf_turn t ON t.account_id = p.account_id AND t.turn_id = p.turn_id
305 WHERE p.phase NOT IN {TERMINAL_PHASES}
306 AND ($1::text IS NULL OR p.account_id = $1)
307 ORDER BY t.received_at, p.turn_id
308 LIMIT $2"
309 );
310 let rows = sqlx::query(&statement)
311 .bind(scoped)
312 .bind(limit_to_sql(limit))
313 .fetch_all(conn)
314 .await
315 .map_err(|error| store_error(&error))?;
316 rows.iter()
317 .map(|row| {
318 let owner: String = column(row, "account_id")?;
319 let turn_id: Uuid = column(row, "turn_id")?;
320 decode_marker(row, &AccountId::from(owner), TurnId::from(turn_id))
321 })
322 .collect()
323}
324
325fn decode_turn(row: &sqlx::postgres::PgRow) -> Result<StoredTurn, StoreError> {
327 let assistant: Option<serde_json::Value> = column(row, "assistant_turn")?;
328 let phase: String = column(row, "phase")?;
329 Ok(StoredTurn {
330 user: StoredUserTurn::new(
331 from_json(column(row, "user_turn")?)?,
332 column(row, "received_at")?,
333 ),
334 assistant: assistant.map(from_json).transpose()?,
335 phase: from_label(&phase)?,
336 })
337}
338
339fn decode_marker(
341 row: &sqlx::postgres::PgRow,
342 account: &AccountId,
343 turn_id: TurnId,
344) -> Result<TurnPhaseMarker, StoreError> {
345 let phase: String = column(row, "phase")?;
346 let conversation: Uuid = column(row, "conversation_id")?;
347 Ok(TurnPhaseMarker {
348 account_id: account.clone(),
349 conversation_id: ConversationId::from(conversation),
350 turn_id,
351 phase: from_label(&phase)?,
352 updated_at: column(row, "updated_at")?,
353 })
354}
355
356#[async_trait]
357impl ConversationReader for PgStores {
358 async fn load_conversation(
359 &self,
360 account: &AccountId,
361 id: &ConversationId,
362 ) -> Result<ConversationRecord, StoreError> {
363 let mut conn = self.connection().await?;
364 load_conversation(&mut conn, account, id).await
365 }
366
367 async fn load_recent_turns(
368 &self,
369 account: &AccountId,
370 conversation: &ConversationId,
371 limit: usize,
372 ) -> Result<Vec<StoredTurn>, StoreError> {
373 let mut conn = self.connection().await?;
374 load_recent_turns(&mut conn, account, conversation, limit).await
375 }
376
377 async fn load_turn(
378 &self,
379 account: &AccountId,
380 turn_id: &TurnId,
381 ) -> Result<StoredTurn, StoreError> {
382 let mut conn = self.connection().await?;
383 load_turn(&mut conn, account, turn_id).await
384 }
385
386 async fn turn_phase(
387 &self,
388 account: &AccountId,
389 turn_id: &TurnId,
390 ) -> Result<TurnPhaseMarker, StoreError> {
391 let mut conn = self.connection().await?;
392 turn_phase(&mut conn, account, turn_id).await
393 }
394
395 async fn list_unfinished_turns(
396 &self,
397 scope: RecoveryScope,
398 limit: usize,
399 ) -> Result<Vec<TurnPhaseMarker>, StoreError> {
400 let mut conn = self.connection().await?;
401 list_unfinished_turns(&mut conn, scope, limit).await
402 }
403}
404
405#[async_trait]
406impl ConversationWriter for PgStores {
407 async fn create_conversation(&self, record: ConversationRecord) -> Result<(), StoreError> {
408 let mut transaction = self.transaction().await?;
409 create_conversation(&mut transaction, record).await?;
410 commit(transaction).await
411 }
412
413 async fn append_user_turn(&self, turn: StoredUserTurn) -> Result<(), StoreError> {
414 let mut transaction = self.transaction().await?;
415 append_user_turn(&mut transaction, turn, now()).await?;
416 commit(transaction).await
417 }
418
419 async fn append_assistant_turn(
420 &self,
421 account: &AccountId,
422 turn: AssistantTurn,
423 ) -> Result<(), StoreError> {
424 let mut transaction = self.transaction().await?;
425 append_assistant_turn(&mut transaction, account, turn).await?;
426 commit(transaction).await
427 }
428
429 async fn set_turn_phase(
430 &self,
431 account: &AccountId,
432 turn_id: &TurnId,
433 phase: TurnPhase,
434 ) -> Result<TurnPhaseMarker, StoreError> {
435 let mut transaction = self.transaction().await?;
436 let marker = set_turn_phase(&mut transaction, account, turn_id, phase, now()).await?;
437 commit(transaction).await?;
438 Ok(marker)
439 }
440}