Skip to main content

cpd_semantic/extract/
mod.rs

1//! Functions for `--semantic`.
2//!
3//! JavaScript and TypeScript (and the scripts of Vue, Svelte and Astro
4//! components) come from cpd-tokenizer's oxc extractor, the one
5//! `--similarity` compares. This module adds the languages only `--semantic`
6//! reads: Rust (a source scanner), Python (the ruff parser) and, through
7//! tree-sitter grammars, C, C++, C#, Go, Java, Kotlin, PHP, Ruby, Scala and
8//! Swift. They find where functions are and what they are called; their
9//! `kinds` stay empty.
10
11mod grammars;
12
13use cpd_tokenizer::functions::{FunctionExtractor, RawFunction};
14use cpd_tokenizer::line_index::LineIndex;
15
16/// The extractors of this module, consulted after cpd-tokenizer's.
17static EXTRACTORS: &[&dyn FunctionExtractor] = &[
18    &RustExtractor,
19    &PythonExtractor,
20    &grammars::C,
21    &grammars::CPP,
22    &grammars::CSHARP,
23    &grammars::GO,
24    &grammars::JAVA,
25    &grammars::KOTLIN,
26    &grammars::PHP,
27    &grammars::RUBY,
28    &grammars::SCALA,
29    &grammars::SWIFT,
30];
31
32/// The extractor serving `format` for `--semantic`: cpd-tokenizer's own
33/// first, then this module's.
34pub fn extractor_for(format: &str) -> Option<&'static dyn FunctionExtractor> {
35    cpd_tokenizer::functions::extractor_for(format).or_else(|| {
36        EXTRACTORS
37            .iter()
38            .copied()
39            .find(|e| e.formats().contains(&format))
40    })
41}
42
43/// Every function of a source. Empty for formats without an extractor and
44/// for sources that fail to parse.
45pub fn extract_functions(source: &str, format: &str) -> Vec<RawFunction> {
46    cpd_tokenizer::functions::extract_with(extractor_for(format), source, format)
47}
48
49/// Python through the ruff parser: every `def` and `async def`, methods and
50/// nested functions included. The span starts at `def` (or `async`), so
51/// decorators stay out as Rust attributes do.
52pub struct PythonExtractor;
53
54impl FunctionExtractor for PythonExtractor {
55    fn grammar(&self) -> &'static str {
56        "python"
57    }
58
59    fn formats(&self) -> &'static [&'static str] {
60        &["python"]
61    }
62
63    fn extract(&self, source: &str, _format: &str) -> Vec<RawFunction> {
64        use ruff_python_ast::visitor::source_order::SourceOrderVisitor;
65        let Ok(parsed) = ruff_python_parser::parse_module(source) else {
66            return Vec::new();
67        };
68        let line_index = LineIndex::new(source.as_bytes());
69        let mut visitor = PythonFunctions {
70            source,
71            line_index: &line_index,
72            out: Vec::new(),
73        };
74        visitor.visit_body(&parsed.syntax().body);
75        visitor.out.sort_by_key(|f| f.start.offset);
76        visitor.out
77    }
78}
79
80struct PythonFunctions<'s> {
81    source: &'s str,
82    line_index: &'s LineIndex,
83    out: Vec<RawFunction>,
84}
85
86impl<'a> ruff_python_ast::visitor::source_order::SourceOrderVisitor<'a> for PythonFunctions<'_> {
87    fn visit_stmt(&mut self, stmt: &'a ruff_python_ast::Stmt) {
88        if let ruff_python_ast::Stmt::FunctionDef(f) = stmt {
89            let name_start = f.name.range.start().to_usize();
90            let end = f.range.end().to_usize();
91            // `def` is the last keyword before the name; `async` precedes it.
92            let head = &self.source[..name_start];
93            if let Some(def) = head.rfind("def") {
94                let before = head[..def].trim_end();
95                let start = match f.is_async && before.ends_with("async") {
96                    true => before.len() - "async".len(),
97                    false => def,
98                };
99                self.out.push(RawFunction {
100                    grammar: "python",
101                    name: f.name.to_string(),
102                    start: self.line_index.location(start),
103                    end: self.line_index.location(end),
104                    head: self.line_index.location(start),
105                    kinds: Vec::new(),
106                });
107            }
108        }
109        ruff_python_ast::visitor::source_order::walk_stmt(self, stmt);
110    }
111}
112
113/// Rust by scanning the source: every `fn` with a body — free functions,
114/// methods, trait methods with a default body, and functions nested in
115/// them.
116///
117/// A scan, not a parse: it knows Rust's comments (nested block comments
118/// included), strings, raw strings and character literals — so a brace or a
119/// `fn` inside one of them is not code, and a lifetime `'a` is not the start
120/// of a character — and that is all it needs to find where a function's
121/// item starts (its first keyword, after doc comments and attributes) and
122/// where its body's braces close. `fn` followed by anything but a name (a
123/// function pointer type, a macro fragment) opens nothing.
124pub struct RustExtractor;
125
126impl FunctionExtractor for RustExtractor {
127    fn grammar(&self) -> &'static str {
128        "rust"
129    }
130
131    fn formats(&self) -> &'static [&'static str] {
132        &["rust"]
133    }
134
135    fn extract(&self, source: &str, _format: &str) -> Vec<RawFunction> {
136        let line_index = LineIndex::new(source.as_bytes());
137        let mut out: Vec<RawFunction> = scan_rust_functions(source)
138            .into_iter()
139            .map(|(name, start, end)| RawFunction {
140                grammar: "rust",
141                name,
142                start: line_index.location(start),
143                end: line_index.location(end),
144                head: line_index.location(start),
145                kinds: Vec::new(),
146            })
147            .collect();
148        out.sort_by_key(|f| f.start.offset);
149        out
150    }
151}
152
153/// `(name, item start, end after the closing brace)` of every Rust function
154/// with a body.
155fn scan_rust_functions(source: &str) -> Vec<(String, usize, usize)> {
156    let b = source.as_bytes();
157    let mut out = Vec::new();
158    // Functions whose body is open: (name, item start, depth inside body).
159    let mut open: Vec<(String, usize, usize)> = Vec::new();
160    let mut depth = 0usize;
161    // Start of the item being read: the first code byte after the last
162    // `;`, `{`, `}` or attribute `]`.
163    let mut item_start: Option<usize> = None;
164    // A `fn NAME` whose body has not started: (name, item start).
165    let mut pending: Option<(String, usize)> = None;
166    // Parentheses and brackets inside a pending signature.
167    let mut nesting = 0usize;
168    let mut i = 0;
169    while i < b.len() {
170        let c = b[i];
171        if c.is_ascii_whitespace() {
172            i += 1;
173            continue;
174        }
175        if let Some(next) = skip_rust_trivia(b, i) {
176            i = next;
177            continue;
178        }
179        if item_start.is_none() {
180            item_start = Some(i);
181        }
182        if let Some(next) = skip_rust_literal(b, i) {
183            i = next;
184            continue;
185        }
186        if c.is_ascii_alphabetic() || c == b'_' || c >= 0x80 {
187            let word_end = word_end(b, i);
188            if &b[i..word_end] == b"fn" && pending.is_none() {
189                let name_start = skip_space_and_trivia(b, word_end);
190                let name_end = word_end_if_ident(b, name_start);
191                if name_end > name_start {
192                    let name = source[name_start..name_end].to_string();
193                    pending = Some((name, item_start.unwrap_or(i)));
194                    nesting = 0;
195                    i = name_end;
196                    continue;
197                }
198            }
199            i = word_end;
200            continue;
201        }
202        match c {
203            b'(' | b'[' if pending.is_some() => nesting += 1,
204            b')' | b']' if pending.is_some() => nesting = nesting.saturating_sub(1),
205            b';' if pending.is_some() && nesting == 0 => {
206                // A declaration without a body.
207                pending = None;
208                item_start = None;
209            }
210            b'{' => {
211                depth += 1;
212                if nesting == 0
213                    && let Some((name, start)) = pending.take()
214                {
215                    open.push((name, start, depth));
216                }
217                if pending.is_none() {
218                    item_start = None;
219                }
220            }
221            b'}' => {
222                if let Some((_, _, d)) = open.last()
223                    && *d == depth
224                {
225                    let (name, start, _) = open.pop().unwrap_or_default();
226                    out.push((name, start, i + 1));
227                }
228                depth = depth.saturating_sub(1);
229                if pending.is_none() {
230                    item_start = None;
231                }
232            }
233            b';' | b']' if pending.is_none() => item_start = None,
234            _ => {}
235        }
236        i += 1;
237    }
238    out
239}
240
241/// The end of a comment starting at `i`, if one does.
242fn skip_rust_trivia(b: &[u8], i: usize) -> Option<usize> {
243    if b[i] != b'/' {
244        return None;
245    }
246    match b.get(i + 1) {
247        Some(b'/') => Some(
248            b[i..]
249                .iter()
250                .position(|&c| c == b'\n')
251                .map_or(b.len(), |p| i + p + 1),
252        ),
253        Some(b'*') => {
254            // Block comments nest in Rust.
255            let mut level = 0usize;
256            let mut j = i;
257            while j < b.len() {
258                if b[j] == b'/' && b.get(j + 1) == Some(&b'*') {
259                    level += 1;
260                    j += 2;
261                } else if b[j] == b'*' && b.get(j + 1) == Some(&b'/') {
262                    level -= 1;
263                    j += 2;
264                    if level == 0 {
265                        return Some(j);
266                    }
267                } else {
268                    j += 1;
269                }
270            }
271            Some(b.len())
272        }
273        _ => None,
274    }
275}
276
277/// The end of a string, raw string, byte string or character literal
278/// starting at `i`, if one does. A lifetime or label (`'a`) is not one.
279fn skip_rust_literal(b: &[u8], i: usize) -> Option<usize> {
280    let mut j = i;
281    // Prefixes: b"..", r"..", br"..", c"..", r#".."#.
282    while j < b.len() && matches!(b[j], b'b' | b'r' | b'c') && j - i < 2 {
283        j += 1;
284    }
285    let raw = b[i..j].contains(&b'r');
286    let mut hashes = 0;
287    if raw {
288        while b.get(j) == Some(&b'#') {
289            hashes += 1;
290            j += 1;
291        }
292    }
293    match b.get(j) {
294        Some(b'"') if j == i || raw || b[i..j].iter().all(|&p| p == b'b' || p == b'c') => {
295            j += 1;
296            while j < b.len() {
297                if !raw && b[j] == b'\\' {
298                    j += 2;
299                    continue;
300                }
301                if b[j] == b'"'
302                    && b.len() >= j + 1 + hashes
303                    && b[j + 1..j + 1 + hashes].iter().all(|&h| h == b'#')
304                {
305                    return Some(j + 1 + hashes);
306                }
307                j += 1;
308            }
309            Some(b.len())
310        }
311        Some(b'\'') if !raw && (j == i || b[i..j] == *b"b") => {
312            // `'x'`, `'\n'`, `'\u{..}'`, `'é'`; anything else is a lifetime.
313            let k = j + 1;
314            if b.get(k) == Some(&b'\\') {
315                let close = b[k..].iter().position(|&c| c == b'\'').map(|p| k + p);
316                // Skip the escaped quote of `'\''`.
317                let close = match close {
318                    Some(p) if p == k + 1 => b[k + 2..]
319                        .iter()
320                        .position(|&c| c == b'\'')
321                        .map(|q| k + 2 + q),
322                    other => other,
323                };
324                return close.map(|p| p + 1);
325            }
326            let ch_len = utf8_len(b.get(k).copied()?);
327            (b.get(k + ch_len) == Some(&b'\'')).then_some(k + ch_len + 1)
328        }
329        _ => None,
330    }
331}
332
333fn utf8_len(first: u8) -> usize {
334    match first {
335        0xF0..=0xFF => 4,
336        0xE0..=0xEF => 3,
337        0xC0..=0xDF => 2,
338        _ => 1,
339    }
340}
341
342fn word_end(b: &[u8], i: usize) -> usize {
343    let mut j = i;
344    while j < b.len() && (b[j].is_ascii_alphanumeric() || b[j] == b'_' || b[j] >= 0x80) {
345        j += 1;
346    }
347    j
348}
349
350/// End of the identifier at `i` (a raw identifier `r#name` included), or
351/// `i` when none starts there.
352fn word_end_if_ident(b: &[u8], i: usize) -> usize {
353    match b.get(i) {
354        Some(c) if c.is_ascii_alphabetic() || *c == b'_' || *c >= 0x80 => {
355            if b[i] == b'r' && b.get(i + 1) == Some(&b'#') {
356                word_end(b, i + 2)
357            } else {
358                word_end(b, i)
359            }
360        }
361        _ => i,
362    }
363}
364
365fn skip_space_and_trivia(b: &[u8], mut i: usize) -> usize {
366    loop {
367        while i < b.len() && b[i].is_ascii_whitespace() {
368            i += 1;
369        }
370        match (i < b.len()).then(|| skip_rust_trivia(b, i)).flatten() {
371            Some(next) => i = next,
372            None => return i,
373        }
374    }
375}
376
377#[cfg(test)]
378mod tests {
379    use super::*;
380    use cpd_tokenizer::functions::supports_functions;
381
382    /// Name, first line and last line of each function.
383    fn spans(fns: &[RawFunction]) -> Vec<(&str, u32, u32)> {
384        fns.iter()
385            .map(|f| (f.name.as_str(), f.start.line, f.end.line))
386            .collect()
387    }
388
389    #[test]
390    fn rust_functions_methods_and_default_trait_methods_are_found() {
391        let src = "use std::fmt;\n\n/// Adds.\npub fn add(a: i32, b: i32) -> i32 {\n    a + b\n}\n\nimpl Cart {\n    pub(crate) async fn total(&self) -> i64 {\n        fn cents(x: i64) -> i64 { x * 100 }\n        cents(self.sum)\n    }\n}\n\ntrait Named {\n    fn name(&self) -> String { String::new() }\n    fn id(&self) -> u32;\n}\n";
392        let fns = extract_functions(src, "rust");
393        assert_eq!(
394            spans(&fns),
395            vec![
396                ("add", 4, 6),
397                ("total", 9, 12),
398                ("cents", 10, 10),
399                ("name", 16, 16)
400            ]
401        );
402        // The span starts at the visibility, after the doc comment, and ends
403        // on the closing brace.
404        let add = &fns[0];
405        assert_eq!(
406            &src[add.start.offset as usize..add.end.offset as usize],
407            "pub fn add(a: i32, b: i32) -> i32 {\n    a + b\n}"
408        );
409        assert!(
410            fns.iter()
411                .all(|f| f.grammar == "rust" && f.kinds.is_empty())
412        );
413    }
414
415    #[test]
416    fn rust_braces_and_fn_inside_literals_and_comments_are_not_code() {
417        let src = r####"fn lifetimes<'a>(x: &'a str) -> &'a str { if x == "{" { x } else { "}" } }
418fn chars() -> [char; 3] { ['{', '\'', '}'] }
419fn raw() -> &'static str { r##"fn fake() { "# }"## }
420/* outer /* fn nested() { */ still comment */
421fn pointer(f: fn(u32) -> u32) -> u32 { f(1) }
422macro_rules! make { ($n:ident) => { fn $n() {} }; }
423fn last() {}
424"####;
425        let names: Vec<(String, u32)> = extract_functions(src, "rust")
426            .into_iter()
427            .map(|f| (f.name, f.end.line))
428            .collect();
429        let expected = [
430            ("lifetimes", 1),
431            ("chars", 2),
432            ("raw", 3),
433            ("pointer", 5),
434            ("last", 7),
435        ];
436        assert_eq!(
437            names,
438            expected
439                .iter()
440                .map(|(n, l)| (n.to_string(), *l))
441                .collect::<Vec<_>>()
442        );
443    }
444
445    #[test]
446    fn python_functions_methods_and_nested_ones_start_at_def() {
447        let src = "import os\n\n@cache\ndef load(path):\n    return open(path).read()\n\nclass Cart:\n    async def total(self):\n        def cents(x):\n            return x * 100\n        return cents(self.sum)\n";
448        let fns = extract_functions(src, "python");
449        assert_eq!(
450            spans(&fns),
451            vec![("load", 4, 5), ("total", 8, 11), ("cents", 9, 10)]
452        );
453        assert_eq!(&src[fns[0].start.offset as usize..][..8], "def load");
454        assert_eq!(&src[fns[1].start.offset as usize..][..9], "async def");
455        assert!(
456            fns.iter()
457                .all(|f| f.grammar == "python" && f.kinds.is_empty())
458        );
459        assert!(!supports_functions("python"));
460        assert!(extract_functions("def broken(:\n", "python").is_empty());
461    }
462
463    #[test]
464    fn rust_that_does_not_parse_still_yields_what_closes() {
465        assert!(extract_functions("fn broken( {", "rust").is_empty());
466        let fns = extract_functions("fn ok() { 1 }\nfn open() {", "rust");
467        assert_eq!(fns.len(), 1);
468        assert_eq!(fns[0].name, "ok");
469    }
470
471    #[test]
472    fn similarity_does_not_compare_what_only_semantic_reads() {
473        for format in ["rust", "python", "go"] {
474            assert!(extractor_for(format).is_some(), "{format}");
475            assert!(!supports_functions(format), "{format}");
476        }
477        assert_eq!(extractor_for("typescript").unwrap().grammar(), "oxc");
478        assert!(extractor_for("haskell").is_none());
479    }
480}