use std::str::FromStr;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use crate::{AgentError, AgentSnapshot, AgentStorage, Result};
use ai_agents_core::traits::storage::StorageCapability;
use ai_agents_core::{
FactCategory, FactFilter, KeyFact, SessionFilter, SessionMetadata, SessionSummary,
};
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SqliteMetadata {
#[serde(default)]
pub tags: Vec<String>,
#[serde(default)]
pub user_id: Option<String>,
#[serde(default)]
pub custom: std::collections::HashMap<String, serde_json::Value>,
#[serde(default)]
pub priority: Option<i32>,
#[serde(default)]
pub expires_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionInfo {
pub session_id: String,
pub agent_id: String,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
pub message_count: usize,
pub current_state: Option<String>,
#[serde(default)]
pub metadata: SqliteMetadata,
}
type SessionInfoRow = (
String,
String,
String,
String,
i64,
Option<String>,
Option<String>,
);
#[derive(Debug, Clone, Default)]
pub struct SessionQuery {
pub agent_id: Option<String>,
pub state: Option<String>,
pub tag: Option<String>,
pub user_id: Option<String>,
pub created_after: Option<DateTime<Utc>>,
pub created_before: Option<DateTime<Utc>>,
pub updated_after: Option<DateTime<Utc>>,
pub limit: Option<u32>,
pub offset: Option<u32>,
pub order_by: SessionOrderBy,
}
#[derive(Debug, Clone, Default)]
pub enum SessionOrderBy {
#[default]
UpdatedAtDesc,
UpdatedAtAsc,
CreatedAtDesc,
CreatedAtAsc,
MessageCountDesc,
}
#[cfg(feature = "sqlite")]
pub struct SqliteStorage {
pool: sqlx::SqlitePool,
}
#[cfg(feature = "sqlite")]
impl SqliteStorage {
pub async fn new(path: &str) -> Result<Self> {
let pool = Self::connect(path).await?;
let storage = Self { pool };
storage.run_migrations().await?;
Ok(storage)
}
pub async fn in_memory() -> Result<Self> {
Self::new(":memory:").await
}
pub async fn close(&self) {
self.pool.close().await;
}
async fn connect(path: &str) -> Result<sqlx::SqlitePool> {
if path != ":memory:"
&& let Some(parent) = std::path::Path::new(path).parent()
&& !parent.as_os_str().is_empty()
{
std::fs::create_dir_all(parent).map_err(|e| {
AgentError::Persistence(format!(
"failed to create directory {}: {}",
parent.display(),
e
))
})?;
}
let options = sqlx::sqlite::SqliteConnectOptions::from_str(path)
.map_err(|e| AgentError::Persistence(e.to_string()))?
.create_if_missing(true)
.journal_mode(sqlx::sqlite::SqliteJournalMode::Wal);
sqlx::SqlitePool::connect_with(options)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))
}
async fn run_migrations(&self) -> Result<()> {
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS sessions (
session_id TEXT PRIMARY KEY,
agent_id TEXT NOT NULL,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL,
message_count INTEGER NOT NULL DEFAULT 0,
current_state TEXT,
data TEXT NOT NULL,
metadata TEXT
)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_sessions_agent_id ON sessions(agent_id)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_sessions_updated_at ON sessions(updated_at)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS session_tags (
session_id TEXT NOT NULL,
tag TEXT NOT NULL,
PRIMARY KEY (session_id, tag),
FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE
)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS actor_facts (
id TEXT NOT NULL,
agent_id TEXT NOT NULL,
actor_id TEXT NOT NULL,
category TEXT NOT NULL,
content TEXT NOT NULL,
confidence REAL NOT NULL,
salience REAL NOT NULL DEFAULT 1.0,
extracted_at TEXT NOT NULL,
last_accessed TEXT,
source_message_id TEXT,
source_language TEXT,
PRIMARY KEY (agent_id, actor_id, id)
)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_actor_facts_agent_actor
ON actor_facts(agent_id, actor_id)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_actor_facts_category
ON actor_facts(agent_id, actor_id, category)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS actor_relationships (
agent_id TEXT NOT NULL,
actor_id TEXT NOT NULL,
actor_name TEXT,
dimensions_json TEXT NOT NULL,
notable_events_json TEXT NOT NULL,
interaction_count INTEGER NOT NULL,
first_interaction TEXT NOT NULL,
last_interaction TEXT NOT NULL,
metadata_json TEXT NOT NULL,
PRIMARY KEY (agent_id, actor_id)
)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_actor_relationships_last_interaction
ON actor_relationships(agent_id, last_interaction)
"#,
)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
fn extract_session_info(
session_id: String,
agent_id: String,
created_at: String,
updated_at: String,
message_count: i64,
current_state: Option<String>,
metadata_json: Option<String>,
) -> SessionInfo {
let metadata: SqliteMetadata = metadata_json
.and_then(|m| serde_json::from_str(&m).ok())
.unwrap_or_default();
SessionInfo {
session_id,
agent_id,
created_at: DateTime::parse_from_rfc3339(&created_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
updated_at: DateTime::parse_from_rfc3339(&updated_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
message_count: message_count as usize,
current_state,
metadata,
}
}
pub async fn save_with_metadata(
&self,
session_id: &str,
snapshot: &AgentSnapshot,
metadata: &SqliteMetadata,
) -> Result<()> {
self.save_record_with_metadata(session_id, snapshot, metadata, &metadata.tags)
.await
}
async fn save_record_with_metadata<M: Serialize>(
&self,
session_id: &str,
snapshot: &AgentSnapshot,
metadata: &M,
tags: &[String],
) -> Result<()> {
let now = Utc::now().to_rfc3339();
let data =
serde_json::to_string(snapshot).map_err(|e| AgentError::Persistence(e.to_string()))?;
let metadata_json =
serde_json::to_string(metadata).map_err(|e| AgentError::Persistence(e.to_string()))?;
let message_count = snapshot.memory.messages.len() as i64;
let current_state = snapshot
.state_machine
.as_ref()
.map(|state| state.current_state.clone());
let mut transaction = self
.pool
.begin()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query(
r#"
INSERT INTO sessions (session_id, agent_id, created_at, updated_at, message_count, current_state, data, metadata)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(session_id) DO UPDATE SET
agent_id = excluded.agent_id,
updated_at = excluded.updated_at,
message_count = excluded.message_count,
current_state = excluded.current_state,
data = excluded.data,
metadata = excluded.metadata
"#,
)
.bind(session_id)
.bind(&snapshot.agent_id)
.bind(&now)
.bind(&now)
.bind(message_count)
.bind(¤t_state)
.bind(&data)
.bind(&metadata_json)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query("DELETE FROM session_tags WHERE session_id = ?")
.bind(session_id)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
for tag in tags {
sqlx::query("INSERT INTO session_tags (session_id, tag) VALUES (?, ?)")
.bind(session_id)
.bind(tag)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
}
transaction
.commit()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))
}
pub async fn get_metadata(&self, session_id: &str) -> Result<Option<SqliteMetadata>> {
let row: Option<(Option<String>,)> =
sqlx::query_as("SELECT metadata FROM sessions WHERE session_id = ?")
.bind(session_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
match row {
Some((Some(metadata_json),)) => {
let metadata: SqliteMetadata = serde_json::from_str(&metadata_json)
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(Some(metadata))
}
Some((None,)) => Ok(Some(SqliteMetadata::default())),
None => Ok(None),
}
}
pub async fn list_sessions_by_agent(&self, agent_id: &str) -> Result<Vec<SessionInfo>> {
let rows: Vec<SessionInfoRow> = sqlx::query_as(
r#"
SELECT session_id, agent_id, created_at, updated_at, message_count, current_state, metadata
FROM sessions
WHERE agent_id = ?
ORDER BY updated_at DESC
"#,
)
.bind(agent_id)
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows
.into_iter()
.map(
|(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)| {
Self::extract_session_info(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)
},
)
.collect())
}
pub async fn search_sessions(&self, query: &SessionQuery) -> Result<Vec<SessionInfo>> {
let mut sql = String::from(
r#"
SELECT DISTINCT s.session_id, s.agent_id, s.created_at, s.updated_at,
s.message_count, s.current_state, s.metadata
FROM sessions s
LEFT JOIN session_tags t ON s.session_id = t.session_id
WHERE 1=1
"#,
);
if query.agent_id.is_some() {
sql.push_str(" AND s.agent_id = ?");
}
if query.state.is_some() {
sql.push_str(" AND s.current_state = ?");
}
if query.tag.is_some() {
sql.push_str(" AND t.tag = ?");
}
if query.user_id.is_some() {
sql.push_str(" AND json_extract(s.metadata, '$.user_id') = ?");
}
if query.created_after.is_some() {
sql.push_str(" AND s.created_at >= ?");
}
if query.created_before.is_some() {
sql.push_str(" AND s.created_at <= ?");
}
if query.updated_after.is_some() {
sql.push_str(" AND s.updated_at >= ?");
}
sql.push_str(match query.order_by {
SessionOrderBy::UpdatedAtDesc => " ORDER BY s.updated_at DESC",
SessionOrderBy::UpdatedAtAsc => " ORDER BY s.updated_at ASC",
SessionOrderBy::CreatedAtDesc => " ORDER BY s.created_at DESC",
SessionOrderBy::CreatedAtAsc => " ORDER BY s.created_at ASC",
SessionOrderBy::MessageCountDesc => " ORDER BY s.message_count DESC",
});
if let Some(limit) = query.limit {
sql.push_str(&format!(" LIMIT {}", limit));
}
if let Some(offset) = query.offset {
sql.push_str(&format!(" OFFSET {}", offset));
}
let mut q = sqlx::query_as::<
_,
(
String,
String,
String,
String,
i64,
Option<String>,
Option<String>,
),
>(&sql);
if let Some(ref agent_id) = query.agent_id {
q = q.bind(agent_id);
}
if let Some(ref state) = query.state {
q = q.bind(state);
}
if let Some(ref tag) = query.tag {
q = q.bind(tag);
}
if let Some(ref user_id) = query.user_id {
q = q.bind(user_id);
}
if let Some(ref created_after) = query.created_after {
q = q.bind(created_after.to_rfc3339());
}
if let Some(ref created_before) = query.created_before {
q = q.bind(created_before.to_rfc3339());
}
if let Some(ref updated_after) = query.updated_after {
q = q.bind(updated_after.to_rfc3339());
}
let rows = q
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows
.into_iter()
.map(
|(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)| {
Self::extract_session_info(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)
},
)
.collect())
}
pub async fn expire_sessions(&self, before: DateTime<Utc>) -> Result<usize> {
let result = sqlx::query("DELETE FROM sessions WHERE updated_at < ?")
.bind(before.to_rfc3339())
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(result.rows_affected() as usize)
}
pub async fn exists(&self, session_id: &str) -> Result<bool> {
let row: Option<(i64,)> =
sqlx::query_as("SELECT 1 FROM sessions WHERE session_id = ? LIMIT 1")
.bind(session_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(row.is_some())
}
pub async fn get_session_info(&self, session_id: &str) -> Result<Option<SessionInfo>> {
let row: Option<SessionInfoRow> = sqlx::query_as(
r#"
SELECT session_id, agent_id, created_at, updated_at, message_count, current_state, metadata
FROM sessions WHERE session_id = ?
"#,
)
.bind(session_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(row.map(
|(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)| {
Self::extract_session_info(
session_id,
agent_id,
created_at,
updated_at,
message_count,
current_state,
metadata,
)
},
))
}
}
#[cfg(feature = "sqlite")]
#[async_trait]
impl AgentStorage for SqliteStorage {
fn supports(&self, capability: StorageCapability) -> bool {
matches!(
capability,
StorageCapability::Snapshot
| StorageCapability::SessionMetadata
| StorageCapability::SessionFiltering
| StorageCapability::ExpiryCleanup
| StorageCapability::ActorFacts
| StorageCapability::ActorRelationships
| StorageCapability::ActorDataDeletion
)
}
async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
let now = Utc::now().to_rfc3339();
let data =
serde_json::to_string(snapshot).map_err(|e| AgentError::Persistence(e.to_string()))?;
let message_count = snapshot.memory.messages.len() as i64;
let current_state = snapshot
.state_machine
.as_ref()
.map(|state| state.current_state.clone());
sqlx::query(
r#"
INSERT INTO sessions (session_id, agent_id, created_at, updated_at, message_count, current_state, data)
VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(session_id) DO UPDATE SET
agent_id = excluded.agent_id,
updated_at = excluded.updated_at,
message_count = excluded.message_count,
current_state = excluded.current_state,
data = excluded.data
"#,
)
.bind(session_id)
.bind(&snapshot.agent_id)
.bind(&now)
.bind(&now)
.bind(message_count)
.bind(¤t_state)
.bind(&data)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn save_snapshot_with_metadata(
&self,
session_id: &str,
snapshot: &AgentSnapshot,
metadata: &SessionMetadata,
) -> Result<()> {
self.save_record_with_metadata(session_id, snapshot, metadata, &metadata.tags)
.await
}
async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
let row: Option<(String,)> =
sqlx::query_as("SELECT data FROM sessions WHERE session_id = ?")
.bind(session_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
match row {
Some((data,)) => {
let snapshot = serde_json::from_str(&data)
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(Some(snapshot))
}
None => Ok(None),
}
}
async fn delete(&self, session_id: &str) -> Result<()> {
sqlx::query("DELETE FROM sessions WHERE session_id = ?")
.bind(session_id)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn list_sessions(&self) -> Result<Vec<String>> {
let rows: Vec<(String,)> =
sqlx::query_as("SELECT session_id FROM sessions ORDER BY updated_at DESC")
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows.into_iter().map(|(id,)| id).collect())
}
async fn save_metadata(&self, session_id: &str, metadata: &SessionMetadata) -> Result<()> {
let meta_json =
serde_json::to_string(metadata).map_err(|e| AgentError::Persistence(e.to_string()))?;
let mut transaction = self
.pool
.begin()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
let update = sqlx::query("UPDATE sessions SET metadata = ? WHERE session_id = ?")
.bind(&meta_json)
.bind(session_id)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
if update.rows_affected() == 0 {
return Err(AgentError::Persistence(format!(
"session not found: {session_id}"
)));
}
sqlx::query("DELETE FROM session_tags WHERE session_id = ?")
.bind(session_id)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
for tag in &metadata.tags {
sqlx::query("INSERT INTO session_tags (session_id, tag) VALUES (?, ?)")
.bind(session_id)
.bind(tag)
.execute(&mut *transaction)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
}
transaction
.commit()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn load_metadata(&self, session_id: &str) -> Result<Option<SessionMetadata>> {
let row: Option<(Option<String>,)> =
sqlx::query_as("SELECT metadata FROM sessions WHERE session_id = ?")
.bind(session_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
match row {
Some((Some(json),)) => {
let meta: SessionMetadata = serde_json::from_str(&json).unwrap_or_default();
Ok(Some(meta))
}
Some((None,)) => Ok(Some(SessionMetadata::default())),
None => Ok(None),
}
}
async fn list_sessions_filtered(&self, filter: &SessionFilter) -> Result<Vec<SessionSummary>> {
let mut sql = String::from(
"SELECT session_id, agent_id, created_at, updated_at, message_count, metadata FROM sessions WHERE 1=1",
);
let mut binds: Vec<String> = vec![];
if let Some(ref agent_id) = filter.agent_id {
sql.push_str(" AND agent_id = ?");
binds.push(agent_id.clone());
}
if let Some(ref actor_id) = filter.actor_id {
sql.push_str(" AND json_extract(metadata, '$.actor_id') = ?");
binds.push(actor_id.clone());
}
if let Some(ref tags) = filter.tags {
for tag in tags {
sql.push_str(
" AND session_id IN (SELECT session_id FROM session_tags WHERE tag = ?)",
);
binds.push(tag.clone());
}
}
if let Some(ref after) = filter.created_after {
sql.push_str(" AND created_at >= ?");
binds.push(after.to_rfc3339());
}
if let Some(ref before) = filter.created_before {
sql.push_str(" AND created_at <= ?");
binds.push(before.to_rfc3339());
}
sql.push_str(" ORDER BY updated_at DESC");
if let Some(limit) = filter.limit {
sql.push_str(&format!(" LIMIT {}", limit));
}
let mut q =
sqlx::query_as::<_, (String, String, String, String, i64, Option<String>)>(&sql);
for b in &binds {
q = q.bind(b);
}
let rows = q
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows
.into_iter()
.map(
|(session_id, agent_id, created_at, updated_at, message_count, metadata_json)| {
let meta: SessionMetadata = metadata_json
.and_then(|j| serde_json::from_str(&j).ok())
.unwrap_or_default();
SessionSummary {
session_id,
agent_id,
actor_id: meta.actor_id,
tags: meta.tags,
created_at: DateTime::parse_from_rfc3339(&created_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
last_active: DateTime::parse_from_rfc3339(&updated_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
message_count: message_count as usize,
}
},
)
.collect())
}
async fn cleanup_expired(&self) -> Result<usize> {
let now = Utc::now().to_rfc3339();
let result = sqlx::query(
r#"
DELETE FROM sessions
WHERE session_id IN (
SELECT s.session_id FROM sessions s
WHERE json_extract(s.metadata, '$.ttl_seconds') IS NOT NULL
AND datetime(s.updated_at, '+' || json_extract(s.metadata, '$.ttl_seconds') || ' seconds') < ?
)
"#,
)
.bind(&now)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(result.rows_affected() as usize)
}
async fn save_facts(&self, agent_id: &str, actor_id: &str, facts: &[KeyFact]) -> Result<()> {
let mut transaction = self
.pool
.begin()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
for fact in facts {
let category_str = fact.category.to_string();
let extracted_at = fact.extracted_at.to_rfc3339();
let last_accessed = fact.last_accessed.map(|dt| dt.to_rfc3339());
sqlx::query(
r#"
INSERT INTO actor_facts
(id, agent_id, actor_id, category, content, confidence, salience,
extracted_at, last_accessed, source_message_id, source_language)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(agent_id, actor_id, id) DO UPDATE SET
category = excluded.category,
content = excluded.content,
confidence = excluded.confidence,
salience = excluded.salience,
last_accessed = excluded.last_accessed,
source_message_id = excluded.source_message_id,
source_language = excluded.source_language
"#,
)
.bind(&fact.id)
.bind(agent_id)
.bind(actor_id)
.bind(&category_str)
.bind(&fact.content)
.bind(fact.confidence)
.bind(fact.salience)
.bind(&extracted_at)
.bind(&last_accessed)
.bind(&fact.source_message_id)
.bind(&fact.source_language)
.execute(&mut *transaction)
.await
.map_err(|error| AgentError::Persistence(error.to_string()))?;
}
transaction
.commit()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn load_facts(&self, agent_id: &str, actor_id: &str) -> Result<Vec<KeyFact>> {
let rows: Vec<(
String,
String,
String,
f64,
f64,
String,
Option<String>,
Option<String>,
Option<String>,
)> = sqlx::query_as(
r#"
SELECT id, category, content, confidence, salience,
extracted_at, last_accessed, source_message_id, source_language
FROM actor_facts
WHERE agent_id = ? AND actor_id = ?
ORDER BY (salience * confidence) DESC
"#,
)
.bind(agent_id)
.bind(actor_id)
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows
.into_iter()
.map(
|(
id,
category,
content,
confidence,
salience,
extracted_at,
last_accessed,
source_msg,
source_lang,
)| {
KeyFact {
id,
actor_id: Some(actor_id.to_string()),
category: parse_fact_category(&category),
content,
confidence: confidence as f32,
salience: salience as f32,
extracted_at: DateTime::parse_from_rfc3339(&extracted_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
last_accessed: last_accessed.and_then(|ref s| {
DateTime::parse_from_rfc3339(s)
.ok()
.map(|dt| dt.with_timezone(&Utc))
}),
source_message_id: source_msg,
source_language: source_lang,
}
},
)
.collect())
}
async fn query_facts(&self, agent_id: &str, filter: &FactFilter) -> Result<Vec<KeyFact>> {
let mut sql = String::from(
"SELECT id, actor_id, category, content, confidence, salience, extracted_at, last_accessed, source_message_id, source_language FROM actor_facts WHERE agent_id = ?",
);
let mut str_binds: Vec<String> = vec![agent_id.to_string()];
if let Some(ref actor_id) = filter.actor_id {
sql.push_str(" AND actor_id = ?");
str_binds.push(actor_id.clone());
}
if let Some(ref category) = filter.category {
sql.push_str(" AND category = ?");
str_binds.push(category.to_string());
}
if let Some(min_conf) = filter.min_confidence {
sql.push_str(&format!(" AND confidence >= {}", min_conf));
}
if let Some(min_sal) = filter.min_salience {
sql.push_str(&format!(" AND salience >= {}", min_sal));
}
sql.push_str(" ORDER BY (salience * confidence) DESC");
if let Some(limit) = filter.limit {
sql.push_str(&format!(" LIMIT {}", limit));
}
let mut q = sqlx::query_as::<
_,
(
String,
String,
String,
String,
f64,
f64,
String,
Option<String>,
Option<String>,
Option<String>,
),
>(&sql);
for b in &str_binds {
q = q.bind(b);
}
let rows = q
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows
.into_iter()
.map(
|(
id,
actor_id,
category,
content,
confidence,
salience,
extracted_at,
last_accessed,
source_msg,
source_lang,
)| {
KeyFact {
id,
actor_id: Some(actor_id),
category: parse_fact_category(&category),
content,
confidence: confidence as f32,
salience: salience as f32,
extracted_at: DateTime::parse_from_rfc3339(&extracted_at)
.map(|dt| dt.with_timezone(&Utc))
.unwrap_or_else(|_| Utc::now()),
last_accessed: last_accessed.and_then(|ref s| {
DateTime::parse_from_rfc3339(s)
.ok()
.map(|dt| dt.with_timezone(&Utc))
}),
source_message_id: source_msg,
source_language: source_lang,
}
},
)
.collect())
}
async fn delete_fact(&self, agent_id: &str, actor_id: &str, fact_id: &str) -> Result<()> {
sqlx::query("DELETE FROM actor_facts WHERE agent_id = ? AND actor_id = ? AND id = ?")
.bind(agent_id)
.bind(actor_id)
.bind(fact_id)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn delete_actor_data(&self, agent_id: &str, actor_id: &str) -> Result<()> {
let mut transaction = self
.pool
.begin()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
sqlx::query("DELETE FROM actor_facts WHERE agent_id = ? AND actor_id = ?")
.bind(agent_id)
.bind(actor_id)
.execute(&mut *transaction)
.await
.map_err(|error| AgentError::Persistence(error.to_string()))?;
sqlx::query("DELETE FROM actor_relationships WHERE agent_id = ? AND actor_id = ?")
.bind(agent_id)
.bind(actor_id)
.execute(&mut *transaction)
.await
.map_err(|error| AgentError::Persistence(error.to_string()))?;
sqlx::query(
"DELETE FROM sessions WHERE agent_id = ? AND json_extract(metadata, '$.actor_id') = ?",
)
.bind(agent_id)
.bind(actor_id)
.execute(&mut *transaction)
.await
.map_err(|error| AgentError::Persistence(error.to_string()))?;
transaction
.commit()
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn save_relationship(
&self,
agent_id: &str,
actor_id: &str,
relationship: &serde_json::Value,
) -> Result<()> {
let actor_name = relationship
.get("actor_name")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
let dimensions_json = serde_json::to_string(
relationship
.get("dimensions")
.unwrap_or(&serde_json::Value::Object(Default::default())),
)
.map_err(|e| AgentError::Persistence(e.to_string()))?;
let notable_events_json = serde_json::to_string(
relationship
.get("notable_events")
.unwrap_or(&serde_json::Value::Array(vec![])),
)
.map_err(|e| AgentError::Persistence(e.to_string()))?;
let mut metadata_value = relationship
.get("metadata")
.cloned()
.unwrap_or_else(|| serde_json::Value::Object(Default::default()));
if let Some(obj) = metadata_value.as_object_mut() {
if let Some(model) = relationship.get("model") {
obj.insert("__relationship_model".to_string(), model.clone());
}
if let Some(perceived) = relationship.get("perceived_actor_to_agent") {
obj.insert("__perceived_actor_to_agent".to_string(), perceived.clone());
}
}
let metadata_json = serde_json::to_string(&metadata_value)
.map_err(|e| AgentError::Persistence(e.to_string()))?;
let interaction_count = relationship
.get("interaction_count")
.and_then(|v| v.as_u64())
.unwrap_or(0) as i64;
let first_interaction = relationship
.get("first_interaction")
.and_then(|v| v.as_str())
.unwrap_or("1970-01-01T00:00:00Z")
.to_string();
let last_interaction = relationship
.get("last_interaction")
.and_then(|v| v.as_str())
.unwrap_or("1970-01-01T00:00:00Z")
.to_string();
sqlx::query(
r#"
INSERT INTO actor_relationships
(agent_id, actor_id, actor_name, dimensions_json, notable_events_json,
interaction_count, first_interaction, last_interaction, metadata_json)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(agent_id, actor_id) DO UPDATE SET
actor_name = excluded.actor_name,
dimensions_json = excluded.dimensions_json,
notable_events_json = excluded.notable_events_json,
interaction_count = excluded.interaction_count,
first_interaction = excluded.first_interaction,
last_interaction = excluded.last_interaction,
metadata_json = excluded.metadata_json
"#,
)
.bind(agent_id)
.bind(actor_id)
.bind(actor_name)
.bind(dimensions_json)
.bind(notable_events_json)
.bind(interaction_count)
.bind(first_interaction)
.bind(last_interaction)
.bind(metadata_json)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
async fn load_relationship(
&self,
agent_id: &str,
actor_id: &str,
) -> Result<Option<serde_json::Value>> {
let row: Option<(Option<String>, String, String, i64, String, String, String)> =
sqlx::query_as(
r#"
SELECT actor_name, dimensions_json, notable_events_json, interaction_count,
first_interaction, last_interaction, metadata_json
FROM actor_relationships
WHERE agent_id = ? AND actor_id = ?
"#,
)
.bind(agent_id)
.bind(actor_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
let Some((
actor_name,
dimensions_json,
notable_events_json,
interaction_count,
first_interaction,
last_interaction,
metadata_json,
)) = row
else {
return Ok(None);
};
let dimensions: serde_json::Value = serde_json::from_str(&dimensions_json).map_err(|error| {
AgentError::Persistence(format!(
"Malformed relationship dimensions for agent '{agent_id}' actor '{actor_id}': {error}"
))
})?;
let notable_events: serde_json::Value =
serde_json::from_str(¬able_events_json).map_err(|error| {
AgentError::Persistence(format!(
"Malformed relationship events for agent '{agent_id}' actor '{actor_id}': {error}"
))
})?;
let mut metadata: serde_json::Value = serde_json::from_str(&metadata_json).map_err(|error| {
AgentError::Persistence(format!(
"Malformed relationship metadata for agent '{agent_id}' actor '{actor_id}': {error}"
))
})?;
let model = metadata
.as_object_mut()
.and_then(|obj| obj.remove("__relationship_model"))
.unwrap_or_else(|| serde_json::json!("one_sided"));
let perceived_actor_to_agent = metadata
.as_object_mut()
.and_then(|obj| obj.remove("__perceived_actor_to_agent"))
.unwrap_or_else(|| serde_json::json!({}));
Ok(Some(serde_json::json!({
"actor_id": actor_id,
"model": model,
"actor_name": actor_name,
"dimensions": dimensions,
"perceived_actor_to_agent": perceived_actor_to_agent,
"notable_events": notable_events,
"interaction_count": interaction_count,
"first_interaction": first_interaction,
"last_interaction": last_interaction,
"metadata": metadata,
})))
}
async fn list_relationship_actors(&self, agent_id: &str) -> Result<Vec<String>> {
let rows: Vec<(String,)> = sqlx::query_as(
"SELECT actor_id FROM actor_relationships WHERE agent_id = ? ORDER BY actor_id ASC",
)
.bind(agent_id)
.fetch_all(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(rows.into_iter().map(|(actor_id,)| actor_id).collect())
}
async fn delete_relationship(&self, agent_id: &str, actor_id: &str) -> Result<()> {
sqlx::query("DELETE FROM actor_relationships WHERE agent_id = ? AND actor_id = ?")
.bind(agent_id)
.bind(actor_id)
.execute(&self.pool)
.await
.map_err(|e| AgentError::Persistence(e.to_string()))?;
Ok(())
}
}
#[cfg(feature = "sqlite")]
fn parse_fact_category(s: &str) -> FactCategory {
match s {
"preference" | "user_preference" => FactCategory::UserPreference,
"context" | "user_context" => FactCategory::UserContext,
"decision" => FactCategory::Decision,
"agreement" => FactCategory::Agreement,
_ => FactCategory::Custom(s.to_string()),
}
}
#[cfg(all(test, feature = "sqlite"))]
mod tests {
use super::*;
use ai_agents_core::{ChatMessage, MemorySnapshot, Role};
async fn create_test_snapshot() -> AgentSnapshot {
let mut snapshot = AgentSnapshot::new("test-agent".into());
snapshot.memory = MemorySnapshot::new(vec![
ChatMessage {
role: Role::User,
content: "Hello".to_string(),
name: None,
timestamp: None,
},
ChatMessage {
role: Role::Assistant,
content: "Hi there!".to_string(),
name: None,
timestamp: None,
},
]);
snapshot
}
fn create_test_fact(actor_id: &str, content: &str) -> KeyFact {
KeyFact {
id: "shared-fact".to_string(),
actor_id: Some(actor_id.to_string()),
category: FactCategory::UserContext,
content: content.to_string(),
confidence: 0.9,
salience: 0.8,
extracted_at: Utc::now(),
last_accessed: None,
source_message_id: None,
source_language: None,
}
}
#[tokio::test]
async fn reports_implemented_capabilities() {
let storage = SqliteStorage::in_memory().await.unwrap();
for capability in [
StorageCapability::Snapshot,
StorageCapability::SessionMetadata,
StorageCapability::SessionFiltering,
StorageCapability::ExpiryCleanup,
StorageCapability::ActorFacts,
StorageCapability::ActorRelationships,
StorageCapability::ActorDataDeletion,
] {
assert!(storage.supports(capability));
}
}
#[tokio::test]
async fn test_sqlite_crud() {
let storage = SqliteStorage::in_memory().await.unwrap();
let snapshot = create_test_snapshot().await;
storage.save("session-1", &snapshot).await.unwrap();
let loaded = storage.load("session-1").await.unwrap();
assert!(loaded.is_some());
assert_eq!(loaded.unwrap().agent_id, "test-agent");
storage.delete("session-1").await.unwrap();
let loaded = storage.load("session-1").await.unwrap();
assert!(loaded.is_none());
}
#[tokio::test]
async fn test_sqlite_list_sessions() {
let storage = SqliteStorage::in_memory().await.unwrap();
storage
.save("session-1", &create_test_snapshot().await)
.await
.unwrap();
storage
.save("session-2", &create_test_snapshot().await)
.await
.unwrap();
let sessions = storage.list_sessions().await.unwrap();
assert_eq!(sessions.len(), 2);
}
#[tokio::test]
async fn save_metadata_synchronizes_tags_and_rejects_missing_sessions() {
let storage = SqliteStorage::in_memory().await.unwrap();
storage
.save("session-1", &create_test_snapshot().await)
.await
.unwrap();
let mut metadata = SessionMetadata {
tags: vec!["vip".to_string(), "support".to_string()],
..Default::default()
};
storage.save_metadata("session-1", &metadata).await.unwrap();
let vip_filter = SessionFilter {
tags: Some(vec!["vip".to_string()]),
..Default::default()
};
assert_eq!(
storage
.list_sessions_filtered(&vip_filter)
.await
.unwrap()
.len(),
1
);
metadata.tags = vec!["updated".to_string()];
storage.save_metadata("session-1", &metadata).await.unwrap();
assert!(
storage
.list_sessions_filtered(&vip_filter)
.await
.unwrap()
.is_empty()
);
let updated_filter = SessionFilter {
tags: Some(vec!["updated".to_string()]),
..Default::default()
};
assert_eq!(
storage
.list_sessions_filtered(&updated_filter)
.await
.unwrap()
.len(),
1
);
assert_eq!(
storage
.load_metadata("session-1")
.await
.unwrap()
.unwrap()
.tags,
vec!["updated".to_string()]
);
let error = storage
.save_metadata("missing", &SessionMetadata::default())
.await
.unwrap_err();
assert!(matches!(
error,
AgentError::Persistence(message) if message == "session not found: missing"
));
let orphan_tags: (i64,) =
sqlx::query_as("SELECT COUNT(*) FROM session_tags WHERE session_id = ?")
.bind("missing")
.fetch_one(&storage.pool)
.await
.unwrap();
assert_eq!(orphan_tags.0, 0);
}
#[tokio::test]
async fn ordinary_save_preserves_existing_metadata_and_tags() {
let storage = SqliteStorage::in_memory().await.unwrap();
let mut original = create_test_snapshot().await;
original.agent_id = "original".into();
let metadata = SessionMetadata {
tags: vec!["keep".into()],
..Default::default()
};
storage
.save_snapshot_with_metadata("session", &original, &metadata)
.await
.unwrap();
let mut updated = original.clone();
updated.agent_id = "updated".into();
storage.save("session", &updated).await.unwrap();
assert_eq!(
storage.load("session").await.unwrap().unwrap().agent_id,
"updated"
);
assert_eq!(
storage
.load_metadata("session")
.await
.unwrap()
.unwrap()
.tags,
vec!["keep"]
);
assert_eq!(
storage
.list_sessions_filtered(&SessionFilter {
tags: Some(vec!["keep".into()]),
..SessionFilter::default()
})
.await
.unwrap()
.len(),
1
);
}
#[tokio::test]
async fn snapshot_metadata_and_tags_roll_back_together() {
let storage = SqliteStorage::in_memory().await.unwrap();
let mut original = create_test_snapshot().await;
original.agent_id = "original".into();
let original_metadata = SessionMetadata {
tags: vec!["original".into()],
..Default::default()
};
storage
.save_snapshot_with_metadata("session", &original, &original_metadata)
.await
.unwrap();
sqlx::query(
r#"
CREATE TRIGGER reject_session_tag
BEFORE INSERT ON session_tags
WHEN NEW.tag = 'reject'
BEGIN
SELECT RAISE(ABORT, 'rejected tag');
END
"#,
)
.execute(&storage.pool)
.await
.unwrap();
let mut replacement = original.clone();
replacement.agent_id = "replacement".into();
let rejected_metadata = SessionMetadata {
tags: vec!["reject".into()],
..Default::default()
};
assert!(
storage
.save_snapshot_with_metadata("session", &replacement, &rejected_metadata)
.await
.is_err()
);
assert_eq!(
storage.load("session").await.unwrap().unwrap().agent_id,
"original"
);
assert_eq!(
storage
.load_metadata("session")
.await
.unwrap()
.unwrap()
.tags,
vec!["original"]
);
assert_eq!(
storage
.list_sessions_filtered(&SessionFilter {
tags: Some(vec!["original".into()]),
..SessionFilter::default()
})
.await
.unwrap()
.len(),
1
);
}
#[tokio::test]
async fn test_sqlite_with_metadata() {
let storage = SqliteStorage::in_memory().await.unwrap();
let snapshot = create_test_snapshot().await;
let metadata = SqliteMetadata {
tags: vec!["vip".to_string(), "support".to_string()],
user_id: Some("user-123".to_string()),
..Default::default()
};
storage
.save_with_metadata("session-1", &snapshot, &metadata)
.await
.unwrap();
let loaded_metadata = storage.get_metadata("session-1").await.unwrap();
assert!(loaded_metadata.is_some());
let loaded_metadata = loaded_metadata.unwrap();
assert_eq!(loaded_metadata.tags.len(), 2);
assert_eq!(loaded_metadata.user_id, Some("user-123".to_string()));
}
#[tokio::test]
async fn test_sqlite_list_by_agent() {
let storage = SqliteStorage::in_memory().await.unwrap();
let mut snapshot1 = create_test_snapshot().await;
snapshot1.agent_id = "agent-A".to_string();
let mut snapshot2 = create_test_snapshot().await;
snapshot2.agent_id = "agent-B".to_string();
storage.save("session-1", &snapshot1).await.unwrap();
storage.save("session-2", &snapshot2).await.unwrap();
storage.save("session-3", &snapshot1).await.unwrap();
let sessions = storage.list_sessions_by_agent("agent-A").await.unwrap();
assert_eq!(sessions.len(), 2);
}
#[tokio::test]
async fn test_sqlite_search() {
let storage = SqliteStorage::in_memory().await.unwrap();
let snapshot = create_test_snapshot().await;
let metadata = SqliteMetadata {
tags: vec!["vip".to_string()],
user_id: Some("user-123".to_string()),
..Default::default()
};
storage
.save_with_metadata("session-1", &snapshot, &metadata)
.await
.unwrap();
let query = SessionQuery {
tag: Some("vip".to_string()),
..Default::default()
};
let results = storage.search_sessions(&query).await.unwrap();
assert_eq!(results.len(), 1);
}
#[tokio::test]
async fn test_sqlite_exists() {
let storage = SqliteStorage::in_memory().await.unwrap();
assert!(!storage.exists("session-1").await.unwrap());
storage
.save("session-1", &create_test_snapshot().await)
.await
.unwrap();
assert!(storage.exists("session-1").await.unwrap());
}
#[tokio::test]
async fn test_sqlite_expire() {
let storage = SqliteStorage::in_memory().await.unwrap();
storage
.save("session-1", &create_test_snapshot().await)
.await
.unwrap();
let future = Utc::now() + chrono::Duration::hours(1);
let expired = storage.expire_sessions(future).await.unwrap();
assert_eq!(expired, 1);
let sessions = storage.list_sessions().await.unwrap();
assert!(sessions.is_empty());
}
#[tokio::test]
async fn save_facts_rolls_back_the_batch_on_failure() {
let storage = SqliteStorage::in_memory().await.unwrap();
sqlx::query(
r#"
CREATE TRIGGER reject_fact_insert
BEFORE INSERT ON actor_facts
WHEN NEW.content = 'reject'
BEGIN
SELECT RAISE(ABORT, 'rejected fact');
END
"#,
)
.execute(&storage.pool)
.await
.unwrap();
let mut accepted = create_test_fact("actor", "accepted");
accepted.id = "accepted".to_string();
let mut rejected = create_test_fact("actor", "reject");
rejected.id = "rejected".to_string();
assert!(
storage
.save_facts("agent", "actor", &[accepted, rejected])
.await
.is_err()
);
assert!(
storage
.load_facts("agent", "actor")
.await
.unwrap()
.is_empty()
);
}
#[tokio::test]
async fn delete_actor_data_rolls_back_all_tables_on_failure() {
let storage = SqliteStorage::in_memory().await.unwrap();
let mut snapshot = create_test_snapshot().await;
snapshot.agent_id = "agent".to_string();
storage.save("session", &snapshot).await.unwrap();
storage
.save_metadata(
"session",
&SessionMetadata {
actor_id: Some("actor".to_string()),
..Default::default()
},
)
.await
.unwrap();
storage
.save_facts("agent", "actor", &[create_test_fact("actor", "fact")])
.await
.unwrap();
storage
.save_relationship(
"agent",
"actor",
&serde_json::json!({ "actor_name": "Actor" }),
)
.await
.unwrap();
sqlx::query(
r#"
CREATE TRIGGER reject_relationship_delete
BEFORE DELETE ON actor_relationships
WHEN OLD.agent_id = 'agent' AND OLD.actor_id = 'actor'
BEGIN
SELECT RAISE(ABORT, 'rejected relationship deletion');
END
"#,
)
.execute(&storage.pool)
.await
.unwrap();
assert!(storage.delete_actor_data("agent", "actor").await.is_err());
assert_eq!(storage.load_facts("agent", "actor").await.unwrap().len(), 1);
assert!(
storage
.load_relationship("agent", "actor")
.await
.unwrap()
.is_some()
);
assert!(storage.load("session").await.unwrap().is_some());
}
#[tokio::test]
async fn malformed_relationship_data_fails_closed() {
let storage = SqliteStorage::in_memory().await.unwrap();
storage
.save_relationship(
"agent",
"actor",
&serde_json::json!({ "actor_name": "Actor" }),
)
.await
.unwrap();
sqlx::query(
"UPDATE actor_relationships SET dimensions_json = 'not-json' WHERE agent_id = ? AND actor_id = ?",
)
.bind("agent")
.bind("actor")
.execute(&storage.pool)
.await
.unwrap();
assert!(matches!(
storage.load_relationship("agent", "actor").await,
Err(AgentError::Persistence(message)) if message.contains("Malformed relationship dimensions")
));
}
#[tokio::test]
async fn actor_data_remains_isolated_after_reopen() {
let temp_dir = tempfile::TempDir::new().unwrap();
let path = temp_dir.path().join("actor-data.sqlite");
let path = path.to_string_lossy().into_owned();
let records = [
("agent-a", "actor-1", "agent-a actor-1"),
("agent-a", "actor-2", "agent-a actor-2"),
("agent-b", "actor-1", "agent-b actor-1"),
];
let storage = SqliteStorage::new(&path).await.unwrap();
for &(agent_id, actor_id, content) in &records {
storage
.save_facts(agent_id, actor_id, &[create_test_fact(actor_id, content)])
.await
.unwrap();
storage
.save_relationship(
agent_id,
actor_id,
&serde_json::json!({ "actor_name": content }),
)
.await
.unwrap();
}
storage.close().await;
let storage = SqliteStorage::new(&path).await.unwrap();
for &(agent_id, actor_id, content) in &records {
let facts = storage.load_facts(agent_id, actor_id).await.unwrap();
assert_eq!(facts.len(), 1);
assert_eq!(facts[0].content, content);
assert_eq!(facts[0].actor_id.as_deref(), Some(actor_id));
let relationship = storage
.load_relationship(agent_id, actor_id)
.await
.unwrap()
.unwrap();
assert_eq!(relationship["actor_name"], content);
}
assert!(
storage
.load_facts("agent-b", "actor-2")
.await
.unwrap()
.is_empty()
);
assert!(
storage
.load_relationship("agent-b", "actor-2")
.await
.unwrap()
.is_none()
);
assert_eq!(
storage.list_relationship_actors("agent-a").await.unwrap(),
vec!["actor-1".to_string(), "actor-2".to_string()]
);
}
#[tokio::test]
async fn test_sqlite_get_session_info() {
let storage = SqliteStorage::in_memory().await.unwrap();
let snapshot = create_test_snapshot().await;
storage.save("session-1", &snapshot).await.unwrap();
let info = storage.get_session_info("session-1").await.unwrap();
assert!(info.is_some());
let info = info.unwrap();
assert_eq!(info.session_id, "session-1");
assert_eq!(info.agent_id, "test-agent");
assert_eq!(info.message_count, 2);
let not_found = storage.get_session_info("nonexistent").await.unwrap();
assert!(not_found.is_none());
}
}