Sources/OrgHighlight/TreeSitterHighlighter.swift
129 lines · 4861 bytes
1import Foundation
2import OrgPresentation
3import SwiftTreeSitter
4import TreeSitterBash
5import TreeSitterC
6import TreeSitterCPP
7import TreeSitterClojure
8import TreeSitterElisp
9import TreeSitterGo
10import TreeSitterHaskell
11import TreeSitterJava
12import TreeSitterJavaScript
13import TreeSitterJSON
14import TreeSitterLua
15import TreeSitterPython
16import TreeSitterR
17import TreeSitterRuby
18import TreeSitterRust
19import TreeSitterScheme
20import TreeSitterTOML
21import TreeSitterYAML
22
23/// Src block highlighting with tree-sitter grammars and their own highlight queries.
24public final class TreeSitterHighlighter: CodeHighlighter, @unchecked Sendable {
25 private struct Grammar {
26 let language: Language
27 let query: Query
28 }
29
30 private let lock = NSLock()
31 private var grammars: [String: Grammar?] = [:]
32 private var cache: [String: [(range: Range<Int>, category: SyntaxCategory)]] = [:]
33
34 public init() {}
35
36 /// The grammar for an org src block language name, as `org-src-lang-modes` maps them.
37 static func grammarName(for language: String) -> String? {
38 switch language.lowercased() {
39 case "sh", "bash", "shell", "zsh": "bash"
40 case "python", "python3", "py": "python"
41 case "emacs-lisp", "elisp": "elisp"
42 case "c": "c"
43 case "c++", "cpp": "cpp"
44 case "r": "r"
45 case "js", "javascript", "node": "javascript"
46 case "java": "java"
47 case "scheme": "scheme"
48 case "clojure", "clj": "clojure"
49 case "haskell": "haskell"
50 case "rust": "rust"
51 case "go": "go"
52 case "ruby": "ruby"
53 case "json": "json"
54 case "yaml", "yml": "yaml"
55 case "toml", "conf-toml": "toml"
56 case "lua": "lua"
57 default: nil
58 }
59 }
60
61 private static func pointer(for name: String) -> OpaquePointer? {
62 switch name {
63 case "bash": tree_sitter_bash()
64 case "c": tree_sitter_c()
65 case "clojure": tree_sitter_clojure()
66 case "cpp": tree_sitter_cpp()
67 case "elisp": tree_sitter_elisp()
68 case "go": tree_sitter_go()
69 case "haskell": tree_sitter_haskell()
70 case "java": tree_sitter_java()
71 case "javascript": tree_sitter_javascript()
72 case "json": tree_sitter_json()
73 case "lua": tree_sitter_lua()
74 case "python": tree_sitter_python()
75 case "r": tree_sitter_r()
76 case "ruby": tree_sitter_ruby()
77 case "rust": tree_sitter_rust()
78 case "scheme": tree_sitter_scheme()
79 case "toml": tree_sitter_toml()
80 case "yaml": tree_sitter_yaml()
81 default: nil
82 }
83 }
84
85 /// Compiled once per language; nil when the language or its query can't be loaded.
86 private func grammar(_ name: String) -> Grammar? {
87 if let known = grammars[name] { return known }
88 var grammar: Grammar?
89 if let pointer = Self.pointer(for: name), let source = HighlightQueries.source[name] {
90 let language = Language(pointer)
91 if let query = try? Query(language: language, data: Data(source.utf8)) {
92 grammar = Grammar(language: language, query: query)
93 }
94 }
95 grammars[name] = grammar
96 return grammar
97 }
98
99 public func highlights(language: String, code: String) -> [(range: Range<Int>, category: SyntaxCategory)] {
100 guard let name = Self.grammarName(for: language) else { return [] }
101 lock.lock()
102 defer { lock.unlock() }
103 let key = name + "\u{0}" + code
104 if let cached = cache[key] { return cached }
105 guard let grammar = grammar(name) else { return [] }
106 let parser = Parser()
107 guard (try? parser.setLanguage(grammar.language)) != nil, let tree = parser.parse(code) else { return [] }
108 let context = Predicate.Context(string: code)
109 var runs: [(range: Range<Int>, category: SyntaxCategory)] = []
110 // Query patterns come in priority order, first match first; reversed, so the first
111 // pattern for a node is applied last and wins.
112 var matches: [QueryMatch] = []
113 for match in grammar.query.execute(in: tree) where match.allowed(in: context) { matches.append(match) }
114 for match in matches {
115 for capture in match.captures {
116 guard let name = capture.name, let category = SyntaxCategory(capture: name) else { continue }
117 let range = capture.range
118 runs.append((range.location..<NSMaxRange(range), category))
119 }
120 }
121 var seen = Set<Range<Int>>()
122 var result: [(range: Range<Int>, category: SyntaxCategory)] = []
123 for run in runs where seen.insert(run.range).inserted { result.append(run) }
124 result.reverse()
125 if cache.count > 512 { cache.removeAll() }
126 cache[key] = result
127 return result
128 }
129}