1use std::io::{IsTerminal, Write};
9
10use sqlx::{Row as _, TypeInfo, ValueRef};
11use tokio::io::{AsyncBufRead, AsyncBufReadExt, BufReader};
12
13use crate::Result;
14use crate::db::{Db, Dialect, Row, RowInner, script, sql};
15
16const CELL_WIDTH: usize = 60;
18
19pub(crate) async fn run(db: &Db) -> Result {
20 let interactive = std::io::stdin().is_terminal();
21 let mut out = std::io::stdout();
22 run_with(
23 db,
24 BufReader::new(tokio::io::stdin()),
25 &mut out,
26 interactive,
27 )
28 .await
29}
30
31#[doc(hidden)]
34pub async fn run_with(
35 db: &Db,
36 input: impl AsyncBufRead + Unpin,
37 out: &mut impl Write,
38 interactive: bool,
39) -> Result {
40 if interactive {
41 writeln!(
42 out,
43 "{} shell. End statements with `;`. `.tables` lists tables, `.quit` leaves.",
44 match db.dialect() {
45 Dialect::Sqlite => "SQLite",
46 Dialect::Postgres => "PostgreSQL",
47 }
48 )?;
49 }
50 let mut lines = input.lines();
51 let mut statement = String::new();
52 loop {
53 if interactive {
54 write!(
55 out,
56 "{}",
57 if statement.is_empty() {
58 "sql> "
59 } else {
60 "...> "
61 }
62 )?;
63 out.flush()?;
64 }
65 let Some(line) = lines.next_line().await? else {
66 break;
67 };
68 let trimmed = line.trim();
69 if statement.is_empty() {
70 match trimmed {
71 "" => continue,
72 ".quit" | ".exit" => break,
73 ".tables" => {
74 let names: Vec<String> = sql(match db.dialect() {
75 Dialect::Sqlite => {
76 "SELECT name FROM sqlite_master \
77 WHERE type = 'table' AND name NOT LIKE 'sqlite_%' ORDER BY name"
78 }
79 Dialect::Postgres => {
80 "SELECT tablename::text FROM pg_tables \
81 WHERE schemaname = current_schema() ORDER BY tablename"
82 }
83 })
84 .scalars(db)
85 .await?;
86 writeln!(out, "{}", names.join(" "))?;
87 continue;
88 }
89 _ => {}
90 }
91 }
92 statement.push_str(&line);
93 statement.push('\n');
94 if trimmed.ends_with(';') {
95 let sql = std::mem::take(&mut statement);
96 match execute(db, sql.trim()).await {
97 Ok(text) => write!(out, "{text}")?,
98 Err(err) => writeln!(out, "Error: {err}")?,
99 }
100 }
101 }
102 Ok(())
103}
104
105async fn execute(db: &Db, sql: &str) -> std::result::Result<String, crate::db::DbError> {
107 let first = sql
108 .split_whitespace()
109 .next()
110 .unwrap_or_default()
111 .to_ascii_lowercase();
112 if matches!(
113 first.as_str(),
114 "select" | "pragma" | "with" | "explain" | "values" | "show"
115 ) {
116 let rows = self::sql(sql).fetch_all(db).await?;
117 Ok(table(&rows))
118 } else {
119 let done = script(db, sql).await?;
120 Ok(format!("OK ({done} row(s) affected)\n"))
121 }
122}
123
124fn cell(row: &Row, index: usize) -> String {
125 match &row.0 {
126 RowInner::Sqlite(row) => sqlite_cell(row, index),
127 #[cfg(feature = "postgres")]
128 RowInner::Postgres(row) => postgres_cell(row, index),
129 }
130}
131
132fn sqlite_cell(row: &sqlx::sqlite::SqliteRow, index: usize) -> String {
133 let Ok(raw) = row.try_get_raw(index) else {
134 return "?".into();
135 };
136 if raw.is_null() {
137 return "NULL".into();
138 }
139 match raw.type_info().name() {
140 "INTEGER" => row
141 .try_get::<i64, _>(index)
142 .map(|v| v.to_string())
143 .unwrap_or_default(),
144 "REAL" => row
145 .try_get::<f64, _>(index)
146 .map(|v| v.to_string())
147 .unwrap_or_default(),
148 "BLOB" => row
149 .try_get::<Vec<u8>, _>(index)
150 .map(|v| format!("<{} bytes>", v.len()))
151 .unwrap_or_default(),
152 _ => row.try_get::<String, _>(index).unwrap_or_default(),
153 }
154}
155
156#[cfg(feature = "postgres")]
157fn postgres_cell(row: &sqlx::postgres::PgRow, index: usize) -> String {
158 let Ok(raw) = row.try_get_raw(index) else {
159 return "?".into();
160 };
161 if raw.is_null() {
162 return "NULL".into();
163 }
164 fn show<T: ToString>(value: Result<T, sqlx::Error>) -> String {
165 value.map(|v| v.to_string()).unwrap_or_else(|_| "?".into())
166 }
167 match raw.type_info().name() {
168 "INT2" => show(row.try_get::<i16, _>(index)),
169 "INT4" => show(row.try_get::<i32, _>(index)),
170 "INT8" => show(row.try_get::<i64, _>(index)),
171 "FLOAT4" => show(row.try_get::<f32, _>(index)),
172 "FLOAT8" => show(row.try_get::<f64, _>(index)),
173 "BOOL" => show(row.try_get::<bool, _>(index)),
174 "BYTEA" => row
175 .try_get::<Vec<u8>, _>(index)
176 .map(|v| format!("<{} bytes>", v.len()))
177 .unwrap_or_else(|_| "?".into()),
178 "TIMESTAMPTZ" => show(row.try_get::<chrono::DateTime<chrono::Utc>, _>(index)),
179 "TIMESTAMP" => show(row.try_get::<chrono::NaiveDateTime, _>(index)),
180 "DATE" => show(row.try_get::<chrono::NaiveDate, _>(index)),
181 _ => show(row.try_get::<String, _>(index)),
182 }
183}
184
185fn table(rows: &[Row]) -> String {
187 let Some(first) = rows.first() else {
188 return "(no rows)\n".into();
189 };
190 let header: Vec<String> = first.columns().into_iter().map(str::to_owned).collect();
191 let body: Vec<Vec<String>> = rows
192 .iter()
193 .map(|row| {
194 (0..header.len())
195 .map(|i| cell(row, i).chars().take(CELL_WIDTH).collect())
196 .collect()
197 })
198 .collect();
199 let mut widths: Vec<usize> = header.iter().map(|h| h.chars().count()).collect();
200 for row in &body {
201 for (width, value) in widths.iter_mut().zip(row) {
202 *width = (*width).max(value.chars().count());
203 }
204 }
205 let line = |cells: &[String]| {
206 let padded: Vec<String> = cells
207 .iter()
208 .zip(&widths)
209 .map(|(c, w)| format!("{c:<w$}"))
210 .collect();
211 padded.join(" | ").trim_end().to_owned()
212 };
213 let rule: Vec<String> = widths.iter().map(|w| "-".repeat(*w)).collect();
214 let mut out = format!("{}\n{}\n", line(&header), rule.join("-+-"));
215 for row in &body {
216 out.push_str(&line(row));
217 out.push('\n');
218 }
219 out.push_str(&format!("({} row(s))\n", body.len()));
220 out
221}