use std::collections::BTreeSet;
use serde::{Deserialize, Serialize};
use crate::store::{Store, StoreError};
use crate::{Edge, Node, SCHEMA};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ContextNode {
pub key: String,
pub kind: String,
pub name: String,
pub path: Option<String>,
pub lang: Option<String>,
}
impl ContextNode {
fn from_node(node: &Node) -> Self {
Self {
key: node.key.clone(),
kind: node.kind.as_str().to_owned(),
name: node.name.clone(),
path: node.path.clone(),
lang: node.lang.clone(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ContextEdge {
pub kind: String,
pub provenance: String,
pub confidence: Option<f64>,
pub node: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct NodeContext {
pub schema: String,
pub fingerprint: String,
pub node: ContextNode,
pub meta: serde_json::Value,
pub outgoing: Vec<ContextEdge>,
pub incoming: Vec<ContextEdge>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
pub struct ContextRefresh {
pub rebuilt: usize,
pub reused: usize,
pub pruned: usize,
}
fn fnv1a64(bytes: &[u8]) -> u64 {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
for &b in bytes {
hash ^= u64::from(b);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash
}
fn node_signature(node: &Node) -> u64 {
if let Some(blob) = &node.blob_hash {
return fnv1a64(blob.as_bytes());
}
let mut s = String::new();
s.push_str(node.kind.as_str());
s.push('\u{0}');
s.push_str(&node.name);
s.push('\u{0}');
s.push_str(&serde_json::to_string(&node.meta).unwrap_or_default());
fnv1a64(s.as_bytes())
}
fn edge_descriptor(edge: &Edge, direction: &str, neighbour: &str, neighbour_sig: u64) -> String {
let confidence = edge
.confidence
.map_or_else(String::new, |c| format!("{:016x}", c.to_bits()));
format!(
"{direction}|{}|{}|{confidence}|{neighbour}|{neighbour_sig:016x}",
edge.kind.as_str(),
edge.provenance.as_str(),
)
}
fn compute_fingerprint(store: &Store, node: &Node) -> Result<String, StoreError> {
let mut descriptors: Vec<String> = Vec::new();
for edge in store.edges_from(&node.key)? {
let sig = store.get_node(&edge.dst)?.map_or(0, |n| node_signature(&n));
descriptors.push(edge_descriptor(&edge, "out", &edge.dst, sig));
}
for edge in store.edges_to(&node.key)? {
let sig = store.get_node(&edge.src)?.map_or(0, |n| node_signature(&n));
descriptors.push(edge_descriptor(&edge, "in", &edge.src, sig));
}
descriptors.sort();
let mut buf = format!("ctx/v1|{:016x}", node_signature(node));
for d in &descriptors {
buf.push('\n');
buf.push_str(d);
}
Ok(format!("{:016x}", fnv1a64(buf.as_bytes())))
}
fn out_ref(edge: &Edge) -> ContextEdge {
ContextEdge {
kind: edge.kind.as_str().to_owned(),
provenance: edge.provenance.as_str().to_owned(),
confidence: edge.confidence,
node: edge.dst.clone(),
}
}
fn in_ref(edge: &Edge) -> ContextEdge {
ContextEdge {
kind: edge.kind.as_str().to_owned(),
provenance: edge.provenance.as_str().to_owned(),
confidence: edge.confidence,
node: edge.src.clone(),
}
}
fn sort_refs(refs: &mut [ContextEdge]) {
refs.sort_by(|a, b| (&a.kind, &a.node, &a.provenance).cmp(&(&b.kind, &b.node, &b.provenance)));
}
fn fresh_bundle(
store: &Store,
node: &Node,
fingerprint: String,
) -> Result<NodeContext, StoreError> {
let mut outgoing: Vec<ContextEdge> = store.edges_from(&node.key)?.iter().map(out_ref).collect();
let mut incoming: Vec<ContextEdge> = store.edges_to(&node.key)?.iter().map(in_ref).collect();
sort_refs(&mut outgoing);
sort_refs(&mut incoming);
Ok(NodeContext {
schema: SCHEMA.to_owned(),
fingerprint,
node: ContextNode::from_node(node),
meta: node.meta.clone(),
outgoing,
incoming,
})
}
pub fn build_context(store: &Store, key: &str) -> Result<Option<NodeContext>, StoreError> {
let Some(node) = store.get_node(key)? else {
return Ok(None);
};
let fingerprint = compute_fingerprint(store, &node)?;
Ok(Some(fresh_bundle(store, &node, fingerprint)?))
}
pub fn context(store: &Store, key: &str) -> Result<Option<NodeContext>, StoreError> {
let Some(node) = store.get_node(key)? else {
store.context_cache_delete(key)?;
return Ok(None);
};
let fingerprint = compute_fingerprint(store, &node)?;
if let Some((cached_fp, json)) = store.context_cache_get(key)?
&& cached_fp == fingerprint
{
return Ok(Some(serde_json::from_str(&json)?));
}
let bundle = fresh_bundle(store, &node, fingerprint.clone())?;
store.context_cache_put(key, &fingerprint, &serde_json::to_string(&bundle)?)?;
Ok(Some(bundle))
}
pub fn dependents(store: &Store, changed: &[String]) -> Result<BTreeSet<String>, StoreError> {
let mut set = BTreeSet::new();
for key in changed {
set.insert(key.clone());
for edge in store.edges_from(key)? {
set.insert(edge.dst);
}
for edge in store.edges_to(key)? {
set.insert(edge.src);
}
}
Ok(set)
}
pub const TOOL_CONTEXT_EDGE_CAP: usize = 50;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct OmittedEdges {
pub kind: String,
pub omitted: usize,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BoundedEdges {
pub total: usize,
pub truncated: bool,
pub omitted: Vec<OmittedEdges>,
pub edges: Vec<ContextEdge>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolContext {
pub schema: String,
pub fingerprint: String,
pub edge_cap: usize,
pub truncated: bool,
pub node: ContextNode,
pub meta: serde_json::Value,
pub outgoing: BoundedEdges,
pub incoming: BoundedEdges,
}
fn bound_edges(refs: Vec<ContextEdge>, cap: usize) -> BoundedEdges {
let total = refs.len();
if total <= cap {
return BoundedEdges {
total,
truncated: false,
omitted: Vec::new(),
edges: refs,
};
}
let mut by_kind: Vec<(String, Vec<ContextEdge>)> = Vec::new();
for r in refs {
match by_kind.last_mut() {
Some((kind, group)) if *kind == r.kind => group.push(r),
_ => by_kind.push((r.kind.clone(), vec![r])),
}
}
let mut kept: Vec<ContextEdge> = Vec::with_capacity(cap);
let mut round = 0usize;
while kept.len() < cap {
let mut dealt = false;
for (_, group) in &by_kind {
if kept.len() == cap {
break;
}
if round < group.len() {
kept.push(group[round].clone());
dealt = true;
}
}
if !dealt {
break;
}
round += 1;
}
let omitted: Vec<OmittedEdges> = by_kind
.iter()
.filter_map(|(kind, group)| {
let carried = kept.iter().filter(|e| e.kind == *kind).count();
(group.len() > carried).then(|| OmittedEdges {
kind: kind.clone(),
omitted: group.len() - carried,
})
})
.collect();
sort_refs(&mut kept);
BoundedEdges {
total,
truncated: true,
omitted,
edges: kept,
}
}
pub fn tool_context(store: &Store, key: &str) -> Result<Option<ToolContext>, StoreError> {
let Some(bundle) = build_context(store, key)? else {
return Ok(None);
};
let outgoing = bound_edges(bundle.outgoing, TOOL_CONTEXT_EDGE_CAP);
let incoming = bound_edges(bundle.incoming, TOOL_CONTEXT_EDGE_CAP);
Ok(Some(ToolContext {
schema: bundle.schema,
fingerprint: bundle.fingerprint,
edge_cap: TOOL_CONTEXT_EDGE_CAP,
truncated: outgoing.truncated || incoming.truncated,
node: bundle.node,
meta: bundle.meta,
outgoing,
incoming,
}))
}
pub fn refresh_contexts(store: &Store) -> Result<ContextRefresh, StoreError> {
let mut out = ContextRefresh {
rebuilt: 0,
reused: 0,
pruned: 0,
};
for key in store.context_cache_keys()? {
let Some(node) = store.get_node(&key)? else {
store.context_cache_delete(&key)?;
out.pruned += 1;
continue;
};
let fingerprint = compute_fingerprint(store, &node)?;
let fresh = store
.context_cache_fingerprint(&key)?
.is_some_and(|fp| fp == fingerprint);
if fresh {
out.reused += 1;
} else {
context(store, &key)?; out.rebuilt += 1;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::{
TOOL_CONTEXT_EDGE_CAP, build_context, context, dependents, refresh_contexts, tool_context,
};
use crate::{Edge, EdgeKind, FactSet, Node, NodeKind, Store};
fn seeded() -> Store {
let mut store = Store::open_in_memory().expect("store");
let mut caller = Node::new("sym:rust:a.rs#caller", NodeKind::Fn, "caller");
caller.blob_hash = Some("BLOB_A".to_owned());
let mut target = Node::new("sym:rust:a.rs#callee", NodeKind::Fn, "callee");
target.blob_hash = Some("BLOB_A".to_owned());
let mut doc = Node::new("file:docs/g.md", NodeKind::Doc, "g.md");
doc.blob_hash = Some("BLOB_DOC".to_owned());
let facts = FactSet::new()
.with_node(caller)
.with_node(target)
.with_node(doc)
.with_edge(Edge::derived(
"sym:rust:a.rs#caller",
"sym:rust:a.rs#callee",
EdgeKind::Calls,
))
.with_edge(Edge::authored(
"file:docs/g.md",
"sym:rust:a.rs#caller",
EdgeKind::References,
));
store.apply_factset(&facts).expect("apply");
store
}
fn wide(count: usize) -> Store {
let mut store = Store::open_in_memory().expect("store");
let mut facts =
FactSet::new().with_node(Node::new("file:wide.rs", NodeKind::File, "wide.rs"));
for (kind, n) in [
(EdgeKind::Defines, count),
(EdgeKind::Imports, 3),
(EdgeKind::References, 2),
] {
for i in 0..n {
let key = format!("sym:rust:wide.rs#{}{i:04}", kind.as_str());
facts = facts
.with_node(Node::new(key.clone(), NodeKind::Fn, "s"))
.with_edge(Edge::derived("file:wide.rs", key, kind.clone()));
}
}
store.apply_factset(&facts).expect("apply");
store
}
#[test]
fn a_small_bundle_is_carried_whole_and_says_it_was_not_truncated() {
let store = seeded();
let out = tool_context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert!(!out.truncated);
assert_eq!(out.edge_cap, TOOL_CONTEXT_EDGE_CAP);
assert_eq!(out.outgoing.total, 1);
assert!(!out.outgoing.truncated);
assert!(out.outgoing.omitted.is_empty());
assert_eq!(out.outgoing.edges.len(), 1);
assert_eq!(out.incoming.total, 1);
assert_eq!(out.incoming.edges.len(), 1);
let whole = build_context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert_eq!(out.outgoing.edges, whole.outgoing);
assert_eq!(out.incoming.edges, whole.incoming);
assert_eq!(out.fingerprint, whole.fingerprint);
}
#[test]
fn a_truncated_bundle_reports_the_cap_the_totals_and_what_it_dropped() {
let store = wide(200);
let out = tool_context(&store, "file:wide.rs")
.expect("ctx")
.expect("present");
assert!(out.truncated, "205 edges must not fit under the cap");
assert_eq!(out.outgoing.total, 205, "the count before the cap");
assert!(out.outgoing.truncated);
assert_eq!(out.outgoing.edges.len(), TOOL_CONTEXT_EDGE_CAP);
let omitted: usize = out.outgoing.omitted.iter().map(|o| o.omitted).sum();
assert_eq!(omitted + out.outgoing.edges.len(), out.outgoing.total);
assert_eq!(
out.outgoing
.omitted
.iter()
.find(|o| o.kind == "defines")
.map(|o| o.omitted),
Some(205usize.saturating_sub(TOOL_CONTEXT_EDGE_CAP)),
"{:?}",
out.outgoing.omitted
);
assert_eq!(out.incoming.total, 0);
assert!(!out.incoming.truncated);
}
#[test]
fn truncation_keeps_every_edge_kind_rather_than_the_first_fifty() {
let store = wide(200);
let out = tool_context(&store, "file:wide.rs")
.expect("ctx")
.expect("present");
for kind in ["defines", "imports", "references"] {
assert!(
out.outgoing.edges.iter().any(|e| e.kind == kind),
"`{kind}` must survive truncation: {:?}",
out.outgoing
.edges
.iter()
.map(|e| &e.kind)
.collect::<std::collections::BTreeSet<_>>()
);
}
assert!(
!out.outgoing
.omitted
.iter()
.any(|o| o.kind == "imports" || o.kind == "references"),
"{:?}",
out.outgoing.omitted
);
let mut sorted = out.outgoing.edges.clone();
super::sort_refs(&mut sorted);
assert_eq!(sorted, out.outgoing.edges);
}
#[test]
fn tool_context_never_writes_to_the_cache() {
let store = seeded();
assert!(store.context_cache_keys().unwrap().is_empty());
tool_context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert!(
store.context_cache_keys().unwrap().is_empty(),
"a tool read must not populate the cache",
);
store
.context_cache_put("sym:rust:a.rs#ghost", "stale-fingerprint", "{}")
.expect("put");
let missing = tool_context(&store, "sym:rust:a.rs#ghost").expect("ctx");
assert!(missing.is_none(), "no such node");
assert_eq!(
store.context_cache_keys().unwrap(),
vec!["sym:rust:a.rs#ghost".to_owned()],
"a missing node must not prune its cache entry",
);
}
#[test]
fn context_is_cached_then_served_from_cache() {
let store = seeded();
let first = context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert_eq!(first.node.key, "sym:rust:a.rs#caller");
assert_eq!(first.outgoing.len(), 1, "calls callee");
assert_eq!(first.incoming.len(), 1, "referenced by doc");
assert_eq!(
store
.context_cache_get("sym:rust:a.rs#caller")
.expect("get")
.expect("cached")
.0,
first.fingerprint,
);
let second = context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert_eq!(first, second);
}
#[test]
fn changing_a_dependency_invalidates_dependent_context() {
let store = seeded();
let before = refresh_first_read(&store);
assert_eq!(before.rebuilt, 0, "warming reads are misses, not refreshes");
let caller_fp_before = context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present")
.fingerprint;
let mut target = store
.get_node("sym:rust:a.rs#callee")
.expect("get")
.expect("present");
target.blob_hash = Some("BLOB_A2".to_owned());
store.upsert_node(&target).expect("upsert");
let deps = dependents(&store, &["sym:rust:a.rs#callee".to_owned()]).expect("deps");
assert!(deps.contains("sym:rust:a.rs#caller"));
let caller_fp_after = context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present")
.fingerprint;
assert_ne!(
caller_fp_before, caller_fp_after,
"dependent context must be invalidated when a dependency changes",
);
let report = refresh_contexts(&store).expect("refresh");
assert!(report.rebuilt >= 1, "affected contexts rebuilt");
assert_eq!(report.pruned, 0);
}
fn refresh_first_read(store: &Store) -> super::ContextRefresh {
for key in store.all_keys().expect("keys") {
context(store, &key).expect("ctx");
}
refresh_contexts(store).expect("refresh")
}
#[test]
fn deleted_node_context_is_pruned() {
let mut store = seeded();
context(&store, "file:docs/g.md")
.expect("ctx")
.expect("present");
assert!(
store
.context_cache_get("file:docs/g.md")
.expect("get")
.is_some()
);
let mut caller = Node::new("sym:rust:a.rs#caller", NodeKind::Fn, "caller");
caller.blob_hash = Some("BLOB_A".to_owned());
store
.rebuild(&FactSet::new().with_node(caller), None)
.expect("rebuild");
assert!(context(&store, "file:docs/g.md").expect("ctx").is_none());
assert!(
store
.context_cache_get("file:docs/g.md")
.expect("get")
.is_none()
);
}
#[test]
fn build_context_matches_cached() {
let store = seeded();
let built = build_context(&store, "sym:rust:a.rs#caller")
.expect("build")
.expect("present");
let cached = context(&store, "sym:rust:a.rs#caller")
.expect("ctx")
.expect("present");
assert_eq!(built, cached);
assert!(
build_context(&store, "sym:rust:a.rs#ghost")
.expect("b")
.is_none()
);
}
}