use std::time::Duration;
use anyhow::{Context, Result, bail};
use crate::turso as turso_mod;
const ROW_LIMIT: usize = 10_000;
const BLOCKLIST: &[&str] = &[
"INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE", "REPLACE", "BEGIN", "COMMIT",
"ROLLBACK", "VACUUM", "REINDEX", "GRANT", "REVOKE", "ATTACH", "DETACH", "ANALYZE",
];
const SQL_PUNCTUATION: &[char] = &[
'(', ')', ';', ',', '.', '*', '+', '-', '/', '=', '<', '>', '!', '|', '&', '~', '\'', '"', '[',
']', '{', '}', ':',
];
pub async fn run_debug() -> Result<()> {
let args: Vec<String> = std::env::args().collect();
if args.get(2).is_some_and(|a| a == "--help") {
print_usage();
return Ok(());
}
if args.len() < 5 {
print_usage();
bail!("expected: mahbot debug --db <name> \"SQL query\"");
}
if args[2] != "--db" {
eprintln!("Error: expected --db flag, got '{}'", args[2]);
print_usage();
bail!("expected --db flag");
}
let db_name = &args[3];
let sql = &args[4];
let mahbot_home = crate::config::default_config_dir()?;
validate_read_only(sql)?;
let db_list = resolve_db_list(db_name)?;
for (label, db_filename) in &db_list {
if db_name == "all" {
println!("=== {label} ===");
}
let file_path = mahbot_home.join(db_filename);
if !file_path.exists() {
if db_name == "all" {
eprintln!(
"Warning: database not found, skipping: {}",
file_path.display()
);
continue;
}
bail!("database file not found: {}", file_path.display());
}
let path_str = file_path.to_string_lossy().to_string();
let opts = crate::turso::experimental_database_opts();
let db = turso::Builder::new_local(&path_str)
.experimental_index_method(opts.enable_index_method)
.experimental_multiprocess_wal(opts.enable_multiprocess_wal)
.build()
.await
.with_context(|| format!("failed to open database '{}'", file_path.display()))?;
let wrapper = db
.connect()
.with_context(|| format!("failed to connect to database '{}'", file_path.display()))?;
wrapper
.busy_timeout(Duration::from_mins(1))
.with_context(|| format!("failed to set busy timeout for '{}'", file_path.display()))?;
let query_result = execute_query(&wrapper, sql).await;
drop(wrapper);
query_result?;
}
Ok(())
}
fn resolve_db_list(name: &str) -> Result<Vec<(String, String)>> {
let names = turso_mod::store_names();
if name == "all" {
Ok(names
.iter()
.map(|n| (n.to_string(), format!("db/{n}.db")))
.collect())
} else if names.contains(&name) {
Ok(vec![(name.to_string(), format!("db/{name}.db"))])
} else {
let valid = names.join(", ");
bail!("invalid database name '{name}'. Valid names: {valid}, all");
}
}
fn validate_read_only(sql: &str) -> Result<()> {
for token in tokenize_sql(sql) {
let upper = token.to_uppercase();
if BLOCKLIST.contains(&upper.as_str()) {
bail!("query rejected: contains blocked keyword '{token}'");
}
}
Ok(())
}
fn tokenize_sql(sql: &str) -> Vec<String> {
sql.split(|c: char| c.is_whitespace() || SQL_PUNCTUATION.contains(&c))
.filter(|s| !s.is_empty())
.map(String::from)
.collect()
}
async fn execute_query(conn: &::turso::Connection, sql: &str) -> Result<()> {
let mut rows = conn.query(sql, ()).await.context("SQL query failed")?;
let column_names = rows.column_names();
if column_names.is_empty() {
return Ok(());
}
println!("{}", column_names.join("|"));
let col_count = column_names.len();
let mut row_count = 0;
let mut has_more = false;
while let Some(row) = rows.next().await? {
if row_count >= ROW_LIMIT {
has_more = true;
break;
}
print_row(&row, col_count);
row_count += 1;
}
if has_more {
print_truncation_row(col_count);
}
Ok(())
}
fn print_row(row: &::turso::Row, column_count: usize) {
let parts: Vec<String> = (0..column_count)
.map(|idx| format_value(row.get_value(idx)))
.collect();
println!("{}", parts.join("|"));
}
fn format_value(val: ::turso::Result<::turso::Value>) -> String {
match val {
Ok(::turso::Value::Null) | Err(_) => String::new(),
Ok(::turso::Value::Integer(i)) => i.to_string(),
Ok(::turso::Value::Real(f)) => f.to_string(),
Ok(::turso::Value::Text(s)) => s,
Ok(::turso::Value::Blob(b)) => hex_encode(&b),
}
}
fn hex_encode(bytes: &[u8]) -> String {
use std::fmt::Write;
let mut s = String::with_capacity(bytes.len() * 2);
for b in bytes {
let _ = write!(s, "{b:02x}");
}
s
}
fn print_truncation_row(column_count: usize) {
let parts: Vec<&str> = match column_count {
1 => vec!["truncated"],
2 => vec!["truncated", "truncated"],
_ => {
let mut parts = vec!["..."];
parts.extend(std::iter::repeat_n("truncated", column_count - 2));
parts.push("...");
parts
}
};
println!("{}", parts.join("|"));
}
fn print_usage() {
eprintln!("Usage: mahbot debug --db <name> \"SQL query\"");
let names = turso_mod::store_names().join(" | ");
eprintln!(" --db <name> {names} | all");
eprintln!(" SQL query read-only SQL, quoted as a single argument");
eprintln!();
eprintln!("Examples:");
eprintln!(" mahbot debug --db board \"SELECT status, COUNT(*) FROM tickets GROUP BY status\"");
eprintln!(" mahbot debug --db all \"SELECT name FROM sqlite_master WHERE type='table'\"");
}