Skip to main content

greplm_core/
lang.rs

1//! Language detection and tree-sitter grammar wiring.
2
3use std::borrow::Cow;
4
5use serde::{Deserialize, Deserializer, Serialize, Serializer};
6
7/// A source language greplm understands for symbol extraction.
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
9pub enum Language {
10    Rust,
11    Python,
12    JavaScript,
13    TypeScript,
14    Tsx,
15    Go,
16    Java,
17    C,
18    Cpp,
19    CSharp,
20    Ruby,
21    Php,
22    Swift,
23    Dart,
24    /// Recognized as text and indexed, but not parsed for symbols.
25    Other,
26}
27
28impl Language {
29    /// Every language variant, useful for iteration and exhaustive tests.
30    pub const ALL: [Language; 15] = [
31        Language::Rust,
32        Language::Python,
33        Language::JavaScript,
34        Language::TypeScript,
35        Language::Tsx,
36        Language::Go,
37        Language::Java,
38        Language::C,
39        Language::Cpp,
40        Language::CSharp,
41        Language::Ruby,
42        Language::Php,
43        Language::Swift,
44        Language::Dart,
45        Language::Other,
46    ];
47
48    /// Short, stable identifier stored in the index and used for `--lang` filters.
49    ///
50    /// This is the single source of truth for the serialized form; serde
51    /// (de)serialization is implemented in terms of `id`/`from_id`.
52    pub fn id(self) -> &'static str {
53        match self {
54            Language::Rust => "rust",
55            Language::Python => "python",
56            Language::JavaScript => "javascript",
57            Language::TypeScript => "typescript",
58            Language::Tsx => "tsx",
59            Language::Go => "go",
60            Language::Java => "java",
61            Language::C => "c",
62            Language::Cpp => "cpp",
63            Language::CSharp => "csharp",
64            Language::Ruby => "ruby",
65            Language::Php => "php",
66            Language::Swift => "swift",
67            Language::Dart => "dart",
68            Language::Other => "other",
69        }
70    }
71
72    /// Parse a language id back from its stable identifier.
73    pub fn from_id(s: &str) -> Option<Language> {
74        Some(match s {
75            "rust" => Language::Rust,
76            "python" => Language::Python,
77            "javascript" => Language::JavaScript,
78            "typescript" => Language::TypeScript,
79            "tsx" => Language::Tsx,
80            "go" => Language::Go,
81            "java" => Language::Java,
82            "c" => Language::C,
83            "cpp" => Language::Cpp,
84            "csharp" => Language::CSharp,
85            "ruby" => Language::Ruby,
86            "php" => Language::Php,
87            "swift" => Language::Swift,
88            "dart" => Language::Dart,
89            "other" => Language::Other,
90            _ => return None,
91        })
92    }
93
94    /// Detect a language from a file extension (without the dot). Matching is
95    /// case-insensitive; the common already-lowercased input does not allocate.
96    pub fn from_extension(ext: &str) -> Language {
97        let lower;
98        let ext = if ext.bytes().any(|b| b.is_ascii_uppercase()) {
99            lower = ext.to_ascii_lowercase();
100            lower.as_str()
101        } else {
102            ext
103        };
104        match ext {
105            "rs" => Language::Rust,
106            "py" | "pyi" | "pyw" => Language::Python,
107            "js" | "mjs" | "cjs" | "jsx" => Language::JavaScript,
108            "ts" | "mts" | "cts" => Language::TypeScript,
109            "tsx" => Language::Tsx,
110            "go" => Language::Go,
111            "java" => Language::Java,
112            "c" | "h" => Language::C,
113            "cc" | "cpp" | "cxx" | "c++" | "hpp" | "hh" | "hxx" | "h++" | "ipp" | "tpp" => {
114                Language::Cpp
115            }
116            "cs" => Language::CSharp,
117            "rb" | "rake" | "gemspec" => Language::Ruby,
118            "php" | "php5" | "php7" | "phtml" => Language::Php,
119            "swift" => Language::Swift,
120            "dart" => Language::Dart,
121            _ => Language::Other,
122        }
123    }
124
125    /// The tree-sitter grammar for this language, if symbol parsing is supported.
126    pub fn grammar(self) -> Option<tree_sitter::Language> {
127        let lang = match self {
128            Language::Rust => tree_sitter_rust::LANGUAGE.into(),
129            Language::Python => tree_sitter_python::LANGUAGE.into(),
130            Language::JavaScript => tree_sitter_javascript::LANGUAGE.into(),
131            Language::TypeScript => tree_sitter_typescript::LANGUAGE_TYPESCRIPT.into(),
132            Language::Tsx => tree_sitter_typescript::LANGUAGE_TSX.into(),
133            Language::Go => tree_sitter_go::LANGUAGE.into(),
134            Language::Java => tree_sitter_java::LANGUAGE.into(),
135            Language::C => tree_sitter_c::LANGUAGE.into(),
136            Language::Cpp => tree_sitter_cpp::LANGUAGE.into(),
137            Language::CSharp => tree_sitter_c_sharp::LANGUAGE.into(),
138            Language::Ruby => tree_sitter_ruby::LANGUAGE.into(),
139            Language::Php => tree_sitter_php::LANGUAGE_PHP.into(),
140            Language::Swift => tree_sitter_swift::LANGUAGE.into(),
141            Language::Dart => tree_sitter_dart::LANGUAGE.into(),
142            Language::Other => return None,
143        };
144        Some(lang)
145    }
146}
147
148impl Serialize for Language {
149    fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
150        serializer.serialize_str(self.id())
151    }
152}
153
154impl<'de> Deserialize<'de> for Language {
155    fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
156        let s = <Cow<'de, str>>::deserialize(deserializer)?;
157        Language::from_id(&s)
158            .ok_or_else(|| serde::de::Error::custom(format!("unknown language id: {s:?}")))
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165
166    #[test]
167    fn id_roundtrips_through_from_id() {
168        for lang in Language::ALL {
169            assert_eq!(Language::from_id(lang.id()), Some(lang));
170        }
171    }
172
173    #[test]
174    fn serde_matches_id() {
175        for lang in Language::ALL {
176            let json = serde_json::to_string(&lang).unwrap();
177            assert_eq!(json, format!("{:?}", lang.id()));
178            let back: Language = serde_json::from_str(&json).unwrap();
179            assert_eq!(back, lang);
180        }
181    }
182
183    #[test]
184    fn from_extension_is_case_insensitive() {
185        assert_eq!(Language::from_extension("RS"), Language::Rust);
186        assert_eq!(Language::from_extension("Cpp"), Language::Cpp);
187        assert_eq!(Language::from_extension("rs"), Language::Rust);
188        assert_eq!(Language::from_extension("dart"), Language::Dart);
189        assert_eq!(Language::from_extension("unknownext"), Language::Other);
190    }
191
192    #[test]
193    fn parseable_languages_have_grammars() {
194        for lang in Language::ALL {
195            if lang == Language::Other {
196                assert!(lang.grammar().is_none());
197            } else {
198                assert!(lang.grammar().is_some(), "{lang:?} missing grammar");
199            }
200        }
201    }
202}