From 94d13b59d29c5ad3c5f2e59a522b37e437ca75cb Mon Sep 17 00:00:00 2001 From: Kevin Robatel Date: Thu, 1 Oct 2026 15:02:54 +0200 Subject: [PATCH 1/3] fix: delete a file's edges by exact, indexed prefix match delete_edges_for_file used `source_qn LIKE 'file::%' OR target_qn LIKE 'file::%'`. LIKE folds ASCII case and treats `_` as a wildcard, so deleting the edges of `my_file.py` also deleted those of `myXfile.py` and `My_File.py`. The OR also prevented SQLite from using the edge indexes, so every call scanned the whole edges table, which made graph builds quadratic in the number of files. Use two range scans on the indexed columns instead: every qualified name starting with "file::" sorts in ["file::", "file:;"). Co-Authored-By: Claude Opus 5.5 --- src/codetree/graph/store.py | 11 ++++++++--- tests/test_graph_store.py | 18 ++++++++++++++++++ 2 files changed, 26 insertions(+), 3 deletions(-) diff --git a/src/codetree/graph/store.py b/src/codetree/graph/store.py index 06f5bc0..6891ae3 100644 --- a/src/codetree/graph/store.py +++ b/src/codetree/graph/store.py @@ -294,10 +294,15 @@ def edges_to(self, target_qn: str, edge_type: str | None = None) -> list[Edge]: def delete_edges_for_file(self, file_path: str): with self._lock: - prefix = file_path + "::" + # Range scans on the indexed columns: every qualified name starting + # with "file::" sorts in ["file::", "file:;"). Unlike LIKE, this uses + # idx_edges_source/target and is exact (no case folding, no '_' wildcard). + low, high = file_path + "::", file_path + ":;" self._conn.execute( - "DELETE FROM edges WHERE source_qn LIKE ? OR target_qn LIKE ?", - (prefix + "%", prefix + "%"), + "DELETE FROM edges WHERE source_qn >= ? AND source_qn < ?", (low, high) + ) + self._conn.execute( + "DELETE FROM edges WHERE target_qn >= ? AND target_qn < ?", (low, high) ) if not self._in_transaction: self._conn.commit() diff --git a/tests/test_graph_store.py b/tests/test_graph_store.py index 384e200..6f6ff64 100644 --- a/tests/test_graph_store.py +++ b/tests/test_graph_store.py @@ -117,6 +117,23 @@ def test_delete_edges_for_file(self, store): assert len(store.edges_from("a.py::foo")) == 0 assert len(store.edges_from("c.py::baz")) == 1 + def test_delete_edges_for_file_removes_incoming_edges(self, store): + store.upsert_edge(Edge("b.py::caller", "a.py::foo", "CALLS")) + store.delete_edges_for_file("a.py") + assert store.edges_to("a.py::foo") == [] + + def test_delete_edges_for_file_is_exact_not_like_pattern(self, store): + # '_' is a LIKE wildcard and LIKE folds ASCII case — neither may leak. + store.upsert_edge(Edge("my_file.py::foo", "x.py::bar", "CALLS")) + store.upsert_edge(Edge("myXfile.py::foo", "x.py::bar", "CALLS")) + store.upsert_edge(Edge("My_File.py::foo", "x.py::bar", "CALLS")) + store.upsert_edge(Edge("my_file.py.bak::foo", "x.py::bar", "CALLS")) + store.delete_edges_for_file("my_file.py") + assert store.edges_from("my_file.py::foo") == [] + assert len(store.edges_from("myXfile.py::foo")) == 1 + assert len(store.edges_from("My_File.py::foo")) == 1 + assert len(store.edges_from("my_file.py.bak::foo")) == 1 + class TestFileCRUD: def test_upsert_and_get_file(self, store): @@ -159,3 +176,4 @@ def test_stats(self, store): assert stats["files"] == 1 assert stats["symbols"] == 2 assert stats["edges"] == 1 + From 2c89469808ce81e42835edf4ffe4068ef2c7ccca Mon Sep 17 00:00:00 2001 From: Kevin Robatel Date: Thu, 1 Oct 2026 15:03:27 +0200 Subject: [PATCH 2/3] perf: memoize tree-sitter queries and remove quadratic graph build steps - Compile each tree-sitter Query once per (language, pattern) with a bounded, thread-safe cache (`_query()` in languages/base.py). Compiling a query costs several milliseconds, more than parsing a whole file, and every plugin method compiled its queries on each call: ~98% of skeleton extraction time. - `CachedParser` wraps each plugin parser and reuses the tree for consecutive calls on the same source (per-thread, 2 trees, keyed by the source bytes, so a changed file never gets a stale tree). - GraphBuilder: lookup tables replace the O(files^2) scans in import resolution; the content hash uses the source already in memory. - GraphBuilder writes symbols, edges and file rows with multi-row INSERT statements and resolves callees from an in-memory name map loaded once per build. sqlite3 releases the GIL on every row, even with executemany, so row-by-row writes stall behind any CPU-bound thread. The plugin diffs are a mechanical swap (`Query(` -> `_query(`, `Parser(...)` -> `CachedParser(Parser(...))`) with no logic change. On a ~1 100-file repository a cold start drops from 193 s to under 7 s, with an identical graph (symbols, edges, files). The test suite runs about 4x faster. Co-Authored-By: Claude Opus 5.5 --- AGENTS.md | 4 +- CLAUDE.md | 4 +- src/codetree/graph/builder.py | 104 +++++++++++++++++---------- src/codetree/graph/store.py | 75 +++++++++++++++++++ src/codetree/languages/_template.py | 14 ++-- src/codetree/languages/base.py | 61 +++++++++++++++- src/codetree/languages/c.py | 30 ++++---- src/codetree/languages/cpp.py | 36 +++++----- src/codetree/languages/go.py | 38 +++++----- src/codetree/languages/java.py | 40 +++++------ src/codetree/languages/javascript.py | 48 ++++++------- src/codetree/languages/kotlin.py | 36 +++++----- src/codetree/languages/python.py | 30 ++++---- src/codetree/languages/ruby.py | 46 ++++++------ src/codetree/languages/rust.py | 36 +++++----- src/codetree/languages/typescript.py | 40 +++++------ tests/test_graph_builder.py | 25 +++++++ tests/test_graph_store.py | 34 +++++++++ tests/test_parse_cache.py | 67 +++++++++++++++++ 19 files changed, 531 insertions(+), 237 deletions(-) create mode 100644 tests/test_parse_cache.py diff --git a/AGENTS.md b/AGENTS.md index 2995f63..d240df6 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -142,6 +142,7 @@ Each plugin implements: |---|---| | `test_server.py` | Original 4 MCP tools via FastMCP, output format, line accuracy, cross-language | | `test_indexer.py` | Build, skip-dirs, skeleton/symbol/refs/callgraph through indexer layer | +| `test_parse_cache.py` | Memoized query compilation and per-thread parse-tree cache | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -173,7 +174,8 @@ Fixtures in `conftest.py`: `sample_repo` (Python-only), `rich_py_repo` (decorato ## tree-sitter 0.25.x API The tree-sitter Python bindings have breaking changes from older docs: -- Use `Query(LANGUAGE, "...")` not `LANGUAGE.query(...)` +- Use `Query(LANGUAGE, "...")` not `LANGUAGE.query(...)` — in plugins, call the memoized `_query(LANGUAGE, "...")` from `languages/base.py` instead (compiling a Query costs more than parsing a file) +- Wrap module parsers as `_PARSER = CachedParser(Parser(_LANGUAGE))` so repeated calls on the same source reuse the tree - Use `QueryCursor(query).matches(node)` not `query.matches(node)` - Match captures are `list[Node]` — unwrap with `nodes[0]` or use the shared `_matches()` helper from `languages/base.py` - All `.decode()` calls must use `errors="replace"` diff --git a/CLAUDE.md b/CLAUDE.md index 427f4b0..d650910 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -148,6 +148,7 @@ Each plugin implements: |---|---| | `test_server.py` | Original 4 MCP tools via FastMCP, output format, line accuracy, cross-language | | `test_indexer.py` | Build, skip-dirs, skeleton/symbol/refs/callgraph through indexer layer | +| `test_parse_cache.py` | Memoized query compilation and per-thread parse-tree cache | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -181,7 +182,8 @@ Fixtures in `conftest.py`: `sample_repo` (Python-only), `rich_py_repo` (decorato ## tree-sitter 0.25.x API The tree-sitter Python bindings have breaking changes from older docs: -- Use `Query(LANGUAGE, "...")` not `LANGUAGE.query(...)` +- Use `Query(LANGUAGE, "...")` not `LANGUAGE.query(...)` — in plugins, call the memoized `_query(LANGUAGE, "...")` from `languages/base.py` instead (compiling a Query costs more than parsing a file) +- Wrap module parsers as `_PARSER = CachedParser(Parser(_LANGUAGE))` so repeated calls on the same source reuse the tree - Use `QueryCursor(query).matches(node)` not `query.matches(node)` - Match captures are `list[Node]` — unwrap with `nodes[0]` or use the shared `_matches()` helper from `languages/base.py` - All `.decode()` calls must use `errors="replace"` diff --git a/src/codetree/graph/builder.py b/src/codetree/graph/builder.py index ee1e930..0bb2b44 100644 --- a/src/codetree/graph/builder.py +++ b/src/codetree/graph/builder.py @@ -1,19 +1,29 @@ import hashlib import time +from functools import lru_cache from pathlib import Path from .store import GraphStore from .models import SymbolNode, Edge, make_qualified_name from ..indexer import Indexer +@lru_cache(maxsize=65536) +def _stem(path: str) -> str: + return Path(path).stem + + class GraphBuilder: def __init__(self, root: str, store: GraphStore): self._root = Path(root) self._store = store self._file_imports: dict[str, set[str]] = {} # file → set of imported module stems/symbols + # name → symbols, loaded once after pass 1 (all files' symbols are then in + # the store) instead of one SQL lookup per callee. + self._symbols_by_name: dict[str, list[SymbolNode]] = {} - def _hash_file(self, path: Path) -> str: - return hashlib.sha256(path.read_bytes()).hexdigest() + @staticmethod + def _hash_source(source: bytes) -> str: + return hashlib.sha256(source).hexdigest() def _is_test_file(self, rel_path: str) -> bool: name = Path(rel_path).name @@ -46,10 +56,22 @@ def build(self, indexer: Indexer | None = None) -> dict: abs_path = self._root / rel_path if abs_path.exists(): current_files[rel_path] = { - "hash": self._hash_file(abs_path), + # Indexer already holds the file bytes — no second read. + "hash": self._hash_source(entry.source), "entry": entry, } + # Lookup tables replacing per-import scans of every file (O(files²)). + # first_file_by_key: stem or full path → first matching file, in order. + # files_by_stem: stem → every file with that stem. + first_file_by_key: dict[str, str] = {} + files_by_stem: dict[str, list[str]] = {} + for candidate in current_files: + stem = _stem(candidate) + first_file_by_key.setdefault(stem, candidate) + first_file_by_key.setdefault(candidate, candidate) + files_by_stem.setdefault(stem, []).append(candidate) + # Determine which files changed files_indexed = 0 files_skipped = 0 @@ -59,10 +81,18 @@ def build(self, indexer: Indexer | None = None) -> dict: indexed_paths = set() changed_files = [] # Track files that need edge resolution + # Writes are batched (see GraphStore._insert_many). Each pass is flushed + # before the next one reads, so the store sees the same state as with + # row-by-row writes. + stored_files = {f["file_path"]: f for f in self._store.all_files()} + file_rows: list[tuple[str, str, str, bool]] = [] + new_symbols: list[SymbolNode] = [] + new_edges: list[Edge] = [] + # ── Pass 1: Insert all symbols first ────────────────────────────── for rel_path, info in current_files.items(): indexed_paths.add(rel_path) - existing = self._store.get_file(rel_path) + existing = stored_files.get(rel_path) if existing and existing["sha256"] == info["hash"]: files_skipped += 1 continue @@ -72,17 +102,12 @@ def build(self, indexer: Indexer | None = None) -> dict: entry = info["entry"] is_test = self._is_test_file(rel_path) - # Clear old data for this file - self._store.delete_symbols_for_file(rel_path) - self._store.delete_edges_for_file(rel_path) + # Clear old data for this file (a never-indexed file has none) + if existing: + self._store.delete_symbols_for_file(rel_path) + self._store.delete_edges_for_file(rel_path) - # Upsert file record - self._store.upsert_file( - rel_path, - sha256=info["hash"], - language=entry.language, - is_test=is_test, - ) + file_rows.append((rel_path, info["hash"], entry.language, is_test)) # Build symbols from skeleton for item in entry.skeleton: @@ -101,22 +126,28 @@ def build(self, indexer: Indexer | None = None) -> dict: is_test=is_test or item["name"].startswith("test_") or item["name"].startswith("Test"), is_entry_point=is_entry, ) - self._store.upsert_symbol(sym) + new_symbols.append(sym) symbols_created += 1 # CONTAINS edges for methods if item.get("parent"): parent_qn = make_qualified_name(rel_path, item["parent"]) - self._store.upsert_edge(Edge(parent_qn, qn, "CONTAINS")) + new_edges.append(Edge(parent_qn, qn, "CONTAINS")) edges_created += 1 changed_files.append((rel_path, info)) + self._store.upsert_files(file_rows) + self._store.upsert_symbols(new_symbols) + self._store.upsert_edges(new_edges) + new_edges = [] + self._symbols_by_name = self._store.symbols_by_name_map() if changed_files else {} + # ── Pass 2: Build CALLS and IMPORTS edges (all symbols now in store) ── # First, parse imports for all changed files to enable type-aware resolution for rel_path, info in changed_files: entry = info["entry"] - self._file_imports[rel_path] = self._parse_file_imports(entry, current_files) + self._file_imports[rel_path] = self._parse_file_imports(entry, files_by_stem) for rel_path, info in changed_files: entry = info["entry"] @@ -130,7 +161,7 @@ def build(self, indexer: Indexer | None = None) -> dict: for callee_name in callees: resolved = self._resolve_callee(rel_path, callee_name) for target_qn, weight in resolved: - self._store.upsert_edge(Edge(caller_qn, target_qn, "CALLS", weight=weight)) + new_edges.append(Edge(caller_qn, target_qn, "CALLS", weight=weight)) edges_created += 1 # Build IMPORTS edges @@ -140,14 +171,12 @@ def build(self, indexer: Indexer | None = None) -> dict: parts = text.split() if len(parts) >= 2: module = parts[1] if parts[0] in ("import", "from") else parts[0] - for candidate in current_files: - stem = Path(candidate).stem - if stem == module or candidate == module: - self._store.upsert_edge( - Edge(f"{rel_path}::__file__", f"{candidate}::__file__", "IMPORTS") - ) - edges_created += 1 - break + candidate = first_file_by_key.get(module) + if candidate is not None: + new_edges.append( + Edge(f"{rel_path}::__file__", f"{candidate}::__file__", "IMPORTS") + ) + edges_created += 1 # ── Pass 3: Build TESTS edges (link test functions to tested symbols) ── for rel_path, info in changed_files: @@ -166,16 +195,19 @@ def build(self, indexer: Indexer | None = None) -> dict: tested_name = name[4:] # strip Test prefix if not tested_name: continue - targets = self._store.symbols_by_name(tested_name) + targets = self._symbols_by_name.get(tested_name, []) if targets: test_qn = make_qualified_name(rel_path, name, item.get("parent")) for t in targets: if not t.is_test: - self._store.upsert_edge(Edge(test_qn, t.qualified_name, "TESTS")) + new_edges.append(Edge(test_qn, t.qualified_name, "TESTS")) edges_created += 1 + # Flush edges before deleting removed files: their edges must go too. + self._store.upsert_edges(new_edges) + # Delete files that no longer exist - for stored_file in self._store.all_files(): + for stored_file in stored_files.values(): if stored_file["file_path"] not in indexed_paths: fp = stored_file["file_path"] self._store.delete_symbols_for_file(fp) @@ -192,7 +224,7 @@ def build(self, indexer: Indexer | None = None) -> dict: "edges_created": edges_created, } - def _parse_file_imports(self, entry, current_files: dict) -> set[str]: + def _parse_file_imports(self, entry, files_by_stem: dict[str, list[str]]) -> set[str]: """Extract imported module stems and symbol names from a file's imports.""" imported = set() imports = entry.plugin.extract_imports(entry.source) @@ -218,11 +250,9 @@ def _parse_file_imports(self, entry, current_files: dict) -> set[str]: # Also add individual path components for component in Path(cleaned).parts: imported.add(component) - # Also record stems of files that this file imports - for candidate in current_files: - stem = Path(candidate).stem - if stem in imported: - imported.add(candidate) + # Also record paths of files whose stem this file imports + for stem in list(imported): + imported.update(files_by_stem.get(stem, ())) return imported def _resolve_callee(self, caller_file: str, callee_name: str) -> list[tuple[str, float]]: @@ -231,7 +261,7 @@ def _resolve_callee(self, caller_file: str, callee_name: str) -> list[tuple[str, Returns list of (qualified_name, weight) tuples. Weight 1.0 = import-confirmed, 0.5 = name-only match. """ - targets = self._store.symbols_by_name(callee_name) + targets = self._symbols_by_name.get(callee_name, []) if not targets: return [(f"?::{callee_name}", 0.5)] @@ -244,7 +274,7 @@ def _resolve_callee(self, caller_file: str, callee_name: str) -> list[tuple[str, if t.file_path == caller_file: # Same file — always high confidence same_file.append((t.qualified_name, 1.0)) - elif t.file_path in caller_imports or Path(t.file_path).stem in caller_imports: + elif t.file_path in caller_imports or _stem(t.file_path) in caller_imports: import_confirmed.append((t.qualified_name, 1.0)) else: name_only.append((t.qualified_name, 0.5)) diff --git a/src/codetree/graph/store.py b/src/codetree/graph/store.py index 6891ae3..073780e 100644 --- a/src/codetree/graph/store.py +++ b/src/codetree/graph/store.py @@ -142,6 +142,41 @@ def upsert_file(self, file_path: str, sha256: str, language: str, is_test: bool) if not self._in_transaction: self._conn.commit() + def upsert_files(self, rows: list[tuple[str, str, str, bool]]): + """Batch upsert_file: rows of (file_path, sha256, language, is_test).""" + now = time.time() + self._insert_many( + "INSERT OR REPLACE INTO files (file_path, sha256, language, is_test, indexed_at) VALUES ", + [(fp, sha, lang, int(is_test), now) for fp, sha, lang, is_test in rows], + ) + + # -- Batch writes -------------------------------------------------------- + + # Bound parameters per statement; 999 is the lowest SQLITE_MAX_VARIABLE_NUMBER. + _MAX_VARIABLES = 999 + + def _insert_many(self, sql_head: str, rows: list[tuple]): + """Insert rows with multi-row VALUES statements, preserving row order. + + sqlite3 releases the GIL around every statement step — once per row even + with executemany. A background build then re-waits for the GIL behind any + CPU-bound tool call on every row; one statement per chunk avoids that. + """ + if not rows: + return + width = len(rows[0]) + per_statement = max(1, self._MAX_VARIABLES // width) + row_sql = "(" + ",".join("?" * width) + ")" + with self._lock: + for start in range(0, len(rows), per_statement): + chunk = rows[start:start + per_statement] + self._conn.execute( + sql_head + ",".join([row_sql] * len(chunk)), + [value for row in chunk for value in row], + ) + if not self._in_transaction: + self._conn.commit() + def get_file(self, file_path: str) -> dict | None: with self._lock: cur = self._conn.execute( @@ -193,6 +228,39 @@ def upsert_symbol(self, sym: SymbolNode): if not self._in_transaction: self._conn.commit() + def upsert_symbols(self, syms: list[SymbolNode]): + """Batch upsert_symbol; later duplicates replace earlier ones, as one by one.""" + self._insert_many( + "INSERT OR REPLACE INTO symbols " + "(qualified_name, name, kind, parent_qn, file_path, start_line, end_line, " + "doc, params, is_test, is_entry_point) VALUES ", + [ + ( + sym.qualified_name, sym.name, sym.kind, sym.parent_qn, + sym.file_path, sym.start_line, sym.end_line, + sym.doc, sym.params, int(sym.is_test), int(sym.is_entry_point), + ) + for sym in syms + ], + ) + + def symbols_by_name_map(self) -> dict[str, list[SymbolNode]]: + """All symbols grouped by name, each list in the order symbols_by_name() returns.""" + with self._lock: + cur = self._conn.execute( + "SELECT qualified_name, name, kind, parent_qn, file_path, start_line, end_line, " + "doc, params, is_test, is_entry_point FROM symbols ORDER BY name, rowid" + ) + by_name: dict[str, list[SymbolNode]] = {} + for r in cur.fetchall(): + by_name.setdefault(r[1], []).append(SymbolNode( + qualified_name=r[0], name=r[1], kind=r[2], parent_qn=r[3], + file_path=r[4], start_line=r[5], end_line=r[6], + doc=r[7] or "", params=r[8] or "", + is_test=bool(r[9]), is_entry_point=bool(r[10]), + )) + return by_name + def get_symbol(self, qualified_name: str) -> SymbolNode | None: with self._lock: cur = self._conn.execute( @@ -262,6 +330,13 @@ def upsert_edge(self, edge: Edge): if not self._in_transaction: self._conn.commit() + def upsert_edges(self, edges: list[Edge]): + """Batch upsert_edge; later duplicates replace earlier ones, as one by one.""" + self._insert_many( + "INSERT OR REPLACE INTO edges (source_qn, target_qn, type, weight) VALUES ", + [(e.source_qn, e.target_qn, e.type, e.weight) for e in edges], + ) + def edges_from(self, source_qn: str, edge_type: str | None = None) -> list[Edge]: with self._lock: if edge_type: diff --git a/src/codetree/languages/_template.py b/src/codetree/languages/_template.py index 9bd0b2e..da7f017 100644 --- a/src/codetree/languages/_template.py +++ b/src/codetree/languages/_template.py @@ -39,14 +39,14 @@ # See docs/language-nodes.md for a cheatsheet of node types per language. # ============================================================ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser # TODO: replace with your grammar import # import tree_sitter_LANG as tslang # _LANGUAGE = Language(tslang.language()) -# _PARSER = Parser(_LANGUAGE) +# _PARSER = CachedParser(Parser(_LANGUAGE)) # CachedParser from .base -from .base import LanguagePlugin, _matches +from .base import LanguagePlugin, _matches, _query # _query: cached Query compile # Import this from base — do not copy @@ -68,7 +68,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: 3. Top-level functions Example (Python): - q = Query(_LANGUAGE, "(module (class_definition name: (identifier) @name) @def)") + q = _query(_LANGUAGE, "(module (class_definition name: (identifier) @name) @def)") """ # tree = _PARSER.parse(source) results = [] @@ -98,7 +98,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: Step 3 — Capture called function names (usually (identifier) or method name) Example (JavaScript): - q = Query(_LANGUAGE, ''' + q = _query(_LANGUAGE, ''' (call_expression function: [ (identifier) @called (member_expression property: (property_identifier) @called) @@ -116,7 +116,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: than Python-level filtering for large files). This query works for Python, JavaScript, Go, Rust, Java: - q = Query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') Each result dict: {"line": int, "col": int} (both 1-based and 0-based respectively) """ @@ -129,7 +129,7 @@ def extract_imports(self, source: bytes) -> list[dict]: Each result dict: {"line": int, "text": str} Example (Python): - q = Query(_LANGUAGE, "(module (import_statement) @imp)") + q = _query(_LANGUAGE, "(module (import_statement) @imp)") """ # tree = _PARSER.parse(source) return [] diff --git a/src/codetree/languages/base.py b/src/codetree/languages/base.py index e3c85f0..69f93ea 100644 --- a/src/codetree/languages/base.py +++ b/src/codetree/languages/base.py @@ -1,8 +1,67 @@ +import threading from abc import ABC, abstractmethod +from collections import OrderedDict from tree_sitter import Query, QueryCursor +# Compiling a tree-sitter Query costs several milliseconds (more than parsing a +# whole file), so compiled queries are memoized per (language, pattern). +# Bounded because some patterns interpolate a symbol name. +_QUERY_CACHE_MAX = 2048 +_QUERY_CACHE: "OrderedDict[tuple[int, str], tuple[object, Query]]" = OrderedDict() +_QUERY_CACHE_LOCK = threading.Lock() + + +def _query(lang, pattern: str) -> Query: + """Return a compiled Query for (lang, pattern), compiling it at most once.""" + key = (id(lang), pattern) + with _QUERY_CACHE_LOCK: + hit = _QUERY_CACHE.get(key) + if hit is not None: + _QUERY_CACHE.move_to_end(key) + return hit[1] + query = Query(lang, pattern) + with _QUERY_CACHE_LOCK: + # Keep a reference to lang so its id() cannot be reused while cached. + _QUERY_CACHE[key] = (lang, query) + if len(_QUERY_CACHE) > _QUERY_CACHE_MAX: + _QUERY_CACHE.popitem(last=False) + return query + + +class CachedParser: + """Wrap a tree-sitter Parser and reuse trees for recently parsed sources. + + Graph building and call-graph construction ask the same plugin about many + functions of one file in a row; without this every call re-parses the file. + The cache is per thread (trees are not shared across threads) and keyed by + the source bytes, so a changed file can never return a stale tree. It only + needs to cover consecutive calls on one file, so it stays tiny: every + FastMCP worker thread holds its own copy. + """ + + _MAX_TREES = 2 + + def __init__(self, parser): + self._parser = parser + self._local = threading.local() + + def parse(self, source: bytes): + trees = getattr(self._local, "trees", None) + if trees is None: + trees = self._local.trees = OrderedDict() + tree = trees.get(source) + if tree is not None: + trees.move_to_end(source) + return tree + tree = self._parser.parse(source) + trees[source] = tree + if len(trees) > self._MAX_TREES: + trees.popitem(last=False) + return tree + + def _clean_doc(text: str) -> str: """Extract the first meaningful line from a doc comment.""" lines = text.strip().splitlines() @@ -22,7 +81,7 @@ def _fill_docs_from_siblings(results: list[dict], tree_root, lang, queries: list """ comment_types = ("comment", "line_comment", "block_comment") for q_str in queries: - for _, m in _matches(Query(lang, q_str), tree_root): + for _, m in _matches(_query(lang, q_str), tree_root): node = m["def"] name = m["name"].text.decode("utf-8", errors="replace") line = m["name"].start_point[0] + 1 diff --git a/src/codetree/languages/c.py b/src/codetree/languages/c.py index 022c71f..7bd99e0 100644 --- a/src/codetree/languages/c.py +++ b/src/codetree/languages/c.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_c as tsc -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsc.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -24,7 +24,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Named structs with body (struct Foo { ... }) - q = Query(_LANGUAGE, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") + q = _query(_LANGUAGE, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") for _, m in _matches(q, tree.root_node): results.append({ "type": "struct", @@ -35,7 +35,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Typedef structs: typedef struct { ... } Name; - q = Query(_LANGUAGE, "(type_definition declarator: (type_identifier) @name) @def") + q = _query(_LANGUAGE, "(type_definition declarator: (type_identifier) @name) @def") for _, m in _matches(q, tree.root_node): results.append({ "type": "struct", @@ -46,7 +46,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Functions - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (translation_unit (function_definition declarator: (function_declarator @@ -87,21 +87,21 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N tree = _parse(source) # Functions - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Structs - q = Query(_LANGUAGE, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") + q = _query(_LANGUAGE, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Typedef structs - q = Query(_LANGUAGE, "(type_definition declarator: (type_identifier) @name) @def") + q = _query(_LANGUAGE, "(type_definition declarator: (type_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -112,14 +112,14 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: return [] - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (call_expression function: [ (identifier) @called (field_expression field: (field_identifier) @called) @@ -135,7 +135,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: usages = [] seen = set() for node_type in ("identifier", "type_identifier", "field_identifier"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -148,7 +148,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(translation_unit (preproc_include) @imp)") + q = _query(_LANGUAGE, "(translation_unit (preproc_include) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -166,7 +166,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Find function node fn_node = None - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -266,7 +266,7 @@ def walk(node): def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: (identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] diff --git a/src/codetree/languages/cpp.py b/src/codetree/languages/cpp.py index 7de2ec6..55dc6ff 100644 --- a/src/codetree/languages/cpp.py +++ b/src/codetree/languages/cpp.py @@ -1,10 +1,10 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_cpp as tscpp from .c import CPlugin -from .base import _matches, _fill_docs_from_siblings +from .base import _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tscpp.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -28,7 +28,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Classes - q = Query(lang, "(class_specifier name: (type_identifier) @name) @def") + q = _query(lang, "(class_specifier name: (type_identifier) @name) @def") for _, m in _matches(q, tree.root_node): results.append({ "type": "class", @@ -39,7 +39,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside classes (function_definition in field_declaration_list) - q = Query(lang, """ + q = _query(lang, """ (class_specifier name: (type_identifier) @class_name body: (field_declaration_list @@ -58,7 +58,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Structs (C++ uses same struct_specifier as C) - q = Query(lang, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") + q = _query(lang, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") for _, m in _matches(q, tree.root_node): results.append({ "type": "struct", @@ -69,7 +69,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions (translation_unit direct children) - q = Query(lang, """ + q = _query(lang, """ (translation_unit (function_definition declarator: (function_declarator @@ -86,7 +86,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Functions inside namespaces - q = Query(lang, """ + q = _query(lang, """ (namespace_definition body: (declaration_list (function_definition @@ -129,21 +129,21 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N tree = _parse(source) # Functions — check both identifier (top-level) and field_identifier (class methods) - q = Query(lang, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") + q = _query(lang, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Classes - q = Query(lang, "(class_specifier name: (type_identifier) @name) @def") + q = _query(lang, "(class_specifier name: (type_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Structs - q = Query(lang, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") + q = _query(lang, "(struct_specifier name: (type_identifier) @name body: (field_declaration_list)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -155,14 +155,14 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: lang = _LANGUAGE tree = _parse(source) fn_node = None - q = Query(lang, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") + q = _query(lang, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: return [] - q = Query(lang, """ + q = _query(lang, """ (call_expression function: [ (identifier) @called (field_expression field: (field_identifier) @called) @@ -178,7 +178,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: usages = [] seen = set() for node_type in ("identifier", "type_identifier", "field_identifier", "namespace_identifier"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -192,7 +192,7 @@ def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] # #include statements - q = Query(_LANGUAGE, "(translation_unit (preproc_include) @imp)") + q = _query(_LANGUAGE, "(translation_unit (preproc_include) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -200,7 +200,7 @@ def extract_imports(self, source: bytes) -> list[dict]: "text": node.text.decode("utf-8", errors="replace").strip(), }) # using declarations - q = Query(_LANGUAGE, "(translation_unit (using_declaration) @imp)") + q = _query(_LANGUAGE, "(translation_unit (using_declaration) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -218,7 +218,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Find function node — check both identifier (top-level) and field_identifier (class methods) fn_node = None - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -337,7 +337,7 @@ def walk(node): def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") + q = _query(_LANGUAGE, "(function_definition declarator: (function_declarator declarator: [(identifier) @name (field_identifier) @name])) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] diff --git a/src/codetree/languages/go.py b/src/codetree/languages/go.py index b7ce2d4..f655c04 100644 --- a/src/codetree/languages/go.py +++ b/src/codetree/languages/go.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_go as tsgo -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsgo.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -18,7 +18,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Structs (Go's equivalent of classes) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (source_file (type_declaration (type_spec name: (type_identifier) @name @@ -34,7 +34,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Interfaces - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (source_file (type_declaration (type_spec name: (type_identifier) @name @@ -50,7 +50,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods (receiver functions) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (method_declaration receiver: (parameter_list (parameter_declaration @@ -69,7 +69,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (source_file (function_declaration name: (identifier) @name @@ -100,7 +100,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N tree = _parse(source) # Functions - q = Query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -110,7 +110,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N ) # Methods (receiver functions) - q = Query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") + q = _query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -120,7 +120,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N ) # Struct/interface types - q = Query(_LANGUAGE, "(type_declaration (type_spec name: (type_identifier) @name)) @def") + q = _query(_LANGUAGE, "(type_declaration (type_spec name: (type_identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -134,21 +134,21 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break # Also search method declarations if fn_node is None: - q = Query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") + q = _query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: return [] - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (call_expression function: [ (identifier) @called (selector_expression field: (field_identifier) @called) @@ -164,7 +164,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: usages = [] seen = set() for node_type in ("identifier", "type_identifier", "field_identifier"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -177,7 +177,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(source_file (import_declaration) @imp)") + q = _query(_LANGUAGE, "(source_file (import_declaration) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -190,13 +190,13 @@ def extract_imports(self, source: bytes) -> list[dict]: def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: - q = Query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") + q = _query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -230,12 +230,12 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Find function node (top-level function or method) fn_node = None - for _, m in _matches(Query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def"), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, "(function_declaration name: (identifier) @name) @def"), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: - for _, m in _matches(Query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def"), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, "(method_declaration name: (field_identifier) @name) @def"), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/java.py b/src/codetree/languages/java.py index 56bf5e5..a0f6978 100644 --- a/src/codetree/languages/java.py +++ b/src/codetree/languages/java.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_java as tsjava -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsjava.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -18,7 +18,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Top-level classes - q = Query(_LANGUAGE, "(program (class_declaration name: (identifier) @name) @def)") + q = _query(_LANGUAGE, "(program (class_declaration name: (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "class", @@ -29,7 +29,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside classes - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class_declaration name: (identifier) @class_name body: (class_body @@ -47,7 +47,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Constructors inside classes - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class_declaration name: (identifier) @class_name body: (class_body @@ -65,7 +65,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Interfaces (top-level) - q = Query(_LANGUAGE, "(program (interface_declaration name: (identifier) @name) @def)") + q = _query(_LANGUAGE, "(program (interface_declaration name: (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "interface", @@ -76,7 +76,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Enums (top-level) - q = Query(_LANGUAGE, "(program (enum_declaration name: (identifier) @name) @def)") + q = _query(_LANGUAGE, "(program (enum_declaration name: (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "enum", @@ -87,7 +87,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside enums (live under enum_body > enum_body_declarations) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (enum_declaration name: (identifier) @class_name body: (enum_body @@ -106,7 +106,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside interfaces - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (interface_declaration name: (identifier) @class_name body: (interface_body @@ -145,20 +145,20 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(interface_declaration name: (identifier) @name) @def", "(enum_declaration name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Methods - q = Query(_LANGUAGE, "(method_declaration name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(method_declaration name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Constructors - q = Query(_LANGUAGE, "(constructor_declaration name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(constructor_declaration name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -173,7 +173,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: "(method_declaration name: (identifier) @name) @def", "(constructor_declaration name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -183,11 +183,11 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: return [] calls = set() # Method calls - q = Query(_LANGUAGE, "(method_invocation name: (identifier) @called)") + q = _query(_LANGUAGE, "(method_invocation name: (identifier) @called)") for _, m in _matches(q, fn_node): calls.add(m["called"].text.decode("utf-8", errors="replace")) # Object creation (new Calculator()) - q = Query(_LANGUAGE, "(object_creation_expression type: (type_identifier) @called)") + q = _query(_LANGUAGE, "(object_creation_expression type: (type_identifier) @called)") for _, m in _matches(q, fn_node): calls.add(m["called"].text.decode("utf-8", errors="replace")) return sorted(calls) @@ -200,7 +200,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: # type positions (e.g. `Calculator calc = ...` uses type_identifier, # while method names use identifier). for node_type in ("identifier", "type_identifier"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -213,7 +213,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(program (import_declaration) @imp)") + q = _query(_LANGUAGE, "(program (import_declaration) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -230,7 +230,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: "(method_declaration name: (identifier) @name) @def", "(constructor_declaration name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -275,7 +275,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: "(method_declaration name: (identifier) @name) @def", "(constructor_declaration name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/javascript.py b/src/codetree/languages/javascript.py index 150cfef..c41cf56 100644 --- a/src/codetree/languages/javascript.py +++ b/src/codetree/languages/javascript.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_javascript as tsjs -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsjs.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -47,7 +47,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(program (class_declaration name: (identifier) @name) @def)", "(program (export_statement (class_declaration name: (identifier) @name) @def))", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "class", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -57,7 +57,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside classes - q = Query(lang, """ + q = _query(lang, """ (class_declaration name: (identifier) @class_name body: (class_body @@ -75,7 +75,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level function declarations: function foo() {} - q = Query(lang, """ + q = _query(lang, """ (program (function_declaration name: (identifier) @name parameters: (formal_parameters) @params)) @@ -100,7 +100,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (identifier) @name parameters: (formal_parameters) @params) @def))""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "function", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -110,7 +110,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Generator functions: function* gen() {} - q = Query(lang, """ + q = _query(lang, """ (program (generator_function_declaration name: (identifier) @name parameters: (formal_parameters) @params)) @@ -136,7 +136,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (identifier) @name value: [(arrow_function) @fn (function_expression) @fn]))))""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): fn_node = m.get("fn") results.append({ "type": "function", @@ -181,7 +181,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(export_statement (class_declaration name: (identifier) @name) @def)", "(export_statement (generator_function_declaration name: (identifier) @name) @def)", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -200,7 +200,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N name: (identifier) @name value: [(arrow_function) (function_expression)])) @def)""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -209,7 +209,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N ) # Methods inside classes (method_definition) - q = Query(lang, "(method_definition name: (property_identifier) @name) @def") + q = _query(lang, "(method_definition name: (property_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -232,7 +232,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: "(export_statement (function_declaration name: (identifier) @name) @def)", "(export_statement (generator_function_declaration name: (identifier) @name) @def)", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -246,7 +246,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: name: (identifier) @name value: [(arrow_function) @def (function_expression) @def])""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -255,7 +255,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: # Methods inside classes if fn_node is None: - q = Query(lang, "(method_definition name: (property_identifier) @name) @def") + q = _query(lang, "(method_definition name: (property_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -264,13 +264,13 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: if fn_node is None: return [] - q_call = Query(lang, """ + q_call = _query(lang, """ (call_expression function: [ (identifier) @called (member_expression property: (property_identifier) @called) ]) """) - q_new = Query(lang, "(new_expression constructor: (identifier) @called)") + q_new = _query(lang, "(new_expression constructor: (identifier) @called)") calls = set() for _, m in _matches(q_call, fn_node): calls.add(m["called"].text.decode("utf-8", errors="replace")) @@ -281,7 +281,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: lang = self._get_language() tree = self._get_parser().parse(source) - q = Query(lang, f'((identifier) @name (#eq? @name "{name}"))') + q = _query(lang, f'((identifier) @name (#eq? @name "{name}"))') usages = [] for _, m in _matches(q, tree.root_node): node = m["name"] @@ -292,7 +292,7 @@ def extract_imports(self, source: bytes) -> list[dict]: lang = self._get_language() tree = self._get_parser().parse(source) results = [] - q = Query(lang, "(program (import_statement) @imp)") + q = _query(lang, "(program (import_statement) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -313,7 +313,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: "(export_statement (function_declaration name: (identifier) @name) @def)", "(export_statement (generator_function_declaration name: (identifier) @name) @def)", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -326,7 +326,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: name: (identifier) @name value: [(arrow_function) @def (function_expression) @def])""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -334,7 +334,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: break if fn_node is None: - q = Query(lang, "(method_definition name: (property_identifier) @name) @def") + q = _query(lang, "(method_definition name: (property_identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -383,7 +383,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: "(generator_function_declaration name: (identifier) @name) @def", "(method_definition name: (property_identifier) @name) @def", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -393,7 +393,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Also look for arrow/function expression: const foo = () => {} if fn_node is None: q_str = "(variable_declarator name: (identifier) @name value: [(arrow_function) @def (function_expression) @def])" - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/kotlin.py b/src/codetree/languages/kotlin.py index 1178091..580bab5 100644 --- a/src/codetree/languages/kotlin.py +++ b/src/codetree/languages/kotlin.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_kotlin as tskotlin -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tskotlin.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -18,7 +18,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Top-level classes / interfaces (both use class_declaration) - q = Query(_LANGUAGE, "(source_file (class_declaration (identifier) @name) @def)") + q = _query(_LANGUAGE, "(source_file (class_declaration (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): node_def = m["def"] # Check if it's an interface @@ -37,7 +37,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level objects - q = Query(_LANGUAGE, "(source_file (object_declaration (identifier) @name) @def)") + q = _query(_LANGUAGE, "(source_file (object_declaration (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "class", @@ -48,7 +48,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions - q = Query(_LANGUAGE, "(source_file (function_declaration (identifier) @name (function_value_parameters) @params) @def)") + q = _query(_LANGUAGE, "(source_file (function_declaration (identifier) @name (function_value_parameters) @params) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "function", @@ -59,7 +59,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside classes/interfaces - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class_declaration (identifier) @class_name (class_body @@ -77,7 +77,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside objects - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (object_declaration (identifier) @class_name (class_body @@ -114,13 +114,13 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(class_declaration (identifier) @name) @def", "(object_declaration (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 # Functions/Methods - q = Query(_LANGUAGE, "(function_declaration (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -131,7 +131,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_declaration (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -141,7 +141,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: calls = set() # Method calls: foo(), foo.bar() - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (call_expression [ (identifier) @called @@ -164,7 +164,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: usages = [] seen = set() # Kotlin uses identifier for most things - q = Query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -178,7 +178,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(import) @imp") + q = _query(_LANGUAGE, "(import) @imp") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -191,7 +191,7 @@ def extract_imports(self, source: bytes) -> list[dict]: def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_declaration (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -226,7 +226,7 @@ def walk(node): def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_declaration (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_declaration (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -243,7 +243,7 @@ def _add(name, line, var_type="", kind="local"): results.append({"name": name, "line": line, "type": var_type, "kind": kind}) # Parameters - q_params = Query(_LANGUAGE, "(parameter (identifier) @name (user_type)? @type)") + q_params = _query(_LANGUAGE, "(parameter (identifier) @name (user_type)? @type)") for _, m in _matches(q_params, fn_node): type_text = m.get("type").text.decode("utf-8", errors="replace") if m.get("type") else "" _add(m["name"].text.decode("utf-8", errors="replace"), @@ -252,7 +252,7 @@ def _add(name, line, var_type="", kind="local"): kind="parameter") # Local variables (val/var) - q_vars = Query(_LANGUAGE, "(variable_declaration (identifier) @name (user_type)? @type)") + q_vars = _query(_LANGUAGE, "(variable_declaration (identifier) @name (user_type)? @type)") for _, m in _matches(q_vars, fn_node): type_text = m.get("type").text.decode("utf-8", errors="replace") if m.get("type") else "" _add(m["name"].text.decode("utf-8", errors="replace"), diff --git a/src/codetree/languages/python.py b/src/codetree/languages/python.py index ec738d6..6bca4e8 100644 --- a/src/codetree/languages/python.py +++ b/src/codetree/languages/python.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_python as tspython -from .base import LanguagePlugin, _matches, _clean_doc +from .base import LanguagePlugin, _matches, _clean_doc, _query, CachedParser _LANGUAGE = Language(tspython.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -30,7 +30,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(module (class_definition name: (identifier) @name) @def)", "(module (decorated_definition (class_definition name: (identifier) @name)) @def)", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): results.append({ "type": "class", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -55,7 +55,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (identifier) @method_name parameters: (parameters) @params))))""", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): results.append({ "type": "method", "name": m["method_name"].text.decode("utf-8", errors="replace"), @@ -74,7 +74,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (identifier) @name parameters: (parameters) @params)) @def)""", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): results.append({ "type": "function", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -92,7 +92,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: '(class_definition name: (identifier) @name body: (block (expression_statement (string) @doc)))', '(function_definition name: (identifier) @name body: (block (expression_statement (string) @doc)))', ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): name = m["name"].text.decode("utf-8", errors="replace") line = m["name"].start_point[0] + 1 doc_text = m["doc"].text.decode("utf-8", errors="replace") @@ -122,7 +122,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N tree = _parse(source) for node_type in ("function_definition", "class_definition"): # Decorated definition first — return full decorated_definition (includes decorator) - q = Query(_LANGUAGE, f"(decorated_definition ({node_type} name: (identifier) @name)) @def") + q = _query(_LANGUAGE, f"(decorated_definition ({node_type} name: (identifier) @name)) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -131,7 +131,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N node.start_point[0] + 1, ) # Plain definition (not decorated) - q = Query(_LANGUAGE, f"({node_type} name: (identifier) @name) @def") + q = _query(_LANGUAGE, f"({node_type} name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -149,7 +149,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: "(function_definition name: (identifier) @name) @def", "(decorated_definition (function_definition name: (identifier) @name)) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -157,7 +157,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: break if fn_node is None: return [] - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (call function: [ (identifier) @called (attribute attribute: (identifier) @called) @@ -170,7 +170,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: tree = _parse(source) - q = Query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'((identifier) @name (#eq? @name "{name}"))') usages = [] for _, m in _matches(q, tree.root_node): node = m["name"] @@ -185,7 +185,7 @@ def extract_imports(self, source: bytes) -> list[dict]: "(module (import_from_statement) @imp)", "(module (future_import_statement) @imp)", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): node = m["imp"] results.append({ "line": node.start_point[0] + 1, @@ -201,7 +201,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: "(function_definition name: (identifier) @name) @def", "(decorated_definition (function_definition name: (identifier) @name)) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -241,7 +241,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: "(decorated_definition definition: (function_definition name: (identifier) @name) @def)", "(function_definition name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/ruby.py b/src/codetree/languages/ruby.py index 4a45a96..5d513ce 100644 --- a/src/codetree/languages/ruby.py +++ b/src/codetree/languages/ruby.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_ruby as tsruby -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsruby.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -18,7 +18,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Classes - q = Query(_LANGUAGE, "(program (class name: (constant) @name) @def)") + q = _query(_LANGUAGE, "(program (class name: (constant) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "class", @@ -29,7 +29,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Modules (treated as class for skeleton purposes) - q = Query(_LANGUAGE, "(program (module name: (constant) @name) @def)") + q = _query(_LANGUAGE, "(program (module name: (constant) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "class", @@ -40,7 +40,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Instance methods inside classes (with params) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class name: (constant) @class_name (body_statement @@ -58,7 +58,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Instance methods inside classes (no params) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class name: (constant) @class_name (body_statement @@ -78,7 +78,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Singleton methods in classes (def self.foo) — with params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class name: (constant) @class_name (body_statement @@ -96,7 +96,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Singleton methods in classes — no params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (class name: (constant) @class_name (body_statement @@ -115,7 +115,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Singleton methods in modules — with params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (module name: (constant) @class_name (body_statement @@ -133,7 +133,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Singleton methods in modules — no params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (module name: (constant) @class_name (body_statement @@ -152,7 +152,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Instance methods in modules — with params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (module name: (constant) @class_name (body_statement @@ -170,7 +170,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Instance methods in modules — no params - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (module name: (constant) @class_name (body_statement @@ -189,7 +189,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions (method at program level) — with params - q = Query(_LANGUAGE, "(program (method name: (identifier) @name parameters: (method_parameters) @params) @def)") + q = _query(_LANGUAGE, "(program (method name: (identifier) @name parameters: (method_parameters) @params) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "function", @@ -200,7 +200,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions — no params - q = Query(_LANGUAGE, "(program (method name: (identifier) @name) @def)") + q = _query(_LANGUAGE, "(program (method name: (identifier) @name) @def)") for _, m in _matches(q, tree.root_node): method_node = m["def"] if not any(child.type == "method_parameters" for child in method_node.children): @@ -241,7 +241,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(class name: (constant) @name) @def", "(module name: (constant) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 @@ -251,7 +251,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(method name: (identifier) @name) @def", "(singleton_method name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return source[node.start_byte:node.end_byte].decode("utf-8", errors="replace"), node.start_point[0] + 1 @@ -265,7 +265,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: "(method name: (identifier) @name) @def", "(singleton_method name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -274,7 +274,7 @@ def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: if fn_node is None: return [] - q = Query(_LANGUAGE, "(call method: (identifier) @called)") + q = _query(_LANGUAGE, "(call method: (identifier) @called)") calls = set() for _, m in _matches(q, fn_node): calls.add(m["called"].text.decode("utf-8", errors="replace")) @@ -285,7 +285,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: usages = [] seen = set() for node_type in ("identifier", "constant"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -298,7 +298,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(program (call method: (identifier) @method arguments: (argument_list (string) @path)) @imp)") + q = _query(_LANGUAGE, "(program (call method: (identifier) @method arguments: (argument_list (string) @path)) @imp)") for _, m in _matches(q, tree.root_node): method = m["method"].text.decode("utf-8", errors="replace") if method in ("require", "require_relative"): @@ -328,7 +328,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: "(method name: (identifier) @name) @def", "(singleton_method name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -419,7 +419,7 @@ def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: "(method name: (identifier) @name) @def", "(singleton_method name: (identifier) @name) @def", ]: - for _, m in _matches(Query(_LANGUAGE, q_str), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/rust.py b/src/codetree/languages/rust.py index 731f3ff..dd16d44 100644 --- a/src/codetree/languages/rust.py +++ b/src/codetree/languages/rust.py @@ -1,9 +1,9 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_rust as tsrust -from .base import LanguagePlugin, _matches, _fill_docs_from_siblings +from .base import LanguagePlugin, _matches, _fill_docs_from_siblings, _query, CachedParser _LANGUAGE = Language(tsrust.language()) -_PARSER = Parser(_LANGUAGE) +_PARSER = CachedParser(Parser(_LANGUAGE)) def _parse(source: bytes): @@ -18,7 +18,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: results = [] # Structs - q = Query(_LANGUAGE, "(source_file (struct_item name: (type_identifier) @name) @def)") + q = _query(_LANGUAGE, "(source_file (struct_item name: (type_identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "struct", @@ -29,7 +29,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Enums - q = Query(_LANGUAGE, "(source_file (enum_item name: (type_identifier) @name) @def)") + q = _query(_LANGUAGE, "(source_file (enum_item name: (type_identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "enum", @@ -40,7 +40,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Traits - q = Query(_LANGUAGE, "(source_file (trait_item name: (type_identifier) @name) @def)") + q = _query(_LANGUAGE, "(source_file (trait_item name: (type_identifier) @name) @def)") for _, m in _matches(q, tree.root_node): results.append({ "type": "trait", @@ -51,7 +51,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Methods inside impl blocks (both direct impl and trait impl) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (impl_item type: (type_identifier) @class_name body: (declaration_list @@ -69,7 +69,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Method signatures inside traits (fn without body) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (trait_item name: (type_identifier) @trait_name body: (declaration_list @@ -87,7 +87,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Default method implementations inside traits (fn with body) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (trait_item name: (type_identifier) @trait_name body: (declaration_list @@ -105,7 +105,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level functions (direct children of source_file) - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (source_file (function_item name: (identifier) @name @@ -145,7 +145,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N tree = _parse(source) # Functions (top-level and inside impl blocks) - q = Query(_LANGUAGE, "(function_item name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_item name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -157,7 +157,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(enum_item name: (type_identifier) @name) @def", "(trait_item name: (type_identifier) @name) @def", ]: - q = Query(_LANGUAGE, q_str) + q = _query(_LANGUAGE, q_str) for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] @@ -168,14 +168,14 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N def extract_calls_in_function(self, source: bytes, fn_name: str) -> list[str]: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_item name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_item name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break if fn_node is None: return [] - q = Query(_LANGUAGE, """ + q = _query(_LANGUAGE, """ (call_expression function: [ (identifier) @called (field_expression field: (field_identifier) @called) @@ -195,7 +195,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: # type positions (e.g. `let calc = Calculator;` uses identifier, # while `impl Calculator` uses type_identifier). for node_type in ("identifier", "type_identifier"): - q = Query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') + q = _query(_LANGUAGE, f'(({node_type}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) @@ -208,7 +208,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: def extract_imports(self, source: bytes) -> list[dict]: tree = _parse(source) results = [] - q = Query(_LANGUAGE, "(source_file (use_declaration) @imp)") + q = _query(_LANGUAGE, "(source_file (use_declaration) @imp)") for _, m in _matches(q, tree.root_node): node = m["imp"] results.append({ @@ -221,7 +221,7 @@ def extract_imports(self, source: bytes) -> list[dict]: def compute_complexity(self, source: bytes, fn_name: str) -> dict | None: tree = _parse(source) fn_node = None - q = Query(_LANGUAGE, "(function_item name: (identifier) @name) @def") + q = _query(_LANGUAGE, "(function_item name: (identifier) @name) @def") for _, m in _matches(q, tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] @@ -254,7 +254,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Find function node fn_node = None - for _, m in _matches(Query(_LANGUAGE, "(function_item name: (identifier) @name) @def"), tree.root_node): + for _, m in _matches(_query(_LANGUAGE, "(function_item name: (identifier) @name) @def"), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break diff --git a/src/codetree/languages/typescript.py b/src/codetree/languages/typescript.py index 47a3be5..8b5d1a5 100644 --- a/src/codetree/languages/typescript.py +++ b/src/codetree/languages/typescript.py @@ -1,12 +1,12 @@ -from tree_sitter import Language, Parser, Query +from tree_sitter import Language, Parser import tree_sitter_typescript as tsts from .javascript import JavaScriptPlugin, _arrow_params -from .base import _matches, _fill_docs_from_siblings +from .base import _matches, _fill_docs_from_siblings, _query, CachedParser _TS_LANGUAGE = Language(tsts.language_typescript()) -_TS_PARSER = Parser(_TS_LANGUAGE) +_TS_PARSER = CachedParser(Parser(_TS_LANGUAGE)) _TSX_LANGUAGE = Language(tsts.language_tsx()) -_TSX_PARSER = Parser(_TSX_LANGUAGE) +_TSX_PARSER = CachedParser(Parser(_TSX_LANGUAGE)) class TypeScriptPlugin(JavaScriptPlugin): @@ -32,7 +32,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(program (class_declaration name: (type_identifier) @name) @def)", "(program (export_statement (class_declaration name: (type_identifier) @name) @def))", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "class", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -46,7 +46,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(program (abstract_class_declaration name: (type_identifier) @name) @def)", "(program (export_statement (abstract_class_declaration name: (type_identifier) @name) @def))", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "class", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -70,7 +70,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (property_identifier) @method_name parameters: (formal_parameters) @params)))""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "method", "name": m["method_name"].text.decode("utf-8", errors="replace"), @@ -80,7 +80,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Abstract method signatures (no body, e.g. abstract doWork(): void;) - q = Query(lang, """ + q = _query(lang, """ (abstract_class_declaration name: (type_identifier) @class_name body: (class_body @@ -98,7 +98,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # Top-level function declarations (same as JS — uses identifier) - q = Query(lang, """ + q = _query(lang, """ (program (function_declaration name: (identifier) @name parameters: (formal_parameters) @params)) @@ -117,7 +117,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(program (interface_declaration name: (type_identifier) @name) @def)", "(program (export_statement (interface_declaration name: (type_identifier) @name) @def))", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "interface", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -131,7 +131,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: "(program (type_alias_declaration name: (type_identifier) @name) @def)", "(program (export_statement (type_alias_declaration name: (type_identifier) @name) @def))", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): results.append({ "type": "type", "name": m["name"].text.decode("utf-8", errors="replace"), @@ -141,7 +141,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: }) # export default function foo() {} - q = Query(lang, """ + q = _query(lang, """ (program (export_statement (function_declaration name: (identifier) @name @@ -168,7 +168,7 @@ def extract_skeleton(self, source: bytes) -> list[dict]: name: (identifier) @name value: [(arrow_function) @fn (function_expression) @fn]))))""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): fn_node = m.get("fn") results.append({ "type": "function", @@ -220,7 +220,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(export_statement (abstract_class_declaration name: (type_identifier) @name) @def)", "(export_statement (interface_declaration name: (type_identifier) @name) @def)", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -239,7 +239,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N name: (identifier) @name value: [(arrow_function) (function_expression)])) @def)""", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -252,7 +252,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(type_alias_declaration name: (type_identifier) @name) @def", "(export_statement (type_alias_declaration name: (type_identifier) @name) @def)", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -265,7 +265,7 @@ def extract_symbol_source(self, source: bytes, name: str) -> tuple[str, int] | N "(method_definition name: (property_identifier) @name) @def", "(abstract_method_signature name: (property_identifier) @name) @def", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == name: node = m["def"] return ( @@ -286,7 +286,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: "(function_declaration name: (identifier) @name) @def", "(method_definition name: (property_identifier) @name) @def", ]: - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -296,7 +296,7 @@ def extract_variables(self, source: bytes, fn_name: str) -> list[dict]: # Arrow/function expressions if fn_node is None: q_str = "(variable_declarator name: (identifier) @name value: [(arrow_function) @def (function_expression) @def])" - for _, m in _matches(Query(lang, q_str), tree.root_node): + for _, m in _matches(_query(lang, q_str), tree.root_node): if m["name"].text.decode("utf-8", errors="replace") == fn_name: fn_node = m["def"] break @@ -391,7 +391,7 @@ def extract_symbol_usages(self, source: bytes, name: str) -> list[dict]: # Match both plain identifiers and type identifiers (class names appear as both) for node_kind in ("identifier", "type_identifier"): - q = Query(lang, f'(({node_kind}) @name (#eq? @name "{name}"))') + q = _query(lang, f'(({node_kind}) @name (#eq? @name "{name}"))') for _, m in _matches(q, tree.root_node): node = m["name"] key = (node.start_point[0], node.start_point[1]) diff --git a/tests/test_graph_builder.py b/tests/test_graph_builder.py index 7f9c5cc..048d322 100644 --- a/tests/test_graph_builder.py +++ b/tests/test_graph_builder.py @@ -290,3 +290,28 @@ def test_change_impact_min_weight(self, multi_module_repo): if "main.py::process" in all_callers: assert "other.py::process" not in all_callers store.close() + + +def test_build_uses_batched_store_access(tmp_path, monkeypatch): + """Row-by-row store calls stall a background build behind busy tool threads.""" + (tmp_path / "calc.py").write_text( + "class Calc:\n def add(self, a, b):\n return helper(a) + b\n\n" + "def helper(x):\n return x\n" + ) + (tmp_path / "test_calc.py").write_text("from calc import helper\n\ndef test_helper():\n helper(1)\n") + + def row_by_row(*args, **kwargs): + raise AssertionError("per-row store access during build") + + for method in ("upsert_symbol", "upsert_edge", "upsert_file", "symbols_by_name", "get_file"): + monkeypatch.setattr(GraphStore, method, row_by_row) + + store = GraphStore(str(tmp_path)) + store.open() + try: + stats = GraphBuilder(str(tmp_path), store).build() + assert stats["files_indexed"] == 2 + assert {e.type for e in store.edges_from("calc.py::Calc")} == {"CONTAINS"} + assert [e.target_qn for e in store.edges_from("test_calc.py::test_helper", "TESTS")] == ["calc.py::helper"] + finally: + store.close() diff --git a/tests/test_graph_store.py b/tests/test_graph_store.py index 6f6ff64..a9df6de 100644 --- a/tests/test_graph_store.py +++ b/tests/test_graph_store.py @@ -177,3 +177,37 @@ def test_stats(self, store): assert stats["symbols"] == 2 assert stats["edges"] == 1 + +class TestBatchWrites: + def test_upsert_edges_last_duplicate_wins_like_row_by_row(self, store): + store.upsert_edges([ + Edge("a.py::f", "b.py::g", "CALLS", weight=0.5), + Edge("a.py::f", "b.py::g", "CALLS", weight=1.0), + ]) + assert [e.weight for e in store.edges_from("a.py::f")] == [1.0] + + def test_batches_larger_than_one_statement(self, store): + syms = [SymbolNode(f"m.py::f{i}", f"f{i}", "function", "m.py", i + 1, None) for i in range(500)] + edges = [Edge(f"m.py::f{i}", f"m.py::f{i + 1}", "CALLS") for i in range(499)] + store.upsert_symbols(syms) + store.upsert_edges(edges) + assert store.stats()["symbols"] == 500 + assert store.stats()["edges"] == 499 + + def test_symbols_by_name_map_matches_symbols_by_name(self, store): + store.upsert_symbols([ + SymbolNode("b.py::add", "add", "function", "b.py", 1, None), + SymbolNode("a.py::add", "add", "function", "a.py", 3, None), + SymbolNode("a.py::sub", "sub", "function", "a.py", 5, None), + ]) + by_name = store.symbols_by_name_map() + for name in ("add", "sub"): + assert [s.qualified_name for s in by_name[name]] == [ + s.qualified_name for s in store.symbols_by_name(name) + ] + + def test_empty_batches_are_noops(self, store): + store.upsert_symbols([]) + store.upsert_edges([]) + store.upsert_files([]) + assert store.stats() == {"files": 0, "symbols": 0, "edges": 0} diff --git a/tests/test_parse_cache.py b/tests/test_parse_cache.py new file mode 100644 index 0000000..704a833 --- /dev/null +++ b/tests/test_parse_cache.py @@ -0,0 +1,67 @@ +"""Tests for the compiled-query cache and the per-thread parse-tree cache.""" + +import threading +from collections import OrderedDict + +from tree_sitter import Language, Parser +import tree_sitter_python as tspython + +from codetree.languages import base +from codetree.languages.base import CachedParser, _query +from codetree.languages.python import PythonPlugin + +_LANG = Language(tspython.language()) + + +def test_query_is_compiled_once_per_pattern(): + pattern = "(function_definition name: (identifier) @name) @def" + assert _query(_LANG, pattern) is _query(_LANG, pattern) + + +def test_query_cache_is_bounded(monkeypatch): + monkeypatch.setattr(base, "_QUERY_CACHE_MAX", 3) + monkeypatch.setattr(base, "_QUERY_CACHE", OrderedDict()) # leave the shared cache intact + for index in range(10): + _query(_LANG, f'((identifier) @name (#eq? @name "n{index}"))') + assert len(base._QUERY_CACHE) == 3 + + +def test_cached_parser_reuses_tree_for_same_source(): + parser = CachedParser(Parser(_LANG)) + source = b"def foo():\n pass\n" + assert parser.parse(source) is parser.parse(bytes(source)) + + +def test_cached_parser_never_returns_stale_tree(): + parser = CachedParser(Parser(_LANG)) + first = parser.parse(b"def foo():\n pass\n") + second = parser.parse(b"def bar():\n pass\n") + assert first is not second + assert b"bar" in second.root_node.text + + +def test_cached_parser_is_bounded(): + parser = CachedParser(Parser(_LANG)) + for index in range(CachedParser._MAX_TREES + 5): + parser.parse(f"x{index} = 1\n".encode()) + assert len(parser._local.trees) == CachedParser._MAX_TREES + + +def test_cached_parser_does_not_share_trees_across_threads(): + parser = CachedParser(Parser(_LANG)) + source = b"def foo():\n pass\n" + main_tree = parser.parse(source) + other = {} + thread = threading.Thread(target=lambda: other.setdefault("tree", parser.parse(source))) + thread.start() + thread.join() + assert other["tree"] is not main_tree + + +def test_plugin_results_unchanged_when_calls_hit_cache(): + plugin = PythonPlugin() + source = b"def a():\n b()\n c()\n\ndef b():\n c()\n\ndef c():\n pass\n" + assert plugin.extract_calls_in_function(source, "a") == ["b", "c"] + assert plugin.extract_calls_in_function(source, "b") == ["c"] + assert plugin.extract_calls_in_function(source, "c") == [] + assert [item["name"] for item in plugin.extract_skeleton(source)] == ["a", "b", "c"] From 866acb6e0331be6a5c1dbdf55d43563633c18d77 Mon Sep 17 00:00:00 2001 From: Kevin Robatel Date: Thu, 1 Oct 2026 15:04:23 +0200 Subject: [PATCH 3/3] feat: discover files with git ls-files and skip nested worktrees Discovery used rglob("*"), which walked into every directory (including .git and node_modules before filtering) and ignored .gitignore. Nested git worktrees, such as the ones Claude Code creates under .claude/worktrees/, were indexed as part of the repository, so every symbol showed up once per worktree in find_references, resolve_symbol and friends, and the index was several times larger than needed. - In a git work tree, files come from `git ls-files`: tracked files are always indexed (.gitignore decides, so a tracked build/ or env/ is included), and untracked, non-ignored files are also checked against SKIP_DIRS so an un-ignored .venv or node_modules is never crawled. Git does not descend into nested worktrees, repositories or submodules. - Without git (no repository, git missing or refusing the repository, root ignored by a parent repository), os.walk prunes SKIP_DIRS and directories whose .git is a file (worktrees, submodules) before descending. A root that groups several full repositories still indexes them. - Paths are decoded with os.fsdecode (non-UTF-8 names), de-duplicated (unmerged paths are listed once per stage) and sorted. The index keeps that order, so results no longer depend on which files came from the cache. - Only discovered files are re-injected from the skeleton cache, and the cache is rewritten from the index. Stale entries for deleted or newly ignored files never come back. On a repository with eight Claude Code worktrees, this indexes ~9 500 files instead of ~86 000. Co-Authored-By: Claude Opus 5.5 --- AGENTS.md | 5 +- CLAUDE.md | 5 +- README.md | 13 ++ docs/LANDING_PAGE.md | 1 + src/codetree/indexer.py | 164 +++++++++++++++---- src/codetree/server.py | 34 ++-- tests/test_file_discovery.py | 299 +++++++++++++++++++++++++++++++++++ 7 files changed, 468 insertions(+), 53 deletions(-) create mode 100644 tests/test_file_discovery.py diff --git a/AGENTS.md b/AGENTS.md index d240df6..01ac088 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -92,7 +92,7 @@ MCP tool call → server.py → indexer.py → FileEntry.plugin → tree-sitter | File | Responsibility | |---|---| | `server.py` | FastMCP 3.1.0 server — defines all 23 tools, wires cache + indexer + graph at startup. Language-unaware. | -| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Skips `.venv`, `node_modules`, `__pycache__`, `.git`, etc. | +| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Discovers files with `git ls-files` (tracked files always; untracked non-ignored files minus `SKIP_DIRS`; nested worktrees/repos/submodules excluded), falling back to an `os.walk` that prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, etc.) and nested worktrees. | | `cache.py` | `.codetree/index.json` — stores pre-computed skeletons with mtime-based invalidation. Language-unaware. | | `registry.py` | Maps file extensions → plugin instances. The **only** place languages are registered. | @@ -142,6 +142,7 @@ Each plugin implements: |---|---| | `test_server.py` | Original 4 MCP tools via FastMCP, output format, line accuracy, cross-language | | `test_indexer.py` | Build, skip-dirs, skeleton/symbol/refs/callgraph through indexer layer | +| `test_file_discovery.py` | git-based discovery (.gitignore, nested worktrees/repos, submodules), walk fallback, stale cache entries | | `test_parse_cache.py` | Memoized query compilation and per-thread parse-tree cache | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | @@ -205,6 +206,6 @@ The tree-sitter Python bindings have breaking changes from older docs: - Plugin classes: `{Lang}Plugin` (e.g., `PythonPlugin`, `GoPlugin`) - Module-level parser/language globals: `_PARSER`, `_LANGUAGE` - Skeleton results are deduplicated by `(name, line)` and sorted by line number -- Indexer `SKIP_DIRS` includes `.venv`, `node_modules`, `__pycache__`, `.git` — without this, crawling `.venv` causes Codex timeout +- File discovery: in a git work tree, tracked files (`git ls-files --cached`) are always indexed — `.gitignore` decides, so a tracked `build/` or `env/` is indexed — while untracked non-ignored files (`--others --exclude-standard`) also skip `SKIP_DIRS`, so an un-ignored `.venv`/`node_modules` is never crawled. Nested worktrees such as `.claude/worktrees/`, nested repos and submodules never leak in, and `.codetree/` is always excluded. Outside git (no repo, git missing or refusing the repo, root ignored), the walk prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, …) and directories whose `.git` is a file (worktrees, submodules) — without this, crawling `.venv` causes Codex timeout. Only discovered files are re-injected from the cache. - FastMCP tool access in tests: `mcp.local_provider._components[f"tool:{name}@"].fn` - **Doc sync rule**: When tools are added, removed, or changed, update all 5 doc files: `README.md`, `TOOLS_GUIDE.md`, `LANDING_PAGE.md`, `CLAUDE.md`, `AGENTS.md` diff --git a/CLAUDE.md b/CLAUDE.md index d650910..a881988 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -99,7 +99,7 @@ MCP tool call → server.py → indexer.py → FileEntry.plugin → tree-sitter | File | Responsibility | |---|---| | `server.py` | FastMCP 3.1.0 server — defines the 23 tools, wires cache + indexer + graph at startup. Language-unaware. | -| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Skips `.venv`, `node_modules`, `__pycache__`, `.git`, etc. | +| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Discovers files with `git ls-files` (tracked files always; untracked non-ignored files minus `SKIP_DIRS`; nested worktrees/repos/submodules excluded), falling back to an `os.walk` that prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, etc.) and nested worktrees. | | `cache.py` | `.codetree/index.json` — stores pre-computed skeletons with mtime-based invalidation. Language-unaware. | | `registry.py` | Maps file extensions → plugin instances. The **only** place languages are registered. | @@ -148,6 +148,7 @@ Each plugin implements: |---|---| | `test_server.py` | Original 4 MCP tools via FastMCP, output format, line accuracy, cross-language | | `test_indexer.py` | Build, skip-dirs, skeleton/symbol/refs/callgraph through indexer layer | +| `test_file_discovery.py` | git-based discovery (.gitignore, nested worktrees/repos, submodules), walk fallback, stale cache entries | | `test_parse_cache.py` | Memoized query compilation and per-thread parse-tree cache | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | @@ -213,7 +214,7 @@ The tree-sitter Python bindings have breaking changes from older docs: - Plugin classes: `{Lang}Plugin` (e.g., `PythonPlugin`, `GoPlugin`) - Module-level parser/language globals: `_PARSER`, `_LANGUAGE` - Skeleton results are deduplicated by `(name, line)` and sorted by line number -- Indexer `SKIP_DIRS` includes `.venv`, `node_modules`, `__pycache__`, `.git` — without this, crawling `.venv` causes Claude Code timeout +- File discovery: in a git work tree, tracked files (`git ls-files --cached`) are always indexed — `.gitignore` decides, so a tracked `build/` or `env/` is indexed — while untracked non-ignored files (`--others --exclude-standard`) also skip `SKIP_DIRS`, so an un-ignored `.venv`/`node_modules` is never crawled. Nested worktrees such as `.claude/worktrees/`, nested repos and submodules never leak in, and `.codetree/` is always excluded. Outside git (no repo, git missing or refusing the repo, root ignored), the walk prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, …) and directories whose `.git` is a file (worktrees, submodules) — without this, crawling `.venv` causes Claude Code timeout. Only discovered files are re-injected from the cache. - FastMCP tool access in tests: `mcp.local_provider._components[f"tool:{name}@"].fn` - **Doc sync rule**: When tools are added, removed, or changed, update all 5 doc files: `README.md`, `TOOLS_GUIDE.md`, `LANDING_PAGE.md`, `CLAUDE.md`, `AGENTS.md` diff --git a/README.md b/README.md index 4e5821a..4e0ad6c 100644 --- a/README.md +++ b/README.md @@ -238,6 +238,19 @@ Add to `~/Library/Application Support/Claude/claude_desktop_config.json` (macOS) | **SCIP / LSIF indexers** | Slow builds, complex setup, huge indexes | ~1s startup, JSON cache, zero config | | **AST-only tools** | Raw trees are verbose and hard for agents | Pre-structured output designed for agents | +## What Gets Indexed + +- **Git repositories:** every file `git ls-files` reports — tracked files plus + untracked files that `.gitignore` does not exclude. Tracked directories are + indexed even if named `build/`, `dist/` or `env/`; untracked `.venv/`, + `node_modules/`, `__pycache__/` and similar are skipped even when nobody + ignored them. Nested worktrees (e.g. `.claude/worktrees/`), nested + repositories and submodules are not indexed. +- **Without git** (not a repository, `git` missing, or git refusing the + repository, e.g. `safe.directory`): the tree is walked, skipping `.venv`, + `node_modules`, `__pycache__`, `.git`, `dist`, `build`, … and nested worktrees. +- `.codetree/` (codetree's own cache) is never indexed. + ## Architecture ``` diff --git a/docs/LANDING_PAGE.md b/docs/LANDING_PAGE.md index bbf2d04..04122d1 100644 --- a/docs/LANDING_PAGE.md +++ b/docs/LANDING_PAGE.md @@ -365,6 +365,7 @@ claude mcp add codetree -- uvx --from mcp-server-codetree codetree --root . - **FastMCP** for the MCP protocol — stdio transport, zero network config. - **Plugin architecture** — each language is a self-contained class implementing 5 core methods. Adding a language is copying a template. - **Smart caching** — `.codetree/index.json` with mtime-based invalidation. Unchanged files skip parsing entirely. +- **Respects `.gitignore`** — files are discovered via `git ls-files`, so ignored build output and nested git worktrees (e.g. `.claude/worktrees/`) are never indexed. - **Lazy call graph** — only built when tools like `find_dead_code` or `get_blast_radius` are first called. Stored in memory, O(1) lookup. - **PageRank** — standard algorithm (25 iterations, damping 0.85) for ranking symbol importance by reference count. - **Clone detection** — AST normalization (identifiers → `_ID_`, strings → `_STR_`, numbers → `_NUM_`) + SHA-256 hashing. Catches exact copies and renamed-variable copies. diff --git a/src/codetree/indexer.py b/src/codetree/indexer.py index 0fb2256..fe0a158 100644 --- a/src/codetree/indexer.py +++ b/src/codetree/indexer.py @@ -1,3 +1,5 @@ +import os +import subprocess from pathlib import Path from dataclasses import dataclass from .languages.base import LanguagePlugin @@ -43,6 +45,12 @@ def __init__(self, root: str | Path): self._call_graph: dict[str, set[str]] = {} self._reverse_graph: dict[str, set[str]] = {} self._call_graph_built: bool = False + # Rel paths in discovery order; _rebuild_definitions() orders _index by + # it so results do not depend on which files came from the cache. + self._discovery_order: list[str] = [] + # Rel paths that build() discovered but skipped because the caller's + # cache already had them at the same mtime — to be injected from cache. + self.cached_candidates: list[str] = [] @property def files(self) -> list[Path]: @@ -64,6 +72,84 @@ def _should_skip(self, path: Path) -> bool: return True return False + def _git_ls_files(self, *args: str) -> list[Path] | None: + """Run `git ls-files -z ` in root. None if git cannot be used.""" + try: + result = subprocess.run( + ["git", "-C", str(self.root), "ls-files", "-z", *args], + capture_output=True, timeout=30, + ) + except (OSError, subprocess.SubprocessError): + return None + if result.returncode != 0: + return None + # Split raw bytes and fsdecode so non-UTF-8 file names round-trip. + return [self.root / os.fsdecode(raw) for raw in result.stdout.split(b"\0") if raw] + + def _git_files(self) -> tuple[list[Path], list[Path]] | None: + """(tracked, untracked-but-not-ignored) files via git. None if git cannot be used. + + Git does not descend into nested repositories, so worktrees and + submodules living under the root are excluded even when not ignored. + """ + tracked = self._git_ls_files("--cached") + if tracked is None: + return None + untracked = self._git_ls_files("--others", "--exclude-standard") + if untracked is None: + return None + if not tracked and not untracked: + # Not inside a work tree in a useful way (e.g. root itself is ignored). + return None + return tracked, untracked + + def _walk_files(self) -> list[Path]: + """List files by walking the tree, pruning before descending. + + Prunes SKIP_DIRS and, like git, nested worktrees and submodules: a + directory whose `.git` is a file. Directories holding a full `.git` + repository are kept, so a root that groups several repos still works. + """ + files = [] + for dirpath, dirnames, filenames in os.walk(self.root): + dirnames[:] = [ + d for d in dirnames + if d not in self.SKIP_DIRS + and not d.endswith(".egg-info") + and not os.path.isfile(os.path.join(dirpath, d, ".git")) + ] + files.extend(Path(dirpath) / name for name in filenames) + return files + + def discover_files(self) -> list[Path]: + """Return supported source files under root, sorted and de-duplicated. + + In a git work tree, every tracked file is a candidate (.gitignore + decides), while untracked files also skip SKIP_DIRS — a `.venv` or + `node_modules` nobody ignored. Outside git, the tree is walked, + skipping SKIP_DIRS and nested worktrees. `.codetree/` is never indexed. + """ + git_files = self._git_files() + if git_files is not None: + tracked, untracked = git_files + candidates = [(path, False) for path in tracked] + candidates += [(path, True) for path in untracked] + else: + candidates = [(path, True) for path in self._walk_files()] + files = set() + for candidate, apply_skip_dirs in candidates: + if get_plugin(candidate) is None: + continue + if candidate.is_symlink() or not candidate.is_file(): + continue + rel = candidate.relative_to(self.root) + if rel.parts and rel.parts[0] == ".codetree": + continue + if apply_skip_dirs and self._should_skip(rel): + continue + files.add(candidate) # a set: unmerged paths are listed once per stage + return sorted(files) + def _rebuild_definitions(self) -> None: """Rebuild _definitions from current _index using qualified (file::name) keys. @@ -75,6 +161,10 @@ def _rebuild_definitions(self) -> None: Also rebuilds _name_to_qualified secondary index for O(1) callee lookup in _ensure_call_graph(). """ + if self._discovery_order: + ordered = {rel: self._index[rel] for rel in self._discovery_order if rel in self._index} + ordered.update(self._index) # entries injected outside discovery keep their place at the end + self._index = ordered self._definitions = {} self._name_to_qualified = {} for rel_path, entry in self._index.items(): @@ -90,6 +180,36 @@ def _rebuild_definitions(self) -> None: if key not in self._name_to_qualified[bare]: self._name_to_qualified[bare].append(key) + def index_file(self, path: Path, mtime: float | None = None) -> FileEntry | None: + """Parse one file into a FileEntry. None if unsupported or unreadable.""" + plugin = get_plugin(path) + if plugin is None: + return None + try: + if mtime is None: + mtime = path.stat().st_mtime + source = path.read_bytes() + except OSError: + return None # deleted or unreadable since discovery + try: + skeleton = plugin.extract_skeleton(source) + has_errors = plugin.check_syntax(source) + except Exception: + # ROBUST-01: Plugin crash (any exception) skips this file gracefully. + # File is still added to _index with empty skeleton and has_errors=True + # so tools that check entry.has_errors can warn the caller. + skeleton = [] + has_errors = True + return FileEntry( + path=path, + source=source, + skeleton=skeleton, + mtime=mtime, + language=path.suffix.lstrip("."), + plugin=plugin, + has_errors=has_errors, + ) + def build(self, cached_mtimes: dict[str, float] | None = None): """Index all supported files under root, skipping non-project dirs. @@ -97,39 +217,21 @@ def build(self, cached_mtimes: dict[str, float] | None = None): the caller injects them via inject_cached(). """ cached_mtimes = cached_mtimes or {} - for candidate in self.root.rglob("*"): - if candidate.is_symlink(): - continue - if not candidate.is_file(): - continue - plugin = get_plugin(candidate) - if plugin is None: - continue - if self._should_skip(candidate.relative_to(self.root)): - continue + self.cached_candidates = [] + files = self.discover_files() + self._discovery_order = [str(f.relative_to(self.root)) for f in files] + for candidate in files: rel = str(candidate.relative_to(self.root)) - mtime = candidate.stat().st_mtime - if cached_mtimes.get(rel) == mtime: - continue - source = candidate.read_bytes() try: - skeleton = plugin.extract_skeleton(source) - has_errors = plugin.check_syntax(source) - except Exception: - # ROBUST-01: Plugin crash (any exception) skips this file gracefully. - # File is still added to _index with empty skeleton and has_errors=True - # so tools that check entry.has_errors can warn the caller. - skeleton = [] - has_errors = True - self._index[rel] = FileEntry( - path=candidate, - source=source, - skeleton=skeleton, - mtime=mtime, - language=candidate.suffix.lstrip("."), - plugin=plugin, - has_errors=has_errors, - ) + mtime = candidate.stat().st_mtime + except OSError: + mtime = None + if mtime is not None and cached_mtimes.get(rel) == mtime: + self.cached_candidates.append(rel) + elif mtime is not None: + entry = self.index_file(candidate, mtime) + if entry is not None: + self._index[rel] = entry # Build definition index from skeleton data (qualified keys, no duplicates, no ghosts) self._rebuild_definitions() diff --git a/src/codetree/server.py b/src/codetree/server.py index 189e22f..5593990 100644 --- a/src/codetree/server.py +++ b/src/codetree/server.py @@ -33,28 +33,26 @@ def _validate_path(file_path: str | None, _root: Path = root_path) -> str | None indexer = Indexer(root) indexer.build(cached_mtimes=cached_mtimes) - # Inject cached entries for unchanged files (skip ignored dirs) - indexed = {str(f.relative_to(root_path)) for f in indexer.files} - for rel_path, entry_data in (cache._data or {}).items(): - if indexer._should_skip(Path(rel_path)): - continue - if rel_path not in indexed: - py_file = root_path / rel_path - if py_file.exists(): - mtime = py_file.stat().st_mtime - if cache.is_valid(rel_path, mtime): - indexer.inject_cached( - rel_path=rel_path, - py_file=py_file, - source=py_file.read_bytes(), - skeleton=entry_data.get("skeleton", []), - mtime=mtime, - ) + # Inject cached entries for unchanged files. Only files discovered by this + # build qualify, so stale cache entries (deleted or now-ignored files, + # e.g. old worktrees) are never resurrected. + for rel_path in indexer.cached_candidates: + py_file = root_path / rel_path + mtime = py_file.stat().st_mtime + if cache.is_valid(rel_path, mtime): + indexer.inject_cached( + rel_path=rel_path, + py_file=py_file, + source=py_file.read_bytes(), + skeleton=cache.get(rel_path).get("skeleton", []), + mtime=mtime, + ) # Rebuild definition index once after all injections (DATA-01, DATA-02, DATA-03 fix) indexer._rebuild_definitions() - # Save updated cache + # Save updated cache — rebuilt from the index so stale entries are dropped + cache._data = {} for rel_path, file_entry in indexer._index.items(): cache.set(rel_path, { "mtime": file_entry.mtime, diff --git a/tests/test_file_discovery.py b/tests/test_file_discovery.py new file mode 100644 index 0000000..cfcb9e8 --- /dev/null +++ b/tests/test_file_discovery.py @@ -0,0 +1,299 @@ +"""Tests for source file discovery: .gitignore, nested worktrees/repos, walk fallback.""" + +import json +import os +import subprocess +from pathlib import Path + +import pytest + +from codetree.indexer import Indexer +from codetree.server import create_server + + +def _git(cwd, *args): + subprocess.run( + ["git", "-c", "user.email=test@test.com", "-c", "user.name=Test", + "-c", "protocol.file.allow=always", *args], + cwd=cwd, capture_output=True, text=True, timeout=30, check=True, + ) + + +def _rel_files(indexer: Indexer) -> set[str]: + return {str(f.relative_to(indexer.root)) for f in indexer.files} + + +def _make_repo(path: Path) -> Path: + path.mkdir(parents=True, exist_ok=True) + _git(path, "init", "-q") + (path / "app.py").write_text("def app():\n pass\n") + _git(path, "add", "-A") + _git(path, "commit", "-q", "-m", "init") + return path + + +# ── Git mode ───────────────────────────────────────────────────────────────── + +def test_gitignored_files_are_not_indexed(tmp_path): + repo = _make_repo(tmp_path / "repo") + (repo / ".gitignore").write_text("generated/\ntmp_*.py\n") + (repo / "generated").mkdir() + (repo / "generated" / "gen.py").write_text("def gen():\n pass\n") + (repo / "tmp_scratch.py").write_text("def scratch():\n pass\n") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +def test_untracked_but_not_ignored_files_are_indexed(tmp_path): + repo = _make_repo(tmp_path / "repo") + (repo / "new_module.py").write_text("def fresh():\n pass\n") + + indexer = Indexer(repo) + indexer.build() + + assert "new_module.py" in _rel_files(indexer) + + +def test_nested_worktree_is_not_indexed_even_if_not_ignored(tmp_path): + repo = _make_repo(tmp_path / "repo") + _git(repo, "worktree", "add", "-q", ".claude/worktrees/feature", "-b", "feature") + assert (repo / ".claude" / "worktrees" / "feature" / "app.py").exists() + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +def test_nested_repository_is_not_indexed(tmp_path): + repo = _make_repo(tmp_path / "repo") + _make_repo(repo / "vendor_repo") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +def test_submodule_is_not_indexed(tmp_path): + upstream = _make_repo(tmp_path / "upstream") + repo = _make_repo(tmp_path / "repo") + _git(repo, "submodule", "add", "-q", str(upstream), "libs/upstream") + _git(repo, "commit", "-q", "-m", "add submodule") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +def test_root_inside_worktree_indexes_that_worktree(tmp_path): + repo = _make_repo(tmp_path / "repo") + _git(repo, "worktree", "add", "-q", ".claude/worktrees/feature", "-b", "feature") + worktree = repo / ".claude" / "worktrees" / "feature" + (worktree / "feature_only.py").write_text("def feature():\n pass\n") + + indexer = Indexer(worktree) + indexer.build() + + assert _rel_files(indexer) == {"app.py", "feature_only.py"} + + +def test_deleted_but_still_tracked_file_is_skipped(tmp_path): + repo = _make_repo(tmp_path / "repo") + (repo / "app.py").unlink() # deleted in work tree, still in git index + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == set() + + +def test_codetree_dir_never_indexed_in_git_mode(tmp_path): + repo = _make_repo(tmp_path / "repo") + (repo / ".codetree").mkdir() + (repo / ".codetree" / "leftover.py").write_text("def leftover():\n pass\n") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +# ── Walk fallback ──────────────────────────────────────────────────────────── + +def test_walk_fallback_without_git_skips_skip_dirs(tmp_path): + root = tmp_path / "plain" + root.mkdir() + (root / "main.py").write_text("def main():\n pass\n") + for skipped in ("node_modules", ".venv", "pkg.egg-info"): + (root / skipped).mkdir() + (root / skipped / "dep.py").write_text("def dep():\n pass\n") + + indexer = Indexer(root) + assert indexer._git_files() is None + indexer.build() + + assert _rel_files(indexer) == {"main.py"} + + +def test_root_ignored_by_parent_repo_falls_back_to_walk(tmp_path): + parent = _make_repo(tmp_path / "parent") + (parent / ".gitignore").write_text("ignored_area/\n") + root = parent / "ignored_area" + root.mkdir() + (root / "lib.py").write_text("def lib():\n pass\n") + + indexer = Indexer(root) + indexer.build() + + assert _rel_files(indexer) == {"lib.py"} + + +def test_walk_fallback_when_git_is_unavailable(tmp_path, monkeypatch): + repo = _make_repo(tmp_path / "repo") + (repo / "node_modules").mkdir() + (repo / "node_modules" / "dep.js").write_text("function dep() {}\n") + + def _no_git(*args, **kwargs): + raise FileNotFoundError("git") + + monkeypatch.setattr("codetree.indexer.subprocess.run", _no_git) + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py"} + + +def test_discovered_files_are_sorted(tmp_path): + root = tmp_path / "plain" + root.mkdir() + for name in ("zeta.py", "alpha.py", "mid.py"): + (root / name).write_text("x = 1\n") + + files = Indexer(root).discover_files() + + assert files == sorted(files) + + +# ── Cache interaction ──────────────────────────────────────────────────────── + +def test_stale_cache_entries_for_ignored_files_are_not_resurrected(tmp_path): + repo = _make_repo(tmp_path / "repo") + ignored = repo / "old_worktree" / "ghost.py" + ignored.parent.mkdir() + ignored.write_text("def ghost():\n pass\n") + + # First startup indexes ghost.py and caches it + create_server(str(repo)) + cache_file = repo / ".codetree" / "index.json" + assert "old_worktree/ghost.py" in json.loads(cache_file.read_text()) + + # Now the directory is ignored; the valid cache entry must not come back + (repo / ".gitignore").write_text("old_worktree/\n.codetree/\n") + mcp = create_server(str(repo)) + find_refs = mcp.local_provider._components["tool:find_references@"].fn + + assert "ghost.py" not in find_refs(symbol_name="ghost") + assert "old_worktree/ghost.py" not in json.loads(cache_file.read_text()) + + +def test_git_mode_indexes_tracked_files_in_skip_dirs(tmp_path): + repo = _make_repo(tmp_path / "repo") + for tracked in ("build", "dist", "env"): + (repo / tracked).mkdir() + (repo / tracked / "mod.py").write_text("def mod():\n pass\n") + _git(repo, "add", "build", "dist", "env") + _git(repo, "commit", "-q", "-m", "tracked build output") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py", "build/mod.py", "dist/mod.py", "env/mod.py"} + + +def test_git_mode_skips_untracked_unignored_venv_and_node_modules(tmp_path): + repo = _make_repo(tmp_path / "repo") # no .gitignore at all + (repo / ".venv" / "lib" / "site-packages" / "pkg").mkdir(parents=True) + (repo / ".venv" / "lib" / "site-packages" / "pkg" / "m.py").write_text("def m():\n pass\n") + (repo / "venv" / "lib").mkdir(parents=True) + (repo / "venv" / "lib" / "dep.py").write_text("def dep():\n pass\n") + (repo / "node_modules" / "x").mkdir(parents=True) + (repo / "node_modules" / "x" / "i.js").write_text("function i() {}\n") + (repo / "fresh.py").write_text("def fresh():\n pass\n") + + indexer = Indexer(repo) + indexer.build() + + assert _rel_files(indexer) == {"app.py", "fresh.py"} + + +def test_walk_fallback_skips_worktrees_but_keeps_nested_repositories(tmp_path): + root = tmp_path / "plain" # not a repo itself, e.g. a folder grouping several repos + root.mkdir() + (root / "main.py").write_text("def main():\n pass\n") + nested_repo = root / "vendor_repo" + (nested_repo / ".git").mkdir(parents=True) + (nested_repo / ".git" / "hooks.py").write_text("x = 1\n") + (nested_repo / "lib.py").write_text("def lib():\n pass\n") + worktree = root / ".claude" / "worktrees" / "feature" + worktree.mkdir(parents=True) + (worktree / ".git").write_text("gitdir: /elsewhere/.git/worktrees/feature\n") + (worktree / "main.py").write_text("def main():\n pass\n") + + indexer = Indexer(root) + assert indexer._git_files() is None + indexer.build() + + assert _rel_files(indexer) == {"main.py", "vendor_repo/lib.py"} + + +def test_non_utf8_file_names_are_indexed_in_git_mode(tmp_path): + repo = _make_repo(tmp_path / "repo") + name = os.fsdecode(b"caf\xe9.py") + (repo / name).write_text("def cafe():\n pass\n") + + indexer = Indexer(repo) + indexer.build() + + assert name in _rel_files(indexer) + + +def test_conflicted_file_is_discovered_once(tmp_path): + repo = _make_repo(tmp_path / "repo") + _git(repo, "checkout", "-q", "-b", "other") + (repo / "app.py").write_text("def app():\n return 1\n") + _git(repo, "commit", "-q", "-am", "other") + _git(repo, "checkout", "-q", "-") + (repo / "app.py").write_text("def app():\n return 2\n") + _git(repo, "commit", "-q", "-am", "main") + subprocess.run(["git", "merge", "-q", "other"], cwd=repo, capture_output=True) # conflicts + + files = Indexer(repo).discover_files() + + assert [str(f.relative_to(repo)) for f in files] == ["app.py"] + + +def test_index_order_does_not_depend_on_cache(tmp_path): + root = tmp_path / "plain" + root.mkdir() + for name in ("a.py", "b.py", "c.py", "d.py"): + (root / name).write_text("def target():\n pass\n") + first = Indexer(root) + first.build() + cached = {rel: e.mtime for rel, e in first._index.items() if rel != "c.py"} + + # c.py is re-parsed, the others come from the "cache": order must not change + second = Indexer(root) + second.build(cached_mtimes=cached) + for rel in second.cached_candidates: + entry = first._index[rel] + second.inject_cached(rel, entry.path, entry.source, entry.skeleton, entry.mtime) + second._rebuild_definitions() + + assert list(second._index) == ["a.py", "b.py", "c.py", "d.py"] + assert [r["file"] for r in second.find_references("target")] == ["a.py", "b.py", "c.py", "d.py"] +