from dataclasses import dataclass @dataclass class ParsedTree: tree: object source_bytes: bytes class TreeSitterParser: def __init__(self) -> None: self._parser_cache: dict[str, object] = {} def parse(self, source: str, language: str) -> ParsedTree | None: parser = self._get_parser(language) if parser is None: return None source_bytes = source.encode("utf-8", errors="ignore") try: return ParsedTree(tree=parser.parse(source_bytes), source_bytes=source_bytes) except Exception: return None def _get_parser(self, language: str): normalized = "javascript" if language == "jsx" else language if normalized in self._parser_cache: return self._parser_cache[normalized] parser = self._get_parser_from_language_package(normalized) if parser is None: parser = self._get_parser_from_grammar_package(normalized) self._parser_cache[normalized] = parser return parser def _get_parser_from_language_package(self, language: str): try: from tree_sitter_language_pack import get_parser return get_parser(language) except Exception: return None def _get_parser_from_grammar_package(self, language: str): grammar_map = { "python": ("tree_sitter_python", "language"), "javascript": ("tree_sitter_javascript", "language"), "typescript": ("tree_sitter_typescript", "language_typescript"), "tsx": ("tree_sitter_typescript", "language_tsx"), "java": ("tree_sitter_java", "language"), "go": ("tree_sitter_go", "language"), } module_name, function_name = grammar_map.get(language, (None, None)) if not module_name or not function_name: return None try: from importlib import import_module from tree_sitter import Language, Parser module = import_module(module_name) language_capsule = getattr(module, function_name)() tree_sitter_language = Language(language_capsule) parser = Parser() try: parser.language = tree_sitter_language except AttributeError: parser.set_language(tree_sitter_language) return parser except Exception: return None tree_sitter_parser = TreeSitterParser()