use async_trait::async_trait;
use paladin_core::platform::container::garrison::{
ConversationRole, GarrisonConfig, GarrisonEntry,
};
use paladin_ports::output::garrison_port::{GarrisonError, GarrisonPort, GarrisonStats};
use sqlx::Row;
use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions};
use std::path::Path;
use std::str::FromStr;
#[derive(Debug, Clone)]
pub struct SqliteGarrison {
pool: SqlitePool,
config: GarrisonConfig,
paladin_id: String,
}
impl SqliteGarrison {
pub async fn connect(
path: impl AsRef<Path>,
config: GarrisonConfig,
paladin_id: impl Into<String>,
) -> Result<Self, GarrisonError> {
let path_str = path
.as_ref()
.to_str()
.ok_or_else(|| GarrisonError::ConfigurationError("Invalid database path".into()))?;
let options = SqliteConnectOptions::from_str(&format!("sqlite://{}", path_str))
.map_err(|e| GarrisonError::StorageError(format!("Connection options error: {}", e)))?
.create_if_missing(true)
.journal_mode(sqlx::sqlite::SqliteJournalMode::Wal)
.busy_timeout(std::time::Duration::from_secs(30));
let pool = SqlitePoolOptions::new()
.max_connections(5)
.connect_with(options)
.await
.map_err(|e| GarrisonError::StorageError(format!("Connection failed: {}", e)))?;
let garrison = Self {
pool,
config,
paladin_id: paladin_id.into(),
};
garrison.initialize().await?;
Ok(garrison)
}
async fn initialize(&self) -> Result<(), GarrisonError> {
sqlx::migrate::Migrator::new(std::path::Path::new("./migrations"))
.await
.map_err(|e| GarrisonError::StorageError(format!("Migration setup failed: {}", e)))?
.run(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Migration failed: {}", e)))?;
sqlx::query(
r#"
INSERT OR IGNORE INTO garrison_metadata
(paladin_id, max_entries, max_tokens, eviction_strategy, preserve_recent_count)
VALUES (?, ?, ?, ?, ?)
"#,
)
.bind(&self.paladin_id)
.bind(self.config.max_entries as i64)
.bind(self.config.max_tokens.map(|t| t as i64))
.bind(format!("{:?}", self.config.eviction_strategy))
.bind(self.config.preserve_recent_count as i64)
.execute(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Metadata init failed: {}", e)))?;
Ok(())
}
async fn apply_eviction(&self) -> Result<(), GarrisonError> {
let stats = self.stats().await?;
let needs_eviction = if stats.entry_count > self.config.max_entries {
true
} else if let Some(max_tokens) = self.config.max_tokens {
stats.total_tokens > max_tokens
} else {
false
};
if !needs_eviction {
return Ok(());
}
let target_count =
std::cmp::min(self.config.max_entries, self.config.preserve_recent_count);
match self.config.eviction_strategy {
paladin_core::platform::container::garrison::EvictionStrategy::FIFO => {
sqlx::query(
r#"
DELETE FROM garrison_entries
WHERE paladin_id = ?
AND id NOT IN (
SELECT id FROM garrison_entries
WHERE paladin_id = ?
ORDER BY timestamp DESC
LIMIT ?
)
"#,
)
.bind(&self.paladin_id)
.bind(&self.paladin_id)
.bind(target_count as i64)
.execute(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("FIFO eviction failed: {}", e)))?;
}
paladin_core::platform::container::garrison::EvictionStrategy::ImportanceBased => {
sqlx::query(
r#"
DELETE FROM garrison_entries
WHERE paladin_id = ?
AND role != 'system'
AND id NOT IN (
SELECT id FROM garrison_entries
WHERE paladin_id = ?
ORDER BY timestamp DESC
LIMIT ?
)
"#,
)
.bind(&self.paladin_id)
.bind(&self.paladin_id)
.bind(target_count as i64)
.execute(&self.pool)
.await
.map_err(|e| {
GarrisonError::StorageError(format!("Importance eviction failed: {}", e))
})?;
}
paladin_core::platform::container::garrison::EvictionStrategy::SlidingWindow => {
sqlx::query(
r#"
DELETE FROM garrison_entries
WHERE paladin_id = ?
AND id NOT IN (
SELECT id FROM garrison_entries
WHERE paladin_id = ?
ORDER BY timestamp DESC
LIMIT ?
)
"#,
)
.bind(&self.paladin_id)
.bind(&self.paladin_id)
.bind(target_count as i64)
.execute(&self.pool)
.await
.map_err(|e| {
GarrisonError::StorageError(format!("Sliding window eviction failed: {}", e))
})?;
}
}
self.update_metadata().await?;
Ok(())
}
async fn update_metadata(&self) -> Result<(), GarrisonError> {
let row = sqlx::query(
r#"
SELECT
COUNT(*) as entry_count,
COALESCE(SUM(token_count), 0) as total_tokens
FROM garrison_entries
WHERE paladin_id = ?
"#,
)
.bind(&self.paladin_id)
.fetch_one(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Stats query failed: {}", e)))?;
let entry_count: i64 = row
.try_get("entry_count")
.map_err(|e| GarrisonError::SerializationError(format!("Entry count error: {}", e)))?;
let total_tokens: i64 = row
.try_get("total_tokens")
.map_err(|e| GarrisonError::SerializationError(format!("Total tokens error: {}", e)))?;
sqlx::query(
r#"
UPDATE garrison_metadata
SET total_entries = ?,
total_tokens = ?,
last_eviction = datetime('now'),
updated_at = datetime('now')
WHERE paladin_id = ?
"#,
)
.bind(entry_count)
.bind(total_tokens)
.bind(&self.paladin_id)
.execute(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Metadata update failed: {}", e)))?;
Ok(())
}
}
#[async_trait]
impl GarrisonPort for SqliteGarrison {
async fn remember(&self, entry: GarrisonEntry) -> Result<(), GarrisonError> {
entry
.validate()
.map_err(GarrisonError::ConfigurationError)?;
sqlx::query(
r#"
INSERT INTO garrison_entries
(id, paladin_id, role, content, timestamp, token_count, metadata, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, datetime('now'), datetime('now'))
"#,
)
.bind(entry.id.to_string())
.bind(&self.paladin_id)
.bind(format!("{:?}", entry.role).to_lowercase())
.bind(&entry.content)
.bind(entry.timestamp.to_rfc3339())
.bind(entry.token_count.map(|t| t as i64))
.bind(serde_json::to_string(&entry.metadata).ok())
.execute(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Insert failed: {}", e)))?;
self.apply_eviction().await?;
Ok(())
}
async fn recall_recent(&self, limit: usize) -> Result<Vec<GarrisonEntry>, GarrisonError> {
let rows = sqlx::query(
r#"
SELECT id, role, content, timestamp, token_count, metadata
FROM garrison_entries
WHERE paladin_id = ?
ORDER BY timestamp DESC
LIMIT ?
"#,
)
.bind(&self.paladin_id)
.bind(limit as i64)
.fetch_all(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Recall failed: {}", e)))?;
let mut entries = Vec::new();
for row in rows {
let role_str: String = row.try_get("role").map_err(|e| {
GarrisonError::SerializationError(format!("Role parse error: {}", e))
})?;
let role = match role_str.as_str() {
"system" => ConversationRole::System,
"user" => ConversationRole::User,
"assistant" => ConversationRole::Assistant,
"tool" => ConversationRole::Tool,
_ => ConversationRole::User,
};
let id: String = row
.try_get("id")
.map_err(|e| GarrisonError::SerializationError(format!("ID parse error: {}", e)))?;
let content: String = row.try_get("content").map_err(|e| {
GarrisonError::SerializationError(format!("Content parse error: {}", e))
})?;
let timestamp_str: String = row.try_get("timestamp").map_err(|e| {
GarrisonError::SerializationError(format!("Timestamp parse error: {}", e))
})?;
let timestamp = chrono::DateTime::parse_from_rfc3339(×tamp_str)
.map_err(|e| {
GarrisonError::SerializationError(format!("Timestamp conversion error: {}", e))
})?
.with_timezone(&chrono::Utc);
let token_count: Option<i64> = row.try_get("token_count").ok();
let metadata_str: Option<String> = row.try_get("metadata").ok();
let metadata = metadata_str
.and_then(|s| serde_json::from_str(&s).ok())
.unwrap_or_default();
let mut entry = GarrisonEntry::new(role, content);
entry.id = uuid::Uuid::parse_str(&id)
.map_err(|e| GarrisonError::SerializationError(format!("UUID parse: {}", e)))?;
entry.timestamp = timestamp;
entry.token_count = token_count.map(|t| t as u32);
entry.metadata = metadata;
entries.push(entry);
}
entries.reverse();
Ok(entries)
}
async fn search(&self, query: &str, limit: usize) -> Result<Vec<GarrisonEntry>, GarrisonError> {
if query.is_empty() {
return Ok(Vec::new());
}
let rows = sqlx::query(
r#"
SELECT e.id, e.role, e.content, e.timestamp, e.token_count, e.metadata
FROM garrison_entries e
JOIN garrison_search s ON e.rowid = s.rowid
WHERE e.paladin_id = ? AND garrison_search MATCH ?
ORDER BY e.timestamp DESC
LIMIT ?
"#,
)
.bind(&self.paladin_id)
.bind(query)
.bind(limit as i64)
.fetch_all(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Search failed: {}", e)))?;
let mut entries = Vec::new();
for row in rows {
let role_str: String = row.try_get("role").map_err(|e| {
GarrisonError::SerializationError(format!("Role parse error: {}", e))
})?;
let role = match role_str.as_str() {
"system" => ConversationRole::System,
"user" => ConversationRole::User,
"assistant" => ConversationRole::Assistant,
"tool" => ConversationRole::Tool,
_ => ConversationRole::User,
};
let id: String = row
.try_get("id")
.map_err(|e| GarrisonError::SerializationError(format!("ID parse error: {}", e)))?;
let content: String = row.try_get("content").map_err(|e| {
GarrisonError::SerializationError(format!("Content parse error: {}", e))
})?;
let timestamp_str: String = row.try_get("timestamp").map_err(|e| {
GarrisonError::SerializationError(format!("Timestamp parse error: {}", e))
})?;
let timestamp = chrono::DateTime::parse_from_rfc3339(×tamp_str)
.map_err(|e| {
GarrisonError::SerializationError(format!("Timestamp conversion error: {}", e))
})?
.with_timezone(&chrono::Utc);
let token_count: Option<i64> = row.try_get("token_count").ok();
let metadata_str: Option<String> = row.try_get("metadata").ok();
let metadata = metadata_str
.and_then(|s| serde_json::from_str(&s).ok())
.unwrap_or_default();
let mut entry = GarrisonEntry::new(role, content);
entry.id = uuid::Uuid::parse_str(&id)
.map_err(|e| GarrisonError::SerializationError(format!("UUID parse: {}", e)))?;
entry.timestamp = timestamp;
entry.token_count = token_count.map(|t| t as u32);
entry.metadata = metadata;
entries.push(entry);
}
Ok(entries)
}
async fn forget_all(&self) -> Result<(), GarrisonError> {
sqlx::query("DELETE FROM garrison_entries WHERE paladin_id = ?")
.bind(&self.paladin_id)
.execute(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Delete all failed: {}", e)))?;
self.update_metadata().await?;
Ok(())
}
async fn stats(&self) -> Result<GarrisonStats, GarrisonError> {
let row = sqlx::query(
r#"
SELECT
COUNT(*) as entry_count,
COALESCE(SUM(token_count), 0) as total_tokens
FROM garrison_entries
WHERE paladin_id = ?
"#,
)
.bind(&self.paladin_id)
.fetch_one(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Stats query failed: {}", e)))?;
let entry_count: i64 = row
.try_get("entry_count")
.map_err(|e| GarrisonError::SerializationError(format!("Entry count error: {}", e)))?;
let total_tokens: i64 = row
.try_get("total_tokens")
.map_err(|e| GarrisonError::SerializationError(format!("Total tokens error: {}", e)))?;
let size_row = sqlx::query(
"SELECT page_count * page_size as size FROM pragma_page_count(), pragma_page_size()",
)
.fetch_optional(&self.pool)
.await
.map_err(|e| GarrisonError::StorageError(format!("Size query failed: {}", e)))?;
let size_bytes = size_row
.and_then(|r| r.try_get::<i64, _>("size").ok())
.map(|s| s as u64);
Ok(GarrisonStats {
entry_count: entry_count as usize,
total_tokens: total_tokens as u32,
size_bytes: size_bytes.map(|s| s as usize),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[tokio::test]
async fn test_sqlite_garrison_creation() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::default();
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
let stats = garrison.stats().await.unwrap();
assert_eq!(stats.entry_count, 0);
assert_eq!(stats.total_tokens, 0);
}
#[tokio::test]
async fn test_sqlite_remember_and_recall() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::default();
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
let entry = GarrisonEntry::new(ConversationRole::User, "Test message".to_string());
garrison.remember(entry).await.unwrap();
let entries = garrison.recall_recent(10).await.unwrap();
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].content, "Test message");
}
#[tokio::test]
async fn test_sqlite_persistence() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::default();
{
let garrison =
SqliteGarrison::connect(temp_file.path(), config.clone(), "test-paladin")
.await
.unwrap();
let entry =
GarrisonEntry::new(ConversationRole::User, "Persistent message".to_string());
garrison.remember(entry).await.unwrap();
}
{
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
let entries = garrison.recall_recent(10).await.unwrap();
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].content, "Persistent message");
}
}
#[tokio::test]
async fn test_sqlite_search() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::default();
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
garrison
.remember(GarrisonEntry::new(
ConversationRole::User,
"Hello world".to_string(),
))
.await
.unwrap();
garrison
.remember(GarrisonEntry::new(
ConversationRole::User,
"Goodbye world".to_string(),
))
.await
.unwrap();
garrison
.remember(GarrisonEntry::new(
ConversationRole::User,
"Random message".to_string(),
))
.await
.unwrap();
let results = garrison.search("world", 10).await.unwrap();
assert_eq!(results.len(), 2);
}
#[tokio::test]
async fn test_sqlite_eviction() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::new(3, None);
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
for i in 0..5 {
garrison
.remember(GarrisonEntry::new(
ConversationRole::User,
format!("Message {}", i),
))
.await
.unwrap();
}
let stats = garrison.stats().await.unwrap();
assert!(stats.entry_count <= 3);
}
#[tokio::test]
async fn test_sqlite_forget_all() {
let temp_file = NamedTempFile::new().unwrap();
let config = GarrisonConfig::default();
let garrison = SqliteGarrison::connect(temp_file.path(), config, "test-paladin")
.await
.unwrap();
garrison
.remember(GarrisonEntry::new(
ConversationRole::User,
"Test".to_string(),
))
.await
.unwrap();
garrison.forget_all().await.unwrap();
let stats = garrison.stats().await.unwrap();
assert_eq!(stats.entry_count, 0);
}
}