use std::cell::RefCell;
use std::sync::OnceLock;
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use streaming_iterator::StreamingIterator;
use tree_sitter::{Node, Parser, Query, QueryCursor};
use crate::lang::Lang;
const SIG_CAP: usize = 100;
const TERM_CAP: usize = 64;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum SymKind {
Fn,
Method,
Class,
Struct,
Enum,
Trait,
Interface,
Type,
Mod,
Const,
Var,
}
impl SymKind {
fn from_capture(suffix: &str) -> Option<SymKind> {
Some(match suffix {
"fn" => SymKind::Fn,
"method" => SymKind::Method,
"class" => SymKind::Class,
"struct" => SymKind::Struct,
"enum" => SymKind::Enum,
"trait" => SymKind::Trait,
"interface" => SymKind::Interface,
"type" => SymKind::Type,
"mod" => SymKind::Mod,
"const" => SymKind::Const,
"var" => SymKind::Var,
_ => return None,
})
}
pub fn name(self) -> &'static str {
match self {
SymKind::Fn => "fn",
SymKind::Method => "method",
SymKind::Class => "class",
SymKind::Struct => "struct",
SymKind::Enum => "enum",
SymKind::Trait => "trait",
SymKind::Interface => "interface",
SymKind::Type => "type",
SymKind::Mod => "mod",
SymKind::Const => "const",
SymKind::Var => "var",
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum Vis {
Pub,
Priv,
}
impl Vis {
pub fn name(self) -> &'static str {
match self {
Vis::Pub => "pub",
Vis::Priv => "priv",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct Symbol {
pub line: u32,
pub end_line: u32,
pub name: String,
pub kind: SymKind,
pub vis: Vis,
pub sig: String,
pub terms: Vec<u32>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum RefKind {
Call,
Import,
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct RefName {
pub line: u32,
pub name: String,
pub kind: RefKind,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct Extraction {
pub defs: Vec<Symbol>,
pub refs: Vec<RefName>,
}
static QUERIES: [OnceLock<Option<Query>>; Lang::ALL.len()] =
[const { OnceLock::new() }; Lang::ALL.len()];
fn query_for(lang: Lang) -> Option<&'static Query> {
let idx = Lang::ALL.iter().position(|&l| l == lang)?;
QUERIES[idx]
.get_or_init(|| Query::new(&lang.language(), lang.query_source()).ok())
.as_ref()
}
thread_local! {
static TS: RefCell<(Parser, Option<Lang>, Vec<QueryCursor>)> =
RefCell::new((Parser::new(), None, Vec::new()));
}
pub fn extract(lang: Lang, src: &str) -> Extraction {
extract_with_timeout(lang, src, Duration::from_millis(500))
}
fn extract_with_timeout(lang: Lang, src: &str, timeout: Duration) -> Extraction {
let Some(query) = query_for(lang) else {
return Extraction::default();
};
let (tree, mut cursor) = match TS.with(|cell| {
let (parser, current, cursors) = &mut *cell.borrow_mut();
if *current != Some(lang) {
if parser.set_language(&lang.language()).is_err() {
return None;
}
*current = Some(lang);
}
let deadline = Instant::now() + timeout;
let bytes = src.as_bytes();
let mut chunker = |offset: usize, _pos: tree_sitter::Point| -> &[u8] {
&bytes[offset.min(bytes.len())..]
};
let mut progress = |_state: &tree_sitter::ParseState| {
if std::time::Instant::now() > deadline {
std::ops::ControlFlow::Break(())
} else {
std::ops::ControlFlow::Continue(())
}
};
let options = tree_sitter::ParseOptions::new().progress_callback(&mut progress);
let tree = match parser.parse_with_options(&mut chunker, None, Some(options)) {
Some(tree) => tree,
None => {
parser.reset();
return None;
}
};
Some((tree, cursors.pop().unwrap_or_default()))
}) {
Some(x) => x,
None => return Extraction::default(),
};
let mut defs: Vec<Symbol> = Vec::new();
let mut refs: Vec<RefName> = Vec::new();
let mut exported: std::collections::BTreeSet<String> = Default::default();
let mut matches = cursor.matches(query, tree.root_node(), src.as_bytes());
while let Some(m) = matches.next() {
let mut def_node: Option<(Node, SymKind)> = None;
let mut name_node: Option<Node> = None;
let mut span_node: Option<Node> = None;
for cap in m.captures {
let cap_name = &query.capture_names()[cap.index as usize];
if let Some(suffix) = cap_name.strip_prefix("def.") {
if let Some(kind) = SymKind::from_capture(suffix) {
def_node = Some((cap.node, kind));
}
} else if *cap_name == "name" {
name_node = Some(cap.node);
} else if *cap_name == "span" {
span_node = Some(cap.node);
} else if *cap_name == "export" {
let name = node_text(cap.node, src);
if !name.is_empty() {
exported.insert(name);
}
} else if *cap_name == "ref" || *cap_name == "ref.import" {
let name = node_text(cap.node, src);
if !name.is_empty() {
refs.push(RefName {
line: cap.node.start_position().row as u32 + 1,
name,
kind: if *cap_name == "ref.import" {
RefKind::Import
} else {
RefKind::Call
},
});
}
}
}
if let (Some((node, kind)), Some(name_node)) = (def_node, name_node) {
let name = node_text(name_node, src);
if name.is_empty() {
continue;
}
let span = span_node.unwrap_or(node);
let vis = visibility(lang, node, src, &name);
defs.push(Symbol {
line: span.start_position().row as u32 + 1,
end_line: span.end_position().row as u32 + 1,
sig: signature(node, src),
terms: lexical_terms(node, src),
name,
kind,
vis,
});
}
}
if !exported.is_empty() {
for d in &mut defs {
if d.vis == Vis::Priv && exported.contains(&d.name) {
d.vis = Vis::Pub;
}
}
}
defs.sort();
defs.dedup();
refs.sort();
refs.dedup();
drop(matches);
TS.with(|cell| cell.borrow_mut().2.push(cursor)); Extraction { defs, refs }
}
fn node_text(node: Node, src: &str) -> String {
node_slice(node, src).unwrap_or_default().to_string()
}
fn node_slice<'a>(node: Node, src: &'a str) -> Option<&'a str> {
src.get(node.start_byte()..node.end_byte())
}
fn lexical_terms(node: Node, src: &str) -> Vec<u32> {
let text = node_slice(node, src).unwrap_or_default();
let mut ranked: std::collections::BTreeMap<u32, usize> = std::collections::BTreeMap::new();
let mut add = |text: &str, priority: usize| {
for chunk in text.split(|character: char| {
!character.is_alphanumeric() && character != '_' && character != '-'
}) {
for word in crate::routes::split_ident(chunk) {
if is_lexical_noise(&word) {
continue;
}
if let Some(fingerprint) = crate::routes::term_fingerprint(&word) {
let score = priority + word.len();
ranked
.entry(fingerprint)
.and_modify(|existing| *existing = (*existing).max(score))
.or_insert(score);
}
}
}
};
add(text, 0);
if let Some(comments) = leading_comments(node, src) {
add(comments, 100);
}
let mut ranked: Vec<(usize, u32)> = ranked
.into_iter()
.map(|(fingerprint, length)| (length, fingerprint))
.collect();
ranked.sort_by(|left, right| right.0.cmp(&left.0).then_with(|| left.1.cmp(&right.1)));
let mut terms: Vec<u32> = ranked
.into_iter()
.take(TERM_CAP)
.map(|(_, fingerprint)| fingerprint)
.collect();
terms.sort_unstable();
terms
}
fn leading_comments<'a>(node: Node, src: &'a str) -> Option<&'a str> {
let prefix = src.get(..node.start_byte())?;
let mut start = prefix.len();
let mut kept = 0usize;
for line in prefix.lines().rev() {
let trimmed = line.trim();
if trimmed.is_empty() {
if kept == 0 {
continue;
}
break;
}
if !trimmed.starts_with("//")
&& !trimmed.starts_with('#')
&& !trimmed.starts_with("/*")
&& !trimmed.starts_with('*')
&& !trimmed.starts_with("--")
{
break;
}
start = start.saturating_sub(line.len());
if start > 0 && prefix.as_bytes().get(start - 1) == Some(&b'\n') {
start -= 1;
}
kept += 1;
if kept == 8 {
break;
}
}
(kept > 0).then(|| prefix.get(start..).unwrap_or_default())
}
fn is_lexical_noise(word: &str) -> bool {
matches!(
word,
"async"
| "await"
| "bool"
| "class"
| "const"
| "crate"
| "default"
| "else"
| "false"
| "function"
| "impl"
| "interface"
| "into"
| "none"
| "option"
| "public"
| "return"
| "self"
| "some"
| "string"
| "struct"
| "super"
| "this"
| "true"
)
}
fn signature(node: Node, src: &str) -> String {
let text = node_slice(node, src).unwrap_or_default();
let first_line = text.lines().next().unwrap_or_default();
let collapsed: String = first_line.split_whitespace().collect::<Vec<_>>().join(" ");
let brace_free = collapsed.split('{').next().unwrap_or(&collapsed);
let trimmed = brace_free.trim_end_matches(':').trim_end();
canonicalize_signature_order(trimmed)
}
fn canonicalize_signature_order(sig: &str) -> String {
let mut tokens: Vec<&str> = sig.split_whitespace().collect();
if tokens.len() >= 2 {
let vis_idx = tokens
.iter()
.enumerate()
.skip(1)
.take(3)
.find_map(|(idx, tok)| is_visibility_token(tok).then_some(idx));
if let Some(i) = vis_idx
&& i > 0
&& !is_visibility_token(tokens[0])
{
tokens.swap(0, i);
}
}
let normalized = tokens.join(" ");
if normalized.chars().count() > SIG_CAP {
let cut: String = normalized.chars().take(SIG_CAP - 1).collect();
format!("{cut}\u{2026}")
} else {
normalized
}
}
fn is_visibility_token(token: &str) -> bool {
matches!(
token,
"pub" | "pub(crate)" | "pub(super)" | "public" | "private" | "protected" | "internal"
) || token.starts_with("pub(")
}
fn visibility(lang: Lang, node: Node, src: &str, name: &str) -> Vis {
match lang {
Lang::Python => {
if name.starts_with('_') {
Vis::Priv
} else {
Vis::Pub
}
}
Lang::Go => {
if name.chars().next().is_some_and(|c| c.is_uppercase()) {
Vis::Pub
} else {
Vis::Priv
}
}
Lang::Rust => {
let target = ancestor_of_kind(node, "trait_item").unwrap_or(node);
if visibility_modifier_text(target, src).is_some_and(|t| t == "pub") {
Vis::Pub
} else {
Vis::Priv
}
}
Lang::JavaScript | Lang::TypeScript | Lang::Tsx => {
if has_ancestor_of_kind(node, "export_statement") {
Vis::Pub
} else {
Vis::Priv
}
}
Lang::Java | Lang::CSharp => {
if has_ancestor_of_kind(node, "interface_declaration")
|| has_modifier(node, src, "public")
{
Vis::Pub
} else {
Vis::Priv
}
}
Lang::C | Lang::Cpp => {
if has_modifier(node, src, "static") {
Vis::Priv
} else {
Vis::Pub
}
}
Lang::Php => {
if has_modifier(node, src, "private") || has_modifier(node, src, "protected") {
Vis::Priv
} else {
Vis::Pub
}
}
Lang::Apex => {
if node.kind() == "trigger_declaration"
|| has_ancestor_of_kind(node, "interface_declaration")
|| has_modifier(node, src, "public")
|| has_modifier(node, src, "global")
{
Vis::Pub
} else {
Vis::Priv
}
}
Lang::Kotlin => {
if has_modifier(node, src, "private")
|| has_modifier(node, src, "protected")
|| has_modifier(node, src, "internal")
{
Vis::Priv
} else {
Vis::Pub
}
}
Lang::Lua => {
if has_modifier(node, src, "local") {
Vis::Priv
} else {
Vis::Pub
}
}
Lang::Ruby | Lang::Bash | Lang::Html => Vis::Pub,
}
}
fn visibility_modifier_text<'a>(node: Node, src: &'a str) -> Option<&'a str> {
let mut cursor = node.walk();
let child = node
.named_children(&mut cursor)
.find(|c| c.kind() == "visibility_modifier")?;
node_slice(child, src).map(str::trim)
}
fn ancestor_of_kind<'a>(node: Node<'a>, kind: &str) -> Option<Node<'a>> {
let mut cur = node.parent();
while let Some(n) = cur {
if n.kind() == kind {
return Some(n);
}
cur = n.parent();
}
None
}
fn has_ancestor_of_kind(node: Node, kind: &str) -> bool {
ancestor_of_kind(node, kind).is_some()
}
fn has_modifier(node: Node, src: &str, keyword: &str) -> bool {
const WRAPPERS: [&str; 4] = [
"modifiers",
"declaration_specifiers",
"visibility_modifier",
"modifier",
];
const MODIFIER_LEAVES: [&str; 3] =
["storage_class_specifier", "visibility_modifier", "modifier"];
let matches = |n: Node| {
n.kind() == keyword
|| (MODIFIER_LEAVES.contains(&n.kind())
&& node_slice(n, src).is_some_and(|t| t.trim() == keyword))
};
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
if matches(child) {
return true;
}
if WRAPPERS.contains(&child.kind()) {
let mut inner = child.walk();
if child.children(&mut inner).any(matches) {
return true;
}
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn queries_compile_for_all_langs() {
for lang in Lang::ALL {
let compiled = Query::new(&lang.language(), lang.query_source());
assert!(
compiled.is_ok(),
"query for {} failed: {:?}",
lang.name(),
compiled.err()
);
}
}
#[test]
fn signature_is_capped_and_collapsed() {
let long = format!("fn {}(a: u32) -> u64 {{", "x".repeat(200));
let src = format!("pub {long} 0 }}");
let out = extract(Lang::Rust, &src);
assert_eq!(out.defs.len(), 1);
assert!(out.defs[0].sig.chars().count() <= SIG_CAP);
assert!(!out.defs[0].sig.contains(" "), "whitespace collapsed");
}
#[test]
fn cancelled_parse_is_reset_before_the_next_source() {
let pathological = "(".repeat(1_000_000);
let _ = extract_with_timeout(Lang::C, &pathological, Duration::ZERO);
let out = extract(Lang::C, "int recovered(void) { return 1; }");
assert!(
out.defs
.iter()
.any(|definition| definition.name == "recovered")
);
}
#[test]
fn out_of_bounds_node_text_degrades_to_empty() {
let mut parser = Parser::new();
parser.set_language(&Lang::Rust.language()).unwrap();
let source = "pub fn valid() {}";
let tree = parser.parse(source, None).unwrap();
assert_eq!(node_text(tree.root_node(), "x"), "");
assert_eq!(signature(tree.root_node(), "x"), "");
}
#[test]
fn canonicalize_signature_order_keeps_visibility_first() {
assert_eq!(
canonicalize_signature_order("class public GeanWasThere"),
"public class GeanWasThere"
);
assert_eq!(
canonicalize_signature_order("void private run()"),
"private void run()"
);
assert_eq!(
canonicalize_signature_order("public class GeanWasThere"),
"public class GeanWasThere"
);
}
#[test]
fn extraction_keeps_compact_body_evidence() {
let source = r#"
/// Accept compatibility framing from older clients.
fn read_message(first: &str) {
if first.strip_prefix("Content-Length").is_some() {
let parsed = serde_json::from_str(first);
println!("newline message: {parsed:?}");
}
}
"#;
let extraction = extract(Lang::Rust, source);
let terms = &extraction.defs[0].terms;
for word in ["compatibility", "content", "length", "json", "newline"] {
let fingerprint = crate::routes::term_fingerprint(word).expect("fingerprint");
assert!(
terms.binary_search(&fingerprint).is_ok(),
"missing {word} body evidence"
);
}
}
#[test]
fn apex_trigger_is_a_public_entry_point() {
let extraction = extract(
Lang::Apex,
"trigger AccountTrigger on Account (before insert) {\n\
AccountService.handle(Trigger.new);\n\
}\n",
);
assert_eq!(extraction.defs.len(), 1);
assert_eq!(extraction.defs[0].name, "AccountTrigger");
assert_eq!(extraction.defs[0].kind, SymKind::Class);
assert_eq!(extraction.defs[0].vis, Vis::Pub);
assert!(
extraction
.refs
.iter()
.any(|reference| reference.name == "handle"),
"trigger call is indexed"
);
}
#[test]
fn lua_local_global_and_module_surface_visibility() {
let extraction = extract(
Lang::Lua,
"local function helper() end\nfunction M.setup() end\nfunction M:render() end\n",
);
let helper = extraction.defs.iter().find(|d| d.name == "helper").unwrap();
assert_eq!(helper.vis, Vis::Priv);
let setup = extraction.defs.iter().find(|d| d.name == "setup").unwrap();
assert_eq!(setup.vis, Vis::Pub);
let render = extraction.defs.iter().find(|d| d.name == "render").unwrap();
assert_eq!(render.vis, Vis::Pub);
assert_eq!(render.kind, SymKind::Method);
}
#[test]
fn lua_multi_assignment_does_not_index_the_non_function_binding() {
let extraction = extract(Lang::Lua, "local a, b = 1, function() return 2 end\n");
assert!(!extraction.defs.iter().any(|d| d.name == "a"));
assert!(!extraction.defs.iter().any(|d| d.name == "b"));
}
#[test]
fn lua_nested_same_signature_functions_with_different_ranges_survive() {
let extraction = extract(
Lang::Lua,
"local function same(value)\n local function same(value)\n return value\n end\n return value\nend\n",
);
let same: Vec<_> = extraction
.defs
.iter()
.filter(|d| d.name == "same")
.collect();
assert_eq!(same.len(), 2);
assert_ne!(same[0].line, same[1].line);
assert_ne!(same[0].end_line, same[1].end_line);
}
}