Skip to main content

renox_core/
shell.rs

1//! `my-app db:shell`: a small SQL prompt on the app's database, so neither
2//! `sqlite3` nor `psql` is needed (on servers or Windows).
3//!
4//! Statements end with `;` and may span lines. `.tables` lists tables and
5//! `.quit` (or end of input) leaves. Input can be piped:
6//! `echo "SELECT count(*) FROM users;" | my-app db:shell`.
7
8use 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
16/// Longest cell shown; longer values are cut.
17const 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/// The shell over any input and output, e.g. for tests. `interactive` shows
32/// a banner and prompts.
33#[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
105/// Runs one statement and returns what to print.
106async 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
185/// Rows as an aligned text table with a header and a count.
186fn 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}