File size: 2,491 Bytes
b4291bc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
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()