From a316590adc05a167e7398b787ddcf242b0e58132 Mon Sep 17 00:00:00 2001 From: Safi Date: Sat, 16 May 2026 16:20:20 +0100 Subject: [PATCH] Fix query seed scoring: IDF weighting, dynamic K seeds, actionable truncation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- graphify/serve.py | 68 +++++++++++++++++++++++--- tests/test_serve.py | 116 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 178 insertions(+), 6 deletions(-) diff --git a/graphify/serve.py b/graphify/serve.py index 8565e51..3f5531f 100644 --- a/graphify/serve.py +++ b/graphify/serve.py @@ -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) diff --git a/tests/test_serve.py b/tests/test_serve.py index 67b097a..e0298a5 100644 --- a/tests/test_serve.py +++ b/tests/test_serve.py @@ -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