diff --git a/AGENTS.md b/AGENTS.md index 2995f63..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,8 @@ 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 | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -173,7 +175,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"` @@ -203,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 427f4b0..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,8 @@ 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 | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -181,7 +183,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"` @@ -211,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/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 06f5bc0..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: @@ -294,10 +369,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 >= ? AND source_qn < ?", (low, high) + ) self._conn.execute( - "DELETE FROM edges WHERE source_qn LIKE ? OR target_qn LIKE ?", - (prefix + "%", prefix + "%"), + "DELETE FROM edges WHERE target_qn >= ? AND target_qn < ?", (low, high) ) if not self._in_transaction: self._conn.commit() 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/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/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"] + 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 384e200..a9df6de 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,38 @@ def test_stats(self, store): assert stats["files"] == 1 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"]