Skip to main content

turnframe_store_postgres/
conversations.rs

1//! Conversations, turns and the crash-recovery phase marker over `tf_conversation`,
2//! `tf_turn` and `tf_turn_phase`.
3//!
4//! The turn row is written once and never rewritten except to attach the
5//! assistant turn, and both sides are stored as whole `jsonb` documents: a
6//! reload deserializes the exact [`AssistantTurn`] that was returned to the
7//! client, with its blocks in their original order, instead of rebuilding cards
8//! from prose (spec §22.3).
9//!
10//! The phase marker lives in its own table because it moves on every step of a
11//! turn while the turn itself does not, and because recovery reads only the
12//! marker. Its partial index carries just the unfinished turns, so the sweep of
13//! spec §23.1 costs nothing on a system with no interrupted turns.
14
15use 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
32/// The phases a turn never leaves, as the SQL literals the partial index uses.
33const TERMINAL_PHASES: &str = "('delivered', 'failed')";
34
35/// Inserts a conversation.
36pub(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
54/// Loads a conversation of `account`.
55pub(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
78/// Appends a user turn and opens its phase marker at `Received`.
79///
80/// The insert is conditional on the conversation existing, so a turn addressed
81/// to a conversation of another tenant is refused by the same `NotFound` an
82/// unknown conversation gets, without a separate lookup that could race.
83pub(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
125/// Attaches the assistant turn to the user turn it answers.
126pub(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    // Nothing moved: say which of the three reasons it was, in the order the
147    // contract specifies — unknown turn, wrong conversation, already answered.
148    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
165/// The most recent `limit` turns of a conversation, oldest of them first.
166pub(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    // The conversation is looked up first so another tenant's conversation is
173    // `NotFound` rather than an empty window.
174    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
197/// Loads one turn.
198pub(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
218/// Reads the phase marker of a turn.
219pub(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
237/// Moves the phase marker, refusing to leave a terminal phase.
238///
239/// The refusal is the `WHERE` clause, so the check and the write are the same
240/// statement and a concurrent delivery cannot slip between them. Writing the
241/// phase a turn already has is accepted, which is what makes recovery safe to
242/// run more than once.
243pub(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
270/// Whether nothing moved because the turn is unknown or because its phase is
271/// final.
272async 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
291/// The markers of turns that have not finished, oldest first.
292pub(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
325/// Rebuilds a stored turn from its row.
326fn 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
339/// Rebuilds a phase marker from its row.
340fn 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}