From 94d13b59d29c5ad3c5f2e59a522b37e437ca75cb Mon Sep 17 00:00:00 2001 From: Kevin Robatel Date: Thu, 1 Oct 2026 15:02:54 +0200 Subject: [PATCH 1/4] 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/4] 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/4] 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"] + From b817232761e0bf1976ad2c238a3fea481c27b09b Mon Sep 17 00:00:00 2001 From: Kevin Robatel Date: Thu, 1 Oct 2026 15:04:57 +0200 Subject: [PATCH 4/4] feat: answer the MCP handshake immediately and index in the background All indexing ran inside create_server(), before mcp.run() answered the MCP handshake. On a large repository a cold build took longer than the client's startup timeout, so the server never came up. Indexing and the graph build now run in a background thread owned by the new IndexState (src/codetree/index_state.py): - Single-file tools (get_file_skeleton, get_symbol, get_imports, get_skeletons, get_symbols, get_complexity, analyze_dataflow flow/taint) answer right away, parsing the requested files on demand (only files discovery would index). - Repo-wide and graph tools wait up to 20 s (CODETREE_WAIT_TIMEOUT), then return a "still building" message with progress. - index_status never blocks. It adds status, files_discovered, files_indexed, index_ready, graph_ready, startup_seconds and error; during a rebuild it reports the last committed graph. - Failures release every waiter. A graph failure rolls back and leaves structural tools working. create_server(root) stays synchronous by default and re-raises build errors (tests, embedding); `codetree` uses create_server(root, background=True). No tool signature changes. Also: - Require fastmcp>=3.0.0, the first release that runs sync tools in a thread pool. On 2.x a waiting tool would block the event loop. - Cache writes are atomic, with a unique temp file per save, and keep the file mode. A failed save no longer fails indexing. The cache now stores has_errors, so syntax warnings survive warm starts. - Indexer: index_file(), build(files=, progress=), and a locked, build-then-publish lazy call graph for concurrent tool calls. - GraphStore.rollback(). On a ~9 500-file repository the handshake completes in about 1 s (previously more than 25 s). Co-Authored-By: Claude Opus 5.5 --- AGENTS.md | 10 +- CLAUDE.md | 14 +- CONTRIBUTING.md | 2 +- README.md | 10 +- docs/LANDING_PAGE.md | 3 +- docs/TOOLS_GUIDE.md | 22 ++- pyproject.toml | 2 +- src/codetree/cache.py | 22 ++- src/codetree/graph/store.py | 7 + src/codetree/index_state.py | 263 ++++++++++++++++++++++++++ src/codetree/indexer.py | 52 ++++-- src/codetree/server.py | 150 +++++++++------ tests/test_async_startup.py | 350 +++++++++++++++++++++++++++++++++++ tests/test_cache.py | 16 ++ tests/test_file_discovery.py | 26 +++ tests/test_graph_store.py | 11 ++ tests/test_indexer.py | 32 ++++ 17 files changed, 903 insertions(+), 89 deletions(-) create mode 100644 src/codetree/index_state.py create mode 100644 tests/test_async_startup.py diff --git a/AGENTS.md b/AGENTS.md index 01ac088..d9a6b69 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -23,7 +23,7 @@ It exposes **23 tools** over MCP: | `detect_clones(file_path?, min_lines?)` | Duplicate/near-duplicate functions | `Clone group 1 (2 functions, 12 lines each):` | | `search_symbols(query?, type?, parent?, ..., format?)` | Flexible symbol search; `format="compact"` omits doc lines | `calc.py: class Calculator → line 1` | | `find_tests(file_path, symbol_name)` | Find test functions for a symbol | `test_calc.py: test_add() → line 3 (name match)` | -| `index_status()` | Graph index freshness and stats | `{files: 42, symbols: 315, edges: 580}` | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | `{graph_exists, files, symbols, edges, last_indexed_at, status, files_discovered, files_indexed, index_ready, graph_ready, startup_seconds?, error?}` | | `get_repository_map(max_items?)` | Compact repo overview for onboarding | `{languages: {py: 20}, hotspots: [...], start_here: [...]}` | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short name into qualified matches | `calc.py::Calculator.add → line 11` | | `search_graph(query?, kind?, file_pattern?)` | Graph search with degree filters and pagination | `{total: 5, results: [...]}` | @@ -60,7 +60,7 @@ All `file_path` arguments are **relative to the repo root** (e.g., `"src/main.py # Activate venv (required before all commands) source .venv/bin/activate -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file @@ -91,7 +91,8 @@ 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. | +| `server.py` | FastMCP 3.1.0 server — defines all 23 tools; each tool asks `IndexState` for the data it needs (single-file → on-demand parse, repo-wide → full index, graph → SQLite graph). Language-unaware. | +| `index_state.py` | `IndexState` — owns the index lifecycle: discovery → indexing (skeleton cache) → graph build, in a background thread when run as `codetree` (`create_server(root, background=True)`), synchronously by default (tests). Exposes `files_ready`/`index_ready`/`graph_ready` events, bounded waits (`WAIT_TIMEOUT`, env `CODETREE_WAIT_TIMEOUT`), progress and errors for `index_status`. Exposed to tests as `mcp._codetree_state`. | | `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. | @@ -144,6 +145,7 @@ Each plugin implements: | `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_async_startup.py` | Background indexing: non-blocking startup, on-demand single-file tools, "still building" messages, failure handling | | `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 | @@ -206,6 +208,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 -- 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. +- 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, which also stores each file's `has_errors` flag. - 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 a881988..767ecd9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -30,7 +30,7 @@ It exposes **23 tools** over MCP: | Tool | Purpose | Returns | |------|---------|---------| -| `index_status()` | Graph index freshness and stats | `{graph_exists, files, symbols, edges, last_indexed_at}` | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | `{graph_exists, files, symbols, edges, last_indexed_at, status, files_discovered, files_indexed, index_ready, graph_ready, startup_seconds?, error?}` | | `get_repository_map(max_items?)` | Compact repo overview for agent onboarding | `{languages, entry_points, hotspots, start_here, test_roots, stats}` | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short symbol names into qualified matches | `{matches: [{qualified_name, name, kind, file, line}]}` | | `search_graph(query?, kind?, file_pattern?, ...)` | Structured graph search with pagination and degree filtering | `{total, results: [{qualified_name, kind, in_degree, out_degree}]}` | @@ -66,7 +66,7 @@ All `file_path` arguments are **relative to the repo root** (e.g., `"src/main.py # Activate venv (required before all commands) source .venv/bin/activate -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file @@ -98,7 +98,8 @@ 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. | +| `server.py` | FastMCP 3.1.0 server — defines the 23 tools; each tool asks `IndexState` for the data it needs (single-file → on-demand parse, repo-wide → full index, graph → SQLite graph). Language-unaware. | +| `index_state.py` | `IndexState` — owns the index lifecycle: discovery → indexing (skeleton cache) → graph build, in a background thread when run as `codetree` (`create_server(root, background=True)`), synchronously by default (tests). Exposes `files_ready`/`index_ready`/`graph_ready` events, bounded waits (`WAIT_TIMEOUT`, env `CODETREE_WAIT_TIMEOUT`), progress and errors for `index_status`. Exposed to tests as `mcp._codetree_state`. | | `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. | @@ -150,6 +151,7 @@ Each plugin implements: | `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_async_startup.py` | Background indexing: non-blocking startup, on-demand single-file tools, "still building" messages, failure handling | | `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 | @@ -214,7 +216,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 -- 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. +- 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, which also stores each file's `has_errors` flag. - 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` @@ -256,13 +258,13 @@ codetree is a Python MCP server that gives coding agents structured code underst - Optional: `uv` for faster installation (recommended in README for Quick Start) - Lockfile: `.venv/` contains installed packages; no `requirements.txt` or `pyproject.lock` committed ## Frameworks -- FastMCP 3.1.0 (or later `>=2.0.0`) - MCP (Model Context Protocol) server framework +- FastMCP 3.x (`>=3.0.0`, which runs sync tools in a thread pool — required for background indexing) - MCP (Model Context Protocol) server framework - tree-sitter 0.23.0+ - AST parsing library (language-agnostic) - pytest (via GitHub Actions workflow, not explicitly in pyproject.toml dependencies but installed in CI) - hatchling (build backend) ## Key Dependencies - tree-sitter (0.23.0+) - Core AST parsing; blocks everything else -- fastmcp (2.0.0+) - MCP server registration and tool transport +- fastmcp (3.0.0+) - MCP server registration and tool transport - tree-sitter-python, tree-sitter-javascript, tree-sitter-typescript, tree-sitter-go, tree-sitter-rust, tree-sitter-java, tree-sitter-c, tree-sitter-cpp, tree-sitter-ruby ## Configuration - No explicit environment variables required for normal operation diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 9efca4e..e10c449 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -11,7 +11,7 @@ python -m venv .venv source .venv/bin/activate pip install -e . pip install pytest -pytest # 999 tests, ~30s +pytest # ~1200 tests, ~20s ``` ## What to work on diff --git a/README.md b/README.md index 4e0ad6c..d1cb71f 100644 --- a/README.md +++ b/README.md @@ -115,7 +115,7 @@ The agent sees every class, method, and docstring — with line numbers — with | Tool | Purpose | |------|---------| -| `index_status()` | Graph index freshness and stats | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | | `get_repository_map(max_items?)` | Compact repo overview: languages, entry points, hotspots | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short name into ranked qualified matches | | `search_graph(query?, kind?, file_pattern?)` | Graph search with degree filters and pagination | @@ -251,6 +251,11 @@ Add to `~/Library/Application Support/Claude/claude_desktop_config.json` (macOS) `node_modules`, `__pycache__`, `.git`, `dist`, `build`, … and nested worktrees. - `.codetree/` (codetree's own cache) is never indexed. +The MCP handshake is answered immediately and indexing runs in the background. +Single-file tools answer right away; repo-wide and graph tools wait up to 20 s +for the index (set `CODETREE_WAIT_TIMEOUT` to change it), then report progress. +`index_status()` never blocks. + ## Architecture ``` @@ -270,6 +275,7 @@ codetree server (FastMCP) | Module | Responsibility | |--------|---------------| | `server.py` | FastMCP server — defines all 23 tools | +| `index_state.py` | Index lifecycle: background indexing, readiness, progress for `index_status` | | `indexer.py` | File discovery, plugin dispatch, definition index | | `cache.py` | Skeleton cache with mtime invalidation | | `registry.py` | Maps file extensions to language plugins | @@ -298,7 +304,7 @@ source .venv/bin/activate pip install -e . pip install pytest -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file diff --git a/docs/LANDING_PAGE.md b/docs/LANDING_PAGE.md index 04122d1..400a8ba 100644 --- a/docs/LANDING_PAGE.md +++ b/docs/LANDING_PAGE.md @@ -159,7 +159,7 @@ The agent sees every class, method, and docstring — with line numbers — with | Tool | What it does | Example | |------|-------------|---------| -| `index_status` | Graph index freshness and stats | See how many files, symbols, and edges are indexed | +| `index_status` | Indexing progress, graph freshness and stats | See indexing progress and how many files, symbols, and edges are indexed | | `get_repository_map` | Compact repo overview for agent onboarding | Languages, entry points, hotspots, suggested starting points | | `resolve_symbol` | Disambiguate a short name into ranked qualified matches | "add" → `calc.py::Calculator.add`, `math.py::add` | | `search_graph` | Flexible graph search with degree filters and pagination | All functions with >5 inbound calls | @@ -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. +- **Non-blocking startup** — the MCP handshake is answered immediately; indexing runs in a background thread. Single-file tools answer right away (parsing on demand); repo-wide and graph tools wait briefly, then report progress. - **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. diff --git a/docs/TOOLS_GUIDE.md b/docs/TOOLS_GUIDE.md index d145d08..47bfdef 100644 --- a/docs/TOOLS_GUIDE.md +++ b/docs/TOOLS_GUIDE.md @@ -299,7 +299,7 @@ Found 3 tests ### 14. `index_status()` -Is the graph up to date? +Is the graph up to date — and is the server still indexing? ``` Agent: index_status() @@ -310,10 +310,28 @@ Returns: "files": 42, "symbols": 315, "edges": 580, - "last_indexed_at": "1741622400.0" + "last_indexed_at": "1741622400.0", + "status": "ready", + "files_discovered": 42, + "files_indexed": 42, + "index_ready": true, + "graph_ready": true, + "startup_seconds": 0.8 } ``` +The server indexes in the background, so this tool never blocks. While it is +still working, `status` is `starting`, `discovering`, `indexing` (with +`files_indexed` / `files_discovered` progress) or `building_graph`; it ends as +`ready`, or `error` with an `error` message if indexing or the graph build +failed. While a graph rebuild runs, `files` / `symbols` / `edges` / +`last_indexed_at` describe the last committed graph. Single-file tools +(`get_file_skeleton`, `get_symbol`, `get_imports`, `get_skeletons`, +`get_symbols`, `get_complexity`, `analyze_dataflow` flow/taint) answer +immediately; repo-wide and graph tools wait up to 20 s (`CODETREE_WAIT_TIMEOUT` +to change it), then return a "still building" message — retry, or poll +`index_status`. + --- ### 15. `get_repository_map(max_items?)` diff --git a/pyproject.toml b/pyproject.toml index 5b02f76..16219a0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ dependencies = [ "tree-sitter-cpp>=0.23.0", "tree-sitter-ruby>=0.23.0", "tree-sitter-kotlin>=0.23.0", - "fastmcp>=2.0.0", + "fastmcp>=3.0.0", ] [project.scripts] diff --git a/src/codetree/cache.py b/src/codetree/cache.py index b5586bb..f04e183 100644 --- a/src/codetree/cache.py +++ b/src/codetree/cache.py @@ -1,4 +1,7 @@ import json +import os +import stat +import tempfile from pathlib import Path @@ -19,7 +22,24 @@ def load(self): def save(self): """Write cache to disk, creating .codetree/ directory if needed.""" self._cache_file.parent.mkdir(parents=True, exist_ok=True) - self._cache_file.write_text(json.dumps(self._data, indent=2)) + # Write-then-rename so a server killed mid-save never leaves a torn file. + # The temp name is unique so concurrent servers on one repo cannot collide. + fd, tmp_name = tempfile.mkstemp( + dir=self._cache_file.parent, prefix="index.", suffix=".json.tmp" + ) + try: + with os.fdopen(fd, "w") as tmp_file: + tmp_file.write(json.dumps(self._data, indent=2)) + # mkstemp creates 0600; keep the previous mode, or a regular 0644. + try: + mode = stat.S_IMODE(self._cache_file.stat().st_mode) + except OSError: + mode = 0o644 + os.chmod(tmp_name, mode) + os.replace(tmp_name, self._cache_file) + except BaseException: + Path(tmp_name).unlink(missing_ok=True) + raise def get(self, rel_path: str) -> dict | None: return self._data.get(rel_path) diff --git a/src/codetree/graph/store.py b/src/codetree/graph/store.py index 073780e..f014b0a 100644 --- a/src/codetree/graph/store.py +++ b/src/codetree/graph/store.py @@ -93,6 +93,13 @@ def commit(self): self._conn.commit() self._in_transaction = False + def rollback(self): + """Discard the current transaction (e.g. after a failed build).""" + with self._lock: + if self._conn: + self._conn.rollback() + self._in_transaction = False + def _auto_commit(self): """Commit unless inside an explicit transaction. diff --git a/src/codetree/index_state.py b/src/codetree/index_state.py new file mode 100644 index 0000000..55eb1e9 --- /dev/null +++ b/src/codetree/index_state.py @@ -0,0 +1,263 @@ +"""Index lifecycle: build the indexer and graph, optionally in a background thread. + +The MCP handshake must not wait for indexing — on a large repository a cold +build takes longer than an MCP client's startup timeout. IndexState lets the +server start immediately while tools wait (bounded) for the data they need: + +- files_ready — file discovery finished (single-file tools can parse on demand) +- index_ready — every file indexed (repo-wide tools) +- graph_ready — SQLite graph built (graph tools) + +Every event is set even when a phase fails, so waiters never hang; they check +``indexer`` / ``error`` afterwards. +""" + +import os +import threading +import time +from pathlib import Path + +from .cache import Cache +from .graph.builder import GraphBuilder +from .graph.queries import GraphQueries +from .graph.store import GraphStore +from .indexer import Indexer + + + +def _wait_timeout_from_env(default: float = 20.0) -> float: + try: + return float(os.environ.get("CODETREE_WAIT_TIMEOUT", default)) + except ValueError: + return default + + +# How long a tool call waits for indexing before answering "still building". +# Override with CODETREE_WAIT_TIMEOUT (seconds) for clients with short tool timeouts. +WAIT_TIMEOUT = _wait_timeout_from_env() + + +class IndexState: + def __init__(self, root: str | Path): + self.root = Path(root) + self.indexer: Indexer | None = None + self.graph_store = GraphStore(str(self.root)) + self.graph_store.open() + self.graph_queries = GraphQueries(self.graph_store) + # Stats of the last committed graph. index_status reports these while a + # build is running: live queries on the builder's connection would see + # its uncommitted, half-rewritten tables. + self._graph_snapshot = self._read_graph_stats() + + self.phase = "starting" + self.files_total = 0 + self.files_done = 0 + self.error: str | None = None + self.started_at = time.time() + self.ready_at: float | None = None + self._discovered: set[str] = set() + + self.files_ready = threading.Event() + self.index_ready = threading.Event() + self.graph_ready = threading.Event() + self._thread: threading.Thread | None = None + + # ── Building ───────────────────────────────────────────────────────── + + def start_background(self) -> None: + """Run build() in a daemon thread and return immediately.""" + self._thread = threading.Thread(target=self.build, name="codetree-index", daemon=True) + self._thread.start() + + def build(self, raise_errors: bool = False) -> None: + """Discover, index (reusing the skeleton cache), then build the graph. + + Failures are recorded in ``error`` for tools to report. With + raise_errors (synchronous startup) the exception is re-raised too, so + the caller fails loudly with the original traceback. + """ + try: + self._build_index() + except Exception as exc: + self._fail(f"indexing failed: {exc!r}") + self.graph_ready.set() # no graph without an index — release waiters + if raise_errors: + raise + return + finally: + self.files_ready.set() + self.index_ready.set() + + try: + self.phase = "building_graph" + GraphBuilder(str(self.root), self.graph_store).build(indexer=self.indexer) + except Exception as exc: + self._fail(f"graph build failed: {exc!r}") + try: + self.graph_store.rollback() # keep the last committed graph intact + except Exception: + pass # store already closed, e.g. at interpreter exit + if raise_errors: + raise + return + finally: + # Waiters must always be released, whatever failed above. + try: + self._graph_snapshot = self._read_graph_stats() + except Exception: + pass + self.graph_ready.set() + + self.phase = "ready" + self.ready_at = time.time() + + def _build_index(self) -> None: + self.phase = "discovering" + cache = Cache(self.root) + cache.load() + + indexer = Indexer(self.root) + files = indexer.discover_files() + self._discovered = {str(f.relative_to(self.root)) for f in files} + self.files_total = len(files) + self.files_ready.set() + + self.phase = "indexing" + # Entries written before has_errors was cached are re-parsed once. + cached_mtimes = { + k: v["mtime"] for k, v in (cache._data or {}).items() if "has_errors" in v + } + indexer.build(cached_mtimes=cached_mtimes, files=files, progress=self._on_progress) + + # 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: + path = self.root / rel_path + try: + mtime = path.stat().st_mtime + if cache.is_valid(rel_path, mtime): + indexer.inject_cached( + rel_path=rel_path, + py_file=path, + source=path.read_bytes(), + skeleton=cache.get(rel_path).get("skeleton", []), + mtime=mtime, + has_errors=cache.get(rel_path).get("has_errors", False), + ) + continue + except OSError: + continue # deleted since discovery + # Changed between discovery and injection — parse it instead + entry = indexer.index_file(path) + if entry is not None: + indexer._index[rel_path] = entry + + # Rebuild definition index once after all injections (DATA-01, DATA-02, DATA-03 fix) + indexer._rebuild_definitions() + + # 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, + "skeleton": file_entry.skeleton, + "has_errors": file_entry.has_errors, + }) + try: + cache.save() + except OSError: + pass # the cache only speeds up the next start — never fail the index for it + + # Publish only the complete indexer: tools never see a half-built one. + self.indexer = indexer + + def _on_progress(self, done: int, total: int) -> None: + self.files_done = done + self.files_total = total + + def _fail(self, message: str) -> None: + self.phase = "error" + self.error = message + + # ── Access for tools ───────────────────────────────────────────────── + + def wait_for_index(self, timeout: float | None = None) -> Indexer | None: + """Return the full indexer once built; None if still indexing or failed.""" + self.index_ready.wait(WAIT_TIMEOUT if timeout is None else timeout) + return self.indexer + + def wait_for_graph(self, timeout: float | None = None) -> bool: + """True once the graph is built and usable.""" + self.graph_ready.wait(WAIT_TIMEOUT if timeout is None else timeout) + return self.graph_ready.is_set() and self.error is None + + def indexer_for_files(self, rel_paths: list[str], timeout: float | None = None) -> Indexer | None: + """Indexer able to answer about rel_paths, without waiting for the full index. + + Before the full index is ready, returns a throwaway Indexer holding just + those files, parsed on demand — limited to discovered files so results + match what the full index will say. None if discovery is not done yet. + """ + if self.index_ready.is_set(): + return self.indexer + if not self.files_ready.wait(WAIT_TIMEOUT if timeout is None else timeout): + return None + if self.index_ready.is_set(): + return self.indexer + partial = Indexer(self.root) + for rel_path in dict.fromkeys(rel_paths): + if rel_path not in self._discovered: + continue + entry = partial.index_file(self.root / rel_path) + if entry is not None: + partial._index[rel_path] = entry + partial._rebuild_definitions() + return partial + + def not_ready_message(self, need: str = "index") -> str: + """Explain why a tool cannot answer yet (or why it never will).""" + if self.error: + return f"codetree {self.error}. Restart the MCP server to retry." + elapsed = int(time.time() - self.started_at) + if self.phase == "indexing" and self.files_total: + where = f"indexing {self.files_done}/{self.files_total} files" + else: + where = self.phase.replace("_", " ") + what = "code graph" if need == "graph" else "repository index" + return ( + f"codetree is still building the {what} ({where}, {elapsed}s elapsed). " + "Retry in a few seconds; index_status shows progress." + ) + + def _read_graph_stats(self) -> dict: + return { + **self.graph_store.stats(), + "last_indexed_at": self.graph_store.get_meta("last_indexed_at"), + } + + def graph_stats(self) -> dict: + """files/symbols/edges/last_indexed_at of the last committed graph (never blocks).""" + if self.graph_ready.is_set(): + return self._read_graph_stats() + return dict(self._graph_snapshot) + + def status(self) -> dict: + """Lifecycle fields for index_status (never blocks).""" + result = { + "status": self.phase, + "files_discovered": self.files_total, + "files_indexed": self.files_done if self.phase == "indexing" else ( + len(self.indexer._index) if self.indexer else 0 + ), + "index_ready": self.index_ready.is_set() and self.indexer is not None, + "graph_ready": self.graph_ready.is_set() and self.error is None, + } + if self.ready_at is not None: + result["startup_seconds"] = round(self.ready_at - self.started_at, 2) + if self.error: + result["error"] = self.error + return result + + def close(self) -> None: + self.graph_store.close() diff --git a/src/codetree/indexer.py b/src/codetree/indexer.py index fe0a158..1068333 100644 --- a/src/codetree/indexer.py +++ b/src/codetree/indexer.py @@ -1,7 +1,9 @@ import os import subprocess +import threading from pathlib import Path from dataclasses import dataclass +from typing import Callable from .languages.base import LanguagePlugin from .registry import get_plugin @@ -45,6 +47,7 @@ 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 + self._call_graph_lock = threading.Lock() # 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] = [] @@ -210,17 +213,26 @@ def index_file(self, path: Path, mtime: float | None = None) -> FileEntry | None has_errors=has_errors, ) - def build(self, cached_mtimes: dict[str, float] | None = None): + def build(self, cached_mtimes: dict[str, float] | None = None, + files: list[Path] | None = None, + progress: Callable[[int, int], None] | None = None): """Index all supported files under root, skipping non-project dirs. Files whose path+mtime appear in cached_mtimes are skipped; the caller injects them via inject_cached(). + + Args: + cached_mtimes: rel_path → mtime of entries the caller has cached + files: pre-discovered file list (defaults to discover_files()) + progress: called as progress(done, total) after each file """ cached_mtimes = cached_mtimes or {} self.cached_candidates = [] - files = self.discover_files() + if files is None: + files = self.discover_files() self._discovery_order = [str(f.relative_to(self.root)) for f in files] - for candidate in files: + total = len(files) + for done, candidate in enumerate(files, start=1): rel = str(candidate.relative_to(self.root)) try: mtime = candidate.stat().st_mtime @@ -232,12 +244,14 @@ def build(self, cached_mtimes: dict[str, float] | None = None): entry = self.index_file(candidate, mtime) if entry is not None: self._index[rel] = entry + if progress is not None: + progress(done, total) # Build definition index from skeleton data (qualified keys, no duplicates, no ghosts) self._rebuild_definitions() def inject_cached(self, rel_path: str, py_file: Path, source: bytes, - skeleton: list[dict], mtime: float): + skeleton: list[dict], mtime: float, has_errors: bool = False): """Inject a pre-computed entry (from cache) without re-parsing.""" self._call_graph_built = False # invalidate so graph is rebuilt with new entry plugin = get_plugin(py_file) @@ -250,6 +264,7 @@ def inject_cached(self, rel_path: str, py_file: Path, source: bytes, mtime=mtime, language=py_file.suffix.lstrip("."), plugin=plugin, + has_errors=has_errors, ) # Note: _definitions is NOT updated here. After all inject_cached() calls # are complete, the caller must invoke _rebuild_definitions() to rebuild @@ -282,11 +297,24 @@ def get_call_graph(self, rel_path: str, function_name: str) -> dict: return {"calls": calls, "callers": callers} def _ensure_call_graph(self): - """Build repo-wide call graph lazily on first use.""" + """Build repo-wide call graph lazily on first use. + + Tools run concurrently (FastMCP thread pool): the lock makes the first + callers build it once, and the graph is published only when complete. + """ if self._call_graph_built: return - self._call_graph = {} - self._reverse_graph = {} + with self._call_graph_lock: + if self._call_graph_built: + return + call_graph, reverse_graph = self._compute_call_graph() + self._call_graph = call_graph + self._reverse_graph = reverse_graph + self._call_graph_built = True + + def _compute_call_graph(self) -> tuple[dict[str, set[str]], dict[str, set[str]]]: + call_graph: dict[str, set[str]] = {} + reverse_graph: dict[str, set[str]] = {} for rel_path, entry in self._index.items(): for item in entry.skeleton: if item["type"] in ("function", "method"): @@ -305,12 +333,12 @@ def _ensure_call_graph(self): else: # External/unresolved — keep as bare name callee_keys.add(f"?::{callee_name}") - self._call_graph[caller_key] = callee_keys + call_graph[caller_key] = callee_keys for ck in callee_keys: - if ck not in self._reverse_graph: - self._reverse_graph[ck] = set() - self._reverse_graph[ck].add(caller_key) - self._call_graph_built = True + if ck not in reverse_graph: + reverse_graph[ck] = set() + reverse_graph[ck].add(caller_key) + return call_graph, reverse_graph def find_dead_code(self, file_path: str | None = None) -> list[dict]: """Find symbols that are defined but never referenced elsewhere. diff --git a/src/codetree/server.py b/src/codetree/server.py index 5593990..851cf84 100644 --- a/src/codetree/server.py +++ b/src/codetree/server.py @@ -1,10 +1,18 @@ +import atexit from fastmcp import FastMCP from pathlib import Path -from .indexer import Indexer -from .cache import Cache +from .index_state import IndexState -def create_server(root: str) -> FastMCP: +def create_server(root: str, background: bool = False) -> FastMCP: + """Create the codetree MCP server for the repository at root. + + Args: + root: repository root + background: index in a background thread so the server can answer the + MCP handshake immediately (used by `codetree`). When False (default, + used by tests and embedding), indexing completes before returning. + """ mcp = FastMCP("codetree") root_path = Path(root) @@ -22,56 +30,17 @@ def _validate_path(file_path: str | None, _root: Path = root_path) -> str | None except ValueError: return f"Error: path '{file_path}' is outside the repo root — access denied" - # Load cache - cache = Cache(root) - cache.load() - - # Build index, skipping unchanged files - cached_mtimes = { - k: v["mtime"] for k, v in (cache._data or {}).items() - } - indexer = Indexer(root) - indexer.build(cached_mtimes=cached_mtimes) - - # 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 — 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, - "skeleton": file_entry.skeleton, - }) - cache.save() - - # ── Build persistent graph ─────────────────────────────────────────── - import atexit - from .graph.store import GraphStore - from .graph.builder import GraphBuilder - from .graph.queries import GraphQueries - - graph_store = GraphStore(root) - graph_store.open() - atexit.register(graph_store.close) - graph_builder = GraphBuilder(root, graph_store) - graph_builder.build(indexer=indexer) - graph_queries = GraphQueries(graph_store) + # Index + graph lifecycle. In background mode tools wait (bounded) for the + # data they need instead of blocking server startup. + state = IndexState(root) + # Register the store only: a reference to state would pin the whole index. + atexit.register(state.graph_store.close) + if background: + state.start_background() + else: + state.build(raise_errors=True) + mcp._codetree_state = state + graph_queries = state.graph_queries # ── Skeleton formatting helpers ────────────────────────────────────────── _TYPE_ABBREV = { @@ -136,6 +105,9 @@ def get_file_skeleton(file_path: str, format: str = "full") -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() skeleton = indexer.get_skeleton(file_path) if not skeleton: return f"File not found or empty: {file_path}" @@ -154,6 +126,9 @@ def get_symbol(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() result = indexer.get_symbol(file_path, symbol_name) if result is None: return f"Symbol '{symbol_name}' not found in {file_path}" @@ -168,6 +143,9 @@ def find_references(symbol_name: str) -> str: symbol_name: name of the symbol to search for; results include file paths relative to the repo root (e.g., "src/main.py") """ + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() refs = indexer.find_references(symbol_name) if not refs: return f"No references found for '{symbol_name}'" @@ -190,6 +168,9 @@ def get_call_graph(file_path: str, function_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() graph = indexer.get_call_graph(file_path, function_name) lines = [f"Call graph for '{function_name}':"] @@ -218,6 +199,9 @@ def get_imports(file_path: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() entry = indexer._index.get(file_path) if entry is None: return f"File not found: {file_path}" @@ -239,6 +223,9 @@ def get_skeletons(file_paths: list[str], format: str = "full") -> str: """ if not file_paths: return "No files requested." + indexer = state.indexer_for_files(file_paths) + if indexer is None: + return state.not_ready_message() parts = [] for fp in file_paths: parts.append(f"=== {fp} ===") @@ -266,6 +253,9 @@ def get_symbols(symbols: list[dict]) -> str: """ if not symbols: return "No symbols requested." + indexer = state.indexer_for_files([item.get("file_path", "") for item in symbols]) + if indexer is None: + return state.not_ready_message() parts = [] for item in symbols: fp = item.get("file_path", "") @@ -293,6 +283,9 @@ def get_complexity(file_path: str, function_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() entry = indexer._index.get(file_path) if entry is None: return f"File not found: {file_path}" @@ -316,6 +309,9 @@ def find_dead_code(file_path: str | None = None) -> str: if file_path is not None: if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path and file_path not in indexer._index: return f"File not found: {file_path}" dead = indexer.find_dead_code(file_path=file_path) @@ -348,6 +344,9 @@ def get_blast_radius(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path not in indexer._index: return f"File not found: {file_path}" result = indexer.get_blast_radius(file_path, symbol_name) @@ -392,6 +391,9 @@ def detect_clones(file_path: str | None = None, min_lines: int = 5) -> str: if file_path is not None: if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() clones = indexer.detect_clones(file_path=file_path, min_lines=min_lines) if not clones: scope = file_path if file_path else "the repo" @@ -427,6 +429,9 @@ def search_symbols(query: str | None = None, type: str | None = None, language: filter by file extension without dot (e.g., "py", "js", "go") format: "full" (default) or "compact" (abbreviated) """ + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() results = indexer.search_symbols( query=query, type=type, parent=parent, has_doc=has_doc, min_complexity=min_complexity, language=language, @@ -476,6 +481,9 @@ def find_tests(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path not in indexer._index: return f"File not found: {file_path}" tests = indexer.find_tests(file_path, symbol_name) @@ -491,13 +499,15 @@ def find_tests(file_path: str, symbol_name: str) -> str: @mcp.tool() def index_status() -> dict: - """Report on graph index freshness and stats.""" - stats = graph_store.stats() - last = graph_store.get_meta("last_indexed_at") + """Report on index build progress, graph freshness and stats. + + Never blocks: while the server is still indexing, "status" shows the + phase (discovering, indexing, building_graph) and file progress. + """ return { "graph_exists": True, - **stats, - "last_indexed_at": last, + **state.graph_stats(), + **state.status(), } @mcp.tool() @@ -510,6 +520,8 @@ def get_repository_map(max_items: int = 5) -> dict: Args: max_items: maximum items per section (default 5) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.repository_map(max_items=max_items) @mcp.tool() @@ -526,6 +538,8 @@ def resolve_symbol(query: str, kind: str | None = None, path_hint: prefer results from files matching this path limit: max results (default 10) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} results = graph_queries.resolve_symbol(query, kind=kind, path_hint=path_hint, limit=limit) return { "query": query, @@ -560,6 +574,8 @@ def search_graph(query: str | None = None, kind: str | None = None, limit: max results per page (default 10) offset: pagination offset (default 0) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.search_graph( query=query, kind=kind, file_pattern=file_pattern, relationship=relationship, direction=direction, @@ -579,6 +595,8 @@ def get_change_impact(symbol_query: str | None = None, diff_scope: "working" (uncommitted), "staged", or "HEAD~1" for git-based analysis depth: max hop depth (default 3) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.change_impact( symbol_query=symbol_query, diff_scope=diff_scope, @@ -606,10 +624,16 @@ def analyze_dataflow(file_path: str, function_name: str, return {"error": err} if mode == "cross_taint": + indexer = state.wait_for_index() + if indexer is None: + return {"error": state.not_ready_message()} if file_path not in indexer._index: return {"error": f"File not found: {file_path}"} return extract_cross_function_taint(indexer, file_path, function_name, depth=depth) + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return {"error": state.not_ready_message()} entry = indexer._index.get(file_path) if entry is None: return {"error": f"File not found: {file_path}"} @@ -633,6 +657,9 @@ def find_hot_paths(top_n: int = 10) -> str: Args: top_n: max results to return (default 10) """ + if not state.wait_for_graph(): + return state.not_ready_message("graph") + indexer = state.indexer results = graph_queries.find_hot_paths(indexer, top_n=top_n) if not results: return "No hot paths found (no functions with both callers and complexity)." @@ -660,6 +687,8 @@ def get_dependency_graph(file_path: str | None = None, if file_path is not None: if err := _validate_path(file_path): return err + if not state.wait_for_graph(): + return state.not_ready_message("graph") result = graph_queries.get_dependency_graph(file_path=file_path, format=format) summary = f"\n\n{result['nodes']} files, {result['edges']} import edges" return result["content"] + summary @@ -740,6 +769,9 @@ def suggest_docs(file_path: str | None = None, if file_path is not None: if err := _validate_path(file_path): return err + if not state.wait_for_graph(): + return state.not_ready_message("graph") + indexer = state.indexer results = graph_queries.suggest_docs(indexer, file_path=file_path, symbol_name=symbol_name) if not results: return "No undocumented functions found." @@ -761,5 +793,5 @@ def suggest_docs(file_path: str | None = None, def run(root: str): - mcp = create_server(root) + mcp = create_server(root, background=True) mcp.run() diff --git a/tests/test_async_startup.py b/tests/test_async_startup.py new file mode 100644 index 0000000..7990f2c --- /dev/null +++ b/tests/test_async_startup.py @@ -0,0 +1,350 @@ +"""Tests for background indexing: the server answers before the index is built.""" + +import subprocess +import threading +import time + +import pytest + +from codetree import index_state +from codetree.indexer import Indexer +from codetree.server import create_server + + +def _tool(mcp, name): + return mcp.local_provider._components[f"tool:{name}@"].fn + + +@pytest.fixture +def repo(tmp_path): + root = tmp_path / "repo" + root.mkdir() + subprocess.run(["git", "init", "-q"], cwd=root, check=True) + (root / ".gitignore").write_text("ignored/\n.codetree/\n") + (root / "calc.py").write_text( + 'class Calculator:\n def add(self, a, b):\n return a + b\n\n' + 'def helper():\n return Calculator().add(1, 2)\n' + ) + (root / "main.py").write_text("from calc import helper\n\ndef main():\n helper()\n") + (root / "ignored").mkdir() + (root / "ignored" / "secret.py").write_text("def hidden():\n pass\n") + return root + + +@pytest.fixture(autouse=True) +def background_states(monkeypatch): + """Join every background build and close its store when the test ends.""" + started = [] + original = index_state.IndexState.start_background + + def tracking_start(self): + started.append(self) + original(self) + + monkeypatch.setattr(index_state.IndexState, "start_background", tracking_start) + yield started + for state in started: + if state._thread is not None: + state._thread.join(15) + state.close() + + +@pytest.fixture +def gate(monkeypatch, background_states): + """Block background indexing (only the codetree-index thread) until released. + + Depends on background_states so the gate opens before threads are joined. + """ + release = threading.Event() + original = Indexer.index_file + + def gated_index_file(self, path, mtime=None): + if threading.current_thread().name == "codetree-index": + release.wait(10) + return original(self, path, mtime) + + monkeypatch.setattr(Indexer, "index_file", gated_index_file) + monkeypatch.setattr(index_state, "WAIT_TIMEOUT", 0.3) + yield release + release.set() + + +def _wait(event, timeout=10): + assert event.wait(timeout), "background indexing did not finish" + + +# ── Startup does not wait ──────────────────────────────────────────────────── + +def test_background_create_server_returns_before_indexing(repo, gate): + start = time.monotonic() + mcp = create_server(str(repo), background=True) + assert time.monotonic() - start < 5 + state = mcp._codetree_state + assert not state.index_ready.is_set() + gate.set() + _wait(state.graph_ready) + + +def test_index_status_never_blocks_and_reports_progress(repo, gate): + mcp = create_server(str(repo), background=True) + state = mcp._codetree_state + _wait(state.files_ready) + + start = time.monotonic() + status = _tool(mcp, "index_status")() + assert time.monotonic() - start < 1 + assert status["status"] in ("indexing", "discovering") + assert status["files_discovered"] == 2 + assert status["index_ready"] is False + assert status["graph_ready"] is False + + gate.set() + _wait(state.graph_ready) + status = _tool(mcp, "index_status")() + assert status["status"] == "ready" + assert status["index_ready"] is True and status["graph_ready"] is True + assert status["files_indexed"] == 2 + assert "startup_seconds" in status + + +# ── Single-file tools answer during indexing ───────────────────────────────── + +def test_single_file_tools_parse_on_demand_while_indexing(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + skeleton = _tool(mcp, "get_file_skeleton")(file_path="calc.py") + assert "class Calculator" in skeleton + assert "def add" in skeleton + assert "def helper" in _tool(mcp, "get_symbol")(file_path="calc.py", symbol_name="helper") + assert "from calc import helper" in _tool(mcp, "get_imports")(file_path="main.py") + assert "Complexity of helper()" in _tool(mcp, "get_complexity")( + file_path="calc.py", function_name="helper") + multi = _tool(mcp, "get_skeletons")(file_paths=["calc.py", "main.py"]) + assert "class Calculator" in multi and "def main" in multi + flow = _tool(mcp, "analyze_dataflow")(file_path="calc.py", function_name="helper") + assert "error" not in flow + assert not mcp._codetree_state.index_ready.is_set() + + +def test_on_demand_parsing_respects_gitignore(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "get_file_skeleton")(file_path="ignored/secret.py") + assert result == "File not found or empty: ignored/secret.py" + + +# ── Repo-wide and graph tools wait, then explain ───────────────────────────── + +def test_repo_wide_tool_reports_progress_when_index_not_ready(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "find_references")(symbol_name="helper") + assert "still building the repository index" in result + assert "/2 files" in result + assert "index_status" in result + + +def test_graph_tools_report_not_ready(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "get_repository_map")() + assert "still building the code graph" in result["error"] + assert "still building the code graph" in _tool(mcp, "find_hot_paths")() + + +def test_repo_wide_tool_waits_for_index_finishing_in_time(repo, gate, monkeypatch): + monkeypatch.setattr(index_state, "WAIT_TIMEOUT", 10) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + threading.Timer(0.2, gate.set).start() + + result = _tool(mcp, "find_references")(symbol_name="helper") + assert "calc.py" in result and "main.py" in result + + +def test_tools_return_full_results_once_ready(repo, gate): + mcp = create_server(str(repo), background=True) + gate.set() + _wait(mcp._codetree_state.graph_ready) + + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + assert "ignored/secret.py" not in _tool(mcp, "find_references")(symbol_name="hidden") + matches = _tool(mcp, "resolve_symbol")(query="helper")["matches"] + assert [m["file"] for m in matches] == ["calc.py"] + + +# ── Failures never hang tools ──────────────────────────────────────────────── + +def test_indexing_failure_is_reported_not_hung(repo, monkeypatch): + def boom(self): + raise RuntimeError("disk on fire") + + monkeypatch.setattr(Indexer, "discover_files", boom) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + + start = time.monotonic() + result = _tool(mcp, "find_references")(symbol_name="helper") + assert time.monotonic() - start < 3 + assert "indexing failed" in result and "disk on fire" in result + assert "indexing failed" in _tool(mcp, "get_file_skeleton")(file_path="calc.py") + status = _tool(mcp, "index_status")() + assert status["status"] == "error" + assert "disk on fire" in status["error"] + + +def test_graph_failure_keeps_structural_tools_working(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + + def boom(self, indexer=None): + raise RuntimeError("graph exploded") + + monkeypatch.setattr(GraphBuilder, "build", boom) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + assert "graph build failed" in _tool(mcp, "get_repository_map")()["error"] + assert _tool(mcp, "index_status")()["graph_ready"] is False + + +# ── Synchronous mode (default) ─────────────────────────────────────────────── + +def test_default_create_server_is_ready_on_return(repo): + mcp = create_server(str(repo)) + status = _tool(mcp, "index_status")() + assert status["status"] == "ready" + assert status["index_ready"] and status["graph_ready"] + + +def test_warm_start_reuses_cache_in_background(repo): + first = create_server(str(repo)) + first._codetree_state.close() + + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + assert "class Calculator" in _tool(mcp, "get_file_skeleton")(file_path="calc.py") + assert _tool(mcp, "index_status")()["files_indexed"] == 2 + + +def test_sync_create_server_raises_on_indexing_failure(repo, monkeypatch): + def boom(self): + raise RuntimeError("disk on fire") + + monkeypatch.setattr(Indexer, "discover_files", boom) + with pytest.raises(RuntimeError, match="disk on fire"): + create_server(str(repo)) + + +def test_cache_save_failure_does_not_fail_indexing(repo, monkeypatch): + from codetree.cache import Cache + + def readonly(self): + raise PermissionError("read-only file system") + + monkeypatch.setattr(Cache, "save", readonly) + mcp = create_server(str(repo)) + + assert _tool(mcp, "index_status")()["status"] == "ready" + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + + +def test_atexit_does_not_pin_the_index(repo, monkeypatch): + registered = [] + monkeypatch.setattr("codetree.server.atexit.register", registered.append) + mcp = create_server(str(repo)) + + # Only the SQLite store is registered — not IndexState, which owns every + # file's source through its indexer. + assert registered == [mcp._codetree_state.graph_store.close] + + +def test_index_status_reports_committed_graph_during_rebuild(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + + first = create_server(str(repo)) + committed = _tool(first, "index_status")() + first._codetree_state.close() + assert committed["symbols"] > 0 + + # Hold the next build mid-transaction, after it deleted rows. + in_build = threading.Event() + release = threading.Event() + original = GraphBuilder.build + + def held_build(self, indexer=None): + self._store.begin() + self._store.delete_symbols_for_file("calc.py") + self._store.delete_symbols_for_file("main.py") + in_build.set() + release.wait(10) + return original(self, indexer=indexer) + + monkeypatch.setattr(GraphBuilder, "build", held_build) + mcp = create_server(str(repo), background=True) + try: + assert in_build.wait(10) + during = _tool(mcp, "index_status")() + assert during["status"] == "building_graph" + assert during["symbols"] == committed["symbols"] + assert during["last_indexed_at"] == committed["last_indexed_at"] + finally: + release.set() + _wait(mcp._codetree_state.graph_ready) + after = _tool(mcp, "index_status")() + assert after["status"] == "ready" + assert after["last_indexed_at"] != committed["last_indexed_at"] + + +def test_syntax_error_flag_survives_warm_start(repo): + (repo / "broken.py").write_text("def broken(:\n pass\n") + cold = _tool(create_server(str(repo)), "get_file_skeleton")(file_path="broken.py") + warm = _tool(create_server(str(repo)), "get_file_skeleton")(file_path="broken.py") + assert "WARNING: File has syntax errors" in cold + assert warm == cold + + +def test_cache_entries_without_error_flag_are_reparsed(repo): + import json + + (repo / "broken.py").write_text("def broken(:\n pass\n") + create_server(str(repo)) + cache_file = repo / ".codetree" / "index.json" + data = json.loads(cache_file.read_text()) + for entry in data.values(): + del entry["has_errors"] # as written by an older codetree + cache_file.write_text(json.dumps(data)) + + mcp = create_server(str(repo)) + assert "WARNING: File has syntax errors" in _tool(mcp, "get_file_skeleton")(file_path="broken.py") + + +def test_wait_timeout_from_env(monkeypatch): + monkeypatch.setenv("CODETREE_WAIT_TIMEOUT", "3.5") + assert index_state._wait_timeout_from_env() == 3.5 + monkeypatch.setenv("CODETREE_WAIT_TIMEOUT", "soon") + assert index_state._wait_timeout_from_env() == 20.0 + + +def test_graph_failure_with_failing_rollback_still_releases_waiters(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + from codetree.graph.store import GraphStore + + def boom(self, indexer=None): + raise RuntimeError("graph exploded") + + def broken_rollback(self): + raise RuntimeError("connection gone") + + monkeypatch.setattr(GraphBuilder, "build", boom) + monkeypatch.setattr(GraphStore, "rollback", broken_rollback) + mcp = create_server(str(repo), background=True) + + _wait(mcp._codetree_state.graph_ready) + status = _tool(mcp, "index_status")() + assert status["status"] == "error" + assert "graph exploded" in status["error"] diff --git a/tests/test_cache.py b/tests/test_cache.py index 279bff1..51fede3 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -37,3 +37,19 @@ def test_cache_creates_directory_if_missing(tmp_path): cache = Cache(tmp_path) cache.save() assert cache_dir.exists() + + +def test_cache_file_is_world_readable_like_before(tmp_path): + import os + import stat + from codetree.cache import Cache + + cache = Cache(tmp_path) + cache.set("a.py", {"mtime": 1.0, "skeleton": []}) + cache.save() + cache_file = tmp_path / ".codetree" / "index.json" + assert stat.S_IMODE(cache_file.stat().st_mode) == 0o644 + + os.chmod(cache_file, 0o600) # a user's explicit choice is kept + cache.save() + assert stat.S_IMODE(cache_file.stat().st_mode) == 0o600 diff --git a/tests/test_file_discovery.py b/tests/test_file_discovery.py index cfcb9e8..6d8de8b 100644 --- a/tests/test_file_discovery.py +++ b/tests/test_file_discovery.py @@ -297,3 +297,29 @@ def test_index_order_does_not_depend_on_cache(tmp_path): 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"] + +def test_concurrent_cache_saves_never_fail_or_tear(tmp_path): + import threading + from codetree.cache import Cache + + errors = [] + + def writer(tag): + cache = Cache(tmp_path) + cache.set("x.py", {"mtime": 1.0, "skeleton": [{"name": tag * 200}]}) + try: + for _ in range(50): + cache.save() + except Exception as exc: + errors.append(exc) + + threads = [threading.Thread(target=writer, args=(tag,)) for tag in "abcd"] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert errors == [] + data = json.loads((tmp_path / ".codetree" / "index.json").read_text()) + assert data["x.py"]["mtime"] == 1.0 + assert list((tmp_path / ".codetree").glob("*.tmp")) == [] diff --git a/tests/test_graph_store.py b/tests/test_graph_store.py index a9df6de..62e98d6 100644 --- a/tests/test_graph_store.py +++ b/tests/test_graph_store.py @@ -135,6 +135,17 @@ def test_delete_edges_for_file_is_exact_not_like_pattern(self, store): assert len(store.edges_from("my_file.py.bak::foo")) == 1 +class TestTransactions: + def test_rollback_discards_uncommitted_writes(self, store): + store.upsert_edge(Edge("kept.py::a", "x.py::b", "CALLS")) + store.begin() + store.upsert_edge(Edge("dropped.py::a", "x.py::b", "CALLS")) + store.rollback() + assert len(store.edges_from("kept.py::a")) == 1 + assert store.edges_from("dropped.py::a") == [] + assert store._in_transaction is False + + class TestFileCRUD: def test_upsert_and_get_file(self, store): store.upsert_file("calc.py", sha256="abc123", language="py", is_test=False) diff --git a/tests/test_indexer.py b/tests/test_indexer.py index d8feadd..cb889f1 100644 --- a/tests/test_indexer.py +++ b/tests/test_indexer.py @@ -338,3 +338,35 @@ def test_unknown_file_has_empty_calls(self, sample_repo): idx.build() graph = idx.get_call_graph("missing.py", "fn") assert graph["calls"] == [] + + +def test_concurrent_call_graph_requests_build_it_once(tmp_path, monkeypatch): + import threading + import time + from codetree.indexer import Indexer + + (tmp_path / "a.py").write_text("def a():\n b()\n\ndef b():\n pass\n") + indexer = Indexer(tmp_path) + indexer.build() + builds = [] + original = Indexer._compute_call_graph + + def slow_compute(self): + builds.append(1) + time.sleep(0.2) + return original(self) + + monkeypatch.setattr(Indexer, "_compute_call_graph", slow_compute) + results = [] + threads = [ + threading.Thread(target=lambda: results.append(indexer.get_blast_radius("a.py", "b"))) + for _ in range(4) + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert builds == [1] + assert all(r == results[0] for r in results) + assert any(c["name"] == "a" for c in results[0]["callers"])