mod doctor;
mod server;
use clap::{Parser, Subcommand};
#[derive(Parser, Debug)]
#[command(name = "duck-sqllsp", version, about, long_about = None)]
struct Cli {
#[command(subcommand)]
cmd: Option<Cmd>,
}
#[derive(Subcommand, Debug)]
enum Cmd {
Server {
#[arg(long, hide = true)]
stdio: bool,
#[arg(long, hide = true)]
node_ipc: bool,
#[arg(long, hide = true)]
socket: Option<String>,
},
Version,
Rules {
#[arg(long)]
json: bool,
#[arg(long)]
severity: Option<String>,
#[arg(long)]
search: Option<String>,
},
Lint {
files: Vec<String>,
#[arg(long, default_value = "text")]
format: String,
#[arg(long)]
warnings_as_errors: bool,
#[arg(long)]
dialect: Option<String>,
},
Format {
files: Vec<String>,
#[arg(long)]
stdout: bool,
#[arg(long, default_value = "postgresql")]
language: String,
},
Doctor {
path: Option<String>,
},
Introspect {
files: Vec<String>,
#[arg(long)]
url: Option<String>,
#[arg(long)]
dialect: Option<String>,
},
}
#[cfg(unix)]
fn restore_sigpipe() {
unsafe {
libc::signal(libc::SIGPIPE, libc::SIG_DFL);
}
}
#[cfg(not(unix))]
fn restore_sigpipe() {}
fn main() -> anyhow::Result<()> {
restore_sigpipe();
init_tracing();
let argv: Vec<String> = std::env::args().collect();
let cli = match Cli::try_parse_from(&argv) {
Ok(c) => c,
Err(e) => {
use clap::error::ErrorKind;
match e.kind() {
ErrorKind::DisplayHelp | ErrorKind::DisplayHelpOnMissingArgumentOrSubcommand | ErrorKind::DisplayVersion => {
e.exit()
},
_ => {
let filtered: Vec<String> =
argv.into_iter().enumerate().filter(|(i, a)| *i == 0 || !a.starts_with("--")).map(|(_, a)| a).collect();
Cli::try_parse_from(&filtered)
.unwrap_or(Cli { cmd: Some(Cmd::Server { stdio: true, node_ipc: false, socket: None }) })
},
}
},
};
match cli.cmd.unwrap_or(Cmd::Server { stdio: true, node_ipc: false, socket: None }) {
Cmd::Server { .. } => server::run(),
Cmd::Doctor { path } => doctor::run(path),
Cmd::Version => {
println!("duck-sqllsp {}", env!("CARGO_PKG_VERSION"));
println!();
let mut by_sev: std::collections::BTreeMap<&str, usize> = Default::default();
for rule in dsl_analysis::rules::all() {
let sev = match rule.default_severity() {
dsl_analysis::Severity::Error => "error",
dsl_analysis::Severity::Warning => "warning",
dsl_analysis::Severity::Info => "info",
dsl_analysis::Severity::Hint => "hint",
};
*by_sev.entry(sev).or_insert(0) += 1;
}
let total: usize = by_sev.values().sum();
let breakdown = by_sev.iter().map(|(s, n)| format!("{n} {s}")).collect::<Vec<_>>().join(", ");
println!("{:<12} postgresql, mysql, sqlite, mssql", "dialects");
println!("{:<12} {total} ({breakdown})", "lint rules");
println!("{:<12} libpg_query for postgres, sqlparser for the rest", "parser");
match dsl_format::external::locate_binary() {
Some(p) => println!("{:<12} {p}", "formatter"),
None => println!("{:<12} sql-formatter not on PATH (alignment pass only)", "formatter"),
}
println!();
println!("`duck-sqllsp doctor` checks this against a specific project.");
Ok(())
},
Cmd::Rules { json, severity, search } => {
let filter = severity.as_deref().map(|s| s.to_ascii_lowercase());
let needle = search.as_deref().map(|s| s.to_ascii_lowercase());
let mut rules: Vec<(String, &'static str)> = dsl_analysis::rules::all()
.into_iter()
.map(|r| {
(
r.code().to_string(),
match r.default_severity() {
dsl_analysis::Severity::Error => "error",
dsl_analysis::Severity::Warning => "warning",
dsl_analysis::Severity::Info => "info",
dsl_analysis::Severity::Hint => "hint",
},
)
})
.filter(|(_, sev)| filter.as_deref().is_none_or(|f| *sev == f))
.filter(|(code, _)| {
needle.as_deref().is_none_or(|n| {
code.to_ascii_lowercase().contains(n)
|| dsl_analysis::rules::title(code).is_some_and(|t| t.to_ascii_lowercase().contains(n))
})
})
.collect();
rules.sort_by(|a, b| a.0.cmp(&b.0));
if json {
print!("[");
for (i, (code, sev)) in rules.iter().enumerate() {
if i > 0 {
print!(",");
}
let title = dsl_analysis::rules::title(code).unwrap_or_default();
let escaped = title.replace('\\', "\\\\").replace('"', "\\\"");
print!("{{\"code\":\"{code}\",\"default_severity\":\"{sev}\",\"title\":\"{escaped}\"}}");
}
println!("]");
return Ok(());
}
let mut by_sev: std::collections::BTreeMap<&str, usize> = Default::default();
println!("{:6} {:8} summary", "code", "severity");
for (code, sev) in &rules {
println!("{:6} {:8} {}", code, sev, dsl_analysis::rules::title(code).unwrap_or(""));
*by_sev.entry(sev).or_insert(0) += 1;
}
if rules.is_empty() {
println!("no rules matched");
return Ok(());
}
println!();
println!("total: {} rules", rules.len());
for (sev, n) in by_sev {
println!(" {sev}: {n}");
}
Ok(())
},
Cmd::Lint { files, format, warnings_as_errors, dialect } => {
let cfg = files
.iter()
.find(|f| *f != "-")
.map(std::path::Path::new)
.and_then(dsl_server::config::load_project_config)
.unwrap_or_default();
let dialect = resolve_dialect(dialect, &files);
let json = matches!(format.as_str(), "json");
let mut error_count = 0usize;
let mut warning_count = 0usize;
if json {
print!("[");
}
let mut json_first = true;
let inputs: Vec<String> = if files.is_empty() { vec!["-".to_string()] } else { files };
for path in &inputs {
let source = if path == "-" {
use std::io::Read;
let mut buf = String::new();
std::io::stdin().read_to_string(&mut buf).map_err(|e| {
eprintln!("error reading stdin: {e}");
anyhow::anyhow!("stdin read failed")
})?;
buf
} else {
match std::fs::read_to_string(path) {
Ok(s) => s,
Err(e) => {
eprintln!("error reading {path}: {e}");
std::process::exit(2);
},
}
};
let parsed = dsl_parse::parse(&source, dialect);
let scopes = dsl_resolve::resolve_with_source(&parsed.statements, &source);
let mut catalog = dsl_completion::source_tables::from_source(&parsed, &source);
if path != "-"
&& let Some(parent) = std::path::Path::new(path).parent()
&& let Ok(rd) = std::fs::read_dir(parent)
{
for entry in rd.flatten() {
let p = entry.path();
if p.as_os_str() == std::ffi::OsStr::new(path) {
continue;
}
let Some(ext) = p.extension().and_then(|s| s.to_str()) else { continue };
if !matches!(ext.to_ascii_lowercase().as_str(), "sql" | "pgsql" | "psql") {
continue;
}
let Ok(meta) = std::fs::metadata(&p) else { continue };
if meta.len() > 4 * 1024 * 1024 {
continue;
}
let Ok(text) = std::fs::read_to_string(&p) else { continue };
let other = dsl_parse::parse(&text, dialect);
let derived = dsl_completion::source_tables::from_source(&other, &text);
catalog = dsl_completion::source_tables::merge(&catalog, &derived);
}
}
let raw = dsl_analysis::run_with_dialect(&source, &parsed, &scopes, &catalog, dialect);
let diags: Vec<dsl_analysis::Diagnostic> = raw
.into_iter()
.filter_map(|mut d| {
let Some(over) = cfg.rules.get(d.code) else { return Some(d) };
d.severity = match over.to_ascii_lowercase().as_str() {
"off" | "ignore" | "none" => return None,
"error" => dsl_analysis::Severity::Error,
"warning" | "warn" => dsl_analysis::Severity::Warning,
"info" | "information" => dsl_analysis::Severity::Info,
"hint" => dsl_analysis::Severity::Hint,
_ => d.severity,
};
Some(d)
})
.collect();
for d in &diags {
let sev_str = match d.severity {
dsl_analysis::Severity::Error => {
error_count += 1;
"error"
},
dsl_analysis::Severity::Warning => {
warning_count += 1;
"warning"
},
dsl_analysis::Severity::Info => "info",
dsl_analysis::Severity::Hint => "hint",
};
let s: u32 = d.range.start().into();
let e: u32 = d.range.end().into();
let (line, col) = byte_to_line_col(&source, s as usize);
if json {
if !json_first {
print!(",");
}
json_first = false;
let msg_esc = d.message.replace('\\', "\\\\").replace('"', "\\\"");
print!(
"{{\"file\":\"{}\",\"line\":{},\"col\":{},\"start\":{},\"end\":{},\"severity\":\"{}\",\"code\":\"{}\",\"message\":\"{}\"}}",
path,
line + 1,
col + 1,
s,
e,
sev_str,
d.code,
msg_esc,
);
} else {
println!("{path}:{}:{}: {sev_str} [{code}] {msg}", line + 1, col + 1, code = d.code, msg = d.message);
}
}
}
if json {
println!("]");
}
if error_count > 0 || (warnings_as_errors && warning_count > 0) {
std::process::exit(1);
}
Ok(())
},
Cmd::Format { files, stdout, language } => {
let cwd = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
let proj = dsl_server::config::load_project_config(&cwd).unwrap_or_default();
let mut style = proj.style.formatter.clone();
if language != "postgresql" || style.language.is_empty() {
style.language = language;
}
let ct_style = proj.style.create_table.clone();
let inputs: Vec<String> = if files.is_empty() { vec!["-".to_string()] } else { files };
for path in &inputs {
let original = if path == "-" {
use std::io::Read;
let mut buf = String::new();
std::io::stdin().read_to_string(&mut buf).map_err(|e| anyhow::anyhow!("stdin: {e}"))?;
buf
} else {
match std::fs::read_to_string(path) {
Ok(s) => s,
Err(e) => {
eprintln!("error reading {path}: {e}");
std::process::exit(2);
},
}
};
let formatted = dsl_format::format(&original, &style, &ct_style);
if stdout || path == "-" {
print!("{formatted}");
} else if formatted != original
&& let Err(e) = std::fs::write(path, formatted)
{
eprintln!("error writing {path}: {e}");
std::process::exit(2);
}
}
Ok(())
},
Cmd::Introspect { files, url, dialect } => {
if let Some(url) = url {
let spec = dsl_conn::ConnectionSpec { name: "cli".into(), url };
let rt = match tokio::runtime::Builder::new_current_thread().enable_all().build() {
Ok(rt) => rt,
Err(e) => {
eprintln!("error: failed to build tokio runtime: {e}");
std::process::exit(2);
},
};
let cat = rt.block_on(async move {
let driver = dsl_conn::build(&spec).map_err(|e| format!("build driver: {e}"))?;
driver.introspect().await.map_err(|e| format!("introspect: {e}"))
});
match cat {
Ok(c) => {
println!("{}", serde_json::to_string_pretty(&c).map_err(|e| anyhow::anyhow!("json: {e}"))?);
return Ok(());
},
Err(e) => {
eprintln!("error: {e}");
std::process::exit(2);
},
}
}
let mut acc = dsl_catalog::Catalog {
version: dsl_catalog::CATALOG_VERSION,
connection_id: "<cli-introspect>".into(),
schemas: Vec::new(),
functions: Vec::new(),
types: Vec::new(),
roles: Vec::new(),
sequences: Vec::new(),
extensions: Vec::new(),
};
let dialect = resolve_dialect(dialect, &files);
for path in &files {
let Ok(source) = std::fs::read_to_string(path) else {
eprintln!("error reading {path}");
std::process::exit(2);
};
let parsed = dsl_parse::parse(&source, dialect);
let derived = dsl_completion::source_tables::from_source(&parsed, &source);
acc = dsl_completion::source_tables::merge(&acc, &derived);
}
println!("{}", serde_json::to_string_pretty(&acc).map_err(|e| anyhow::anyhow!("json: {e}"))?);
Ok(())
},
}
}
fn resolve_dialect(explicit: Option<String>, files: &[String]) -> dsl_parse::Dialect {
let cfg = files
.iter()
.find(|f| *f != "-")
.map(std::path::Path::new)
.and_then(dsl_server::config::load_project_config)
.unwrap_or_default();
let name = explicit.unwrap_or_else(|| match cfg.effective_dialect() {
dsl_server::config::Dialect::Postgresql => "postgres".into(),
dsl_server::config::Dialect::Mysql => "mysql".into(),
dsl_server::config::Dialect::Sqlite => "sqlite".into(),
dsl_server::config::Dialect::Mssql => "mssql".into(),
});
match name.to_ascii_lowercase().as_str() {
"postgres" | "postgresql" | "pg" => dsl_parse::Dialect::Postgres,
"mysql" | "mariadb" => dsl_parse::Dialect::MySql,
"sqlite" => dsl_parse::Dialect::SQLite,
"mssql" | "tsql" | "sqlserver" => dsl_parse::Dialect::MsSql,
"generic" => dsl_parse::Dialect::Generic,
other => {
eprintln!("error: unknown dialect '{other}'; valid: postgres, mysql, sqlite, mssql, generic");
std::process::exit(2);
},
}
}
fn byte_to_line_col(src: &str, off: usize) -> (usize, usize) {
let mut line = 0usize;
let mut col = 0usize;
for (i, b) in src.bytes().enumerate() {
if i >= off {
break;
}
if b == b'\n' {
line += 1;
col = 0
} else {
col += 1
}
}
(line, col)
}
fn init_tracing() {
use tracing_subscriber::EnvFilter;
let _ = tracing_subscriber::fmt()
.with_env_filter(EnvFilter::try_from_env("DUCK_SQLLSP_LOG").unwrap_or_else(|_| EnvFilter::new("warn")))
.with_writer(std::io::stderr)
.with_ansi(false)
.try_init();
}