Skip to main content

tsift_graph/
complexity.rs

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}