Skip to main content

hearth_graph/
bundled.rs

1use std::borrow::Cow;
2
3use tree_sitter::Language;
4
5use crate::{ImportKind, ImportSpec, LanguageRegistry, LanguageSpec};
6
7const JAVASCRIPT_IMPORTS_QUERY: &str = include_str!("../queries/javascript/imports.scm");
8const TYPESCRIPT_IMPORTS_QUERY: &str = include_str!("../queries/typescript/imports.scm");
9const VUE_SCRIPT_INJECTIONS_QUERY: &str = include_str!("../queries/vue/injections.scm");
10
11fn language_spec(
12    name: &'static str,
13    language: Language,
14    extensions: &[&str],
15    tags_query: Cow<'static, str>,
16) -> LanguageSpec {
17    language_spec_with_imports(name, language, extensions, tags_query, None)
18}
19
20fn language_spec_with_imports(
21    name: &'static str,
22    language: Language,
23    extensions: &[&str],
24    tags_query: Cow<'static, str>,
25    imports: Option<ImportSpec>,
26) -> LanguageSpec {
27    let spec = LanguageSpec::new(name, language, extensions).with_tags_query(tags_query);
28    match imports {
29        Some(imports) => spec.with_imports(imports),
30        None => spec,
31    }
32}
33
34fn language_spec_merging_adjacent_definitions(
35    name: &'static str,
36    language: Language,
37    extensions: &[&str],
38    tags_query: Cow<'static, str>,
39) -> LanguageSpec {
40    language_spec(name, language, extensions, tags_query)
41        .with_merge_adjacent_same_name_definitions(true)
42}
43
44fn import_kind(capture: &str) -> ImportKind {
45    match capture {
46        "import.source.static" => ImportKind::EsStatic,
47        "import.source.reexport" => ImportKind::EsReexport,
48        "import.source.dynamic" => ImportKind::EsDynamic,
49        "import.source.commonjs" => ImportKind::CommonJs,
50        "import.source.tsrequire" => ImportKind::TsImportRequire,
51        _ => unreachable!("unexpected import capture: {capture}"),
52    }
53}
54
55impl LanguageRegistry {
56    /// Creates a registry containing Hearth's bundled language grammars.
57    #[must_use]
58    pub fn bundled() -> Self {
59        let mut registry = Self::empty();
60
61        registry.register(language_spec_with_imports(
62            "rust",
63            tree_sitter_rust::LANGUAGE.into(),
64            &["rs"],
65            Cow::Borrowed(tree_sitter_rust::TAGS_QUERY),
66            Some(ImportSpec::Custom(crate::imports::rust::extract)),
67        ));
68        registry.register(language_spec_with_imports(
69            "typescript",
70            tree_sitter_typescript::LANGUAGE_TYPESCRIPT.into(),
71            &["ts", "mts", "cts"],
72            Cow::Owned(format!(
73                "{}\n{}",
74                tree_sitter_javascript::TAGS_QUERY,
75                tree_sitter_typescript::TAGS_QUERY
76            )),
77            Some(ImportSpec::Query {
78                source: Cow::Owned(format!(
79                    "{JAVASCRIPT_IMPORTS_QUERY}\n{TYPESCRIPT_IMPORTS_QUERY}"
80                )),
81                kind_map: import_kind,
82            }),
83        ));
84        registry.register(language_spec_with_imports(
85            "tsx",
86            tree_sitter_typescript::LANGUAGE_TSX.into(),
87            &["tsx"],
88            Cow::Owned(format!(
89                "{}\n{}",
90                tree_sitter_javascript::TAGS_QUERY,
91                tree_sitter_typescript::TAGS_QUERY
92            )),
93            Some(ImportSpec::Query {
94                source: Cow::Owned(format!(
95                    "{JAVASCRIPT_IMPORTS_QUERY}\n{TYPESCRIPT_IMPORTS_QUERY}"
96                )),
97                kind_map: import_kind,
98            }),
99        ));
100        registry.register(language_spec_with_imports(
101            "javascript",
102            tree_sitter_javascript::LANGUAGE.into(),
103            &["js", "mjs", "cjs"],
104            Cow::Borrowed(tree_sitter_javascript::TAGS_QUERY),
105            Some(ImportSpec::Query {
106                source: Cow::Borrowed(JAVASCRIPT_IMPORTS_QUERY),
107                kind_map: import_kind,
108            }),
109        ));
110        registry.register(language_spec_with_imports(
111            "jsx",
112            tree_sitter_javascript::LANGUAGE.into(),
113            &["jsx"],
114            Cow::Borrowed(tree_sitter_javascript::TAGS_QUERY),
115            Some(ImportSpec::Query {
116                source: Cow::Borrowed(JAVASCRIPT_IMPORTS_QUERY),
117                kind_map: import_kind,
118            }),
119        ));
120        registry.register(language_spec(
121            "go",
122            tree_sitter_go::LANGUAGE.into(),
123            &["go"],
124            Cow::Borrowed(tree_sitter_go::TAGS_QUERY),
125        ));
126        registry.register(language_spec(
127            "python",
128            tree_sitter_python::LANGUAGE.into(),
129            &["py"],
130            Cow::Borrowed(tree_sitter_python::TAGS_QUERY),
131        ));
132        registry.register(language_spec(
133            "ruby",
134            tree_sitter_ruby::LANGUAGE.into(),
135            &["rb", "rake", "gemspec"],
136            Cow::Borrowed(tree_sitter_ruby::TAGS_QUERY),
137        ));
138        registry.register(language_spec(
139            "c",
140            tree_sitter_c::LANGUAGE.into(),
141            &["c", "h"],
142            Cow::Borrowed(tree_sitter_c::TAGS_QUERY),
143        ));
144        registry.register(language_spec(
145            "cpp",
146            tree_sitter_cpp::LANGUAGE.into(),
147            &["cpp", "cc", "cxx", "hpp", "hxx"],
148            Cow::Borrowed(tree_sitter_cpp::TAGS_QUERY),
149        ));
150        registry.register(language_spec(
151            "java",
152            tree_sitter_java::LANGUAGE.into(),
153            &["java"],
154            Cow::Borrowed(tree_sitter_java::TAGS_QUERY),
155        ));
156        registry.register(language_spec(
157            "csharp",
158            tree_sitter_c_sharp::LANGUAGE.into(),
159            &["cs"],
160            Cow::Borrowed(include_str!("../queries/c_sharp/tags.scm")),
161        ));
162        registry.register(language_spec(
163            "zig",
164            tree_sitter_zig::LANGUAGE.into(),
165            &["zig"],
166            Cow::Borrowed(include_str!("../queries/zig/tags.scm")),
167        ));
168        registry.register(language_spec(
169            "bash",
170            tree_sitter_bash::LANGUAGE.into(),
171            &["sh", "bash", "zsh"],
172            Cow::Borrowed(include_str!("../queries/bash/tags.scm")),
173        ));
174        registry.register(language_spec_merging_adjacent_definitions(
175            "haskell",
176            tree_sitter_haskell::LANGUAGE.into(),
177            &["hs", "lhs"],
178            Cow::Borrowed(include_str!("../queries/haskell/tags.scm")),
179        ));
180        registry.register(language_spec(
181            "lua",
182            tree_sitter_lua::LANGUAGE.into(),
183            &["lua"],
184            Cow::Borrowed(tree_sitter_lua::TAGS_QUERY),
185        ));
186        registry.register(language_spec(
187            "php",
188            tree_sitter_php::LANGUAGE_PHP.into(),
189            &["php"],
190            Cow::Borrowed(tree_sitter_php::TAGS_QUERY),
191        ));
192        registry.register(language_spec(
193            "swift",
194            tree_sitter_swift::LANGUAGE.into(),
195            &["swift"],
196            Cow::Borrowed(tree_sitter_swift::TAGS_QUERY),
197        ));
198        registry.register(
199            LanguageSpec::new("vue", tree_sitter_vue3::LANGUAGE.into(), ["vue"])
200                .with_injections_query(VUE_SCRIPT_INJECTIONS_QUERY),
201        );
202        registry.register(language_spec(
203            "markdown",
204            tree_sitter_md::LANGUAGE.into(),
205            &["md", "markdown"],
206            Cow::Borrowed(include_str!("../queries/markdown/tags.scm")),
207        ));
208
209        registry
210    }
211}
212
213#[cfg(test)]
214mod tests {
215    use std::path::Path;
216
217    use super::*;
218    use crate::ParserPool;
219
220    #[test]
221    fn all_bundled_tags_queries_compile_and_are_pool_reachable() {
222        let registry = LanguageRegistry::bundled();
223        assert_eq!(registry.iter().len(), 20);
224
225        let mut pool = ParserPool::new(&registry);
226        let mut query_count = 0;
227        for (id, spec) in registry.iter() {
228            let Some(source) = spec.tags_query.as_deref() else {
229                assert_eq!(spec.name, "vue");
230                assert!(pool.tags_query(id).is_none());
231                continue;
232            };
233            query_count += 1;
234            tree_sitter::Query::new(&spec.language, source)
235                .unwrap_or_else(|error| panic!("{} tags query failed: {error}", spec.name));
236            assert!(
237                pool.tags_query(id).is_some(),
238                "{} is not reachable through ParserPool::tags_query",
239                spec.name
240            );
241        }
242        assert_eq!(query_count, 19);
243    }
244
245    #[test]
246    fn bundled_injection_queries_compile_and_are_pool_reachable() {
247        let registry = LanguageRegistry::bundled();
248        let mut pool = ParserPool::new(&registry);
249        let mut query_count = 0;
250
251        for (id, spec) in registry.iter() {
252            let Some(source) = spec.injections_query.as_deref() else {
253                assert!(pool.injections_query(id).is_none(), "{}", spec.name);
254                continue;
255            };
256            query_count += 1;
257            tree_sitter::Query::new(&spec.language, source)
258                .unwrap_or_else(|error| panic!("{} injections query failed: {error}", spec.name));
259            assert!(
260                pool.injections_query(id).is_some(),
261                "{} is not reachable through ParserPool::injections_query",
262                spec.name
263            );
264        }
265
266        assert_eq!(query_count, 1);
267    }
268
269    #[test]
270    fn all_bundled_import_queries_compile_and_are_pool_reachable() {
271        let registry = LanguageRegistry::bundled();
272        let mut pool = ParserPool::new(&registry);
273        let mut query_count = 0;
274
275        for (id, spec) in registry.iter() {
276            match spec.imports.as_ref() {
277                Some(ImportSpec::Query { source, .. }) => {
278                    query_count += 1;
279                    tree_sitter::Query::new(&spec.language, source).unwrap_or_else(|error| {
280                        panic!("{} imports query failed: {error}", spec.name)
281                    });
282                    assert!(
283                        pool.imports_query(id).is_some(),
284                        "{} is not reachable through ParserPool::imports_query",
285                        spec.name
286                    );
287                }
288                Some(ImportSpec::Custom(_)) | None => {
289                    assert!(pool.imports_query(id).is_none(), "{}", spec.name);
290                }
291            }
292        }
293
294        assert_eq!(query_count, 4);
295    }
296
297    #[test]
298    fn register_uses_last_extension_owner_and_increments_generation() {
299        let mut registry = LanguageRegistry::bundled();
300        let original_id = registry
301            .for_path(Path::new("main.rs"))
302            .expect("bundled Rust extension");
303        let original_generation = registry.generation();
304
305        let replacement_id = registry.register(language_spec(
306            "replacement-rust",
307            tree_sitter_rust::LANGUAGE.into(),
308            &["rs"],
309            Cow::Borrowed(tree_sitter_rust::TAGS_QUERY),
310        ));
311
312        assert_ne!(replacement_id, original_id);
313        assert_eq!(
314            registry.for_path(Path::new("main.rs")),
315            Some(replacement_id)
316        );
317        assert_eq!(registry.generation(), original_generation + 1);
318    }
319
320    #[test]
321    fn supports_symbols_for_bundled_module_extensions() {
322        let registry = LanguageRegistry::bundled();
323
324        for path in [
325            "module.mjs",
326            "module.cjs",
327            "module.mts",
328            "module.cts",
329            "module.ts",
330            "component.tsx",
331            "component.vue",
332            "lib.rs",
333        ] {
334            assert!(registry.supports_symbols(Path::new(path)), "{path}");
335        }
336
337        for path in ["style.css", "Makefile"] {
338            assert!(!registry.supports_symbols(Path::new(path)), "{path}");
339        }
340    }
341
342    #[test]
343    fn supports_imports_for_bundled_module_extensions() {
344        let registry = LanguageRegistry::bundled();
345
346        for path in [
347            "module.mjs",
348            "module.cjs",
349            "module.mts",
350            "module.cts",
351            "module.ts",
352            "component.tsx",
353            "component.jsx",
354            "component.vue",
355            "lib.rs",
356        ] {
357            assert!(registry.supports_imports(Path::new(path)), "{path}");
358        }
359
360        for path in ["main.go", "style.css", "Makefile"] {
361            assert!(!registry.supports_imports(Path::new(path)), "{path}");
362        }
363    }
364}