1use 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}