Skip to content

Commit 6d9bb5b

Browse files
Bound graph retrieval by query relevance
1 parent e3f8858 commit 6d9bb5b

4 files changed

Lines changed: 220 additions & 14 deletions

File tree

engraphis/core/recall.py

Lines changed: 66 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -364,23 +364,39 @@ def _code_arm(
364364
if not symbols:
365365
return {}
366366

367+
aliases: dict[str, str] = {}
368+
for symbol_id, symbol in symbols.items():
369+
for key in ("id", "name", "fqname"):
370+
value = str(symbol.get(key) or "")
371+
if value:
372+
aliases[value] = symbol_id
367373
# Expand one stored code edge to capture callers/callees, bounded by a
368-
# multiple of candidate_k. Both symbol ids and parser-emitted names are
369-
# accepted because language backends use both representations.
374+
# multiple of candidate_k. Query only edges incident to matched aliases
375+
# before applying that cap, so later files cannot be hidden by a global prefix.
376+
edge_kwargs = {
377+
"limit": max(100, min(2000, candidate_k * 20)),
378+
"layers": flt.graph_layers,
379+
}
380+
# ``endpoints`` is a v2 Store optimization. Preserve compatibility with
381+
# external code stores that have not added the optional filter yet.
382+
try:
383+
edge_parameters = inspect.signature(self.store.list_code_edges).parameters.values()
384+
supports_endpoints = any(
385+
parameter.name == "endpoints"
386+
or parameter.kind == inspect.Parameter.VAR_KEYWORD
387+
for parameter in edge_parameters
388+
)
389+
except (TypeError, ValueError):
390+
supports_endpoints = False
391+
if supports_endpoints:
392+
edge_kwargs["endpoints"] = list(aliases)
370393
code_edges = _call_temporal_store(
371394
self.store.list_code_edges,
372395
flt,
373396
flt.repo_id,
374-
limit=max(100, min(2000, candidate_k * 20)),
375-
layers=flt.graph_layers,
376397
requested_historical=historical,
398+
**edge_kwargs,
377399
)
378-
aliases: dict[str, str] = {}
379-
for symbol_id, symbol in symbols.items():
380-
for key in ("id", "name", "fqname"):
381-
value = str(symbol.get(key) or "")
382-
if value:
383-
aliases[value] = symbol_id
384400
related_names: dict[str, float] = {}
385401
for edge in code_edges:
386402
src, dst = str(edge.get("src") or ""), str(edge.get("dst") or "")
@@ -493,8 +509,31 @@ def connect(a: str, b: str, w: float) -> None:
493509
adj.setdefault(a, []).append((b, w))
494510
adj.setdefault(b, []).append((a, w))
495511

496-
edges = self.store.edges_in_scope(flt, at=now, limit=4000)
497-
for e in edges:
512+
# Build a bounded edge set outward from the query entities. A global
513+
# ULID-ordered cap would let old unrelated edges crowd out a new relation
514+
# required by this query before PPR sees it.
515+
edge_cap = 4000
516+
edges_by_id = {}
517+
frontier = set(seeds)
518+
expanded: set[str] = set()
519+
while frontier and len(edges_by_id) < edge_cap:
520+
batch = sorted(frontier - expanded)[:400]
521+
if not batch:
522+
break
523+
frontier.difference_update(batch)
524+
expanded.update(batch)
525+
next_frontier: set[str] = set()
526+
for edge in self.store.neighbors(
527+
batch, at=now, layers=flt.graph_layers, flt=flt,
528+
limit=edge_cap - len(edges_by_id)):
529+
if edge.id in edges_by_id:
530+
continue
531+
edges_by_id[edge.id] = edge
532+
next_frontier.update((edge.src, edge.dst))
533+
if len(edges_by_id) >= edge_cap:
534+
break
535+
frontier.update(next_frontier - expanded)
536+
for e in edges_by_id.values():
498537
connect(ent(e.src), ent(e.dst), max(float(e.weight or 1.0), 1e-6))
499538

500539
incidence = self.store.list_memory_entities(flt, limit=12_000)
@@ -505,9 +544,23 @@ def connect(a: str, b: str, w: float) -> None:
505544
# retrieval arms so PPR can traverse that edge without widening scope. Keep
506545
# the incidence frontier as well when independent caps choose a different
507546
# subset of the scoped memory universe.
508-
memory_ids = sorted({
547+
incidence_memory_ids = {
509548
str(row.get("memory_id") or "")
510549
for row in incidence if row.get("memory_id")
550+
}
551+
frontier_links = self.store.links_touching(
552+
sorted(incidence_memory_ids),
553+
layers=flt.graph_layers,
554+
flt=flt,
555+
limit=20_000,
556+
)
557+
# Expand from the entity-incidence frontier before adding the bounded newest
558+
# memory window. An older unmentioned endpoint can then participate in PPR
559+
# through its visible link instead of being silently dropped by that window.
560+
memory_ids = sorted(incidence_memory_ids | {
561+
endpoint
562+
for link in frontier_links
563+
for endpoint in (link["a"], link["b"])
511564
} | {
512565
memory.id for memory in self.store.list_memories(flt, limit=12_000)
513566
})

engraphis/core/store.py

Lines changed: 72 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3168,9 +3168,68 @@ def links_among(self, ids: list[str], *,
31683168
break
31693169
return rows
31703170

3171+
def links_touching(self, ids: list[str], *,
3172+
layers: Optional[list[GraphLayer]] = None,
3173+
flt: Optional[SearchFilter] = None,
3174+
include_invalid: bool = False,
3175+
limit: Optional[int] = None) -> list[dict]:
3176+
"""Return visible links with at least one endpoint in ``ids``.
3177+
3178+
This bounded frontier expansion is distinct from :meth:`links_among`: graph
3179+
recall uses it to retain an unmentioned endpoint linked to an entity-attached
3180+
memory, without first materializing every memory in a large scope.
3181+
"""
3182+
if not ids:
3183+
return []
3184+
if layers is not None and not layers:
3185+
return []
3186+
row_cap = None if limit is None else max(0, int(limit))
3187+
if row_cap == 0:
3188+
return []
3189+
ordered_ids = sorted(set(ids))
3190+
visibility_sql, visibility_params = _temporal_visibility_sql("", flt)
3191+
rows: list[dict] = []
3192+
seen: set[tuple] = set()
3193+
# Each id appears once for each endpoint predicate; reserve parameters for
3194+
# time/layer filters so this remains under SQLite's portable bind limit.
3195+
chunk_size = max(1, (IN_CLAUSE_CHUNK - 16) // 2)
3196+
for start in range(0, len(ordered_ids), chunk_size):
3197+
if row_cap is not None and len(rows) >= row_cap:
3198+
break
3199+
chunk = ordered_ids[start:start + chunk_size]
3200+
marks = ",".join("?" for _ in chunk)
3201+
sql = (
3202+
"SELECT a, b, relation, layer, reason, created_at, valid_from, valid_to, "
3203+
"valid_to_recorded_at, ingested_at, expired_at FROM mem_links "
3204+
f"WHERE (a IN ({marks}) OR b IN ({marks}))"
3205+
)
3206+
params: list[Any] = [*chunk, *chunk]
3207+
if not include_invalid:
3208+
sql += f" AND {visibility_sql}"
3209+
params.extend(visibility_params)
3210+
if layers is not None:
3211+
layer_marks = ",".join("?" for _ in layers)
3212+
sql += f" AND layer IN ({layer_marks})"
3213+
params.extend(_enum(layer) for layer in layers)
3214+
sql += " ORDER BY a, b, relation, valid_from, ingested_at"
3215+
for row in self.conn.execute(sql, params).fetchall():
3216+
item = dict(row)
3217+
key = (
3218+
item["a"], item["b"], item["relation"], item["layer"],
3219+
item["valid_from"], item["valid_to"], item["ingested_at"],
3220+
)
3221+
if key in seen:
3222+
continue
3223+
seen.add(key)
3224+
rows.append(item)
3225+
if row_cap is not None and len(rows) >= row_cap:
3226+
break
3227+
return rows
3228+
31713229
def neighbors(self, node_ids: list[str], *, at: Optional[float] = None,
31723230
layers: Optional[list[GraphLayer]] = None,
3173-
flt: Optional[SearchFilter] = None) -> list[Edge]:
3231+
flt: Optional[SearchFilter] = None,
3232+
limit: Optional[int] = None) -> list[Edge]:
31743233
if not node_ids:
31753234
return []
31763235
valid_at, known_at = _temporal_anchors(flt, valid_at=at)
@@ -3213,6 +3272,10 @@ def neighbors(self, node_ids: list[str], *, at: Optional[float] = None,
32133272
else:
32143273
sql += " AND repo_id=?"
32153274
params.append(flt.repo_id)
3275+
sql += " ORDER BY id"
3276+
if limit is not None:
3277+
sql += " LIMIT ?"
3278+
params.append(max(0, int(limit)))
32163279
rows = self.conn.execute(sql, params).fetchall()
32173280
return [_row_to_edge(r) for r in rows]
32183281

@@ -3409,6 +3472,7 @@ def list_symbols_page(self, repo_id: str, *,
34093472

34103473
def list_code_edges(self, repo_id: str, *, limit: Optional[int] = None,
34113474
layers: Optional[list[GraphLayer]] = None,
3475+
endpoints: Optional[list[str]] = None,
34123476
flt: Optional[SearchFilter] = None) -> list[dict]:
34133477
temporal, params = _temporal_visibility_sql("", flt)
34143478
sql = "SELECT * FROM code_edges WHERE repo_id=? AND " + temporal
@@ -3419,6 +3483,13 @@ def list_code_edges(self, repo_id: str, *, limit: Optional[int] = None,
34193483
marks = ",".join("?" for _ in layers)
34203484
sql += f" AND layer IN ({marks})"
34213485
params.extend(_enum(layer) for layer in layers)
3486+
if endpoints is not None:
3487+
if not endpoints:
3488+
return []
3489+
marks = ",".join("?" for _ in endpoints)
3490+
sql += f" AND (src IN ({marks}) OR dst IN ({marks}))"
3491+
params.extend(endpoints)
3492+
params.extend(endpoints)
34223493
sql += " ORDER BY file, line, id"
34233494
if limit is not None:
34243495
sql += " LIMIT ?"

tests/test_core_store.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -491,6 +491,22 @@ def test_code_listing_helpers_honor_limit(store):
491491
assert store.list_code_files("repo_x", languages={"rust"}, limit=3) == []
492492

493493

494+
def test_code_edge_endpoint_filter_applies_before_limit(store):
495+
for index in range(4):
496+
store.add_code_edge(
497+
repo_id="repo_x", src=f"noise_{index}", dst=f"other_{index}",
498+
relation="calls", file="a_noise.py", line=index,
499+
)
500+
store.add_code_edge(
501+
repo_id="repo_x", src="target", dst="caller", relation="calls",
502+
file="z_target.py", line=1,
503+
)
504+
505+
edges = store.list_code_edges("repo_x", endpoints=["target"], limit=1)
506+
507+
assert [(edge["src"], edge["dst"]) for edge in edges] == [("target", "caller")]
508+
509+
494510
def test_memory_links_infer_and_filter_graph_layers(store):
495511
wid = store.get_or_create_workspace("w")
496512
a = store.add_memory(MemoryRecord(id="", content="cause", workspace_id=wid))

tests/test_recall.py

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -95,6 +95,37 @@ def test_graph_arm_pulls_related_via_entities():
9595
assert any("checkout" in c["content"].lower() for c in res.chunks)
9696

9797

98+
def test_graph_arm_selects_seed_frontier_edges_before_global_edge_cap(monkeypatch):
99+
from engraphis.core.interfaces import Edge, Node
100+
101+
store, emb, eng = _engine()
102+
wid = store.get_or_create_workspace("w")
103+
rid = store.get_or_create_repo(wid, "r")
104+
redis = store.upsert_entity(Node(
105+
id="", name="Redis", ntype="tech", workspace_id=wid, repo_id=rid))
106+
checkout = store.upsert_entity(Node(
107+
id="", name="checkout", ntype="module", workspace_id=wid, repo_id=rid))
108+
store.upsert_edge(Edge(
109+
id="", src=redis, dst=checkout, relation="used_by",
110+
workspace_id=wid, repo_id=rid))
111+
memory_id = _add(store, emb, wid, rid, "Checkout uses the Redis cache.")
112+
store.link_memory_entity(
113+
memory_id=memory_id, entity_id=checkout, workspace_id=wid, repo_id=rid,
114+
source_kind="test", confidence=1.0,
115+
)
116+
def no_global_edges(*_args, **_kwargs):
117+
raise AssertionError("PPR must traverse from query seeds, not global edges")
118+
119+
monkeypatch.setattr(store, "edges_in_scope", no_global_edges)
120+
121+
scores = eng._graph_arm_ppr(
122+
"How does Redis relate to the checkout service?",
123+
SearchFilter(workspace_id=wid, repo_id=rid), now=10**12,
124+
)
125+
126+
assert memory_id in scores
127+
128+
98129
def test_graph_arm_backfills_text_memory_when_its_entity_is_added_later():
99130
from engraphis.core.interfaces import Edge, Node
100131

@@ -169,6 +200,41 @@ def test_graph_arm_backfills_workspace_mentions_for_a_later_repo_entity():
169200
)
170201

171202

203+
def test_graph_arm_expands_an_older_unmentioned_link_endpoint_from_incidence(monkeypatch):
204+
from engraphis.core.interfaces import Node
205+
206+
store, emb, eng = _engine()
207+
wid = store.get_or_create_workspace("w")
208+
rid = store.get_or_create_repo(wid, "r")
209+
redis = store.upsert_entity(Node(
210+
id="", name="Redis", ntype="technology", workspace_id=wid, repo_id=rid,
211+
))
212+
attached = _add(store, emb, wid, rid, "The cache migration is attached evidence.")
213+
older_unmentioned = _add(store, emb, wid, rid, "The old rollout required a staged cutover.")
214+
store.link_memory_entity(
215+
memory_id=attached, entity_id=redis, workspace_id=wid, repo_id=rid,
216+
source_kind="test", confidence=1.0,
217+
)
218+
store.add_link(attached, older_unmentioned, relation="supports")
219+
220+
# Simulate a full scope whose bounded newest-memory window excludes the older
221+
# endpoint. The incidence frontier still contains ``attached``.
222+
monkeypatch.setattr(
223+
store,
224+
"list_memories",
225+
lambda *_args, **_kwargs: [
226+
MemoryRecord(id=f"mem_new_{i}", content="") for i in range(12_000)
227+
],
228+
)
229+
230+
scores = eng._graph_arm_ppr(
231+
"How does Redis relate to the rollout?",
232+
SearchFilter(workspace_id=wid, repo_id=rid), now=10**12,
233+
)
234+
235+
assert older_unmentioned in scores
236+
237+
172238
def test_entity_backfill_preserves_closed_workspace_memory_history():
173239
from engraphis.core.interfaces import Node
174240

0 commit comments

Comments
 (0)