mod grammars;
use cpd_tokenizer::functions::{FunctionExtractor, RawFunction};
use cpd_tokenizer::line_index::LineIndex;
static EXTRACTORS: &[&dyn FunctionExtractor] = &[
&RustExtractor,
&PythonExtractor,
&grammars::C,
&grammars::CPP,
&grammars::CSHARP,
&grammars::GO,
&grammars::JAVA,
&grammars::KOTLIN,
&grammars::PHP,
&grammars::RUBY,
&grammars::SCALA,
&grammars::SWIFT,
];
pub fn extractor_for(format: &str) -> Option<&'static dyn FunctionExtractor> {
cpd_tokenizer::functions::extractor_for(format).or_else(|| {
EXTRACTORS
.iter()
.copied()
.find(|e| e.formats().contains(&format))
})
}
pub fn extract_functions(source: &str, format: &str) -> Vec<RawFunction> {
cpd_tokenizer::functions::extract_with(extractor_for(format), source, format)
}
pub struct PythonExtractor;
impl FunctionExtractor for PythonExtractor {
fn grammar(&self) -> &'static str {
"python"
}
fn formats(&self) -> &'static [&'static str] {
&["python"]
}
fn extract(&self, source: &str, _format: &str) -> Vec<RawFunction> {
use ruff_python_ast::visitor::source_order::SourceOrderVisitor;
let Ok(parsed) = ruff_python_parser::parse_module(source) else {
return Vec::new();
};
let line_index = LineIndex::new(source.as_bytes());
let mut visitor = PythonFunctions {
source,
line_index: &line_index,
out: Vec::new(),
};
visitor.visit_body(&parsed.syntax().body);
visitor.out.sort_by_key(|f| f.start.offset);
visitor.out
}
}
struct PythonFunctions<'s> {
source: &'s str,
line_index: &'s LineIndex,
out: Vec<RawFunction>,
}
impl<'a> ruff_python_ast::visitor::source_order::SourceOrderVisitor<'a> for PythonFunctions<'_> {
fn visit_stmt(&mut self, stmt: &'a ruff_python_ast::Stmt) {
if let ruff_python_ast::Stmt::FunctionDef(f) = stmt {
let name_start = f.name.range.start().to_usize();
let end = f.range.end().to_usize();
let head = &self.source[..name_start];
if let Some(def) = head.rfind("def") {
let before = head[..def].trim_end();
let start = match f.is_async && before.ends_with("async") {
true => before.len() - "async".len(),
false => def,
};
self.out.push(RawFunction {
grammar: "python",
name: f.name.to_string(),
start: self.line_index.location(start),
end: self.line_index.location(end),
head: self.line_index.location(start),
kinds: Vec::new(),
});
}
}
ruff_python_ast::visitor::source_order::walk_stmt(self, stmt);
}
}
pub struct RustExtractor;
impl FunctionExtractor for RustExtractor {
fn grammar(&self) -> &'static str {
"rust"
}
fn formats(&self) -> &'static [&'static str] {
&["rust"]
}
fn extract(&self, source: &str, _format: &str) -> Vec<RawFunction> {
let line_index = LineIndex::new(source.as_bytes());
let mut out: Vec<RawFunction> = scan_rust_functions(source)
.into_iter()
.map(|(name, start, end)| RawFunction {
grammar: "rust",
name,
start: line_index.location(start),
end: line_index.location(end),
head: line_index.location(start),
kinds: Vec::new(),
})
.collect();
out.sort_by_key(|f| f.start.offset);
out
}
}
fn scan_rust_functions(source: &str) -> Vec<(String, usize, usize)> {
let b = source.as_bytes();
let mut out = Vec::new();
let mut open: Vec<(String, usize, usize)> = Vec::new();
let mut depth = 0usize;
let mut item_start: Option<usize> = None;
let mut pending: Option<(String, usize)> = None;
let mut nesting = 0usize;
let mut i = 0;
while i < b.len() {
let c = b[i];
if c.is_ascii_whitespace() {
i += 1;
continue;
}
if let Some(next) = skip_rust_trivia(b, i) {
i = next;
continue;
}
if item_start.is_none() {
item_start = Some(i);
}
if let Some(next) = skip_rust_literal(b, i) {
i = next;
continue;
}
if c.is_ascii_alphabetic() || c == b'_' || c >= 0x80 {
let word_end = word_end(b, i);
if &b[i..word_end] == b"fn" && pending.is_none() {
let name_start = skip_space_and_trivia(b, word_end);
let name_end = word_end_if_ident(b, name_start);
if name_end > name_start {
let name = source[name_start..name_end].to_string();
pending = Some((name, item_start.unwrap_or(i)));
nesting = 0;
i = name_end;
continue;
}
}
i = word_end;
continue;
}
match c {
b'(' | b'[' if pending.is_some() => nesting += 1,
b')' | b']' if pending.is_some() => nesting = nesting.saturating_sub(1),
b';' if pending.is_some() && nesting == 0 => {
pending = None;
item_start = None;
}
b'{' => {
depth += 1;
if nesting == 0
&& let Some((name, start)) = pending.take()
{
open.push((name, start, depth));
}
if pending.is_none() {
item_start = None;
}
}
b'}' => {
if let Some((_, _, d)) = open.last()
&& *d == depth
{
let (name, start, _) = open.pop().unwrap_or_default();
out.push((name, start, i + 1));
}
depth = depth.saturating_sub(1);
if pending.is_none() {
item_start = None;
}
}
b';' | b']' if pending.is_none() => item_start = None,
_ => {}
}
i += 1;
}
out
}
fn skip_rust_trivia(b: &[u8], i: usize) -> Option<usize> {
if b[i] != b'/' {
return None;
}
match b.get(i + 1) {
Some(b'/') => Some(
b[i..]
.iter()
.position(|&c| c == b'\n')
.map_or(b.len(), |p| i + p + 1),
),
Some(b'*') => {
let mut level = 0usize;
let mut j = i;
while j < b.len() {
if b[j] == b'/' && b.get(j + 1) == Some(&b'*') {
level += 1;
j += 2;
} else if b[j] == b'*' && b.get(j + 1) == Some(&b'/') {
level -= 1;
j += 2;
if level == 0 {
return Some(j);
}
} else {
j += 1;
}
}
Some(b.len())
}
_ => None,
}
}
fn skip_rust_literal(b: &[u8], i: usize) -> Option<usize> {
let mut j = i;
while j < b.len() && matches!(b[j], b'b' | b'r' | b'c') && j - i < 2 {
j += 1;
}
let raw = b[i..j].contains(&b'r');
let mut hashes = 0;
if raw {
while b.get(j) == Some(&b'#') {
hashes += 1;
j += 1;
}
}
match b.get(j) {
Some(b'"') if j == i || raw || b[i..j].iter().all(|&p| p == b'b' || p == b'c') => {
j += 1;
while j < b.len() {
if !raw && b[j] == b'\\' {
j += 2;
continue;
}
if b[j] == b'"'
&& b.len() >= j + 1 + hashes
&& b[j + 1..j + 1 + hashes].iter().all(|&h| h == b'#')
{
return Some(j + 1 + hashes);
}
j += 1;
}
Some(b.len())
}
Some(b'\'') if !raw && (j == i || b[i..j] == *b"b") => {
let k = j + 1;
if b.get(k) == Some(&b'\\') {
let close = b[k..].iter().position(|&c| c == b'\'').map(|p| k + p);
let close = match close {
Some(p) if p == k + 1 => b[k + 2..]
.iter()
.position(|&c| c == b'\'')
.map(|q| k + 2 + q),
other => other,
};
return close.map(|p| p + 1);
}
let ch_len = utf8_len(b.get(k).copied()?);
(b.get(k + ch_len) == Some(&b'\'')).then_some(k + ch_len + 1)
}
_ => None,
}
}
fn utf8_len(first: u8) -> usize {
match first {
0xF0..=0xFF => 4,
0xE0..=0xEF => 3,
0xC0..=0xDF => 2,
_ => 1,
}
}
fn word_end(b: &[u8], i: usize) -> usize {
let mut j = i;
while j < b.len() && (b[j].is_ascii_alphanumeric() || b[j] == b'_' || b[j] >= 0x80) {
j += 1;
}
j
}
fn word_end_if_ident(b: &[u8], i: usize) -> usize {
match b.get(i) {
Some(c) if c.is_ascii_alphabetic() || *c == b'_' || *c >= 0x80 => {
if b[i] == b'r' && b.get(i + 1) == Some(&b'#') {
word_end(b, i + 2)
} else {
word_end(b, i)
}
}
_ => i,
}
}
fn skip_space_and_trivia(b: &[u8], mut i: usize) -> usize {
loop {
while i < b.len() && b[i].is_ascii_whitespace() {
i += 1;
}
match (i < b.len()).then(|| skip_rust_trivia(b, i)).flatten() {
Some(next) => i = next,
None => return i,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use cpd_tokenizer::functions::supports_functions;
fn spans(fns: &[RawFunction]) -> Vec<(&str, u32, u32)> {
fns.iter()
.map(|f| (f.name.as_str(), f.start.line, f.end.line))
.collect()
}
#[test]
fn rust_functions_methods_and_default_trait_methods_are_found() {
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";
let fns = extract_functions(src, "rust");
assert_eq!(
spans(&fns),
vec![
("add", 4, 6),
("total", 9, 12),
("cents", 10, 10),
("name", 16, 16)
]
);
let add = &fns[0];
assert_eq!(
&src[add.start.offset as usize..add.end.offset as usize],
"pub fn add(a: i32, b: i32) -> i32 {\n a + b\n}"
);
assert!(
fns.iter()
.all(|f| f.grammar == "rust" && f.kinds.is_empty())
);
}
#[test]
fn rust_braces_and_fn_inside_literals_and_comments_are_not_code() {
let src = r####"fn lifetimes<'a>(x: &'a str) -> &'a str { if x == "{" { x } else { "}" } }
fn chars() -> [char; 3] { ['{', '\'', '}'] }
fn raw() -> &'static str { r##"fn fake() { "# }"## }
/* outer /* fn nested() { */ still comment */
fn pointer(f: fn(u32) -> u32) -> u32 { f(1) }
macro_rules! make { ($n:ident) => { fn $n() {} }; }
fn last() {}
"####;
let names: Vec<(String, u32)> = extract_functions(src, "rust")
.into_iter()
.map(|f| (f.name, f.end.line))
.collect();
let expected = [
("lifetimes", 1),
("chars", 2),
("raw", 3),
("pointer", 5),
("last", 7),
];
assert_eq!(
names,
expected
.iter()
.map(|(n, l)| (n.to_string(), *l))
.collect::<Vec<_>>()
);
}
#[test]
fn python_functions_methods_and_nested_ones_start_at_def() {
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";
let fns = extract_functions(src, "python");
assert_eq!(
spans(&fns),
vec![("load", 4, 5), ("total", 8, 11), ("cents", 9, 10)]
);
assert_eq!(&src[fns[0].start.offset as usize..][..8], "def load");
assert_eq!(&src[fns[1].start.offset as usize..][..9], "async def");
assert!(
fns.iter()
.all(|f| f.grammar == "python" && f.kinds.is_empty())
);
assert!(!supports_functions("python"));
assert!(extract_functions("def broken(:\n", "python").is_empty());
}
#[test]
fn rust_that_does_not_parse_still_yields_what_closes() {
assert!(extract_functions("fn broken( {", "rust").is_empty());
let fns = extract_functions("fn ok() { 1 }\nfn open() {", "rust");
assert_eq!(fns.len(), 1);
assert_eq!(fns[0].name, "ok");
}
#[test]
fn similarity_does_not_compare_what_only_semantic_reads() {
for format in ["rust", "python", "go"] {
assert!(extractor_for(format).is_some(), "{format}");
assert!(!supports_functions(format), "{format}");
}
assert_eq!(extractor_for("typescript").unwrap().grammar(), "oxc");
assert!(extractor_for("haskell").is_none());
}
}