use std::collections::BTreeSet;
use std::io::IsTerminal;
use std::path::Path;
use crate::error::CoraError;
use glob::Pattern;
use ignore::WalkBuilder;
use indicatif::{ProgressBar, ProgressDrawTarget, ProgressStyle};
use tracing::debug;
#[derive(Debug, Clone)]
pub struct FileEntry {
pub path: String,
pub content: String,
pub lines: usize,
}
const DEFAULT_EXTENSIONS: &[&str] = &[
"rs", "py", "js", "ts", "tsx", "jsx", "go", "java", "kt", "rb", "c", "cpp", "h", "hpp", "cs",
"php", "swift", "scala", "vue", "svelte", "sh", "bash", "zsh", "ps1", "toml", "yaml", "yml",
"json", "sql", "graphql", "proto", "md", "rst", "txt", "html", "css", "scss", "less",
];
#[allow(clippy::unnecessary_wraps, clippy::format_push_string)]
pub fn walk_project(
root: &Path,
include_patterns: &[String],
exclude_patterns: &[String],
extra_extensions: &[String],
) -> std::result::Result<Vec<FileEntry>, CoraError> {
debug!(
root = %root.display(),
"walking project directory"
);
let mut extensions: BTreeSet<String> = DEFAULT_EXTENSIONS
.iter()
.map(std::string::ToString::to_string)
.collect();
for ext in extra_extensions {
extensions.insert(ext.trim_start_matches('.').to_lowercase());
}
let include_globs: Vec<Pattern> = include_patterns
.iter()
.filter_map(|p| Pattern::new(p).ok())
.collect();
let exclude_globs: Vec<Pattern> = exclude_patterns
.iter()
.filter_map(|p| Pattern::new(p).ok())
.collect();
let mut entries = Vec::new();
let spinner = ProgressBar::new_spinner();
if std::io::stderr().is_terminal() {
spinner.enable_steady_tick(std::time::Duration::from_millis(80));
spinner.set_style(
ProgressStyle::with_template("{spinner:.cyan} {msg}")
.unwrap()
.tick_chars("⠁⠂⠄⡀⢀⠠⠐⠈ "),
);
spinner.set_message("Scanning files…");
} else {
spinner.set_draw_target(ProgressDrawTarget::hidden());
}
let walker = WalkBuilder::new(root)
.hidden(true) .git_ignore(true) .git_global(true) .git_exclude(true) .require_git(false) .build();
for result in walker {
let entry = match result {
Ok(e) => e,
Err(err) => {
debug!(error = %err, "error during directory walk");
continue;
}
};
let path = entry.path();
if path.is_dir() {
continue;
}
let relative = path
.strip_prefix(root)
.unwrap_or(path)
.to_string_lossy()
.to_string();
if exclude_globs.iter().any(|g| g.matches(&relative)) {
continue;
}
let has_include = !include_globs.is_empty();
if has_include && !include_globs.iter().any(|g| g.matches(&relative)) {
continue;
}
let has_extension = path
.extension()
.and_then(|e| e.to_str())
.map(|e| extensions.contains(&e.to_lowercase()));
if has_extension == Some(false) {
continue;
}
let content = match std::fs::read_to_string(path) {
Ok(c) => c,
Err(e) => {
debug!(file = %relative, error = %e, "skipping unreadable file");
continue;
}
};
if content.trim().is_empty() {
continue;
}
if content.len() > 200_000 {
debug!(file = %relative, "skipping large file");
continue;
}
let lines = content.lines().count();
entries.push(FileEntry {
path: relative,
content,
lines,
});
}
spinner.finish_and_clear();
debug!(files = entries.len(), "found scannable files");
Ok(entries)
}
pub fn batch_files(files: &[FileEntry], max_chars: usize, max_files: usize) -> Vec<Vec<FileEntry>> {
let mut batches = Vec::new();
let mut current_batch = Vec::new();
let mut current_size: usize = 0;
for file in files {
let file_size = file.content.len() + file.path.len() + 20;
if (current_batch.len() >= max_files || current_size + file_size > max_chars)
&& !current_batch.is_empty()
{
batches.push(std::mem::take(&mut current_batch));
current_size = 0;
}
current_size += file_size;
current_batch.push(file.clone());
}
if !current_batch.is_empty() {
batches.push(current_batch);
}
batches
}
#[allow(clippy::format_push_string)]
pub fn format_batch_for_prompt(files: &[FileEntry]) -> String {
let mut output = String::new();
for file in files {
output.push_str(&format!("=== {} ===\n", file.path));
for (i, line) in file.content.lines().enumerate() {
output.push_str(&format!("{:>5} | {}\n", i + 1, line));
}
output.push('\n');
}
output
}