use std::collections::{HashMap, HashSet};
use std::path::Path;
use std::sync::Arc;
use serde_json::{json, Map, Value};
use velesdb_core::agent::AgentMemory;
use velesdb_core::{Database, SearchResult};
pub type Metadata = Map<String, Value>;
use crate::embedder::Embedder;
use crate::error::MemoryError;
use crate::extract::Extractor;
use crate::id;
use crate::model::{ColumnFilter, Explanation, Link, MemoryEdge, MemoryNode, Recollection};
const HUB_FIELD: &str = "_veles_hub";
const HUB_ID_SALT: &str = "\u{0}_veles_entity_hub\u{0}";
pub struct MemoryService<E: Embedder> {
memory: AgentMemory,
embedder: E,
}
impl<E: Embedder> MemoryService<E> {
pub fn open<P: AsRef<Path>>(path: P, embedder: E) -> Result<Self, MemoryError> {
let db = Arc::new(Database::open(path)?);
let memory = AgentMemory::with_dimension(db, embedder.dimension())?;
Ok(Self { memory, embedder })
}
pub fn remember(
&self,
fact: &str,
links: &[Link],
metadata: Option<&Metadata>,
) -> Result<u64, MemoryError> {
self.remember_with_ttl(fact, links, metadata, None)
}
pub fn remember_with_ttl(
&self,
fact: &str,
links: &[Link],
metadata: Option<&Metadata>,
ttl_seconds: Option<u64>,
) -> Result<u64, MemoryError> {
let fact = fact.trim();
if fact.is_empty() {
return Err(MemoryError::EmptyFact);
}
reject_reserved_keys(metadata)?;
self.ensure_link_targets_exist(links)?;
let fact_id = id::stable_id(fact);
let embedding = self.embedder.embed(fact)?;
self.store(
fact_id,
fact,
&embedding,
metadata,
positive_ttl(ttl_seconds),
)?;
self.relate_links(fact_id, links)?;
Ok(fact_id)
}
fn relate_links(&self, fact_id: u64, links: &[Link]) -> Result<(), MemoryError> {
for link in links {
validate_relation(&link.relation)?;
self.memory
.semantic()
.relate(fact_id, link.target, &link.relation, None)?;
}
Ok(())
}
pub fn remember_extracted<X: Extractor>(
&self,
text: &str,
extractor: &X,
metadata: Option<&Metadata>,
) -> Result<Vec<u64>, MemoryError> {
let text = text.trim();
if text.is_empty() {
return Err(MemoryError::EmptyFact);
}
let facts = extractor.extract(text)?;
let mut fact_ids = Vec::with_capacity(facts.len());
let mut entity_ids: HashMap<String, u64> = HashMap::new();
let mut edges: HashSet<(u64, u64)> = HashSet::new();
let mut seeded: HashSet<u64> = HashSet::new();
for fact in &facts {
let content = fact.text.trim();
if content.is_empty() {
continue;
}
let fact_id = self.remember(content, &[], metadata)?;
fact_ids.push(fact_id);
self.wire_entities(
fact_id,
&fact.entities,
&mut entity_ids,
&mut edges,
&mut seeded,
)?;
}
Ok(fact_ids)
}
fn wire_entities(
&self,
fact_id: u64,
entities: &[String],
entity_ids: &mut HashMap<String, u64>,
edges: &mut HashSet<(u64, u64)>,
seeded: &mut HashSet<u64>,
) -> Result<(), MemoryError> {
for entity in entities {
if entity.chars().any(char::is_alphanumeric) {
self.wire_entity(fact_id, entity, entity_ids, edges, seeded)?;
}
}
Ok(())
}
fn wire_entity(
&self,
fact_id: u64,
entity: &str,
entity_ids: &mut HashMap<String, u64>,
edges: &mut HashSet<(u64, u64)>,
seeded: &mut HashSet<u64>,
) -> Result<(), MemoryError> {
let entity_id = self.entity_hub(entity, entity_ids)?;
if entity_id == fact_id {
return Ok(());
}
self.seed_existing_edges(fact_id, edges, seeded)?;
self.seed_existing_edges(entity_id, edges, seeded)?;
self.add_edge(fact_id, entity_id, "about", edges)?;
self.add_edge(entity_id, fact_id, "mentions", edges)?;
Ok(())
}
fn add_edge(
&self,
from: u64,
to: u64,
label: &str,
edges: &mut HashSet<(u64, u64)>,
) -> Result<(), MemoryError> {
if edges.insert((from, to)) {
self.relate(from, to, label)?;
}
Ok(())
}
fn seed_existing_edges(
&self,
node: u64,
edges: &mut HashSet<(u64, u64)>,
seeded: &mut HashSet<u64>,
) -> Result<(), MemoryError> {
if !seeded.insert(node) {
return Ok(());
}
for edge in self.memory.semantic().relations(node)? {
edges.insert((node, edge.target()));
}
Ok(())
}
fn entity_hub(
&self,
entity: &str,
entity_ids: &mut HashMap<String, u64>,
) -> Result<u64, MemoryError> {
let key = entity.trim().to_lowercase();
if let Some(&id) = entity_ids.get(&key) {
return Ok(id);
}
let id = self.remember_hub(&key)?;
entity_ids.insert(key, id);
Ok(id)
}
fn remember_hub(&self, key: &str) -> Result<u64, MemoryError> {
let id = id::stable_id(&format!("{HUB_ID_SALT}{key}"));
let content = format!("Entity: {key}");
let embedding = self.embedder.embed(&content)?;
let mut meta = Map::new();
meta.insert(HUB_FIELD.to_string(), Value::Bool(true));
self.store(id, &content, &embedding, Some(&meta), None)?;
Ok(id)
}
fn ensure_exists(&self, id: u64) -> Result<(), MemoryError> {
if self.memory.semantic().get(id)?.is_none() {
return Err(MemoryError::UnknownMemory(id));
}
Ok(())
}
fn ensure_link_targets_exist(&self, links: &[Link]) -> Result<(), MemoryError> {
for link in links {
self.ensure_exists(link.target)?;
}
Ok(())
}
fn store(
&self,
id: u64,
fact: &str,
embedding: &[f32],
metadata: Option<&Metadata>,
ttl_seconds: Option<u64>,
) -> Result<(), MemoryError> {
let semantic = self.memory.semantic();
match (metadata, ttl_seconds) {
(Some(meta), Some(ttl)) => {
semantic.store_with_ttl(id, fact, embedding, ttl)?;
semantic.update_metadata(id, meta)?;
}
(Some(meta), None) => semantic.store_with_metadata(id, fact, embedding, meta)?,
(None, Some(ttl)) => semantic.store_with_ttl(id, fact, embedding, ttl)?,
(None, None) => semantic.store(id, fact, embedding)?,
}
Ok(())
}
pub fn recall(
&self,
query: &str,
k: usize,
filter: Option<&Metadata>,
) -> Result<Vec<Recollection>, MemoryError> {
let query = query.trim();
if query.is_empty() {
return Ok(Vec::new());
}
reject_reserved_keys(filter)?;
let embedding = self.embedder.embed(query)?;
let hits = self.search(&embedding, k, filter)?;
Ok(hits
.into_iter()
.map(|(id, score, content)| Recollection { id, score, content })
.collect())
}
fn search(
&self,
embedding: &[f32],
k: usize,
filter: Option<&Metadata>,
) -> Result<Vec<(u64, f32, String)>, MemoryError> {
match filter {
Some(meta) => self
.memory
.semantic()
.query_filtered(embedding, k, meta, 0)
.map_err(MemoryError::from),
None => self
.memory
.semantic()
.query_excluding(embedding, k, &hub_exclude_filter())
.map_err(MemoryError::from),
}
}
pub fn recall_where(
&self,
query: &str,
k: usize,
filters: &[ColumnFilter],
) -> Result<Vec<Recollection>, MemoryError> {
let query = query.trim();
if query.is_empty() || k == 0 {
return Ok(Vec::new());
}
let embedding = self.embedder.embed(query)?;
let (sql, params) = self.build_fused_query(&embedding, k, filters)?;
for filter in filters {
self.memory
.semantic()
.ensure_index(&filter.field)
.map_err(MemoryError::from)?;
}
let results = self
.memory
.query_semantic(&sql, ¶ms)
.map_err(MemoryError::from)?;
Ok(results.iter().map(to_recollection).collect())
}
fn build_fused_query(
&self,
embedding: &[f32],
k: usize,
filters: &[ColumnFilter],
) -> Result<(String, HashMap<String, Value>), MemoryError> {
use std::fmt::Write as _;
let mut params: HashMap<String, Value> = HashMap::new();
params.insert("q".to_string(), json!(embedding));
let mut predicate = String::from("vector NEAR $q");
for (index, filter) in filters.iter().enumerate() {
validate_field(&filter.field)?;
validate_scalar(&filter.value)?;
let key = format!("p{index}");
let _ = write!(
predicate,
" AND {} {} ${key}",
filter.field,
filter.op.as_sql()
);
params.insert(key, filter.value.clone());
}
let sql = format!(
"SELECT * FROM {} WHERE {predicate} LIMIT {k}",
self.memory.semantic().collection_name()
);
Ok((sql, params))
}
pub fn relate(&self, from: u64, to: u64, relation: &str) -> Result<u64, MemoryError> {
validate_relation(relation)?;
self.ensure_exists(from)?;
self.ensure_exists(to)?;
self.memory
.semantic()
.relate(from, to, relation, None)
.map_err(MemoryError::from)
}
pub fn forget(&self, fact_id: u64) -> Result<(), MemoryError> {
self.memory
.semantic()
.delete(fact_id)
.map_err(MemoryError::from)
}
pub fn why(
&self,
decision: &str,
max_hops: usize,
filter: Option<&Metadata>,
) -> Result<Explanation, MemoryError> {
let decision = decision.trim();
if decision.is_empty() {
return Ok(Explanation::default());
}
reject_reserved_keys(filter)?;
let embedding = self.embedder.embed(decision)?;
let seeds = self.search(&embedding, 1, filter)?;
let Some((seed_id, _score, seed_content)) = seeds.into_iter().next() else {
return Ok(Explanation::default());
};
self.traverse(seed_id, seed_content, max_hops)
}
fn traverse(
&self,
seed_id: u64,
seed_content: String,
max_hops: usize,
) -> Result<Explanation, MemoryError> {
let mut explanation = Explanation {
nodes: vec![MemoryNode {
id: seed_id,
content: seed_content,
hop: 0,
}],
edges: Vec::new(),
};
let mut visited: HashSet<u64> = HashSet::from([seed_id]);
let mut frontier = vec![seed_id];
let mut next: Vec<u64> = Vec::new();
for hop in 1..=max_hops {
next.clear();
for node_id in frontier.drain(..) {
self.expand(node_id, hop, &mut explanation, &mut visited, &mut next)?;
}
if next.is_empty() {
break;
}
std::mem::swap(&mut frontier, &mut next);
}
Ok(explanation)
}
fn expand(
&self,
node_id: u64,
hop: usize,
explanation: &mut Explanation,
visited: &mut HashSet<u64>,
next: &mut Vec<u64>,
) -> Result<(), MemoryError> {
for edge in self.memory.semantic().relations(node_id)? {
let target = edge.target();
if !visited.contains(&target) {
let Some((content, _embedding)) = self.memory.semantic().get(target)? else {
continue; };
visited.insert(target);
explanation.nodes.push(MemoryNode {
id: target,
content,
hop,
});
next.push(target);
}
explanation.edges.push(MemoryEdge {
from: edge.source(),
to: target,
relation: edge.label().to_owned(),
});
}
Ok(())
}
}
fn hub_exclude_filter() -> Metadata {
let mut exclude = Map::new();
exclude.insert(HUB_FIELD.to_string(), Value::Bool(true));
exclude
}
fn is_reserved_key(key: &str) -> bool {
key == "content" || key.starts_with("_veles_")
}
fn reject_reserved_keys(metadata: Option<&Metadata>) -> Result<(), MemoryError> {
let Some(meta) = metadata else {
return Ok(());
};
for key in meta.keys() {
if is_reserved_key(key) {
return Err(MemoryError::ReservedKey(key.clone()));
}
}
Ok(())
}
fn positive_ttl(ttl_seconds: Option<u64>) -> Option<u64> {
ttl_seconds.filter(|&seconds| seconds > 0)
}
fn to_recollection(result: &SearchResult) -> Recollection {
let content = result
.point
.payload
.as_ref()
.and_then(|payload| payload.get("content"))
.and_then(Value::as_str)
.unwrap_or_default()
.to_owned();
Recollection {
id: result.point.id,
score: result.score,
content,
}
}
fn validate_field(field: &str) -> Result<(), MemoryError> {
let plain = !field.is_empty() && field.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
let reserved = field == "content" || field.starts_with("_veles_");
if plain && !reserved {
Ok(())
} else {
Err(MemoryError::InvalidFilter(field.to_owned()))
}
}
fn validate_scalar(value: &Value) -> Result<(), MemoryError> {
match value {
Value::String(_) | Value::Number(_) | Value::Bool(_) => Ok(()),
_ => Err(MemoryError::InvalidFilter(format!(
"value must be a string, number, or boolean, got {value}"
))),
}
}
const MAX_RELATION_BYTES: usize = 512;
fn validate_relation(label: &str) -> Result<(), MemoryError> {
if label.is_empty() {
return Err(MemoryError::InvalidRelation(
"relation label must not be empty".to_owned(),
));
}
if label.len() > MAX_RELATION_BYTES {
return Err(MemoryError::InvalidRelation(format!(
"relation label exceeds maximum of {MAX_RELATION_BYTES} bytes ({} given)",
label.len()
)));
}
if label.chars().any(|c| c.is_ascii() && c.is_ascii_control()) {
return Err(MemoryError::InvalidRelation(
"relation label must not contain ASCII control characters".to_owned(),
));
}
Ok(())
}