Skip to main content

runifold_store_postgres/conversation/
store.rs

1//! `ConversationStore` persistence implementation.
2
3use std::time::Instant;
4
5use runifold_agent::{
6    ConversationAppend, ConversationCreateOutcome, ConversationId, ConversationSequence,
7    ConversationStore, ConversationStoreError, ConversationStoreFuture, ConversationSummary,
8    ConversationSummaryBatch, ConversationSummaryCommit, ConversationTranscriptEntry,
9    ConversationVersion, ConversationView, ConversationWindow, MemoryNamespace, SemanticMemory,
10    SemanticMemoryQuery, SemanticMemorySearchOutcome, SemanticMemoryUpsert,
11    SemanticMemoryUpsertOutcome,
12};
13use runifold_core::CheckpointId;
14use runifold_retrieval::{EmbeddingRequest, EmbeddingTask, RetrievalContext};
15
16use super::{
17    PostgresConversationStore,
18    support::{
19        combine_usage, conversation_uuid, database_usage, decode_memory, decode_required_summary,
20        decode_summary, decode_transcript_entry, decode_version, encode_error, invalid_input,
21        memory_uuid, namespace_mismatch, not_found, retrieval_error, storage_error, to_i64,
22        to_pgvector, validate_append, validate_memory, validate_summary,
23    },
24};
25
26impl ConversationStore for PostgresConversationStore {
27    fn create(
28        &self,
29        conversation_id: ConversationId,
30        namespace: MemoryNamespace,
31    ) -> ConversationStoreFuture<'_, Result<ConversationCreateOutcome, ConversationStoreError>>
32    {
33        Box::pin(async move {
34            let sql = format!(
35                r"
36                WITH inserted AS (
37                    INSERT INTO {} (conversation_id, namespace)
38                    VALUES ($1, $2)
39                    ON CONFLICT (conversation_id) DO NOTHING
40                    RETURNING namespace
41                )
42                SELECT namespace, TRUE AS created FROM inserted
43                UNION ALL
44                SELECT namespace, FALSE AS created FROM {}
45                WHERE conversation_id = $1 AND NOT EXISTS (SELECT 1 FROM inserted)
46                ",
47                self.table, self.table
48            );
49            let row = self
50                .client
51                .query_one(
52                    &sql,
53                    &[&conversation_uuid(conversation_id), &namespace.as_str()],
54                )
55                .await
56                .map_err(storage_error)?;
57            let actual: String = row.get("namespace");
58            if actual != namespace.as_str() {
59                return Err(namespace_mismatch());
60            }
61            Ok(if row.get("created") {
62                ConversationCreateOutcome::Created
63            } else {
64                ConversationCreateOutcome::Duplicate
65            })
66        })
67    }
68
69    fn load_view(
70        &self,
71        conversation_id: ConversationId,
72        namespace: MemoryNamespace,
73        window: ConversationWindow,
74        summary_batch: ConversationSummaryBatch,
75    ) -> ConversationStoreFuture<'_, Result<ConversationView, ConversationStoreError>> {
76        Box::pin(async move {
77            let metadata_sql = format!(
78                r"
79                SELECT namespace, version, summary_id, summary_content, summary_through,
80                    summary_transcript_version,
81                    (EXTRACT(EPOCH FROM summary_created_at) * 1000)::BIGINT
82                        AS summary_created_at_ms
83                FROM {} WHERE conversation_id = $1
84                ",
85                self.table
86            );
87            let Some(row) = self
88                .client
89                .query_opt(&metadata_sql, &[&conversation_uuid(conversation_id)])
90                .await
91                .map_err(storage_error)?
92            else {
93                return Err(not_found());
94            };
95            let actual: String = row.get("namespace");
96            if actual != namespace.as_str() {
97                return Err(namespace_mismatch());
98            }
99            let summary = decode_summary(&row)?;
100            let summarized_through = summary
101                .as_ref()
102                .map_or(0, |summary| summary.through_sequence.get());
103            let transcript = self
104                .load_bounded_transcript(conversation_id, summarized_through, window, summary_batch)
105                .await?;
106            Ok(ConversationView {
107                conversation_id,
108                namespace,
109                version: decode_version(row.get("version"))?,
110                summary,
111                summary_buffer: transcript.summary_buffer,
112                summary_backlog: transcript.summary_backlog,
113                window: transcript.window,
114            })
115        })
116    }
117
118    fn list_transcript(
119        &self,
120        conversation_id: ConversationId,
121        namespace: MemoryNamespace,
122        after: Option<ConversationSequence>,
123        limit: ConversationWindow,
124    ) -> ConversationStoreFuture<'_, Result<Vec<ConversationTranscriptEntry>, ConversationStoreError>>
125    {
126        Box::pin(async move {
127            let Some(actual) = self.conversation_namespace(conversation_id).await? else {
128                return Err(not_found());
129            };
130            if actual != namespace.as_str() {
131                return Err(namespace_mismatch());
132            }
133            let sql = format!(
134                r"
135                SELECT sequence, message FROM {}_transcript
136                WHERE conversation_id = $1 AND sequence > $2
137                ORDER BY sequence ASC LIMIT $3
138                ",
139                self.table
140            );
141            self.client
142                .query(
143                    &sql,
144                    &[
145                        &conversation_uuid(conversation_id),
146                        &to_i64(after.map_or(0, ConversationSequence::get))?,
147                        &i64::from(limit.get()),
148                    ],
149                )
150                .await
151                .map_err(storage_error)?
152                .iter()
153                .map(decode_transcript_entry)
154                .collect()
155        })
156    }
157
158    fn append(
159        &self,
160        namespace: MemoryNamespace,
161        command: ConversationAppend,
162    ) -> ConversationStoreFuture<'_, Result<ConversationVersion, ConversationStoreError>> {
163        Box::pin(async move {
164            validate_append(&command)?;
165            let messages = serde_json::to_value(&command.messages).map_err(encode_error)?;
166            let sql = format!(
167                r"
168                WITH updated AS (
169                    UPDATE {table}
170                    SET version = version + 1, updated_at = clock_timestamp()
171                    WHERE conversation_id = $1 AND namespace = $2
172                        AND version = $3 AND version < 9223372036854775807
173                    RETURNING conversation_id, version
174                ),
175                base AS (
176                    SELECT COALESCE(MAX(sequence), 0) AS last_sequence
177                    FROM {table}_transcript WHERE conversation_id = $1
178                ),
179                inserted AS (
180                    INSERT INTO {table}_transcript (conversation_id, sequence, message)
181                    SELECT updated.conversation_id,
182                        base.last_sequence + payload.ordinality,
183                        payload.message
184                    FROM updated CROSS JOIN base
185                    CROSS JOIN LATERAL
186                        jsonb_array_elements($4::JSONB)
187                        WITH ORDINALITY AS payload(message, ordinality)
188                    RETURNING sequence
189                )
190                SELECT updated.version FROM updated
191                WHERE (SELECT COUNT(*) FROM inserted) = jsonb_array_length($4::JSONB)
192                ",
193                table = self.table
194            );
195            let row = self
196                .client
197                .query_opt(
198                    &sql,
199                    &[
200                        &conversation_uuid(command.conversation_id),
201                        &namespace.as_str(),
202                        &to_i64(command.expected_version.get())?,
203                        &messages,
204                    ],
205                )
206                .await
207                .map_err(storage_error)?;
208            match row {
209                Some(row) => decode_version(row.get("version")),
210                None => Err(self
211                    .diagnose_conversation(
212                        command.conversation_id,
213                        &namespace,
214                        "conversation transcript version precondition failed",
215                    )
216                    .await),
217            }
218        })
219    }
220
221    fn commit_summary(
222        &self,
223        namespace: MemoryNamespace,
224        command: ConversationSummaryCommit,
225    ) -> ConversationStoreFuture<'_, Result<ConversationSummary, ConversationStoreError>> {
226        Box::pin(async move {
227            validate_summary(&command)?;
228            let summary_id = CheckpointId::new();
229            let sql = format!(
230                r"
231                UPDATE {table}
232                SET summary_id = $4, summary_content = $5, summary_through = $6,
233                    summary_transcript_version = version,
234                    summary_created_at = clock_timestamp(),
235                    updated_at = clock_timestamp()
236                WHERE conversation_id = $1 AND namespace = $2 AND version = $3
237                    AND $6 > COALESCE(summary_through, 0)
238                    AND $6 <= (
239                        SELECT COALESCE(MAX(sequence), 0)
240                        FROM {table}_transcript WHERE conversation_id = $1
241                    )
242                RETURNING summary_id, summary_content, summary_through,
243                    summary_transcript_version,
244                    (EXTRACT(EPOCH FROM summary_created_at) * 1000)::BIGINT
245                        AS summary_created_at_ms
246                ",
247                table = self.table
248            );
249            let row = self
250                .client
251                .query_opt(
252                    &sql,
253                    &[
254                        &conversation_uuid(command.conversation_id),
255                        &namespace.as_str(),
256                        &to_i64(command.expected_version.get())?,
257                        &summary_id.as_uuid(),
258                        &command.content,
259                        &to_i64(command.through_sequence.get())?,
260                    ],
261                )
262                .await
263                .map_err(storage_error)?;
264            match row {
265                Some(row) => decode_required_summary(&row),
266                None => Err(self
267                    .diagnose_conversation(
268                        command.conversation_id,
269                        &namespace,
270                        "conversation summary precondition failed",
271                    )
272                    .await),
273            }
274        })
275    }
276
277    fn upsert_memory(
278        &self,
279        command: SemanticMemoryUpsert,
280    ) -> ConversationStoreFuture<'_, Result<SemanticMemory, ConversationStoreError>> {
281        if self.semantic_embedder.is_some() {
282            return Box::pin(async move {
283                self.upsert_memory_scoped(command, RetrievalContext::new())
284                    .await
285                    .map(|outcome| outcome.memory)
286            });
287        }
288        Box::pin(async move {
289            validate_memory(&command)?;
290            self.validate_memory_sources(&command).await?;
291            let sources = serde_json::to_value(&command.sources).map_err(encode_error)?;
292            let metadata = serde_json::to_value(&command.metadata).map_err(encode_error)?;
293            let sql = format!(
294                r"
295                INSERT INTO {table}_memory (
296                    memory_id, namespace, content, sources, metadata, revision
297                )
298                VALUES ($1, $2, $3, $4, $5, 0)
299                ON CONFLICT (memory_id) DO UPDATE
300                SET content = EXCLUDED.content, sources = EXCLUDED.sources,
301                    metadata = EXCLUDED.metadata,
302                    revision = {table}_memory.revision + 1,
303                    updated_at = clock_timestamp()
304                WHERE {table}_memory.namespace = EXCLUDED.namespace
305                    AND $6::BIGINT IS NOT NULL
306                    AND {table}_memory.revision = $6
307                    AND {table}_memory.revision < 9223372036854775807
308                RETURNING memory_id, namespace, content, sources, metadata, revision,
309                    (EXTRACT(EPOCH FROM created_at) * 1000)::BIGINT AS created_at_ms,
310                    (EXTRACT(EPOCH FROM updated_at) * 1000)::BIGINT AS updated_at_ms
311                ",
312                table = self.table
313            );
314            let expected = command.expected_revision.map(to_i64).transpose()?;
315            let row = self
316                .client
317                .query_opt(
318                    &sql,
319                    &[
320                        &memory_uuid(command.memory_id),
321                        &command.namespace.as_str(),
322                        &command.content,
323                        &sources,
324                        &metadata,
325                        &expected,
326                    ],
327                )
328                .await
329                .map_err(storage_error)?;
330            match row {
331                Some(row) => decode_memory(&row),
332                None => self.diagnose_memory(&command).await,
333            }
334        })
335    }
336
337    fn search_memory(
338        &self,
339        query: SemanticMemoryQuery,
340    ) -> ConversationStoreFuture<'_, Result<Vec<SemanticMemory>, ConversationStoreError>> {
341        if self.semantic_embedder.is_some() {
342            return Box::pin(async move {
343                self.search_memory_scoped(query, RetrievalContext::new())
344                    .await
345                    .map(|outcome| outcome.memories)
346            });
347        }
348        Box::pin(async move {
349            let sql = format!(
350                r"
351                SELECT memory_id, namespace, content, sources, metadata, revision,
352                    (EXTRACT(EPOCH FROM created_at) * 1000)::BIGINT AS created_at_ms,
353                    (EXTRACT(EPOCH FROM updated_at) * 1000)::BIGINT AS updated_at_ms
354                FROM {table}_memory
355                WHERE namespace = $1
356                    AND to_tsvector('simple', content) @@ plainto_tsquery('simple', $2)
357                ORDER BY ts_rank_cd(
358                    to_tsvector('simple', content),
359                    plainto_tsquery('simple', $2)
360                ) DESC, updated_at DESC, memory_id ASC
361                LIMIT $3
362                ",
363                table = self.table
364            );
365            self.client
366                .query(
367                    &sql,
368                    &[
369                        &query.namespace.as_str(),
370                        &query.text,
371                        &i64::from(query.limit.get()),
372                    ],
373                )
374                .await
375                .map_err(storage_error)?
376                .iter()
377                .map(decode_memory)
378                .collect()
379        })
380    }
381
382    fn upsert_memory_scoped(
383        &self,
384        command: SemanticMemoryUpsert,
385        context: RetrievalContext,
386    ) -> ConversationStoreFuture<'_, Result<SemanticMemoryUpsertOutcome, ConversationStoreError>>
387    {
388        Box::pin(async move {
389            context
390                .check_live()
391                .map_err(|error| retrieval_error(&error))?;
392            let Some(embedder) = &self.semantic_embedder else {
393                let started = Instant::now();
394                let memory = self.upsert_memory(command).await?;
395                return Ok(SemanticMemoryUpsertOutcome {
396                    memory,
397                    usage: database_usage(started),
398                });
399            };
400            validate_memory(&command)?;
401            self.validate_memory_sources(&command).await?;
402            let batch = embedder
403                .embed(
404                    EmbeddingRequest::new(
405                        vec![command.content.clone()],
406                        EmbeddingTask::RetrievalDocument,
407                    )
408                    .map_err(|error| retrieval_error(&error))?,
409                    context.child_attempt(),
410                )
411                .await
412                .map_err(|error| retrieval_error(&error))?
413                .validate_count(1)
414                .map_err(|error| retrieval_error(&error))?;
415            let embedding_usage = batch.usage;
416            let embedding = batch
417                .embeddings
418                .into_iter()
419                .next()
420                .ok_or_else(|| invalid_input("semantic memory embedding response was empty"))?;
421            let started = Instant::now();
422            let memory = self
423                .upsert_vector_memory(command, to_pgvector(&embedding)?)
424                .await?;
425            Ok(SemanticMemoryUpsertOutcome {
426                memory,
427                usage: combine_usage(embedding_usage, database_usage(started))?,
428            })
429        })
430    }
431
432    fn search_memory_scoped(
433        &self,
434        query: SemanticMemoryQuery,
435        context: RetrievalContext,
436    ) -> ConversationStoreFuture<'_, Result<SemanticMemorySearchOutcome, ConversationStoreError>>
437    {
438        Box::pin(async move {
439            context
440                .check_live()
441                .map_err(|error| retrieval_error(&error))?;
442            let Some(embedder) = &self.semantic_embedder else {
443                let started = Instant::now();
444                let memories = self.search_memory(query).await?;
445                return Ok(SemanticMemorySearchOutcome {
446                    memories,
447                    usage: database_usage(started),
448                });
449            };
450            let batch = embedder
451                .embed(
452                    EmbeddingRequest::new(vec![query.text.clone()], EmbeddingTask::RetrievalQuery)
453                        .map_err(|error| retrieval_error(&error))?,
454                    context.child_attempt(),
455                )
456                .await
457                .map_err(|error| retrieval_error(&error))?
458                .validate_count(1)
459                .map_err(|error| retrieval_error(&error))?;
460            let embedding_usage = batch.usage;
461            let embedding = batch
462                .embeddings
463                .into_iter()
464                .next()
465                .ok_or_else(|| invalid_input("semantic memory query embedding was empty"))?;
466            let started = Instant::now();
467            let memories = self
468                .search_vector_memory(&query, to_pgvector(&embedding)?)
469                .await?;
470            Ok(SemanticMemorySearchOutcome {
471                memories,
472                usage: combine_usage(embedding_usage, database_usage(started))?,
473            })
474        })
475    }
476}