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