Skip to main content

cpd_semantic/
units.rs

1//! Functions to embed for semantic clones (`--semantic`).
2//!
3//! A unit is a function found by a
4//! [`FunctionExtractor`](cpd_tokenizer::functions::FunctionExtractor)
5//! plus the text an embedding model sees: the function's own code, from the
6//! name it is declared under (a method's key, the variable an arrow is
7//! assigned to) to its end, with its comments removed and its indentation
8//! kept relative to its first line, so two copies that differ only in
9//! comments or nesting depth embed alike.
10//! Comments go because a jscpd tokenizer drops them in weak mode, which
11//! makes the rule the same for every language with an extractor.
12//!
13//! Vue, Svelte and Astro files contribute the functions of their `<script>`
14//! blocks (and Astro frontmatter), positioned in the host file and grouped
15//! by the block's format, the same format the block's detection source has.
16
17use crate::extract::{extract_functions, extractor_for};
18use crate::test_code::{inline_test, rust_test_modules};
19use cpd_core::models::{Location, Token};
20use cpd_tokenizer::line_index::LineIndex;
21use cpd_tokenizer::tokenizer::{Mode, tokenize};
22
23/// A function ready to embed, positioned in the file it was found in.
24#[derive(Debug, Clone, PartialEq)]
25pub struct RawUnit {
26    /// Grammar of the extractor that found it.
27    pub grammar: &'static str,
28    pub name: String,
29    pub start: Location,
30    pub end: Location,
31    /// The function's code without comments.
32    pub text: String,
33    /// A test that lives among the code: a Rust test or a JavaScript test
34    /// case (see [`crate::test_code`]). Tests in files of their own are
35    /// told by their paths instead.
36    pub test: bool,
37}
38
39/// The units of one detection source of a file. For a single-format file
40/// `format` is the file's own; for a component it is a script block's.
41#[derive(Debug, Clone, PartialEq)]
42pub struct UnitMap {
43    pub format: String,
44    pub units: Vec<RawUnit>,
45}
46
47const COMPONENT_FORMATS: &[&str] = &["vue", "svelte", "astro"];
48
49/// Formats [`extract_units`] finds functions in.
50pub fn supports_units(format: &str) -> bool {
51    COMPONENT_FORMATS.contains(&format) || extractor_for(format).is_some()
52}
53
54/// Every function of `source`, a file of `format`, as units to embed.
55/// Empty for formats without an extractor and for sources that do not
56/// parse.
57pub fn extract_units(source: &str, format: &str) -> Vec<UnitMap> {
58    if COMPONENT_FORMATS.contains(&format) {
59        return component_units(source, format);
60    }
61    let units = units_in(source, format, 0, None);
62    if units.is_empty() {
63        return Vec::new();
64    }
65    vec![UnitMap {
66        format: format.to_string(),
67        units,
68    }]
69}
70
71fn component_units(source: &str, file_format: &str) -> Vec<UnitMap> {
72    let host = LineIndex::new(source.as_bytes());
73    let mut maps: Vec<UnitMap> = Vec::new();
74    for (format, range) in cpd_tokenizer::sfc::script_blocks(source, file_format) {
75        if extractor_for(&format).is_none() {
76            continue;
77        }
78        let units = units_in(&source[range.clone()], &format, range.start, Some(&host));
79        if units.is_empty() {
80            continue;
81        }
82        match maps.iter_mut().find(|m| m.format == format) {
83            Some(map) => map.units.extend(units),
84            None => maps.push(UnitMap { format, units }),
85        }
86    }
87    maps
88}
89
90/// Units of `code`, which starts at byte `shift` of the file whose line
91/// index is `host` (`None` when `code` is the whole file).
92fn units_in(code: &str, format: &str, shift: usize, host: Option<&LineIndex>) -> Vec<RawUnit> {
93    let functions = extract_functions(code, format);
94    if functions.is_empty() {
95        return Vec::new();
96    }
97    let tokens = tokenize(format, code, Mode::Weak);
98    let rust_modules = match functions.iter().any(|f| f.grammar == "rust") {
99        true => rust_test_modules(code, &tokens),
100        false => Vec::new(),
101    };
102    let place = |loc: &Location| match host {
103        Some(index) => index.location(shift + loc.offset as usize),
104        None => loc.clone(),
105    };
106    functions
107        .into_iter()
108        .filter_map(|f| {
109            let text = code_text(code, &tokens, f.head.offset as usize, f.end.offset as usize);
110            let test = inline_test(
111                f.grammar,
112                code,
113                f.head.offset as usize,
114                f.start.offset as usize,
115                &rust_modules,
116            );
117            (!text.is_empty()).then(|| RawUnit {
118                test,
119                grammar: f.grammar,
120                start: place(&f.head),
121                end: place(&f.end),
122                name: f.name,
123                text,
124            })
125        })
126        .collect()
127}
128
129/// The source of the tokens inside `start..end`, joined by the whitespace
130/// between them: a line break (with the next line's indentation) where the
131/// gap has one, a single space where it has anything else. Whatever the
132/// tokenizer skipped — comments in weak mode — is gone. Lines lose the
133/// indentation of the function's first line, so the text reads as if the
134/// function stood at the top level.
135fn code_text(code: &str, tokens: &[Token], start: usize, end: usize) -> String {
136    let first = tokens.partition_point(|t| (t.start.offset as usize) < start);
137    let mut out = String::new();
138    let mut prev_end: Option<usize> = None;
139    for token in &tokens[first..] {
140        let (from, to) = (token.start.offset as usize, token.end.offset as usize);
141        if to > end {
142            break;
143        }
144        if from < to && to <= code.len() && code.is_char_boundary(from) && code.is_char_boundary(to)
145        {
146            if let Some(prev) = prev_end.filter(|&p| p <= from) {
147                let gap = &code[prev..from];
148                match gap.rfind('\n') {
149                    Some(nl) => {
150                        out.push('\n');
151                        let line_start = &gap[nl + 1..];
152                        let indent = line_start.len() - line_start.trim_start().len();
153                        out.push_str(&line_start[..indent]);
154                    }
155                    None if !gap.is_empty() => out.push(' '),
156                    None => {}
157                }
158            }
159            out.push_str(&code[from..to]);
160            prev_end = Some(to);
161        }
162    }
163    let line_start = code[..start.min(code.len())]
164        .rfind('\n')
165        .map_or(0, |nl| nl + 1);
166    let base = code[line_start..start.min(code.len())]
167        .chars()
168        .take_while(|c| c.is_whitespace())
169        .count();
170    dedent(&out, base)
171}
172
173/// Remove up to `base` leading whitespace characters from every line after
174/// the first (which starts at the function's first token).
175fn dedent(text: &str, base: usize) -> String {
176    let mut out = String::with_capacity(text.len());
177    for (i, line) in text.split('\n').enumerate() {
178        if i > 0 {
179            out.push('\n');
180            // `base` counts characters, the slice takes bytes: a no-break
181            // space is two bytes, an ideographic space three.
182            let strip: usize = line
183                .chars()
184                .take(base)
185                .take_while(|c| c.is_whitespace())
186                .map(char::len_utf8)
187                .sum();
188            out.push_str(&line[strip..]);
189        } else {
190            out.push_str(line);
191        }
192    }
193    out
194}
195
196#[cfg(test)]
197mod tests {
198    use super::*;
199
200    #[test]
201    fn rust_units_lose_comments_and_shared_indentation() {
202        let src = "impl Cart {\n    /// Sum of the lines.\n    pub fn total(&self) -> i64 {\n        // cents\n        self.lines.iter().map(|l| l.price /* each */ * l.qty).sum()\n    }\n}\n";
203        let maps = extract_units(src, "rust");
204        assert_eq!(maps.len(), 1);
205        assert_eq!(maps[0].format, "rust");
206        let unit = &maps[0].units[0];
207        assert_eq!(unit.name, "total");
208        assert_eq!(unit.grammar, "rust");
209        assert_eq!((unit.start.line, unit.end.line), (3, 6));
210        assert_eq!(
211            unit.text,
212            "pub fn total(&self) -> i64 {\n    self.lines.iter().map(|l| l.price * l.qty).sum()\n}"
213        );
214    }
215
216    #[test]
217    fn typescript_units_keep_code_and_drop_comments() {
218        let src = "// helpers\nexport function slug(t: string): string {\n  /* lower */\n  return t.toLowerCase().replace(/[^a-z0-9]+/g, '-'); // dash\n}\n";
219        let maps = extract_units(src, "typescript");
220        let unit = &maps[0].units[0];
221        assert_eq!(
222            (unit.name.as_str(), unit.start.line, unit.end.line),
223            ("slug", 2, 5)
224        );
225        // The function node starts after `export`, as `--similarity` reports it.
226        assert_eq!(
227            unit.text,
228            "function slug(t: string): string {\n  return t.toLowerCase().replace(/[^a-z0-9]+/g, '-');\n}"
229        );
230    }
231
232    #[test]
233    fn typescript_units_start_at_the_name_they_are_declared_under() {
234        let src = "class Stream {\n  async segment(n: bigint): Promise<void> {\n    return run(n);\n  }\n  handle = (e: Event) => log(e);\n}\nexport const total = (xs: number[]) => xs.reduce((a, b) => a + b, 0);\nconst api = { fetchAll(url: string) { return get(url); }, save: async (x: number) => put(x) };\nconst later = debounce(() => refresh(), 300);\n";
235        let maps = extract_units(src, "typescript");
236        let texts: Vec<(&str, &str)> = maps[0]
237            .units
238            .iter()
239            .map(|u| (u.name.as_str(), u.text.as_str()))
240            .collect();
241        assert_eq!(
242            texts,
243            vec![
244                (
245                    "segment",
246                    "segment(n: bigint): Promise<void> {\n  return run(n);\n}"
247                ),
248                ("handle", "handle = (e: Event) => log(e)"),
249                ("<arrow>", "(a, b) => a + b"),
250                (
251                    "total",
252                    "total = (xs: number[]) => xs.reduce((a, b) => a + b, 0)"
253                ),
254                ("fetchAll", "fetchAll(url: string) { return get(url); }"),
255                ("save", "save: async (x: number) => put(x)"),
256                // Not the variable's value, so not under its name.
257                ("later", "() => refresh()"),
258            ]
259        );
260        let start = &maps[0].units[0].start;
261        assert_eq!(start.line, 2);
262        assert!(src[start.offset as usize..].starts_with("segment("));
263    }
264
265    #[test]
266    fn svelte_script_functions_are_placed_in_the_host_file() {
267        let src = "<script lang=\"ts\">\n  let n = $state(0);\n\n  function double(x: number): number {\n    return x * 2;\n  }\n</script>\n\n<button onclick={() => (n = double(n))}>{n}</button>\n";
268        let maps = extract_units(src, "svelte");
269        assert_eq!(maps.len(), 1);
270        assert_eq!(maps[0].format, "typescript");
271        let unit = &maps[0].units[0];
272        assert_eq!(unit.name, "double");
273        assert_eq!((unit.start.line, unit.end.line), (4, 6));
274        assert_eq!(
275            &src[unit.start.offset as usize..unit.start.offset as usize + 8],
276            "function"
277        );
278        assert_eq!(
279            unit.text,
280            "function double(x: number): number {\n  return x * 2;\n}"
281        );
282    }
283
284    #[test]
285    fn vue_blocks_of_one_format_share_a_map() {
286        let src = "<template><p>{{ a }}</p></template>\n<script>\nexport function one() { return 1; }\n</script>\n<script setup>\nfunction two() { return 2; }\n</script>\n";
287        let maps = extract_units(src, "vue");
288        assert_eq!(maps.len(), 1);
289        assert_eq!(maps[0].format, "javascript");
290        let names: Vec<&str> = maps[0].units.iter().map(|u| u.name.as_str()).collect();
291        assert_eq!(names, vec!["one", "two"]);
292        assert_eq!(maps[0].units[1].start.line, 6);
293    }
294
295    #[test]
296    fn formats_without_an_extractor_have_no_units() {
297        assert!(supports_units("rust") && supports_units("svelte") && supports_units("tsx"));
298        assert!(supports_units("python") && supports_units("ruby") && supports_units("go"));
299        assert!(!supports_units("haskell") && !supports_units("markdown"));
300        assert!(extract_units("f x = x + 1\n", "haskell").is_empty());
301        assert!(extract_units("<style>p { color: red }</style>", "svelte").is_empty());
302    }
303
304    #[test]
305    fn crlf_sources_give_the_same_text() {
306        let lf = "pub fn alpha(x: u32) -> u32 {\n    let y = x + 1; // one\n    y * 2\n}\n";
307        let crlf = lf.replace('\n', "\r\n");
308        let text = |src: &str| extract_units(src, "rust")[0].units[0].text.clone();
309        assert_eq!(text(&crlf), text(lf));
310        assert_eq!(
311            text(lf),
312            "pub fn alpha(x: u32) -> u32 {\n    let y = x + 1;\n    y * 2\n}"
313        );
314    }
315
316    #[test]
317    fn python_units_drop_comments_but_keep_docstrings() {
318        let src = "class Cart:\n    def total(self):\n        \"\"\"Sum of the lines.\"\"\"\n        # cents\n        return sum(l.price * l.qty for l in self.lines)  # all\n";
319        let unit = &extract_units(src, "python")[0].units[0];
320        assert_eq!(unit.grammar, "python");
321        assert_eq!(
322            unit.text,
323            "def total(self):\n    \"\"\"Sum of the lines.\"\"\"\n    return sum(l.price * l.qty for l in self.lines)"
324        );
325    }
326
327    #[test]
328    fn grammar_languages_give_units_without_comments() {
329        let src = "package cart\n\n// Total sums the lines.\nfunc (c *Cart) Total() int {\n\t// cents\n\tsum := 0\n\tfor _, l := range c.lines {\n\t\tsum += l.price * l.qty // each\n\t}\n\treturn sum\n}\n";
330        let maps = extract_units(src, "go");
331        assert_eq!(maps[0].format, "go");
332        let unit = &maps[0].units[0];
333        assert_eq!((unit.name.as_str(), unit.grammar), ("Total", "go"));
334        assert_eq!((unit.start.line, unit.end.line), (4, 11));
335        assert_eq!(
336            unit.text,
337            "func (c *Cart) Total() int {\n\tsum := 0\n\tfor _, l := range c.lines {\n\t\tsum += l.price * l.qty\n\t}\n\treturn sum\n}"
338        );
339    }
340
341    #[test]
342    fn dedent_strips_the_first_line_indentation() {
343        assert_eq!(
344            dedent(
345                "fn f() {\n        if x {\n            y\n        }\n    }",
346                4
347            ),
348            "fn f() {\n    if x {\n        y\n    }\n}"
349        );
350        assert_eq!(dedent("one line", 8), "one line");
351        assert_eq!(
352            dedent("def f():\n  x\n\ty", 4),
353            "def f():\nx\ny",
354            "never more than the line has"
355        );
356        assert_eq!(
357            dedent("f() {\n\u{a0}\u{a0}x\n\u{3000}y\n}", 1),
358            "f() {\n\u{a0}x\ny\n}",
359            "whitespace wider than a byte is stripped by character"
360        );
361    }
362}