tree-sitting
Symbol-level navigation of a local checkout using tree-sitter ASTs. Answers where a symbol is defined, what lines it spans, which symbols a file exposes, what a directory holds, and where a name is referenced — every answer carries exact line ranges to feed straight into a scoped
Install
npx skills add https://github.com/oaustegard/claude-skills/tree/main/plugins/code-intelligence/skills/tree-sitting
claude plugin marketplace add https://llmmart.ai/marketplace.json && claude plugin install oaustegard-claude-skills@llmmart
git clone https://github.com/oaustegard/claude-skills.git
The skills CLI installs just this skill, for any of its supported agents. Claude Code installs the whole oaustegard/claude-skills collection as a plugin from our marketplace. Git is the plain clone.
README
tree-sitting
AST-powered code navigation using tree-sitter. Parses all source files in a codebase into in-memory syntax trees, then provides fast query tools for symbol search, file/directory overview, source retrieval, and reference finding.
Features
- Fast scanning — parses ~250 files in ~700ms, then all queries are sub-millisecond from cache
- Symbol search — find by exact name, substring, or glob pattern across the entire codebase
- Directory and file overview — structural summaries with symbol counts, signatures, and doc comments
- Source retrieval — fetch implementation of any symbol, preferring definitions over declarations
- Reference finding — locate all textual references to a symbol via fast grep against cached source
- 12 grammars — Python, JavaScript, TypeScript, TSX, Go, Rust, Ruby, Java, C, HTML, Markdown, Mojo. Bundled as
parsers/*.sofor Linux x86_64. Other platforms installtree-sitter-<lang>wheels. - Three-tier extraction — custom extractors (richest), community tags.scm queries, and generic heuristic fallback
- Dual deployment — direct Python calls in Claude.ai, or long-lived MCP server in Claude Code
Dependencies
- tree-sitter — bare parser runtime
- tree-sitter- wheels — grammars off Linux x86_64 (macOS, arm64)
- fastmcp — required only for MCP server mode (Claude Code)
Skill manifest
tree-sitting
Tree-sitter symbol index over a local checkout. Every result carries line
ranges, so the next step is a scoped Read.
Setup
Grammars load from the first source that works: $TREESIT_PARSERS_DIR
(default ~/.cache/tree-sitting/parsers/), then the bundled parsers/*.so
(Linux x86_64 only), then PyPI wheels.
pip install tree-sitter # Linux x86_64: done
pip install tree-sitter tree-sitter-{python,javascript,typescript,go,rust,ruby,java,c,html,markdown} # macOS, arm64
Use the pip that belongs to the python3 you run the CLI with. Mojo has no
wheel. To get it, compile src/parser.c src/scanner.c from
oaustegard/tree-sitter-mojo with cc -shared -fPIC -I src into
~/.cache/tree-sitting/parsers/libtree_sitter_mojo.dylib.
A file whose grammar didn't load is skipped, and stderr prints
WARNING: no grammar for … with the install command. Fix that before
trusting an empty result.
Use
scripts/treesit.py in this skill's directory scans the repo on every call
(cached on disk; invalidated when files or grammars change), prints a tree
overview, then answers the queries. Batch queries so one scan serves them all.
python3 scripts/treesit.py REPO # overview, depth 1
python3 scripts/treesit.py REPO --path=src/core --detail=full # drill in
python3 scripts/treesit.py REPO --no-tree 'find:Parser*' 'source:parse_input' 'refs:ParseState'
| Query | Returns |
|---|---|
find:PATTERN[:KIND[:LIMIT]] |
symbols by name, substring or glob |
symbols:FILE |
every symbol in a file |
source:SYMBOL[:FILE] |
the symbol's source |
refs:SYMBOL[:LIMIT] |
textual references |
imports:FILE |
a file's imports |
dir:PATH |
directory overview |
Flags: --depth N (-1 = all), --detail sparse|normal|full, --path DIR,
--skip DIRS, --no-tree, --stats, --no-cache, --rebuild-cache.
scripts/engine.py exposes CodeCache for use inside one Python process.
Languages: Python, JavaScript, TypeScript/TSX, Go, Rust, Ruby, Java, C, HTML,
Markdown (heading outline), Mojo. Other languages get generic extraction if
their tree-sitter-<lang> wheel is installed.
Use something else for
- a repo that isn't on disk:
accessing-github-repos - what a codebase does:
featuring - a first look at an unfamiliar repo:
exploring-codebases - binding-resolved Python callers:
searching-codebases --refs - literal text or regex:
rg
Files (claude-skills)
-
scripts
-
engine.py 76.8 KB
""" tree-sitting engine: AST cache + symbol extraction using tree-sitter. Parses source files, caches ASTs in memory, and provides query APIs. Designed to be held in a long-lived process (MCP server) for fast queries. """ import fnmatch import hashlib import json import os import tempfile from dataclasses import dataclass, field from pathlib import Path # Grammar sources, tried in order per language (first that loads wins): # 1. user-built $TREESIT_PARSERS_DIR or ~/.cache/tree-sitting/parsers, # libtree_sitter_<lang>.{dylib,so} compiled on this host # 2. bundled parsers/libtree_sitter_<lang>.so — Linux x86_64 only # 3. PyPI wheel `pip install tree-sitter-<lang>` — every platform with a # wheel, which is how macOS and Linux arm64 get grammars # We deliberately do NOT depend on tree-sitter-language-pack: its 1.6.x # wheel layout is broken in the Claude.ai container and it downloads grammars # at runtime from a domain outside the network allowlist. _parsers: dict = {} _languages: dict = {} _grammar_sources: dict = {} # lang -> 'user' | 'bundled' | 'wheel' | None _PARSERS_DIR = Path(__file__).parent.parent / 'parsers' # The languages this skill extracts symbols for; a missing grammar for one of # these is worth a warning. Other EXT_TO_LANG entries load if a wheel exists. CORE_LANGS = ('python', 'javascript', 'typescript', 'tsx', 'go', 'rust', 'ruby', 'java', 'c', 'html', 'markdown', 'mojo') # lang -> (module, function) where the wheel's naming isn't tree_sitter_<lang>.language _WHEEL_OVERRIDES = { 'typescript': ('tree_sitter_typescript', 'language_typescript'), 'tsx': ('tree_sitter_typescript', 'language_tsx'), } def _wheel_package(lang: str) -> str | None: """PyPI package that provides lang's grammar, or None if there isn't one.""" if lang == 'mojo': return None module = _WHEEL_OVERRIDES.get(lang, (f'tree_sitter_{lang}',))[0] return module.replace('_', '-') def _user_parsers_dir() -> Path: env = os.environ.get('TREESIT_PARSERS_DIR') return Path(env) if env else Path.home() / '.cache' / 'tree-sitting' / 'parsers' def _load_shared_lib(path: Path, lang: str): """Load a tree_sitter.Language from a compiled grammar via ctypes, or None.""" try: from ctypes import CDLL, c_char_p, c_void_p, py_object, pythonapi from tree_sitter import Language except ImportError: return None try: lib = CDLL(str(path)) fn = getattr(lib, f'tree_sitter_{lang}') fn.restype = c_void_p # Wrap the pointer in a PyCapsule — the forward-compatible # tree_sitter.Language API since 0.23. pythonapi.PyCapsule_New.restype = py_object pythonapi.PyCapsule_New.argtypes = [c_void_p, c_char_p, c_void_p] capsule = pythonapi.PyCapsule_New(fn(), b"tree_sitter.Language", None) return Language(capsule) except Exception: # Wrong OS/arch (an ELF .so on macOS raises OSError here), missing # symbol, or an ABI the binding rejects. return None def _load_wheel(lang: str): """Load a tree_sitter.Language from an installed tree-sitter-<lang> wheel, or None.""" module, func = _WHEEL_OVERRIDES.get(lang, (f'tree_sitter_{lang}', 'language')) try: import importlib from tree_sitter import Language return Language(getattr(importlib.import_module(module), func)()) except Exception: return None def _load_language(lang: str): """Load lang's grammar from the first source that works; None if none does.""" for source, candidates in ( ('user', [_user_parsers_dir() / f'libtree_sitter_{lang}{ext}' for ext in ('.dylib', '.so')]), ('bundled', [_PARSERS_DIR / f'libtree_sitter_{lang}.so']), ): for path in candidates: if path.is_file(): language = _load_shared_lib(path, lang) if language is not None: _grammar_sources[lang] = source return language language = _load_wheel(lang) _grammar_sources[lang] = 'wheel' if language is not None else None return language def _get_language(lang: str): """Memoized _load_language.""" if lang not in _languages: _languages[lang] = _load_language(lang) return _languages[lang] def grammar_source(lang: str) -> str | None: """Where lang's grammar loaded from ('user'/'bundled'/'wheel'), or None.""" _get_language(lang) return _grammar_sources.get(lang) def missing_grammar_hint(langs) -> str | None: """One-line install hint for CORE_LANGS in langs whose grammar didn't load.""" missing = sorted(lang for lang in set(langs) if lang in CORE_LANGS and _get_language(lang) is None) if not missing: return None pkgs = sorted({pkg for lang in missing if (pkg := _wheel_package(lang))}) parts = [f"no grammar for {', '.join(missing)} — those files were skipped."] if pkgs: parts.append(f"Fix: pip install tree-sitter {' '.join(pkgs)}") if 'mojo' in missing: parts.append(f"Mojo has no wheel: compile oaustegard/tree-sitter-mojo's src/ " f"into {_user_parsers_dir()}/libtree_sitter_mojo.dylib") return ' '.join(parts) EXT_TO_LANG = { '.py': 'python', '.pyi': 'python', '.js': 'javascript', '.jsx': 'javascript', '.mjs': 'javascript', '.ts': 'typescript', '.tsx': 'tsx', '.go': 'go', '.rs': 'rust', '.rb': 'ruby', '.java': 'java', '.c': 'c', '.h': 'c', '.cpp': 'cpp', '.cc': 'cpp', '.cxx': 'cpp', '.hpp': 'cpp', '.hh': 'cpp', '.cs': 'c_sharp', '.swift': 'swift', '.kt': 'kotlin', '.kts': 'kotlin', '.scala': 'scala', '.html': 'html', '.htm': 'html', '.css': 'css', '.md': 'markdown', '.json': 'json', '.yaml': 'yaml', '.yml': 'yaml', '.toml': 'toml', '.lua': 'lua', '.sh': 'bash', '.bash': 'bash', '.el': 'elisp', '.zig': 'zig', '.ex': 'elixir', '.exs': 'elixir', '.mojo': 'mojo', '.🔥': 'mojo', } DEFAULT_SKIP = { '.git', 'node_modules', '__pycache__', '.venv', 'venv', 'dist', 'build', '.next', '.mypy_cache', '.pytest_cache', '.tox', 'target', '.cache', 'vendor', 'coverage', '.eggs', '*.egg-info', } # Cache format version — bump this to invalidate all existing caches CACHE_FORMAT_VERSION = 1 def cache_path_for(root: str) -> Path: """Determine deterministic cache file path for a root directory. Honors TREESIT_CACHE_DIR environment variable if set, otherwise uses system temp directory. Filename derived from SHA256 of resolved abspath. Args: root: Source root path (may contain symlinks or be relative) Returns: pathlib.Path to cache file (may not exist) """ # Resolve to absolute canonical path root_resolved = str(Path(root).resolve()) # Determine cache directory cache_dir_env = os.environ.get('TREESIT_CACHE_DIR') if cache_dir_env: cache_dir = Path(cache_dir_env) else: cache_dir = Path(tempfile.gettempdir()) / 'treesit-cache' # Ensure cache directory exists cache_dir.mkdir(parents=True, exist_ok=True) # Derive filename from SHA256 of resolved path cache_hash = hashlib.sha256(root_resolved.encode()).hexdigest() return cache_dir / f'{cache_hash}.json' @dataclass class Symbol: name: str kind: str file: str # relative path line: int end_line: int signature: str | None = None doc: str | None = None children: list['Symbol'] = field(default_factory=list) def to_dict(self, include_children=True) -> dict: d = { 'name': self.name, 'kind': self.kind, 'file': self.file, 'line': self.line, 'end_line': self.end_line, } if self.signature: d['signature'] = self.signature if self.doc: d['doc'] = self.doc if include_children and self.children: d['children'] = [c.to_dict(include_children=False) for c in self.children] return d @staticmethod def from_dict(d: dict) -> 'Symbol': """Deserialize a Symbol from a dict (from cache JSON). Reconstructs one level of children (to_dict stores one level). """ children = [] if d.get('children'): for child_dict in d['children']: children.append(Symbol.from_dict(child_dict)) return Symbol( name=d['name'], kind=d['kind'], file=d['file'], line=d['line'], end_line=d['end_line'], signature=d.get('signature'), doc=d.get('doc'), children=children, ) def format_oneline(self) -> str: """Format as a concise one-line string.""" parts = [f"{self.name} ({self.kind})"] if self.signature: parts.append(f"`{self.signature}`") parts.append(f":{self.line}-{self.end_line}") if self.doc: parts.append(f"— {self.doc}") return ' '.join(parts) def _get_parser(lang: str): """Get or create a cached parser for the given language. Returns None if the grammar isn't bundled (or fails to load). Callers handle None by skipping the file — the same behaviour the scan loop already expects. """ if lang not in _parsers: try: from tree_sitter import Parser except ImportError: _parsers[lang] = None return None language = _get_language(lang) if language is None: _parsers[lang] = None else: try: _parsers[lang] = Parser(language) except Exception: _parsers[lang] = None return _parsers[lang] def _get_text(node, source: bytes) -> str: return source[node.start_byte:node.end_byte].decode('utf-8', errors='replace') def _first_doc_line(text: str) -> str: """Extract first meaningful line from a comment.""" text = text.strip() # Strip comment markers for prefix in ('/**', '/*', '///', '//', '#'): text = text.removeprefix(prefix) text = text.rstrip('*/').strip() for line in text.split('\n'): line = line.strip().lstrip('*#/').strip() if line and not line.startswith('@') and not line.startswith('\\'): return line return '' def _preceding_doc(siblings: list, idx: int, source: bytes) -> str | None: """Get doc comment preceding siblings[idx].""" if idx <= 0: return None target_line = siblings[idx].start_point[0] prev = siblings[idx - 1] if prev.type != 'comment': return None if target_line - prev.end_point[0] > 1: return None text = _get_text(prev, source) result = _first_doc_line(text) return result if result else None def _python_docstring(node, source: bytes) -> str | None: """Extract docstring from Python function/class body.""" for child in node.children: if child.type == 'block': for stmt in child.children: # New grammar: string directly in block if stmt.type == 'string': text = _get_text(stmt, source).strip('"""').strip("'''").strip() return text.split('\n')[0].strip() or None # Old grammar: string inside expression_statement if stmt.type == 'expression_statement': for expr in stmt.children: if expr.type == 'string': text = _get_text(expr, source).strip('"""').strip("'''").strip() return text.split('\n')[0].strip() or None elif stmt.type != 'comment': break break return None # ─── Extractors ─────────────────────────────────────────────────────────── def _unwrap_decorated(node): """Return the def/class inside a ``decorated_definition``, else the node. Python's grammar wraps ``@deco\\ndef f(): ...`` in a ``decorated_definition`` whose payload is the real ``function_definition`` / ``class_definition``. A walk that matches only the bare types drops every decorated symbol, which in practice is much of a module's public API surface. The unwrapped node is deliberately what gets recorded, so ``line`` points at the ``def`` rather than at the first decorator. Callers that seed a language server from these positions (searching-codebases ``--refs``/``--def``) need the line the symbol token is actually on. The other extractors survive wrapper nodes already: ``_extract_generic`` recurses two levels, and the JS path unwraps ``export_statement`` explicitly. Python's extractor is the only non-recursing one, so it needs this. """ if node.type == 'decorated_definition': for c in node.children: if c.type in ('function_definition', 'class_definition'): return c return node def _extract_python(tree, source: bytes, relpath: str) -> list[Symbol]: symbols = [] module = tree.root_node children = [_unwrap_decorated(n) for n in module.children] for i, node in enumerate(children): if node.type == 'function_definition': name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') if name and not name.startswith('_'): sig = next((_get_text(c, source) for c in node.children if c.type == 'parameters'), None) doc = _python_docstring(node, source) sym = Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc) # Extract methods if it's weirdly at module level (rare) symbols.append(sym) elif node.type == 'class_definition': name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') if name and not name.startswith('_'): doc = _python_docstring(node, source) sym = Symbol(name=name, kind='class', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc) # Extract methods for child in node.children: if child.type == 'block': for sc in (_unwrap_decorated(c) for c in child.children): if sc.type == 'function_definition': mname = next((_get_text(c, source) for c in sc.children if c.type == 'identifier'), '') if mname: msig = next((_get_text(c, source) for c in sc.children if c.type == 'parameters'), None) mdoc = _python_docstring(sc, source) sym.children.append(Symbol( name=mname, kind='method', file=relpath, line=sc.start_point[0]+1, end_line=sc.end_point[0]+1, signature=msig, doc=mdoc)) symbols.append(sym) return symbols def _extract_c(tree, source: bytes, relpath: str) -> list[Symbol]: symbols = [] containers = {'preproc_ifdef', 'preproc_ifndef', 'preproc_if', 'preproc_else', 'preproc_elif', 'linkage_specification', 'declaration_list'} def collect(node): for child in node.children: if child.type in containers: yield from collect(child) else: yield child nodes = list(collect(tree.root_node)) def find_func_name(node): for c in node.children: if c.type == 'function_declarator': for cc in c.children: if cc.type == 'identifier': return _get_text(cc, source) if c.type == 'pointer_declarator': result = find_func_name(c) if result: return result return '' def return_type(node): parts = [] for c in node.children: if c.type in ('function_declarator', 'pointer_declarator', 'compound_statement'): break if c.type == 'storage_class_specifier': continue if c.type in ('primitive_type', 'type_identifier', 'sized_type_specifier', 'type_qualifier'): parts.append(_get_text(c, source)) has_ptr = any(c.type == 'pointer_declarator' for c in node.children) rt = ' '.join(parts) return (rt + ' *').strip() if has_ptr else rt def params(node): def find_fd(n): for c in n.children: if c.type == 'function_declarator': return c if c.type == 'pointer_declarator': r = find_fd(c) if r: return r return None fd = find_fd(node) if fd: for c in fd.children: if c.type == 'parameter_list': return ' '.join(_get_text(c, source).split()) return '' def is_static(node): return any(c.type == 'storage_class_specifier' and _get_text(c, source) == 'static' for c in node.children) for i, node in enumerate(nodes): if node.type == 'function_definition' and not is_static(node): name = find_func_name(node) if name: rt = return_type(node) p = params(node) sig = f"{p} -> {rt}" if rt else p doc = _preceding_doc(nodes, i, source) symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) elif node.type == 'declaration' and not is_static(node): has_fd = any(c.type in ('function_declarator', 'pointer_declarator') for c in node.children) if has_fd: name = find_func_name(node) if name: rt = return_type(node) p = params(node) sig = f"{p} -> {rt}" if rt else p doc = _preceding_doc(nodes, i, source) symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) elif node.type == 'type_definition': name = next((_get_text(c, source) for c in node.children if c.type == 'type_identifier'), '') kind = 'typedef' for c in node.children: if c.type == 'struct_specifier': kind = 'struct' elif c.type == 'enum_specifier': kind = 'enum' elif c.type == 'union_specifier': kind = 'union' if name: doc = _preceding_doc(nodes, i, source) symbols.append(Symbol(name=name, kind=kind, file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'struct_specifier': name = next((_get_text(c, source) for c in node.children if c.type == 'type_identifier'), '') if name: doc = _preceding_doc(nodes, i, source) symbols.append(Symbol(name=name, kind='struct', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'enum_specifier': name = next((_get_text(c, source) for c in node.children if c.type == 'type_identifier'), '') if name: doc = _preceding_doc(nodes, i, source) symbols.append(Symbol(name=name, kind='enum', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'preproc_def': name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') value = next((_get_text(c, source).strip() for c in node.children if c.type == 'preproc_arg'), '') if name and name.isupper() and value: symbols.append(Symbol(name=name, kind='define', file=relpath, line=node.start_point[0]+1, end_line=node.start_point[0]+1, signature=value)) return symbols def _extract_go(tree, source: bytes, relpath: str) -> list[Symbol]: """Extract symbols from Go AST with signatures and receiver method grouping.""" symbols = [] receiver_methods: dict[str, list[Symbol]] = {} # type_name -> [method Symbols] def func_sig(node, skip_receiver: bool = False) -> str | None: params = None result = None seen_pl = 0 for c in node.children: if c.type == 'parameter_list': seen_pl += 1 if skip_receiver and seen_pl == 1: continue params = _get_text(c, source) elif params is not None and c.type in ( 'type_identifier', 'pointer_type', 'qualified_type', 'slice_type', 'map_type', 'interface_type', 'parameter_list', ): if c.type == 'parameter_list': result = _get_text(c, source) # multi-return else: result = _get_text(c, source) if params: return f"{params} {result}" if result else params return None top = list(tree.root_node.children) def visit(node, siblings=None, idx=None): if node.type == 'function_declaration': name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') if name: sig = func_sig(node) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) elif node.type == 'method_declaration': recv_type = None mname = None for c in node.children: if c.type == 'parameter_list' and recv_type is None: for p in c.children: if p.type == 'parameter_declaration': for s in p.children: if s.type == 'pointer_type': for inner in s.children: if inner.type == 'type_identifier': recv_type = _get_text(inner, source) elif s.type == 'type_identifier': recv_type = _get_text(s, source) elif c.type == 'field_identifier': mname = _get_text(c, source) if recv_type and mname: sig = func_sig(node, skip_receiver=True) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None msym = Symbol(name=mname, kind='method', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc) receiver_methods.setdefault(recv_type, []).append(msym) elif node.type == 'type_declaration': for c in node.children: if c.type == 'type_spec': name = next((_get_text(sc, source) for sc in c.children if sc.type == 'type_identifier'), '') if name: # Determine kind from type body kind = 'type' for sc in c.children: if sc.type == 'struct_type': kind = 'struct' elif sc.type == 'interface_type': kind = 'interface' doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind=kind, file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'const_declaration': doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None for c in node.children: if c.type == 'const_spec': name = next((_get_text(sc, source) for sc in c.children if sc.type == 'identifier'), '') if name: symbols.append(Symbol(name=name, kind='constant', file=relpath, line=c.start_point[0]+1, end_line=c.end_point[0]+1, doc=doc)) doc = None # only first const in group gets doc elif node.type == 'var_declaration': for c in node.children: if c.type == 'var_spec': name = next((_get_text(sc, source) for sc in c.children if sc.type == 'identifier'), '') if name: symbols.append(Symbol(name=name, kind='variable', file=relpath, line=c.start_point[0]+1, end_line=c.end_point[0]+1)) children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i) for i, child in enumerate(top): visit(child, siblings=top, idx=i) # Attach receiver methods to their types for sym in symbols: if sym.kind in ('struct', 'interface', 'type') and sym.name in receiver_methods: sym.children = receiver_methods.pop(sym.name) for type_name, methods in receiver_methods.items(): symbols.append(Symbol(name=type_name, kind='type', file=relpath, line=methods[0].line, end_line=methods[-1].end_line, children=methods)) return symbols def _extract_rust(tree, source: bytes, relpath: str) -> list[Symbol]: """Extract symbols from Rust AST with signatures and impl method grouping.""" symbols = [] impl_methods: dict[str, list[Symbol]] = {} # type_name -> [method Symbols] def func_sig(node) -> str | None: params = None ret = None for c in node.children: if c.type == 'parameters': params = _get_text(c, source) elif params is not None and c.type in ( 'type_identifier', 'generic_type', 'reference_type', 'scoped_type_identifier', 'primitive_type', 'tuple_type', ): ret = _get_text(c, source) if params: return f"{params} -> {ret}" if ret else params return None def is_pub(node) -> bool: return any(c.type == 'visibility_modifier' and 'pub' in _get_text(c, source) for c in node.children) top = list(tree.root_node.children) def visit(node, siblings=None, idx=None): if node.type in ('function_item', 'struct_item', 'enum_item', 'trait_item'): if is_pub(node): name = next((_get_text(c, source) for c in node.children if c.type in ('identifier', 'type_identifier')), '') if name: kind = node.type.replace('_item', '') sig = func_sig(node) if node.type == 'function_item' else None doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind=kind, file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) elif node.type in ('const_item', 'static_item', 'type_item'): if is_pub(node): name = next((_get_text(c, source) for c in node.children if c.type in ('identifier', 'type_identifier')), '') if name: kind_map = {'const_item': 'constant', 'static_item': 'static', 'type_item': 'type'} doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind=kind_map[node.type], file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'mod_item': if is_pub(node): name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') if name: doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='module', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) elif node.type == 'impl_item': impl_type = None is_trait_impl = False # Detect trait impl: `impl Trait for Type` # Children: type_identifier(Trait), 'for', type_identifier(Type), declaration_list saw_for = False for c in node.children: if c.type == 'for': saw_for = True is_trait_impl = True elif c.type == 'type_identifier': if saw_for: impl_type = _get_text(c, source) # type being implemented elif impl_type is None and not saw_for: impl_type = _get_text(c, source) # might be overwritten if 'for' comes later elif c.type == 'generic_type': for sc in c.children: if sc.type == 'type_identifier': if saw_for or impl_type is None: impl_type = _get_text(sc, source) break elif c.type == 'declaration_list' and impl_type: decl_children = list(c.children) for di, dc in enumerate(decl_children): if dc.type == 'function_item': # Trait impl methods don't need pub; inherent impl methods do if is_trait_impl or is_pub(dc): fname = next((_get_text(p, source) for p in dc.children if p.type == 'identifier'), '') if fname: sig = func_sig(dc) doc = _preceding_doc(decl_children, di, source) msym = Symbol(name=fname, kind='method', file=relpath, line=dc.start_point[0]+1, end_line=dc.end_point[0]+1, signature=sig, doc=doc) impl_methods.setdefault(impl_type, []).append(msym) return # don't recurse into impl children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i) for i, child in enumerate(top): visit(child, siblings=top, idx=i) # Attach impl methods to their types for sym in symbols: if sym.kind in ('struct', 'enum', 'trait') and sym.name in impl_methods: sym.children = impl_methods.pop(sym.name) for type_name, methods in impl_methods.items(): symbols.append(Symbol(name=type_name, kind='type', file=relpath, line=methods[0].line, end_line=methods[-1].end_line, children=methods)) return symbols def _extract_typescript(tree, source: bytes, relpath: str) -> list[Symbol]: """Extract symbols from TypeScript/JavaScript AST with signatures and class hierarchy.""" symbols = [] def return_type(node) -> str | None: for c in node.children: if c.type in ('type_annotation', 'type_predicate_annotation'): return _get_text(c, source).lstrip(': ').strip() return None def func_sig(node) -> str | None: params = next((_get_text(c, source) for c in node.children if c.type == 'formal_parameters'), None) if params: rt = return_type(node) return f"{params}: {rt}" if rt else params return None def class_methods(node) -> list[Symbol]: methods = [] for c in node.children: if c.type == 'class_body': body_kids = list(c.children) for i, sc in enumerate(body_kids): if sc.type == 'method_definition': name = next((_get_text(p, source) for p in sc.children if p.type == 'property_identifier'), '') if name: sig = func_sig(sc) doc = _preceding_doc(body_kids, i, source) methods.append(Symbol(name=name, kind='method', file=relpath, line=sc.start_point[0]+1, end_line=sc.end_point[0]+1, signature=sig, doc=doc)) return methods top = list(tree.root_node.children) def visit(node, siblings=None, idx=None): # Functions (named declarations) if node.type in ('function_declaration', 'generator_function_declaration'): name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') if name: sig = func_sig(node) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) # Classes elif node.type in ('class_declaration', 'class'): name = next((_get_text(c, source) for c in node.children if c.type in ('type_identifier', 'identifier')), '') if name: methods = class_methods(node) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='class', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc, children=methods)) # Interfaces (TS) elif node.type == 'interface_declaration': name = next((_get_text(c, source) for c in node.children if c.type == 'type_identifier'), '') if name: doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='interface', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) # Abstract classes (TS) elif node.type == 'abstract_class_declaration': name = next((_get_text(c, source) for c in node.children if c.type == 'type_identifier'), '') if name: methods = class_methods(node) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='class', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc, children=methods)) # const/let = arrow function elif node.type == 'lexical_declaration': doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None for c in node.children: if c.type == 'variable_declarator': name = next((_get_text(p, source) for p in c.children if p.type == 'identifier'), '') value = next((p for p in c.children if p.type in ('arrow_function', 'function_expression')), None) if name and value: sig = func_sig(value) symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) # Export wrapping — recurse into exported declarations elif node.type == 'export_statement': children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i) return # don't double-recurse children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i) for i, child in enumerate(top): visit(child, siblings=top, idx=i) return symbols def _extract_ruby(tree, source: bytes, relpath: str) -> list[Symbol]: """Extract symbols from Ruby AST with signatures and class/module hierarchy.""" symbols = [] def body_children(node) -> list[Symbol]: """Extract methods and nested classes/modules from a class/module body.""" children = [] for c in node.children: if c.type == 'body_statement': kids = list(c.children) for i, sc in enumerate(kids): if sc.type == 'method': name = next((_get_text(p, source) for p in sc.children if p.type == 'identifier'), '') sig = next((_get_text(p, source) for p in sc.children if p.type == 'method_parameters'), None) if name: doc = _preceding_doc(kids, i, source) children.append(Symbol(name=name, kind='method', file=relpath, line=sc.start_point[0]+1, end_line=sc.end_point[0]+1, signature=sig, doc=doc)) elif sc.type == 'singleton_method': name = next((_get_text(p, source) for p in sc.children if p.type == 'identifier'), '') sig = next((_get_text(p, source) for p in sc.children if p.type == 'method_parameters'), None) if name: doc = _preceding_doc(kids, i, source) children.append(Symbol(name=f"self.{name}", kind='method', file=relpath, line=sc.start_point[0]+1, end_line=sc.end_point[0]+1, signature=sig, doc=doc)) elif sc.type in ('class', 'module'): cname = next((_get_text(p, source) for p in sc.children if p.type in ('constant', 'scope_resolution')), '') if cname: doc = _preceding_doc(kids, i, source) nested = body_children(sc) children.append(Symbol(name=cname, kind=sc.type, file=relpath, line=sc.start_point[0]+1, end_line=sc.end_point[0]+1, doc=doc, children=nested)) return children top = list(tree.root_node.children) def visit(node, siblings=None, idx=None, depth=0): if node.type in ('class', 'module'): name = next((_get_text(c, source) for c in node.children if c.type in ('constant', 'scope_resolution')), '') if name: children = body_children(node) doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind=node.type, file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc, children=children)) return elif node.type == 'method' and depth == 0: name = next((_get_text(c, source) for c in node.children if c.type == 'identifier'), '') sig = next((_get_text(c, source) for c in node.children if c.type == 'method_parameters'), None) if name: doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind='function', file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, signature=sig, doc=doc)) children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i, depth=depth+1) for i, child in enumerate(top): visit(child, siblings=top, idx=i) return symbols _HEADING_MARKERS = { 'atx_h1_marker': 1, 'atx_h2_marker': 2, 'atx_h3_marker': 3, 'atx_h4_marker': 4, 'atx_h5_marker': 5, 'atx_h6_marker': 6, } def _extract_markdown(tree, source: bytes, relpath: str) -> list[Symbol]: """Extract heading outline from Markdown AST as hierarchical symbols.""" symbols = [] def extract_section(node) -> Symbol | None: heading = None children = [] for c in node.children: if c.type == 'atx_heading' and heading is None: level = 0 text = '' for sc in c.children: if sc.type in _HEADING_MARKERS: level = _HEADING_MARKERS[sc.type] elif sc.type == 'inline': text = source[sc.start_byte:sc.end_byte].decode('utf-8', errors='replace').strip() if text: heading = Symbol(name=text, kind=f'h{level}', file=relpath, line=c.start_point[0]+1, end_line=node.end_point[0]+1) elif c.type == 'section': child_sym = extract_section(c) if child_sym: children.append(child_sym) if heading: heading.children = children return heading for child in tree.root_node.children: if child.type == 'section': sym = extract_section(child) if sym: symbols.append(sym) return symbols def _extract_generic(tree, source: bytes, relpath: str, lang: str) -> list[Symbol]: """Generic extractor using node type heuristics. Works for many languages.""" symbols = [] # Walk top-level children looking for common patterns def visit(node, siblings=None, idx=None, depth=0): kind = None name = '' # Function/method definitions if node.type in ('function_definition', 'function_declaration', 'function_item', 'method_definition', 'method_declaration', 'function_signature_item'): kind = 'function' for c in node.children: if c.type in ('identifier', 'name', 'field_identifier', 'property_identifier'): name = _get_text(c, source) break # Class/struct/type definitions elif node.type in ('class_definition', 'class_declaration', 'struct_item', 'enum_item', 'trait_item', 'interface_declaration', 'type_declaration', 'type_spec'): kind = node.type.split('_')[0] # 'class', 'struct', 'enum', etc. for c in node.children: if c.type in ('identifier', 'type_identifier', 'name'): name = _get_text(c, source) break if kind and name: doc = _preceding_doc(siblings, idx, source) if siblings and idx is not None else None symbols.append(Symbol(name=name, kind=kind, file=relpath, line=node.start_point[0]+1, end_line=node.end_point[0]+1, doc=doc)) # Recurse (limited depth) if depth < 2: children = list(node.children) for i, child in enumerate(children): visit(child, siblings=children, idx=i, depth=depth+1) top = list(tree.root_node.children) for i, child in enumerate(top): visit(child, siblings=top, idx=i, depth=0) return symbols # ─── tags.scm registry ─────────────────────────────────────────────────── # Community-maintained tree-sitter tag queries. Each returns @name + @definition.{kind} # captures. Some include @doc for doc-comment extraction. Predicates like #strip! are # parsed but treated as no-ops by the Python binding — we handle stripping ourselves. TAGS_SCM: dict[str, str] = { 'rust': ''' (struct_item name: (type_identifier) @name) @definition.class (enum_item name: (type_identifier) @name) @definition.class (union_item name: (type_identifier) @name) @definition.class (type_item name: (type_identifier) @name) @definition.class (declaration_list (function_item name: (identifier) @name) @definition.method) (function_item name: (identifier) @name) @definition.function (trait_item name: (type_identifier) @name) @definition.interface (mod_item name: (identifier) @name) @definition.module (macro_definition name: (identifier) @name) @definition.macro ''', 'go': ''' ( (comment)* @doc . (function_declaration name: (identifier) @name) @definition.function (#strip! @doc "^//\\\\s*") (#set-adjacent! @doc @definition.function) ) ( (comment)* @doc . (method_declaration name: (field_identifier) @name) @definition.method (#strip! @doc "^//\\\\s*") (#set-adjacent! @doc @definition.method) ) (type_spec name: (type_identifier) @name) @definition.type (type_declaration (type_spec name: (type_identifier) @name type: (interface_type))) @definition.interface (type_declaration (type_spec name: (type_identifier) @name type: (struct_type))) @definition.class (var_declaration (var_spec name: (identifier) @name)) @definition.constant (const_declaration (const_spec name: (identifier) @name)) @definition.constant ''', 'javascript': ''' ( (comment)* @doc . (method_definition name: (property_identifier) @name) @definition.method (#not-eq? @name "constructor") (#strip! @doc "^[\\\\s\\\\*/]+|^[\\\\s\\\\*/]$") (#select-adjacent! @doc @definition.method) ) ( (comment)* @doc . [ (class name: (_) @name) (class_declaration name: (_) @name) ] @definition.class (#strip! @doc "^[\\\\s\\\\*/]+|^[\\\\s\\\\*/]$") (#select-adjacent! @doc @definition.class) ) ( (comment)* @doc . [ (function_expression name: (identifier) @name) (function_declaration name: (identifier) @name) (generator_function name: (identifier) @name) (generator_function_declaration name: (identifier) @name) ] @definition.function (#strip! @doc "^[\\\\s\\\\*/]+|^[\\\\s\\\\*/]$") (#select-adjacent! @doc @definition.function) ) ( (comment)* @doc . (lexical_declaration (variable_declarator name: (identifier) @name value: [(arrow_function) (function_expression)]) @definition.function) (#strip! @doc "^[\\\\s\\\\*/]+|^[\\\\s\\\\*/]$") (#select-adjacent! @doc @definition.function) ) ( (comment)* @doc . (variable_declaration (variable_declarator name: (identifier) @name value: [(arrow_function) (function_expression)]) @definition.function) (#strip! @doc "^[\\\\s\\\\*/]+|^[\\\\s\\\\*/]$") (#select-adjacent! @doc @definition.function) ) (assignment_expression left: [ (identifier) @name (member_expression property: (property_identifier) @name) ] right: [(arrow_function) (function_expression)] ) @definition.function (pair key: (property_identifier) @name value: [(arrow_function) (function_expression)]) @definition.function ''', 'typescript': ''' (function_signature name: (identifier) @name) @definition.function (method_signature name: (property_identifier) @name) @definition.method (abstract_method_signature name: (property_identifier) @name) @definition.method (abstract_class_declaration name: (type_identifier) @name) @definition.class (module name: (identifier) @name) @definition.module (interface_declaration name: (type_identifier) @name) @definition.interface ''', 'tsx': ''' (function_signature name: (identifier) @name) @definition.function (method_signature name: (property_identifier) @name) @definition.method (abstract_method_signature name: (property_identifier) @name) @definition.method (abstract_class_declaration name: (type_identifier) @name) @definition.class (module name: (identifier) @name) @definition.module (interface_declaration name: (type_identifier) @name) @definition.interface ''', 'ruby': ''' ( (comment)* @doc . [ (method name: (_) @name) @definition.method (singleton_method name: (_) @name) @definition.method ] (#strip! @doc "^#\\\\s*") (#select-adjacent! @doc @definition.method) ) (alias name: (_) @name) @definition.method ( (comment)* @doc . [ (class name: [(constant) @name (scope_resolution name: (_) @name)]) @definition.class (singleton_class value: [(constant) @name (scope_resolution name: (_) @name)]) @definition.class ] (#strip! @doc "^#\\\\s*") (#select-adjacent! @doc @definition.class) ) (module name: [(constant) @name (scope_resolution name: (_) @name)]) @definition.module ''', 'java': ''' (class_declaration name: (identifier) @name) @definition.class (method_declaration name: (identifier) @name) @definition.method (interface_declaration name: (identifier) @name) @definition.interface ''', 'cpp': ''' (struct_specifier name: (type_identifier) @name body:(_)) @definition.class (declaration type: (union_specifier name: (type_identifier) @name)) @definition.class (function_declarator declarator: (identifier) @name) @definition.function (function_declarator declarator: (field_identifier) @name) @definition.function (function_declarator declarator: (qualified_identifier scope: (namespace_identifier) @scope name: (identifier) @name)) @definition.method (type_definition declarator: (type_identifier) @name) @definition.type (enum_specifier name: (type_identifier) @name) @definition.type (class_specifier name: (type_identifier) @name) @definition.class ''', 'c_sharp': ''' (class_declaration name: (identifier) @name) @definition.class (interface_declaration name: (identifier) @name) @definition.interface (method_declaration name: (identifier) @name) @definition.method (namespace_declaration name: (identifier) @name) @definition.module ''', # tree-sitter-mojo ships its own queries/tags.scm; this mirrors it # exactly (plus a method capture for fn defs nested in a class/struct/trait # body, matching how the JS/TS extractor distinguishes function vs method). 'mojo': ''' (class_definition name: (identifier) @name) @definition.class (struct_definition name: (identifier) @name) @definition.class (trait_definition name: (identifier) @name) @definition.interface (function_definition name: (identifier) @name) @definition.function (block (function_definition name: (identifier) @name) @definition.method) (alias_declaration name: (identifier) @name) @definition.constant (variable_declaration name: (identifier) @name) @definition.variable ''', } # TS/TSX inherit JS patterns for runtime definitions (tags.scm only has TS-specific extras) # Combine: JS base + TS extras for _ts_lang in ('typescript', 'tsx'): TAGS_SCM[_ts_lang] = TAGS_SCM['javascript'] + '\n' + TAGS_SCM[_ts_lang] _query_cache: dict[str, object] = {} # lang -> compiled Query def _get_tags_query(lang: str): """Get or compile cached tags.scm query for a language.""" if lang in _query_cache: return _query_cache[lang] scm = TAGS_SCM.get(lang) if not scm: return None try: from tree_sitter import Query parser = _get_parser(lang) if parser is None: _query_cache[lang] = None return None q = Query(parser.language, scm) _query_cache[lang] = q return q except Exception: _query_cache[lang] = None # Don't retry on error return None def _strip_doc_comment(text: str) -> str: """Strip comment markers from a doc comment node's text.""" text = text.strip() # Try each comment style for prefix in ('/**', '/*', '///', '//', '#'): text = text.removeprefix(prefix) text = text.rstrip('*/').strip() # Return first meaningful line for line in text.split('\n'): line = line.strip().lstrip('*#/').strip() if line and not line.startswith('@') and not line.startswith('\\'): return line return '' def _extract_via_tags(tree, source: bytes, relpath: str, lang: str) -> list[Symbol]: """Extract symbols using tags.scm queries. Returns empty list if no tags.scm available.""" query = _get_tags_query(lang) if query is None: return [] from tree_sitter import QueryCursor cursor = QueryCursor(query) # Collect definitions: (start_line, name) -> Symbol, preferring specific kinds KIND_PRIORITY = {'method': 3, 'interface': 2, 'class': 2, 'module': 2, 'macro': 2, 'type': 1, 'function': 1, 'constant': 0} seen: dict[tuple[int, str], Symbol] = {} # (start_line, name) -> Symbol for _pat_idx, captures in cursor.matches(tree.root_node): # Find the definition capture def_key = None for k in captures: if k.startswith('definition.'): def_key = k break if not def_key: continue name_nodes = captures.get('name', []) def_nodes = captures.get(def_key, []) if not name_nodes or not def_nodes: continue name_text = source[name_nodes[0].start_byte:name_nodes[0].end_byte].decode('utf-8', errors='replace') def_node = def_nodes[0] kind = def_key.split('.')[-1] # function, method, class, etc. start_line = def_node.start_point[0] + 1 end_line = def_node.end_point[0] + 1 # Extract doc from @doc capture if present doc = None doc_nodes = captures.get('doc', []) if doc_nodes: doc_text = source[doc_nodes[0].start_byte:doc_nodes[0].end_byte].decode('utf-8', errors='replace') doc = _strip_doc_comment(doc_text) or None # Dedup: prefer higher-priority kind for same (line, name) key = (start_line, name_text) existing = seen.get(key) if existing: if KIND_PRIORITY.get(kind, 0) > KIND_PRIORITY.get(existing.kind, 0): existing.kind = kind # Upgrade in place if doc and not existing.doc: existing.doc = doc continue sym = Symbol(name=name_text, kind=kind, file=relpath, line=start_line, end_line=end_line, doc=doc) seen[key] = sym return list(seen.values()) # ─── Extractor dispatch ─────────────────────────────────────────────────── EXTRACTORS = { 'python': _extract_python, 'c': _extract_c, 'go': _extract_go, 'rust': _extract_rust, 'javascript': _extract_typescript, # same grammar family 'typescript': _extract_typescript, 'tsx': _extract_typescript, 'ruby': _extract_ruby, 'markdown': _extract_markdown, } def extract_symbols(tree, source: bytes, relpath: str, lang: str) -> list[Symbol]: """Extract symbols from a parsed tree. Priority: custom extractor → tags.scm → generic heuristic. """ # 1. Custom extractor (language-specific, richest output) extractor = EXTRACTORS.get(lang) if extractor: return extractor(tree, source, relpath) # 2. tags.scm query (community-maintained, good coverage) tags_result = _extract_via_tags(tree, source, relpath, lang) if tags_result: return tags_result # 3. Generic heuristic fallback return _extract_generic(tree, source, relpath, lang) def extract_imports(tree, source: bytes, lang: str) -> list[str]: """Extract import/include paths from a file.""" imports = [] for node in tree.root_node.children: if lang == 'python': if node.type == 'import_statement': for c in node.children: if c.type == 'dotted_name': imports.append(_get_text(c, source)) elif node.type == 'import_from_statement': for c in node.children: if c.type in ('dotted_name', 'relative_import'): imports.append(_get_text(c, source)) break elif lang in ('c', 'cpp'): if node.type == 'preproc_include': for c in node.children: if c.type in ('system_lib_string', 'string_literal'): imports.append(_get_text(c, source).strip('"<>')) elif lang in ('javascript', 'typescript', 'tsx'): if node.type == 'import_statement': for c in node.children: if c.type == 'string': imports.append(_get_text(c, source).strip('"\'')) elif lang == 'go': if node.type == 'import_declaration': def find_imports(n): if n.type == 'interpreted_string_literal': imports.append(_get_text(n, source).strip('"')) for c in n.children: find_imports(c) find_imports(node) elif lang == 'rust': if node.type == 'use_declaration': for c in node.children: -
server.py 6.6 KB
""" tree-sitting: MCP server for AST-powered code navigation. Parses codebases with tree-sitter, caches ASTs in memory, and exposes query tools for fast symbol lookup, navigation, and source retrieval. Install: uv pip install fastmcp tree-sitter-language-pack Run: fastmcp run server.py python server.py # stdio python server.py --port 8080 # SSE """ from typing import Annotated from engine import cache from fastmcp import FastMCP from pydantic import Field mcp = FastMCP( name="tree-sitting", instructions=( "AST-powered code navigation. Call `scan` first with a repo path, " "then use find_symbol, file_symbols, dir_overview, tree_overview, " "get_source, and references to explore the codebase. " "All queries run against in-memory parsed ASTs — sub-millisecond." ), ) @mcp.tool(annotations={"title": "Scan Codebase", "destructiveHint": False}) def scan( path: Annotated[str, Field(description="Absolute path to codebase root")], skip: Annotated[str | None, Field( description="Comma-separated dirs to skip (default: .git,node_modules,__pycache__,...)", default=None )] = None, ) -> str: """Parse all source files under path into ASTs. Must be called first. Fast: ~700ms for a 3MB/250-file repo. Results cached for all subsequent queries.""" import time t0 = time.perf_counter() skip_set = set(skip.split(',')) if skip else None stats = cache.scan(path, skip=skip_set) elapsed = (time.perf_counter() - t0) * 1000 stats['elapsed_ms'] = round(elapsed) return ( f"Scanned {stats['files']} files ({stats['bytes']//1024} KB) in {stats['elapsed_ms']}ms\n" f"Symbols: {stats['symbols']} | Languages: {', '.join(stats['languages'])}\n" f"Errors: {stats['errors']}" + (f"\nWARNING: {stats['grammar_hint']}" if stats.get('grammar_hint') else "") ) @mcp.tool(annotations={"title": "Tree Overview", "readOnlyHint": True}) def tree_overview() -> str: """High-level directory tree with file and symbol counts per directory. Call after scan. Good first orientation tool.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." return cache.tree_overview() @mcp.tool(annotations={"title": "Directory Overview", "readOnlyHint": True}) def dir_overview( path: Annotated[str, Field(description="Directory path relative to repo root ('' for root)")] = '', ) -> str: """List files and their top-level symbols for a specific directory. A structural digest of one directory.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." return cache.dir_overview(path) @mcp.tool(annotations={"title": "Find Symbol", "readOnlyHint": True}) def find_symbol( query: Annotated[str, Field(description="Symbol name, substring, or wildcard pattern (e.g. 'ts_parser_*')")], kind: Annotated[str | None, Field( description="Filter by kind: function, class, struct, enum, method, const, define, type", default=None )] = None, limit: Annotated[int, Field(description="Max results", default=20)] = 20, ) -> str: """Search for symbols across the entire codebase by name. Supports exact match, substring, and glob patterns.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." results = cache.find_symbol(query, kind=kind, limit=limit) if not results: return f"No symbols matching '{query}'" lines = [f"Found {len(results)} symbol(s) matching '{query}':\n"] for sym in results: lines.append(f" {sym.file}:{sym.line} — {sym.format_oneline()}") return '\n'.join(lines) @mcp.tool(annotations={"title": "File Symbols", "readOnlyHint": True}) def file_symbols( path: Annotated[str, Field(description="File path relative to repo root (or partial match)")], ) -> str: """List all symbols in a specific file with signatures and doc comments. The structural digest for a single file.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." syms = cache.file_symbols(path) if not syms: return f"No symbols found for '{path}'" imps = cache.file_imports(path) lines = [] if imps: preview = ', '.join(imps[:6]) if len(imps) > 6: preview += f', ... +{len(imps)-6}' lines.append(f"Imports: {preview}\n") for sym in syms: lines.append(sym.format_oneline()) for child in sym.children: lines.append(f" {child.format_oneline()}") return '\n'.join(lines) @mcp.tool(annotations={"title": "Get Source", "readOnlyHint": True}) def get_source( symbol: Annotated[str, Field(description="Symbol name to get source for")], file: Annotated[str | None, Field( description="File path to disambiguate if symbol exists in multiple files", default=None )] = None, ) -> str: """Get the source code of a specific symbol (function, class, struct). Finds the symbol, then returns its source lines.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." results = cache.find_symbol(symbol, limit=5) if file: results = [s for s in results if file in s.file] if not results: return f"Symbol '{symbol}' not found" # Prefer implementation (largest span) over declaration sym = max(results, key=lambda s: s.end_line - s.line) header = f"# {sym.name} ({sym.kind}) in {sym.file}:{sym.line}-{sym.end_line}" if sym.doc: header += f"\n# {sym.doc}" source = cache.get_source_range(sym.file, sym.line, sym.end_line) return f"{header}\n\n{source}" @mcp.tool(annotations={"title": "References", "readOnlyHint": True}) def references( symbol: Annotated[str, Field(description="Symbol name to find references for")], limit: Annotated[int, Field(description="Max results", default=20)] = 20, ) -> str: """Find all textual references to a symbol across the codebase. Fast grep-like search against cached source.""" if not cache.is_loaded: return "No codebase scanned. Call scan() first." refs = cache.references(symbol, limit=limit) if not refs: return f"No references to '{symbol}'" lines = [f"Found {len(refs)} reference(s) to '{symbol}':\n"] for ref in refs: lines.append(f" {ref['file']}:{ref['line']} | {ref['text']}") return '\n'.join(lines) if __name__ == "__main__": import sys if '--port' in sys.argv: idx = sys.argv.index('--port') port = int(sys.argv[idx + 1]) mcp.run(transport='sse', port=port) else: mcp.run() -
treesit.py 14.3 KB
#!/usr/bin/env python3 """ treesit.py — AST-powered code navigation CLI. Auto-scans on every invocation (~700ms), then runs queries. Designed for environments where each call is a separate process. Always prints a tree overview first (progressive disclosure context), then any query results. Usage: treesit.py REPO [OPTIONS] [QUERIES...] Options: --depth N Directory depth: -1=all, 0=root only, 1=one level (default: 1) --detail LEVEL Node detail: sparse|normal|full (default: normal) --path DIR Scope to subdirectory (relative to repo root) --skip DIRS Comma-separated dirs to skip (added to defaults) --no-tree Suppress tree overview, show only query results --stats Show scan statistics Queries: find:PATTERN[:KIND[:LIMIT]] Search symbols by name/glob symbols:FILE_PATH All symbols in a file source:SYMBOL[:FILE] Source code of a symbol refs:SYMBOL[:LIMIT] Find references across codebase imports:FILE_PATH Imports for a file dir:DIR_PATH Directory overview No queries = tree overview only. Detail levels (tree overview rows show per-file symbol lists with line ranges): sparse — name:start-end [featuring: complete shape] normal — name(kind_initial):start-end [exploring: orientation] full — per-symbol formatter + children + imports [exploring: deep dive] """ import argparse import os import sys import time def find_engine(): """Locate tree-sitting engine.""" candidates = [ os.path.join(os.path.dirname(os.path.abspath(__file__)), '.'), '/mnt/skills/user/tree-sitting/scripts', ] for p in candidates: if os.path.exists(os.path.join(p, 'engine.py')): if p not in sys.path: sys.path.insert(0, p) return p return None def format_symbol_sparse(sym, indent=''): """name (kind) :line-end""" return f"{indent}{sym.name} ({sym.kind}) :{sym.line}-{sym.end_line}" def format_symbol_normal(sym, indent=''): """name (kind) signature :line-end — doc""" parts = [f"{indent}{sym.name} ({sym.kind})"] if sym.signature: parts.append(f"{sym.signature}") parts.append(f":{sym.line}-{sym.end_line}") if sym.doc: parts.append(f"— {sym.doc}") return ' '.join(parts) def format_symbol_full(sym, indent=''): """Normal + children.""" lines = [format_symbol_normal(sym, indent)] for child in sym.children: lines.append(format_symbol_normal(child, indent + ' ')) return '\n'.join(lines) FORMATTERS = { 'sparse': format_symbol_sparse, 'normal': format_symbol_normal, 'full': format_symbol_full, } def tree_overview(cache, depth, detail, scope_path=''): """Progressive-disclosure tree overview with depth/detail control.""" if not cache.root: return "No codebase scanned." fmt = FORMATTERS[detail] # Collect directory stats dir_entries = {} # dirpath -> {files: [...], subdirs: set()} for relpath, entry in sorted(cache.files.items()): # Apply scope filter if scope_path: if not relpath.startswith(scope_path.rstrip('/') + '/') and relpath != scope_path: continue # Compute directory relative to scope if scope_path: rest = relpath[len(scope_path.rstrip('/')) + 1:] else: rest = relpath parts = rest.split('/') dirpart = '/'.join(parts[:-1]) if len(parts) > 1 else '' if dirpart not in dir_entries: dir_entries[dirpart] = {'files': [], 'subdirs': set()} dir_entries[dirpart]['files'].append(entry) # Register parent dirs for i in range(1, len(parts) - 1): parent = '/'.join(parts[:i]) child_dir = parts[i] if parent not in dir_entries: dir_entries[parent] = {'files': [], 'subdirs': set()} dir_entries[parent]['subdirs'].add(child_dir) # Root-level subdirs if len(parts) > 1: if '' not in dir_entries: dir_entries[''] = {'files': [], 'subdirs': set()} dir_entries['']['subdirs'].add(parts[0]) if not dir_entries: return f"No files found under '{scope_path}'" if scope_path else "No files scanned." lines = [] display_root = scope_path or (cache.root.name if cache.root else '.') total_files = sum(len(d['files']) for d in dir_entries.values()) total_symbols = sum( len(e.symbols) for d in dir_entries.values() for e in d['files'] ) lines.append(f"# {display_root}/ ({total_files} files, {total_symbols} symbols)\n") def render_dir(dirpath, current_depth): info = dir_entries.get(dirpath, {'files': [], 'subdirs': set()}) dir_indent = ' ' * current_depth # Show files at this level if detail == 'full': for entry in info['files']: fname = os.path.basename(entry.path) lines.append(f"{dir_indent} {fname}") if entry.imports: preview = ', '.join(entry.imports[:6]) if len(entry.imports) > 6: preview += f' +{len(entry.imports) - 6}' lines.append(f"{dir_indent} imports: {preview}") for sym in entry.symbols: lines.append(fmt(sym, dir_indent + ' ')) elif detail == 'normal': for entry in info['files']: fname = os.path.basename(entry.path) sym_summary = ', '.join( f"{s.name}({s.kind[0]}):{s.line}-{s.end_line}" for s in entry.symbols[:6] ) if len(entry.symbols) > 6: sym_summary += f' +{len(entry.symbols) - 6}' lines.append(f"{dir_indent} {fname}: {sym_summary}" if sym_summary else f"{dir_indent} {fname}") else: # sparse for entry in info['files']: fname = os.path.basename(entry.path) sym_names = ', '.join(f"{s.name}:{s.line}-{s.end_line}" for s in entry.symbols[:8]) if len(entry.symbols) > 8: sym_names += f' +{len(entry.symbols) - 8}' lines.append(f"{dir_indent} {fname}: {sym_names}" if sym_names else f"{dir_indent} {fname}") # Show subdirs (respecting depth) if depth == -1 or current_depth < depth: for subdir in sorted(info['subdirs']): child_path = f"{dirpath}/{subdir}" if dirpath else subdir child_info = dir_entries.get(child_path, {'files': [], 'subdirs': set()}) # Count total files recursively under this subdir prefix = child_path + '/' file_count = sum( len(d['files']) for dp, d in dir_entries.items() if dp == child_path or dp.startswith(prefix) ) sym_count = sum( len(e.symbols) for dp, d in dir_entries.items() if dp == child_path or dp.startswith(prefix) for e in d['files'] ) langs = set() for dp, d in dir_entries.items(): if dp == child_path or dp.startswith(prefix): for e in d['files']: langs.add(e.lang) lang_str = ','.join(sorted(langs)) lines.append(f"{dir_indent}{subdir}/ — {file_count} files, {sym_count} symbols [{lang_str}]") render_dir(child_path, current_depth + 1) else: # At depth limit — show collapsed subdirs for subdir in sorted(info['subdirs']): child_path = f"{dirpath}/{subdir}" if dirpath else subdir prefix = child_path + '/' file_count = sum( len(d['files']) for dp, d in dir_entries.items() if dp == child_path or dp.startswith(prefix) ) sym_count = sum( len(e.symbols) for dp, d in dir_entries.items() if dp == child_path or dp.startswith(prefix) for e in d['files'] ) lines.append(f"{dir_indent}{subdir}/ — {file_count} files, {sym_count} symbols [...]") render_dir('', 0) return '\n'.join(lines) def run_query(cache, query_str, detail='normal'): """Parse and execute a query string.""" fmt = FORMATTERS[detail] if ':' in query_str: cmd, _, args = query_str.partition(':') else: return f"Unknown query format: {query_str}\nExpected: find:PATTERN, symbols:FILE, source:SYMBOL, refs:SYMBOL, imports:FILE, dir:PATH" cmd = cmd.strip().lower() args = args.strip() if cmd == 'find': # find:PATTERN[:KIND[:LIMIT]] parts = args.split(':') pattern = parts[0] kind = parts[1] if len(parts) > 1 and parts[1] else None limit = int(parts[2]) if len(parts) > 2 else 20 results = cache.find_symbol(pattern, kind=kind, limit=limit) if not results: return f"No symbols matching '{pattern}'" lines = [f"Found {len(results)} symbol(s) matching '{pattern}':\n"] for sym in results: lines.append(f" {sym.file}:{fmt(sym)}") return '\n'.join(lines) elif cmd == 'symbols': syms = cache.file_symbols(args) if not syms: return f"No symbols found for '{args}'" lines = [f"Symbols in {args}:\n"] for sym in syms: lines.append(fmt(sym, ' ')) if detail == 'full': for child in sym.children: lines.append(fmt(child, ' ')) return '\n'.join(lines) elif cmd == 'source': # source:SYMBOL[:FILE] parts = args.split(':', 1) symbol_name = parts[0] file_filter = parts[1] if len(parts) > 1 else None results = cache.find_symbol(symbol_name, limit=5) if file_filter: results = [s for s in results if file_filter in s.file] if not results: return f"Symbol '{symbol_name}' not found" sym = max(results, key=lambda s: s.end_line - s.line) header = f"# {sym.name} ({sym.kind}) in {sym.file}:{sym.line}-{sym.end_line}" if sym.doc: header += f"\n# {sym.doc}" source = cache.get_source_range(sym.file, sym.line, sym.end_line) return f"{header}\n\n{source}" elif cmd == 'refs': # refs:SYMBOL[:LIMIT] parts = args.split(':', 1) symbol_name = parts[0] limit = int(parts[1]) if len(parts) > 1 else 20 refs = cache.references(symbol_name, limit=limit) if not refs: return f"No references to '{symbol_name}'" lines = [f"Found {len(refs)} reference(s) to '{symbol_name}':\n"] for ref in refs: lines.append(f" {ref['file']}:{ref['line']} | {ref['text']}") return '\n'.join(lines) elif cmd == 'imports': imps = cache.file_imports(args) if not imps: return f"No imports found for '{args}'" return f"Imports in {args}:\n " + '\n '.join(imps) elif cmd == 'dir': return cache.dir_overview(args) else: return f"Unknown command: {cmd}\nAvailable: find, symbols, source, refs, imports, dir" def main(): parser = argparse.ArgumentParser( description='AST-powered code navigation. Auto-scans, then queries.', epilog='Queries: find:PATTERN symbols:FILE source:SYMBOL refs:SYMBOL imports:FILE dir:PATH' ) parser.add_argument('repo', help='Path to codebase root') parser.add_argument('queries', nargs='*', help='Queries to run after scan') parser.add_argument('--depth', type=int, default=1, help='Directory depth: -1=all, 0=root, 1=one level (default: 1)') parser.add_argument('--detail', choices=['sparse', 'normal', 'full'], default='normal', help='Node detail level (default: normal)') parser.add_argument('--path', default='', help='Scope to subdirectory (relative to repo root)') parser.add_argument('--skip', default='', help='Comma-separated dirs to skip (added to defaults)') parser.add_argument('--no-tree', action='store_true', help='Suppress tree overview, show only query results') parser.add_argument('--stats', action='store_true', help='Show scan statistics') parser.add_argument('--no-cache', action='store_true', help='Disable persistent cache (never read or write)') parser.add_argument('--rebuild-cache', action='store_true', help='Ignore existing cache, re-parse, and overwrite cache') args, extra = parser.parse_known_args() # Extra args are queries (argparse struggles with nargs='*' after options) args.queries = list(args.queries) + extra # Find and import engine engine_path = find_engine() if not engine_path: print("ERROR: tree-sitting engine not found.", file=sys.stderr) sys.exit(1) from engine import CodeCache # Scan cache = CodeCache() skip = set(args.skip.split(',')) if args.skip else None t0 = time.perf_counter() use_cache = not args.no_cache rebuild_cache = args.rebuild_cache stats = cache.scan(args.repo, skip=skip, use_cache=use_cache, rebuild_cache=rebuild_cache) elapsed = (time.perf_counter() - t0) * 1000 if stats.get('grammar_hint'): print(f"WARNING: {stats['grammar_hint']}", file=sys.stderr) if args.stats: cached_marker = " (cached)" if stats.get('loaded_from_cache', False) else "" print(f"Scanned {stats['files']} files ({stats['bytes']//1024} KB) in {elapsed:.0f}ms{cached_marker}") print(f"Symbols: {stats['symbols']} | Languages: {', '.join(stats['languages'])}") if stats['errors']: print(f"Errors: {stats['errors']}") print() # Tree overview (unless suppressed) if not args.no_tree: print(tree_overview(cache, args.depth, args.detail, args.path)) # Run queries for q in args.queries: print(f"\n{'─' * 60}") print(run_query(cache, q, args.detail)) if __name__ == '__main__': main()
-
-
tests
-
test_cache_fingerprint.py 11.5 KB
"""Tests for tree-sitting cache fingerprinting and cache path resolution. This test suite covers the persistent scan-cache contract (v1, cheap tier): - _fingerprint stability and invalidation - cache_path_for determinism and env override behavior Run: python -m pytest tests/test_cache_fingerprint.py -v """ import os import sys import time from pathlib import Path # Bootstrap parsers before importing engine sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) from engine import CodeCache # cache_path_for will be imported from engine after implementation # ── Fixtures ──────────────────────────────────────────────────────────── def create_test_repo(tmp_path: Path) -> Path: """Create a minimal test repository with real source files.""" repo = tmp_path / "test_repo" repo.mkdir(parents=True, exist_ok=True) # Create a Python source file py_file = repo / "example.py" py_file.write_bytes(b''' def greet(name: str) -> str: """Greet someone.""" return f"Hello, {name}!" class Service: def start(self) -> None: pass ''') # Create a Rust source file rs_file = repo / "lib.rs" rs_file.write_bytes(b''' pub struct Config { pub name: String, } pub fn create() -> Config { Config { name: String::new() } } impl Config { pub fn new(name: &str) -> Self { Config { name: name.to_string() } } } ''') return repo # ── Fingerprint: Stability ────────────────────────────────────────────── def test_fingerprint_is_stable_unchanged_tree(tmp_path: Path): """Fingerprint is STABLE: two calls on an unchanged tree return the same value.""" repo = create_test_repo(tmp_path) cache = CodeCache() # First fingerprint fp1 = cache._fingerprint(str(repo), skip=None) assert isinstance(fp1, str), "Fingerprint should return a string" assert len(fp1) > 0, "Fingerprint should not be empty" # Second fingerprint on unchanged tree fp2 = cache._fingerprint(str(repo), skip=None) assert fp2 == fp1, "Fingerprint should be stable across rescans of unchanged tree" # ── Fingerprint: Content Changes ──────────────────────────────────────── def test_fingerprint_changes_on_file_content_modify(tmp_path: Path): """Fingerprint CHANGES when a file's content is modified (mtime/size changes).""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint before modification fp_before = cache._fingerprint(str(repo), skip=None) # Modify file content (use os.utime to force distinct mtime_ns) py_file = repo / "example.py" py_file.write_bytes(b''' def greet(name: str) -> str: """Greet someone.""" return f"Hello, {name}!" class Service: def start(self) -> None: pass def stop(self) -> None: pass ''') # Force mtime to be distinct mtime_ns = (time.time_ns() + 1000000000) os.utime(py_file, ns=(mtime_ns, mtime_ns)) # Get fingerprint after modification fp_after = cache._fingerprint(str(repo), skip=None) assert fp_after != fp_before, "Fingerprint should change when file content is modified" def test_fingerprint_changes_on_file_size_change(tmp_path: Path): """Fingerprint CHANGES when file size changes.""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint before size change fp_before = cache._fingerprint(str(repo), skip=None) # Modify file size (truncate to make size noticeably different) py_file = repo / "example.py" original_size = py_file.stat().st_size py_file.write_bytes(b'# Much shorter file\n') assert py_file.stat().st_size != original_size # Get fingerprint after size change fp_after = cache._fingerprint(str(repo), skip=None) assert fp_after != fp_before, "Fingerprint should change when file size changes" # ── Fingerprint: File Operations ──────────────────────────────────────── def test_fingerprint_changes_on_file_added(tmp_path: Path): """Fingerprint CHANGES when a file is added.""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint before adding file fp_before = cache._fingerprint(str(repo), skip=None) # Add a new Python file new_file = repo / "new_module.py" new_file.write_bytes(b'def new_function():\n pass\n') # Get fingerprint after adding file fp_after = cache._fingerprint(str(repo), skip=None) assert fp_after != fp_before, "Fingerprint should change when a file is added" def test_fingerprint_changes_on_file_removed(tmp_path: Path): """Fingerprint CHANGES when a file is removed.""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint before removing file fp_before = cache._fingerprint(str(repo), skip=None) # Remove a file py_file = repo / "example.py" py_file.unlink() # Get fingerprint after removing file fp_after = cache._fingerprint(str(repo), skip=None) assert fp_after != fp_before, "Fingerprint should change when a file is removed" # ── Fingerprint: Skip-set Changes ─────────────────────────────────────── def test_fingerprint_changes_on_skip_set_change(tmp_path: Path): """Fingerprint CHANGES when the skip-set changes.""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint with no skip set fp_no_skip = cache._fingerprint(str(repo), skip=None) # Get fingerprint with a skip set (even if it doesn't exclude anything) fp_with_skip = cache._fingerprint(str(repo), skip={'nonexistent_dir'}) # The fingerprints should differ because skip-set is part of the hash assert fp_with_skip != fp_no_skip, "Fingerprint should change when skip-set changes" def test_fingerprint_changes_with_different_skip_sets(tmp_path: Path): """Fingerprint differs for different skip-sets (included in hash).""" repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint with skip set A fp_skip_a = cache._fingerprint(str(repo), skip={'dir_a'}) # Get fingerprint with skip set B fp_skip_b = cache._fingerprint(str(repo), skip={'dir_b'}) assert fp_skip_a != fp_skip_b, "Fingerprint should differ for different skip-sets" # ── Fingerprint: Cache Format Version ─────────────────────────────────── def test_fingerprint_changes_on_cache_format_version_bump(tmp_path: Path, monkeypatch): """Fingerprint CHANGES when engine.CACHE_FORMAT_VERSION changes.""" import engine repo = create_test_repo(tmp_path) cache = CodeCache() # Get fingerprint with current version original_version = engine.CACHE_FORMAT_VERSION fp_v1 = cache._fingerprint(str(repo), skip=None) # Bump the cache format version monkeypatch.setattr(engine, 'CACHE_FORMAT_VERSION', original_version + 1) # Get fingerprint with new version fp_v2 = cache._fingerprint(str(repo), skip=None) # Restore original version monkeypatch.setattr(engine, 'CACHE_FORMAT_VERSION', original_version) assert fp_v2 != fp_v1, "Fingerprint should change when CACHE_FORMAT_VERSION is bumped" # ── cache_path_for: Determinism ──────────────────────────────────────── def test_cache_path_for_is_deterministic(tmp_path: Path): """cache_path_for(root) is deterministic for the same resolved abspath.""" import engine repo = create_test_repo(tmp_path) repo_path = str(repo) # Call cache_path_for twice with the same path path1 = engine.cache_path_for(repo_path) path2 = engine.cache_path_for(repo_path) assert path1 == path2, "cache_path_for should be deterministic for the same root" def test_cache_path_for_resolves_symlinks(tmp_path: Path): """cache_path_for resolves symlinks to canonical path.""" import engine repo = create_test_repo(tmp_path) # Create a symlink to the repo symlink = tmp_path / "link_to_repo" symlink.symlink_to(repo) # Get cache paths for both the real and symlinked paths path_real = engine.cache_path_for(str(repo)) path_symlink = engine.cache_path_for(str(symlink)) # They should be the same (both resolve to the real path) assert path_real == path_symlink, "cache_path_for should resolve symlinks to the same canonical path" def test_cache_path_for_differs_for_different_roots(tmp_path: Path): """cache_path_for differs for different roots.""" import engine repo1 = create_test_repo(tmp_path / "root_a") repo2 = create_test_repo(tmp_path / "root_b") # a genuinely distinct root path1 = engine.cache_path_for(str(repo1)) path2 = engine.cache_path_for(str(repo2)) assert path1 != path2, "cache_path_for should return different paths for different roots" # ── cache_path_for: Environment Override ──────────────────────────────── def test_cache_path_for_honors_treesit_cache_dir_env(tmp_path: Path, monkeypatch): """cache_path_for honors TREESIT_CACHE_DIR environment variable.""" import engine repo = create_test_repo(tmp_path) cache_dir = tmp_path / "my_cache" cache_dir.mkdir() # Set TREESIT_CACHE_DIR env var monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) # Get cache path cache_path = engine.cache_path_for(str(repo)) # Verify the cache path is under the specified cache dir assert cache_path.parent == cache_dir, ( f"cache_path_for should use TREESIT_CACHE_DIR; " f"got {cache_path.parent}, expected {cache_dir}" ) def test_cache_path_for_uses_temp_dir_when_no_env(tmp_path: Path, monkeypatch): """cache_path_for uses system temp dir when TREESIT_CACHE_DIR is not set.""" import engine repo = create_test_repo(tmp_path) # Ensure TREESIT_CACHE_DIR is not set monkeypatch.delenv('TREESIT_CACHE_DIR', raising=False) # Get cache path cache_path = engine.cache_path_for(str(repo)) # Verify it's a Path assert isinstance(cache_path, Path), "cache_path_for should return a Path object" # Verify parent directory exists (temp dir) assert cache_path.parent.exists(), "cache_path_for should use an existing directory (temp)" # Verify parent is writable assert os.access(cache_path.parent, os.W_OK), "cache path parent should be writable" def test_cache_path_for_filename_derived_from_abspath(tmp_path: Path): """cache_path_for filename is deterministically derived from resolved abspath.""" import engine repo1 = create_test_repo(tmp_path / "subdir1") repo2 = create_test_repo(tmp_path / "subdir2") path1 = engine.cache_path_for(str(repo1)) path2 = engine.cache_path_for(str(repo2)) # Filenames should be different (based on different paths) assert path1.name != path2.name, ( "cache filenames should differ for different root paths" ) # Both filenames should look like SHA256 hashes or similar deterministic names # (exact format depends on implementation, but should be hex/base64-ish) assert len(path1.name) > 8, "cache filename should be long enough to be unique" assert len(path2.name) > 8, "cache filename should be long enough to be unique" -
test_cache_lazy_source.py 10.3 KB
"""Tests for lazy source correctness on cache hits — T4. This module tests that FileEntry objects loaded from the persistent cache have source=None and tree=None, and that code paths requiring source bytes (get_source_range, references) lazily read from disk and return identical results to a fresh parse. Run: python -m pytest tests/test_cache_lazy_source.py -v """ import json import os import shutil import sys import tempfile from pathlib import Path # Bootstrap parsers before importing engine sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) from engine import CodeCache, FileEntry # ── Helpers ────────────────────────────────────────────────────────────── def _cache_dir_for_root(root: str) -> Path: """Return the cache file path for a root, using TREESIT_CACHE_DIR if set.""" # This mirrors the engine.cache_path_for() logic that will be implemented cache_dir = os.environ.get('TREESIT_CACHE_DIR') if cache_dir: return Path(cache_dir) # Fallback: use system temp return Path(tempfile.gettempdir()) / 'treesit-cache' def _read_cache_json(root: str) -> dict: """Read raw cache JSON from disk (for inspection).""" import hashlib root_resolved = str(Path(root).resolve()) cache_hash = hashlib.sha256(root_resolved.encode()).hexdigest() cache_dir = _cache_dir_for_root(root) cache_file = cache_dir / f'{cache_hash}.json' if not cache_file.exists(): return {} return json.loads(cache_file.read_text()) # ── T4.1: Cache hit, get_source_range returns same as fresh parse ─────── def test_cache_hit_get_source_range_matches_fresh(): """After cache-hit scan, get_source_range returns same text as fresh parse.""" tmpdir = tempfile.mkdtemp() cache_dir = tempfile.mkdtemp() try: os.environ['TREESIT_CACHE_DIR'] = cache_dir # Create a test file with recognizable content test_file = Path(tmpdir) / 'module.py' test_file.write_text( 'def alpha():\n' ' """Line 2 docstring."""\n' ' return 42\n' 'def beta():\n' ' pass\n' ) # Scan fresh (populates cache) cache_fresh = CodeCache() cache_fresh.scan(tmpdir) src_fresh = cache_fresh.get_source_range('module.py', 2, 4) # Clear in-memory cache, scan from disk cache (cache hit) cache_hit = CodeCache() hit_stats = cache_hit.scan(tmpdir) # must be a cache hit assert hit_stats.get('loaded_from_cache') is True, \ "second scan must load from cache, else this test is vacuous" src_hit = cache_hit.get_source_range('module.py', 2, 4) # Results must be identical assert src_fresh == src_hit, \ f"Fresh: {src_fresh!r} != Hit: {src_hit!r}" assert 'Line 2 docstring' in src_hit finally: if 'TREESIT_CACHE_DIR' in os.environ: del os.environ['TREESIT_CACHE_DIR'] shutil.rmtree(tmpdir) shutil.rmtree(cache_dir) # ── T4.2: Cache hit, references returns same as fresh parse ──────────── def test_cache_hit_references_matches_fresh(): """After cache-hit scan, references() returns same results as fresh parse.""" tmpdir = tempfile.mkdtemp() cache_dir = tempfile.mkdtemp() try: os.environ['TREESIT_CACHE_DIR'] = cache_dir # Create test files with a symbol used in multiple places Path(tmpdir, 'def.py').write_text( 'class UserConfig:\n' ' """Holds user settings."""\n' ' def __init__(self):\n' ' pass\n' ) Path(tmpdir, 'use.py').write_text( 'from def import UserConfig\n' 'cfg = UserConfig()\n' 'settings = UserConfig()\n' ) # Scan fresh cache_fresh = CodeCache() cache_fresh.scan(tmpdir) refs_fresh = cache_fresh.references('UserConfig', limit=10) # Scan from cache hit cache_hit = CodeCache() hit_stats = cache_hit.scan(tmpdir) assert hit_stats.get('loaded_from_cache') is True, \ "second scan must load from cache, else this test is vacuous" refs_hit = cache_hit.references('UserConfig', limit=10) # Results must be identical: same files, same line numbers, same text assert len(refs_fresh) == len(refs_hit), \ f"Fresh found {len(refs_fresh)} refs, hit found {len(refs_hit)}" for ref_f, ref_h in zip(refs_fresh, refs_hit): assert ref_f == ref_h, \ f"Fresh: {ref_f} != Hit: {ref_h}" # Verify we found references in both files files = [r['file'] for r in refs_hit] assert any('def.py' in f for f in files) assert any('use.py' in f for f in files) finally: if 'TREESIT_CACHE_DIR' in os.environ: del os.environ['TREESIT_CACHE_DIR'] shutil.rmtree(tmpdir) shutil.rmtree(cache_dir) # ── T4.3: Cache file does NOT contain raw source text ─────────────────── def test_cache_file_omits_raw_source(): """Cache JSON does not contain the raw source bytes of fixture files.""" tmpdir = tempfile.mkdtemp() cache_dir = tempfile.mkdtemp() try: os.environ['TREESIT_CACHE_DIR'] = cache_dir # Create test file with a distinctive marker comment marker = '### DISTINCTIVE_CACHE_TEST_MARKER_12345 ###' test_file = Path(tmpdir) / 'sample.py' test_file.write_text( f'# {marker}\n' 'def function():\n' ' return "value"\n' ) # Scan to write cache cache = CodeCache() cache.scan(tmpdir) # Read raw cache JSON from disk — it must actually exist, else the # marker-absence check below is vacuously true. cache_json = _read_cache_json(tmpdir) assert cache_json, "cache file must exist and be non-empty" cache_text = json.dumps(cache_json) # Verify the distinctive marker is NOT in the cache JSON assert marker not in cache_text, \ f"Marker found in cache JSON: {cache_text[:500]}" assert 'function' not in cache_text or 'function' in str(cache_json.get('files', [])), \ "Source code should not be in cache (only symbol names)" finally: if 'TREESIT_CACHE_DIR' in os.environ: del os.environ['TREESIT_CACHE_DIR'] shutil.rmtree(tmpdir) shutil.rmtree(cache_dir) # ── T4.4: Graceful degradation when source file deleted ──────────────── def test_cache_hit_missing_file_degrades_gracefully(): """After cache written, delete source file, cache-hit scan with get_source_range and references must degrade gracefully (no crash, empty/skipped result).""" tmpdir = tempfile.mkdtemp() cache_dir = tempfile.mkdtemp() try: os.environ['TREESIT_CACHE_DIR'] = cache_dir # Create and scan test files Path(tmpdir, 'exists.py').write_text('def keep_me(): pass\n') Path(tmpdir, 'delete_me.py').write_text( 'def symbol_in_deleted():\n' ' return 42\n' ) cache = CodeCache() cache.scan(tmpdir) # writes cache initial_refs = cache.references('symbol_in_deleted', limit=5) assert len(initial_refs) > 0, "Should find symbol before deletion" # Load from cache (entry.source lazily None), THEN delete the file so the # lazy source read must cope with a now-missing file on a cached entry. cache2 = CodeCache() hit_stats = cache2.scan(tmpdir) assert hit_stats.get('loaded_from_cache') is True, \ "must be a cache hit so the lazy-read path is exercised" (Path(tmpdir) / 'delete_me.py').unlink() # get_source_range on deleted file should not crash result = cache2.get_source_range('delete_me.py', 1, 2) # Should either be empty, or return a "not found" message, not crash assert isinstance(result, str) # references should not crash even if file is missing refs = cache2.references('symbol_in_deleted', limit=5) # Should either be empty list or skip the missing file gracefully assert isinstance(refs, list) finally: if 'TREESIT_CACHE_DIR' in os.environ: del os.environ['TREESIT_CACHE_DIR'] shutil.rmtree(tmpdir) shutil.rmtree(cache_dir) # ── T4.5: FileEntry from cache has source=None, tree=None immediately ─── def test_file_entry_lazy_fields_null_after_cache_hit(): """FileEntry objects loaded from cache have source=None and tree=None immediately after scan (before any lazy read).""" tmpdir = tempfile.mkdtemp() cache_dir = tempfile.mkdtemp() try: os.environ['TREESIT_CACHE_DIR'] = cache_dir # Create and scan test file Path(tmpdir, 'lazy.py').write_text( 'def lazy_function():\n' ' pass\n' ) # First scan: writes cache cache1 = CodeCache() cache1.scan(tmpdir) entry1 = cache1.files.get('lazy.py') assert entry1 is not None # First scan may have source loaded (depends on implementation) # Just verify the structure exists assert isinstance(entry1, FileEntry) # Second scan: should hit cache cache2 = CodeCache() cache2.scan(tmpdir) # Check FileEntry from cache hit entry2 = cache2.files.get('lazy.py') assert entry2 is not None, "FileEntry should be populated from cache" # CRITICAL: source and tree must be None on cache hit # to force lazy loading behavior assert entry2.source is None, \ "FileEntry.source must be None on cache hit (lazy loading required)" assert entry2.tree is None, \ "FileEntry.tree must be None on cache hit (lazy loading required)" # But symbols should be populated from cache assert entry2.symbols is not None assert len(entry2.symbols) > 0, "Symbols should be loaded from cache" finally: if 'TREESIT_CACHE_DIR' in os.environ: del os.environ['TREESIT_CACHE_DIR'] shutil.rmtree(tmpdir) shutil.rmtree(cache_dir) -
test_cache_lifecycle.py 30.3 KB
"""Tests for tree-sitting cache lifecycle (T3: cache lifecycle, flags, atomicity, robustness). Run: python -m pytest tests/test_cache_lifecycle.py -v Or: python tests/test_cache_lifecycle.py These tests are for NOT-YET-IMPLEMENTED cache features. TDD red phase. All tests WILL FAIL until the cache implementation is complete. """ import json import os import shutil import sys import tempfile import time from pathlib import Path from unittest.mock import patch # Bootstrap parsers before importing engine sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) import engine from engine import CodeCache # Not-yet-implemented names: resolve lazily so the module COLLECTS (red-fails per # test) instead of erroring at import during TDD red phase. cache_path_for = getattr(engine, "cache_path_for", None) CACHE_FORMAT_VERSION = getattr(engine, "CACHE_FORMAT_VERSION", None) # ── Fixtures and helpers ───────────────────────────────────────────────────── def create_test_tree(tmpdir, files_dict: dict[str, str]) -> Path: """Create a test directory tree with given files. Args: tmpdir: Path to create files in files_dict: dict of relpath -> content (as string) Returns: Path to tmpdir """ root = Path(tmpdir) for relpath, content in files_dict.items(): filepath = root / relpath filepath.parent.mkdir(parents=True, exist_ok=True) filepath.write_text(content) return root def read_cache_json(cache_path: Path) -> dict: """Read cache JSON file safely.""" if not cache_path.exists(): raise FileNotFoundError(f"Cache file not found: {cache_path}") with open(cache_path, 'r') as f: return json.load(f) def find_temp_files_in_dir(dirpath: Path, pattern: str = '.tmp') -> list[Path]: """Find temporary files matching pattern in directory.""" return list(dirpath.glob(f'*{pattern}*')) # ── T3.1: First scan(use_cache=True) creates cache, loaded_from_cache=False ── def test_first_scan_creates_cache(tmp_path, monkeypatch): """First scan(use_cache=True) CREATES the cache file at cache_path_for(root). Assertions: - Cache file exists after scan - loaded_from_cache is False (no prior cache) - Cache contains JSON with CACHE_FORMAT_VERSION header """ cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) # Create test tree root = create_test_tree(tmp_path / 'project', { 'main.py': 'def hello(): pass\n', 'lib.py': 'class Foo:\n def bar(self): pass\n', }) # First scan cache = CodeCache() stats = cache.scan(str(root), use_cache=True) # Assert cache was created expected_cache_path = cache_path_for(str(root)) assert expected_cache_path.exists(), f"Cache file not created at {expected_cache_path}" # Assert loaded_from_cache is False assert 'loaded_from_cache' in stats, "loaded_from_cache key missing from stats" assert stats['loaded_from_cache'] is False, "First scan should have loaded_from_cache=False" # Assert cache contains valid JSON with version header cache_data = read_cache_json(expected_cache_path) assert 'cache_format_version' in cache_data, "Cache missing cache_format_version" assert cache_data['cache_format_version'] == CACHE_FORMAT_VERSION def test_first_scan_cache_has_fingerprint(tmp_path, monkeypatch): """Cache file includes fingerprint of scanned tree.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) cache = CodeCache() cache.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) cache_data = read_cache_json(cache_path) assert 'fingerprint' in cache_data, "Cache missing fingerprint field" assert isinstance(cache_data['fingerprint'], str), "Fingerprint should be a string" assert len(cache_data['fingerprint']) > 0, "Fingerprint should not be empty" def test_first_scan_cache_stores_symbols(tmp_path, monkeypatch): """Cache file stores symbols with required fields from Symbol.to_dict().""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def process(x): """Process x."""\n pass\n', }) cache = CodeCache() cache.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) cache_data = read_cache_json(cache_path) assert 'files' in cache_data, "Cache missing files field" assert isinstance(cache_data['files'], dict), "files should be a dict" # Check that at least one file is stored with required fields for relpath, file_data in cache_data['files'].items(): assert 'lang' in file_data, f"File {relpath} missing lang" assert 'symbols' in file_data, f"File {relpath} missing symbols" # Symbols should have fields from Symbol.to_dict() for sym in file_data['symbols']: assert 'name' in sym, "Symbol missing name" assert 'kind' in sym, "Symbol missing kind" assert 'file' in sym, "Symbol missing file" assert 'line' in sym, "Symbol missing line" assert 'end_line' in sym, "Symbol missing end_line" # ── T3.2: Second scan(use_cache=True) loads cache, loaded_from_cache=True ── def test_second_scan_loads_cache(tmp_path, monkeypatch): """Second scan(use_cache=True) on unchanged tree loads it: loaded_from_cache=True.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'main.py': 'def hello(): pass\n', }) # First scan cache1 = CodeCache() stats1 = cache1.scan(str(root), use_cache=True) assert stats1['loaded_from_cache'] is False files_count_1 = stats1['files'] symbols_count_1 = stats1['symbols'] # Second scan (should load from cache) cache2 = CodeCache() stats2 = cache2.scan(str(root), use_cache=True) # Assert cache was loaded assert 'loaded_from_cache' in stats2, "loaded_from_cache key missing" assert stats2['loaded_from_cache'] is True, "Second scan should have loaded_from_cache=True" # Assert stats match (same content) assert stats2['files'] == files_count_1, "File count should match" assert stats2['symbols'] == symbols_count_1, "Symbol count should match" def test_cache_load_reconstructs_symbol_index(tmp_path, monkeypatch): """Loaded cache reconstructs symbol index correctly.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'lib.py': 'def helper(): pass\nclass Service:\n def process(self): pass\n', }) # First scan cache1 = CodeCache() cache1.scan(str(root), use_cache=True) results_1 = cache1.find_symbol('helper') assert len(results_1) > 0, "First scan should find 'helper'" # Second scan (loads cache) cache2 = CodeCache() cache2.scan(str(root), use_cache=True) results_2 = cache2.find_symbol('helper') # Assert symbol index is rebuilt assert len(results_2) > 0, "Loaded cache should find 'helper' via symbol index" assert results_2[0].name == 'helper', "Symbol name should match" def test_cache_load_no_parse(tmp_path, monkeypatch): """When cache is loaded, no parsing happens (tree=None in entries).""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) # First scan cache1 = CodeCache() cache1.scan(str(root), use_cache=True) # Second scan (loads from cache) cache2 = CodeCache() cache2.scan(str(root), use_cache=True) # Check that loaded entries have tree=None and source=None for entry in cache2.files.values(): assert entry.tree is None, "Loaded cache entry should have tree=None" assert entry.source is None, "Loaded cache entry should have source=None" # ── T3.3: use_cache=False neither reads nor writes cache ── def test_use_cache_false_skips_read(tmp_path, monkeypatch): """use_cache=False: existing cache is NOT consulted, always parses.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) # Create cache first with use_cache=True cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) assert cache_path.exists(), "Cache should be created" # Now scan with use_cache=False cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=False) # Assert cache was not used assert stats['loaded_from_cache'] is False, "use_cache=False should never load from cache" # Assert entries have source and tree (parsed fresh) for entry in cache2.files.values(): assert entry.source is not None, "use_cache=False should parse and populate source" assert entry.tree is not None, "use_cache=False should parse and populate tree" def test_use_cache_false_does_not_write_cache(tmp_path, monkeypatch): """use_cache=False: cache is NOT written even if it doesn't exist.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def bar(): pass\n', }) # Ensure no cache exists cache_path = cache_path_for(str(root)) assert not cache_path.exists(), "Cache should not exist initially" # Scan with use_cache=False cache = CodeCache() cache.scan(str(root), use_cache=False) # Assert cache was not created assert not cache_path.exists(), "use_cache=False should not create cache" def test_use_cache_false_delete_existing_cache(tmp_path, monkeypatch): """Verify that use_cache=False doesn't consult pre-existing cache.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def original(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Delete the source file (root / 'code.py').unlink() # Create new tree with different content create_test_tree(root, { 'code.py': 'def updated(): pass\n', }) # Scan with use_cache=False should parse the updated file cache2 = CodeCache() cache2.scan(str(root), use_cache=False) # Should find 'updated', not 'original' results = cache2.find_symbol('updated') assert len(results) > 0, "Should find updated function" # ── T3.4: rebuild_cache=True overwrites cache, loaded_from_cache=False ── def test_rebuild_cache_overwrites_existing(tmp_path, monkeypatch): """rebuild_cache=True: parse fresh and overwrite cache.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def original(): pass\n', }) # First scan to create cache cache1 = CodeCache() stats1 = cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) cache_mtime_1 = cache_path.stat().st_mtime # Wait a bit to ensure mtime will differ time.sleep(0.01) # Scan with rebuild_cache=True cache2 = CodeCache() stats2 = cache2.scan(str(root), rebuild_cache=True) # Assert full parse happened assert stats2['loaded_from_cache'] is False, "rebuild_cache=True should have loaded_from_cache=False" # Assert cache was overwritten (mtime advanced) cache_mtime_2 = cache_path.stat().st_mtime assert cache_mtime_2 >= cache_mtime_1, "Cache file should be updated/rewritten" # Assert entries have source and tree (parsed) for entry in cache2.files.values(): assert entry.source is not None, "rebuild_cache=True should parse fresh" assert entry.tree is not None, "rebuild_cache=True should parse fresh" def test_rebuild_cache_true_ignores_existing_cache(tmp_path, monkeypatch): """rebuild_cache=True ignores any existing valid cache.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) # Scan again with rebuild_cache=True (should not load, should parse fresh) cache2 = CodeCache() stats = cache2.scan(str(root), rebuild_cache=True) assert stats['loaded_from_cache'] is False, "rebuild_cache=True should parse fresh, not load" # Entries should have tree and source for entry in cache2.files.values(): assert entry.tree is not None assert entry.source is not None # ── T3.5: Robust — corrupt cache falls back to parse, no crash ── def test_corrupt_cache_json_fallback(tmp_path, monkeypatch): """Corrupt JSON in cache -> falls back to full parse, loaded_from_cache=False.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def process(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Corrupt the cache file cache_path.write_text('{ INVALID JSON HERE }') # Scan should NOT crash and should fall back to parse cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) # Assert full parse happened (no crash, correct fallback) assert stats['loaded_from_cache'] is False, "Corrupt cache should fall back to parse" assert stats['files'] > 0, "Should still parse the tree" assert stats['errors'] == 0, "Should complete successfully" def test_corrupt_cache_unreadable_fallback(tmp_path, monkeypatch): """Unreadable cache file -> falls back to parse, produces correct results.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def helper(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Write garbage bytes cache_path.write_bytes(b'\x00\x01\x02\x03 garbage') # Scan should not crash cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) # Verify fallback succeeded assert stats['loaded_from_cache'] is False assert stats['files'] == 1 results = cache2.find_symbol('helper') assert len(results) > 0, "Should find symbol despite cache corruption" def test_corrupt_cache_missing_fields_fallback(tmp_path, monkeypatch): """Cache with missing required fields -> falls back to parse.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def func(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Write incomplete cache (missing 'files' field) invalid_cache = { 'cache_format_version': CACHE_FORMAT_VERSION, 'fingerprint': 'some_fp', # Missing 'files' field } cache_path.write_text(json.dumps(invalid_cache)) # Should not crash and should parse fresh cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) assert stats['loaded_from_cache'] is False assert stats['files'] > 0 # ── T3.6: Robustness — version mismatch cache -> full parse, no crash ── def test_version_mismatch_fallback(tmp_path, monkeypatch): """Cache with stale CACHE_FORMAT_VERSION -> full parse, no crash.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def legacy(): pass\n', }) # Create cache with current version cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Manually change version in cache file cache_data = read_cache_json(cache_path) cache_data['cache_format_version'] = CACHE_FORMAT_VERSION - 1 # Stale version cache_path.write_text(json.dumps(cache_data)) # Scan with use_cache=True should detect version mismatch and parse fresh cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) # Assert full parse happened assert stats['loaded_from_cache'] is False, "Version mismatch should fall back to parse" assert stats['files'] > 0, "Should successfully parse" # Verify correct results results = cache2.find_symbol('legacy') assert len(results) > 0, "Should find symbols despite version mismatch" def test_version_mismatch_updates_cache(tmp_path, monkeypatch): """Version mismatch cache is rewritten with new version.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def func(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Change version to simulate staleness cache_data = read_cache_json(cache_path) old_version = CACHE_FORMAT_VERSION - 1 cache_data['cache_format_version'] = old_version cache_path.write_text(json.dumps(cache_data)) # Scan should rewrite with new version cache2 = CodeCache() cache2.scan(str(root), use_cache=True) # Check that cache was updated updated_cache_data = read_cache_json(cache_path) assert updated_cache_data['cache_format_version'] == CACHE_FORMAT_VERSION # ── T3.7: Robustness — fingerprint mismatch -> full parse, cache refreshed ── def test_fingerprint_mismatch_fallback(tmp_path, monkeypatch): """Valid cache exists, file modified -> fingerprint mismatch -> full parse.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def original(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) cache_data_1 = read_cache_json(cache_path) orig_fingerprint = cache_data_1['fingerprint'] # Modify source file (root / 'code.py').write_text('def modified(): pass\n') # Scan with use_cache=True cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) # Assert full parse happened (fingerprint changed) assert stats['loaded_from_cache'] is False, "Fingerprint mismatch should trigger full parse" # Assert new fingerprint in cache cache_data_2 = read_cache_json(cache_path) assert cache_data_2['fingerprint'] != orig_fingerprint, "Cache should have new fingerprint" # Verify new content is in cache results = cache2.find_symbol('modified') assert len(results) > 0, "Cache should contain updated symbol" def test_fingerprint_mismatch_detects_file_size_change(tmp_path, monkeypatch): """Fingerprint includes file size; size change -> cache miss.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def func(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) # Modify file (add content, changing size) (root / 'code.py').write_text('def func(): pass\n\n# comment\n') # Scan should detect size change cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) assert stats['loaded_from_cache'] is False, "File size change should invalidate cache" def test_fingerprint_mismatch_detects_mtime_change(tmp_path, monkeypatch): """Fingerprint includes mtime_ns; mtime change -> cache miss.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) # Touch file to change mtime time.sleep(0.01) # Ensure mtime differs (root / 'code.py').touch() # Scan should detect mtime change cache2 = CodeCache() stats = cache2.scan(str(root), use_cache=True) assert stats['loaded_from_cache'] is False, "mtime change should invalidate cache" # ── T3.8: Atomic write via temp file + os.replace ── def test_atomic_write_uses_temp_file(tmp_path, monkeypatch): """Cache write uses temp file + os.replace, no leftover temp files.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) # Scan to create cache cache = CodeCache() cache.scan(str(root), use_cache=True) # Assert no temp files left in cache dir temp_files = find_temp_files_in_dir(cache_dir, '.tmp') assert len(temp_files) == 0, f"Should not have leftover temp files, found: {temp_files}" def test_atomic_write_replaces_atomically(tmp_path, monkeypatch): """Verify cache write uses os.replace for atomic replacement.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def bar(): pass\n', }) # Create cache cache1 = CodeCache() cache1.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Patch os.replace to verify it's called original_replace = os.replace replace_called = [] def tracked_replace(src, dst): replace_called.append((src, dst)) return original_replace(src, dst) with patch('os.replace', side_effect=tracked_replace): # Create new cache instance and scan (should write) cache2 = CodeCache() cache2.scan(str(root), rebuild_cache=True) # Assert os.replace was called with appropriate args assert len(replace_called) > 0, "os.replace should be called for atomic write" # Verify destination is the cache file _, dest_file = replace_called[0] assert Path(dest_file) == cache_path, "Replace destination should be cache_path" def test_atomic_write_no_half_written_cache(tmp_path, monkeypatch): """On successful scan, cache file is complete (not half-written).""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def func(): pass\n', 'lib.py': 'class Service: pass\n', }) # Scan cache = CodeCache() cache.scan(str(root), use_cache=True) cache_path = cache_path_for(str(root)) # Verify cache file is valid JSON (not truncated/half-written) try: cache_data = read_cache_json(cache_path) assert 'cache_format_version' in cache_data assert 'fingerprint' in cache_data assert 'files' in cache_data except json.JSONDecodeError: raise AssertionError("Cache file is not valid JSON (half-written?)") # ── Backward compatibility ────────────────────────────────────────────────── def test_scan_backward_compat_no_params(tmp_path, monkeypatch): """scan(root) with no extra params maintains backward compatibility.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def hello(): pass\n', }) # Call scan with no use_cache or rebuild_cache params cache = CodeCache() stats = cache.scan(str(root)) # Should work without errors, defaults to use_cache=True assert 'files' in stats assert 'symbols' in stats assert 'loaded_from_cache' in stats def test_scan_backward_compat_skip_param(tmp_path, monkeypatch): """scan(root, skip=...) maintains backward compatibility.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', 'node_modules/pkg/lib.js': 'function bar() {}', }) # Call scan with skip parameter cache = CodeCache() skip_set = {'node_modules'} stats = cache.scan(str(root), skip=skip_set) # Should work and respect skip assert 'files' in stats assert stats['files'] == 1 # Only code.py, node_modules skipped # ── Cache path generation ────────────────────────────────────────────────── def test_cache_path_for_deterministic(tmp_path): """cache_path_for() is deterministic: same root -> same path.""" root = str((tmp_path / 'project').resolve()) path1 = cache_path_for(root) path2 = cache_path_for(root) assert path1 == path2, "cache_path_for should be deterministic" def test_cache_path_for_honors_env(tmp_path, monkeypatch): """cache_path_for() honors TREESIT_CACHE_DIR env var.""" cache_dir = tmp_path / 'custom_cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = str((tmp_path / 'project').resolve()) cache_path = cache_path_for(root) # Cache should be under custom dir assert str(cache_path).startswith(str(cache_dir)), \ f"Cache path {cache_path} should be under {cache_dir}" def test_cache_path_for_derives_from_root_hash(tmp_path, monkeypatch): """cache_path_for() filename derived from sha256 of root.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root1 = str((tmp_path / 'project1').resolve()) root2 = str((tmp_path / 'project2').resolve()) path1 = cache_path_for(root1) path2 = cache_path_for(root2) # Paths should be different (derived from different roots) assert path1 != path2, "Different roots should have different cache paths" assert path1.parent == path2.parent, "Should be in same cache dir" # ── Fingerprint generation ────────────────────────────────────────────────── def test_fingerprint_changes_on_file_add(tmp_path, monkeypatch): """CodeCache._fingerprint() changes when file is added.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) cache = CodeCache() # First fingerprint fp1 = cache._fingerprint(str(root), None) # Add a file (root / 'new.py').write_text('def bar(): pass\n') # Second fingerprint should differ fp2 = cache._fingerprint(str(root), None) assert fp1 != fp2, "Fingerprint should change when file is added" def test_fingerprint_changes_on_skip_change(tmp_path, monkeypatch): """_fingerprint() changes when skip set changes.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', 'test.py': 'def test_foo(): pass\n', }) cache = CodeCache() # Fingerprint with no skip fp1 = cache._fingerprint(str(root), None) # Fingerprint with skip={'test.py'} fp2 = cache._fingerprint(str(root), {'test.py'}) assert fp1 != fp2, "Fingerprint should change when skip set changes" def test_fingerprint_stable_on_unchanged_tree(tmp_path, monkeypatch): """_fingerprint() is stable across re-scans of unchanged tree.""" cache_dir = tmp_path / 'cache' cache_dir.mkdir() monkeypatch.setenv('TREESIT_CACHE_DIR', str(cache_dir)) root = create_test_tree(tmp_path / 'project', { 'code.py': 'def foo(): pass\n', }) cache = CodeCache() # Get fingerprints at different times fp1 = cache._fingerprint(str(root), None) time.sleep(0.01) fp2 = cache._fingerprint(str(root), None) assert fp1 == fp2, "Fingerprint should be stable for unchanged tree" # ── Standalone runner ─────────────────────────────────────────────────────── if __name__ == '__main__': import traceback # Collect tests tests = [v for k, v in sorted(globals().items()) if k.startswith('test_') and callable(v)] # Run with fixtures support (simplified: just tmp_path via tempfile) passed = failed = 0 for test in tests: # Create tmp_path for each test tmp_path = Path(tempfile.mkdtemp()) try: # Create monkeypatch mock class Monkeypatch: def setenv(self, key, val): os.environ[key] = val monkeypatch = Monkeypatch() # Run test test(tmp_path, monkeypatch) passed += 1 print(f" ✓ {test.__name__}") except Exception as e: failed += 1 print(f" ✗ {test.__name__}: {e}") traceback.print_exc() finally: shutil.rmtree(tmp_path, ignore_errors=True) print(f"\n{passed} passed, {failed} failed") sys.exit(1 if failed else 0) -
test_cache_roundtrip.py 18.9 KB
"""Tests for cache roundtrip equivalence: fresh parse vs cache hit. This tests the core cache invariant: query results must be byte-identical whether served from a fresh parse or a cache hit. Run: python -m pytest tests/test_cache_roundtrip.py -v Or: python tests/test_cache_roundtrip.py (standalone) """ import shutil import sys import tempfile from pathlib import Path from unittest.mock import patch # Bootstrap parsers before importing engine sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) from engine import ( CodeCache, _get_parser, ) # These are not yet implemented; tests will fail trying to use them. # They're listed here for reference of what contract they must fulfill. try: from engine import CACHE_FORMAT_VERSION, cache_path_for except ImportError: CACHE_FORMAT_VERSION = None cache_path_for = None # ── Helpers ────────────────────────────────────────────────────────────── def create_fixture_repo(tmpdir: str) -> str: """Create a multi-language fixture repo for testing. Returns the root path. """ root = Path(tmpdir) # Python files (root / 'py_module.py').write_text(''' class UserManager: """Manages users.""" def find_by_id(self, user_id: int): """Find a user by ID.""" return None def create_user(self, name: str): pass def process_data(items): """Process data items.""" return [] ''') (root / 'utils.py').write_text(''' import os import sys from pathlib import Path def read_config(): return {} class Config: def __init__(self): self.debug = False ''') # JavaScript/TypeScript files (root / 'app.js').write_text(''' /** Main application. */ class App { /** Initialize the app. */ init(config) { this.config = config; } /** Start the server. */ start() { console.log('Starting...'); } } function createApp(name) { return new App(); } const shutdown = async () => {}; ''') (root / 'types.ts').write_text(''' interface User { id: number; name: string; } interface Config { debug: boolean; port: number; } function createUser(name: string): User { return { id: 1, name }; } export const DEFAULT_CONFIG: Config = { debug: false, port: 3000, }; ''') # C file (root / 'utils.c').write_text(''' #include <stdio.h> int add(int a, int b) { return a + b; } typedef struct { int x; int y; } Point; enum Color { RED, GREEN, BLUE }; ''') # Create a subdirectory with more files sub = root / 'src' sub.mkdir() (sub / 'core.py').write_text(''' class Engine: def __init__(self): self.state = None def run(self): """Run the engine.""" pass def stop(self): pass ''') # Markdown file (root / 'README.md').write_text(''' # My Project ## Overview This is a test project. ### Features - Feature 1 - Feature 2 ## Usage See below. ''') return str(root) # ── Test: Cache Roundtrip Equivalence ──────────────────────────────────── def test_cache_roundtrip_fresh_then_hit(): """Verify cache hit returns identical results to fresh parse. 1. Scan fresh repo with use_cache=True (writes cache) 2. Create new CodeCache, scan same repo (should hit cache) 3. Assert all query outputs are identical """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # ── Fresh scan (writes cache) ── cache_fresh = CodeCache() stats_fresh = cache_fresh.scan(repo_root, use_cache=True) assert 'loaded_from_cache' in stats_fresh, \ "scan() must return 'loaded_from_cache' key" assert stats_fresh['loaded_from_cache'] is False, \ "First scan should not load from cache" assert stats_fresh['files'] > 0, "Fixture should have files" # Capture fresh results from all query functions fresh_results = { 'tree_overview': cache_fresh.tree_overview(), 'dir_overview': cache_fresh.dir_overview(''), 'find_symbol_create': cache_fresh.find_symbol('create*'), 'find_symbol_exact': cache_fresh.find_symbol('UserManager'), 'file_symbols_py': cache_fresh.file_symbols('py_module.py'), 'file_imports_utils': cache_fresh.file_imports('utils.py'), 'references_Config': cache_fresh.references('Config'), 'get_source_range': cache_fresh.get_source_range('utils.py', 6, 8), } # Verify fresh results are non-empty/valid assert fresh_results['tree_overview'], "tree_overview should return content" assert fresh_results['dir_overview'], "dir_overview should return content" assert len(fresh_results['find_symbol_create']) > 0, \ "Should find symbols matching 'create*'" assert len(fresh_results['find_symbol_exact']) > 0, \ "Should find exact symbol 'UserManager'" # ── Cache hit scan ── cache_hit = CodeCache() stats_hit = cache_hit.scan(repo_root, use_cache=True) assert stats_hit['loaded_from_cache'] is True, \ "Second scan should load from cache" # Capture results from cache hit hit_results = { 'tree_overview': cache_hit.tree_overview(), 'dir_overview': cache_hit.dir_overview(''), 'find_symbol_create': cache_hit.find_symbol('create*'), 'find_symbol_exact': cache_hit.find_symbol('UserManager'), 'file_symbols_py': cache_hit.file_symbols('py_module.py'), 'file_imports_utils': cache_hit.file_imports('utils.py'), 'references_Config': cache_hit.references('Config'), 'get_source_range': cache_hit.get_source_range('utils.py', 6, 8), } # ── Equivalence: all results must be identical ── assert fresh_results['tree_overview'] == hit_results['tree_overview'], \ "tree_overview must return identical results" assert fresh_results['dir_overview'] == hit_results['dir_overview'], \ "dir_overview must return identical results" # find_symbol results: compare by symbol names and properties fresh_names = {s.name for s in fresh_results['find_symbol_create']} hit_names = {s.name for s in hit_results['find_symbol_create']} assert fresh_names == hit_names, \ "find_symbol('create*') must return same symbol names" fresh_exact = fresh_results['find_symbol_exact'] hit_exact = hit_results['find_symbol_exact'] assert len(fresh_exact) == len(hit_exact), \ "find_symbol exact match must return same count" if fresh_exact: s_fresh = fresh_exact[0] s_hit = hit_exact[0] assert s_fresh.name == s_hit.name, "Symbol names must match" assert s_fresh.kind == s_hit.kind, "Symbol kinds must match" assert s_fresh.file == s_hit.file, "Symbol files must match" assert s_fresh.line == s_hit.line, "Symbol lines must match" # file_symbols comparison fresh_file_syms = fresh_results['file_symbols_py'] hit_file_syms = hit_results['file_symbols_py'] assert len(fresh_file_syms) == len(hit_file_syms), \ "file_symbols must return same count" for fs, hs in zip(fresh_file_syms, hit_file_syms): assert fs.name == hs.name, "file_symbols names must match" assert fs.kind == hs.kind, "file_symbols kinds must match" # file_imports comparison fresh_imps = fresh_results['file_imports_utils'] hit_imps = hit_results['file_imports_utils'] assert fresh_imps == hit_imps, \ "file_imports must return identical list" # references comparison fresh_refs = fresh_results['references_Config'] hit_refs = hit_results['references_Config'] assert len(fresh_refs) == len(hit_refs), \ "references must return same count" assert fresh_refs == hit_refs, \ "references must return identical results" # get_source_range comparison assert fresh_results['get_source_range'] == hit_results['get_source_range'], \ "get_source_range must return identical source" finally: shutil.rmtree(tmpdir) def test_cache_no_parsing_on_hit(): """Verify parser factory is NOT called on cache hit. Monkeypatch _get_parser to track calls: - Fresh scan: _get_parser MUST be called - Cache hit scan: _get_parser MUST NOT be called """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # ── Fresh scan with spy ── parse_call_count_fresh = 0 original_get_parser = None def spy_get_parser(lang): nonlocal parse_call_count_fresh parse_call_count_fresh += 1 return original_get_parser(lang) # Fresh scan original_get_parser = _get_parser.__wrapped__ if hasattr(_get_parser, '__wrapped__') else _get_parser with patch('engine._get_parser', side_effect=spy_get_parser) as mock_parser: cache_fresh = CodeCache() stats_fresh = cache_fresh.scan(repo_root, use_cache=True) fresh_parse_calls = mock_parser.call_count assert stats_fresh['loaded_from_cache'] is False assert fresh_parse_calls > 0, \ "Fresh scan should call _get_parser for language detection" # ── Cache hit with spy ── parse_call_count_hit = 0 def spy_get_parser_hit(lang): nonlocal parse_call_count_hit parse_call_count_hit += 1 return original_get_parser(lang) with patch('engine._get_parser', side_effect=spy_get_parser_hit) as mock_parser_hit: cache_hit = CodeCache() stats_hit = cache_hit.scan(repo_root, use_cache=True) hit_parse_calls = mock_parser_hit.call_count assert stats_hit['loaded_from_cache'] is True, \ "Second scan should load from cache" # On cache hit, _get_parser should NOT be called at all # (the cache contains already-extracted symbols, no parsing needed) assert hit_parse_calls == 0, \ f"Cache hit should not call _get_parser, but got {hit_parse_calls} calls" finally: shutil.rmtree(tmpdir) def test_cache_use_cache_false(): """Verify use_cache=False never reads or writes cache. Scan with use_cache=False twice: - Both scans return loaded_from_cache=False - Results still correct """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # First scan with use_cache=False cache1 = CodeCache() stats1 = cache1.scan(repo_root, use_cache=False) assert stats1['loaded_from_cache'] is False, \ "use_cache=False should never load from cache" assert stats1['files'] > 0, "Should parse files" results1 = { 'tree_overview': cache1.tree_overview(), 'symbols': cache1.find_symbol('create*'), } # Second scan with use_cache=False cache2 = CodeCache() stats2 = cache2.scan(repo_root, use_cache=False) assert stats2['loaded_from_cache'] is False, \ "use_cache=False should never load from cache" results2 = { 'tree_overview': cache2.tree_overview(), 'symbols': cache2.find_symbol('create*'), } # Results should still be identical assert results1['tree_overview'] == results2['tree_overview'] assert len(results1['symbols']) == len(results2['symbols']) finally: shutil.rmtree(tmpdir) def test_cache_rebuild_cache_flag(): """Verify rebuild_cache=True forces fresh parse and overwrites cache. Scan with rebuild_cache=True should: - Return loaded_from_cache=False - Parse fresh - Overwrite the cache file """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # Initial scan cache1 = CodeCache() stats1 = cache1.scan(repo_root, use_cache=True) assert stats1['loaded_from_cache'] is False results1 = cache1.find_symbol('UserManager') # Rebuild cache cache2 = CodeCache() stats2 = cache2.scan(repo_root, use_cache=True, rebuild_cache=True) assert stats2['loaded_from_cache'] is False, \ "rebuild_cache=True should parse fresh, not load" results2 = cache2.find_symbol('UserManager') # Results should match assert len(results1) == len(results2), \ "rebuild_cache should produce same results as original parse" if results1: assert results1[0].name == results2[0].name finally: shutil.rmtree(tmpdir) def test_cache_format_version_defined(): """Verify CACHE_FORMAT_VERSION constant exists and is an int.""" assert hasattr(__import__('engine'), 'CACHE_FORMAT_VERSION'), \ "engine.CACHE_FORMAT_VERSION must be defined" version = __import__('engine').CACHE_FORMAT_VERSION assert isinstance(version, int), \ f"CACHE_FORMAT_VERSION must be int, got {type(version)}" def test_cache_path_for_function(): """Verify cache_path_for() function exists and returns Path. Should: - Return a pathlib.Path - Be deterministic (same root -> same path) - Honor TREESIT_CACHE_DIR env var if set """ assert hasattr(__import__('engine'), 'cache_path_for'), \ "engine.cache_path_for must be defined" # Test determinism root = '/some/project/root' path1 = cache_path_for(root) path2 = cache_path_for(root) assert isinstance(path1, Path), "cache_path_for must return pathlib.Path" assert path1 == path2, "cache_path_for must be deterministic" def test_cache_fingerprint_method(): """Verify CodeCache._fingerprint() method exists and works. Should: - Return a string - Be stable across rescans of unchanged tree - Change on file add/remove/modify """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) cache = CodeCache() assert hasattr(cache, '_fingerprint'), \ "CodeCache must have _fingerprint method" # Get fingerprint before any changes fp1 = cache._fingerprint(repo_root, skip=None) assert isinstance(fp1, str), "_fingerprint must return string" assert len(fp1) > 0, "_fingerprint must return non-empty string" # Get fingerprint again (should be identical) fp2 = cache._fingerprint(repo_root, skip=None) assert fp1 == fp2, "_fingerprint must be stable for unchanged tree" # Modify a file mod_file = Path(repo_root) / 'utils.py' mod_file.write_text(mod_file.read_text() + '\n# comment') # Fingerprint should change fp3 = cache._fingerprint(repo_root, skip=None) assert fp3 != fp1, "_fingerprint must change when file is modified" # Add a file new_file = Path(repo_root) / 'new.py' new_file.write_text('# new file') fp4 = cache._fingerprint(repo_root, skip=None) assert fp4 != fp3, "_fingerprint must change when file is added" finally: shutil.rmtree(tmpdir) def test_cache_file_entry_lazy_source(): """Verify lazy source loading: tree=None and source=None on cache hit. After cache hit, FileEntry should have: - tree: None (not loaded) - source: None (not loaded) - symbols: populated from cache - But get_source_range() still works by lazily reading disk """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # Fresh scan cache_fresh = CodeCache() cache_fresh.scan(repo_root, use_cache=True) fresh_entry = cache_fresh.files.get('utils.py') assert fresh_entry is not None assert fresh_entry.source is not None, "Fresh parse should have source" assert fresh_entry.tree is not None, "Fresh parse should have tree" # Cache hit cache_hit = CodeCache() cache_hit.scan(repo_root, use_cache=True) hit_entry = cache_hit.files.get('utils.py') assert hit_entry is not None # On cache hit, tree and source should be None (lazy loaded) if cache_hit.root: # Only check if cache was actually loaded assert hit_entry.tree is None, \ "Cache hit should not load tree (lazy loading)" assert hit_entry.source is None, \ "Cache hit should not load source (lazy loading)" # But symbols should be populated assert len(hit_entry.symbols) > 0, \ "Cache hit should populate symbols" # And get_source_range() should still work (lazy loading) source = cache_hit.get_source_range('utils.py', 6, 8) assert 'read_config' in source or 'return {}' in source, \ "get_source_range should work with lazy source loading" finally: shutil.rmtree(tmpdir) def test_cache_symbol_children_roundtrip(): """Verify Symbol.to_dict(include_children=True) roundtrips correctly. Cache stores symbols via to_dict(), must deserialize children properly. """ tmpdir = tempfile.mkdtemp() try: repo_root = create_fixture_repo(tmpdir) # Fresh scan cache_fresh = CodeCache() cache_fresh.scan(repo_root, use_cache=True) # Find a symbol with children (class) user_mgr = cache_fresh.find_symbol('UserManager') assert len(user_mgr) > 0 fresh_sym = user_mgr[0] assert fresh_sym.kind == 'class' fresh_children_names = {c.name for c in fresh_sym.children} # Cache hit cache_hit = CodeCache() cache_hit.scan(repo_root, use_cache=True) user_mgr_hit = cache_hit.find_symbol('UserManager') assert len(user_mgr_hit) > 0 hit_sym = user_mgr_hit[0] hit_children_names = {c.name for c in hit_sym.children} # Children must match assert fresh_children_names == hit_children_names, \ "Symbol children must roundtrip correctly through cache" finally: shutil.rmtree(tmpdir) # ── Standalone runner ──────────────────────────────────────────────────── if __name__ == '__main__': import traceback tests = [v for k, v in sorted(globals().items()) if k.startswith('test_') and callable(v)] passed = failed = 0 for test in tests: try: test() passed += 1 print(f" ✓ {test.__name__}") except Exception as e: failed += 1 print(f" ✗ {test.__name__}: {e}") traceback.print_exc() print(f"\n{passed} passed, {failed} failed") sys.exit(1 if failed else 0) -
test_cli_cache.py 14.6 KB
"""Tests for tree-sitting CLI cache integration. Tests the --no-cache and --rebuild-cache flags, cache file creation/loading, and the "(cached)" marker in stats output. Run: python -m pytest tests/test_cli_cache.py -v Or: cd /home/user/claude-skills/tree-sitting && python -m pytest tests/test_cli_cache.py -v """ import os import subprocess import sys from pathlib import Path class TestCLICacheIntegration: """CLI invocations with caching behavior.""" @staticmethod def _fixture_repo(tmp_path: Path) -> Path: """Create a minimal test repository with Python and Rust files.""" repo = tmp_path / "fixture_repo" repo.mkdir() # Python file with named symbols py_file = repo / "example.py" py_file.write_text("""\ def greet(name: str) -> str: \"\"\"Greet someone.\"\"\" return f"Hello, {name}!" class UserService: \"\"\"Manages users.\"\"\" def find_by_id(self, user_id: int): \"\"\"Find a user by ID.\"\"\" return None def create(self, name: str): \"\"\"Create a new user.\"\"\" return {"name": name} def utility_function(): pass """) # Rust file with named symbols rs_file = repo / "lib.rs" rs_file.write_text("""\ pub struct Config { pub name: String, } pub trait Handler { fn handle(&self); } pub fn create_config(name: &str) -> Config { Config { name: name.to_string(), } } pub enum Status { Active, Inactive, } impl Config { pub fn new(name: String) -> Self { Config { name } } } """) return repo @staticmethod def _run_cli(repo_path: str, cache_dir: str, queries: list = None, flags: list = None) -> tuple[str, str, int]: """Run treesit.py CLI and return (stdout, stderr, returncode). repo_path: absolute path to repo to scan cache_dir: absolute path to cache directory (set as TREESIT_CACHE_DIR env) queries: list of query strings (e.g., ['find:greet', 'find:Config']) flags: list of flag strings (e.g., ['--no-cache', '--stats']) """ python = sys.executable treesit_py = str(Path(__file__).resolve().parent.parent / "scripts" / "treesit.py") cmd = [python, treesit_py, repo_path] if flags: cmd.extend(flags) if queries: cmd.extend(queries) env = os.environ.copy() env["TREESIT_CACHE_DIR"] = cache_dir result = subprocess.run( cmd, capture_output=True, text=True, env=env, cwd="/home/user/claude-skills/tree-sitting" ) return result.stdout, result.stderr, result.returncode def test_first_invocation_creates_cache_no_marker(self, tmp_path: Path): """First CLI invocation parses and creates cache; no '(cached)' in stats.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation: should parse (no cache yet) stdout, stderr, rc = self._run_cli( str(repo), str(cache_dir), queries=[], flags=["--stats"] ) assert rc == 0, f"CLI failed: {stderr}" assert "(cached)" not in stdout, "First run should NOT show '(cached)' marker" assert "Scanned" in stdout, "Stats output should be present" # Cache file should now exist cache_files = list(cache_dir.glob("*")) assert len(cache_files) > 0, "Cache file should be created after first invocation" def test_second_invocation_serves_from_cache(self, tmp_path: Path): """Second identical invocation served from cache; stats show '(cached)'.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) # Second invocation: should hit cache stdout, stderr, rc = self._run_cli( str(repo), str(cache_dir), queries=[], flags=["--stats"] ) assert rc == 0, f"CLI failed: {stderr}" assert "(cached)" in stdout, "Second run should show '(cached)' marker in stats" def test_no_cache_flag_disables_cache(self, tmp_path: Path): """--no-cache flag: never read/write cache; no '(cached)' even on repeat.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation with --no-cache stdout1, stderr1, rc1 = self._run_cli( str(repo), str(cache_dir), queries=[], flags=["--no-cache", "--stats"] ) assert rc1 == 0, f"First --no-cache run failed: {stderr1}" assert "(cached)" not in stdout1, "First --no-cache run should NOT have '(cached)'" # Verify no cache file created cache_files = list(cache_dir.glob("*")) assert len(cache_files) == 0, "Cache file should NOT be created with --no-cache" # Second invocation with --no-cache (identical to first) stdout2, stderr2, rc2 = self._run_cli( str(repo), str(cache_dir), queries=[], flags=["--no-cache", "--stats"] ) assert rc2 == 0, f"Second --no-cache run failed: {stderr2}" assert "(cached)" not in stdout2, "Repeat --no-cache run should NOT have '(cached)'" # Still no cache file cache_files = list(cache_dir.glob("*")) assert len(cache_files) == 0, "Cache file should still NOT exist after second --no-cache run" # Positive control: prove caching is real, so the negatives above aren't # vacuously green. A normal run creates a cache; a normal repeat is a hit. self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) pc_out, _, pc_rc = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert pc_rc == 0 assert list(cache_dir.glob("*")), "positive control: a cache-enabled run must create a cache file" assert "(cached)" in pc_out, "positive control: a cache-enabled repeat must show '(cached)'" def test_rebuild_cache_flag_refreshes_cache(self, tmp_path: Path): """--rebuild-cache: ignore existing cache, re-parse, and overwrite cache.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation: normal (creates cache) self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) # Positive control: a normal repeat must be a cache hit, so the # "not (cached)" assertion below is meaningful rather than vacuous. pc_out, _, pc_rc = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert pc_rc == 0 assert "(cached)" in pc_out, "positive control: normal repeat must be a cache hit" # Second invocation with --rebuild-cache stdout, stderr, rc = self._run_cli( str(repo), str(cache_dir), queries=[], flags=["--rebuild-cache", "--stats"] ) assert rc == 0, f"--rebuild-cache run failed: {stderr}" assert "(cached)" not in stdout, "--rebuild-cache should NOT show '(cached)' marker" assert "Scanned" in stdout, "Stats should be present (fresh parse)" def test_multi_query_with_cache(self, tmp_path: Path): """Batched multi-query invocation with caching: scans once, runs all queries.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # Multi-query invocation with find and source stdout, stderr, rc = self._run_cli( str(repo), str(cache_dir), queries=["find:greet", "find:Config", "source:UserService"], flags=["--stats"] ) assert rc == 0, f"Multi-query run failed: {stderr}" assert "Scanned" in stdout, "Stats should be present" # All three queries should appear in output # Note: the exact output depends on implementation, but should include results assert stdout, "Query results should be in output" # A second identical batched run must be a cache hit and return the same # results (proving the cache path serves batched queries correctly). stdout2, stderr2, rc2 = self._run_cli( str(repo), str(cache_dir), queries=["find:greet", "find:Config", "source:UserService"], flags=["--stats"], ) assert rc2 == 0, f"Second multi-query run failed: {stderr2}" assert "(cached)" in stdout2, "batched multi-query repeat must be a cache hit" def test_cache_hit_and_no_cache_identical_output(self, tmp_path: Path): """Cache-hit and --no-cache runs produce identical query output (except stats).""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() query = "find:greet" flags_normal = ["--stats"] flags_no_cache = ["--no-cache", "--stats"] # First: populate cache self._run_cli(str(repo), str(cache_dir), queries=[query], flags=flags_normal) # Second: hit cache stdout_cached, stderr_cached, rc_cached = self._run_cli( str(repo), str(cache_dir), queries=[query], flags=flags_normal ) assert rc_cached == 0 # Third: no-cache (equivalent of re-parsing) stdout_no_cache, stderr_no_cache, rc_no_cache = self._run_cli( str(repo), str(cache_dir), queries=[query], flags=flags_no_cache ) assert rc_no_cache == 0 # Strip the stats lines (they'll differ due to "(cached)" marker and timing) # and compare the rest def extract_query_output(output: str) -> str: """Extract query results, stripping stats and timing info.""" lines = output.split('\n') # Skip lines with "Scanned", "Symbols:", "Errors:", "(cached)", timing filtered = [ line for line in lines if line and not any(x in line for x in ["Scanned", "Symbols:", "Errors:", "(cached)"]) ] return '\n'.join(filtered) query_out_cached = extract_query_output(stdout_cached) query_out_no_cache = extract_query_output(stdout_no_cache) assert query_out_cached, "Cache hit should produce query output" assert query_out_no_cache, "No-cache run should produce query output" # The essential query results should be identical assert query_out_cached == query_out_no_cache, \ f"Query output differs between cache-hit and no-cache.\nCached:\n{query_out_cached}\n\nNo-cache:\n{query_out_no_cache}" def test_cache_survives_tree_unchanged(self, tmp_path: Path): """Cache is reused across invocations when files are unchanged.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation stdout1, _, rc1 = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert rc1 == 0 assert "(cached)" not in stdout1 # Second invocation (same repo, same files) stdout2, _, rc2 = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert rc2 == 0 assert "(cached)" in stdout2, "Cache should be reused when files unchanged" def test_cache_invalidated_on_file_change(self, tmp_path: Path): """Cache is invalidated when files are modified.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) # Positive control: unchanged repeat must be a cache hit before we can # meaningfully assert that a modification invalidates it. pc_out, _, _ = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert "(cached)" in pc_out, "cache must engage on an unchanged repeat first" # Modify a file py_file = repo / "example.py" py_file.write_text(py_file.read_text() + "\n\ndef new_function():\n pass\n") # Third invocation: should detect change and re-parse (no "(cached)") stdout3, _, rc3 = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert rc3 == 0 assert "(cached)" not in stdout3, "Cache should be invalidated after file modification" def test_cache_invalidated_on_file_add(self, tmp_path: Path): """Cache is invalidated when new files are added.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) # Positive control: unchanged repeat must be a cache hit first. pc_out, _, _ = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert "(cached)" in pc_out, "cache must engage on an unchanged repeat first" # Add a new file new_file = repo / "new_module.py" new_file.write_text("def new_func():\n pass\n") # Third invocation: should detect change and re-parse stdout3, _, rc3 = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert rc3 == 0 assert "(cached)" not in stdout3, "Cache should be invalidated after adding file" def test_cache_invalidated_on_file_delete(self, tmp_path: Path): """Cache is invalidated when files are deleted.""" repo = self._fixture_repo(tmp_path) cache_dir = tmp_path / "cache" cache_dir.mkdir() # First invocation self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) # Positive control: unchanged repeat must be a cache hit first. pc_out, _, _ = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert "(cached)" in pc_out, "cache must engage on an unchanged repeat first" # Delete a file py_file = repo / "example.py" py_file.unlink() # Third invocation: should detect change and re-parse stdout3, _, rc3 = self._run_cli(str(repo), str(cache_dir), queries=[], flags=["--stats"]) assert rc3 == 0 assert "(cached)" not in stdout3, "Cache should be invalidated after deleting file" if __name__ == '__main__': import pytest pytest.main([__file__, '-v']) -
test_engine.py 19 KB
"""Tests for tree-sitting engine. Run: python -m pytest tests/test_engine.py -v Or: python tests/test_engine.py (standalone) """ import shutil import sys import tempfile from pathlib import Path # Bootstrap parsers before importing engine sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) from engine import ( EXTRACTORS, TAGS_SCM, CodeCache, Symbol, _get_parser, extract_symbols, ) # ── Helpers ────────────────────────────────────────────────────────────── def parse_and_extract(lang: str, code: bytes, relpath: str = 'test') -> list[Symbol]: parser = _get_parser(lang) assert parser is not None, f"Parser for {lang} not available" tree = parser.parse(code) return extract_symbols(tree, code, relpath, lang) def names(symbols: list[Symbol]) -> list[str]: return [s.name for s in symbols] def find(symbols: list[Symbol], name: str) -> Symbol: for s in symbols: if s.name == name: return s for c in s.children: if c.name == name: return c raise AssertionError(f"Symbol {name!r} not found in {names(symbols)}") class AssertionError(AssertionError if False else AssertionError): pass # ── Bootstrap ──────────────────────────────────────────────────────────── def test_parser_available(): """Parsers for bundled languages are available.""" for lang in ('python', 'javascript', 'go', 'rust', 'ruby', 'java', 'c', 'markdown', 'mojo'): parser = _get_parser(lang) assert parser is not None, f"Parser for {lang} should be available" def test_parser_unavailable_graceful(): """Unknown language returns None, not an exception.""" parser = _get_parser('nonexistent_lang_xyz') assert parser is None # ── Extraction routing ─────────────────────────────────────────────────── def test_routing_custom(): """Languages with custom extractors use them.""" for lang in ('python', 'c', 'go', 'rust', 'javascript', 'typescript', 'tsx', 'ruby', 'markdown'): assert lang in EXTRACTORS, f"{lang} should have a custom extractor" def test_routing_tags_scm(): """Languages without custom extractors but with tags.scm use those.""" for lang in ('java', 'cpp', 'c_sharp', 'mojo'): assert lang not in EXTRACTORS, f"{lang} should NOT have a custom extractor" assert lang in TAGS_SCM, f"{lang} should have tags.scm" def test_routing_fallback(): """Languages with neither use generic extractor (returns something).""" code = b'function hello() {}\nclass Foo {}' # Lua uses generic — just verify no crash parser = _get_parser('lua') if parser: tree = parser.parse(b'function hello() end') syms = extract_symbols(tree, b'function hello() end', 'test.lua', 'lua') # Generic may or may not find symbols, but shouldn't crash assert isinstance(syms, list) # ── Python extractor ──────────────────────────────────────────────────── def test_python_functions(): syms = parse_and_extract('python', b''' def hello(name: str) -> str: """Greet someone.""" return f"Hello, {name}!" def _private(): pass ''') assert 'hello' in names(syms) hello = find(syms, 'hello') assert hello.kind == 'function' assert hello.signature == '(name: str)' assert hello.doc == 'Greet someone.' def test_python_class_hierarchy(): syms = parse_and_extract('python', b''' class UserService: """Handles users.""" def find(self, user_id: int) -> dict: """Find a user.""" return {} def _internal(self): pass ''') svc = find(syms, 'UserService') assert svc.kind == 'class' assert svc.doc == 'Handles users.' method_names = [c.name for c in svc.children] assert 'find' in method_names assert '_internal' in method_names find_method = next(c for c in svc.children if c.name == 'find') assert find_method.signature == '(self, user_id: int)' assert find_method.doc == 'Find a user.' def test_python_decorated_functions(): """Decorated defs must not vanish. Python wraps ``@deco``-ed definitions in a ``decorated_definition`` node. A walk matching only ``function_definition`` drops them silently, which is how ``recall``/``remember`` went missing from an indexed module whose whole public API carries a decorator. """ syms = parse_and_extract('python', b''' import functools @functools.cache def cached(name: str) -> str: """Cached greeting.""" return name @app.route("/x") @auth_required def multi_decorated(a, b): """Two decorators.""" return a def plain(): pass @decorator def _private_decorated(): pass ''') assert 'cached' in names(syms) assert 'multi_decorated' in names(syms) assert 'plain' in names(syms) # The private filter still applies through the wrapper. assert '_private_decorated' not in names(syms) cached = find(syms, 'cached') assert cached.kind == 'function' assert cached.signature == '(name: str)' assert cached.doc == 'Cached greeting.' # Position is the `def` line, not the decorator line — callers seed a # language server from it and need the line the token is on. assert cached.line == 5 def test_python_decorated_class_and_methods(): syms = parse_and_extract('python', b''' @dataclass class Config: """Settings.""" @property def name(self) -> str: """The name.""" return self._name @staticmethod def build(raw: dict) -> "Config": return Config() def plain(self): pass ''') cfg = find(syms, 'Config') assert cfg.kind == 'class' assert cfg.doc == 'Settings.' method_names = [c.name for c in cfg.children] assert 'name' in method_names assert 'build' in method_names assert 'plain' in method_names name_method = next(c for c in cfg.children if c.name == 'name') assert name_method.signature == '(self)' assert name_method.doc == 'The name.' # ── C extractor ────────────────────────────────────────────────────────── def test_c_functions(): syms = parse_and_extract('c', b''' int add(int a, int b) { return a + b; } static void internal() {} ''') assert 'add' in names(syms) add = find(syms, 'add') assert add.kind == 'function' assert 'int' in add.signature # Static functions should be excluded assert 'internal' not in names(syms) def test_c_structs(): syms = parse_and_extract('c', b''' typedef struct { int x, y; } Point; enum Color { RED, GREEN, BLUE }; ''') assert 'Point' in names(syms) assert 'Color' in names(syms) # ── Go extractor ───────────────────────────────────────────────────────── def test_go_functions_and_types(): syms = parse_and_extract('go', b'''package main // Config holds settings. type Config struct { Name string } // NewConfig creates a Config. func NewConfig(name string) *Config { return &Config{Name: name} } ''') assert 'Config' in names(syms) config = find(syms, 'Config') assert config.kind == 'struct' assert config.doc == 'Config holds settings.' nc = find(syms, 'NewConfig') assert nc.kind == 'function' assert '(name string)' in nc.signature assert nc.doc == 'NewConfig creates a Config.' def test_go_receiver_methods(): syms = parse_and_extract('go', b'''package main type Server struct{} func (s *Server) Start() error { return nil } func (s *Server) Stop() {} ''') server = find(syms, 'Server') child_names = [c.name for c in server.children] assert 'Start' in child_names assert 'Stop' in child_names # ── Rust extractor ─────────────────────────────────────────────────────── def test_rust_pub_items(): syms = parse_and_extract('rust', b''' pub struct Config { pub name: String } pub enum Status { Active, Inactive } pub trait Serializable { fn serialize(&self) -> String; } pub fn create() -> Config { Config { name: String::new() } } fn private_fn() {} ''') assert 'Config' in names(syms) assert 'Status' in names(syms) assert 'Serializable' in names(syms) assert 'create' in names(syms) assert 'private_fn' not in names(syms) create = find(syms, 'create') assert '-> Config' in create.signature def test_rust_trait_impl_methods(): """Trait impl methods don't require pub.""" syms = parse_and_extract('rust', b''' pub trait Runnable { fn run(&self); } pub struct Engine {} impl Runnable for Engine { fn run(&self) {} } impl Engine { pub fn new() -> Engine { Engine {} } fn private_helper(&self) {} } ''') engine = find(syms, 'Engine') child_names = [c.name for c in engine.children] assert 'run' in child_names, "Trait impl method should be included" assert 'new' in child_names, "pub inherent method should be included" assert 'private_helper' not in child_names, "Private inherent method should be excluded" # ── TypeScript/JavaScript extractor ────────────────────────────────────── def test_js_classes_and_functions(): syms = parse_and_extract('javascript', b''' /** App controller. */ class App { /** Initialize. */ init(config) { this.ready = true; } } function createApp(name) { return new App(); } const shutdown = async () => {}; ''') app = find(syms, 'App') assert app.kind == 'class' assert app.doc == 'App controller.' assert any(c.name == 'init' for c in app.children) create = find(syms, 'createApp') assert create.kind == 'function' assert create.signature == '(name)' sd = find(syms, 'shutdown') assert sd.kind == 'function' def test_ts_interfaces_and_signatures(): syms = parse_and_extract('typescript', b''' interface Config { name: string; } function create(name: string): Config { return { name }; } const process = (items: string[]): number => items.length; ''') assert 'Config' in names(syms) config = find(syms, 'Config') assert config.kind == 'interface' create = find(syms, 'create') assert '(name: string)' in create.signature assert 'Config' in create.signature # return type # ── Ruby extractor ─────────────────────────────────────────────────────── def test_ruby_module_hierarchy(): syms = parse_and_extract('ruby', b''' module Auth class User def find(email) end def self.all end end class Admin < User def permissions end end end ''') auth = find(syms, 'Auth') assert auth.kind == 'module' child_names = [c.name for c in auth.children] assert 'User' in child_names assert 'Admin' in child_names user = next(c for c in auth.children if c.name == 'User') method_names = [m.name for m in user.children] assert 'find' in method_names assert 'self.all' in method_names def test_ruby_method_signatures(): syms = parse_and_extract('ruby', b''' def process(data, format: :json) end ''') proc = find(syms, 'process') assert proc.signature is not None assert 'data' in proc.signature # ── Markdown extractor ─────────────────────────────────────────────────── def test_markdown_heading_hierarchy(): syms = parse_and_extract('markdown', b'''# Title ## Chapter 1 ### Section 1.1 ### Section 1.2 ## Chapter 2 ### Section 2.1 ''') assert len(syms) == 1 # one h1 title = syms[0] assert title.name == 'Title' assert title.kind == 'h1' assert len(title.children) == 2 # two h2s ch1 = title.children[0] assert ch1.name == 'Chapter 1' assert len(ch1.children) == 2 # two h3s assert ch1.children[0].name == 'Section 1.1' # ── tags.scm extractor ────────────────────────────────────────────────── def test_tags_scm_java(): """Java uses tags.scm (no custom extractor).""" syms = parse_and_extract('java', b''' public class UserService { public User find(int id) { return null; } } public interface Repository { User find(int id); } ''') assert 'UserService' in names(syms) assert 'Repository' in names(syms) svc = find(syms, 'UserService') assert svc.kind == 'class' repo = find(syms, 'Repository') assert repo.kind == 'interface' def test_tags_scm_mojo(): """Mojo uses tags.scm (no custom extractor) — retires the rename-to-.py workaround documented in Muninn memory 77cf6b41.""" syms = parse_and_extract('mojo', b''' struct Point: var x: Float64 var y: Float64 fn __init__(out self, x: Float64, y: Float64): self.x = x self.y = y fn distance(self) -> Float64: return (self.x * self.x + self.y * self.y) ** 0.5 trait Drawable: fn draw(self): ... alias PI: Float64 = 3.14159 fn polynomial(x: Float64) -> Float64: return x * x + 2 * x + 1 ''') # struct captured as a class-like definition pt = find(syms, 'Point') assert pt.kind == 'class' # trait captured as interface drw = find(syms, 'Drawable') assert drw.kind == 'interface' # alias declaration captured as constant pi = find(syms, 'PI') assert pi.kind == 'constant' # top-level fn captured as function poly = find(syms, 'polynomial') assert poly.kind == 'function' # nested fn captured as method (distinct from top-level function) dist = find(syms, 'distance') assert dist.kind == 'method' # ── CodeCache ──────────────────────────────────────────────────────────── def test_cache_scan(): """scan() parses files and builds index.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'main.py').write_text('def hello(): pass\ndef world(): pass\n') Path(tmpdir, 'README.md').write_text('# Hello\n## World\n') cache = CodeCache() stats = cache.scan(tmpdir) assert stats['files'] == 2 assert stats['symbols'] >= 3 # 2 python + at least 1 markdown assert stats['errors'] == 0 assert 'python' in stats['languages'] assert 'markdown' in stats['languages'] finally: shutil.rmtree(tmpdir) def test_cache_find_symbol(): """find_symbol() works across files.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'a.py').write_text('def create_user(): pass\n') Path(tmpdir, 'b.py').write_text('def create_order(): pass\n') cache = CodeCache() cache.scan(tmpdir) results = cache.find_symbol('create*') result_names = [s.name for s in results] assert 'create_user' in result_names assert 'create_order' in result_names finally: shutil.rmtree(tmpdir) def test_cache_file_symbols(): """file_symbols() returns symbols for a specific file.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'lib.py').write_text('class Foo:\n def bar(self): pass\n') cache = CodeCache() cache.scan(tmpdir) syms = cache.file_symbols('lib.py') assert len(syms) == 1 assert syms[0].name == 'Foo' assert any(c.name == 'bar' for c in syms[0].children) finally: shutil.rmtree(tmpdir) def test_cache_get_source_range(): """get_source_range() returns correct lines.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'code.py').write_text('line1\nline2\nline3\nline4\n') cache = CodeCache() cache.scan(tmpdir) src = cache.get_source_range('code.py', 2, 3) assert 'line2' in src assert 'line3' in src assert 'line1' not in src finally: shutil.rmtree(tmpdir) def test_cache_references(): """references() finds text occurrences across files.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'a.py').write_text('class Config: pass\n') Path(tmpdir, 'b.py').write_text('def make():\n c = Config()\n return c\n') cache = CodeCache() cache.scan(tmpdir) refs = cache.references('Config') files = [r['file'] for r in refs] assert 'a.py' in files assert 'b.py' in files finally: shutil.rmtree(tmpdir) def test_cache_imports(): """Import extraction works for Python.""" tmpdir = tempfile.mkdtemp() try: Path(tmpdir, 'app.py').write_text('import os\nfrom pathlib import Path\n') cache = CodeCache() cache.scan(tmpdir) imps = cache.file_imports('app.py') assert 'os' in imps assert 'pathlib' in imps finally: shutil.rmtree(tmpdir) def test_cache_tree_overview(): """tree_overview() produces a readable summary.""" tmpdir = tempfile.mkdtemp() try: sub = Path(tmpdir, 'src') sub.mkdir() Path(sub, 'main.py').write_text('def main(): pass\n') Path(tmpdir, 'README.md').write_text('# Hello\n') cache = CodeCache() cache.scan(tmpdir) overview = cache.tree_overview() assert 'src/' in overview assert 'files' in overview.lower() finally: shutil.rmtree(tmpdir) def test_cache_skip_dirs(): """scan() skips node_modules and similar.""" tmpdir = tempfile.mkdtemp() try: nm = Path(tmpdir, 'node_modules', 'pkg') nm.mkdir(parents=True) Path(nm, 'index.js').write_text('function internal() {}\n') Path(tmpdir, 'app.js').write_text('function main() {}\n') cache = CodeCache() stats = cache.scan(tmpdir) assert stats['files'] == 1 # only app.js, not node_modules finally: shutil.rmtree(tmpdir) # ── Standalone runner ──────────────────────────────────────────────────── if __name__ == '__main__': import traceback tests = [v for k, v in sorted(globals().items()) if k.startswith('test_') and callable(v)] passed = failed = 0 for test in tests: try: test() passed += 1 print(f" ✓ {test.__name__}") except Exception as e: failed += 1 print(f" ✗ {test.__name__}: {e}") traceback.print_exc() print(f"\n{passed} passed, {failed} failed") sys.exit(1 if failed else 0) -
test_grammar_sources.py 2.5 KB
"""Grammar loading off Linux x86_64: fallbacks, the missing-grammar hint, and cache invalidation when a grammar becomes available. Run: python -m pytest tests/test_grammar_sources.py -v """ import sys from pathlib import Path import pytest sys.path.insert(0, str(Path(__file__).parent.parent / 'scripts')) import engine from engine import CodeCache @pytest.fixture def no_grammars(tmp_path, monkeypatch): """Every grammar source empty — what a Mac sees before installing wheels.""" empty = tmp_path / 'empty' empty.mkdir() monkeypatch.setattr(engine, '_PARSERS_DIR', empty) monkeypatch.setenv('TREESIT_PARSERS_DIR', str(empty)) monkeypatch.setattr(engine, '_load_wheel', lambda lang: None) _reset_memos(monkeypatch) return empty def _reset_memos(monkeypatch): for name in ('_parsers', '_languages', '_grammar_sources'): monkeypatch.setattr(engine, name, {}) def _repo(tmp_path): repo = tmp_path / 'repo' repo.mkdir() (repo / 'a.py').write_text('def f():\n return 1\n') (repo / 'data.json').write_text('{}\n') return repo def test_missing_core_grammar_warns(tmp_path, monkeypatch, no_grammars): monkeypatch.setenv('TREESIT_CACHE_DIR', str(tmp_path / 'cache')) stats = CodeCache().scan(str(_repo(tmp_path))) assert stats['files'] == 0 hint = stats['grammar_hint'] assert 'python' in hint and 'pip install tree-sitter tree-sitter-python' in hint # .json has no bundled grammar and isn't a core language: no nagging. assert 'json' not in hint def test_cache_invalidates_when_grammar_appears(tmp_path, monkeypatch, no_grammars): monkeypatch.setenv('TREESIT_CACHE_DIR', str(tmp_path / 'cache')) repo = _repo(tmp_path) empty = CodeCache().scan(str(repo)) assert empty['files'] == 0 # A grammar becomes available (wheel installed); the cached empty scan # must not be served. monkeypatch.undo() monkeypatch.setenv('TREESIT_CACHE_DIR', str(tmp_path / 'cache')) stats = CodeCache().scan(str(repo)) assert stats['loaded_from_cache'] is False assert stats['files'] == 1 and stats['grammar_hint'] is None def test_wheel_fallback_without_bundled(tmp_path, monkeypatch): pytest.importorskip('tree_sitter_python') monkeypatch.setattr(engine, '_PARSERS_DIR', tmp_path) monkeypatch.setenv('TREESIT_PARSERS_DIR', str(tmp_path)) _reset_memos(monkeypatch) assert engine._get_parser('python') is not None assert engine.grammar_source('python') == 'wheel'
-
-
CHANGELOG.md 4.5 KB
# tree-sitting - Changelog All notable changes to the `tree-sitting` skill are documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/). ## [0.9.0] - 2026-09-30 ### Added - load grammars on macOS; cut SKILL.md to essentials ## [0.8.0] - 2026-08-25 ### Other - Applicability boundaries, real failure signals, findable descriptions (#774) - tree-sitting: extract decorated defs, which were silently dropped (#753) - Deprecate mapping-codebases; adopt ruff 0.16.0 baseline (#747) ## [0.7.0] - 2026-07-16 ### Added - persistent scan-cache (0.7.0) (#732) ### Other - tree-sitting: note tree-sitter v0.26.10 is core-only (no Python binding yet) (#719) ## [0.7.0] - 2026-07-16 ### Added - Persistent filesystem scan-cache: scans are cached to disk, keyed on fileset fingerprint (mtime + size of all candidate files, combined with skip-set and cache format version). Auto-invalidates when files change, are added, or removed. Atomic writes prevent corrupt cache files. Cache hits return byte-identical results vs. fresh parse. New flags: `--no-cache` (disable cache), `--rebuild-cache` (force rewrite). New env var: `TREESIT_CACHE_DIR` to relocate cache from system temp. ### Changed - Workflow guidance: batched multi-query drills are now the default approach for structural exploration. Users should combine `find:`, `source:`, and `refs:` queries in a single invocation instead of separate calls. Added explicit note NOT to fall back to grep/sed for symbol structure lookups. ## [0.6.0] - 2026-05-17 ### Other - tree-sitting v0.6.0: bundle Mojo grammar, retire rename-to-.py workaround (#654) - tree-sitting: pin language-pack <1.6.3 in fallback install hint (#579) ## [0.6.0] - 2026-05-17 ### Added - Bundled Mojo grammar (`parsers/libtree_sitter_mojo.so`) built from [`oaustegard/tree-sitter-mojo`](https://github.com/oaustegard/tree-sitter-mojo) @ `1fe537c` (post-v1.0, 78/78 corpus + 29/29 acceptance suite passing). Registers `.mojo` and `.🔥` in `EXT_TO_LANG` and adds a tags.scm entry mirroring the grammar's own `queries/tags.scm` (structs/classes → class, traits → interface, `fn` → function or method depending on nesting, alias → constant, var → variable). Retires the rename-to-`.py` workaround documented in Muninn memory `77cf6b41`. Smoke-tested on a 31-file Mojo corpus: 982 symbols in 54ms. ## [0.5.0] - 2026-04-24 ### Other - tree-sitting: drop tree_sitter_language_pack, load bundled .so directly (fixes #572) (#573) ## [0.5.0] - 2026-04-23 ### Fixed - Drop `tree-sitter-language-pack` dependency (#572). The 1.6.x wheels install into `_native/` with no top-level package directory, making imports fail in Claude.ai containers. Even if the import is patched, the pack tries to download grammars at runtime from a domain outside the network allowlist. Grammars are now loaded directly from bundled `parsers/*.so` via `ctypes`, against the bare `tree-sitter` package (which installs cleanly). Setup is simpler (no venv) and install is ~1s. ### Changed - Setup command is now `uv pip install --system --break-system-packages tree-sitter` — no venv required. - Supported-languages list narrowed to the 11 bundled grammars (Python, JavaScript, TypeScript, TSX, Go, Rust, Ruby, Java, C, HTML, Markdown). Previously advertised languages without bundled parsers silently returned empty before this change anyway; now they're documented honestly with instructions for adding a grammar. ## [0.4.0] - 2026-04-21 ### Other - tree-sitting: show line ranges in sparse/normal tree overviews (#568) - Remove _MAP.md files, direct agents to tree-sitting for code navigation (#545) ## [0.4.0] - 2026-04-21 ### Added - Tree overview now shows `:start-end` line ranges per symbol in `sparse` and `normal` detail levels, not just `full`. The default orientation output (used by `exploring-codebases` step 2) becomes actionable: pick a symbol's line window and feed it directly to `Read` via `offset`/`limit` without a second treesit call. ## [0.3.0] - 2026-04-08 ### Added - add treesit.py CLI, fix cross-process cache loss, fix Symbol dict bug (#536) ### Other - marketplace: restructure as category-based plugins for Claude Code discovery (#530) - Add missing READMEs for searching-codebases, featuring, tree-sitting (#521) ## [0.2.0] - 2026-03-31 ### Added - tree-sitting v0.2.0 — AST navigation + tags.scm extraction (#511) ## 0.7.1 — 2026-07-26 - Docstrings no longer describe output by analogy to `_MAP.md`, which no longer exists (mapping-codebases deprecated). No behaviour change. -
README.md 1.4 KB
# tree-sitting AST-powered code navigation using tree-sitter. Parses all source files in a codebase into in-memory syntax trees, then provides fast query tools for symbol search, file/directory overview, source retrieval, and reference finding. ## Features - **Fast scanning** — parses ~250 files in ~700ms, then all queries are sub-millisecond from cache - **Symbol search** — find by exact name, substring, or glob pattern across the entire codebase - **Directory and file overview** — structural summaries with symbol counts, signatures, and doc comments - **Source retrieval** — fetch implementation of any symbol, preferring definitions over declarations - **Reference finding** — locate all textual references to a symbol via fast grep against cached source - **12 grammars** — Python, JavaScript, TypeScript, TSX, Go, Rust, Ruby, Java, C, HTML, Markdown, Mojo. Bundled as `parsers/*.so` for Linux x86_64. Other platforms install `tree-sitter-<lang>` wheels. - **Three-tier extraction** — custom extractors (richest), community tags.scm queries, and generic heuristic fallback - **Dual deployment** — direct Python calls in Claude.ai, or long-lived MCP server in Claude Code ## Dependencies - **tree-sitter** — bare parser runtime - **tree-sitter-<lang>** wheels — grammars off Linux x86_64 (macOS, arm64) - **fastmcp** — required only for MCP server mode (Claude Code) -
SKILL.md 3.3 KB
--- name: tree-sitting description: Symbol-level navigation of a local checkout using tree-sitter ASTs. Answers where a symbol is defined, what lines it spans, which symbols a file exposes, what a directory holds, and where a name is referenced — every answer carries exact line ranges to feed straight into a scoped read. Use for "where is X defined", "who calls X", "find the function/class named", "what's in this file", "give me the line range for", "show me the source of", "list the symbols in", or before editing a file you have not read. Each invocation auto-scans and is self-contained. Not for first-encounter repo orientation (use exploring-codebases), for what a codebase DOES rather than what it contains (featuring), for binding-resolved Python caller sets (searching-codebases), or for literal text and regex matching (plain ripgrep). metadata: version: 0.9.0 --- # tree-sitting Tree-sitter symbol index over a local checkout. Every result carries line ranges, so the next step is a scoped `Read`. ## Setup Grammars load from the first source that works: `$TREESIT_PARSERS_DIR` (default `~/.cache/tree-sitting/parsers/`), then the bundled `parsers/*.so` (Linux x86_64 only), then PyPI wheels. ```bash pip install tree-sitter # Linux x86_64: done pip install tree-sitter tree-sitter-{python,javascript,typescript,go,rust,ruby,java,c,html,markdown} # macOS, arm64 ``` Use the pip that belongs to the `python3` you run the CLI with. Mojo has no wheel. To get it, compile `src/parser.c src/scanner.c` from `oaustegard/tree-sitter-mojo` with `cc -shared -fPIC -I src` into `~/.cache/tree-sitting/parsers/libtree_sitter_mojo.dylib`. A file whose grammar didn't load is skipped, and stderr prints `WARNING: no grammar for …` with the install command. Fix that before trusting an empty result. ## Use `scripts/treesit.py` in this skill's directory scans the repo on every call (cached on disk; invalidated when files or grammars change), prints a tree overview, then answers the queries. Batch queries so one scan serves them all. ```bash python3 scripts/treesit.py REPO # overview, depth 1 python3 scripts/treesit.py REPO --path=src/core --detail=full # drill in python3 scripts/treesit.py REPO --no-tree 'find:Parser*' 'source:parse_input' 'refs:ParseState' ``` | Query | Returns | |---|---| | `find:PATTERN[:KIND[:LIMIT]]` | symbols by name, substring or glob | | `symbols:FILE` | every symbol in a file | | `source:SYMBOL[:FILE]` | the symbol's source | | `refs:SYMBOL[:LIMIT]` | textual references | | `imports:FILE` | a file's imports | | `dir:PATH` | directory overview | Flags: `--depth N` (-1 = all), `--detail sparse|normal|full`, `--path DIR`, `--skip DIRS`, `--no-tree`, `--stats`, `--no-cache`, `--rebuild-cache`. `scripts/engine.py` exposes `CodeCache` for use inside one Python process. Languages: Python, JavaScript, TypeScript/TSX, Go, Rust, Ruby, Java, C, HTML, Markdown (heading outline), Mojo. Other languages get generic extraction if their `tree-sitter-<lang>` wheel is installed. ## Use something else for - a repo that isn't on disk: `accessing-github-repos` - what a codebase does: `featuring` - a first look at an unfamiliar repo: `exploring-codebases` - binding-resolved Python callers: `searching-codebases --refs` - literal text or regex: `rg`
Comments (0)
Sign in to join the conversation.
Reviews (0)
No reviews yet.
No comments yet.