use std::collections::{BTreeMap, HashMap, HashSet};
use std::sync::Arc;
use async_trait::async_trait;
use trusty_common::memory_core::palace::{Drawer, DrawerType, RoomType};
use trusty_common::memory_core::retrieval::{shared_embedder, PalaceHandle, RememberOptions};
use trusty_common::memory_core::store::{Triple, VectorStore as _};
use uuid::Uuid;
use super::bridge::KuzuExport;
use super::ledger::{Ledger, MemoryPlan};
use super::mapping::{
drawer_subject, entity_subject, entity_triples, map_memory, merge_tags, relates_to_predicate,
source_key, triple, MappedMemory,
};
use super::retract::{retract_stale, Retraction};
use super::screen::{screen, tally, RuleTally, SecretRule};
use super::KuzuImportError;
pub const SECRET_SHAPED_ID: &str = "(secret-shaped id)";
#[async_trait]
pub trait PalaceView: Send + Sync {
fn drawers(&self) -> Vec<Drawer>;
async fn active_triples(&self, subject: &str) -> Result<Vec<Triple>, KuzuImportError>;
async fn triple_is_active(&self, t: &Triple) -> Result<bool, KuzuImportError> {
let active = self.active_triples(&t.subject).await?;
Ok(active
.iter()
.any(|a| a.predicate == t.predicate && a.object == t.object))
}
}
#[async_trait]
pub trait PalaceSink: PalaceView {
async fn insert_drawer(&self, m: &MappedMemory) -> Result<Uuid, KuzuImportError>;
async fn stamp_drawer(&self, id: Uuid, m: &MappedMemory) -> Result<(), KuzuImportError>;
async fn update_memory(&self, id: Uuid, m: &MappedMemory) -> Result<(), KuzuImportError>;
async fn assert_triple(&self, t: Triple) -> Result<(), KuzuImportError>;
async fn retract_triple(&self, t: &Triple) -> Result<(), KuzuImportError>;
}
#[derive(Debug, Default, Clone, PartialEq, Eq)]
pub struct StoreCounts {
pub memories: usize,
pub edges: usize,
pub entities: usize,
pub skipped_empty: usize,
pub new_memories: usize,
pub unchanged: usize,
pub changed: usize,
pub updated: usize,
pub new_triples: usize,
pub existing_triples: usize,
pub dangling_edges: usize,
pub skipped_edges: usize,
pub failed_writes: usize,
pub refused_ids: Vec<String>,
pub shared_ids: Vec<String>,
pub refused_tags: usize,
pub refused_triples: usize,
pub refusal_rules: RuleTally,
pub retracted: Vec<Retraction>,
pub unsupported_edges: BTreeMap<String, usize>,
pub unsupported_error: Option<String>,
}
impl StoreCounts {
fn refuse_under(&mut self, rule: SecretRule, bucket: fn(&mut Self) -> &mut usize) {
*bucket(self) += 1;
tally(&mut self.refusal_rules, rule);
}
}
#[derive(Debug, Default)]
pub struct StorePlan {
pub memories: Vec<(MappedMemory, MemoryPlan)>,
pub refused_memory_ids: Vec<String>,
pub retract_scope: HashSet<String>,
pub entity_triples: Vec<Triple>,
pub mentions: Vec<(String, String, Option<f64>)>,
pub relates: Vec<(String, String, String, Option<f64>)>,
pub counts: StoreCounts,
}
pub fn plan_store(
export: &KuzuExport,
store: &str,
ledger: &Ledger,
store_is_live: &dyn Fn(&str) -> bool,
) -> StorePlan {
let mut plan = StorePlan::default();
let c = &mut plan.counts;
c.memories = export.memories.len();
c.edges = export.edge_count();
c.entities = export.entities.len();
(c.unsupported_edges, c.unsupported_error) = unsupported_edges(export);
for row in &export.memories {
if let Some(rule) = row.id.as_deref().and_then(screen) {
tally(&mut c.refusal_rules, rule);
c.refused_ids.push(SECRET_SHAPED_ID.to_string());
plan.refused_memory_ids.extend(row.id.clone());
continue;
}
let Some(m) = map_memory(row, store) else {
c.skipped_empty += 1;
continue;
};
for rule in &m.refused_tags {
c.refuse_under(*rule, |c| &mut c.refused_tags);
}
let own = ledger
.get(&source_key(&m.memory_id))
.and_then(|d| d.store.as_deref());
if own.is_none_or(|s| s == store || !store_is_live(s)) {
plan.retract_scope.insert(m.memory_id.clone());
}
let p = ledger.plan(&m, store_is_live);
plan.memories.push((m, p));
}
for e in &export.entities {
let (triples, refused) = entity_triples(e);
plan.entity_triples.extend(triples);
for rule in refused {
c.refuse_under(rule, |c| &mut c.refused_triples);
}
}
for m in &export.mentions {
let (Some(mem), Some(ent)) = (&m.memory_id, &m.entity_id) else {
c.dangling_edges += 1;
continue;
};
match screen(ent) {
Some(rule) => c.refuse_under(rule, |c| &mut c.refused_triples),
None => plan.mentions.push((mem.clone(), ent.clone(), m.confidence)),
}
}
for r in &export.relates_to {
let (Some(a), Some(b)) = (&r.from_id, &r.to_id) else {
c.dangling_edges += 1;
continue;
};
match r.relationship_type.as_deref().and_then(screen) {
Some(rule) => c.refuse_under(rule, |c| &mut c.refused_triples),
None => {
let pred = relates_to_predicate(r.relationship_type.as_deref());
plan.relates.push((a.clone(), b.clone(), pred, r.strength));
}
}
}
plan
}
fn unsupported_edges(export: &KuzuExport) -> (BTreeMap<String, usize>, Option<String>) {
let tables = export
.other_edges
.iter()
.filter(|(_, n)| **n > 0)
.map(|(t, n)| (t.clone(), *n))
.collect();
(tables, export.other_edges_error.clone())
}
#[derive(Debug, Clone, Copy)]
enum Endpoint {
Drawer(Uuid),
Pending,
Failed,
Skipped,
Dangling,
}
fn content_rule(m: &MappedMemory) -> Option<SecretRule> {
screen(m.content.trim())
}
fn refuse(c: &mut StoreCounts, m: &MappedMemory, rule: SecretRule) {
tracing::warn!(memory_id = %m.memory_id, rule = rule.label(), "kuzu import: refused a secret-shaped memory");
c.refused_ids.push(m.memory_id.clone());
tally(&mut c.refusal_rules, rule);
}
async fn insert_new(
s: &dyn PalaceSink,
m: &MappedMemory,
) -> Result<Uuid, (Option<Uuid>, KuzuImportError)> {
let id = s.insert_drawer(m).await.map_err(|e| (None, e))?;
s.stamp_drawer(id, m).await.map_err(|e| (Some(id), e))?;
Ok(id)
}
pub async fn execute(
mut plan: StorePlan,
view: &dyn PalaceView,
sink: Option<&dyn PalaceSink>,
update: bool,
) -> StoreCounts {
let mut ids: HashMap<String, Endpoint> = HashMap::new();
for id in &plan.refused_memory_ids {
ids.insert(id.clone(), Endpoint::Skipped);
}
let memories = std::mem::take(&mut plan.memories);
let c = &mut plan.counts;
for (m, p) in &memories {
let rule = content_rule(m);
let end = match *p {
MemoryPlan::Unchanged(id) => {
c.unchanged += 1;
Endpoint::Drawer(id)
}
MemoryPlan::SharedId(_) => {
c.shared_ids.push(m.memory_id.clone());
Endpoint::Skipped
}
MemoryPlan::Changed(id) => {
c.changed += 1;
if let (true, Some(r)) = (update, rule) {
refuse(c, m, r);
} else if update {
match sink {
Some(s) => match s.update_memory(id, m).await {
Ok(()) => c.updated += 1,
Err(e) => fail(c, "update memory", &e),
},
None => c.updated += 1,
}
}
Endpoint::Drawer(id)
}
_ if rule.is_some() => {
if let Some(r) = rule {
refuse(c, m, r);
}
Endpoint::Skipped
}
MemoryPlan::Resume(id) => match sink {
Some(s) => match s.update_memory(id, m).await {
Ok(()) => {
c.new_memories += 1;
Endpoint::Drawer(id)
}
Err(e) => {
fail(c, "finish pending memory", &e);
Endpoint::Failed
}
},
None => {
c.new_memories += 1;
Endpoint::Drawer(id)
}
},
MemoryPlan::New => match sink {
Some(s) => match insert_new(s, m).await {
Ok(id) => {
c.new_memories += 1;
Endpoint::Drawer(id)
}
Err((_, e)) => {
fail(c, "insert memory", &e);
Endpoint::Failed
}
},
None => {
c.new_memories += 1;
Endpoint::Pending
}
},
};
ids.insert(m.memory_id.clone(), end);
}
let resolve = |memory_id: &str| ids.get(memory_id).copied().unwrap_or(Endpoint::Dangling);
let mut edges: Vec<Result<Triple, Endpoint>> = Vec::new();
for t in std::mem::take(&mut plan.entity_triples) {
edges.push(Ok(t));
}
for (mem, ent, conf) in &plan.mentions {
edges.push(match resolve(mem) {
Endpoint::Drawer(id) => Ok(triple(
drawer_subject(id),
"mentions",
entity_subject(ent),
*conf,
)),
other => Err(other),
});
}
for (a, b, pred, conf) in &plan.relates {
edges.push(match (resolve(a), resolve(b)) {
(Endpoint::Drawer(x), Endpoint::Drawer(y)) => {
Ok(triple(drawer_subject(x), pred, drawer_subject(y), *conf))
}
(Endpoint::Dangling, _) | (_, Endpoint::Dangling) => Err(Endpoint::Dangling),
(Endpoint::Skipped, _) | (_, Endpoint::Skipped) => Err(Endpoint::Skipped),
(Endpoint::Failed, _) | (_, Endpoint::Failed) => Err(Endpoint::Failed),
_ => Err(Endpoint::Pending),
});
}
let keep: HashSet<(String, String, String)> = edges
.iter()
.flatten()
.map(|t| (t.subject.clone(), t.predicate.clone(), t.object.clone()))
.collect();
let scope: Vec<(String, Uuid)> = memories
.iter()
.filter(|(m, _)| plan.retract_scope.contains(&m.memory_id))
.filter_map(|(m, _)| match ids.get(&m.memory_id) {
Some(Endpoint::Drawer(id)) => Some((m.memory_id.clone(), *id)),
_ => None,
})
.collect();
let c = &mut plan.counts;
for edge in edges {
match edge {
Err(Endpoint::Pending) => c.new_triples += 1,
Err(Endpoint::Dangling) => c.dangling_edges += 1,
Err(Endpoint::Skipped) => c.skipped_edges += 1,
Err(_) => c.failed_writes += 1,
Ok(t) => match view.triple_is_active(&t).await {
Ok(true) => c.existing_triples += 1,
Ok(false) => match sink {
None => c.new_triples += 1,
Some(s) => match s.assert_triple(t).await {
Ok(()) => c.new_triples += 1,
Err(e) => fail(c, "assert triple", &e),
},
},
Err(e) => fail(c, "read triple", &e),
},
}
}
if update {
retract_stale(&scope, &keep, view, sink, c).await;
}
plan.counts
}
pub(super) fn fail(c: &mut StoreCounts, what: &str, e: &KuzuImportError) {
tracing::warn!(error_kind = %e.kind(), "kuzu import: {what} failed");
c.failed_writes += 1;
}
pub struct HandleSink {
pub handle: Arc<PalaceHandle>,
}
fn sink_err(e: anyhow::Error) -> KuzuImportError {
KuzuImportError::Palace(format!("{e:#}"))
}
#[async_trait]
impl PalaceView for HandleSink {
fn drawers(&self) -> Vec<Drawer> {
self.handle.drawers.read().clone()
}
async fn active_triples(&self, subject: &str) -> Result<Vec<Triple>, KuzuImportError> {
self.handle.kg.query_active(subject).await.map_err(sink_err)
}
}
#[async_trait]
impl PalaceSink for HandleSink {
async fn insert_drawer(&self, m: &MappedMemory) -> Result<Uuid, KuzuImportError> {
let opts = RememberOptions {
force: true,
enforce_min_tokens: false,
classify_as: Some(DrawerType::Unknown),
..RememberOptions::default()
};
self.handle
.remember_with_options(
m.content.clone(),
RoomType::General,
m.staging_tags(),
m.importance,
opts,
)
.await
.map_err(sink_err)
}
async fn stamp_drawer(&self, id: Uuid, m: &MappedMemory) -> Result<(), KuzuImportError> {
let (tags, created) = (m.tags.clone(), m.created_at);
self.rewrite(id, move |d| {
d.tags = merge_tags(&d.tags, &tags);
if let Some(c) = created {
d.created_at = c;
}
})
.await
}
async fn update_memory(&self, id: Uuid, m: &MappedMemory) -> Result<(), KuzuImportError> {
let embedder = shared_embedder().await.map_err(sink_err)?;
let vectors = embedder
.embed_batch(std::slice::from_ref(&m.content))
.await
.map_err(sink_err)?;
if let Some(v) = vectors.into_iter().next() {
self.handle
.vector_store
.upsert(id, v)
.await
.map_err(sink_err)?;
}
let (content, tags, importance, created) = (
m.content.clone(),
m.tags.clone(),
m.importance,
m.created_at,
);
self.rewrite(id, move |d| {
d.set_content(content);
d.tags = merge_tags(&d.tags, &tags);
d.importance = importance;
if let Some(c) = created {
d.created_at = c;
}
})
.await
}
async fn assert_triple(&self, t: Triple) -> Result<(), KuzuImportError> {
self.handle.kg.assert(t).await.map_err(sink_err)
}
async fn retract_triple(&self, t: &Triple) -> Result<(), KuzuImportError> {
self.handle
.kg
.retract_triple(&t.subject, &t.predicate, &t.object)
.await
.map(drop)
.map_err(sink_err)
}
}
impl HandleSink {
async fn rewrite(
&self,
id: Uuid,
edit: impl FnOnce(&mut Drawer) + Send,
) -> Result<(), KuzuImportError> {
let current = self
.handle
.drawers
.read()
.iter()
.find(|d| d.id == id)
.cloned();
let mut drawer = match current {
Some(d) => d,
None => self
.handle
.kg
.load_drawer(id)
.map_err(sink_err)?
.ok_or_else(|| KuzuImportError::Palace(format!("drawer {id} not found")))?,
};
edit(&mut drawer);
self.handle
.kg
.upsert_drawer(&drawer)
.await
.map_err(sink_err)?;
let mut table = self.handle.drawers.write();
if let Some(slot) = table.iter_mut().find(|d| d.id == id) {
*slot = drawer;
}
Ok(())
}
}
pub struct SnapshotView {
pub drawers: Vec<Drawer>,
pub kg: Option<trusty_common::memory_core::store::KnowledgeGraph>,
pub snapshot_dir: Option<tempfile::TempDir>,
}
#[async_trait]
impl PalaceView for SnapshotView {
fn drawers(&self) -> Vec<Drawer> {
self.drawers.clone()
}
async fn active_triples(&self, subject: &str) -> Result<Vec<Triple>, KuzuImportError> {
match &self.kg {
Some(kg) => kg.query_active(subject).await.map_err(sink_err),
None => Ok(Vec::new()),
}
}
}