1mod grammars;
12
13use cpd_tokenizer::functions::{FunctionExtractor, RawFunction};
14use cpd_tokenizer::line_index::LineIndex;
15
16static 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
32pub 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
43pub fn extract_functions(source: &str, format: &str) -> Vec<RawFunction> {
46 cpd_tokenizer::functions::extract_with(extractor_for(format), source, format)
47}
48
49pub 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 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
113pub 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
153fn scan_rust_functions(source: &str) -> Vec<(String, usize, usize)> {
156 let b = source.as_bytes();
157 let mut out = Vec::new();
158 let mut open: Vec<(String, usize, usize)> = Vec::new();
160 let mut depth = 0usize;
161 let mut item_start: Option<usize> = None;
164 let mut pending: Option<(String, usize)> = None;
166 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 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
241fn 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 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
277fn skip_rust_literal(b: &[u8], i: usize) -> Option<usize> {
280 let mut j = i;
281 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 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 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
350fn 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 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 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}