use async_trait::async_trait;
use awaken_contract::contract::config_store::{
ConfigChangeEvent, ConfigChangeKind, ConfigChangeNotifier, ConfigChangeSubscriber, ConfigStore,
};
use awaken_contract::contract::message::Message;
use awaken_contract::contract::storage::{
RunPage, RunQuery, RunRecord, RunStore, StorageError, ThreadRunStore, ThreadStore,
};
use awaken_contract::thread::Thread;
use sqlx::postgres::{PgListener, PgRow};
use sqlx::{PgPool, Row};
use tokio::sync::Mutex;
pub struct PostgresStore {
pool: PgPool,
threads_table: String,
runs_table: String,
messages_table: String,
configs_table: String,
config_notify_channel: String,
schema_ready: Mutex<bool>,
}
impl PostgresStore {
pub fn new(pool: PgPool) -> Self {
Self {
pool,
threads_table: "awaken_threads".to_string(),
runs_table: "awaken_runs".to_string(),
messages_table: "awaken_messages".to_string(),
configs_table: "awaken_configs".to_string(),
config_notify_channel: "awaken_config_changes".to_string(),
schema_ready: Mutex::new(false),
}
}
pub fn with_prefix(pool: PgPool, prefix: impl Into<String>) -> Self {
let prefix = prefix.into();
Self {
pool,
threads_table: format!("{prefix}_threads"),
runs_table: format!("{prefix}_runs"),
messages_table: format!("{prefix}_messages"),
configs_table: format!("{prefix}_configs"),
config_notify_channel: format!("{prefix}_config_changes"),
schema_ready: Mutex::new(false),
}
}
pub async fn ensure_schema(&self) -> Result<(), StorageError> {
let mut ready = self.schema_ready.lock().await;
if *ready {
return Ok(());
}
let statements = vec![
format!(
"CREATE TABLE IF NOT EXISTS {} (
id TEXT PRIMARY KEY,
data JSONB NOT NULL,
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
)",
self.threads_table
),
format!(
"CREATE TABLE IF NOT EXISTS {} (
thread_id TEXT NOT NULL,
data JSONB NOT NULL,
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
)",
self.messages_table
),
format!(
"CREATE TABLE IF NOT EXISTS {} (
run_id TEXT PRIMARY KEY,
thread_id TEXT NOT NULL,
agent_id TEXT NOT NULL DEFAULT '',
parent_run_id TEXT,
request JSONB,
run_input JSONB,
run_output JSONB,
status TEXT NOT NULL,
termination_reason JSONB,
final_output TEXT,
error_payload JSONB,
dispatch_id TEXT,
session_id TEXT,
transport_request_id TEXT,
waiting JSONB,
outcome JSONB,
created_at BIGINT NOT NULL,
started_at BIGINT,
finished_at BIGINT,
updated_at BIGINT NOT NULL,
steps INTEGER NOT NULL DEFAULT 0,
input_tokens BIGINT NOT NULL DEFAULT 0,
output_tokens BIGINT NOT NULL DEFAULT 0,
state JSONB
)",
self.runs_table
),
format!(
"CREATE INDEX IF NOT EXISTS idx_{}_thread_id ON {} (thread_id)",
self.runs_table, self.runs_table
),
format!(
"CREATE INDEX IF NOT EXISTS idx_{}_thread_created ON {} (thread_id, created_at DESC)",
self.runs_table, self.runs_table
),
format!(
"CREATE INDEX IF NOT EXISTS idx_{}_thread_id ON {} (thread_id)",
self.messages_table, self.messages_table
),
format!(
"CREATE TABLE IF NOT EXISTS {} (
namespace TEXT NOT NULL,
id TEXT NOT NULL,
data JSONB NOT NULL,
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
PRIMARY KEY (namespace, id)
)",
self.configs_table
),
format!(
"CREATE INDEX IF NOT EXISTS idx_{}_namespace_id ON {} (namespace, id)",
self.configs_table, self.configs_table
),
];
for stmt in statements {
sqlx::query(&stmt)
.execute(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
}
let run_migrations = [
("request", "JSONB"),
("run_input", "JSONB"),
("run_output", "JSONB"),
("termination_reason", "JSONB"),
("final_output", "TEXT"),
("error_payload", "JSONB"),
("dispatch_id", "TEXT"),
("session_id", "TEXT"),
("transport_request_id", "TEXT"),
("waiting", "JSONB"),
("outcome", "JSONB"),
("started_at", "BIGINT"),
("finished_at", "BIGINT"),
];
for (column, ty) in run_migrations {
let stmt = format!(
"ALTER TABLE {} ADD COLUMN IF NOT EXISTS {} {}",
self.runs_table, column, ty
);
sqlx::query(&stmt)
.execute(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
}
*ready = true;
Ok(())
}
}
struct PostgresConfigChangeSubscriber {
listener: PgListener,
}
#[async_trait]
impl ConfigChangeSubscriber for PostgresConfigChangeSubscriber {
async fn next(&mut self) -> Result<ConfigChangeEvent, StorageError> {
let notification = self
.listener
.recv()
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
serde_json::from_str(notification.payload())
.map_err(|error| StorageError::Serialization(error.to_string()))
}
}
#[async_trait]
impl ThreadStore for PostgresStore {
async fn load_thread(&self, thread_id: &str) -> Result<Option<Thread>, StorageError> {
self.ensure_schema().await?;
let sql = format!("SELECT data FROM {} WHERE id = $1", self.threads_table);
let row: Option<(serde_json::Value,)> = sqlx::query_as(&sql)
.bind(thread_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
match row {
Some((data,)) => {
let thread: Thread = serde_json::from_value(data)
.map_err(|e| StorageError::Serialization(e.to_string()))?;
Ok(Some(thread))
}
None => Ok(None),
}
}
async fn save_thread(&self, thread: &Thread) -> Result<(), StorageError> {
self.ensure_schema().await?;
let data =
serde_json::to_value(thread).map_err(|e| StorageError::Serialization(e.to_string()))?;
let sql = format!(
"INSERT INTO {} (id, data) VALUES ($1, $2)
ON CONFLICT (id) DO UPDATE SET data = $2, updated_at = now()",
self.threads_table
);
sqlx::query(&sql)
.bind(&thread.id)
.bind(&data)
.execute(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(())
}
async fn delete_thread(&self, thread_id: &str) -> Result<(), StorageError> {
self.ensure_schema().await?;
let mut tx = self
.pool
.begin()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
let delete_messages = format!("DELETE FROM {} WHERE thread_id = $1", self.messages_table);
sqlx::query(&delete_messages)
.bind(thread_id)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
let delete_thread = format!("DELETE FROM {} WHERE id = $1", self.threads_table);
sqlx::query(&delete_thread)
.bind(thread_id)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
tx.commit()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(())
}
async fn list_threads(&self, offset: usize, limit: usize) -> Result<Vec<String>, StorageError> {
self.ensure_schema().await?;
let sql = format!(
"SELECT id FROM {} ORDER BY updated_at DESC, id ASC LIMIT $1 OFFSET $2",
self.threads_table
);
let rows: Vec<(String,)> = sqlx::query_as(&sql)
.bind(limit as i64)
.bind(offset as i64)
.fetch_all(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(rows.into_iter().map(|(id,)| id).collect())
}
async fn load_messages(&self, thread_id: &str) -> Result<Option<Vec<Message>>, StorageError> {
self.ensure_schema().await?;
let sql = format!(
"SELECT data FROM {} WHERE thread_id = $1 ORDER BY updated_at DESC LIMIT 1",
self.messages_table
);
let row: Option<(serde_json::Value,)> = sqlx::query_as(&sql)
.bind(thread_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
match row {
Some((data,)) => {
let messages: Vec<Message> = serde_json::from_value(data)
.map_err(|e| StorageError::Serialization(e.to_string()))?;
Ok(Some(messages))
}
None => Ok(None),
}
}
async fn save_messages(
&self,
thread_id: &str,
messages: &[Message],
) -> Result<(), StorageError> {
self.ensure_schema().await?;
let msg_data = serde_json::to_value(messages)
.map_err(|e| StorageError::Serialization(e.to_string()))?;
let delete_sql = format!("DELETE FROM {} WHERE thread_id = $1", self.messages_table);
let insert_sql = format!(
"INSERT INTO {} (thread_id, data) VALUES ($1, $2)",
self.messages_table
);
let mut tx = self
.pool
.begin()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
sqlx::query(&delete_sql)
.bind(thread_id)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
sqlx::query(&insert_sql)
.bind(thread_id)
.bind(&msg_data)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
tx.commit()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(())
}
async fn delete_messages(&self, thread_id: &str) -> Result<(), StorageError> {
self.ensure_schema().await?;
let check_sql = format!("SELECT 1 FROM {} WHERE id = $1", self.threads_table);
let exists: Option<(i32,)> = sqlx::query_as(&check_sql)
.bind(thread_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
if exists.is_none() {
return Err(StorageError::NotFound(thread_id.to_owned()));
}
let sql = format!("DELETE FROM {} WHERE thread_id = $1", self.messages_table);
sqlx::query(&sql)
.bind(thread_id)
.execute(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(())
}
async fn update_thread_metadata(
&self,
id: &str,
metadata: awaken_contract::thread::ThreadMetadata,
) -> Result<(), StorageError> {
self.ensure_schema().await?;
let thread = self
.load_thread(id)
.await?
.ok_or_else(|| StorageError::NotFound(id.to_owned()))?;
let mut updated = thread;
updated.metadata = metadata;
self.save_thread(&updated).await
}
}
#[async_trait]
impl RunStore for PostgresStore {
async fn create_run(&self, record: &RunRecord) -> Result<(), StorageError> {
self.ensure_schema().await?;
let state_json = record
.state
.as_ref()
.and_then(|s| serde_json::to_value(s).ok());
let termination_reason_json = record
.termination_reason
.as_ref()
.and_then(|reason| serde_json::to_value(reason).ok());
let request_json = record
.request
.as_ref()
.and_then(|request| serde_json::to_value(request).ok());
let input_json = record
.input
.as_ref()
.and_then(|input| serde_json::to_value(input).ok());
let output_json = record
.output
.as_ref()
.and_then(|output| serde_json::to_value(output).ok());
let waiting_json = record
.waiting
.as_ref()
.and_then(|waiting| serde_json::to_value(waiting).ok());
let outcome_json = record
.outcome
.as_ref()
.and_then(|outcome| serde_json::to_value(outcome).ok());
let sql = format!(
"INSERT INTO {} (run_id, thread_id, agent_id, parent_run_id, request, run_input, run_output, status, termination_reason, final_output, error_payload, dispatch_id, session_id, transport_request_id, waiting, outcome, created_at, started_at, finished_at, updated_at, steps, input_tokens, output_tokens, state)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, $24)",
self.runs_table
);
sqlx::query(&sql)
.bind(&record.run_id)
.bind(&record.thread_id)
.bind(&record.agent_id)
.bind(&record.parent_run_id)
.bind(&request_json)
.bind(&input_json)
.bind(&output_json)
.bind(format!("{:?}", record.status).to_lowercase())
.bind(&termination_reason_json)
.bind(&record.final_output)
.bind(&record.error_payload)
.bind(&record.dispatch_id)
.bind(&record.session_id)
.bind(&record.transport_request_id)
.bind(&waiting_json)
.bind(&outcome_json)
.bind(record.created_at as i64)
.bind(record.started_at.map(|value| value as i64))
.bind(record.finished_at.map(|value| value as i64))
.bind(record.updated_at as i64)
.bind(record.steps as i32)
.bind(record.input_tokens as i64)
.bind(record.output_tokens as i64)
.bind(&state_json)
.execute(&self.pool)
.await
.map_err(|e| {
if e.to_string().contains("duplicate key")
|| e.to_string().contains("unique constraint")
{
StorageError::AlreadyExists(record.run_id.clone())
} else {
StorageError::Io(e.to_string())
}
})?;
Ok(())
}
async fn load_run(&self, run_id: &str) -> Result<Option<RunRecord>, StorageError> {
self.ensure_schema().await?;
let sql = format!(
"SELECT run_id, thread_id, agent_id, parent_run_id, request, run_input, run_output, status, termination_reason, final_output, error_payload, dispatch_id, session_id, transport_request_id, waiting, outcome, created_at, started_at, finished_at, updated_at, steps, input_tokens, output_tokens, state FROM {} WHERE run_id = $1",
self.runs_table
);
let row = sqlx::query(&sql)
.bind(run_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(row.map(run_record_from_pg_row))
}
async fn latest_run(&self, thread_id: &str) -> Result<Option<RunRecord>, StorageError> {
self.ensure_schema().await?;
let sql = format!(
"SELECT run_id, thread_id, agent_id, parent_run_id, request, run_input, run_output, status, termination_reason, final_output, error_payload, dispatch_id, session_id, transport_request_id, waiting, outcome, created_at, started_at, finished_at, updated_at, steps, input_tokens, output_tokens, state FROM {} WHERE thread_id = $1 ORDER BY updated_at DESC LIMIT 1",
self.runs_table
);
let row = sqlx::query(&sql)
.bind(thread_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(row.map(run_record_from_pg_row))
}
async fn list_runs(&self, query: &RunQuery) -> Result<RunPage, StorageError> {
self.ensure_schema().await?;
let mut conditions = Vec::new();
if query.thread_id.is_some() {
conditions.push("thread_id = $1".to_string());
}
if query.status.is_some() {
let idx = if query.thread_id.is_some() { 2 } else { 1 };
conditions.push(format!("status = ${idx}"));
}
let where_clause = if conditions.is_empty() {
String::new()
} else {
format!(" WHERE {}", conditions.join(" AND "))
};
let count_sql = format!("SELECT COUNT(*) FROM {}{}", self.runs_table, where_clause);
let list_sql = format!(
"SELECT run_id, thread_id, agent_id, parent_run_id, request, run_input, run_output, status, termination_reason, final_output, error_payload, dispatch_id, session_id, transport_request_id, waiting, outcome, created_at, started_at, finished_at, updated_at, steps, input_tokens, output_tokens, state FROM {}{} ORDER BY created_at ASC LIMIT {} OFFSET {}",
self.runs_table,
where_clause,
query.limit.clamp(1, 200),
query.offset
);
let (total,): (i64,) = {
let mut q = sqlx::query_as(&count_sql);
if let Some(ref tid) = query.thread_id {
q = q.bind(tid);
}
if let Some(status) = query.status {
q = q.bind(format!("{status:?}").to_lowercase());
}
q.fetch_one(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?
};
let rows = {
let mut q = sqlx::query(&list_sql);
if let Some(ref tid) = query.thread_id {
q = q.bind(tid);
}
if let Some(status) = query.status {
q = q.bind(format!("{status:?}").to_lowercase());
}
q.fetch_all(&self.pool)
.await
.map_err(|e| StorageError::Io(e.to_string()))?
};
let items: Vec<RunRecord> = rows.into_iter().map(run_record_from_pg_row).collect();
let has_more = (query.offset + items.len()) < total as usize;
Ok(RunPage {
items,
total: total as usize,
has_more,
})
}
}
#[async_trait]
impl ThreadRunStore for PostgresStore {
async fn checkpoint(
&self,
thread_id: &str,
messages: &[Message],
run: &RunRecord,
) -> Result<(), StorageError> {
self.ensure_schema().await?;
let msg_data = serde_json::to_value(messages)
.map_err(|e| StorageError::Serialization(e.to_string()))?;
let delete_sql = format!("DELETE FROM {} WHERE thread_id = $1", self.messages_table);
let insert_sql = format!(
"INSERT INTO {} (thread_id, data) VALUES ($1, $2)",
self.messages_table
);
let mut tx = self
.pool
.begin()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
let load_thread_sql = format!("SELECT data FROM {} WHERE id = $1", self.threads_table);
let existing_thread: Option<(serde_json::Value,)> = sqlx::query_as(&load_thread_sql)
.bind(thread_id)
.fetch_optional(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock before UNIX epoch")
.as_millis() as u64;
let mut thread = match existing_thread {
Some((data,)) => serde_json::from_value(data)
.map_err(|e| StorageError::Serialization(e.to_string()))?,
None => Thread::with_id(thread_id),
};
thread.metadata.created_at.get_or_insert(now);
thread.metadata.updated_at = Some(now);
thread.apply_run_projection(run);
let thread_data = serde_json::to_value(&thread)
.map_err(|e| StorageError::Serialization(e.to_string()))?;
let thread_sql = format!(
"INSERT INTO {} (id, data) VALUES ($1, $2)
ON CONFLICT (id) DO UPDATE SET data = $2, updated_at = now()",
self.threads_table
);
sqlx::query(&thread_sql)
.bind(thread_id)
.bind(&thread_data)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
sqlx::query(&delete_sql)
.bind(thread_id)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
sqlx::query(&insert_sql)
.bind(thread_id)
.bind(&msg_data)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
let state_json = run
.state
.as_ref()
.and_then(|s| serde_json::to_value(s).ok());
let termination_reason_json = run
.termination_reason
.as_ref()
.and_then(|reason| serde_json::to_value(reason).ok());
let request_json = run
.request
.as_ref()
.and_then(|request| serde_json::to_value(request).ok());
let input_json = run
.input
.as_ref()
.and_then(|input| serde_json::to_value(input).ok());
let output_json = run
.output
.as_ref()
.and_then(|output| serde_json::to_value(output).ok());
let waiting_json = run
.waiting
.as_ref()
.and_then(|waiting| serde_json::to_value(waiting).ok());
let outcome_json = run
.outcome
.as_ref()
.and_then(|outcome| serde_json::to_value(outcome).ok());
let run_sql = format!(
"INSERT INTO {} (run_id, thread_id, agent_id, parent_run_id, request, run_input, run_output, status, termination_reason, final_output, error_payload, dispatch_id, session_id, transport_request_id, waiting, outcome, created_at, started_at, finished_at, updated_at, steps, input_tokens, output_tokens, state)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, $24)
ON CONFLICT (run_id) DO UPDATE SET
request = $5, run_input = $6, run_output = $7, status = $8,
termination_reason = $9, final_output = $10,
error_payload = $11, dispatch_id = $12, session_id = $13,
transport_request_id = $14, waiting = $15, outcome = $16,
started_at = $18, finished_at = $19, updated_at = $20,
steps = $21, input_tokens = $22, output_tokens = $23, state = $24",
self.runs_table
);
sqlx::query(&run_sql)
.bind(&run.run_id)
.bind(&run.thread_id)
.bind(&run.agent_id)
.bind(&run.parent_run_id)
.bind(&request_json)
.bind(&input_json)
.bind(&output_json)
.bind(format!("{:?}", run.status).to_lowercase())
.bind(&termination_reason_json)
.bind(&run.final_output)
.bind(&run.error_payload)
.bind(&run.dispatch_id)
.bind(&run.session_id)
.bind(&run.transport_request_id)
.bind(&waiting_json)
.bind(&outcome_json)
.bind(run.created_at as i64)
.bind(run.started_at.map(|value| value as i64))
.bind(run.finished_at.map(|value| value as i64))
.bind(run.updated_at as i64)
.bind(run.steps as i32)
.bind(run.input_tokens as i64)
.bind(run.output_tokens as i64)
.bind(&state_json)
.execute(&mut *tx)
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
tx.commit()
.await
.map_err(|e| StorageError::Io(e.to_string()))?;
Ok(())
}
}
fn run_record_from_pg_row(row: PgRow) -> RunRecord {
let status: String = row.get("status");
let state: Option<serde_json::Value> = row.get("state");
let request: Option<serde_json::Value> = row.get("request");
let input: Option<serde_json::Value> = row.get("run_input");
let output: Option<serde_json::Value> = row.get("run_output");
let termination_reason: Option<serde_json::Value> = row.get("termination_reason");
let waiting: Option<serde_json::Value> = row.get("waiting");
let outcome: Option<serde_json::Value> = row.get("outcome");
let created_at: i64 = row.get("created_at");
let started_at: Option<i64> = row.get("started_at");
let finished_at: Option<i64> = row.get("finished_at");
let updated_at: i64 = row.get("updated_at");
let steps: i32 = row.get("steps");
let input_tokens: i64 = row.get("input_tokens");
let output_tokens: i64 = row.get("output_tokens");
RunRecord {
run_id: row.get("run_id"),
thread_id: row.get("thread_id"),
agent_id: row.get("agent_id"),
parent_run_id: row.get("parent_run_id"),
request: request.and_then(|value| serde_json::from_value(value).ok()),
input: input.and_then(|value| serde_json::from_value(value).ok()),
output: output.and_then(|value| serde_json::from_value(value).ok()),
status: parse_run_status(&status),
termination_reason: termination_reason.and_then(|value| serde_json::from_value(value).ok()),
final_output: row.get("final_output"),
error_payload: row.get("error_payload"),
dispatch_id: row.get("dispatch_id"),
session_id: row.get("session_id"),
transport_request_id: row.get("transport_request_id"),
waiting: waiting.and_then(|value| serde_json::from_value(value).ok()),
outcome: outcome.and_then(|value| serde_json::from_value(value).ok()),
created_at: created_at as u64,
started_at: started_at.map(|value| value as u64),
finished_at: finished_at.map(|value| value as u64),
updated_at: updated_at as u64,
steps: steps as usize,
input_tokens: input_tokens as u64,
output_tokens: output_tokens as u64,
state: state.and_then(|value| serde_json::from_value(value).ok()),
}
}
fn parse_run_status(s: &str) -> awaken_contract::contract::lifecycle::RunStatus {
use awaken_contract::contract::lifecycle::RunStatus;
match s {
"created" => RunStatus::Created,
"running" => RunStatus::Running,
"waiting" => RunStatus::Waiting,
"done" => RunStatus::Done,
_ => RunStatus::Running,
}
}
#[async_trait]
impl ConfigStore for PostgresStore {
async fn get(
&self,
namespace: &str,
id: &str,
) -> Result<Option<serde_json::Value>, StorageError> {
self.ensure_schema().await?;
let sql = format!(
"SELECT data FROM {} WHERE namespace = $1 AND id = $2",
self.configs_table
);
let row: Option<(serde_json::Value,)> = sqlx::query_as(&sql)
.bind(namespace)
.bind(id)
.fetch_optional(&self.pool)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
Ok(row.map(|(value,)| value))
}
async fn list(
&self,
namespace: &str,
offset: usize,
limit: usize,
) -> Result<Vec<(String, serde_json::Value)>, StorageError> {
self.ensure_schema().await?;
let limit = limit.min(i64::MAX as usize) as i64;
let offset = offset.min(i64::MAX as usize) as i64;
let sql = format!(
"SELECT id, data FROM {} WHERE namespace = $1 ORDER BY id ASC LIMIT $2 OFFSET $3",
self.configs_table
);
sqlx::query_as(&sql)
.bind(namespace)
.bind(limit)
.bind(offset)
.fetch_all(&self.pool)
.await
.map_err(|error| StorageError::Io(error.to_string()))
}
async fn put(
&self,
namespace: &str,
id: &str,
value: &serde_json::Value,
) -> Result<(), StorageError> {
self.ensure_schema().await?;
let mut tx = self
.pool
.begin()
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
let sql = format!(
"INSERT INTO {} (namespace, id, data) VALUES ($1, $2, $3)
ON CONFLICT (namespace, id) DO UPDATE SET data = $3, updated_at = now()",
self.configs_table
);
sqlx::query(&sql)
.bind(namespace)
.bind(id)
.bind(value)
.execute(&mut *tx)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
let payload = serde_json::to_string(&ConfigChangeEvent {
namespace: namespace.to_string(),
id: id.to_string(),
kind: ConfigChangeKind::Put,
})
.map_err(|error| StorageError::Serialization(error.to_string()))?;
sqlx::query("SELECT pg_notify($1, $2)")
.bind(&self.config_notify_channel)
.bind(payload)
.execute(&mut *tx)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
tx.commit()
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
Ok(())
}
async fn delete(&self, namespace: &str, id: &str) -> Result<(), StorageError> {
self.ensure_schema().await?;
let mut tx = self
.pool
.begin()
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
let sql = format!(
"DELETE FROM {} WHERE namespace = $1 AND id = $2",
self.configs_table
);
let result = sqlx::query(&sql)
.bind(namespace)
.bind(id)
.execute(&mut *tx)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
if result.rows_affected() > 0 {
let payload = serde_json::to_string(&ConfigChangeEvent {
namespace: namespace.to_string(),
id: id.to_string(),
kind: ConfigChangeKind::Delete,
})
.map_err(|error| StorageError::Serialization(error.to_string()))?;
sqlx::query("SELECT pg_notify($1, $2)")
.bind(&self.config_notify_channel)
.bind(payload)
.execute(&mut *tx)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
}
tx.commit()
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
Ok(())
}
}
#[async_trait]
impl ConfigChangeNotifier for PostgresStore {
async fn subscribe(&self) -> Result<Box<dyn ConfigChangeSubscriber>, StorageError> {
self.ensure_schema().await?;
let mut listener = PgListener::connect_with(&self.pool)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
listener
.listen(&self.config_notify_channel)
.await
.map_err(|error| StorageError::Io(error.to_string()))?;
Ok(Box::new(PostgresConfigChangeSubscriber { listener }))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_run_status_known_values() {
use awaken_contract::contract::lifecycle::RunStatus;
assert!(matches!(parse_run_status("created"), RunStatus::Created));
assert!(matches!(parse_run_status("running"), RunStatus::Running));
assert!(matches!(parse_run_status("waiting"), RunStatus::Waiting));
assert!(matches!(parse_run_status("done"), RunStatus::Done));
}
#[test]
fn parse_run_status_unknown_defaults_to_running() {
use awaken_contract::contract::lifecycle::RunStatus;
assert!(matches!(parse_run_status("unknown"), RunStatus::Running));
assert!(matches!(parse_run_status(""), RunStatus::Running));
}
#[test]
fn postgres_store_default_table_names() {
let prefix = "test_prefix";
assert_eq!(format!("{prefix}_threads"), "test_prefix_threads");
assert_eq!(format!("{prefix}_runs"), "test_prefix_runs");
assert_eq!(format!("{prefix}_configs"), "test_prefix_configs");
assert_eq!(
format!("{prefix}_config_changes"),
"test_prefix_config_changes"
);
}
#[tokio::test]
#[ignore]
async fn schema_initialization() {
let pool = PgPool::connect("postgres://localhost/awaken_test")
.await
.unwrap();
let store = PostgresStore::with_prefix(pool, "test_schema_init");
store.ensure_schema().await.unwrap();
store.ensure_schema().await.unwrap();
}
#[tokio::test]
#[ignore]
async fn connection_error_handling() {
let pool = PgPool::connect("postgres://localhost:19999/nonexistent")
.await
.unwrap_err();
let _ = pool;
}
#[tokio::test]
#[ignore]
async fn thread_crud_operations() {
let pool = PgPool::connect("postgres://localhost/awaken_test")
.await
.unwrap();
let store = PostgresStore::with_prefix(pool, "test_crud");
store.ensure_schema().await.unwrap();
let thread = Thread::new();
store.save_thread(&thread).await.unwrap();
let loaded = store.load_thread(&thread.id).await.unwrap().unwrap();
assert_eq!(loaded.id, thread.id);
store.delete_thread(&thread.id).await.unwrap();
assert!(store.load_thread(&thread.id).await.unwrap().is_none());
}
#[tokio::test]
#[ignore]
async fn run_create_duplicate_returns_already_exists() {
use awaken_contract::contract::lifecycle::RunStatus;
let pool = PgPool::connect("postgres://localhost/awaken_test")
.await
.unwrap();
let store = PostgresStore::with_prefix(pool, "test_dup_run");
store.ensure_schema().await.unwrap();
let run = RunRecord {
run_id: format!("dup-{}", uuid::Uuid::now_v7()),
thread_id: "t-1".to_string(),
agent_id: "agent".to_string(),
parent_run_id: None,
request: None,
input: None,
output: None,
status: RunStatus::Running,
termination_reason: None,
final_output: None,
error_payload: None,
dispatch_id: None,
session_id: None,
transport_request_id: None,
waiting: None,
outcome: None,
created_at: 100,
started_at: None,
finished_at: None,
updated_at: 100,
steps: 0,
input_tokens: 0,
output_tokens: 0,
state: None,
};
store.create_run(&run).await.unwrap();
let err = store.create_run(&run).await.unwrap_err();
assert!(matches!(err, StorageError::AlreadyExists(_)));
}
#[tokio::test]
#[ignore]
async fn checkpoint_atomicity() {
use awaken_contract::contract::lifecycle::RunStatus;
use awaken_contract::contract::message::Message;
let pool = PgPool::connect("postgres://localhost/awaken_test")
.await
.unwrap();
let store = PostgresStore::with_prefix(pool, "test_checkpoint");
store.ensure_schema().await.unwrap();
let thread_id = format!("t-{}", uuid::Uuid::now_v7());
let msgs = vec![Message::user("checkpoint test")];
let run = RunRecord {
run_id: format!("r-{}", uuid::Uuid::now_v7()),
thread_id: thread_id.clone(),
agent_id: "agent".to_string(),
parent_run_id: None,
request: None,
input: None,
output: None,
status: RunStatus::Running,
termination_reason: None,
final_output: None,
error_payload: None,
dispatch_id: None,
session_id: None,
transport_request_id: None,
waiting: None,
outcome: None,
created_at: 100,
started_at: None,
finished_at: None,
updated_at: 100,
steps: 1,
input_tokens: 10,
output_tokens: 20,
state: None,
};
store.checkpoint(&thread_id, &msgs, &run).await.unwrap();
let loaded_msgs = store.load_messages(&thread_id).await.unwrap().unwrap();
assert_eq!(loaded_msgs.len(), 1);
let loaded_run = store.load_run(&run.run_id).await.unwrap().unwrap();
assert_eq!(loaded_run.thread_id, thread_id);
}
}