diff --git a/AGENTS.md b/AGENTS.md index 2995f63..d9a6b69 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -23,7 +23,7 @@ It exposes **23 tools** over MCP: | `detect_clones(file_path?, min_lines?)` | Duplicate/near-duplicate functions | `Clone group 1 (2 functions, 12 lines each):` | | `search_symbols(query?, type?, parent?, ..., format?)` | Flexible symbol search; `format="compact"` omits doc lines | `calc.py: class Calculator → line 1` | | `find_tests(file_path, symbol_name)` | Find test functions for a symbol | `test_calc.py: test_add() → line 3 (name match)` | -| `index_status()` | Graph index freshness and stats | `{files: 42, symbols: 315, edges: 580}` | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | `{graph_exists, files, symbols, edges, last_indexed_at, status, files_discovered, files_indexed, index_ready, graph_ready, startup_seconds?, error?}` | | `get_repository_map(max_items?)` | Compact repo overview for onboarding | `{languages: {py: 20}, hotspots: [...], start_here: [...]}` | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short name into qualified matches | `calc.py::Calculator.add → line 11` | | `search_graph(query?, kind?, file_pattern?)` | Graph search with degree filters and pagination | `{total: 5, results: [...]}` | @@ -60,7 +60,7 @@ All `file_path` arguments are **relative to the repo root** (e.g., `"src/main.py # Activate venv (required before all commands) source .venv/bin/activate -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file @@ -91,8 +91,9 @@ 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. | +| `server.py` | FastMCP 3.1.0 server — defines all 23 tools; each tool asks `IndexState` for the data it needs (single-file → on-demand parse, repo-wide → full index, graph → SQLite graph). Language-unaware. | +| `index_state.py` | `IndexState` — owns the index lifecycle: discovery → indexing (skeleton cache) → graph build, in a background thread when run as `codetree` (`create_server(root, background=True)`), synchronously by default (tests). Exposes `files_ready`/`index_ready`/`graph_ready` events, bounded waits (`WAIT_TIMEOUT`, env `CODETREE_WAIT_TIMEOUT`), progress and errors for `index_status`. Exposed to tests as `mcp._codetree_state`. | +| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Discovers files with `git ls-files` (tracked files always; untracked non-ignored files minus `SKIP_DIRS`; nested worktrees/repos/submodules excluded), falling back to an `os.walk` that prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, etc.) and nested worktrees. | | `cache.py` | `.codetree/index.json` — stores pre-computed skeletons with mtime-based invalidation. Language-unaware. | | `registry.py` | Maps file extensions → plugin instances. The **only** place languages are registered. | @@ -142,6 +143,9 @@ 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_async_startup.py` | Background indexing: non-blocking startup, on-demand single-file tools, "still building" messages, failure handling | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -173,7 +177,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 +208,6 @@ The tree-sitter Python bindings have breaking changes from older docs: - Plugin classes: `{Lang}Plugin` (e.g., `PythonPlugin`, `GoPlugin`) - Module-level parser/language globals: `_PARSER`, `_LANGUAGE` - Skeleton results are deduplicated by `(name, line)` and sorted by line number -- 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, which also stores each file's `has_errors` flag. - FastMCP tool access in tests: `mcp.local_provider._components[f"tool:{name}@"].fn` - **Doc sync rule**: When tools are added, removed, or changed, update all 5 doc files: `README.md`, `TOOLS_GUIDE.md`, `LANDING_PAGE.md`, `CLAUDE.md`, `AGENTS.md` diff --git a/CLAUDE.md b/CLAUDE.md index 427f4b0..767ecd9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -30,7 +30,7 @@ It exposes **23 tools** over MCP: | Tool | Purpose | Returns | |------|---------|---------| -| `index_status()` | Graph index freshness and stats | `{graph_exists, files, symbols, edges, last_indexed_at}` | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | `{graph_exists, files, symbols, edges, last_indexed_at, status, files_discovered, files_indexed, index_ready, graph_ready, startup_seconds?, error?}` | | `get_repository_map(max_items?)` | Compact repo overview for agent onboarding | `{languages, entry_points, hotspots, start_here, test_roots, stats}` | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short symbol names into qualified matches | `{matches: [{qualified_name, name, kind, file, line}]}` | | `search_graph(query?, kind?, file_pattern?, ...)` | Structured graph search with pagination and degree filtering | `{total, results: [{qualified_name, kind, in_degree, out_degree}]}` | @@ -66,7 +66,7 @@ All `file_path` arguments are **relative to the repo root** (e.g., `"src/main.py # Activate venv (required before all commands) source .venv/bin/activate -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file @@ -98,8 +98,9 @@ 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. | +| `server.py` | FastMCP 3.1.0 server — defines the 23 tools; each tool asks `IndexState` for the data it needs (single-file → on-demand parse, repo-wide → full index, graph → SQLite graph). Language-unaware. | +| `index_state.py` | `IndexState` — owns the index lifecycle: discovery → indexing (skeleton cache) → graph build, in a background thread when run as `codetree` (`create_server(root, background=True)`), synchronously by default (tests). Exposes `files_ready`/`index_ready`/`graph_ready` events, bounded waits (`WAIT_TIMEOUT`, env `CODETREE_WAIT_TIMEOUT`), progress and errors for `index_status`. Exposed to tests as `mcp._codetree_state`. | +| `indexer.py` | Discovers files, stores a `FileEntry` per file (with its plugin + `has_errors` flag), routes all queries through the stored plugin. Builds a definition index and lazy call graph for dead code, blast radius, and clone detection. Discovers files with `git ls-files` (tracked files always; untracked non-ignored files minus `SKIP_DIRS`; nested worktrees/repos/submodules excluded), falling back to an `os.walk` that prunes `SKIP_DIRS` (`.venv`, `node_modules`, `__pycache__`, `.git`, etc.) and nested worktrees. | | `cache.py` | `.codetree/index.json` — stores pre-computed skeletons with mtime-based invalidation. Language-unaware. | | `registry.py` | Maps file extensions → plugin instances. The **only** place languages are registered. | @@ -148,6 +149,9 @@ 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_async_startup.py` | Background indexing: non-blocking startup, on-demand single-file tools, "still building" messages, failure handling | | `test_cache.py` | Cache load/save/invalidation | | `tests/languages/test_.py` | Per-language core tests | | `tests/languages/test__comprehensive.py` | Exhaustive code pattern coverage per language | @@ -181,7 +185,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 +216,7 @@ The tree-sitter Python bindings have breaking changes from older docs: - Plugin classes: `{Lang}Plugin` (e.g., `PythonPlugin`, `GoPlugin`) - Module-level parser/language globals: `_PARSER`, `_LANGUAGE` - Skeleton results are deduplicated by `(name, line)` and sorted by line number -- 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, which also stores each file's `has_errors` flag. - FastMCP tool access in tests: `mcp.local_provider._components[f"tool:{name}@"].fn` - **Doc sync rule**: When tools are added, removed, or changed, update all 5 doc files: `README.md`, `TOOLS_GUIDE.md`, `LANDING_PAGE.md`, `CLAUDE.md`, `AGENTS.md` @@ -253,13 +258,13 @@ codetree is a Python MCP server that gives coding agents structured code underst - Optional: `uv` for faster installation (recommended in README for Quick Start) - Lockfile: `.venv/` contains installed packages; no `requirements.txt` or `pyproject.lock` committed ## Frameworks -- FastMCP 3.1.0 (or later `>=2.0.0`) - MCP (Model Context Protocol) server framework +- FastMCP 3.x (`>=3.0.0`, which runs sync tools in a thread pool — required for background indexing) - MCP (Model Context Protocol) server framework - tree-sitter 0.23.0+ - AST parsing library (language-agnostic) - pytest (via GitHub Actions workflow, not explicitly in pyproject.toml dependencies but installed in CI) - hatchling (build backend) ## Key Dependencies - tree-sitter (0.23.0+) - Core AST parsing; blocks everything else -- fastmcp (2.0.0+) - MCP server registration and tool transport +- fastmcp (3.0.0+) - MCP server registration and tool transport - tree-sitter-python, tree-sitter-javascript, tree-sitter-typescript, tree-sitter-go, tree-sitter-rust, tree-sitter-java, tree-sitter-c, tree-sitter-cpp, tree-sitter-ruby ## Configuration - No explicit environment variables required for normal operation diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 9efca4e..e10c449 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -11,7 +11,7 @@ python -m venv .venv source .venv/bin/activate pip install -e . pip install pytest -pytest # 999 tests, ~30s +pytest # ~1200 tests, ~20s ``` ## What to work on diff --git a/README.md b/README.md index 4e5821a..d1cb71f 100644 --- a/README.md +++ b/README.md @@ -115,7 +115,7 @@ The agent sees every class, method, and docstring — with line numbers — with | Tool | Purpose | |------|---------| -| `index_status()` | Graph index freshness and stats | +| `index_status()` | Indexing progress, graph freshness and stats (never blocks) | | `get_repository_map(max_items?)` | Compact repo overview: languages, entry points, hotspots | | `resolve_symbol(query, kind?, path_hint?)` | Disambiguate short name into ranked qualified matches | | `search_graph(query?, kind?, file_pattern?)` | Graph search with degree filters and pagination | @@ -238,6 +238,24 @@ 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. + +The MCP handshake is answered immediately and indexing runs in the background. +Single-file tools answer right away; repo-wide and graph tools wait up to 20 s +for the index (set `CODETREE_WAIT_TIMEOUT` to change it), then report progress. +`index_status()` never blocks. + ## Architecture ``` @@ -257,6 +275,7 @@ codetree server (FastMCP) | Module | Responsibility | |--------|---------------| | `server.py` | FastMCP server — defines all 23 tools | +| `index_state.py` | Index lifecycle: background indexing, readiness, progress for `index_status` | | `indexer.py` | File discovery, plugin dispatch, definition index | | `cache.py` | Skeleton cache with mtime invalidation | | `registry.py` | Maps file extensions to language plugins | @@ -285,7 +304,7 @@ source .venv/bin/activate pip install -e . pip install pytest -# Run all tests (~1058 tests, ~35s) +# Run all tests (~1200 tests, ~20s) pytest # Run a single test file diff --git a/docs/LANDING_PAGE.md b/docs/LANDING_PAGE.md index bbf2d04..400a8ba 100644 --- a/docs/LANDING_PAGE.md +++ b/docs/LANDING_PAGE.md @@ -159,7 +159,7 @@ The agent sees every class, method, and docstring — with line numbers — with | Tool | What it does | Example | |------|-------------|---------| -| `index_status` | Graph index freshness and stats | See how many files, symbols, and edges are indexed | +| `index_status` | Indexing progress, graph freshness and stats | See indexing progress and how many files, symbols, and edges are indexed | | `get_repository_map` | Compact repo overview for agent onboarding | Languages, entry points, hotspots, suggested starting points | | `resolve_symbol` | Disambiguate a short name into ranked qualified matches | "add" → `calc.py::Calculator.add`, `math.py::add` | | `search_graph` | Flexible graph search with degree filters and pagination | All functions with >5 inbound calls | @@ -365,6 +365,8 @@ claude mcp add codetree -- uvx --from mcp-server-codetree codetree --root . - **FastMCP** for the MCP protocol — stdio transport, zero network config. - **Plugin architecture** — each language is a self-contained class implementing 5 core methods. Adding a language is copying a template. - **Smart caching** — `.codetree/index.json` with mtime-based invalidation. Unchanged files skip parsing entirely. +- **Non-blocking startup** — the MCP handshake is answered immediately; indexing runs in a background thread. Single-file tools answer right away (parsing on demand); repo-wide and graph tools wait briefly, then report progress. +- **Respects `.gitignore`** — files are discovered via `git ls-files`, so ignored build output and nested git worktrees (e.g. `.claude/worktrees/`) are never indexed. - **Lazy call graph** — only built when tools like `find_dead_code` or `get_blast_radius` are first called. Stored in memory, O(1) lookup. - **PageRank** — standard algorithm (25 iterations, damping 0.85) for ranking symbol importance by reference count. - **Clone detection** — AST normalization (identifiers → `_ID_`, strings → `_STR_`, numbers → `_NUM_`) + SHA-256 hashing. Catches exact copies and renamed-variable copies. diff --git a/docs/TOOLS_GUIDE.md b/docs/TOOLS_GUIDE.md index d145d08..47bfdef 100644 --- a/docs/TOOLS_GUIDE.md +++ b/docs/TOOLS_GUIDE.md @@ -299,7 +299,7 @@ Found 3 tests ### 14. `index_status()` -Is the graph up to date? +Is the graph up to date — and is the server still indexing? ``` Agent: index_status() @@ -310,10 +310,28 @@ Returns: "files": 42, "symbols": 315, "edges": 580, - "last_indexed_at": "1741622400.0" + "last_indexed_at": "1741622400.0", + "status": "ready", + "files_discovered": 42, + "files_indexed": 42, + "index_ready": true, + "graph_ready": true, + "startup_seconds": 0.8 } ``` +The server indexes in the background, so this tool never blocks. While it is +still working, `status` is `starting`, `discovering`, `indexing` (with +`files_indexed` / `files_discovered` progress) or `building_graph`; it ends as +`ready`, or `error` with an `error` message if indexing or the graph build +failed. While a graph rebuild runs, `files` / `symbols` / `edges` / +`last_indexed_at` describe the last committed graph. Single-file tools +(`get_file_skeleton`, `get_symbol`, `get_imports`, `get_skeletons`, +`get_symbols`, `get_complexity`, `analyze_dataflow` flow/taint) answer +immediately; repo-wide and graph tools wait up to 20 s (`CODETREE_WAIT_TIMEOUT` +to change it), then return a "still building" message — retry, or poll +`index_status`. + --- ### 15. `get_repository_map(max_items?)` diff --git a/pyproject.toml b/pyproject.toml index 5b02f76..16219a0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ dependencies = [ "tree-sitter-cpp>=0.23.0", "tree-sitter-ruby>=0.23.0", "tree-sitter-kotlin>=0.23.0", - "fastmcp>=2.0.0", + "fastmcp>=3.0.0", ] [project.scripts] diff --git a/src/codetree/cache.py b/src/codetree/cache.py index b5586bb..f04e183 100644 --- a/src/codetree/cache.py +++ b/src/codetree/cache.py @@ -1,4 +1,7 @@ import json +import os +import stat +import tempfile from pathlib import Path @@ -19,7 +22,24 @@ def load(self): def save(self): """Write cache to disk, creating .codetree/ directory if needed.""" self._cache_file.parent.mkdir(parents=True, exist_ok=True) - self._cache_file.write_text(json.dumps(self._data, indent=2)) + # Write-then-rename so a server killed mid-save never leaves a torn file. + # The temp name is unique so concurrent servers on one repo cannot collide. + fd, tmp_name = tempfile.mkstemp( + dir=self._cache_file.parent, prefix="index.", suffix=".json.tmp" + ) + try: + with os.fdopen(fd, "w") as tmp_file: + tmp_file.write(json.dumps(self._data, indent=2)) + # mkstemp creates 0600; keep the previous mode, or a regular 0644. + try: + mode = stat.S_IMODE(self._cache_file.stat().st_mode) + except OSError: + mode = 0o644 + os.chmod(tmp_name, mode) + os.replace(tmp_name, self._cache_file) + except BaseException: + Path(tmp_name).unlink(missing_ok=True) + raise def get(self, rel_path: str) -> dict | None: return self._data.get(rel_path) diff --git a/src/codetree/graph/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..f014b0a 100644 --- a/src/codetree/graph/store.py +++ b/src/codetree/graph/store.py @@ -93,6 +93,13 @@ def commit(self): self._conn.commit() self._in_transaction = False + def rollback(self): + """Discard the current transaction (e.g. after a failed build).""" + with self._lock: + if self._conn: + self._conn.rollback() + self._in_transaction = False + def _auto_commit(self): """Commit unless inside an explicit transaction. @@ -142,6 +149,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 +235,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 +337,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 +376,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/index_state.py b/src/codetree/index_state.py new file mode 100644 index 0000000..55eb1e9 --- /dev/null +++ b/src/codetree/index_state.py @@ -0,0 +1,263 @@ +"""Index lifecycle: build the indexer and graph, optionally in a background thread. + +The MCP handshake must not wait for indexing — on a large repository a cold +build takes longer than an MCP client's startup timeout. IndexState lets the +server start immediately while tools wait (bounded) for the data they need: + +- files_ready — file discovery finished (single-file tools can parse on demand) +- index_ready — every file indexed (repo-wide tools) +- graph_ready — SQLite graph built (graph tools) + +Every event is set even when a phase fails, so waiters never hang; they check +``indexer`` / ``error`` afterwards. +""" + +import os +import threading +import time +from pathlib import Path + +from .cache import Cache +from .graph.builder import GraphBuilder +from .graph.queries import GraphQueries +from .graph.store import GraphStore +from .indexer import Indexer + + + +def _wait_timeout_from_env(default: float = 20.0) -> float: + try: + return float(os.environ.get("CODETREE_WAIT_TIMEOUT", default)) + except ValueError: + return default + + +# How long a tool call waits for indexing before answering "still building". +# Override with CODETREE_WAIT_TIMEOUT (seconds) for clients with short tool timeouts. +WAIT_TIMEOUT = _wait_timeout_from_env() + + +class IndexState: + def __init__(self, root: str | Path): + self.root = Path(root) + self.indexer: Indexer | None = None + self.graph_store = GraphStore(str(self.root)) + self.graph_store.open() + self.graph_queries = GraphQueries(self.graph_store) + # Stats of the last committed graph. index_status reports these while a + # build is running: live queries on the builder's connection would see + # its uncommitted, half-rewritten tables. + self._graph_snapshot = self._read_graph_stats() + + self.phase = "starting" + self.files_total = 0 + self.files_done = 0 + self.error: str | None = None + self.started_at = time.time() + self.ready_at: float | None = None + self._discovered: set[str] = set() + + self.files_ready = threading.Event() + self.index_ready = threading.Event() + self.graph_ready = threading.Event() + self._thread: threading.Thread | None = None + + # ── Building ───────────────────────────────────────────────────────── + + def start_background(self) -> None: + """Run build() in a daemon thread and return immediately.""" + self._thread = threading.Thread(target=self.build, name="codetree-index", daemon=True) + self._thread.start() + + def build(self, raise_errors: bool = False) -> None: + """Discover, index (reusing the skeleton cache), then build the graph. + + Failures are recorded in ``error`` for tools to report. With + raise_errors (synchronous startup) the exception is re-raised too, so + the caller fails loudly with the original traceback. + """ + try: + self._build_index() + except Exception as exc: + self._fail(f"indexing failed: {exc!r}") + self.graph_ready.set() # no graph without an index — release waiters + if raise_errors: + raise + return + finally: + self.files_ready.set() + self.index_ready.set() + + try: + self.phase = "building_graph" + GraphBuilder(str(self.root), self.graph_store).build(indexer=self.indexer) + except Exception as exc: + self._fail(f"graph build failed: {exc!r}") + try: + self.graph_store.rollback() # keep the last committed graph intact + except Exception: + pass # store already closed, e.g. at interpreter exit + if raise_errors: + raise + return + finally: + # Waiters must always be released, whatever failed above. + try: + self._graph_snapshot = self._read_graph_stats() + except Exception: + pass + self.graph_ready.set() + + self.phase = "ready" + self.ready_at = time.time() + + def _build_index(self) -> None: + self.phase = "discovering" + cache = Cache(self.root) + cache.load() + + indexer = Indexer(self.root) + files = indexer.discover_files() + self._discovered = {str(f.relative_to(self.root)) for f in files} + self.files_total = len(files) + self.files_ready.set() + + self.phase = "indexing" + # Entries written before has_errors was cached are re-parsed once. + cached_mtimes = { + k: v["mtime"] for k, v in (cache._data or {}).items() if "has_errors" in v + } + indexer.build(cached_mtimes=cached_mtimes, files=files, progress=self._on_progress) + + # Inject cached entries for unchanged files. Only files discovered by this + # build qualify, so stale cache entries (deleted or now-ignored files, + # e.g. old worktrees) are never resurrected. + for rel_path in indexer.cached_candidates: + path = self.root / rel_path + try: + mtime = path.stat().st_mtime + if cache.is_valid(rel_path, mtime): + indexer.inject_cached( + rel_path=rel_path, + py_file=path, + source=path.read_bytes(), + skeleton=cache.get(rel_path).get("skeleton", []), + mtime=mtime, + has_errors=cache.get(rel_path).get("has_errors", False), + ) + continue + except OSError: + continue # deleted since discovery + # Changed between discovery and injection — parse it instead + entry = indexer.index_file(path) + if entry is not None: + indexer._index[rel_path] = entry + + # Rebuild definition index once after all injections (DATA-01, DATA-02, DATA-03 fix) + indexer._rebuild_definitions() + + # Save updated cache — rebuilt from the index so stale entries are dropped + cache._data = {} + for rel_path, file_entry in indexer._index.items(): + cache.set(rel_path, { + "mtime": file_entry.mtime, + "skeleton": file_entry.skeleton, + "has_errors": file_entry.has_errors, + }) + try: + cache.save() + except OSError: + pass # the cache only speeds up the next start — never fail the index for it + + # Publish only the complete indexer: tools never see a half-built one. + self.indexer = indexer + + def _on_progress(self, done: int, total: int) -> None: + self.files_done = done + self.files_total = total + + def _fail(self, message: str) -> None: + self.phase = "error" + self.error = message + + # ── Access for tools ───────────────────────────────────────────────── + + def wait_for_index(self, timeout: float | None = None) -> Indexer | None: + """Return the full indexer once built; None if still indexing or failed.""" + self.index_ready.wait(WAIT_TIMEOUT if timeout is None else timeout) + return self.indexer + + def wait_for_graph(self, timeout: float | None = None) -> bool: + """True once the graph is built and usable.""" + self.graph_ready.wait(WAIT_TIMEOUT if timeout is None else timeout) + return self.graph_ready.is_set() and self.error is None + + def indexer_for_files(self, rel_paths: list[str], timeout: float | None = None) -> Indexer | None: + """Indexer able to answer about rel_paths, without waiting for the full index. + + Before the full index is ready, returns a throwaway Indexer holding just + those files, parsed on demand — limited to discovered files so results + match what the full index will say. None if discovery is not done yet. + """ + if self.index_ready.is_set(): + return self.indexer + if not self.files_ready.wait(WAIT_TIMEOUT if timeout is None else timeout): + return None + if self.index_ready.is_set(): + return self.indexer + partial = Indexer(self.root) + for rel_path in dict.fromkeys(rel_paths): + if rel_path not in self._discovered: + continue + entry = partial.index_file(self.root / rel_path) + if entry is not None: + partial._index[rel_path] = entry + partial._rebuild_definitions() + return partial + + def not_ready_message(self, need: str = "index") -> str: + """Explain why a tool cannot answer yet (or why it never will).""" + if self.error: + return f"codetree {self.error}. Restart the MCP server to retry." + elapsed = int(time.time() - self.started_at) + if self.phase == "indexing" and self.files_total: + where = f"indexing {self.files_done}/{self.files_total} files" + else: + where = self.phase.replace("_", " ") + what = "code graph" if need == "graph" else "repository index" + return ( + f"codetree is still building the {what} ({where}, {elapsed}s elapsed). " + "Retry in a few seconds; index_status shows progress." + ) + + def _read_graph_stats(self) -> dict: + return { + **self.graph_store.stats(), + "last_indexed_at": self.graph_store.get_meta("last_indexed_at"), + } + + def graph_stats(self) -> dict: + """files/symbols/edges/last_indexed_at of the last committed graph (never blocks).""" + if self.graph_ready.is_set(): + return self._read_graph_stats() + return dict(self._graph_snapshot) + + def status(self) -> dict: + """Lifecycle fields for index_status (never blocks).""" + result = { + "status": self.phase, + "files_discovered": self.files_total, + "files_indexed": self.files_done if self.phase == "indexing" else ( + len(self.indexer._index) if self.indexer else 0 + ), + "index_ready": self.index_ready.is_set() and self.indexer is not None, + "graph_ready": self.graph_ready.is_set() and self.error is None, + } + if self.ready_at is not None: + result["startup_seconds"] = round(self.ready_at - self.started_at, 2) + if self.error: + result["error"] = self.error + return result + + def close(self) -> None: + self.graph_store.close() diff --git a/src/codetree/indexer.py b/src/codetree/indexer.py index 0fb2256..1068333 100644 --- a/src/codetree/indexer.py +++ b/src/codetree/indexer.py @@ -1,5 +1,9 @@ +import os +import subprocess +import threading from pathlib import Path from dataclasses import dataclass +from typing import Callable from .languages.base import LanguagePlugin from .registry import get_plugin @@ -43,6 +47,13 @@ def __init__(self, root: str | Path): self._call_graph: dict[str, set[str]] = {} self._reverse_graph: dict[str, set[str]] = {} self._call_graph_built: bool = False + self._call_graph_lock = threading.Lock() + # Rel paths in discovery order; _rebuild_definitions() orders _index by + # it so results do not depend on which files came from the cache. + self._discovery_order: list[str] = [] + # 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 +75,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 +164,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,52 +183,75 @@ def _rebuild_definitions(self) -> None: if key not in self._name_to_qualified[bare]: self._name_to_qualified[bare].append(key) - def build(self, cached_mtimes: dict[str, float] | None = None): + 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, + files: list[Path] | None = None, + progress: Callable[[int, int], None] | None = None): """Index all supported files under root, skipping non-project dirs. Files whose path+mtime appear in cached_mtimes are skipped; the caller injects them via inject_cached(). + + Args: + cached_mtimes: rel_path → mtime of entries the caller has cached + files: pre-discovered file list (defaults to discover_files()) + progress: called as progress(done, total) after each file """ cached_mtimes = cached_mtimes or {} - 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 = [] + if files is None: + files = self.discover_files() + self._discovery_order = [str(f.relative_to(self.root)) for f in files] + total = len(files) + for done, candidate in enumerate(files, start=1): 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 + if progress is not None: + progress(done, total) # Build definition index from skeleton data (qualified keys, no duplicates, no ghosts) self._rebuild_definitions() def inject_cached(self, rel_path: str, py_file: Path, source: bytes, - skeleton: list[dict], mtime: float): + skeleton: list[dict], mtime: float, has_errors: bool = False): """Inject a pre-computed entry (from cache) without re-parsing.""" self._call_graph_built = False # invalidate so graph is rebuilt with new entry plugin = get_plugin(py_file) @@ -148,6 +264,7 @@ def inject_cached(self, rel_path: str, py_file: Path, source: bytes, mtime=mtime, language=py_file.suffix.lstrip("."), plugin=plugin, + has_errors=has_errors, ) # Note: _definitions is NOT updated here. After all inject_cached() calls # are complete, the caller must invoke _rebuild_definitions() to rebuild @@ -180,11 +297,24 @@ def get_call_graph(self, rel_path: str, function_name: str) -> dict: return {"calls": calls, "callers": callers} def _ensure_call_graph(self): - """Build repo-wide call graph lazily on first use.""" + """Build repo-wide call graph lazily on first use. + + Tools run concurrently (FastMCP thread pool): the lock makes the first + callers build it once, and the graph is published only when complete. + """ if self._call_graph_built: return - self._call_graph = {} - self._reverse_graph = {} + with self._call_graph_lock: + if self._call_graph_built: + return + call_graph, reverse_graph = self._compute_call_graph() + self._call_graph = call_graph + self._reverse_graph = reverse_graph + self._call_graph_built = True + + def _compute_call_graph(self) -> tuple[dict[str, set[str]], dict[str, set[str]]]: + call_graph: dict[str, set[str]] = {} + reverse_graph: dict[str, set[str]] = {} for rel_path, entry in self._index.items(): for item in entry.skeleton: if item["type"] in ("function", "method"): @@ -203,12 +333,12 @@ def _ensure_call_graph(self): else: # External/unresolved — keep as bare name callee_keys.add(f"?::{callee_name}") - self._call_graph[caller_key] = callee_keys + call_graph[caller_key] = callee_keys for ck in callee_keys: - if ck not in self._reverse_graph: - self._reverse_graph[ck] = set() - self._reverse_graph[ck].add(caller_key) - self._call_graph_built = True + if ck not in reverse_graph: + reverse_graph[ck] = set() + reverse_graph[ck].add(caller_key) + return call_graph, reverse_graph def find_dead_code(self, file_path: str | None = None) -> list[dict]: """Find symbols that are defined but never referenced elsewhere. diff --git a/src/codetree/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..851cf84 100644 --- a/src/codetree/server.py +++ b/src/codetree/server.py @@ -1,10 +1,18 @@ +import atexit from fastmcp import FastMCP from pathlib import Path -from .indexer import Indexer -from .cache import Cache +from .index_state import IndexState -def create_server(root: str) -> FastMCP: +def create_server(root: str, background: bool = False) -> FastMCP: + """Create the codetree MCP server for the repository at root. + + Args: + root: repository root + background: index in a background thread so the server can answer the + MCP handshake immediately (used by `codetree`). When False (default, + used by tests and embedding), indexing completes before returning. + """ mcp = FastMCP("codetree") root_path = Path(root) @@ -22,58 +30,17 @@ def _validate_path(file_path: str | None, _root: Path = root_path) -> str | None except ValueError: return f"Error: path '{file_path}' is outside the repo root — access denied" - # Load cache - cache = Cache(root) - cache.load() - - # Build index, skipping unchanged files - cached_mtimes = { - k: v["mtime"] for k, v in (cache._data or {}).items() - } - indexer = Indexer(root) - indexer.build(cached_mtimes=cached_mtimes) - - # Inject cached entries for unchanged files (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, - ) - - # Rebuild definition index once after all injections (DATA-01, DATA-02, DATA-03 fix) - indexer._rebuild_definitions() - - # Save updated cache - for rel_path, file_entry in indexer._index.items(): - cache.set(rel_path, { - "mtime": file_entry.mtime, - "skeleton": file_entry.skeleton, - }) - cache.save() - - # ── Build persistent graph ─────────────────────────────────────────── - import atexit - from .graph.store import GraphStore - from .graph.builder import GraphBuilder - from .graph.queries import GraphQueries - - graph_store = GraphStore(root) - graph_store.open() - atexit.register(graph_store.close) - graph_builder = GraphBuilder(root, graph_store) - graph_builder.build(indexer=indexer) - graph_queries = GraphQueries(graph_store) + # Index + graph lifecycle. In background mode tools wait (bounded) for the + # data they need instead of blocking server startup. + state = IndexState(root) + # Register the store only: a reference to state would pin the whole index. + atexit.register(state.graph_store.close) + if background: + state.start_background() + else: + state.build(raise_errors=True) + mcp._codetree_state = state + graph_queries = state.graph_queries # ── Skeleton formatting helpers ────────────────────────────────────────── _TYPE_ABBREV = { @@ -138,6 +105,9 @@ def get_file_skeleton(file_path: str, format: str = "full") -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() skeleton = indexer.get_skeleton(file_path) if not skeleton: return f"File not found or empty: {file_path}" @@ -156,6 +126,9 @@ def get_symbol(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() result = indexer.get_symbol(file_path, symbol_name) if result is None: return f"Symbol '{symbol_name}' not found in {file_path}" @@ -170,6 +143,9 @@ def find_references(symbol_name: str) -> str: symbol_name: name of the symbol to search for; results include file paths relative to the repo root (e.g., "src/main.py") """ + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() refs = indexer.find_references(symbol_name) if not refs: return f"No references found for '{symbol_name}'" @@ -192,6 +168,9 @@ def get_call_graph(file_path: str, function_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() graph = indexer.get_call_graph(file_path, function_name) lines = [f"Call graph for '{function_name}':"] @@ -220,6 +199,9 @@ def get_imports(file_path: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() entry = indexer._index.get(file_path) if entry is None: return f"File not found: {file_path}" @@ -241,6 +223,9 @@ def get_skeletons(file_paths: list[str], format: str = "full") -> str: """ if not file_paths: return "No files requested." + indexer = state.indexer_for_files(file_paths) + if indexer is None: + return state.not_ready_message() parts = [] for fp in file_paths: parts.append(f"=== {fp} ===") @@ -268,6 +253,9 @@ def get_symbols(symbols: list[dict]) -> str: """ if not symbols: return "No symbols requested." + indexer = state.indexer_for_files([item.get("file_path", "") for item in symbols]) + if indexer is None: + return state.not_ready_message() parts = [] for item in symbols: fp = item.get("file_path", "") @@ -295,6 +283,9 @@ def get_complexity(file_path: str, function_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return state.not_ready_message() entry = indexer._index.get(file_path) if entry is None: return f"File not found: {file_path}" @@ -318,6 +309,9 @@ def find_dead_code(file_path: str | None = None) -> str: if file_path is not None: if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path and file_path not in indexer._index: return f"File not found: {file_path}" dead = indexer.find_dead_code(file_path=file_path) @@ -350,6 +344,9 @@ def get_blast_radius(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path not in indexer._index: return f"File not found: {file_path}" result = indexer.get_blast_radius(file_path, symbol_name) @@ -394,6 +391,9 @@ def detect_clones(file_path: str | None = None, min_lines: int = 5) -> str: if file_path is not None: if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() clones = indexer.detect_clones(file_path=file_path, min_lines=min_lines) if not clones: scope = file_path if file_path else "the repo" @@ -429,6 +429,9 @@ def search_symbols(query: str | None = None, type: str | None = None, language: filter by file extension without dot (e.g., "py", "js", "go") format: "full" (default) or "compact" (abbreviated) """ + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() results = indexer.search_symbols( query=query, type=type, parent=parent, has_doc=has_doc, min_complexity=min_complexity, language=language, @@ -478,6 +481,9 @@ def find_tests(file_path: str, symbol_name: str) -> str: """ if err := _validate_path(file_path): return err + indexer = state.wait_for_index() + if indexer is None: + return state.not_ready_message() if file_path not in indexer._index: return f"File not found: {file_path}" tests = indexer.find_tests(file_path, symbol_name) @@ -493,13 +499,15 @@ def find_tests(file_path: str, symbol_name: str) -> str: @mcp.tool() def index_status() -> dict: - """Report on graph index freshness and stats.""" - stats = graph_store.stats() - last = graph_store.get_meta("last_indexed_at") + """Report on index build progress, graph freshness and stats. + + Never blocks: while the server is still indexing, "status" shows the + phase (discovering, indexing, building_graph) and file progress. + """ return { "graph_exists": True, - **stats, - "last_indexed_at": last, + **state.graph_stats(), + **state.status(), } @mcp.tool() @@ -512,6 +520,8 @@ def get_repository_map(max_items: int = 5) -> dict: Args: max_items: maximum items per section (default 5) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.repository_map(max_items=max_items) @mcp.tool() @@ -528,6 +538,8 @@ def resolve_symbol(query: str, kind: str | None = None, path_hint: prefer results from files matching this path limit: max results (default 10) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} results = graph_queries.resolve_symbol(query, kind=kind, path_hint=path_hint, limit=limit) return { "query": query, @@ -562,6 +574,8 @@ def search_graph(query: str | None = None, kind: str | None = None, limit: max results per page (default 10) offset: pagination offset (default 0) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.search_graph( query=query, kind=kind, file_pattern=file_pattern, relationship=relationship, direction=direction, @@ -581,6 +595,8 @@ def get_change_impact(symbol_query: str | None = None, diff_scope: "working" (uncommitted), "staged", or "HEAD~1" for git-based analysis depth: max hop depth (default 3) """ + if not state.wait_for_graph(): + return {"error": state.not_ready_message("graph")} return graph_queries.change_impact( symbol_query=symbol_query, diff_scope=diff_scope, @@ -608,10 +624,16 @@ def analyze_dataflow(file_path: str, function_name: str, return {"error": err} if mode == "cross_taint": + indexer = state.wait_for_index() + if indexer is None: + return {"error": state.not_ready_message()} if file_path not in indexer._index: return {"error": f"File not found: {file_path}"} return extract_cross_function_taint(indexer, file_path, function_name, depth=depth) + indexer = state.indexer_for_files([file_path]) + if indexer is None: + return {"error": state.not_ready_message()} entry = indexer._index.get(file_path) if entry is None: return {"error": f"File not found: {file_path}"} @@ -635,6 +657,9 @@ def find_hot_paths(top_n: int = 10) -> str: Args: top_n: max results to return (default 10) """ + if not state.wait_for_graph(): + return state.not_ready_message("graph") + indexer = state.indexer results = graph_queries.find_hot_paths(indexer, top_n=top_n) if not results: return "No hot paths found (no functions with both callers and complexity)." @@ -662,6 +687,8 @@ def get_dependency_graph(file_path: str | None = None, if file_path is not None: if err := _validate_path(file_path): return err + if not state.wait_for_graph(): + return state.not_ready_message("graph") result = graph_queries.get_dependency_graph(file_path=file_path, format=format) summary = f"\n\n{result['nodes']} files, {result['edges']} import edges" return result["content"] + summary @@ -742,6 +769,9 @@ def suggest_docs(file_path: str | None = None, if file_path is not None: if err := _validate_path(file_path): return err + if not state.wait_for_graph(): + return state.not_ready_message("graph") + indexer = state.indexer results = graph_queries.suggest_docs(indexer, file_path=file_path, symbol_name=symbol_name) if not results: return "No undocumented functions found." @@ -763,5 +793,5 @@ def suggest_docs(file_path: str | None = None, def run(root: str): - mcp = create_server(root) + mcp = create_server(root, background=True) mcp.run() diff --git a/tests/test_async_startup.py b/tests/test_async_startup.py new file mode 100644 index 0000000..7990f2c --- /dev/null +++ b/tests/test_async_startup.py @@ -0,0 +1,350 @@ +"""Tests for background indexing: the server answers before the index is built.""" + +import subprocess +import threading +import time + +import pytest + +from codetree import index_state +from codetree.indexer import Indexer +from codetree.server import create_server + + +def _tool(mcp, name): + return mcp.local_provider._components[f"tool:{name}@"].fn + + +@pytest.fixture +def repo(tmp_path): + root = tmp_path / "repo" + root.mkdir() + subprocess.run(["git", "init", "-q"], cwd=root, check=True) + (root / ".gitignore").write_text("ignored/\n.codetree/\n") + (root / "calc.py").write_text( + 'class Calculator:\n def add(self, a, b):\n return a + b\n\n' + 'def helper():\n return Calculator().add(1, 2)\n' + ) + (root / "main.py").write_text("from calc import helper\n\ndef main():\n helper()\n") + (root / "ignored").mkdir() + (root / "ignored" / "secret.py").write_text("def hidden():\n pass\n") + return root + + +@pytest.fixture(autouse=True) +def background_states(monkeypatch): + """Join every background build and close its store when the test ends.""" + started = [] + original = index_state.IndexState.start_background + + def tracking_start(self): + started.append(self) + original(self) + + monkeypatch.setattr(index_state.IndexState, "start_background", tracking_start) + yield started + for state in started: + if state._thread is not None: + state._thread.join(15) + state.close() + + +@pytest.fixture +def gate(monkeypatch, background_states): + """Block background indexing (only the codetree-index thread) until released. + + Depends on background_states so the gate opens before threads are joined. + """ + release = threading.Event() + original = Indexer.index_file + + def gated_index_file(self, path, mtime=None): + if threading.current_thread().name == "codetree-index": + release.wait(10) + return original(self, path, mtime) + + monkeypatch.setattr(Indexer, "index_file", gated_index_file) + monkeypatch.setattr(index_state, "WAIT_TIMEOUT", 0.3) + yield release + release.set() + + +def _wait(event, timeout=10): + assert event.wait(timeout), "background indexing did not finish" + + +# ── Startup does not wait ──────────────────────────────────────────────────── + +def test_background_create_server_returns_before_indexing(repo, gate): + start = time.monotonic() + mcp = create_server(str(repo), background=True) + assert time.monotonic() - start < 5 + state = mcp._codetree_state + assert not state.index_ready.is_set() + gate.set() + _wait(state.graph_ready) + + +def test_index_status_never_blocks_and_reports_progress(repo, gate): + mcp = create_server(str(repo), background=True) + state = mcp._codetree_state + _wait(state.files_ready) + + start = time.monotonic() + status = _tool(mcp, "index_status")() + assert time.monotonic() - start < 1 + assert status["status"] in ("indexing", "discovering") + assert status["files_discovered"] == 2 + assert status["index_ready"] is False + assert status["graph_ready"] is False + + gate.set() + _wait(state.graph_ready) + status = _tool(mcp, "index_status")() + assert status["status"] == "ready" + assert status["index_ready"] is True and status["graph_ready"] is True + assert status["files_indexed"] == 2 + assert "startup_seconds" in status + + +# ── Single-file tools answer during indexing ───────────────────────────────── + +def test_single_file_tools_parse_on_demand_while_indexing(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + skeleton = _tool(mcp, "get_file_skeleton")(file_path="calc.py") + assert "class Calculator" in skeleton + assert "def add" in skeleton + assert "def helper" in _tool(mcp, "get_symbol")(file_path="calc.py", symbol_name="helper") + assert "from calc import helper" in _tool(mcp, "get_imports")(file_path="main.py") + assert "Complexity of helper()" in _tool(mcp, "get_complexity")( + file_path="calc.py", function_name="helper") + multi = _tool(mcp, "get_skeletons")(file_paths=["calc.py", "main.py"]) + assert "class Calculator" in multi and "def main" in multi + flow = _tool(mcp, "analyze_dataflow")(file_path="calc.py", function_name="helper") + assert "error" not in flow + assert not mcp._codetree_state.index_ready.is_set() + + +def test_on_demand_parsing_respects_gitignore(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "get_file_skeleton")(file_path="ignored/secret.py") + assert result == "File not found or empty: ignored/secret.py" + + +# ── Repo-wide and graph tools wait, then explain ───────────────────────────── + +def test_repo_wide_tool_reports_progress_when_index_not_ready(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "find_references")(symbol_name="helper") + assert "still building the repository index" in result + assert "/2 files" in result + assert "index_status" in result + + +def test_graph_tools_report_not_ready(repo, gate): + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + + result = _tool(mcp, "get_repository_map")() + assert "still building the code graph" in result["error"] + assert "still building the code graph" in _tool(mcp, "find_hot_paths")() + + +def test_repo_wide_tool_waits_for_index_finishing_in_time(repo, gate, monkeypatch): + monkeypatch.setattr(index_state, "WAIT_TIMEOUT", 10) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.files_ready) + threading.Timer(0.2, gate.set).start() + + result = _tool(mcp, "find_references")(symbol_name="helper") + assert "calc.py" in result and "main.py" in result + + +def test_tools_return_full_results_once_ready(repo, gate): + mcp = create_server(str(repo), background=True) + gate.set() + _wait(mcp._codetree_state.graph_ready) + + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + assert "ignored/secret.py" not in _tool(mcp, "find_references")(symbol_name="hidden") + matches = _tool(mcp, "resolve_symbol")(query="helper")["matches"] + assert [m["file"] for m in matches] == ["calc.py"] + + +# ── Failures never hang tools ──────────────────────────────────────────────── + +def test_indexing_failure_is_reported_not_hung(repo, monkeypatch): + def boom(self): + raise RuntimeError("disk on fire") + + monkeypatch.setattr(Indexer, "discover_files", boom) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + + start = time.monotonic() + result = _tool(mcp, "find_references")(symbol_name="helper") + assert time.monotonic() - start < 3 + assert "indexing failed" in result and "disk on fire" in result + assert "indexing failed" in _tool(mcp, "get_file_skeleton")(file_path="calc.py") + status = _tool(mcp, "index_status")() + assert status["status"] == "error" + assert "disk on fire" in status["error"] + + +def test_graph_failure_keeps_structural_tools_working(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + + def boom(self, indexer=None): + raise RuntimeError("graph exploded") + + monkeypatch.setattr(GraphBuilder, "build", boom) + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + assert "graph build failed" in _tool(mcp, "get_repository_map")()["error"] + assert _tool(mcp, "index_status")()["graph_ready"] is False + + +# ── Synchronous mode (default) ─────────────────────────────────────────────── + +def test_default_create_server_is_ready_on_return(repo): + mcp = create_server(str(repo)) + status = _tool(mcp, "index_status")() + assert status["status"] == "ready" + assert status["index_ready"] and status["graph_ready"] + + +def test_warm_start_reuses_cache_in_background(repo): + first = create_server(str(repo)) + first._codetree_state.close() + + mcp = create_server(str(repo), background=True) + _wait(mcp._codetree_state.graph_ready) + assert "class Calculator" in _tool(mcp, "get_file_skeleton")(file_path="calc.py") + assert _tool(mcp, "index_status")()["files_indexed"] == 2 + + +def test_sync_create_server_raises_on_indexing_failure(repo, monkeypatch): + def boom(self): + raise RuntimeError("disk on fire") + + monkeypatch.setattr(Indexer, "discover_files", boom) + with pytest.raises(RuntimeError, match="disk on fire"): + create_server(str(repo)) + + +def test_cache_save_failure_does_not_fail_indexing(repo, monkeypatch): + from codetree.cache import Cache + + def readonly(self): + raise PermissionError("read-only file system") + + monkeypatch.setattr(Cache, "save", readonly) + mcp = create_server(str(repo)) + + assert _tool(mcp, "index_status")()["status"] == "ready" + assert "main.py" in _tool(mcp, "find_references")(symbol_name="helper") + + +def test_atexit_does_not_pin_the_index(repo, monkeypatch): + registered = [] + monkeypatch.setattr("codetree.server.atexit.register", registered.append) + mcp = create_server(str(repo)) + + # Only the SQLite store is registered — not IndexState, which owns every + # file's source through its indexer. + assert registered == [mcp._codetree_state.graph_store.close] + + +def test_index_status_reports_committed_graph_during_rebuild(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + + first = create_server(str(repo)) + committed = _tool(first, "index_status")() + first._codetree_state.close() + assert committed["symbols"] > 0 + + # Hold the next build mid-transaction, after it deleted rows. + in_build = threading.Event() + release = threading.Event() + original = GraphBuilder.build + + def held_build(self, indexer=None): + self._store.begin() + self._store.delete_symbols_for_file("calc.py") + self._store.delete_symbols_for_file("main.py") + in_build.set() + release.wait(10) + return original(self, indexer=indexer) + + monkeypatch.setattr(GraphBuilder, "build", held_build) + mcp = create_server(str(repo), background=True) + try: + assert in_build.wait(10) + during = _tool(mcp, "index_status")() + assert during["status"] == "building_graph" + assert during["symbols"] == committed["symbols"] + assert during["last_indexed_at"] == committed["last_indexed_at"] + finally: + release.set() + _wait(mcp._codetree_state.graph_ready) + after = _tool(mcp, "index_status")() + assert after["status"] == "ready" + assert after["last_indexed_at"] != committed["last_indexed_at"] + + +def test_syntax_error_flag_survives_warm_start(repo): + (repo / "broken.py").write_text("def broken(:\n pass\n") + cold = _tool(create_server(str(repo)), "get_file_skeleton")(file_path="broken.py") + warm = _tool(create_server(str(repo)), "get_file_skeleton")(file_path="broken.py") + assert "WARNING: File has syntax errors" in cold + assert warm == cold + + +def test_cache_entries_without_error_flag_are_reparsed(repo): + import json + + (repo / "broken.py").write_text("def broken(:\n pass\n") + create_server(str(repo)) + cache_file = repo / ".codetree" / "index.json" + data = json.loads(cache_file.read_text()) + for entry in data.values(): + del entry["has_errors"] # as written by an older codetree + cache_file.write_text(json.dumps(data)) + + mcp = create_server(str(repo)) + assert "WARNING: File has syntax errors" in _tool(mcp, "get_file_skeleton")(file_path="broken.py") + + +def test_wait_timeout_from_env(monkeypatch): + monkeypatch.setenv("CODETREE_WAIT_TIMEOUT", "3.5") + assert index_state._wait_timeout_from_env() == 3.5 + monkeypatch.setenv("CODETREE_WAIT_TIMEOUT", "soon") + assert index_state._wait_timeout_from_env() == 20.0 + + +def test_graph_failure_with_failing_rollback_still_releases_waiters(repo, monkeypatch): + from codetree.graph.builder import GraphBuilder + from codetree.graph.store import GraphStore + + def boom(self, indexer=None): + raise RuntimeError("graph exploded") + + def broken_rollback(self): + raise RuntimeError("connection gone") + + monkeypatch.setattr(GraphBuilder, "build", boom) + monkeypatch.setattr(GraphStore, "rollback", broken_rollback) + mcp = create_server(str(repo), background=True) + + _wait(mcp._codetree_state.graph_ready) + status = _tool(mcp, "index_status")() + assert status["status"] == "error" + assert "graph exploded" in status["error"] diff --git a/tests/test_cache.py b/tests/test_cache.py index 279bff1..51fede3 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -37,3 +37,19 @@ def test_cache_creates_directory_if_missing(tmp_path): cache = Cache(tmp_path) cache.save() assert cache_dir.exists() + + +def test_cache_file_is_world_readable_like_before(tmp_path): + import os + import stat + from codetree.cache import Cache + + cache = Cache(tmp_path) + cache.set("a.py", {"mtime": 1.0, "skeleton": []}) + cache.save() + cache_file = tmp_path / ".codetree" / "index.json" + assert stat.S_IMODE(cache_file.stat().st_mode) == 0o644 + + os.chmod(cache_file, 0o600) # a user's explicit choice is kept + cache.save() + assert stat.S_IMODE(cache_file.stat().st_mode) == 0o600 diff --git a/tests/test_file_discovery.py b/tests/test_file_discovery.py new file mode 100644 index 0000000..6d8de8b --- /dev/null +++ b/tests/test_file_discovery.py @@ -0,0 +1,325 @@ +"""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"] + + +def test_concurrent_cache_saves_never_fail_or_tear(tmp_path): + import threading + from codetree.cache import Cache + + errors = [] + + def writer(tag): + cache = Cache(tmp_path) + cache.set("x.py", {"mtime": 1.0, "skeleton": [{"name": tag * 200}]}) + try: + for _ in range(50): + cache.save() + except Exception as exc: + errors.append(exc) + + threads = [threading.Thread(target=writer, args=(tag,)) for tag in "abcd"] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert errors == [] + data = json.loads((tmp_path / ".codetree" / "index.json").read_text()) + assert data["x.py"]["mtime"] == 1.0 + assert list((tmp_path / ".codetree").glob("*.tmp")) == [] diff --git a/tests/test_graph_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..62e98d6 100644 --- a/tests/test_graph_store.py +++ b/tests/test_graph_store.py @@ -117,6 +117,34 @@ 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 TestTransactions: + def test_rollback_discards_uncommitted_writes(self, store): + store.upsert_edge(Edge("kept.py::a", "x.py::b", "CALLS")) + store.begin() + store.upsert_edge(Edge("dropped.py::a", "x.py::b", "CALLS")) + store.rollback() + assert len(store.edges_from("kept.py::a")) == 1 + assert store.edges_from("dropped.py::a") == [] + assert store._in_transaction is False + class TestFileCRUD: def test_upsert_and_get_file(self, store): @@ -159,3 +187,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_indexer.py b/tests/test_indexer.py index d8feadd..cb889f1 100644 --- a/tests/test_indexer.py +++ b/tests/test_indexer.py @@ -338,3 +338,35 @@ def test_unknown_file_has_empty_calls(self, sample_repo): idx.build() graph = idx.get_call_graph("missing.py", "fn") assert graph["calls"] == [] + + +def test_concurrent_call_graph_requests_build_it_once(tmp_path, monkeypatch): + import threading + import time + from codetree.indexer import Indexer + + (tmp_path / "a.py").write_text("def a():\n b()\n\ndef b():\n pass\n") + indexer = Indexer(tmp_path) + indexer.build() + builds = [] + original = Indexer._compute_call_graph + + def slow_compute(self): + builds.append(1) + time.sleep(0.2) + return original(self) + + monkeypatch.setattr(Indexer, "_compute_call_graph", slow_compute) + results = [] + threads = [ + threading.Thread(target=lambda: results.append(indexer.get_blast_radius("a.py", "b"))) + for _ in range(4) + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert builds == [1] + assert all(r == results[0] for r in results) + assert any(c["name"] == "a" for c in results[0]["callers"]) 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"]