use crossterm::event::{KeyCode, KeyEvent, KeyModifiers};
use ratatui::layout::{Constraint, Layout, Rect};
use ratatui::style::{Modifier, Style};
use ratatui::text::{Line, Span};
use ratatui::widgets::{Block, Borders, Clear, List, ListItem, ListState, Paragraph};
use ratatui::Frame;
use crate::theme::Theme;
const MAX_TRANSCRIPT: usize = 400;
const MAX_HISTORY: usize = 200;
pub const RESULT_ROWS: usize = 500;
const MAX_CANDIDATES: usize = 12;
#[derive(Debug, Clone, PartialEq, Eq)]
enum Ctx {
Start,
Table,
Column(Scope),
Qualified(usize),
Pragma,
Dot,
Nothing,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct Scope {
tables: Vec<usize>,
aliases: Vec<(String, usize)>,
}
impl Scope {
fn resolve(&self, name: &str, schema: &[(String, Vec<String>)]) -> Option<usize> {
if let Some((_, idx)) = self
.aliases
.iter()
.find(|(a, _)| a.eq_ignore_ascii_case(name))
{
return Some(*idx);
}
schema
.iter()
.position(|(t, _)| t.eq_ignore_ascii_case(name))
}
}
const CLAUSES: &[&str] = &[
"SELECT", "FROM", "WHERE", "SET", "VALUES", "JOIN", "ON", "BY", "HAVING", "LIMIT", "OFFSET",
"INTO", "UPDATE", "DELETE", "INSERT", "CREATE", "DROP", "ALTER", "PRAGMA", "WITH", "UNION",
];
const STATEMENTS: &[&str] = &[
"SELECT",
"INSERT INTO",
"UPDATE",
"DELETE FROM",
"CREATE TABLE",
"CREATE INDEX",
"CREATE VIEW",
"DROP TABLE",
"DROP INDEX",
"DROP VIEW",
"ALTER TABLE",
"PRAGMA",
"EXPLAIN",
"EXPLAIN QUERY PLAN",
"WITH",
"BEGIN",
"COMMIT",
"ROLLBACK",
"VACUUM",
"ANALYZE",
"ATTACH",
"DETACH",
"REINDEX",
];
const CLAUSE_WORDS: &[&str] = &[
"AS",
"AND",
"OR",
"NOT",
"IS NULL",
"IS NOT NULL",
"IN",
"LIKE",
"GLOB",
"BETWEEN",
"NULL",
"FROM",
"WHERE",
"GROUP BY",
"ORDER BY",
"HAVING",
"LIMIT",
"OFFSET",
"ASC",
"DESC",
"DISTINCT",
"JOIN",
"LEFT JOIN",
"ON",
"UNION ALL",
"COLLATE NOCASE",
];
const FUNCTIONS: &[&str] = &[
"COUNT(*)",
"SUM",
"AVG",
"MIN",
"MAX",
"TOTAL",
"GROUP_CONCAT",
"LENGTH",
"LOWER",
"UPPER",
"SUBSTR",
"REPLACE",
"COALESCE",
"IFNULL",
"CAST",
"ROUND",
"ABS",
"HEX",
"QUOTE",
"DATETIME",
"STRFTIME",
];
pub const DOT_COMMANDS: &[&str] = &[
".help",
".tables",
".schema",
".indexes",
".dump",
".databases",
".attach",
".detach",
".mode",
".headers",
".output",
".once",
".timer",
".eqp",
".import",
".read",
".backup",
".expert",
".recover",
".vacuum",
".analyze",
".reindex",
".quit",
];
const PRAGMAS: &[&str] = &[
"table_info",
"table_list",
"index_list",
"index_info",
"foreign_key_list",
"journal_mode",
"wal_checkpoint",
"synchronous",
"page_size",
"page_count",
"freelist_count",
"cache_size",
"integrity_check",
"quick_check",
"optimize",
"user_version",
"schema_version",
"auto_vacuum",
"encoding",
"busy_timeout",
"compile_options",
"database_list",
];
const KEYWORD_WORDS: &[&str] = &[
"SELECT",
"FROM",
"WHERE",
"AND",
"OR",
"NOT",
"NULL",
"IN",
"LIKE",
"GLOB",
"BETWEEN",
"ORDER",
"GROUP",
"BY",
"HAVING",
"LIMIT",
"OFFSET",
"ASC",
"DESC",
"DISTINCT",
"AS",
"ON",
"JOIN",
"LEFT",
"INNER",
"CROSS",
"UNION",
"ALL",
"EXCEPT",
"INTERSECT",
"SET",
"VALUES",
"INTO",
"UPDATE",
"DELETE",
"INSERT",
"CREATE",
"DROP",
"ALTER",
"TABLE",
"INDEX",
"VIEW",
"PRAGMA",
"EXPLAIN",
"WITH",
"BEGIN",
"COMMIT",
"ROLLBACK",
"VACUUM",
"ANALYZE",
"USING",
];
fn words_before(text: &[char]) -> Vec<String> {
let s: String = text.iter().collect();
s.split(|c: char| !c.is_alphanumeric() && c != '_' && c != '*')
.filter(|w| !w.is_empty())
.map(|w| w.to_string())
.collect()
}
#[derive(Debug, PartialEq, Eq)]
pub enum Action {
None,
Close,
Execute(String),
Explain(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Entry {
Sql(String),
Rows {
columns: Vec<String>,
rows: Vec<Vec<String>>,
truncated: bool,
},
Changed(usize),
Timing(std::time::Duration),
Plan(Vec<String>),
Note(Vec<String>),
Error(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct Completion {
at: usize,
prefix: String,
items: Vec<String>,
selected: usize,
}
pub struct SqlEdit {
input: Vec<char>,
cursor: usize,
transcript: Vec<Entry>,
scroll: usize,
follow: bool,
history: Vec<String>,
hist_idx: Option<usize>,
stash: Vec<char>,
completion: Option<Completion>,
schema: Vec<(String, Vec<String>)>,
pub running: bool,
}
impl SqlEdit {
pub fn new(schema: Vec<(String, Vec<String>)>) -> Self {
SqlEdit {
input: Vec::new(),
cursor: 0,
transcript: Vec::new(),
scroll: 0,
follow: true,
history: load_history(),
hist_idx: None,
stash: Vec::new(),
completion: None,
schema,
running: false,
}
}
pub fn text(&self) -> String {
self.input.iter().collect()
}
pub fn last_statement(&self) -> Option<String> {
self.transcript.iter().rev().find_map(|e| match e {
Entry::Sql(sql) if !sql.starts_with('.') => Some(sql.clone()),
_ => None,
})
}
pub fn push(&mut self, entry: Entry) {
self.transcript.push(entry);
while self.transcript.len() > MAX_TRANSCRIPT {
self.transcript.remove(0);
}
self.running = false;
self.follow = true;
}
#[cfg(test)]
pub fn transcript(&self) -> &[Entry] {
&self.transcript
}
fn insert_char(&mut self, c: char) {
self.input.insert(self.cursor, c);
self.cursor += 1;
self.hist_idx = None;
self.completion = None;
}
fn submit(&mut self) -> Action {
let sql = self.text().trim().to_string();
if sql.is_empty() {
return Action::None;
}
self.transcript.push(Entry::Sql(sql.clone()));
if self.history.last().map(String::as_str) != Some(sql.as_str()) {
self.history.push(sql.clone());
while self.history.len() > MAX_HISTORY {
self.history.remove(0);
}
save_history(&self.history);
}
self.input.clear();
self.cursor = 0;
self.hist_idx = None;
self.completion = None;
self.running = true;
self.follow = true;
Action::Execute(sql)
}
fn history_move(&mut self, back: bool) {
if self.history.is_empty() {
return;
}
match (self.hist_idx, back) {
(None, true) => {
self.stash = std::mem::take(&mut self.input);
self.hist_idx = Some(self.history.len() - 1);
}
(Some(0), true) => return,
(Some(i), true) => self.hist_idx = Some(i - 1),
(Some(i), false) if i + 1 < self.history.len() => self.hist_idx = Some(i + 1),
(Some(_), false) => {
self.hist_idx = None;
self.input = std::mem::take(&mut self.stash);
self.cursor = self.input.len();
return;
}
(None, false) => return,
}
if let Some(i) = self.hist_idx {
self.input = self.history[i].chars().collect();
self.cursor = self.input.len();
}
self.completion = None;
}
fn kill_word(&mut self) {
let mut i = self.cursor;
while i > 0 && self.input[i - 1].is_whitespace() {
i -= 1;
}
while i > 0 && !self.input[i - 1].is_whitespace() {
i -= 1;
}
self.input.drain(i..self.cursor);
self.cursor = i;
self.hist_idx = None;
self.completion = None;
}
fn explain(&mut self) -> Action {
let sql = self.text().trim().to_string();
if sql.is_empty() {
return Action::None;
}
self.transcript
.push(Entry::Sql(format!("EXPLAIN QUERY PLAN {sql}")));
self.follow = true;
Action::Explain(sql)
}
fn complete(&mut self, back: bool) {
if let Some(c) = self.completion.as_mut() {
let n = c.items.len();
c.selected = if back {
(c.selected + n - 1) % n
} else {
(c.selected + 1) % n
};
return;
}
let (at, prefix) = self.word_before_cursor();
let items = self.candidates_at(&prefix, at);
if items.is_empty() {
return;
}
if items.len() == 1 {
self.replace_word(at, &items[0]);
return;
}
self.completion = Some(Completion {
at,
prefix,
items,
selected: 0,
});
}
fn accept_completion(&mut self) {
if let Some(c) = self.completion.take() {
let pick = c.items[c.selected].clone();
self.replace_word(c.at, &pick);
}
}
fn replace_word(&mut self, at: usize, word: &str) {
self.input.drain(at..self.cursor);
for (i, ch) in word.chars().enumerate() {
self.input.insert(at + i, ch);
}
self.cursor = at + word.chars().count();
self.completion = None;
self.hist_idx = None;
}
fn word_before_cursor(&self) -> (usize, String) {
let mut at = self.cursor;
while at > 0 {
let c = self.input[at - 1];
if c.is_alphanumeric() || c == '_' {
at -= 1;
} else {
break;
}
}
if at > 0
&& self.input[at - 1] == '.'
&& self.input[..at - 1]
.iter()
.rev()
.take_while(|c| **c != '\n')
.all(|c| c.is_whitespace())
{
at -= 1;
}
(at, self.input[at..self.cursor].iter().collect())
}
#[cfg(test)]
fn candidates(&self, prefix: &str) -> Vec<String> {
let at = self.cursor - prefix.chars().count();
self.candidates_at(prefix, at)
}
fn candidates_at(&self, prefix: &str, at: usize) -> Vec<String> {
let p = prefix.to_lowercase();
let mut out: Vec<String> = Vec::new();
let push = |s: &str, out: &mut Vec<String>| {
if (p.is_empty() || s.to_lowercase().starts_with(&p))
&& !out.iter().any(|o| o.eq_ignore_ascii_case(s))
{
out.push(s.to_string());
}
};
let tables = |out: &mut Vec<String>| {
for (t, _) in &self.schema {
push(t, out);
}
};
match self.context(at) {
Ctx::Nothing => {}
Ctx::Start => {
for k in STATEMENTS {
push(k, &mut out);
}
}
Ctx::Pragma => {
for k in PRAGMAS {
push(k, &mut out);
}
}
Ctx::Dot => {
for k in DOT_COMMANDS {
push(k, &mut out);
}
}
Ctx::Table => tables(&mut out),
Ctx::Qualified(t) => {
for c in &self.schema[t].1 {
push(c.as_str(), &mut out);
}
}
Ctx::Column(scope) => {
for &t in &scope.tables {
for c in &self.schema[t].1 {
push(c.as_str(), &mut out);
}
}
for (alias, _) in &scope.aliases {
push(alias, &mut out);
}
if scope.tables.is_empty() {
for (_, cols) in &self.schema {
for c in cols {
push(c.as_str(), &mut out);
}
}
}
for k in CLAUSE_WORDS {
push(k, &mut out);
}
tables(&mut out);
for k in FUNCTIONS {
push(k, &mut out);
}
}
}
out.truncate(MAX_CANDIDATES);
out
}
fn context(&self, at: usize) -> Ctx {
if at > 0 && self.input[at - 1] == '.' {
let mut start = at - 1;
while start > 0 {
let c = self.input[start - 1];
if c.is_alphanumeric() || c == '_' {
start -= 1;
} else {
break;
}
}
let name: String = self.input[start..at - 1].iter().collect();
let scope = self.scope(at);
if let Some(t) = scope.resolve(&name, &self.schema) {
return Ctx::Qualified(t);
}
return Ctx::Nothing;
}
let typed: String = self.input[..at].iter().collect();
if typed.trim_start().starts_with('.') && !typed.trim().contains(' ') {
return Ctx::Dot;
}
let words = words_before(&self.input[..at]);
let last = words.last().map(|w| w.to_ascii_uppercase());
let last = last.as_deref().unwrap_or("");
let clause = words
.iter()
.rev()
.map(|w| w.to_ascii_uppercase())
.find(|w| CLAUSES.contains(&w.as_str()));
let after_semicolon = self.input[..at]
.iter()
.rev()
.find(|c| !c.is_whitespace())
.map(|&c| c == ';')
.unwrap_or(true);
if words.is_empty() || after_semicolon {
return Ctx::Start;
}
if last == "PRAGMA" {
return Ctx::Pragma;
}
if matches!(last, "FROM" | "JOIN" | "INTO" | "UPDATE" | "TABLE") {
return Ctx::Table;
}
let scope = self.scope(at);
match clause.as_deref() {
Some("SELECT") | Some("WHERE") | Some("SET") | Some("ON") | Some("BY")
| Some("HAVING") | Some("VALUES") | None => Ctx::Column(scope),
_ => Ctx::Column(scope),
}
}
fn scope(&self, at: usize) -> Scope {
let words = words_before(&self.input[..at]);
let mut scope = Scope::default();
let mut expect = false;
for (i, w) in words.iter().enumerate() {
let upper = w.to_ascii_uppercase();
if matches!(upper.as_str(), "FROM" | "JOIN" | "UPDATE" | "INTO") {
expect = true;
continue;
}
if !expect {
continue;
}
expect = false;
if let Some(idx) = self
.schema
.iter()
.position(|(t, _)| t.eq_ignore_ascii_case(w))
{
if !scope.tables.contains(&idx) {
scope.tables.push(idx);
}
let mut n = i + 1;
if words.get(n).map(|w| w.eq_ignore_ascii_case("AS")) == Some(true) {
n += 1;
}
if let Some(alias) = words.get(n) {
let upper = alias.to_ascii_uppercase();
if !CLAUSES.contains(&upper.as_str())
&& !KEYWORD_WORDS.contains(&upper.as_str())
&& !self
.schema
.iter()
.any(|(t, _)| t.eq_ignore_ascii_case(alias))
{
scope.aliases.push((alias.to_string(), idx));
}
}
}
}
scope
}
pub fn on_key(&mut self, key: KeyEvent, page: usize) -> Action {
let ctrl = key.modifiers.contains(KeyModifiers::CONTROL);
let alt = key.modifiers.contains(KeyModifiers::ALT);
if self.completion.is_some() {
match key.code {
KeyCode::Tab => {
self.complete(false);
return Action::None;
}
KeyCode::BackTab => {
self.complete(true);
return Action::None;
}
KeyCode::Down => {
self.complete(false);
return Action::None;
}
KeyCode::Up => {
self.complete(true);
return Action::None;
}
KeyCode::Enter => {
self.accept_completion();
return Action::None;
}
KeyCode::Esc => {
self.completion = None;
return Action::None;
}
_ => {}
}
}
match key.code {
KeyCode::Esc => return Action::Close,
KeyCode::Char('c') if ctrl => return Action::Close,
KeyCode::Char('g') if ctrl => {
self.input.clear();
self.cursor = 0;
self.hist_idx = None;
self.completion = None;
}
KeyCode::Tab => self.complete(false),
KeyCode::BackTab => self.complete(true),
KeyCode::Char('e') if alt => return self.explain(),
KeyCode::F(5) => return self.explain(),
KeyCode::Enter if alt => self.insert_char('\n'),
KeyCode::Char('j') if ctrl => self.insert_char('\n'),
KeyCode::Enter => return self.submit(),
KeyCode::Up => self.history_move(true),
KeyCode::Down => self.history_move(false),
KeyCode::Char('p') if ctrl => self.history_move(true),
KeyCode::Char('n') if ctrl => self.history_move(false),
KeyCode::Char('l') if ctrl => {
self.transcript.clear();
self.scroll = 0;
self.follow = true;
}
KeyCode::Left => self.cursor = self.cursor.saturating_sub(1),
KeyCode::Char('b') if ctrl => self.cursor = self.cursor.saturating_sub(1),
KeyCode::Right => self.cursor = (self.cursor + 1).min(self.input.len()),
KeyCode::Char('f') if ctrl => self.cursor = (self.cursor + 1).min(self.input.len()),
KeyCode::Home => self.cursor = 0,
KeyCode::Char('a') if ctrl => self.cursor = 0,
KeyCode::End => self.cursor = self.input.len(),
KeyCode::Char('e') if ctrl => self.cursor = self.input.len(),
KeyCode::Char('k') if ctrl => {
self.input.truncate(self.cursor);
self.hist_idx = None;
}
KeyCode::Char('u') if ctrl => {
self.input.drain(0..self.cursor);
self.cursor = 0;
self.hist_idx = None;
}
KeyCode::Char('w') if ctrl => self.kill_word(),
KeyCode::Backspace => {
if self.cursor > 0 {
self.cursor -= 1;
self.input.remove(self.cursor);
self.hist_idx = None;
self.completion = None;
}
}
KeyCode::Delete => {
if self.cursor < self.input.len() {
self.input.remove(self.cursor);
self.hist_idx = None;
self.completion = None;
}
}
KeyCode::PageUp => {
self.scroll = self.scroll.saturating_sub(page);
self.follow = false;
}
KeyCode::PageDown => {
self.scroll += page;
self.follow = true;
}
KeyCode::Char(c) if !ctrl => self.insert_char(c),
_ => {}
}
Action::None
}
fn transcript_lines(&self, t: &Theme, width: usize) -> Vec<Line<'static>> {
let mut out: Vec<Line> = Vec::new();
for e in &self.transcript {
match e {
Entry::Sql(sql) => {
for (i, l) in sql.lines().enumerate() {
out.push(Line::from(vec![
Span::styled(
if i == 0 { "sql> " } else { " . " }.to_string(),
Style::default().fg(t.accent),
),
Span::styled(l.to_string(), Style::default().fg(t.primary)),
]));
}
}
Entry::Changed(n) => out.push(Line::from(Span::styled(
format!(" {} row{} changed", n, if *n == 1 { "" } else { "s" }),
Style::default().fg(t.label),
))),
Entry::Timing(d) => out.push(Line::from(Span::styled(
format!(" {:.3?}", d),
Style::default().fg(t.dim),
))),
Entry::Plan(steps) => {
for s in steps {
out.push(Line::from(Span::styled(
format!(" {s}"),
Style::default().fg(t.label),
)));
}
}
Entry::Note(lines) => {
for l in lines {
out.push(Line::from(Span::styled(
format!(" {l}"),
Style::default().fg(t.primary),
)));
}
}
Entry::Error(msg) => {
for l in msg.lines() {
out.push(Line::from(Span::styled(
format!(" {l}"),
Style::default().fg(t.alt).add_modifier(Modifier::BOLD),
)));
}
}
Entry::Rows {
columns,
rows,
truncated,
} => {
let per = (width.saturating_sub(6) / columns.len().max(1)).clamp(6, 28);
let head: String = columns
.iter()
.map(|c| format!("{:<per$}", crate::app::truncate(c, per), per = per))
.collect::<Vec<_>>()
.join(" ");
out.push(Line::from(Span::styled(
format!(" {head}"),
Style::default().fg(t.label).add_modifier(Modifier::BOLD),
)));
for r in rows {
let line: String = r
.iter()
.map(|c| format!("{:<per$}", crate::app::truncate(c, per), per = per))
.collect::<Vec<_>>()
.join(" ");
out.push(Line::from(Span::styled(
format!(" {line}"),
Style::default().fg(t.primary),
)));
}
out.push(Line::from(Span::styled(
format!(
" {} row{}{}",
rows.len(),
if rows.len() == 1 { "" } else { "s" },
if *truncated {
format!(" (first {RESULT_ROWS})")
} else {
String::new()
}
),
Style::default().fg(t.dim),
)));
}
}
}
out
}
pub fn render(&mut self, f: &mut Frame, area: Rect, t: &Theme) {
let input_lines = (self.input.iter().filter(|&&c| c == '\n').count() + 1).min(8) as u16;
let rows =
Layout::vertical([Constraint::Min(3), Constraint::Length(input_lines + 2)]).split(area);
let lines = self.transcript_lines(t, rows[0].width as usize);
let view = rows[0].height.saturating_sub(2) as usize;
let max_scroll = lines.len().saturating_sub(view);
if self.follow {
self.scroll = max_scroll;
}
self.scroll = self.scroll.min(max_scroll);
let shown: Vec<Line> = lines
.into_iter()
.skip(self.scroll)
.take(view.max(1))
.collect();
f.render_widget(
Paragraph::new(shown).block(
Block::default()
.borders(Borders::ALL)
.border_style(Style::default().fg(t.accent))
.title(format!(
" SQL — {} statement{} · Tab completes · ^j newline · Enter runs · Esc back ",
self.history.len(),
if self.history.len() == 1 { "" } else { "s" }
)),
),
rows[0],
);
let text: String = self.input.iter().collect();
let mut body: Vec<Line> = Vec::new();
let mut idx = 0usize;
for (n, l) in text.split('\n').enumerate() {
let len = l.chars().count();
let mut spans = vec![Span::styled(
if n == 0 { "sql> " } else { " . " }.to_string(),
Style::default().fg(t.accent),
)];
if self.cursor >= idx && self.cursor <= idx + len {
let at = self.cursor - idx;
let (before, rest) =
l.split_at(l.char_indices().nth(at).map_or(l.len(), |(i, _)| i));
let mut chars = rest.chars();
let under = chars.next();
spans.push(Span::raw(before.to_string()));
spans.push(Span::styled(
under.map(|c| c.to_string()).unwrap_or_else(|| " ".into()),
Style::default().add_modifier(Modifier::REVERSED),
));
spans.push(Span::raw(chars.as_str().to_string()));
} else {
spans.push(Span::raw(l.to_string()));
}
body.push(Line::from(spans));
idx += len + 1;
}
f.render_widget(
Paragraph::new(body).block(
Block::default()
.borders(Borders::ALL)
.border_style(Style::default().fg(if self.running { t.label } else { t.dim }))
.title(if self.running {
" running… ".to_string()
} else {
format!(" input · {} chars ", self.input.len())
}),
),
rows[1],
);
if let Some(c) = &self.completion {
let h = (c.items.len() as u16 + 2).min(area.height.saturating_sub(3));
let w = c
.items
.iter()
.map(|i| i.chars().count())
.max()
.unwrap_or(10)
.clamp(12, 40) as u16
+ 4;
let x = area.x + 5;
let y = rows[1].y.saturating_sub(h + 1);
let popup = Rect::new(x, y, w.min(area.width.saturating_sub(5)), h.max(3));
f.render_widget(Clear, popup);
let items: Vec<ListItem> = c.items.iter().map(|i| ListItem::new(i.clone())).collect();
let mut st = ListState::default();
st.select(Some(c.selected));
f.render_stateful_widget(
List::new(items)
.block(
Block::default()
.borders(Borders::ALL)
.border_style(Style::default().fg(t.accent))
.title(format!(" {} ", c.prefix)),
)
.highlight_style(
Style::default()
.fg(t.accent)
.add_modifier(Modifier::REVERSED),
),
popup,
&mut st,
);
}
}
pub fn page_rows(area: Rect) -> usize {
(area.height as usize).saturating_sub(6).max(1)
}
}
fn history_file() -> Option<std::path::PathBuf> {
let base = std::env::var_os("XDG_CACHE_HOME")
.map(std::path::PathBuf::from)
.or_else(|| std::env::var_os("HOME").map(|h| std::path::PathBuf::from(h).join(".cache")))?;
Some(base.join("zdbview").join("sql_history"))
}
fn load_history() -> Vec<String> {
let path = match history_file() {
Some(p) => p,
None => return Vec::new(),
};
std::fs::read_to_string(path)
.map(|s| {
s.lines()
.filter(|l| !l.trim().is_empty())
.map(|l| l.replace("\\n", "\n"))
.collect()
})
.unwrap_or_default()
}
fn save_history(history: &[String]) {
let path = match history_file() {
Some(p) => p,
None => return,
};
if let Some(dir) = path.parent() {
let _ = std::fs::create_dir_all(dir);
}
let body: String = history
.iter()
.map(|s| format!("{}\n", s.replace('\n', "\\n")))
.collect();
let tmp = path.with_extension("tmp");
if std::fs::write(&tmp, body).is_ok() {
let _ = std::fs::rename(&tmp, path);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::theme::ThemeName;
use ratatui::backend::TestBackend;
use ratatui::Terminal;
fn schema() -> Vec<(String, Vec<String>)> {
vec![
(
"users".into(),
vec!["user_id".into(), "email".into(), "created_at".into()],
),
("orders".into(), vec!["order_id".into(), "total".into()]),
]
}
fn ed() -> SqlEdit {
let mut e = SqlEdit::new(schema());
e.history.clear();
e
}
fn key(c: char) -> KeyEvent {
KeyEvent::from(KeyCode::Char(c))
}
fn code(c: KeyCode) -> KeyEvent {
KeyEvent::from(c)
}
fn ctrl(c: char) -> KeyEvent {
KeyEvent::new(KeyCode::Char(c), KeyModifiers::CONTROL)
}
fn type_str(e: &mut SqlEdit, s: &str) {
for c in s.chars() {
e.on_key(key(c), 10);
}
}
fn rows_of(e: &mut SqlEdit, w: u16, h: u16) -> Vec<String> {
let theme = Theme::from_name(ThemeName::NeonSprawl);
let mut term = Terminal::new(TestBackend::new(w, h)).unwrap();
term.draw(|f| e.render(f, f.area(), &theme)).unwrap();
let buf = term.backend().buffer().clone();
(0..buf.area().height)
.map(|y| {
(0..buf.area().width)
.map(|x| buf[(x, y)].symbol())
.collect::<String>()
})
.collect()
}
#[test]
fn enter_submits_and_the_result_comes_back() {
let mut e = ed();
type_str(&mut e, "select 1");
assert_eq!(e.text(), "select 1");
let action = e.on_key(code(KeyCode::Enter), 10);
assert_eq!(action, Action::Execute("select 1".into()));
assert!(e.running, "the host has it now");
assert!(e.text().is_empty(), "input cleared for the next statement");
assert_eq!(e.transcript()[0], Entry::Sql("select 1".into()));
e.push(Entry::Rows {
columns: vec!["1".into()],
rows: vec![vec!["1".into()]],
truncated: false,
});
assert!(!e.running);
assert_eq!(e.transcript().len(), 2);
assert_eq!(e.on_key(code(KeyCode::Enter), 10), Action::None);
assert_eq!(e.transcript().len(), 2);
}
#[test]
fn multi_line_statements() {
let mut e = ed();
type_str(&mut e, "select *");
e.on_key(ctrl('j'), 10);
type_str(&mut e, "from users");
assert_eq!(e.text(), "select *\nfrom users");
e.on_key(KeyEvent::new(KeyCode::Enter, KeyModifiers::ALT), 10);
type_str(&mut e, "where 1");
assert_eq!(e.text(), "select *\nfrom users\nwhere 1");
assert_eq!(
e.on_key(code(KeyCode::Enter), 10),
Action::Execute("select *\nfrom users\nwhere 1".into())
);
}
#[test]
fn completion_follows_the_clause() {
let mut e = ed();
assert_eq!(e.candidates(""), {
let mut v: Vec<String> = STATEMENTS.iter().map(|s| s.to_string()).collect();
v.truncate(MAX_CANDIDATES);
v
});
type_str(&mut e, "sel");
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "SELECT", "the only statement starting with sel");
let mut e = ed();
type_str(&mut e, "select * from ");
assert_eq!(e.candidates(""), vec!["users".to_string(), "orders".into()]);
type_str(&mut e, "us");
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "select * from users");
type_str(&mut e, " where us");
let cands = e.candidates("us");
assert_eq!(
cands.first().map(String::as_str),
Some("user_id"),
"the scoped column must come before the table name: {cands:?}"
);
e.on_key(code(KeyCode::Tab), 10);
let c = e.completion.as_ref().expect("menu");
assert_eq!(c.items[0], "user_id");
e.on_key(code(KeyCode::Enter), 10);
assert_eq!(e.text(), "select * from users where user_id");
let mut e = ed();
type_str(&mut e, "select tot");
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(
e.text(),
"select total",
"orders.total, table not yet named"
);
}
#[test]
fn qualified_completion_uses_the_alias() {
let mut e = ed();
type_str(&mut e, "select * from orders o where o.");
let cands = e.candidates("");
assert_eq!(
cands,
vec!["order_id".to_string(), "total".into()],
"{cands:?}"
);
type_str(&mut e, "tot");
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "select * from orders o where o.total");
let mut e = ed();
type_str(&mut e, "select users.");
assert!(e.candidates("").contains(&"email".to_string()));
let mut e = ed();
type_str(&mut e, "select nosuch.");
assert!(e.candidates("").is_empty());
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "select nosuch.", "nothing was inserted");
}
#[test]
fn a_join_scopes_both_tables() {
let mut e = ed();
type_str(
&mut e,
"select * from users u join orders o on u.user_id = o.",
);
assert_eq!(
e.candidates(""),
vec!["order_id".to_string(), "total".into()],
"the qualifier decides"
);
let mut e = ed();
type_str(&mut e, "select * from users join orders where ");
let cands = e.candidates("");
let users_first = cands.iter().position(|c| c == "user_id").unwrap();
let orders_next = cands.iter().position(|c| c == "order_id").unwrap();
assert!(users_first < orders_next, "FROM order preserved: {cands:?}");
}
#[test]
fn pragma_completes_pragma_names() {
let mut e = ed();
type_str(&mut e, "pragma jour");
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "pragma journal_mode");
let mut e = ed();
type_str(&mut e, "pragma ");
let cands = e.candidates("");
assert!(cands.contains(&"table_info".to_string()), "{cands:?}");
assert!(
!cands.iter().any(|c| c == "users"),
"no tables here: {cands:?}"
);
}
#[test]
fn completion_replaces_only_the_current_word() {
let mut e = ed();
type_str(&mut e, "select ema from users");
for _ in 0.." from users".len() - 1 {
e.on_key(code(KeyCode::Left), 10);
}
e.on_key(code(KeyCode::Tab), 10);
assert_eq!(e.text(), "select email from users");
}
#[test]
fn history_browses_and_restores_the_stash() {
let mut e = ed();
for sql in ["select 1", "select 2"] {
type_str(&mut e, sql);
e.on_key(code(KeyCode::Enter), 10);
e.push(Entry::Changed(0));
}
type_str(&mut e, "half typed");
e.on_key(code(KeyCode::Up), 10);
assert_eq!(e.text(), "select 2", "newest first");
e.on_key(ctrl('p'), 10);
assert_eq!(e.text(), "select 1");
e.on_key(ctrl('p'), 10);
assert_eq!(e.text(), "select 1", "stops at the oldest");
e.on_key(code(KeyCode::Down), 10);
assert_eq!(e.text(), "select 2");
e.on_key(ctrl('n'), 10);
assert_eq!(e.text(), "half typed", "past the newest restores the stash");
let before = e.history.len();
e.on_key(ctrl('g'), 10);
type_str(&mut e, "select 2");
e.on_key(code(KeyCode::Enter), 10);
assert_eq!(
e.history.len(),
before,
"a repeat of the newest entry is not recorded again"
);
e.push(Entry::Changed(0));
type_str(&mut e, "select 3");
e.on_key(code(KeyCode::Enter), 10);
assert_eq!(e.history.len(), before + 1, "a new statement is recorded");
}
#[test]
fn editing_chords() {
let mut e = ed();
type_str(&mut e, "select * from users");
e.on_key(ctrl('w'), 10);
assert_eq!(e.text(), "select * from ");
e.on_key(ctrl('a'), 10);
assert_eq!(e.cursor, 0);
e.on_key(ctrl('e'), 10);
assert_eq!(e.cursor, e.text().chars().count());
e.on_key(ctrl('u'), 10);
assert!(e.text().is_empty(), "^u killed to the start");
type_str(&mut e, "abc");
e.on_key(code(KeyCode::Left), 10);
e.on_key(ctrl('k'), 10);
assert_eq!(e.text(), "ab", "^k killed to the end");
e.on_key(code(KeyCode::Backspace), 10);
assert_eq!(e.text(), "a");
e.on_key(ctrl('g'), 10);
assert!(e.text().is_empty(), "^g cleared the line");
e.push(Entry::Changed(1));
assert!(!e.transcript().is_empty());
e.on_key(ctrl('l'), 10);
assert!(e.transcript().is_empty(), "^l cleared the transcript");
}
#[test]
fn render_shows_results_and_errors() {
let mut e = ed();
e.push(Entry::Sql("select * from users".into()));
e.push(Entry::Rows {
columns: vec!["user_id".into(), "email".into()],
rows: vec![
vec!["1".into(), "a@example.com".into()],
vec!["2".into(), "b@example.com".into()],
],
truncated: true,
});
let r = rows_of(&mut e, 100, 18);
assert!(r.iter().any(|l| l.contains("sql> select * from users")));
assert!(r.iter().any(|l| l.contains("user_id")), "header missing");
assert!(r.iter().any(|l| l.contains("a@example.com")), "row missing");
assert!(
r.iter().any(|l| l.contains("2 rows (first 500)")),
"count/truncation missing: {r:#?}"
);
assert!(
r.iter().any(|l| l.contains("Tab completes")),
"no key hints"
);
e.push(Entry::Changed(3));
assert!(rows_of(&mut e, 100, 18)
.iter()
.any(|l| l.contains("3 rows changed")));
e.push(Entry::Error("no such table: nope".into()));
assert!(rows_of(&mut e, 100, 18)
.iter()
.any(|l| l.contains("no such table: nope")));
type_str(&mut e, "select");
let r = rows_of(&mut e, 100, 18);
assert!(r.iter().any(|l| l.contains("sql> select")), "{r:#?}");
assert!(r.iter().any(|l| l.contains("6 chars")));
}
#[test]
fn render_shows_the_completion_menu() {
let mut e = ed();
type_str(&mut e, "select * from users where us");
e.on_key(code(KeyCode::Tab), 10);
assert!(e.completion.is_some(), "a menu should be open");
let r = rows_of(&mut e, 100, 18);
assert!(
r.iter().any(|l| l.contains("user_id")),
"the menu is not on screen: {r:#?}"
);
assert!(r.iter().any(|l| l.contains(" us ")), "prefix not titled");
}
#[test]
fn render_survives_a_tiny_terminal() {
let mut e = ed();
e.push(Entry::Changed(1));
type_str(&mut e, "select 1");
e.on_key(code(KeyCode::Tab), 10);
for (w, h) in [(20u16, 6u16), (8, 4), (1, 1), (200, 60)] {
rows_of(&mut e, w, h);
}
assert!(SqlEdit::page_rows(Rect::new(0, 0, 80, 24)) >= 1);
assert_eq!(SqlEdit::page_rows(Rect::new(0, 0, 80, 2)), 1, "never zero");
}
#[test]
fn transcript_is_bounded() {
let mut e = ed();
for i in 0..MAX_TRANSCRIPT + 50 {
e.push(Entry::Changed(i));
}
assert_eq!(e.transcript().len(), MAX_TRANSCRIPT);
assert_eq!(e.transcript()[0], Entry::Changed(50));
}
}