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:
co-authored by
Claude Sonnet 4.6
parent
b1ade00ece
commit
a316590adc
+62
-6
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user