1use anyhow::Result;
2use std::collections::HashMap;
3use tree_sitter::{Parser, Query, QueryCursor, StreamingIterator};
4
5use crate::lang::Lang;
6
7#[derive(Debug, Clone, Default)]
8pub struct ComplexityMetrics {
9 pub branches: i64,
10 pub loops: i64,
11 pub returns: i64,
12 pub max_nesting: i64,
13 pub unsafe_blocks: i64,
14}
15
16impl ComplexityMetrics {
17 pub fn total_complexity(&self) -> i64 {
18 self.branches + self.loops + self.returns
19 }
20
21 pub fn from_raw_fields(
22 branches: i64,
23 loops: i64,
24 returns: i64,
25 max_nesting: i64,
26 unsafe_blocks: i64,
27 ) -> Self {
28 Self {
29 branches,
30 loops,
31 returns,
32 max_nesting,
33 unsafe_blocks,
34 }
35 }
36}
37
38pub trait LanguageExtractor: Send + Sync {
39 fn lang(&self) -> Lang;
40 fn extract_complexity(&self, source: &[u8]) -> Result<ComplexityMetrics>;
41}
42
43struct BuiltinExtractor {
44 lang: Lang,
45}
46
47impl BuiltinExtractor {
48 fn complexity_query(&self) -> Option<&'static str> {
49 match self.lang {
50 #[cfg(feature = "lang-rust")]
51 Lang::Rust => Some(
52 r#"
53 (if_expression) @branch
54 (match_expression) @branch
55 (for_expression) @loop
56 (while_expression) @loop
57 (loop_expression) @loop
58 (return_expression) @return
59 (unsafe_block) @unsafe
60 "#,
61 ),
62 #[cfg(feature = "lang-python")]
63 Lang::Python => Some(
64 r#"
65 (if_statement) @branch
66 (elif_clause) @branch
67 (for_statement) @loop
68 (while_statement) @loop
69 (return_statement) @return
70 "#,
71 ),
72 #[cfg(feature = "lang-typescript")]
73 Lang::TypeScript | Lang::Tsx => Some(
74 r#"
75 (if_statement) @branch
76 (switch_statement) @branch
77 (ternary_expression) @branch
78 (for_statement) @loop
79 (for_in_statement) @loop
80 (while_statement) @loop
81 (do_statement) @loop
82 (return_statement) @return
83 "#,
84 ),
85 #[cfg(feature = "lang-javascript")]
86 Lang::JavaScript | Lang::Jsx => Some(
87 r#"
88 (if_statement) @branch
89 (switch_statement) @branch
90 (ternary_expression) @branch
91 (for_statement) @loop
92 (for_in_statement) @loop
93 (while_statement) @loop
94 (do_statement) @loop
95 (return_statement) @return
96 "#,
97 ),
98 #[cfg(feature = "lang-kotlin")]
99 Lang::Kotlin => Some(
100 r#"
101 (if_expression) @branch
102 (when_expression) @branch
103 (for_statement) @loop
104 (while_statement) @loop
105 (do_while_statement) @loop
106 (return_expression) @return
107 "#,
108 ),
109 #[cfg(feature = "lang-gdscript")]
110 Lang::GdScript => Some(
111 r#"
112 (if_statement) @branch
113 (elif_clause) @branch
114 (match_statement) @branch
115 (conditional_expression) @branch
116 (for_statement) @loop
117 (while_statement) @loop
118 (return_statement) @return
119 "#,
120 ),
121 _ => None,
122 }
123 }
124
125 fn compute_max_nesting(&self, source: &[u8]) -> i64 {
126 let ts_lang = self.lang.tree_sitter_language();
127 let mut parser = Parser::new();
128 if parser.set_language(&ts_lang).is_err() {
129 return 0;
130 }
131 let tree = match parser.parse(source, None) {
132 Some(t) => t,
133 None => return 0,
134 };
135 let mut max_depth: i64 = 0;
136 fn walk(node: tree_sitter::Node, depth: i64, max_depth: &mut i64) {
137 let kind = node.kind();
138 let is_scope = matches!(
139 kind,
140 "function_item"
141 | "function_definition"
142 | "function_declaration"
143 | "class_definition"
144 | "class_declaration"
145 | "impl_item"
146 | "if_expression"
147 | "if_statement"
148 | "for_expression"
149 | "for_statement"
150 | "while_expression"
151 | "while_statement"
152 | "loop_expression"
153 | "match_expression"
154 | "switch_statement"
155 | "when_expression"
156 | "block"
157 | "expression_list"
158 );
159 let child_depth = if is_scope { depth + 1 } else { depth };
160 if child_depth > *max_depth {
161 *max_depth = child_depth;
162 }
163 let mut cursor = node.walk();
164 for child in node.children(&mut cursor) {
165 walk(child, child_depth, max_depth);
166 }
167 }
168 walk(tree.root_node(), 0, &mut max_depth);
169 max_depth.max(0)
170 }
171}
172
173impl LanguageExtractor for BuiltinExtractor {
174 fn lang(&self) -> Lang {
175 self.lang
176 }
177
178 fn extract_complexity(&self, source: &[u8]) -> Result<ComplexityMetrics> {
179 let query_str = match self.complexity_query() {
180 Some(q) => q,
181 None => return Ok(ComplexityMetrics::default()),
182 };
183 let ts_lang = self.lang.tree_sitter_language();
184 let mut parser = Parser::new();
185 parser.set_language(&ts_lang)?;
186 let tree = parser
187 .parse(source, None)
188 .ok_or_else(|| anyhow::anyhow!("parse failed"))?;
189 let query = Query::new(&ts_lang, query_str)?;
190 let mut cursor = QueryCursor::new();
191 let mut metrics = ComplexityMetrics::default();
192
193 let capture_names: Vec<String> = query
194 .capture_names()
195 .iter()
196 .map(|s| s.to_string())
197 .collect();
198
199 let mut matches = cursor.matches(&query, tree.root_node(), source);
200 while let Some(m) = matches.next() {
201 for capture in m.captures {
202 let name = &capture_names[capture.index as usize];
203 match name.as_str() {
204 "branch" => metrics.branches += 1,
205 "loop" => metrics.loops += 1,
206 "return" => metrics.returns += 1,
207 "unsafe" => metrics.unsafe_blocks += 1,
208 _ => {}
209 }
210 }
211 }
212
213 metrics.max_nesting = self.compute_max_nesting(source);
214 Ok(metrics)
215 }
216}
217
218pub struct LanguageRegistry {
219 extractors: HashMap<String, Box<dyn LanguageExtractor>>,
220}
221
222impl LanguageRegistry {
223 pub fn new() -> Self {
224 let mut registry = Self {
225 extractors: HashMap::new(),
226 };
227 registry.register_builtins();
228 registry
229 }
230
231 fn register_builtins(&mut self) {
232 for lang in Lang::all() {
233 let ext = lang.name().to_string();
234 let extractor = BuiltinExtractor { lang };
235 self.extractors.insert(ext, Box::new(extractor));
236 }
237 }
238
239 pub fn register(&mut self, name: String, extractor: Box<dyn LanguageExtractor>) {
240 self.extractors.insert(name, extractor);
241 }
242
243 pub fn get(&self, lang_name: &str) -> Option<&dyn LanguageExtractor> {
244 self.extractors.get(lang_name).map(|e| e.as_ref())
245 }
246
247 pub fn extractor_for_extension(&self, ext: &str) -> Option<&dyn LanguageExtractor> {
248 let lang = Lang::from_extension(ext)?;
249 self.get(lang.name())
250 }
251
252 pub fn complexity_for_source(&self, lang: Lang, source: &[u8]) -> Result<ComplexityMetrics> {
253 let extractor = self.get(lang.name()).ok_or_else(|| {
254 anyhow::anyhow!("no extractor registered for language: {}", lang.name())
255 })?;
256 extractor.extract_complexity(source)
257 }
258
259 pub fn registered_languages(&self) -> Vec<&str> {
260 let mut names: Vec<&str> = self.extractors.keys().map(|s| s.as_str()).collect();
261 names.sort();
262 names
263 }
264}
265
266impl Default for LanguageRegistry {
267 fn default() -> Self {
268 Self::new()
269 }
270}
271
272#[cfg(test)]
273mod tests {
274 use super::*;
275
276 #[test]
277 fn registry_has_all_builtin_languages() {
278 let registry = LanguageRegistry::new();
279 let languages = registry.registered_languages();
280 for lang in Lang::all() {
281 assert!(
282 languages.contains(&lang.name()),
283 "missing builtin language: {}",
284 lang.name()
285 );
286 }
287 }
288
289 #[cfg(feature = "lang-rust")]
290 #[test]
291 fn rust_complexity_counting() {
292 let registry = LanguageRegistry::new();
293 let source = br#"fn example(x: i32) -> i32 {
294 if x > 0 {
295 return x;
296 }
297 for i in 0..x {
298 if i % 2 == 0 {
299 continue;
300 }
301 }
302 0
303}
304"#;
305 let metrics = registry.complexity_for_source(Lang::Rust, source).unwrap();
306 assert!(
307 metrics.branches >= 2,
308 "expected >=2 branches, got {}",
309 metrics.branches
310 );
311 assert!(
312 metrics.loops >= 1,
313 "expected >=1 loop, got {}",
314 metrics.loops
315 );
316 assert!(
317 metrics.returns >= 1,
318 "expected >=1 return, got {}",
319 metrics.returns
320 );
321 }
322
323 #[cfg(feature = "lang-python")]
324 #[test]
325 fn python_complexity_counting() {
326 let registry = LanguageRegistry::new();
327 let source = br#"def example(x):
328 if x > 0:
329 return x
330 for i in range(x):
331 if i % 2 == 0:
332 continue
333 return 0
334"#;
335 let metrics = registry
336 .complexity_for_source(Lang::Python, source)
337 .unwrap();
338 assert!(
339 metrics.branches >= 2,
340 "expected >=2 branches, got {}",
341 metrics.branches
342 );
343 assert!(
344 metrics.loops >= 1,
345 "expected >=1 loop, got {}",
346 metrics.loops
347 );
348 assert!(
349 metrics.returns >= 2,
350 "expected >=2 returns, got {}",
351 metrics.returns
352 );
353 }
354
355 #[cfg(feature = "lang-typescript")]
356 #[test]
357 fn typescript_complexity_counting() {
358 let registry = LanguageRegistry::new();
359 let source = br#"function example(x: number): number {
360 if (x > 0) {
361 return x;
362 }
363 for (let i = 0; i < x; i++) {
364 if (i % 2 === 0) continue;
365 }
366 return 0;
367}
368"#;
369 let metrics = registry
370 .complexity_for_source(Lang::TypeScript, source)
371 .unwrap();
372 assert!(
373 metrics.branches >= 2,
374 "expected >=2 branches, got {}",
375 metrics.branches
376 );
377 assert!(
378 metrics.loops >= 1,
379 "expected >=1 loop, got {}",
380 metrics.loops
381 );
382 assert!(
383 metrics.returns >= 2,
384 "expected >=2 returns, got {}",
385 metrics.returns
386 );
387 }
388
389 #[test]
390 fn total_complexity_sums_metrics() {
391 let metrics = ComplexityMetrics::from_raw_fields(3, 2, 1, 4, 0);
392 assert_eq!(metrics.total_complexity(), 6);
393 }
394
395 #[test]
396 fn extractor_for_extension_works() {
397 let registry = LanguageRegistry::new();
398 assert!(registry.extractor_for_extension("rs").is_some());
399 assert!(registry.extractor_for_extension("py").is_some());
400 assert!(registry.extractor_for_extension("xyz").is_none());
401 }
402}