Fix query seed scoring: IDF weighting, dynamic K seeds, actionable truncation

Common terms like 'error'/'exception' were stealing BFS seed slots from
rare identifiers like 'FooBarService', burning the token budget on noise.

- _compute_idf: weights query terms by inverse document frequency, cached
  on G.graph so cost is paid once per graph load not per query
- _score_nodes: multiplies each tier bonus by IDF weight
- _pick_seeds: replaces fixed top-3 with gap-ratio selection — stops adding
  seeds when score drops below 20% of the top match
- _subgraph_to_text: truncation hint now tells Claude to narrow with
  context_filter or use get_node instead of just saying 'truncated'

Fixes #897

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Safi
2026-05-16 16:20:20 +01:00
co-authored by Claude Sonnet 4.6
parent b1ade00ece
commit a316590adc
2 changed files with 178 additions and 6 deletions
+62 -6
View File
@@ -1,6 +1,7 @@
# MCP stdio server - exposes graph query tools to Claude and other agents
from __future__ import annotations
import json
import math
import sys
from pathlib import Path
import networkx as nx
@@ -55,30 +56,76 @@ _SUBSTRING_MATCH_BONUS = 1.0
_SOURCE_MATCH_BONUS = 0.5
def _compute_idf(G: nx.Graph, terms: list[str]) -> dict[str, float]:
"""IDF weights for query terms, cached in G.graph['_idf_cache'].
Common terms like 'error' or 'exception' that match hundreds of nodes get
low weights; rare identifiers like 'FooBarService' get high weights.
Cache is stored on the graph object itself so it auto-invalidates when
_maybe_reload() replaces G with a new object.
"""
cache: dict[str, float] = G.graph.setdefault("_idf_cache", {})
N = G.number_of_nodes() or 1
uncached = [t for t in terms if t not in cache]
if uncached:
df: dict[str, int] = {t: 0 for t in uncached}
for _, data in G.nodes(data=True):
norm_label = (
data.get("norm_label") or _strip_diacritics(data.get("label") or "")
).lower()
for t in uncached:
if t in norm_label:
df[t] += 1
for t in uncached:
cache[t] = math.log(1 + N / (1 + df[t]))
return {t: cache.get(t, math.log(1 + N)) for t in terms}
def _score_nodes(G: nx.Graph, terms: list[str]) -> list[tuple[float, str]]:
scored = []
norm_terms = [_strip_diacritics(t).lower() for t in terms]
idf = _compute_idf(G, norm_terms)
for nid, data in G.nodes(data=True):
norm_label = data.get("norm_label") or _strip_diacritics(data.get("label") or "").lower()
bare_label = norm_label.rstrip("()")
source = (data.get("source_file") or "").lower()
score = 0.0
for t in norm_terms:
w = idf.get(t, 1.0)
# Three-tier precedence: exact > prefix > substring (take the
# strongest tier per term so a single term cannot double-count).
if t == norm_label or t == bare_label:
score += _EXACT_MATCH_BONUS
score += _EXACT_MATCH_BONUS * w
elif norm_label.startswith(t) or bare_label.startswith(t):
score += _PREFIX_MATCH_BONUS
score += _PREFIX_MATCH_BONUS * w
elif t in norm_label:
score += _SUBSTRING_MATCH_BONUS
score += _SUBSTRING_MATCH_BONUS * w
if t in source:
score += _SOURCE_MATCH_BONUS
score += _SOURCE_MATCH_BONUS * w
if score > 0:
scored.append((score, nid))
return sorted(scored, reverse=True)
def _pick_seeds(scored: list[tuple[float, str]], max_k: int = 3, gap_ratio: float = 0.2) -> list[str]:
"""Select BFS seed nodes, stopping when score drops too far below the top.
Prevents high-frequency noise terms (error, exception) from stealing seed
slots from a dominant identifier match. When FooBarService scores 1000 and
error nodes score 1.0, only FooBarService is seeded — the score gap is 99.9%
which is well above the 20% threshold that would allow additional seeds.
"""
if not scored:
return []
top_score = scored[0][0]
seeds = []
for score, nid in scored[:max_k]:
if seeds and score < top_score * gap_ratio:
break
seeds.append(nid)
return seeds
_CONTEXT_HINTS: tuple[tuple[str, tuple[str, ...]], ...] = (
("call", ("call", "calls", "called", "invoke", "invokes", "invoked")),
("import", ("import", "imports", "imported", "module", "modules")),
@@ -237,7 +284,16 @@ def _subgraph_to_text(G: nx.Graph, nodes: set[str], edges: list[tuple], token_bu
lines.append(line)
output = "\n".join(lines)
if len(output) > char_budget:
output = output[:char_budget] + f"\n... (truncated to ~{token_budget} token budget)"
cut_at = output[:char_budget].rfind("\n")
cut_at = cut_at if cut_at > 0 else char_budget
total_nodes = sum(1 for l in lines if l.startswith("NODE "))
shown_nodes = output[:cut_at].count("\nNODE ") + (1 if output.startswith("NODE ") else 0)
cut_count = total_nodes - shown_nodes
output = (
output[:cut_at]
+ f"\n... (truncated — {cut_count} more nodes cut by ~{token_budget}-token budget."
f" Narrow with context_filter=['call'] or use get_node for a specific symbol)"
)
return output
@@ -252,7 +308,7 @@ def _query_graph_text(
) -> str:
terms = [t.lower() for t in question.split() if len(t) > 2]
scored = _score_nodes(G, terms)
start_nodes = [nid for _, nid in scored[:3]]
start_nodes = _pick_seeds(scored)
if not start_nodes:
return "No matching nodes found."
resolved_filters, filter_source = _resolve_context_filters(question, context_filters)
+116
View File
@@ -7,6 +7,8 @@ from networkx.readwrite import json_graph
from graphify.serve import (
_communities_from_graph,
_score_nodes,
_compute_idf,
_pick_seeds,
_bfs,
_dfs,
_filter_graph_by_context,
@@ -250,3 +252,117 @@ def test_load_graph_cache_key_changes_with_content(tmp_path):
key2 = (s2.st_mtime_ns, s2.st_size)
assert key1 != key2, "stat key must change when file content changes"
# --- IDF weighting tests (#897) ---
def _make_noisy_graph() -> nx.Graph:
"""20 error-handler nodes + 1 rare identifier: FooBarService."""
G = nx.Graph()
for i in range(20):
G.add_node(f"err{i}", label=f"error_handler_{i}", source_file=f"err{i}.py", community=0)
if i > 0:
G.add_edge(f"err{i-1}", f"err{i}", relation="calls", confidence="EXTRACTED")
G.add_node("fbs", label="FooBarService", source_file="service.py", community=1)
G.add_node("fbs_dep", label="ServiceClient", source_file="client.py", community=1)
G.add_edge("fbs", "fbs_dep", relation="uses", confidence="EXTRACTED")
return G
def test_idf_downweights_common_terms():
"""'error' matches 20 nodes, 'foobarservice' matches 1 — IDF should make
FooBarService rank first despite error's higher raw frequency."""
G = _make_noisy_graph()
scored = _score_nodes(G, ["foobarservice", "error"])
assert scored, "should have results"
assert scored[0][1] == "fbs", (
f"FooBarService should rank first, got {scored[0][1]}"
)
def test_idf_cached_on_graph():
"""IDF results are stored in G.graph so repeated queries don't recompute."""
G = _make_graph()
_score_nodes(G, ["extract"])
assert "_idf_cache" in G.graph
assert "extract" in G.graph["_idf_cache"]
def test_idf_new_graph_starts_fresh():
"""Two separate graph instances must not share an IDF cache."""
G1 = _make_graph()
G2 = _make_graph()
_score_nodes(G1, ["extract"])
assert "_idf_cache" not in G2.graph
def test_idf_rare_term_gets_high_weight():
"""A term matching only 1 of N nodes should get IDF > 1."""
import math
G = _make_graph() # 5 nodes
idf = _compute_idf(G, ["extract"])
# extract matches only n1: IDF = log(1 + 5/2) ≈ 1.25
assert idf["extract"] > 1.0
def test_idf_common_term_gets_low_weight():
"""A term matching most nodes should get IDF < 1."""
import math
G = nx.Graph()
# 'handle' in every node label
for i in range(20):
G.add_node(f"n{i}", label=f"handle_{i}", source_file=f"f{i}.py")
idf = _compute_idf(G, ["handle"])
assert idf["handle"] < 1.0
# --- _pick_seeds tests (#897) ---
def test_pick_seeds_dominant_identifier_gives_one_seed():
"""FooBarService at 1000 vs error nodes at 1.0 → only 1 seed chosen."""
scored = [(1000.0, "fbs"), (1.0, "err1"), (0.9, "err2")]
seeds = _pick_seeds(scored)
assert seeds == ["fbs"]
def test_pick_seeds_close_scores_keeps_multiple():
"""When all scores are within 20% of the top, keep up to 3 seeds."""
scored = [(10.0, "a"), (9.0, "b"), (8.5, "c")]
seeds = _pick_seeds(scored)
assert len(seeds) == 3
def test_pick_seeds_empty():
assert _pick_seeds([]) == []
def test_pick_seeds_single():
assert _pick_seeds([(5.0, "x")]) == ["x"]
def test_pick_seeds_respects_max_k():
"""Never return more than max_k seeds even when all scores are close."""
scored = [(10.0, f"n{i}") for i in range(10)]
seeds = _pick_seeds(scored, max_k=3)
assert len(seeds) == 3
# --- actionable truncation hint (#897) ---
def test_subgraph_to_text_truncation_hint_is_actionable():
"""Truncation message must tell Claude what to do, not just say truncated."""
G = _make_graph()
text = _subgraph_to_text(G, {"n1", "n2", "n3", "n4"}, [("n1", "n2")], token_budget=1)
assert "truncated" in text
assert "get_node" in text or "context_filter" in text
# --- integration: identifier + noise query seeds from identifier (#897) ---
def test_query_seeds_from_identifier_not_noise():
"""'FooBarService error handling' should expand from FooBarService,
not from error-handler nodes, so ServiceClient appears in results."""
G = _make_noisy_graph()
text = _query_graph_text(G, "FooBarService error handling", mode="bfs", depth=2)
assert "FooBarService" in text
assert "ServiceClient" in text