feat: grammar index for runtime lookup
Add GrammarIndex class that loads grammars.yml and provides fast lookup by (package, context_symbol). Enables agents to get the right GBNF grammar at code generation time. API: load_grammar_index(project_root) → GrammarIndex idx.get(file_path, context_symbol=None) → gbnf_string idx.get_package(file_path) → [(symbol, grammar, score, methods)] 14 tests, all passing.
This commit is contained in:
parent
fe3261b9de
commit
7b54a5ef77
2 changed files with 301 additions and 0 deletions
139
bex/grammar_index.py
Normal file
139
bex/grammar_index.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
"""Grammar index for runtime lookup.
|
||||
|
||||
Loads grammars.yml and builds a lookup index mapping
|
||||
(package, context_symbol) → GBNF grammar string.
|
||||
|
||||
Usage:
|
||||
from bex.grammar_index import load_grammar_index
|
||||
|
||||
idx = load_grammar_index("/path/to/project")
|
||||
grammar = idx.get(file_path="src/services/Chat.kt", context_symbol="return")
|
||||
all_grammars = idx.get_package("src/services")
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
import yaml
|
||||
|
||||
|
||||
_LABEL_RE = re.compile(r"^(.+)\s+\[([^\]]+)\]$")
|
||||
|
||||
|
||||
class GrammarIndex:
|
||||
"""In-memory index of grammars keyed by (package, context_symbol)."""
|
||||
|
||||
def __init__(self, project_root, entries):
|
||||
"""
|
||||
Args:
|
||||
project_root: absolute path to the project root.
|
||||
entries: list of dicts with keys: package, grammar, score, methods, algorithm.
|
||||
"""
|
||||
self.project_root = project_root
|
||||
self._by_package = {} # package → [(symbol, grammar, score)]
|
||||
self._all = entries
|
||||
|
||||
for e in entries:
|
||||
label = e["package"]
|
||||
m = _LABEL_RE.match(label)
|
||||
if m:
|
||||
pkg = m.group(1).rstrip("/")
|
||||
symbol = m.group(2)
|
||||
else:
|
||||
pkg = label.rstrip("/")
|
||||
symbol = ""
|
||||
|
||||
self._by_package.setdefault(pkg, []).append((
|
||||
symbol,
|
||||
e["grammar"],
|
||||
e.get("score", 0),
|
||||
e.get("methods", 0),
|
||||
))
|
||||
|
||||
# Sort each package's entries by score descending (best first)
|
||||
for pkg in self._by_package:
|
||||
self._by_package[pkg].sort(key=lambda x: -x[2])
|
||||
|
||||
def get(self, file_path, context_symbol=None):
|
||||
"""Get the best grammar for a file, optionally filtered by context.
|
||||
|
||||
Args:
|
||||
file_path: path to the source file (absolute or relative to project_root).
|
||||
context_symbol: if provided, match the leaf grammar whose first symbol
|
||||
is this. If None, return the best grammar for the package.
|
||||
|
||||
Returns:
|
||||
GBNF grammar string, or None if no match.
|
||||
"""
|
||||
pkg = self._resolve_package(file_path)
|
||||
entries = self._by_package.get(pkg, [])
|
||||
if not entries:
|
||||
return None
|
||||
|
||||
if context_symbol:
|
||||
for sym, grammar, score, methods in entries:
|
||||
if sym == context_symbol:
|
||||
return grammar
|
||||
|
||||
# Fall back to best grammar for the package
|
||||
return entries[0][1] if entries else None
|
||||
|
||||
def get_package(self, file_path):
|
||||
"""Get all grammars for a file's package.
|
||||
|
||||
Returns:
|
||||
list of (context_symbol, grammar, score, methods) tuples,
|
||||
sorted by score descending. Empty list if no match.
|
||||
"""
|
||||
pkg = self._resolve_package(file_path)
|
||||
return list(self._by_package.get(pkg, []))
|
||||
|
||||
def get_all(self):
|
||||
"""Return all entries as a flat list."""
|
||||
return list(self._all)
|
||||
|
||||
def packages(self):
|
||||
"""Return sorted list of all indexed packages."""
|
||||
return sorted(self._by_package.keys())
|
||||
|
||||
def _resolve_package(self, file_path):
|
||||
"""Map a file path to its package."""
|
||||
# Make absolute if relative
|
||||
if not os.path.isabs(file_path):
|
||||
file_path = os.path.join(self.project_root, file_path)
|
||||
|
||||
# Get directory of file relative to project root
|
||||
try:
|
||||
rel = os.path.relpath(os.path.dirname(file_path), self.project_root)
|
||||
except ValueError:
|
||||
# Different drives on Windows
|
||||
return ""
|
||||
if rel == ".":
|
||||
return ""
|
||||
return rel
|
||||
|
||||
def __repr__(self):
|
||||
return f"GrammarIndex({self.project_root}, {len(self._all)} entries, {len(self._by_package)} packages)"
|
||||
|
||||
|
||||
def load_grammar_index(project_root):
|
||||
"""Load grammars.yml from {project_root}/.dervish/grammars.yml.
|
||||
|
||||
Returns GrammarIndex, or empty index if no file found.
|
||||
"""
|
||||
yml_path = os.path.join(project_root, ".dervish", "grammars.yml")
|
||||
if not os.path.exists(yml_path):
|
||||
return GrammarIndex(project_root, [])
|
||||
|
||||
with open(yml_path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
|
||||
entries = []
|
||||
if isinstance(data, dict):
|
||||
for module, items in data.items():
|
||||
if not isinstance(items, list):
|
||||
continue
|
||||
for item in items:
|
||||
if isinstance(item, dict) and "package" in item and "grammar" in item:
|
||||
entries.append(item)
|
||||
|
||||
return GrammarIndex(project_root, entries)
|
||||
162
tests/test_grammar_index.py
Normal file
162
tests/test_grammar_index.py
Normal file
|
|
@ -0,0 +1,162 @@
|
|||
"""Tests for bex.grammar_index — runtime grammar lookup."""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from bex.grammar_index import GrammarIndex, load_grammar_index
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_yml(tmp_path, entries):
|
||||
"""Write a grammars.yml with the given entries grouped by module."""
|
||||
dervish = tmp_path / ".dervish"
|
||||
dervish.mkdir()
|
||||
data = {"test_module": entries}
|
||||
with open(dervish / "grammars.yml", "w") as f:
|
||||
yaml.dump(data, f, default_flow_style=False)
|
||||
|
||||
|
||||
SAMPLE_ENTRIES = [
|
||||
{
|
||||
"package": "src/services [return]",
|
||||
"methods": 10,
|
||||
"grammar": "return.ok+?.error?",
|
||||
"score": 0.8,
|
||||
"algorithm": "CRX",
|
||||
"mdl": 12.0,
|
||||
},
|
||||
{
|
||||
"package": "src/services [if]",
|
||||
"methods": 5,
|
||||
"grammar": "if.cond+.then+.else?",
|
||||
"score": 0.6,
|
||||
"algorithm": "CRX",
|
||||
"mdl": 15.0,
|
||||
},
|
||||
{
|
||||
"package": "src/utils [return]",
|
||||
"methods": 8,
|
||||
"grammar": "return.value+",
|
||||
"score": 1.0,
|
||||
"algorithm": "CRX",
|
||||
"mdl": 5.0,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GrammarIndex construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGrammarIndex:
|
||||
|
||||
def test_parse_label_with_symbol(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
pkgs = idx.packages()
|
||||
assert "src/services" in pkgs
|
||||
assert "src/utils" in pkgs
|
||||
|
||||
def test_parse_label_without_symbol(self, tmp_path):
|
||||
entries = [{"package": "src/root", "grammar": "a.b", "score": 0.5}]
|
||||
idx = GrammarIndex(str(tmp_path), entries)
|
||||
assert "src/root" in idx.packages()
|
||||
|
||||
def test_sorted_by_score_desc(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
entries = idx.get_package("src/services")
|
||||
scores = [s for _, _, s, _ in entries]
|
||||
assert scores == sorted(scores, reverse=True)
|
||||
|
||||
def test_repr(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
assert "3 entries" in repr(idx)
|
||||
assert "2 packages" in repr(idx)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGet:
|
||||
|
||||
def test_get_best_no_context(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
g = idx.get("src/services/main.kt")
|
||||
assert g == "return.ok+?.error?"
|
||||
|
||||
def test_get_by_context_symbol(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
g = idx.get("src/services/main.kt", context_symbol="if")
|
||||
assert g == "if.cond+.then+.else?"
|
||||
|
||||
def test_get_missing_context_falls_back_to_best(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
g = idx.get("src/services/main.kt", context_symbol="nonexistent")
|
||||
assert g == "return.ok+?.error?"
|
||||
|
||||
def test_get_unknown_package_returns_none(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
g = idx.get("src/unknown/file.kt")
|
||||
assert g is None
|
||||
|
||||
def test_get_absolute_path(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
abs_path = str(tmp_path / "src" / "services" / "main.kt")
|
||||
g = idx.get(abs_path, context_symbol="return")
|
||||
assert g == "return.ok+?.error?"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_package()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGetPackage:
|
||||
|
||||
def test_returns_all_entries(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
entries = idx.get_package("src/services/main.kt")
|
||||
assert len(entries) == 2
|
||||
|
||||
def test_unknown_package_returns_empty(self, tmp_path):
|
||||
idx = GrammarIndex(str(tmp_path), SAMPLE_ENTRIES)
|
||||
entries = idx.get_package("src/unknown/file.kt")
|
||||
assert entries == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_grammar_index()
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestLoadGrammarIndex:
|
||||
|
||||
def test_load_existing(self, tmp_path):
|
||||
_make_yml(tmp_path, SAMPLE_ENTRIES)
|
||||
idx = load_grammar_index(str(tmp_path))
|
||||
assert len(idx.get_all()) == 3
|
||||
assert "src/services" in idx.packages()
|
||||
|
||||
def test_load_missing_file(self, tmp_path):
|
||||
idx = load_grammar_index(str(tmp_path))
|
||||
assert len(idx.get_all()) == 0
|
||||
assert idx.packages() == []
|
||||
|
||||
def test_roundtrip_persisted_ragsak(self):
|
||||
"""Integration test: grammars.yml generated by analyze_directory."""
|
||||
ragsak_yml = "/home/tobi/Desktop/kesai/RAGSAK/.dervish/grammars.yml"
|
||||
if not os.path.exists(ragsak_yml):
|
||||
pytest.skip("RAGSAK grammars.yml not generated yet")
|
||||
idx = load_grammar_index("/home/tobi/Desktop/kesai/RAGSAK")
|
||||
assert len(idx.get_all()) > 0
|
||||
# Should be able to resolve a file in the storage package
|
||||
g = idx.get(
|
||||
"/home/tobi/Desktop/kesai/RAGSAK/infrastructure/adapters/search/src/main/kotlin/eu/corentic/springrag/service/storage/GcsStorage.kt",
|
||||
context_symbol="return",
|
||||
)
|
||||
assert g is not None
|
||||
assert "return" in g
|
||||
Loading…
Add table
Reference in a new issue