1use std::time::{Duration, SystemTime, UNIX_EPOCH};
37
38use agent_effects_store::{
39 EffectEvent, EffectId, EffectKey, EffectKind, EffectName, EffectRecord, EffectStatus,
40 EffectStore, ErrorRecord, InsertOutcome, Lease, ListQuery, LogicalKey, NewEffect, PruneQuery,
41 StoreError, Transition, TransitionRequest, WorkerId,
42};
43use serde_json::Value;
44use sqlx::postgres::{PgConnectOptions, PgPool, PgPoolOptions, PgRow};
45use sqlx::{AssertSqlSafe, Postgres, Row, Transaction};
46use time::OffsetDateTime;
47
48static MIGRATOR: sqlx::migrate::Migrator = sqlx::migrate!("./migrations");
49
50const COLUMNS: &str = "id, effect_name, logical_key, kind, status, input, input_fingerprint, \
51 output, last_error, created_by, attempt_count, may_have_applied, compensation_attempts, \
52 approved, next_attempt_at, attempt_started_at, attempt_ended_at, lease_owner, lease_epoch, \
53 lease_expires_at, version, created_at, updated_at, committed_at";
54
55#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
57pub enum ClockSource {
58 #[default]
60 Database,
61 Caller,
63}
64
65#[derive(Clone, Debug)]
69pub struct PostgresStore {
70 pool: PgPool,
71 clock: ClockSource,
72}
73
74impl PostgresStore {
75 pub async fn connect(url: &str) -> Result<Self, StoreError> {
81 let options: PgConnectOptions = url.parse().map_err(StoreError::backend)?;
82 Self::connect_with(options).await
83 }
84
85 pub async fn connect_with(options: PgConnectOptions) -> Result<Self, StoreError> {
92 let pool = PgPoolOptions::new()
93 .max_connections(10)
94 .connect_with(options)
95 .await
96 .map_err(StoreError::backend)?;
97 Self::from_pool(pool).await
98 }
99
100 pub async fn from_pool(pool: PgPool) -> Result<Self, StoreError> {
107 MIGRATOR.run(&pool).await.map_err(StoreError::backend)?;
108 Ok(Self {
109 pool,
110 clock: ClockSource::Database,
111 })
112 }
113
114 #[must_use]
116 pub fn with_clock_source(mut self, clock: ClockSource) -> Self {
117 self.clock = clock;
118 self
119 }
120
121 pub fn pool(&self) -> &PgPool {
123 &self.pool
124 }
125
126 async fn database_now(&self) -> Result<SystemTime, StoreError> {
127 let now: OffsetDateTime = sqlx::query_scalar("SELECT clock_timestamp()")
128 .fetch_one(&self.pool)
129 .await
130 .map_err(StoreError::backend)?;
131 Ok(SystemTime::from(now))
132 }
133
134 fn effective(&self, database: SystemTime, caller: SystemTime) -> SystemTime {
136 match self.clock {
137 ClockSource::Database => database,
138 ClockSource::Caller => caller,
139 }
140 }
141
142 async fn modify<T>(
146 &self,
147 id: EffectId,
148 change: impl FnOnce(&mut EffectRecord, SystemTime) -> Result<T, StoreError>,
149 ) -> Result<(T, EffectRecord, Transaction<'static, Postgres>), StoreError> {
150 let mut tx = self.pool.begin().await.map_err(StoreError::backend)?;
151 let row = sqlx::query(AssertSqlSafe(format!(
152 "SELECT {COLUMNS} FROM effects WHERE id = $1 FOR UPDATE"
153 )))
154 .bind(*id.as_uuid())
155 .fetch_optional(&mut *tx)
156 .await
157 .map_err(StoreError::backend)?
158 .ok_or(StoreError::NotFound(id))?;
159 let mut record = decode_record(&row)?;
160 let db_now: OffsetDateTime = sqlx::query_scalar("SELECT clock_timestamp()")
163 .fetch_one(&mut *tx)
164 .await
165 .map_err(StoreError::backend)?;
166 let read_version = record.version;
167 let result = change(&mut record, SystemTime::from(db_now))?;
168 normalize(&mut record);
169 save(&mut tx, &record, read_version).await?;
170 Ok((result, record, tx))
171 }
172}
173
174impl EffectStore for PostgresStore {
175 async fn insert_or_get(&self, mut new: NewEffect) -> Result<InsertOutcome, StoreError> {
176 if self.clock == ClockSource::Database {
177 new.now = self.database_now().await?;
178 }
179 let mut record = EffectRecord::new(new);
180 normalize(&mut record);
181 let inserted = sqlx::query(AssertSqlSafe(format!(
182 "INSERT INTO effects ({COLUMNS}) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, \
183 $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, $24) \
184 ON CONFLICT (effect_name, logical_key) DO NOTHING"
185 )))
186 .bind(*record.id.as_uuid())
187 .bind(record.key.name.as_str())
188 .bind(record.key.key.as_str())
189 .bind(record.kind.as_str())
190 .bind(record.status.as_str())
191 .bind(record.input.clone())
192 .bind(record.input_fingerprint.clone())
193 .bind(record.output.clone())
194 .bind(json(record.last_error.as_ref())?)
195 .bind(record.created_by.clone())
196 .bind(to_i32(record.attempt_count)?)
197 .bind(record.may_have_applied)
198 .bind(to_i32(record.compensation_attempts)?)
199 .bind(record.approved)
200 .bind(record.next_attempt_at.map(timestamp))
201 .bind(record.attempt_started_at.map(timestamp))
202 .bind(record.attempt_ended_at.map(timestamp))
203 .bind(record.lease_owner.as_ref().map(|w| w.as_str().to_owned()))
204 .bind(to_i64(record.lease_epoch)?)
205 .bind(record.lease_expires_at.map(timestamp))
206 .bind(to_i64(record.version)?)
207 .bind(timestamp(record.created_at))
208 .bind(timestamp(record.updated_at))
209 .bind(record.committed_at.map(timestamp))
210 .execute(&self.pool)
211 .await
212 .map_err(StoreError::backend)?
213 .rows_affected()
214 == 1;
215 if inserted {
216 return Ok(InsertOutcome {
217 record,
218 inserted: true,
219 });
220 }
221 let existing = self
222 .get_by_key(&record.key)
223 .await?
224 .ok_or_else(|| StoreError::backend("conflicting insert, but no existing record"))?;
225 Ok(InsertOutcome {
226 record: existing,
227 inserted: false,
228 })
229 }
230
231 async fn get(&self, id: EffectId) -> Result<Option<EffectRecord>, StoreError> {
232 sqlx::query(AssertSqlSafe(format!(
233 "SELECT {COLUMNS} FROM effects WHERE id = $1"
234 )))
235 .bind(*id.as_uuid())
236 .fetch_optional(&self.pool)
237 .await
238 .map_err(StoreError::backend)?
239 .as_ref()
240 .map(decode_record)
241 .transpose()
242 }
243
244 async fn get_by_key(&self, key: &EffectKey) -> Result<Option<EffectRecord>, StoreError> {
245 sqlx::query(AssertSqlSafe(format!(
246 "SELECT {COLUMNS} FROM effects WHERE effect_name = $1 AND logical_key = $2"
247 )))
248 .bind(key.name.as_str())
249 .bind(key.key.as_str())
250 .fetch_optional(&self.pool)
251 .await
252 .map_err(StoreError::backend)?
253 .as_ref()
254 .map(decode_record)
255 .transpose()
256 }
257
258 async fn acquire_lease(
259 &self,
260 id: EffectId,
261 owner: &WorkerId,
262 now: SystemTime,
263 ttl: Duration,
264 ) -> Result<Lease, StoreError> {
265 let (mut lease, _, tx) = self
266 .modify(id, |record, db_now| {
267 record.acquire_lease(owner, self.effective(db_now, now), ttl)
268 })
269 .await?;
270 tx.commit().await.map_err(StoreError::backend)?;
271 lease.expires_at = round_trip(lease.expires_at);
272 Ok(lease)
273 }
274
275 async fn renew_lease(
276 &self,
277 lease: &Lease,
278 now: SystemTime,
279 ttl: Duration,
280 ) -> Result<Lease, StoreError> {
281 let (mut renewed, _, tx) = self
282 .modify(lease.effect_id, |record, db_now| {
283 record.renew_lease(lease, self.effective(db_now, now), ttl)
284 })
285 .await?;
286 tx.commit().await.map_err(StoreError::backend)?;
287 renewed.expires_at = round_trip(renewed.expires_at);
288 Ok(renewed)
289 }
290
291 async fn release_lease(&self, lease: &Lease) -> Result<(), StoreError> {
292 let (released, _, tx) = self
293 .modify(lease.effect_id, |record, _| Ok(record.release_lease(lease)))
294 .await?;
295 if released {
296 tx.commit().await.map_err(StoreError::backend)?;
297 } else {
298 tx.rollback().await.map_err(StoreError::backend)?;
299 }
300 Ok(())
301 }
302
303 async fn transition(&self, mut request: TransitionRequest) -> Result<EffectRecord, StoreError> {
304 let id = request.id;
305 let (mut event, record, mut tx) = self
306 .modify(id, |record, db_now| {
307 request.now = self.effective(db_now, request.now);
308 record.apply(request)
309 })
310 .await?;
311 event.at = round_trip(event.at);
312 sqlx::query(
313 "INSERT INTO effect_events \
314 (effect_id, sequence, transition, from_status, to_status, attempt, actor, payload, at) \
315 VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)",
316 )
317 .bind(*event.effect_id.as_uuid())
318 .bind(to_i64(event.sequence)?)
319 .bind(event.transition.as_str())
320 .bind(event.from.as_str())
321 .bind(event.to.as_str())
322 .bind(to_i32(event.attempt)?)
323 .bind(event.actor.clone())
324 .bind(event.payload.clone())
325 .bind(timestamp(event.at))
326 .execute(&mut *tx)
327 .await
328 .map_err(StoreError::backend)?;
329 tx.commit().await.map_err(StoreError::backend)?;
330 Ok(record)
331 }
332
333 async fn list(&self, query: ListQuery) -> Result<Vec<EffectRecord>, StoreError> {
334 let mut sql = format!("SELECT {COLUMNS} FROM effects WHERE TRUE");
335 let mut next = 1;
336 let mut placeholder = || {
337 let p = format!("${next}");
338 next += 1;
339 p
340 };
341 if !query.statuses.is_empty() {
342 let marks: Vec<String> = query.statuses.iter().map(|_| placeholder()).collect();
343 sql.push_str(" AND status IN (");
344 sql.push_str(&marks.join(", "));
345 sql.push(')');
346 }
347 let lease_filter = query.lease_expired_at.is_some();
348 if lease_filter {
349 let now = match self.clock {
350 ClockSource::Database => "clock_timestamp()".to_owned(),
351 ClockSource::Caller => placeholder(),
352 };
353 sql.push_str(
354 " AND (lease_owner IS NULL OR lease_expires_at IS NULL OR lease_expires_at <= ",
355 );
356 sql.push_str(&now);
357 sql.push(')');
358 }
359 if query.after.is_some() {
360 sql.push_str(" AND id > ");
361 sql.push_str(&placeholder());
362 }
363 sql.push_str(" ORDER BY id LIMIT ");
364 sql.push_str(&placeholder());
365 if lease_filter {
366 sql.push_str(" FOR UPDATE SKIP LOCKED");
368 }
369
370 let mut statement = sqlx::query(AssertSqlSafe(sql));
372 for status in &query.statuses {
373 statement = statement.bind(status.as_str());
374 }
375 if let (Some(now), ClockSource::Caller) = (query.lease_expired_at, self.clock) {
376 statement = statement.bind(timestamp(now));
377 }
378 if let Some(after) = query.after {
379 statement = statement.bind(*after.as_uuid());
380 }
381 statement = statement.bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
382 statement
383 .fetch_all(&self.pool)
384 .await
385 .map_err(StoreError::backend)?
386 .iter()
387 .map(decode_record)
388 .collect()
389 }
390
391 async fn events(&self, id: EffectId) -> Result<Vec<EffectEvent>, StoreError> {
392 if self.get(id).await?.is_none() {
393 return Err(StoreError::NotFound(id));
394 }
395 sqlx::query(
396 "SELECT effect_id, sequence, transition, from_status, to_status, attempt, actor, \
397 payload, at FROM effect_events WHERE effect_id = $1 ORDER BY sequence",
398 )
399 .bind(*id.as_uuid())
400 .fetch_all(&self.pool)
401 .await
402 .map_err(StoreError::backend)?
403 .iter()
404 .map(decode_event)
405 .collect()
406 }
407
408 async fn prune(&self, query: PruneQuery) -> Result<u64, StoreError> {
409 let Some(cutoff) = query.cutoff().filter(|_| query.status.is_settled()) else {
410 return Ok(0);
411 };
412 let (cutoff_sql, now_sql, limit) = match self.clock {
415 ClockSource::Database => (
416 "clock_timestamp() - $2 * interval '1 microsecond'",
417 "clock_timestamp()",
418 "$3",
419 ),
420 ClockSource::Caller => ("$2", "$3", "$4"),
421 };
422 let sql = format!(
423 "WITH doomed AS (\
424 SELECT id FROM effects WHERE status = $1 AND updated_at <= {cutoff_sql} \
425 AND (lease_owner IS NULL OR lease_expires_at IS NULL \
426 OR lease_expires_at <= {now_sql}) \
427 ORDER BY id LIMIT {limit} FOR UPDATE SKIP LOCKED), \
428 events AS (DELETE FROM effect_events WHERE effect_id IN (SELECT id FROM doomed)) \
429 DELETE FROM effects WHERE id IN (SELECT id FROM doomed)"
430 );
431 let mut statement = sqlx::query(AssertSqlSafe(sql)).bind(query.status.as_str());
433 statement = match self.clock {
434 ClockSource::Database => statement
435 .bind(i64::try_from(query.older_than.as_micros()).map_err(StoreError::backend)?),
436 ClockSource::Caller => statement.bind(timestamp(cutoff)).bind(timestamp(query.now)),
437 };
438 let deleted = statement
439 .bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
440 .execute(&self.pool)
441 .await
442 .map_err(StoreError::backend)?;
443 Ok(deleted.rows_affected())
444 }
445}
446
447async fn save(
449 tx: &mut Transaction<'static, Postgres>,
450 record: &EffectRecord,
451 read_version: u64,
452) -> Result<(), StoreError> {
453 let updated = sqlx::query(
454 "UPDATE effects SET status = $1, output = $2, last_error = $3, attempt_count = $4, \
455 may_have_applied = $5, compensation_attempts = $6, approved = $7, \
456 next_attempt_at = $8, attempt_started_at = $9, attempt_ended_at = $10, \
457 lease_owner = $11, lease_epoch = $12, lease_expires_at = $13, version = $14, \
458 updated_at = $15, committed_at = $16 \
459 WHERE id = $17 AND version = $18",
460 )
461 .bind(record.status.as_str())
462 .bind(record.output.clone())
463 .bind(json(record.last_error.as_ref())?)
464 .bind(to_i32(record.attempt_count)?)
465 .bind(record.may_have_applied)
466 .bind(to_i32(record.compensation_attempts)?)
467 .bind(record.approved)
468 .bind(record.next_attempt_at.map(timestamp))
469 .bind(record.attempt_started_at.map(timestamp))
470 .bind(record.attempt_ended_at.map(timestamp))
471 .bind(record.lease_owner.as_ref().map(|w| w.as_str().to_owned()))
472 .bind(to_i64(record.lease_epoch)?)
473 .bind(record.lease_expires_at.map(timestamp))
474 .bind(to_i64(record.version)?)
475 .bind(timestamp(record.updated_at))
476 .bind(record.committed_at.map(timestamp))
477 .bind(*record.id.as_uuid())
478 .bind(to_i64(read_version)?)
479 .execute(&mut **tx)
480 .await
481 .map_err(StoreError::backend)?
482 .rows_affected();
483 if updated == 1 {
484 Ok(())
485 } else {
486 Err(StoreError::backend(format!(
488 "effect {} changed under a row lock",
489 record.id
490 )))
491 }
492}
493
494fn decode_record(row: &PgRow) -> Result<EffectRecord, StoreError> {
495 let kind: String = get(row, "kind")?;
496 let status: String = get(row, "status")?;
497 Ok(EffectRecord {
498 id: EffectId::from_uuid(get(row, "id")?),
499 key: EffectKey::new(
500 EffectName::new(get::<String>(row, "effect_name")?).map_err(StoreError::backend)?,
501 LogicalKey::new(get::<String>(row, "logical_key")?).map_err(StoreError::backend)?,
502 ),
503 kind: EffectKind::parse(&kind).ok_or_else(|| corrupt("kind", &kind))?,
504 status: EffectStatus::parse(&status).ok_or_else(|| corrupt("status", &status))?,
505 input: get(row, "input")?,
506 input_fingerprint: get(row, "input_fingerprint")?,
507 output: get(row, "output")?,
508 last_error: get::<Option<Value>>(row, "last_error")?
509 .map(serde_json::from_value::<ErrorRecord>)
510 .transpose()
511 .map_err(StoreError::backend)?,
512 created_by: get(row, "created_by")?,
513 attempt_count: to_u32(get(row, "attempt_count")?)?,
514 may_have_applied: get(row, "may_have_applied")?,
515 compensation_attempts: to_u32(get(row, "compensation_attempts")?)?,
516 approved: get(row, "approved")?,
517 next_attempt_at: get::<Option<OffsetDateTime>>(row, "next_attempt_at")?
518 .map(SystemTime::from),
519 attempt_started_at: get::<Option<OffsetDateTime>>(row, "attempt_started_at")?
520 .map(SystemTime::from),
521 attempt_ended_at: get::<Option<OffsetDateTime>>(row, "attempt_ended_at")?
522 .map(SystemTime::from),
523 lease_owner: get::<Option<String>>(row, "lease_owner")?.map(WorkerId::new),
524 lease_epoch: to_u64(get(row, "lease_epoch")?)?,
525 lease_expires_at: get::<Option<OffsetDateTime>>(row, "lease_expires_at")?
526 .map(SystemTime::from),
527 version: to_u64(get(row, "version")?)?,
528 created_at: SystemTime::from(get::<OffsetDateTime>(row, "created_at")?),
529 updated_at: SystemTime::from(get::<OffsetDateTime>(row, "updated_at")?),
530 committed_at: get::<Option<OffsetDateTime>>(row, "committed_at")?.map(SystemTime::from),
531 })
532}
533
534fn decode_event(row: &PgRow) -> Result<EffectEvent, StoreError> {
535 let transition: String = get(row, "transition")?;
536 let from: String = get(row, "from_status")?;
537 let to: String = get(row, "to_status")?;
538 Ok(EffectEvent {
539 effect_id: EffectId::from_uuid(get(row, "effect_id")?),
540 sequence: to_u64(get(row, "sequence")?)?,
541 transition: Transition::parse(&transition)
542 .ok_or_else(|| corrupt("transition", &transition))?,
543 from: EffectStatus::parse(&from).ok_or_else(|| corrupt("from_status", &from))?,
544 to: EffectStatus::parse(&to).ok_or_else(|| corrupt("to_status", &to))?,
545 attempt: to_u32(get(row, "attempt")?)?,
546 actor: get(row, "actor")?,
547 payload: get(row, "payload")?,
548 at: SystemTime::from(get::<OffsetDateTime>(row, "at")?),
549 })
550}
551
552fn get<'r, T>(row: &'r PgRow, column: &str) -> Result<T, StoreError>
553where
554 T: sqlx::Decode<'r, Postgres> + sqlx::Type<Postgres>,
555{
556 row.try_get(column).map_err(StoreError::backend)
557}
558
559fn corrupt(column: &str, value: &str) -> StoreError {
560 StoreError::backend(format!(
561 "unreadable {column} in effects database: {value:?}"
562 ))
563}
564
565fn json(error: Option<&ErrorRecord>) -> Result<Option<Value>, StoreError> {
566 error
567 .map(serde_json::to_value)
568 .transpose()
569 .map_err(StoreError::backend)
570}
571
572fn to_i32(value: u32) -> Result<i32, StoreError> {
573 i32::try_from(value).map_err(StoreError::backend)
574}
575
576fn to_u32(value: i32) -> Result<u32, StoreError> {
577 u32::try_from(value).map_err(StoreError::backend)
578}
579
580fn to_i64(value: u64) -> Result<i64, StoreError> {
581 i64::try_from(value).map_err(StoreError::backend)
582}
583
584fn to_u64(value: i64) -> Result<u64, StoreError> {
585 u64::try_from(value).map_err(StoreError::backend)
586}
587
588fn timestamp(time: SystemTime) -> OffsetDateTime {
589 OffsetDateTime::from(round_trip(time))
590}
591
592fn round_trip(time: SystemTime) -> SystemTime {
594 match time.duration_since(UNIX_EPOCH) {
595 Ok(after) => {
596 UNIX_EPOCH + Duration::from_millis(u64::try_from(after.as_millis()).unwrap_or(u64::MAX))
597 }
598 Err(before) => {
599 let millis = u64::try_from(before.duration().as_millis()).unwrap_or(u64::MAX);
600 UNIX_EPOCH - Duration::from_millis(millis)
601 }
602 }
603}
604
605fn normalize(record: &mut EffectRecord) {
608 for time in [
609 &mut record.next_attempt_at,
610 &mut record.attempt_started_at,
611 &mut record.attempt_ended_at,
612 &mut record.lease_expires_at,
613 &mut record.committed_at,
614 ]
615 .into_iter()
616 .flatten()
617 {
618 *time = round_trip(*time);
619 }
620 record.created_at = round_trip(record.created_at);
621 record.updated_at = round_trip(record.updated_at);
622}