use crate::error::RragResult;
use crate::storage::{Memory, MemoryValue};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct KnowledgeEntry {
pub id: String,
pub key: String,
pub value: MemoryValue,
pub created_by: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub updated_by: String,
pub updated_at: chrono::DateTime<chrono::Utc>,
pub tags: Vec<String>,
pub acl: Option<Vec<String>>,
pub metadata: std::collections::HashMap<String, String>,
}
impl KnowledgeEntry {
pub fn new(
key: impl Into<String>,
value: impl Into<MemoryValue>,
created_by: impl Into<String>,
) -> Self {
let now = chrono::Utc::now();
let created_by = created_by.into();
Self {
id: uuid::Uuid::new_v4().to_string(),
key: key.into(),
value: value.into(),
created_by: created_by.clone(),
created_at: now,
updated_by: created_by,
updated_at: now,
tags: Vec::new(),
acl: None,
metadata: std::collections::HashMap::new(),
}
}
pub fn with_tags(mut self, tags: Vec<String>) -> Self {
self.tags = tags;
self
}
pub fn with_acl(mut self, acl: Vec<String>) -> Self {
self.acl = Some(acl);
self
}
pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.metadata.insert(key.into(), value.into());
self
}
pub fn has_access(&self, agent_id: &str) -> bool {
match &self.acl {
None => true, Some(acl) => acl.contains(&agent_id.to_string()) || agent_id == self.created_by,
}
}
}
pub struct SharedKnowledgeBase {
storage: Arc<dyn Memory>,
agent_id: String,
namespace: String,
}
impl SharedKnowledgeBase {
pub fn new(storage: Arc<dyn Memory>, agent_id: String) -> Self {
Self {
storage,
agent_id,
namespace: "global::knowledge".to_string(),
}
}
pub async fn store(
&self,
key: impl Into<String>,
value: impl Into<MemoryValue>,
) -> RragResult<KnowledgeEntry> {
let entry = KnowledgeEntry::new(key, value, self.agent_id.clone());
self.store_entry(entry.clone()).await?;
Ok(entry)
}
pub async fn store_with_tags(
&self,
key: impl Into<String>,
value: impl Into<MemoryValue>,
tags: Vec<String>,
) -> RragResult<KnowledgeEntry> {
let entry = KnowledgeEntry::new(key, value, self.agent_id.clone()).with_tags(tags);
self.store_entry(entry.clone()).await?;
Ok(entry)
}
pub async fn store_entry(&self, mut entry: KnowledgeEntry) -> RragResult<()> {
entry.updated_by = self.agent_id.clone();
entry.updated_at = chrono::Utc::now();
let storage_key = self.entry_key(&entry.key);
let value = serde_json::to_value(&entry).map_err(|e| {
crate::error::RragError::storage(
"serialize_entry",
std::io::Error::new(std::io::ErrorKind::Other, e),
)
})?;
self.storage
.set(&storage_key, MemoryValue::Json(value))
.await
}
pub async fn get(&self, key: &str) -> RragResult<Option<KnowledgeEntry>> {
let storage_key = self.entry_key(key);
if let Some(value) = self.storage.get(&storage_key).await? {
if let Some(json) = value.as_json() {
let entry: KnowledgeEntry = serde_json::from_value(json.clone()).map_err(|e| {
crate::error::RragError::storage(
"deserialize_entry",
std::io::Error::new(std::io::ErrorKind::Other, e),
)
})?;
if entry.has_access(&self.agent_id) {
return Ok(Some(entry));
}
}
}
Ok(None)
}
pub async fn get_value(&self, key: &str) -> RragResult<Option<MemoryValue>> {
if let Some(entry) = self.get(key).await? {
Ok(Some(entry.value))
} else {
Ok(None)
}
}
pub async fn delete(&self, key: &str) -> RragResult<bool> {
if let Some(entry) = self.get(key).await? {
if entry.created_by != self.agent_id {
return Ok(false);
}
}
let storage_key = self.entry_key(key);
self.storage.delete(&storage_key).await
}
pub async fn exists(&self, key: &str) -> RragResult<bool> {
Ok(self.get(key).await?.is_some())
}
pub async fn find_by_tag(&self, tag: &str) -> RragResult<Vec<KnowledgeEntry>> {
let all_entries = self.get_all_entries().await?;
let matching = all_entries
.into_iter()
.filter(|e| e.has_access(&self.agent_id) && e.tags.contains(&tag.to_string()))
.collect();
Ok(matching)
}
pub async fn find_by_creator(&self, creator_agent_id: &str) -> RragResult<Vec<KnowledgeEntry>> {
let all_entries = self.get_all_entries().await?;
let matching = all_entries
.into_iter()
.filter(|e| e.has_access(&self.agent_id) && e.created_by == creator_agent_id)
.collect();
Ok(matching)
}
pub async fn get_all_entries(&self) -> RragResult<Vec<KnowledgeEntry>> {
let all_keys = self.list_entry_keys().await?;
let mut entries = Vec::new();
for key in all_keys {
if let Some(entry) = self.get(&key).await? {
entries.push(entry);
}
}
Ok(entries)
}
pub async fn count(&self) -> RragResult<usize> {
self.storage.count(Some(&self.namespace)).await
}
pub async fn clear(&self) -> RragResult<()> {
self.storage.clear(Some(&self.namespace)).await
}
fn entry_key(&self, key: &str) -> String {
format!("{}::{}", self.namespace, key)
}
async fn list_entry_keys(&self) -> RragResult<Vec<String>> {
use crate::storage::MemoryQuery;
let query = MemoryQuery::new().with_namespace(self.namespace.clone());
let all_keys = self.storage.keys(&query).await?;
let prefix = format!("{}::", self.namespace);
let keys = all_keys
.into_iter()
.filter_map(|k| k.strip_prefix(&prefix).map(String::from))
.collect();
Ok(keys)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::InMemoryStorage;
#[tokio::test]
async fn test_shared_knowledge_store_and_retrieve() {
let storage = Arc::new(InMemoryStorage::new());
let kb = SharedKnowledgeBase::new(storage, "agent1".to_string());
kb.store("api_key", MemoryValue::from("secret123"))
.await
.unwrap();
let value = kb.get_value("api_key").await.unwrap().unwrap();
assert_eq!(value.as_string(), Some("secret123"));
}
#[tokio::test]
async fn test_shared_knowledge_cross_agent_access() {
let storage = Arc::new(InMemoryStorage::new());
let kb1 = SharedKnowledgeBase::new(storage.clone(), "agent1".to_string());
let kb2 = SharedKnowledgeBase::new(storage.clone(), "agent2".to_string());
kb1.store("shared_config", MemoryValue::from("config_value"))
.await
.unwrap();
let value = kb2.get_value("shared_config").await.unwrap().unwrap();
assert_eq!(value.as_string(), Some("config_value"));
}
#[tokio::test]
async fn test_shared_knowledge_with_acl() {
let storage = Arc::new(InMemoryStorage::new());
let kb1 = SharedKnowledgeBase::new(storage.clone(), "agent1".to_string());
let kb2 = SharedKnowledgeBase::new(storage.clone(), "agent2".to_string());
let kb3 = SharedKnowledgeBase::new(storage.clone(), "agent3".to_string());
let entry = KnowledgeEntry::new("private_data", MemoryValue::from("secret"), "agent1")
.with_acl(vec!["agent1".to_string(), "agent2".to_string()]);
kb1.store_entry(entry).await.unwrap();
assert!(kb2.get("private_data").await.unwrap().is_some());
assert!(kb3.get("private_data").await.unwrap().is_none());
}
#[tokio::test]
async fn test_shared_knowledge_with_tags() {
let storage = Arc::new(InMemoryStorage::new());
let kb = SharedKnowledgeBase::new(storage, "agent1".to_string());
kb.store_with_tags(
"config1",
MemoryValue::from("value1"),
vec!["config".to_string(), "production".to_string()],
)
.await
.unwrap();
kb.store_with_tags(
"config2",
MemoryValue::from("value2"),
vec!["config".to_string(), "development".to_string()],
)
.await
.unwrap();
let config_entries = kb.find_by_tag("config").await.unwrap();
assert_eq!(config_entries.len(), 2);
let prod_entries = kb.find_by_tag("production").await.unwrap();
assert_eq!(prod_entries.len(), 1);
}
#[tokio::test]
async fn test_shared_knowledge_delete_permissions() {
let storage = Arc::new(InMemoryStorage::new());
let kb1 = SharedKnowledgeBase::new(storage.clone(), "agent1".to_string());
let kb2 = SharedKnowledgeBase::new(storage.clone(), "agent2".to_string());
kb1.store("data", MemoryValue::from("value")).await.unwrap();
let deleted = kb2.delete("data").await.unwrap();
assert!(!deleted);
let deleted = kb1.delete("data").await.unwrap();
assert!(deleted);
}
}