use std::sync::Arc;
use parking_lot::Mutex;
pub use zkr::{
Error as SelfImproveError, MemoryDb, PersonId, SelfImprove as ZkrSelfImprove, TenantId,
};
#[derive(Clone)]
pub struct SelfImprove {
inner: Arc<Mutex<ZkrSelfImprove>>,
}
impl SelfImprove {
pub fn new(db: MemoryDb, tenant_id: TenantId, person_id: PersonId) -> Self {
Self {
inner: Arc::new(Mutex::new(ZkrSelfImprove::new(db, tenant_id, person_id))),
}
}
pub async fn record(
&self,
context: &str,
action: &str,
outcome: &str,
lesson: &str,
) -> Result<(), SelfImproveError> {
let inner = Arc::clone(&self.inner);
let context = context.to_string();
let action = action.to_string();
let outcome = outcome.to_string();
let lesson = lesson.to_string();
tokio::task::spawn_blocking(move || {
let mut guard = inner.lock();
guard
.record(&context, &action, &outcome, &lesson)
.map(|_| ())
})
.await
.map_err(|e| SelfImproveError::Invalid(e.to_string()))?
}
pub async fn augment(&self, query: &str, base: &str) -> Result<String, SelfImproveError> {
let inner = Arc::clone(&self.inner);
let query = query.to_string();
let base = base.to_string();
tokio::task::spawn_blocking(move || inner.lock().augment(&query, &base))
.await
.map_err(|e| SelfImproveError::Invalid(e.to_string()))?
}
pub async fn lessons(&self, query: &str, limit: u32) -> Result<Vec<String>, SelfImproveError> {
let inner = Arc::clone(&self.inner);
let query = query.to_string();
tokio::task::spawn_blocking(move || inner.lock().lessons(&query, limit))
.await
.map_err(|e| SelfImproveError::Invalid(e.to_string()))?
}
}
#[cfg(test)]
mod tests {
use super::*;
use uuid::Uuid;
use zkr::MemoryDb;
#[tokio::test]
async fn test_self_improve_record_lessons_and_augment() -> Result<(), Box<dyn std::error::Error>>
{
let tmp = tempfile::tempdir()?;
let db = MemoryDb::open(tmp.path().join("self_improve.db"))?;
let tenant_id = TenantId::new(Uuid::new_v4())?;
let person_id = PersonId::new(Uuid::new_v4())?;
let self_improve = SelfImprove::new(db, tenant_id, person_id);
self_improve
.record("context 1", "action 1", "outcome 1", "lesson 1")
.await?;
let lessons = self_improve.lessons("context", 10).await?;
assert_eq!(lessons.len(), 1);
assert!(lessons[0].contains("lesson 1"));
let augmented = self_improve.augment("context", "base prompt").await?;
assert!(augmented.contains("base prompt"));
assert!(augmented.contains("lesson 1"));
Ok(())
}
}