feature/treesitter-tag-queries #2

Open
tobi wants to merge 78 commits from feature/treesitter-tag-queries into main
2 changed files with 109 additions and 4 deletions
Showing only changes of commit 9be44c2964 - Show all commits

View file

@ -224,3 +224,98 @@ def grammar_noise_ratio(node, noise_tokens=None):
return grammar_noise_ratio(node.child, noise_tokens)
return 0, 0
def grammar_quality_score(node):
"""Score grammar quality (0.0 = useless, 1.0 = excellent).
Criteria:
- Has ordering (not a pure bag): +0.3
- Has multiple alternation groups: +0.2 per group (max 0.4)
- Has enough symbols (>=3 domain tokens): +0.2
- Not too short (>=2 concat parts): +0.1
"""
from .grammar import Concat, Alt, Optional, Plus, Star
if node is None or isinstance(node, (Epsilon, Empty)):
return 0.0
score = 0.0
# Check for ordering (Concat with multiple parts)
if isinstance(node, Concat) and len(node.parts) >= 2:
score += 0.3
# Check for alternation groups
n_groups = _count_alt_groups(node)
score += min(0.4, n_groups * 0.2)
# Check symbol count
n_symbols = _count_symbols(node)
if n_symbols >= 3:
score += 0.2
# Check concat depth
n_concat = _count_concat_parts(node)
if n_concat >= 2:
score += 0.1
return min(1.0, score)
def _count_alt_groups(node):
"""Count alternation groups in AST."""
from .grammar import Concat, Alt, Optional, Plus, Star
if node is None or isinstance(node, (Epsilon, Empty, Symbol)):
return 0
if isinstance(node, Alt):
return 1 + sum(_count_alt_groups(p) for p in node.parts)
if isinstance(node, Concat):
return sum(_count_alt_groups(p) for p in node.parts)
if isinstance(node, (Plus, Optional, Star)):
return _count_alt_groups(node.child)
return 0
def _count_symbols(node):
"""Count total symbols in AST."""
from .grammar import Concat, Alt, Optional, Plus, Star
if node is None or isinstance(node, (Epsilon, Empty)):
return 0
if isinstance(node, Symbol):
return 1
if isinstance(node, Concat):
return sum(_count_symbols(p) for p in node.parts)
if isinstance(node, Alt):
return sum(_count_symbols(p) for p in node.parts)
if isinstance(node, (Plus, Optional, Star)):
return _count_symbols(node.child)
return 0
def _count_concat_parts(node):
"""Count top-level concat parts."""
from .grammar import Concat, Optional, Plus, Star
if isinstance(node, Concat):
return len(node.parts)
if isinstance(node, (Plus, Optional, Star)):
return _count_concat_parts(node.child)
return 1
def is_useful_grammar(node, min_quality=0.3):
"""Check if grammar is useful for LLM constraining.
A grammar is useful if:
1. Not empty after noise filtering
2. Has some structure (not a pure bag)
3. Has enough symbols to be constraining
"""
if node is None or isinstance(node, (Epsilon, Empty)):
return False
quality = grammar_quality_score(node)
return quality >= min_quality

View file

@ -20,7 +20,7 @@ import pathspec
from .code import preprocess_by_method, extract_arg_info, _summarize_arg_info
from bex.ensemble import infer_ensemble
from bex.gbnf import grammar_structure_score, to_gbnf, filter_noise, grammar_noise_ratio
from bex.gbnf import grammar_structure_score, to_gbnf, filter_noise, grammar_noise_ratio, is_useful_grammar, grammar_quality_score
from bex.grammar import Empty
from bex.distributional import distributional_split
from bex.decompose import decompose_with_coverage, get_decomposition_stats
@ -878,7 +878,7 @@ def analyze_directory(
return results
def _build_json_output(results):
def _build_json_output(results, min_quality=0.3):
"""Convert results dict to a compact JSON structure for prompt injection."""
output = []
for ext, clusters in results.items():
@ -896,14 +896,19 @@ def _build_json_output(results):
filtered_grammar = filter_noise(grammar)
if filtered_grammar and not isinstance(filtered_grammar, Empty):
n_noise, n_total = grammar_noise_ratio(grammar)
quality = grammar_quality_score(filtered_grammar)
entry["grammar"] = to_gbnf(filtered_grammar)
entry["grammar_clean"] = to_gbnf(filtered_grammar)
entry["noise_ratio"] = round(n_noise / n_total, 2) if n_total > 0 else 1.0
entry["symbols_before"] = n_total
entry["symbols_after"] = n_total - n_noise
entry["quality"] = round(quality, 2)
entry["useful"] = quality >= min_quality
else:
entry["grammar"] = to_gbnf(grammar)
entry["noise_ratio"] = 1.0
entry["quality"] = 0.0
entry["useful"] = False
entry["algorithm"] = result["best"]["algorithm"]
entry["mdl_score"] = round(result['best']['mdl_score'], 1)
entry["imports"] = meta.get("imports", [])
@ -914,11 +919,11 @@ def _build_json_output(results):
return json.dumps(output, indent=2)
def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filter_grammar_noise=True):
def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filter_grammar_noise=True, min_quality=0.3):
"""Build YAML output grouped by top-level module, sorted by MDL.
Filters out (other), no-grammar groups, groups above max_mdl,
and groups below min_structure.
groups below min_structure, and low-quality grammars.
Returns YAML string.
"""
import yaml
@ -952,6 +957,10 @@ def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filt
else:
continue # Skip grammars that become empty after filtering
# Apply quality gate
if not is_useful_grammar(grammar, min_quality):
continue
# Extract top-level module from package path
parts = label.replace(os.sep, "/").split("/")
module = parts[0] if len(parts) > 1 else "(root)"
@ -964,6 +973,7 @@ def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filt
"score": round(best.get("mdl_score", 0), 3),
"algorithm": best["algorithm"],
"mdl": round(best["mdl_score"], 1),
"quality": round(grammar_quality_score(grammar), 2),
}
modules.setdefault(module, []).append(entry)