use async_trait::async_trait;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::{DatabaseConfig, DatabaseError, MemoryPool, MemoryPoolManager, PoolStats};
pub struct SqliteMemoryPools {
pools: Vec<Arc<SqliteMemoryPool>>,
config: DatabaseConfig,
}
impl SqliteMemoryPools {
pub async fn new(config: DatabaseConfig) -> Result<Self, DatabaseError> {
tokio::fs::create_dir_all(&config.sqlite.pool_directory).await?;
let mut pools = Vec::with_capacity(config.memory_pools.pool_count);
for i in 0..config.memory_pools.pool_count {
let pool = SqliteMemoryPool::new(
config.sqlite.pool_directory.join(format!("pool_{}.db", i)),
config.memory_pools.max_pool_size_mb,
).await?;
pools.push(Arc::new(pool));
}
Ok(Self { pools, config })
}
}
#[async_trait]
impl MemoryPoolManager for SqliteMemoryPools {
async fn get_pool(&self, index: usize) -> Result<Box<dyn MemoryPool>, DatabaseError> {
if index >= self.pools.len() {
return Err(DatabaseError::InvalidPoolIndex(index));
}
Ok(Box::new(self.pools[index].clone()))
}
async fn get_pool_for_key(&self, key: &str) -> Result<Box<dyn MemoryPool>, DatabaseError> {
let hash = seahash::hash(key.as_bytes());
let index = (hash as usize) % self.pools.len();
self.get_pool(index).await
}
async fn all_stats(&self) -> Result<Vec<PoolStats>, DatabaseError> {
let mut stats = Vec::new();
for pool in &self.pools {
stats.push(pool.stats().await?);
}
Ok(stats)
}
async fn rebalance(&self) -> Result<(), DatabaseError> {
Ok(())
}
}
#[derive(Clone)]
pub struct SqliteMemoryPool {
pool: Arc<sqlx::SqlitePool>,
stats: Arc<RwLock<PoolStats>>,
max_size_bytes: usize,
}
impl SqliteMemoryPool {
async fn new(path: PathBuf, max_size_mb: usize) -> Result<Self, DatabaseError> {
let pool = sqlx::SqlitePool::connect(&format!("sqlite:{}", path.display())).await?;
sqlx::query(
"CREATE TABLE IF NOT EXISTS memory_pool (
key TEXT PRIMARY KEY,
value BLOB NOT NULL,
created_at INTEGER NOT NULL,
accessed_at INTEGER NOT NULL,
access_count INTEGER DEFAULT 1
)"
)
.execute(&pool)
.await?;
sqlx::query("CREATE INDEX IF NOT EXISTS idx_accessed_at ON memory_pool(accessed_at)")
.execute(&pool)
.await?;
let stats = PoolStats {
size: 0,
items: 0,
hits: 0,
misses: 0,
evictions: 0,
};
Ok(Self {
pool: Arc::new(pool),
stats: Arc::new(RwLock::new(stats)),
max_size_bytes: max_size_mb * 1024 * 1024,
})
}
}
#[async_trait]
impl MemoryPool for SqliteMemoryPool {
async fn store(&self, key: &str, value: &[u8]) -> Result<(), DatabaseError> {
let now = chrono::Utc::now().timestamp();
let current_size = self.get_total_size().await?;
if current_size + value.len() > self.max_size_bytes {
self.evict_lru().await?;
}
sqlx::query(
"INSERT OR REPLACE INTO memory_pool (key, value, created_at, accessed_at)
VALUES (?, ?, ?, ?)"
)
.bind(key)
.bind(value)
.bind(now)
.bind(now)
.execute(&*self.pool)
.await?;
let mut stats = self.stats.write().await;
stats.items += 1;
stats.size += value.len();
Ok(())
}
async fn retrieve(&self, key: &str) -> Result<Option<Vec<u8>>, DatabaseError> {
let now = chrono::Utc::now().timestamp();
let result: Option<(Vec<u8>,)> = sqlx::query_as(
"UPDATE memory_pool
SET accessed_at = ?, access_count = access_count + 1
WHERE key = ?
RETURNING value"
)
.bind(now)
.bind(key)
.fetch_optional(&*self.pool)
.await?;
let mut stats = self.stats.write().await;
if result.is_some() {
stats.hits += 1;
} else {
stats.misses += 1;
}
Ok(result.map(|r| r.0))
}
async fn delete(&self, key: &str) -> Result<(), DatabaseError> {
let result = sqlx::query("DELETE FROM memory_pool WHERE key = ?")
.bind(key)
.execute(&*self.pool)
.await?;
if result.rows_affected() > 0 {
let mut stats = self.stats.write().await;
stats.items -= 1;
}
Ok(())
}
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, DatabaseError> {
let pattern = format!("{}%", prefix);
let keys: Vec<(String,)> = sqlx::query_as(
"SELECT key FROM memory_pool WHERE key LIKE ? ORDER BY key"
)
.bind(pattern)
.fetch_all(&*self.pool)
.await?;
Ok(keys.into_iter().map(|k| k.0).collect())
}
async fn stats(&self) -> Result<PoolStats, DatabaseError> {
Ok(self.stats.read().await.clone())
}
}
impl SqliteMemoryPool {
async fn get_total_size(&self) -> Result<usize, DatabaseError> {
let size: (i64,) = sqlx::query_as(
"SELECT COALESCE(SUM(LENGTH(value)), 0) FROM memory_pool"
)
.fetch_one(&*self.pool)
.await?;
Ok(size.0 as usize)
}
async fn evict_lru(&self) -> Result<(), DatabaseError> {
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM memory_pool")
.fetch_one(&*self.pool)
.await?;
let to_evict = (count.0 as f64 * 0.1).ceil() as i64;
sqlx::query(
"DELETE FROM memory_pool
WHERE key IN (
SELECT key FROM memory_pool
ORDER BY accessed_at ASC
LIMIT ?
)"
)
.bind(to_evict)
.execute(&*self.pool)
.await?;
let mut stats = self.stats.write().await;
stats.evictions += to_evict as u64;
stats.items -= to_evict as usize;
Ok(())
}
}
#[cfg(feature = "astra")]
pub struct AstraMemoryPools {
pools: Vec<Arc<AstraMemoryPool>>,
config: DatabaseConfig,
}
#[cfg(feature = "astra")]
impl AstraMemoryPools {
pub async fn new(config: DatabaseConfig) -> Result<Self, DatabaseError> {
Err(DatabaseError::NotImplemented)
}
}
#[cfg(feature = "astra")]
#[async_trait]
impl MemoryPoolManager for AstraMemoryPools {
async fn get_pool(&self, index: usize) -> Result<Box<dyn MemoryPool>, DatabaseError> {
if index >= self.pools.len() {
return Err(DatabaseError::InvalidPoolIndex(index));
}
Ok(Box::new(self.pools[index].clone()))
}
async fn get_pool_for_key(&self, key: &str) -> Result<Box<dyn MemoryPool>, DatabaseError> {
let hash = seahash::hash(key.as_bytes());
let index = (hash as usize) % self.pools.len();
self.get_pool(index).await
}
async fn all_stats(&self) -> Result<Vec<PoolStats>, DatabaseError> {
let mut stats = Vec::new();
for pool in &self.pools {
stats.push(pool.stats().await?);
}
Ok(stats)
}
async fn rebalance(&self) -> Result<(), DatabaseError> {
Ok(())
}
}
#[cfg(feature = "astra")]
#[derive(Clone)]
pub struct AstraMemoryPool {
collection_name: String,
stats: Arc<RwLock<PoolStats>>,
}
#[cfg(feature = "astra")]
#[async_trait]
impl MemoryPool for AstraMemoryPool {
async fn store(&self, _key: &str, _value: &[u8]) -> Result<(), DatabaseError> {
Err(DatabaseError::NotImplemented)
}
async fn retrieve(&self, _key: &str) -> Result<Option<Vec<u8>>, DatabaseError> {
Err(DatabaseError::NotImplemented)
}
async fn delete(&self, _key: &str) -> Result<(), DatabaseError> {
Err(DatabaseError::NotImplemented)
}
async fn list_keys(&self, _prefix: &str) -> Result<Vec<String>, DatabaseError> {
Err(DatabaseError::NotImplemented)
}
async fn stats(&self) -> Result<PoolStats, DatabaseError> {
Ok(self.stats.read().await.clone())
}
}
#[async_trait]
impl MemoryPool for Arc<SqliteMemoryPool> {
async fn store(&self, key: &str, value: &[u8]) -> Result<(), DatabaseError> {
self.as_ref().store(key, value).await
}
async fn retrieve(&self, key: &str) -> Result<Option<Vec<u8>>, DatabaseError> {
self.as_ref().retrieve(key).await
}
async fn delete(&self, key: &str) -> Result<(), DatabaseError> {
self.as_ref().delete(key).await
}
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, DatabaseError> {
self.as_ref().list_keys(prefix).await
}
async fn stats(&self) -> Result<PoolStats, DatabaseError> {
self.as_ref().stats().await
}
}
mod seahash {
pub fn hash(data: &[u8]) -> u64 {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
hasher.write(data);
hasher.finish()
}
use std::hash::Hasher;
}