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