diff --git a/bex/grammar_index.py b/bex/grammar_index.py new file mode 100644 index 0000000..9561be1 --- /dev/null +++ b/bex/grammar_index.py @@ -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) diff --git a/tests/test_grammar_index.py b/tests/test_grammar_index.py new file mode 100644 index 0000000..8154e23 --- /dev/null +++ b/tests/test_grammar_index.py @@ -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