use super::embedded::{
EmbeddedLanguage as SearchLanguage, MAX_FILES, MAX_SCAN_BYTES, MAX_TREE_NODES, read_source,
};
use super::*;
use ast_grep_core::{
AstGrep as Root, Language, Pattern,
matcher::MatcherExt,
replacer::{Replacer, TemplateFix},
};
use ast_grep_language::SupportLang;
use std::collections::HashMap;
const MAX_QUERY_BYTES: usize = 4096;
const MAX_REPLACEMENT_BYTES: usize = 1024 * 1024;
struct Compiled {
pattern: Pattern,
rewrite: Option<TemplateFix>,
}
fn compile(args: &AstGrepArgs, lang: &SearchLanguage) -> anyhow::Result<Compiled> {
let pattern = Pattern::try_new(args.pattern.as_deref().unwrap_or_default(), lang.clone())?;
anyhow::ensure!(
!pattern.has_error(),
"invalid ast-grep pattern for {}",
lang.language
);
let rewrite = args
.rewrite
.as_deref()
.map(|text| TemplateFix::try_new(text, lang))
.transpose()?;
if let Some(rewrite) = &rewrite {
let defined = pattern.defined_vars();
anyhow::ensure!(
rewrite.used_vars().iter().all(|var| defined.contains(var)),
"rewrite references a variable not captured by the pattern"
);
}
Ok(Compiled { pattern, rewrite })
}
pub(super) fn run(
runtime: &ToolRuntime,
args: &AstGrepArgs,
resolved: &Path,
cancellation: &AgentCancellation,
timeout: Duration,
) -> anyhow::Result<ToolResult> {
anyhow::ensure!(
args.pattern
.as_ref()
.is_some_and(|p| p.len() <= MAX_QUERY_BYTES)
&& args
.rewrite
.as_ref()
.is_none_or(|r| r.len() <= MAX_QUERY_BYTES),
"ast-grep pattern and rewrite must each be at most {MAX_QUERY_BYTES} bytes"
);
let explicit = args
.language
.as_deref()
.map(str::parse::<SupportLang>)
.transpose()
.map_err(|_| anyhow::anyhow!("unsupported ast-grep language"))?;
let started = Instant::now();
let language = |language| SearchLanguage {
language,
cancellation: cancellation.clone(),
started,
timeout,
};
let mut compiled = HashMap::new();
let mut state = AstGrepRunState {
exit_code: 1,
..Default::default()
};
let limit = args.limit.unwrap_or(AST_GREP_DEFAULT_LIMIT);
let mut lines = 0;
let mut files = 0;
let mut bytes = 0;
let mut skipped = 0;
let allow_outside = runtime.ast_grep_absolute_paths
&& args
.path
.as_deref()
.is_some_and(|path| Path::new(path).is_absolute());
if let Some(lang) = explicit {
match compile(args, &language(lang)) {
Ok(pattern) => {
compiled.insert(lang, Ok(pattern));
}
Err(error) => {
cancellation.check()?;
if started.elapsed() < timeout {
return Err(error);
}
state.timed_out = true;
}
}
}
let mut visitor = |path: PathBuf| -> anyhow::Result<bool> {
cancellation.check()?;
if started.elapsed() >= timeout {
state.timed_out = true;
state.exit_code = -1;
return Ok(false);
}
files += 1;
if files > MAX_FILES || bytes >= MAX_SCAN_BYTES {
state.workset_truncated = true;
return Ok(false);
}
let Some(lang) = explicit.or_else(|| SupportLang::from_path(&path)) else {
return Ok(true);
};
let Some(path_text) = path.to_str() else {
skipped += 1;
state.workset_truncated = true;
return Ok(true);
};
let path = runtime
.resolve_existing_path(path_text, ExistingPathPolicy::ast_grep(allow_outside))?;
let Some(src) = read_source(&path, &mut bytes) else {
skipped += 1;
state.workset_truncated = true;
return Ok(true);
};
if src.contains('\0') {
return Ok(true);
}
let lang = language(lang);
let doc = match lang.parse(src, MAX_TREE_NODES) {
Ok(doc) => doc,
Err(_) => {
cancellation.check()?;
skipped += 1;
state.workset_truncated = true;
state.timed_out |= lang.stopped();
return Ok(!state.timed_out);
}
};
let matcher = compiled
.entry(lang.language)
.or_insert_with(|| compile(args, &lang).map_err(|error| error.to_string()));
cancellation.check()?;
if lang.stopped() {
state.timed_out = true;
return Ok(false);
}
let matcher = match matcher {
Ok(matcher) => matcher,
Err(_) => {
skipped += 1;
state.workset_truncated = true;
return Ok(true);
}
};
let replacement_bound =
replacement_bound(&doc.src, args.rewrite.as_deref().unwrap_or_default());
let root = Root::doc(doc);
let node = root.root();
let display_path = path
.strip_prefix(&runtime.cwd)
.unwrap_or(&path)
.display()
.to_string();
for candidate in node.dfs() {
cancellation.check()?;
if lang.stopped() {
state.timed_out = true;
return Ok(false);
}
let Some(found) = matcher.pattern.match_node(candidate) else {
continue;
};
cancellation.check()?;
if lang.stopped() {
state.timed_out = true;
return Ok(false);
}
state.exit_code = 0;
let line = found.start_pos().line() + 1;
let mut append = |text: &str| {
if lines >= limit
|| text.len() + 1 > AST_GREP_STDOUT_MAX_BYTES.saturating_sub(state.stdout.len())
{
state.workset_truncated = true;
state.stdout_truncated |= lines < limit;
return false;
}
state.stdout.push_str(text);
state.stdout.push('\n');
lines += 1;
true
};
if let Some(rewrite) = &matcher.rewrite {
if replacement_bound > MAX_REPLACEMENT_BYTES {
skipped += 1;
state.workset_truncated = true;
return Ok(true);
}
let replacement = rewrite.generate_replacement(&found);
let replacement = String::from_utf8(replacement)?;
let range = rewrite.get_replaced_range(&found, &matcher.pattern);
let removed = &node.get_doc().src[range.clone()];
if !append(&format!(
"@@ {display_path}:{line} bytes {}..{} (end exclusive; surrounding text unchanged) @@",
range.start, range.end
)) {
return Ok(false);
}
for text in removed.lines() {
if !append(&format!("-{text}")) {
return Ok(false);
}
}
for text in replacement.lines() {
if !append(&format!("+{text}")) {
return Ok(false);
}
}
} else {
for (offset, text) in found.text().lines().enumerate() {
if !append(&format!("{display_path}:{}:{text}", line + offset)) {
return Ok(false);
}
}
}
if lines >= limit {
state.workset_truncated = true;
return Ok(false);
}
}
Ok(true)
};
if resolved.is_dir() {
let status = runtime.workspace_walker.visit_files_with_deadline(
WorkspaceWalkOptions {
root: resolved,
skip_dirs: &[".git"],
cancel_interval: 32,
},
Some(cancellation),
Some(started + timeout),
&mut visitor,
)?;
state.workset_truncated |= status.walk_errors > 0 || status.entries_omitted > 0;
} else {
visitor(resolved.to_path_buf())?;
}
cancellation.check()?;
state.timed_out |= started.elapsed() >= timeout;
if !state.timed_out
&& compiled.values().all(Result::is_err)
&& let Some(error) = compiled.values().find_map(|value| value.as_ref().err())
{
return Err(anyhow::anyhow!(
"pattern/rewrite could not compile for the searched languages: {error}"
));
}
if state.timed_out {
state.exit_code = -1;
}
let mut result = ast_grep_result(
AstGrepResultInput {
stdout: &state.stdout,
exit_code: state.exit_code,
timed_out: state.timed_out,
stdout_truncated: state.stdout_truncated,
workset_truncated: state.workset_truncated,
},
args,
limit,
);
result.metadata["execution"] = json!("embedded");
result.metadata["files_skipped"] = json!(skipped);
result.metadata["partial"] = json!(state.workset_truncated || state.timed_out);
Ok(result)
}
fn replacement_bound(source: &str, template: &str) -> usize {
let occurrences = template.chars().filter(|c| *c == '$' || *c == '#').count();
let source_lines = source.bytes().filter(|b| *b == b'\n').count();
let template_lines = template.bytes().filter(|b| *b == b'\n').count();
let max_indent = source
.lines()
.map(|line| line.len() - line.trim_start_matches([' ', '\t']).len())
.max()
.unwrap_or(0);
let lines = source_lines
.saturating_mul(occurrences)
.saturating_add(template_lines);
source
.len()
.saturating_mul(occurrences)
.saturating_add(template.len())
.saturating_add(lines.saturating_mul(max_indent.saturating_add(template.len())))
}