diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..5251da8 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,42 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +env: + CARGO_TERM_COLOR: always + RUSTFLAGS: "-D warnings" + +jobs: + fmt: + name: rustfmt + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt + - run: cargo fmt --all -- --check + + clippy: + name: clippy + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + with: + components: clippy + - uses: Swatinem/rust-cache@v2 + - run: cargo clippy --workspace --all-targets -- -D warnings + + test: + name: test + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + - uses: Swatinem/rust-cache@v2 + - run: cargo test --workspace diff --git a/README.md b/README.md index f9bc91a..5efd417 100644 --- a/README.md +++ b/README.md @@ -31,7 +31,7 @@ libsql SDK (Rust, JS, Python, Go) # Start with a MySQL frontend litewire --mysql-listen 127.0.0.1:3306 --db app.db -# Start with all frontends +# Start with all frontends (postgres + tds require --features postgres,tds at build time) litewire --mysql-listen 127.0.0.1:3306 --postgres-listen 127.0.0.1:5432 --tds-listen 127.0.0.1:1433 --hrana-listen 127.0.0.1:8080 --db app.db # Connect from any MySQL client @@ -110,17 +110,29 @@ litewire translates MySQL and PostgreSQL SQL dialects to SQLite on the fly: | `SET NAMES utf8mb4` / `SET NOCOUNT ON` | No-op | | Backtick / `[bracket]` quoting | Passed through or converted | -See [docs/architecture.md](docs/architecture.md) for the full translation reference. - -## Tested With - -- WordPress (via `pdo_mysql`) -- Laravel (via `pdo_mysql` / `pdo_pgsql` / `pdo_sqlsrv`) -- Drupal -- `mysql` CLI -- `psql` CLI -- `sqlcmd` CLI -- DBeaver, pgAdmin, SSMS, TablePlus +See [docs/architecture.md](docs/architecture.md) for the full architecture and translation reference. + +## Compatibility + +The MySQL frontend is exercised end-to-end by an in-process test suite +(`crates/litewire/tests/mysql_e2e.rs`) that drives the wire protocol via +`mysql_async` -- CRUD, prepared statements, transactions (`START TRANSACTION` / +`BEGIN` / `COMMIT` / `ROLLBACK`), `LAST_INSERT_ID()`, `SHOW TABLES`, +`DESCRIBE`, `INFORMATION_SCHEMA` probes, `SET NAMES` / `SET autocommit`, and +the metadata queries used at connection setup. + +The PostgreSQL and TDS frontends are wire-compatible enough for basic CRUD +against `psql` / `sqlcmd` and the extended-query flow used by `pdo_pgsql` / +`pdo_sqlsrv`; the TDS frontend is **experimental** -- authentication is +simplified, the type coverage is a subset (BigInt / Float8 / NVARCHAR / +VarBinary), and the SSL handshake is not implemented. Real SQL Server tools +(SSMS, sqlcmd with encryption) will not connect until those land. + +Anywhere you would normally point at MySQL/PG/SQL Server -- PHP PDO drivers, +`mysql` / `psql` / `sqlcmd` CLIs, DBeaver, pgAdmin -- should work for standard +CRUD workloads. Anything that depends on server-side features SQLite doesn't +have (stored procedures, `SQL_CALC_FOUND_ROWS`, `LOCK TABLES` isolation, +row-level locking semantics, dollar-quoted PL/pgSQL bodies, etc.) will not. ## Limitations diff --git a/crates/litewire-backend/src/hrana_client.rs b/crates/litewire-backend/src/hrana_client.rs index fa061f4..922cb80 100644 --- a/crates/litewire-backend/src/hrana_client.rs +++ b/crates/litewire-backend/src/hrana_client.rs @@ -236,9 +236,7 @@ fn value_to_hrana(val: &Value) -> HranaValue { value: i.to_string(), }, Value::Float(f) => HranaValue::Float { value: *f }, - Value::Text(s) => HranaValue::Text { - value: s.clone(), - }, + Value::Text(s) => HranaValue::Text { value: s.clone() }, Value::Blob(b) => { use base64::Engine; HranaValue::Blob { @@ -375,17 +373,13 @@ mod tests { ], rows: vec![ vec![ - ResponseValue::Integer { - value: "1".into(), - }, + ResponseValue::Integer { value: "1".into() }, ResponseValue::Text { value: "alice".into(), }, ], vec![ - ResponseValue::Integer { - value: "2".into(), - }, + ResponseValue::Integer { value: "2".into() }, ResponseValue::Text { value: "bob".into(), }, diff --git a/crates/litewire-backend/src/lib.rs b/crates/litewire-backend/src/lib.rs index d1e617c..c391fd3 100644 --- a/crates/litewire-backend/src/lib.rs +++ b/crates/litewire-backend/src/lib.rs @@ -191,8 +191,14 @@ mod tests { fn empty_result_set() { let rs = ResultSet { columns: vec![ - Column { name: "a".into(), decltype: None }, - Column { name: "b".into(), decltype: None }, + Column { + name: "a".into(), + decltype: None, + }, + Column { + name: "b".into(), + decltype: None, + }, ], rows: vec![], }; diff --git a/crates/litewire-backend/src/rusqlite_backend.rs b/crates/litewire-backend/src/rusqlite_backend.rs index de086d2..ea9be27 100644 --- a/crates/litewire-backend/src/rusqlite_backend.rs +++ b/crates/litewire-backend/src/rusqlite_backend.rs @@ -42,8 +42,7 @@ impl Rusqlite { /// /// Returns an error if the database cannot be opened. pub fn memory() -> Result { - let conn = - Connection::open_in_memory().map_err(|e| BackendError::Sqlite(e.to_string()))?; + let conn = Connection::open_in_memory().map_err(|e| BackendError::Sqlite(e.to_string()))?; Ok(Self { conn: Arc::new(Mutex::new(conn)), }) @@ -73,9 +72,7 @@ fn extract_value(row: &rusqlite::Row<'_>, idx: usize) -> Result Ok(Value::Null), ValueRef::Integer(i) => Ok(Value::Integer(i)), ValueRef::Real(f) => Ok(Value::Float(f)), - ValueRef::Text(s) => Ok(Value::Text( - String::from_utf8_lossy(s).into_owned(), - )), + ValueRef::Text(s) => Ok(Value::Text(String::from_utf8_lossy(s).into_owned())), ValueRef::Blob(b) => Ok(Value::Blob(b.to_vec())), } } @@ -110,12 +107,14 @@ impl Backend for Rusqlite { .query(param_refs.as_slice()) .map_err(|e| BackendError::Sqlite(e.to_string()))?; - while let Some(row) = rows.next().map_err(|e| BackendError::Sqlite(e.to_string()))? { + while let Some(row) = rows + .next() + .map_err(|e| BackendError::Sqlite(e.to_string()))? + { let mut values = Vec::with_capacity(col_count); for i in 0..col_count { values.push( - extract_value(row, i) - .map_err(|e| BackendError::Sqlite(e.to_string()))?, + extract_value(row, i).map_err(|e| BackendError::Sqlite(e.to_string()))?, ); } result_rows.push(values); @@ -193,7 +192,10 @@ mod tests { .unwrap(); assert_eq!(result.last_insert_rowid, Some(2)); - let rs = backend.query("SELECT id, name FROM users ORDER BY id", &[]).await.unwrap(); + let rs = backend + .query("SELECT id, name FROM users ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rs.columns.len(), 2); assert_eq!(rs.columns[0].name, "id"); assert_eq!(rs.columns[1].name, "name"); @@ -243,7 +245,10 @@ mod tests { #[tokio::test] async fn null_handling() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (v TEXT)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (v TEXT)", &[]) + .await + .unwrap(); backend .execute("INSERT INTO t VALUES (?1)", &[Value::Null]) .await @@ -286,10 +291,10 @@ mod tests { .unwrap(); let rs = backend - .query("SELECT * FROM t WHERE a = ?1 AND b = ?2", &[ - Value::Integer(1), - Value::Text("hello".into()), - ]) + .query( + "SELECT * FROM t WHERE a = ?1 AND b = ?2", + &[Value::Integer(1), Value::Text("hello".into())], + ) .await .unwrap(); assert_eq!(rs.rows.len(), 1); @@ -386,11 +391,17 @@ mod tests { async fn column_names_preserved() { let backend = Rusqlite::memory().unwrap(); backend - .execute("CREATE TABLE users (id INTEGER, name TEXT, email TEXT)", &[]) + .execute( + "CREATE TABLE users (id INTEGER, name TEXT, email TEXT)", + &[], + ) .await .unwrap(); - let rs = backend.query("SELECT id, name, email FROM users", &[]).await.unwrap(); + let rs = backend + .query("SELECT id, name, email FROM users", &[]) + .await + .unwrap(); assert_eq!(rs.columns[0].name, "id"); assert_eq!(rs.columns[1].name, "name"); assert_eq!(rs.columns[2].name, "email"); @@ -399,7 +410,10 @@ mod tests { #[tokio::test] async fn query_with_alias() { let backend = Rusqlite::memory().unwrap(); - let rs = backend.query("SELECT 1 AS num, 'hello' AS greeting", &[]).await.unwrap(); + let rs = backend + .query("SELECT 1 AS num, 'hello' AS greeting", &[]) + .await + .unwrap(); assert_eq!(rs.columns[0].name, "num"); assert_eq!(rs.columns[1].name, "greeting"); assert_eq!(rs.rows[0][0], Value::Integer(1)); diff --git a/crates/litewire-backend/tests/hrana_client_integration.rs b/crates/litewire-backend/tests/hrana_client_integration.rs index 4061844..176762a 100644 --- a/crates/litewire-backend/tests/hrana_client_integration.rs +++ b/crates/litewire-backend/tests/hrana_client_integration.rs @@ -56,7 +56,10 @@ async fn create_table_and_insert() { // Create table let result = client - .execute("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)", &[]) + .execute( + "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)", + &[], + ) .await .expect("CREATE TABLE failed"); assert_eq!(result.affected_rows, 0); @@ -89,7 +92,10 @@ async fn query_rows() { let (client, _server) = start_server().await; client - .execute("CREATE TABLE items (id INTEGER PRIMARY KEY, label TEXT, price REAL)", &[]) + .execute( + "CREATE TABLE items (id INTEGER PRIMARY KEY, label TEXT, price REAL)", + &[], + ) .await .unwrap(); client @@ -183,7 +189,10 @@ async fn blob_roundtrip() { let (client, _server) = start_server().await; client - .execute("CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", &[]) + .execute( + "CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", + &[], + ) .await .unwrap(); @@ -234,9 +243,7 @@ async fn null_values() { async fn sql_error_returns_backend_error() { let (client, _server) = start_server().await; - let result = client - .query("SELECT * FROM nonexistent_table", &[]) - .await; + let result = client.query("SELECT * FROM nonexistent_table", &[]).await; assert!(result.is_err()); let err = result.unwrap_err().to_string(); @@ -255,18 +262,12 @@ async fn update_and_delete() { .await .unwrap(); client - .execute( - "INSERT INTO counters VALUES ('hits', 0)", - &[], - ) + .execute("INSERT INTO counters VALUES ('hits', 0)", &[]) .await .unwrap(); let result = client - .execute( - "UPDATE counters SET val = val + 1 WHERE name = 'hits'", - &[], - ) + .execute("UPDATE counters SET val = val + 1 WHERE name = 'hits'", &[]) .await .unwrap(); assert_eq!(result.affected_rows, 1); diff --git a/crates/litewire-hrana/src/http.rs b/crates/litewire-hrana/src/http.rs index e71bad7..b164170 100644 --- a/crates/litewire-hrana/src/http.rs +++ b/crates/litewire-hrana/src/http.rs @@ -74,11 +74,8 @@ async fn execute_stmt( backend: &SharedBackend, stmt: &StmtRequest, ) -> Result { - let params: Vec = stmt - .args - .iter() - .map(|a| a.to_backend_value()) - .collect(); + let params: Vec = + stmt.args.iter().map(|a| a.to_backend_value()).collect(); // Hrana sends SQLite SQL natively -- no translation needed. let sql_upper = stmt.sql.trim().to_ascii_uppercase(); @@ -257,10 +254,7 @@ mod tests { .await .unwrap(); backend - .execute( - "INSERT INTO t VALUES (1, 'Alice')", - &[], - ) + .execute("INSERT INTO t VALUES (1, 'Alice')", &[]) .await .unwrap(); @@ -356,10 +350,12 @@ mod tests { let body = axum::body::to_bytes(resp.into_body(), 8192).await.unwrap(); let resp: serde_json::Value = serde_json::from_slice(&body).unwrap(); assert_eq!(resp["results"][0]["type"], "error"); - assert!(resp["results"][0]["error"]["message"] - .as_str() - .unwrap() - .contains("nonexistent_table")); + assert!( + resp["results"][0]["error"]["message"] + .as_str() + .unwrap() + .contains("nonexistent_table") + ); } #[tokio::test] diff --git a/crates/litewire-hrana/src/lib.rs b/crates/litewire-hrana/src/lib.rs index 46d36f3..4b72470 100644 --- a/crates/litewire-hrana/src/lib.rs +++ b/crates/litewire-hrana/src/lib.rs @@ -7,8 +7,8 @@ mod http; mod types; -use std::net::SocketAddr; use litewire_backend::SharedBackend; +use std::net::SocketAddr; use tracing::info; /// Configuration for the Hrana HTTP frontend. diff --git a/crates/litewire-hrana/src/types.rs b/crates/litewire-hrana/src/types.rs index 934a017..b887954 100644 --- a/crates/litewire-hrana/src/types.rs +++ b/crates/litewire-hrana/src/types.rs @@ -47,9 +47,7 @@ impl HranaValue { pub fn to_backend_value(&self) -> Value { match self { Self::Null => Value::Null, - Self::Integer { value } => { - Value::Integer(value.parse().unwrap_or(0)) - } + Self::Integer { value } => Value::Integer(value.parse().unwrap_or(0)), Self::Float { value } => Value::Float(*value), Self::Text { value } => Value::Text(value.clone()), Self::Blob { base64 } => { @@ -119,9 +117,7 @@ impl ResponseValue { value: i.to_string(), }, Value::Float(f) => Self::Float { value: *f }, - Value::Text(s) => Self::Text { - value: s.clone(), - }, + Value::Text(s) => Self::Text { value: s.clone() }, Value::Blob(b) => { use base64::Engine; Self::Blob { @@ -153,9 +149,7 @@ mod tests { #[test] fn integer_to_backend() { - let v = HranaValue::Integer { - value: "42".into(), - }; + let v = HranaValue::Integer { value: "42".into() }; assert!(matches!(v.to_backend_value(), Value::Integer(42))); } @@ -297,9 +291,7 @@ mod tests { name: "id".into(), decltype: Some("INTEGER".into()), }], - rows: vec![vec![ResponseValue::Integer { - value: "1".into(), - }]], + rows: vec![vec![ResponseValue::Integer { value: "1".into() }]], affected_row_count: 0, last_insert_rowid: None, }, diff --git a/crates/litewire-mysql/src/error_map.rs b/crates/litewire-mysql/src/error_map.rs new file mode 100644 index 0000000..e2aa2d8 --- /dev/null +++ b/crates/litewire-mysql/src/error_map.rs @@ -0,0 +1,149 @@ +//! Map litewire-backend error messages to real MySQL error codes. +//! +//! `BackendError` is a stringly-typed pass-through of the underlying +//! rusqlite error (see `litewire-backend`), so this module works by matching +//! substrings of the message text against the shape of `rusqlite::Error`'s +//! Display impl and the SQLite constraint-failure message conventions. +//! +//! This is deliberately conservative: any error we can't classify falls back +//! to `ER_UNKNOWN_ERROR` (1105 / HY000) so callers still see the raw text. +//! +//! Reference: + +use opensrv_mysql::ErrorKind; + +/// The full MySQL error triple: code + SQLSTATE + message. +#[derive(Debug, Clone)] +pub struct MysqlError { + /// MySQL error code (e.g. 1062 for duplicate entry). + pub code: ErrorKind, + /// SQLSTATE (5 chars, e.g. "23000"). Not read by production code -- the + /// wire packet SQLSTATE is derived from `ErrorKind::sqlstate()` inside + /// `opensrv-mysql`. Retained on the struct so tests can pin down the exact + /// SQLSTATE we intend each mapping to produce and so future callers that + /// want to log both the code and the SQLSTATE can do so from one place. + #[cfg_attr(not(test), allow(dead_code))] + pub sqlstate: [u8; 5], + /// Human-readable message, forwarded verbatim from the backend. + pub message: String, +} + +/// Classify a backend error string into a `MysqlError`. +/// +/// This function is pure and infallible; unknown errors return +/// `ER_UNKNOWN_ERROR` (MySQL 1105). +#[must_use] +pub fn classify(err_msg: &str) -> MysqlError { + let lower = err_msg.to_ascii_lowercase(); + + // -- Locking / busy -------------------------------------------------------- + // SQLITE_BUSY / SQLITE_LOCKED -> MySQL 1205 "Lock wait timeout exceeded" + // (SQLSTATE HY000). This is the closest analogue clients will actually + // retry on. + if lower.contains("database is locked") + || lower.contains("database table is locked") + || lower.contains("sqlite_busy") + || lower.contains("sqlite_locked") + { + return MysqlError { + code: ErrorKind::ER_LOCK_WAIT_TIMEOUT, + sqlstate: *b"HY000", + message: err_msg.to_string(), + }; + } + + // -- Constraint violations ------------------------------------------------- + // Unique / primary key -> 1062 (SQLSTATE 23000). + if lower.contains("unique constraint failed") || lower.contains("primary key constraint failed") + { + return MysqlError { + code: ErrorKind::ER_DUP_ENTRY, + sqlstate: *b"23000", + message: err_msg.to_string(), + }; + } + + // Foreign key -> 1452 (SQLSTATE 23000). + if lower.contains("foreign key constraint failed") { + return MysqlError { + code: ErrorKind::ER_NO_REFERENCED_ROW_2, + sqlstate: *b"23000", + message: err_msg.to_string(), + }; + } + + // -- Read-only ------------------------------------------------------------ + // SQLITE_READONLY -> 1290 "The MySQL server is running with the ... + // --read-only option so it cannot execute this statement" (HY000). + if lower.contains("attempt to write a readonly database") + || lower.contains("readonly database") + || lower.contains("sqlite_readonly") + { + return MysqlError { + code: ErrorKind::ER_OPTION_PREVENTS_STATEMENT, + sqlstate: *b"HY000", + message: err_msg.to_string(), + }; + } + + // Fallback. + MysqlError { + code: ErrorKind::ER_UNKNOWN_ERROR, + sqlstate: *b"HY000", + message: err_msg.to_string(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn unique_constraint_maps_to_1062() { + let e = classify("UNIQUE constraint failed: users.email"); + assert!(matches!(e.code, ErrorKind::ER_DUP_ENTRY)); + assert_eq!(&e.sqlstate, b"23000"); + } + + #[test] + fn primary_key_constraint_maps_to_1062() { + let e = classify("PRIMARY KEY constraint failed: users.id"); + assert!(matches!(e.code, ErrorKind::ER_DUP_ENTRY)); + assert_eq!(&e.sqlstate, b"23000"); + } + + #[test] + fn foreign_key_maps_to_1452() { + let e = classify("FOREIGN KEY constraint failed"); + assert!(matches!(e.code, ErrorKind::ER_NO_REFERENCED_ROW_2)); + assert_eq!(&e.sqlstate, b"23000"); + } + + #[test] + fn busy_maps_to_1205() { + let e = classify("database is locked"); + assert!(matches!(e.code, ErrorKind::ER_LOCK_WAIT_TIMEOUT)); + assert_eq!(&e.sqlstate, b"HY000"); + } + + #[test] + fn readonly_maps_to_1290() { + let e = classify("attempt to write a readonly database"); + assert!(matches!(e.code, ErrorKind::ER_OPTION_PREVENTS_STATEMENT)); + assert_eq!(&e.sqlstate, b"HY000"); + } + + #[test] + fn unknown_falls_back_to_1105() { + let e = classify("no such table: sprockets"); + assert!(matches!(e.code, ErrorKind::ER_UNKNOWN_ERROR)); + assert_eq!(&e.sqlstate, b"HY000"); + } + + #[test] + fn classify_preserves_message() { + let msg = "UNIQUE constraint failed: users.email"; + let e = classify(msg); + assert_eq!(e.message, msg); + } +} diff --git a/crates/litewire-mysql/src/handler.rs b/crates/litewire-mysql/src/handler.rs index ccb8b40..06ed08c 100644 --- a/crates/litewire-mysql/src/handler.rs +++ b/crates/litewire-mysql/src/handler.rs @@ -11,6 +11,16 @@ use opensrv_mysql::*; use tokio::io::AsyncWrite; use tracing::{debug, warn}; +use crate::error_map; + +/// Maximum number of prepared statements a single connection may hold at once. +/// +/// MySQL's default `max_prepared_stmt_count` is 16382 (global, not per-connection), +/// but here it's per-connection because litewire has no global registry. 1024 +/// per connection is generous for real workloads and prevents a runaway client +/// from exhausting memory via COM_STMT_PREPARE without matching COM_STMT_CLOSE. +const MAX_PREPARED_STMTS_PER_CONN: usize = 1024; + /// Build an `OkResponse` with the correct transaction status flag. fn ok_response(affected_rows: u64, last_insert_id: u64, in_transaction: bool) -> OkResponse { let status_flags = if in_transaction { @@ -34,8 +44,6 @@ struct PreparedStmt { sqlite_sql: String, /// Whether this is a query (SELECT) or mutation (INSERT/UPDATE/DELETE). kind: StatementKind, - /// Number of `?` parameters. - param_count: usize, } /// Handler for a single MySQL client connection. @@ -96,11 +104,7 @@ impl LiteWireHandler { rw.finish().await } - Err(e) => { - results - .error(ErrorKind::ER_UNKNOWN_ERROR, e.to_string().as_bytes()) - .await - } + Err(e) => write_backend_error(results, &e.to_string()).await, } } @@ -113,18 +117,17 @@ impl LiteWireHandler { ) -> Result<(), std::io::Error> { match self.backend.execute(sql, params).await { Ok(r) => { - let resp = ok_response( - r.affected_rows, - r.last_insert_rowid.unwrap_or(0) as u64, - self.in_transaction, - ); + // last_insert_rowid comes back as i64 -- clamp negatives (should + // never happen; SQLite rowids are always >= 1 for a real insert) + // to 0 rather than reinterpret via `as u64`. + let last_id_u64: u64 = r + .last_insert_rowid + .and_then(|v| u64::try_from(v.max(0)).ok()) + .unwrap_or(0); + let resp = ok_response(r.affected_rows, last_id_u64, self.in_transaction); results.completed(resp).await } - Err(e) => { - results - .error(ErrorKind::ER_UNKNOWN_ERROR, e.to_string().as_bytes()) - .await - } + Err(e) => write_backend_error(results, &e.to_string()).await, } } @@ -145,18 +148,14 @@ impl LiteWireHandler { let resp = ok_response(0, 0, self.in_transaction); results.completed(resp).await } - Err(e) => { - results - .error(ErrorKind::ER_UNKNOWN_ERROR, e.to_string().as_bytes()) - .await - } + Err(e) => write_backend_error(results, &e.to_string()).await, } } /// Translate SQL and return the first translated result, or an error string. fn translate_sql(&self, query: &str) -> Result<(String, StatementKind), String> { - let translated = litewire_translate::translate(query, Dialect::MySQL) - .map_err(|e| e.to_string())?; + let translated = + litewire_translate::translate(query, Dialect::MySQL).map_err(|e| e.to_string())?; let Some(result) = translated.into_iter().next() else { return Ok((String::new(), StatementKind::Other)); @@ -176,6 +175,16 @@ impl LiteWireHandler { } } +/// Convert a backend error string into a MySQL error packet with a specific +/// error code + SQLSTATE (via [`crate::error_map::classify`]) and send it. +async fn write_backend_error( + results: QueryResultWriter<'_, W>, + err_msg: &str, +) -> Result<(), std::io::Error> { + let mapped = error_map::classify(err_msg); + results.error(mapped.code, mapped.message.as_bytes()).await +} + /// Convert an opensrv-mysql parameter value to our backend Value type. fn param_to_value(param: ParamValue<'_>) -> Value { match param.value.into_inner() { @@ -208,9 +217,7 @@ impl AsyncMysqlShim for LiteWireHandler { let (sqlite_sql, kind) = match self.translate_sql(query) { Ok(r) => r, Err(e) => { - return info - .error(ErrorKind::ER_PARSE_ERROR, e.as_bytes()) - .await; + return info.error(ErrorKind::ER_PARSE_ERROR, e.as_bytes()).await; } }; @@ -246,18 +253,32 @@ impl AsyncMysqlShim for LiteWireHandler { vec![] }; + // Bound the per-connection prepared-statement cache so a client that + // never sends COM_STMT_CLOSE can't wedge the process. Return the same + // error code (1461) real MySQL uses when max_prepared_stmt_count is hit. + if self.stmts.len() >= MAX_PREPARED_STMTS_PER_CONN { + warn!( + stmts = self.stmts.len(), + "prepared-statement cap hit ({MAX_PREPARED_STMTS_PER_CONN}); rejecting COM_STMT_PREPARE" + ); + return info + .error( + ErrorKind::ER_MAX_PREPARED_STMT_COUNT_REACHED, + format!( + "Can't create more than {MAX_PREPARED_STMTS_PER_CONN} prepared statements \ + on this connection" + ) + .as_bytes(), + ) + .await; + } + // Assign a statement ID and cache it. let stmt_id = self.next_stmt_id; self.next_stmt_id += 1; - self.stmts.insert( - stmt_id, - PreparedStmt { - sqlite_sql, - kind, - param_count, - }, - ); + self.stmts + .insert(stmt_id, PreparedStmt { sqlite_sql, kind }); info.reply(stmt_id, ¶ms, &columns).await } @@ -335,9 +356,7 @@ impl AsyncMysqlShim for LiteWireHandler { let kind = classify(&sqlite_sql); match kind { StatementKind::Query => self.do_query(&sqlite_sql, &[], results).await, - StatementKind::Transaction => { - self.do_transaction(&sqlite_sql, results).await - } + StatementKind::Transaction => self.do_transaction(&sqlite_sql, results).await, _ => self.do_execute(&sqlite_sql, &[], results).await, } } diff --git a/crates/litewire-mysql/src/lib.rs b/crates/litewire-mysql/src/lib.rs index cf9f15c..a224b79 100644 --- a/crates/litewire-mysql/src/lib.rs +++ b/crates/litewire-mysql/src/lib.rs @@ -4,6 +4,7 @@ //! incoming SQL from MySQL dialect to SQLite, executes against the backend, //! and returns results in MySQL wire format. +mod error_map; mod handler; mod resultset; mod types; diff --git a/crates/litewire-postgres/src/error_map.rs b/crates/litewire-postgres/src/error_map.rs new file mode 100644 index 0000000..1709255 --- /dev/null +++ b/crates/litewire-postgres/src/error_map.rs @@ -0,0 +1,146 @@ +//! Map litewire-backend error messages to PostgreSQL SQLSTATE codes. +//! +//! `BackendError` is a stringly-typed pass-through of the underlying +//! rusqlite error; this module matches on the same substrings the MySQL +//! frontend's `error_map` uses and produces PG SQLSTATE codes instead. +//! +//! Reference: + +/// A pgwire-shaped error identifier: SQLSTATE + human-readable message. +#[derive(Debug, Clone)] +pub struct PgError { + /// 5-char SQLSTATE, e.g. `"23505"`. + pub sqlstate: &'static str, + /// Human-readable message, forwarded verbatim from the backend. + pub message: String, +} + +/// Classify a backend error string into a `PgError`. +/// +/// Unknown errors return SQLSTATE `XX000` ("internal error"). +#[must_use] +pub fn classify(err_msg: &str) -> PgError { + let lower = err_msg.to_ascii_lowercase(); + + // Unique / primary key violation -> 23505 unique_violation + if lower.contains("unique constraint failed") || lower.contains("primary key constraint failed") + { + return PgError { + sqlstate: "23505", + message: err_msg.to_string(), + }; + } + + // Foreign key violation -> 23503 foreign_key_violation + if lower.contains("foreign key constraint failed") { + return PgError { + sqlstate: "23503", + message: err_msg.to_string(), + }; + } + + // NOT NULL violation -> 23502 not_null_violation + if lower.contains("not null constraint failed") { + return PgError { + sqlstate: "23502", + message: err_msg.to_string(), + }; + } + + // CHECK constraint -> 23514 check_violation + if lower.contains("check constraint failed") { + return PgError { + sqlstate: "23514", + message: err_msg.to_string(), + }; + } + + // SQLITE_BUSY / SQLITE_LOCKED -> 55P03 lock_not_available + if lower.contains("database is locked") + || lower.contains("database table is locked") + || lower.contains("sqlite_busy") + || lower.contains("sqlite_locked") + { + return PgError { + sqlstate: "55P03", + message: err_msg.to_string(), + }; + } + + // SQLITE_READONLY -> 25006 read_only_sql_transaction + if lower.contains("attempt to write a readonly database") + || lower.contains("readonly database") + || lower.contains("sqlite_readonly") + { + return PgError { + sqlstate: "25006", + message: err_msg.to_string(), + }; + } + + // Fallback: internal_error + PgError { + sqlstate: "XX000", + message: err_msg.to_string(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn unique_maps_to_23505() { + let e = classify("UNIQUE constraint failed: users.email"); + assert_eq!(e.sqlstate, "23505"); + } + + #[test] + fn primary_key_maps_to_23505() { + let e = classify("PRIMARY KEY constraint failed: users.id"); + assert_eq!(e.sqlstate, "23505"); + } + + #[test] + fn foreign_key_maps_to_23503() { + let e = classify("FOREIGN KEY constraint failed"); + assert_eq!(e.sqlstate, "23503"); + } + + #[test] + fn not_null_maps_to_23502() { + let e = classify("NOT NULL constraint failed: users.name"); + assert_eq!(e.sqlstate, "23502"); + } + + #[test] + fn check_constraint_maps_to_23514() { + let e = classify("CHECK constraint failed: users_age_check"); + assert_eq!(e.sqlstate, "23514"); + } + + #[test] + fn busy_maps_to_55p03() { + let e = classify("database is locked"); + assert_eq!(e.sqlstate, "55P03"); + } + + #[test] + fn readonly_maps_to_25006() { + let e = classify("attempt to write a readonly database"); + assert_eq!(e.sqlstate, "25006"); + } + + #[test] + fn unknown_maps_to_xx000() { + let e = classify("something random"); + assert_eq!(e.sqlstate, "XX000"); + } + + #[test] + fn classify_preserves_message() { + let msg = "UNIQUE constraint failed: users.email"; + let e = classify(msg); + assert_eq!(e.message, msg); + } +} diff --git a/crates/litewire-postgres/src/handler.rs b/crates/litewire-postgres/src/handler.rs index 203a2a2..db9bc89 100644 --- a/crates/litewire-postgres/src/handler.rs +++ b/crates/litewire-postgres/src/handler.rs @@ -11,8 +11,8 @@ use futures::stream; use pgwire::api::portal::{Format, Portal}; use pgwire::api::query::{ExtendedQueryHandler, SimpleQueryHandler}; use pgwire::api::results::{ - DataRowEncoder, DescribePortalResponse, DescribeStatementResponse, FieldInfo, - QueryResponse, Response, Tag, + DataRowEncoder, DescribePortalResponse, DescribeStatementResponse, FieldInfo, QueryResponse, + Response, Tag, }; use pgwire::api::stmt::{NoopQueryParser, StoredStatement}; use pgwire::api::{ClientInfo, Type}; @@ -23,6 +23,7 @@ use tracing::{debug, warn}; use litewire_backend::{SharedBackend, Value}; use litewire_translate::{self, Dialect, StatementKind, TranslateResult, classify}; +use crate::error_map; use crate::types::sqlite_to_pg_type; /// Handler for a single PostgreSQL client connection. @@ -41,8 +42,8 @@ impl PostgresHandler { /// Translate SQL from PostgreSQL dialect to SQLite and classify it. fn translate_sql(&self, query: &str) -> Result<(String, StatementKind), String> { - let translated = litewire_translate::translate(query, Dialect::PostgreSQL) - .map_err(|e| e.to_string())?; + let translated = + litewire_translate::translate(query, Dialect::PostgreSQL).map_err(|e| e.to_string())?; let Some(result) = translated.into_iter().next() else { return Ok((String::new(), StatementKind::Other)); @@ -86,7 +87,7 @@ impl PostgresHandler { .backend .query(sql, params) .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + .map_err(|e| pg_backend_error(&e))?; // Build field info, inferring types from declared type first, // then falling back to the first row's actual values for expressions @@ -146,7 +147,7 @@ impl PostgresHandler { self.backend .execute(sql, params) .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + .map_err(|e| pg_backend_error(&e))?; let upper = sql.trim().to_ascii_uppercase(); if upper.starts_with("BEGIN") || upper.starts_with("START") { @@ -165,7 +166,7 @@ impl PostgresHandler { .backend .execute(sql, params) .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + .map_err(|e| pg_backend_error(&e))?; let tag_name = match kind { StatementKind::Mutation => { @@ -207,11 +208,7 @@ impl PostgresHandler { /// /// Uses LIMIT 1 (not LIMIT 0) so we can infer types from actual data /// when columns lack declared types (e.g. `SELECT 1 + 2`). - async fn probe_columns( - &self, - sql: &str, - format: &Format, - ) -> PgWireResult> { + async fn probe_columns(&self, sql: &str, format: &Format) -> PgWireResult> { let probe = format!("{sql} LIMIT 1"); match self.backend.query(&probe, &[]).await { Ok(rs) => Ok(rs @@ -256,11 +253,7 @@ fn value_to_pg_type(val: &Value) -> Type { } /// Encode a single backend `Value` into a pgwire `DataRowEncoder`. -fn encode_value( - encoder: &mut DataRowEncoder, - val: &Value, - field: &FieldInfo, -) -> PgWireResult<()> { +fn encode_value(encoder: &mut DataRowEncoder, val: &Value, field: &FieldInfo) -> PgWireResult<()> { match val { Value::Null => encoder.encode_field(&None::), Value::Integer(i) => { @@ -312,7 +305,7 @@ fn extract_params(portal: &Portal) -> Vec { .parameter::(i, ¶m_type) .ok() .flatten() - .map_or(Value::Null, |v| Value::Integer(v)), + .map_or(Value::Null, Value::Integer), t if *t == Type::FLOAT4 => portal .parameter::(i, ¶m_type) @@ -324,7 +317,7 @@ fn extract_params(portal: &Portal) -> Vec { .parameter::(i, ¶m_type) .ok() .flatten() - .map_or(Value::Null, |v| Value::Float(v)), + .map_or(Value::Null, Value::Float), t if *t == Type::BYTEA => portal .parameter::>(i, ¶m_type) @@ -344,7 +337,9 @@ fn extract_params(portal: &Portal) -> Vec { values } -/// Build a PgWireError from an error string. +/// Build a PgWireError from an arbitrary error string with SQLSTATE `XX000` +/// (internal_error). Prefer [`pg_backend_error`] for messages that came from +/// the backend so they get classified into a specific SQLSTATE. fn pg_error(msg: &str) -> PgWireError { PgWireError::UserError(Box::new(ErrorInfo::new( "ERROR".to_owned(), @@ -353,6 +348,17 @@ fn pg_error(msg: &str) -> PgWireError { ))) } +/// Build a PgWireError from a litewire-backend error, classifying it into a +/// real PostgreSQL SQLSTATE (see [`crate::error_map`]). +fn pg_backend_error(err: &litewire_backend::BackendError) -> PgWireError { + let mapped = error_map::classify(&err.to_string()); + PgWireError::UserError(Box::new(ErrorInfo::new( + "ERROR".to_owned(), + mapped.sqlstate.to_owned(), + mapped.message, + ))) +} + #[async_trait] impl SimpleQueryHandler for PostgresHandler { async fn do_query<'a, C>( @@ -365,8 +371,8 @@ impl SimpleQueryHandler for PostgresHandler { { debug!(sql = %query, "PG simple query"); - let translated = litewire_translate::translate(query, Dialect::PostgreSQL) - .map_err(|e| { + let translated = + litewire_translate::translate(query, Dialect::PostgreSQL).map_err(|e| { warn!("SQL translation error: {e}"); pg_error(&e.to_string()) })?; @@ -375,9 +381,7 @@ impl SimpleQueryHandler for PostgresHandler { for result in translated { let resp = match result { - TranslateResult::Noop => { - Response::Execution(Tag::new("SET")) - } + TranslateResult::Noop => Response::Execution(Tag::new("SET")), TranslateResult::Metadata(meta) => { let sqlite_sql = meta.to_sqlite_sql(); self.exec_query(&sqlite_sql, &[], &Format::UnifiedText) diff --git a/crates/litewire-postgres/src/lib.rs b/crates/litewire-postgres/src/lib.rs index ac0f037..13640ea 100644 --- a/crates/litewire-postgres/src/lib.rs +++ b/crates/litewire-postgres/src/lib.rs @@ -4,6 +4,7 @@ //! incoming SQL from PostgreSQL dialect to SQLite, executes against the //! backend, and returns results in PostgreSQL wire format. +mod error_map; mod handler; mod types; @@ -11,10 +12,10 @@ use std::net::SocketAddr; use std::sync::Arc; use litewire_backend::SharedBackend; +use pgwire::api::NoopErrorHandler; +use pgwire::api::PgWireServerHandlers; use pgwire::api::auth::noop::NoopStartupHandler; use pgwire::api::copy::NoopCopyHandler; -use pgwire::api::PgWireServerHandlers; -use pgwire::api::NoopErrorHandler; use pgwire::tokio::process_socket; use tokio::net::TcpListener; use tracing::{debug, info, warn}; diff --git a/crates/litewire-tds/src/handler.rs b/crates/litewire-tds/src/handler.rs index 3e6b36f..d6aec33 100644 --- a/crates/litewire-tds/src/handler.rs +++ b/crates/litewire-tds/src/handler.rs @@ -10,7 +10,7 @@ use tracing::{debug, warn}; use litewire_backend::{SharedBackend, Value}; use litewire_translate::{self, Dialect, StatementKind, TranslateResult, classify}; -use crate::packet::{self, PacketType, DEFAULT_PACKET_SIZE}; +use crate::packet::{self, DEFAULT_PACKET_SIZE, PacketType}; use crate::token; /// Per-connection transaction state for TDS. @@ -167,7 +167,11 @@ async fn handle_login7( token::write_loginack(&mut resp, "litewire"); token::write_envchange_database(&mut resp, &db_name); token::write_envchange_packet_size(&mut resp, DEFAULT_PACKET_SIZE as u32); - token::write_info(&mut resp, 5701, &format!("Changed database context to '{db_name}'.")); + token::write_info( + &mut resp, + 5701, + &format!("Changed database context to '{db_name}'."), + ); token::write_done(&mut resp, token::DONE_FINAL, 0); packet::write_message(stream, PacketType::Response, &resp, DEFAULT_PACKET_SIZE).await?; @@ -245,8 +249,7 @@ fn skip_all_headers(payload: &[u8]) -> usize { if payload.len() < 4 { return 0; } - let total_len = - u32::from_le_bytes([payload[0], payload[1], payload[2], payload[3]]) as usize; + let total_len = u32::from_le_bytes([payload[0], payload[1], payload[2], payload[3]]) as usize; if total_len >= 4 && total_len <= payload.len() { total_len } else { @@ -408,10 +411,8 @@ async fn execute_sql( write_query_result(&mut resp, backend, &sqlite_sql, params).await; } StatementKind::Transaction => { - write_transaction_result( - &mut resp, backend, &sqlite_sql, session, - ) - .await; + write_transaction_result(&mut resp, backend, &sqlite_sql, session) + .await; } _ => { write_exec_result(&mut resp, backend, &sqlite_sql, params).await; diff --git a/crates/litewire-tds/src/token.rs b/crates/litewire-tds/src/token.rs index f3b5b73..9498dba 100644 --- a/crates/litewire-tds/src/token.rs +++ b/crates/litewire-tds/src/token.rs @@ -20,7 +20,6 @@ pub const TOKEN_ERROR: u8 = 0xAA; // ── DONE status flags ────────────────────────────────────────────────────── pub const DONE_FINAL: u16 = 0x0000; -pub const DONE_MORE: u16 = 0x0001; pub const DONE_COUNT: u16 = 0x0010; // ── TDS type IDs for COLMETADATA ─────────────────────────────────────────── @@ -185,7 +184,8 @@ pub fn write_error(buf: &mut BytesMut, number: u32, message: &str) { write_info_or_error(buf, TOKEN_ERROR, number, 14, message, "", "", 0); } -/// Shared writer for INFO (0xAB) and ERROR (0xAA) tokens — same format. +/// Shared writer for INFO (0xAB) and ERROR (0xAA) tokens -- same format. +#[allow(clippy::too_many_arguments)] fn write_info_or_error( buf: &mut BytesMut, token: u8, @@ -354,10 +354,7 @@ fn write_value(buf: &mut BytesMut, val: &Value, tds_type: TdsType) { /// Write a UTF-16LE NVARCHAR value with u16 byte-length prefix. fn write_nvarchar(buf: &mut BytesMut, s: &str) { - let utf16: Vec = s - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let utf16: Vec = s.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); buf.put_u16_le(utf16.len() as u16); buf.put_slice(&utf16); } diff --git a/crates/litewire-translate/src/common.rs b/crates/litewire-translate/src/common.rs index 6353711..88d2c77 100644 --- a/crates/litewire-translate/src/common.rs +++ b/crates/litewire-translate/src/common.rs @@ -4,8 +4,8 @@ //! type casts, and parameter placeholders. use sqlparser::ast::{ - Expr, Function, FunctionArg, FunctionArgExpr, FunctionArgumentList, - FunctionArguments, Ident, ObjectName, Statement, Value, ValueWithSpan, + Expr, Function, FunctionArg, FunctionArgExpr, FunctionArgumentList, FunctionArguments, Ident, + ObjectName, Statement, Value, ValueWithSpan, }; use crate::TranslateError; @@ -230,6 +230,46 @@ fn rewrite_function(func: &mut Function) { "ISNULL" => { func.name = func_name("IFNULL"); } + // MySQL LAST_INSERT_ID() -> SQLite last_insert_rowid(). + // Note: the MySQL 2-arg form LAST_INSERT_ID(expr) sets the session's + // last-insert-id; SQLite has no equivalent. We rewrite the name and let + // SQLite's function-arity check reject the 2-arg form loudly rather + // than silently changing semantics. + "LAST_INSERT_ID" => { + func.name = func_name("last_insert_rowid"); + } + // MySQL ROW_COUNT() -> SQLite changes(). + "ROW_COUNT" => { + func.name = func_name("changes"); + } + // Session-identity built-ins. SQLite has no analogue, so we return + // constant placeholders that mirror the values used by the metadata + // fast path in `metadata::system_variable_value`. + // DATABASE() / SCHEMA() -> 'main' + // VERSION() -> '8.0.0-litewire' + // USER() / CURRENT_USER() / + // SESSION_USER() / SYSTEM_USER() -> 'root@localhost' + // CONNECTION_ID() -> 0 (SQLite has no per-connection ID) + "DATABASE" | "SCHEMA" => { + func.name = func_name("coalesce"); + func.args = func_args(vec![value_expr(Value::SingleQuotedString("main".into()))]); + } + "VERSION" => { + func.name = func_name("coalesce"); + func.args = func_args(vec![value_expr(Value::SingleQuotedString( + "8.0.0-litewire".into(), + ))]); + } + "USER" | "CURRENT_USER" | "SESSION_USER" | "SYSTEM_USER" => { + func.name = func_name("coalesce"); + func.args = func_args(vec![value_expr(Value::SingleQuotedString( + "root@localhost".into(), + ))]); + } + "CONNECTION_ID" => { + func.name = func_name("coalesce"); + func.args = func_args(vec![value_expr(Value::Number("0".into(), false))]); + } "NEWID" => { // NEWID() -> lower(hex(randomblob(16))) func.name = func_name("lower"); @@ -277,7 +317,7 @@ fn rewrite_value(val: &mut Value) { #[cfg(test)] mod tests { use super::*; - use crate::{translate, Dialect, TranslateResult}; + use crate::{Dialect, TranslateResult, translate}; #[test] fn boolean_rewrite() { @@ -341,30 +381,22 @@ mod tests { #[test] fn boolean_in_where_clause() { - let results = - translate("SELECT * FROM t WHERE active = TRUE", Dialect::MySQL).unwrap(); + let results = translate("SELECT * FROM t WHERE active = TRUE", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains('1'), "got: {sql}"); } #[test] fn function_in_insert_values() { - let results = translate( - "INSERT INTO t (created) VALUES (NOW())", - Dialect::MySQL, - ) - .unwrap(); + let results = translate("INSERT INTO t (created) VALUES (NOW())", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); assert!(sql.to_ascii_lowercase().contains("datetime"), "got: {sql}"); } #[test] fn function_in_update_set() { - let results = translate( - "UPDATE t SET updated = NOW() WHERE id = 1", - Dialect::MySQL, - ) - .unwrap(); + let results = + translate("UPDATE t SET updated = NOW() WHERE id = 1", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); assert!(sql.to_ascii_lowercase().contains("datetime"), "got: {sql}"); } @@ -380,7 +412,67 @@ mod tests { assert!(sql.contains('1'), "got: {sql}"); } - // ── NEWID rewrite ──────────────────────────────────────────────────── + // -- LAST_INSERT_ID / ROW_COUNT / DATABASE / VERSION / USER / CONNECTION_ID -- + + #[test] + fn last_insert_id_rewrite() { + let results = translate("SELECT LAST_INSERT_ID()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + assert!( + sql.to_ascii_lowercase().contains("last_insert_rowid"), + "got: {sql}" + ); + assert!( + !sql.to_ascii_uppercase().contains("LAST_INSERT_ID("), + "got: {sql}" + ); + } + + #[test] + fn row_count_rewrite() { + let results = translate("SELECT ROW_COUNT()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + assert!(sql.to_ascii_lowercase().contains("changes"), "got: {sql}"); + assert!( + !sql.to_ascii_uppercase().contains("ROW_COUNT("), + "got: {sql}" + ); + } + + #[test] + fn database_rewrite() { + let results = translate("SELECT DATABASE()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + assert!(sql.contains("'main'"), "got: {sql}"); + } + + #[test] + fn version_rewrite() { + let results = translate("SELECT VERSION()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + assert!(sql.contains("8.0.0-litewire"), "got: {sql}"); + } + + #[test] + fn current_user_rewrite() { + let results = translate("SELECT CURRENT_USER()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + assert!(sql.contains("'root@localhost'"), "got: {sql}"); + } + + #[test] + fn connection_id_rewrite() { + let results = translate("SELECT CONNECTION_ID()", Dialect::MySQL).unwrap(); + let sql = extract_sql(&results[0]); + // Must produce a numeric literal 0 in the emitted SQL. + assert!(sql.contains('0'), "got: {sql}"); + assert!( + !sql.to_ascii_uppercase().contains("CONNECTION_ID("), + "got: {sql}" + ); + } + + // -- NEWID rewrite -------------------------------------------------------- #[test] fn newid_rewrite() { diff --git a/crates/litewire-translate/src/lib.rs b/crates/litewire-translate/src/lib.rs index 0939ea2..4e971bd 100644 --- a/crates/litewire-translate/src/lib.rs +++ b/crates/litewire-translate/src/lib.rs @@ -11,7 +11,7 @@ pub mod postgres; pub mod tds; use sqlparser::ast::Statement; -use sqlparser::dialect::{MySqlDialect, PostgreSqlDialect, MsSqlDialect}; +use sqlparser::dialect::{MsSqlDialect, MySqlDialect, PostgreSqlDialect}; use sqlparser::parser::Parser; /// Source SQL dialect for translation. @@ -57,6 +57,14 @@ pub fn translate(sql: &str, dialect: Dialect) -> Result, Tr return Ok(vec![TranslateResult::Metadata(meta)]); } + // Session/transaction statements that we handle textually before hitting + // sqlparser -- these either don't parse cleanly across dialects + // (`START TRANSACTION WITH CONSISTENT SNAPSHOT`, `BEGIN TRANSACTION named`) + // or are simple keyword substitutions (`LAST_INSERT_ID()` -> `last_insert_rowid()`). + if let Some(rewritten) = rewrite_transaction_statement(sql) { + return Ok(vec![TranslateResult::Sql(rewritten)]); + } + // Check for no-op statements. if is_noop(sql, dialect) { return Ok(vec![TranslateResult::Noop]); @@ -68,8 +76,8 @@ pub fn translate(sql: &str, dialect: Dialect) -> Result, Tr Dialect::TDS => Box::new(MsSqlDialect {}), }; - let statements = - Parser::parse_sql(parser_dialect.as_ref(), sql).map_err(|e| TranslateError::Parse(e.to_string()))?; + let statements = Parser::parse_sql(parser_dialect.as_ref(), sql) + .map_err(|e| TranslateError::Parse(e.to_string()))?; let mut results = Vec::with_capacity(statements.len()); for stmt in statements { @@ -81,18 +89,96 @@ pub fn translate(sql: &str, dialect: Dialect) -> Result, Tr Ok(results) } +/// Rewrite MySQL/T-SQL-flavored transaction control statements to a form +/// SQLite accepts. Returns the rewritten SQL if the input matched a known +/// transaction shape, or `None` to fall through to the parser / no-op check. +/// +/// Handled: +/// * `START TRANSACTION`, `START TRANSACTION READ ONLY`, +/// `START TRANSACTION READ WRITE`, `START TRANSACTION WITH CONSISTENT SNAPSHOT` +/// -> `BEGIN` +/// * `BEGIN`, `BEGIN WORK`, `BEGIN TRANSACTION`, `BEGIN TRANSACTION ` +/// (T-SQL named transactions) -> `BEGIN` (name is stripped; SQLite does not +/// support named transactions -- callers should use SAVEPOINT if they want +/// a nameable rollback point) +/// * `COMMIT`, `COMMIT WORK`, `COMMIT TRANSACTION [name]` -> `COMMIT` +/// * `ROLLBACK`, `ROLLBACK WORK`, `ROLLBACK TRANSACTION [name]` -> `ROLLBACK` +/// (but *not* `ROLLBACK TO [SAVEPOINT] name`, which is passed through so +/// SQLite's savepoint machinery handles it) +fn rewrite_transaction_statement(sql: &str) -> Option { + let trimmed = sql.trim().trim_end_matches(';').trim(); + let upper = trimmed.to_ascii_uppercase(); + + // START TRANSACTION ... (MySQL / PG) + if upper == "START TRANSACTION" || upper.starts_with("START TRANSACTION ") { + return Some("BEGIN".to_string()); + } + + // BEGIN WORK / BEGIN TRANSACTION [name] -- T-SQL / SQL standard + // Do NOT eat plain "BEGIN" or "BEGIN;" (SQLite accepts it as-is; passing it + // through keeps `classify()` reporting Transaction and preserves any future + // dialect-specific handling in the caller). + if upper == "BEGIN WORK" { + return Some("BEGIN".to_string()); + } + if upper == "BEGIN TRANSACTION" || upper.starts_with("BEGIN TRANSACTION ") { + return Some("BEGIN".to_string()); + } + + // COMMIT WORK / COMMIT TRANSACTION [name] + if upper == "COMMIT WORK" { + return Some("COMMIT".to_string()); + } + if upper == "COMMIT TRANSACTION" || upper.starts_with("COMMIT TRANSACTION ") { + return Some("COMMIT".to_string()); + } + + // ROLLBACK WORK / ROLLBACK TRANSACTION [name] + // Careful: ROLLBACK TO SAVEPOINT name must pass through unchanged. + if upper == "ROLLBACK WORK" { + return Some("ROLLBACK".to_string()); + } + if upper == "ROLLBACK TRANSACTION" { + return Some("ROLLBACK".to_string()); + } + if let Some(rest) = upper.strip_prefix("ROLLBACK TRANSACTION ") { + // `ROLLBACK TRANSACTION TO SAVEPOINT foo` and `ROLLBACK TRANSACTION TO foo` + // are T-SQL savepoint rollback forms -- pass those through to SQLite as + // `ROLLBACK TO SAVEPOINT foo` / `ROLLBACK TO foo`. + let rest = rest.trim_start(); + if let Some(after_to) = rest.strip_prefix("TO ") { + let after_to = after_to.trim_start(); + let name_upper = if let Some(after_sp) = after_to.strip_prefix("SAVEPOINT ") { + after_sp.trim() + } else { + after_to.trim() + }; + // Extract the original-case name by taking the same trailing chars from `trimmed`. + let orig_name = &trimmed[trimmed.len() - name_upper.len()..]; + return Some(format!("ROLLBACK TO SAVEPOINT {orig_name}")); + } + // Plain `ROLLBACK TRANSACTION some_name` -- strip the name. + return Some("ROLLBACK".to_string()); + } + + None +} + /// Check if a SQL statement is a no-op for SQLite. fn is_noop(sql: &str, _dialect: Dialect) -> bool { let upper = sql.trim().to_ascii_uppercase(); // SET statements that have no SQLite equivalent. - if upper.starts_with("SET ") { - let rest = upper["SET ".len()..].trim_start(); - // SET NAMES, SET CHARACTER SET, SET SESSION, SET GLOBAL, SET time_zone, SET sql_mode + if let Some(after_set) = upper.strip_prefix("SET ") { + let rest = after_set.trim_start(); + // Normalize `SET @@[session|global.]name` -> plain name, so autocommit-style + // rules downstream match regardless of the `@@`/`SESSION.`/`GLOBAL.` prefix. + let rest_norm = normalize_session_prefix(rest); + let rest = rest_norm.as_deref().unwrap_or(rest); + + // Explicitly ignored no-ops (session/global toggles + T-SQL SET switches). if rest.starts_with("NAMES") || rest.starts_with("CHARACTER SET") - || rest.starts_with("SESSION") - || rest.starts_with("GLOBAL") || rest.starts_with("TIME_ZONE") || rest.starts_with("SQL_MODE") || rest.starts_with("NOCOUNT") @@ -100,18 +186,97 @@ fn is_noop(sql: &str, _dialect: Dialect) -> bool { || rest.starts_with("QUOTED_IDENTIFIER") || rest.starts_with("XACT_ABORT") { + tracing::debug!(sql = %sql.trim(), "SET: treating as noop"); + return true; + } + + // SET autocommit = 1|0|ON|OFF|TRUE|FALSE + if let Some(val) = rest.strip_prefix("AUTOCOMMIT") { + let val = val.trim_start_matches(|c: char| c == '=' || c.is_whitespace()); + if val.starts_with('0') || val.starts_with("OFF") || val.starts_with("FALSE") { + tracing::warn!( + "SET autocommit=0 requested but litewire does not emulate MySQL implicit \ + transactions -- statements will still auto-commit unless wrapped in BEGIN/COMMIT" + ); + } else { + tracing::debug!("SET autocommit: noop (SQLite default matches)"); + } return true; } + + // SET [SESSION|GLOBAL] TRANSACTION ISOLATION LEVEL ... and bare + // SET TRANSACTION ... (post-normalization the SESSION/GLOBAL prefix is + // already stripped, so we only need to look for the TRANSACTION keyword here). + if rest.starts_with("TRANSACTION") { + tracing::debug!( + "SET TRANSACTION: noop (SQLite runs at serializable isolation by default)" + ); + return true; + } + + // Any remaining `SET SESSION ...` / `SET GLOBAL ...` that we didn't + // rewrite above (e.g. `SET SESSION wait_timeout=28800`) -- also noop. + if rest.starts_with("SESSION") || rest.starts_with("GLOBAL") { + tracing::debug!("SET SESSION/GLOBAL: noop"); + return true; + } + } + + // LOCK TABLES / UNLOCK TABLES -- SQLite has no equivalent; warn once so the + // user knows the requested locking semantics are not being enforced. + if upper.starts_with("LOCK TABLES") || upper.starts_with("LOCK TABLE ") { + tracing::warn!("LOCK TABLES: noop (SQLite provides file-level locking only)"); + return true; + } + if upper == "UNLOCK TABLES" || upper.starts_with("UNLOCK TABLES ") { + tracing::warn!("UNLOCK TABLES: noop"); + return true; } false } +/// Strip a leading `@@` / `@@SESSION.` / `@@GLOBAL.` / `SESSION ` / `GLOBAL ` +/// prefix from an uppercased SET body, returning the trimmed remainder. +/// Returns `None` if no prefix applied. +fn normalize_session_prefix(rest: &str) -> Option { + if let Some(after) = rest.strip_prefix("@@SESSION.") { + return Some(after.to_string()); + } + if let Some(after) = rest.strip_prefix("@@GLOBAL.") { + return Some(after.to_string()); + } + if let Some(after) = rest.strip_prefix("@@") { + return Some(after.to_string()); + } + if let Some(after) = rest.strip_prefix("SESSION.") { + return Some(after.to_string()); + } + if let Some(after) = rest.strip_prefix("GLOBAL.") { + return Some(after.to_string()); + } + // `SET SESSION TRANSACTION ...` / `SET GLOBAL TRANSACTION ...` + if let Some(after) = rest.strip_prefix("SESSION ") { + // Preserve the leading keyword so downstream `.starts_with("TRANSACTION")` + // and `.starts_with("SESSION")` checks both remain accurate. + // For `SET SESSION TRANSACTION ISOLATION LEVEL ...` we want to fall + // through and match TRANSACTION; for `SET SESSION wait_timeout=...` we + // want to fall through and match SESSION. Return the tail with SESSION + // stripped only if TRANSACTION follows. + if after.trim_start().starts_with("TRANSACTION") { + return Some(after.trim_start().to_string()); + } + } + if let Some(after) = rest.strip_prefix("GLOBAL ") { + if after.trim_start().starts_with("TRANSACTION") { + return Some(after.trim_start().to_string()); + } + } + None +} + /// Rewrite a parsed statement from the source dialect to SQLite-compatible form. -fn rewrite_statement( - mut stmt: Statement, - dialect: Dialect, -) -> Result { +fn rewrite_statement(mut stmt: Statement, dialect: Dialect) -> Result { // Apply common rewrites (expressions, types). common::rewrite_statement(&mut stmt)?; @@ -217,10 +382,7 @@ mod tests { #[test] fn classify_create() { - assert_eq!( - classify("CREATE TABLE users (id INT)"), - StatementKind::Ddl - ); + assert_eq!(classify("CREATE TABLE users (id INT)"), StatementKind::Ddl); } #[test] @@ -276,11 +438,10 @@ mod tests { #[test] fn translate_empty_returns_empty() { // An empty input should still parse (to zero statements). - let results = translate("", Dialect::MySQL); - // sqlparser may return an error or empty vec — either is fine. - match results { - Ok(v) => assert!(v.is_empty()), - Err(_) => {} // parse error on empty string is acceptable + // sqlparser may return an error or empty vec -- either is fine; a parse + // error on empty input is acceptable. + if let Ok(v) = translate("", Dialect::MySQL) { + assert!(v.is_empty()); } } @@ -289,4 +450,178 @@ mod tests { let result = translate("NOT VALID SQL !!! @@@ {{{}}", Dialect::MySQL); assert!(result.is_err()); } + + // -- Transaction rewrites -------------------------------------------------- + + fn expect_sql(sql: &str, dialect: Dialect) -> String { + let results = translate(sql, dialect).unwrap_or_else(|e| panic!("translate({sql:?}): {e}")); + match &results[0] { + TranslateResult::Sql(s) => s.clone(), + other => panic!("expected Sql, got: {other:?}"), + } + } + + fn expect_noop(sql: &str, dialect: Dialect) { + let results = translate(sql, dialect).unwrap_or_else(|e| panic!("translate({sql:?}): {e}")); + assert!( + matches!(results[0], TranslateResult::Noop), + "expected Noop for {sql:?}, got: {:?}", + results[0] + ); + } + + #[test] + fn start_transaction_becomes_begin() { + assert_eq!(expect_sql("START TRANSACTION", Dialect::MySQL), "BEGIN"); + assert_eq!(expect_sql("start transaction", Dialect::MySQL), "BEGIN"); + assert_eq!(expect_sql("START TRANSACTION;", Dialect::MySQL), "BEGIN"); + } + + #[test] + fn start_transaction_read_only_becomes_begin() { + assert_eq!( + expect_sql("START TRANSACTION READ ONLY", Dialect::MySQL), + "BEGIN" + ); + } + + #[test] + fn start_transaction_with_consistent_snapshot_becomes_begin() { + assert_eq!( + expect_sql("START TRANSACTION WITH CONSISTENT SNAPSHOT", Dialect::MySQL), + "BEGIN" + ); + } + + #[test] + fn tsql_begin_transaction_becomes_begin() { + assert_eq!(expect_sql("BEGIN TRANSACTION", Dialect::TDS), "BEGIN"); + } + + #[test] + fn tsql_named_begin_transaction_strips_name() { + assert_eq!( + expect_sql("BEGIN TRANSACTION my_txn", Dialect::TDS), + "BEGIN" + ); + } + + #[test] + fn begin_work_becomes_begin() { + assert_eq!(expect_sql("BEGIN WORK", Dialect::MySQL), "BEGIN"); + } + + #[test] + fn commit_work_and_named_become_commit() { + assert_eq!(expect_sql("COMMIT WORK", Dialect::MySQL), "COMMIT"); + assert_eq!( + expect_sql("COMMIT TRANSACTION my_txn", Dialect::TDS), + "COMMIT" + ); + } + + #[test] + fn rollback_work_and_named_become_rollback() { + assert_eq!(expect_sql("ROLLBACK WORK", Dialect::MySQL), "ROLLBACK"); + assert_eq!( + expect_sql("ROLLBACK TRANSACTION my_txn", Dialect::TDS), + "ROLLBACK" + ); + } + + #[test] + fn rollback_to_savepoint_passthrough() { + // Plain `ROLLBACK TO SAVEPOINT foo` is not a transaction control statement -- + // it's a savepoint rollback and must reach the parser so SQLite gets it. + let results = translate("ROLLBACK TO SAVEPOINT foo", Dialect::MySQL).unwrap(); + // Whatever emit produces, it must still be a Sql result containing SAVEPOINT and foo. + match &results[0] { + TranslateResult::Sql(s) => { + assert!(s.to_ascii_uppercase().contains("SAVEPOINT"), "got: {s}"); + assert!(s.contains("foo"), "got: {s}"); + } + other => panic!("expected Sql, got: {other:?}"), + } + } + + #[test] + fn tsql_rollback_transaction_to_savepoint_rewritten() { + // T-SQL: ROLLBACK TRANSACTION TO SAVEPOINT foo -> ROLLBACK TO SAVEPOINT foo + let sql = expect_sql("ROLLBACK TRANSACTION TO SAVEPOINT foo", Dialect::TDS); + let up = sql.to_ascii_uppercase(); + assert!(up.contains("ROLLBACK"), "got: {sql}"); + assert!(up.contains("SAVEPOINT"), "got: {sql}"); + assert!(sql.contains("foo"), "got: {sql}"); + } + + #[test] + fn savepoint_passthrough() { + // SAVEPOINT foo should reach the parser and come out as valid SAVEPOINT SQL. + let results = translate("SAVEPOINT foo", Dialect::MySQL).unwrap(); + match &results[0] { + TranslateResult::Sql(s) => { + assert!(s.to_ascii_uppercase().contains("SAVEPOINT"), "got: {s}"); + assert!(s.contains("foo"), "got: {s}"); + } + other => panic!("expected Sql for SAVEPOINT, got: {other:?}"), + } + } + + #[test] + fn release_savepoint_passthrough() { + let results = translate("RELEASE SAVEPOINT foo", Dialect::MySQL).unwrap(); + match &results[0] { + TranslateResult::Sql(s) => { + assert!(s.to_ascii_uppercase().contains("RELEASE"), "got: {s}"); + assert!(s.contains("foo"), "got: {s}"); + } + other => panic!("expected Sql for RELEASE SAVEPOINT, got: {other:?}"), + } + } + + // -- SET session/global/autocommit noops ----------------------------------- + + #[test] + fn set_autocommit_1_is_noop() { + expect_noop("SET autocommit = 1", Dialect::MySQL); + expect_noop("SET AUTOCOMMIT=ON", Dialect::MySQL); + expect_noop("SET autocommit = true", Dialect::MySQL); + } + + #[test] + fn set_autocommit_0_is_noop_with_warning() { + // We don't emulate implicit-transaction mode, but we do return Noop + // rather than an error -- log will be a WARN at runtime. + expect_noop("SET autocommit = 0", Dialect::MySQL); + expect_noop("SET autocommit = OFF", Dialect::MySQL); + } + + #[test] + fn set_transaction_isolation_level_is_noop() { + expect_noop( + "SET TRANSACTION ISOLATION LEVEL SERIALIZABLE", + Dialect::MySQL, + ); + expect_noop( + "SET SESSION TRANSACTION ISOLATION LEVEL READ COMMITTED", + Dialect::MySQL, + ); + expect_noop( + "SET GLOBAL TRANSACTION ISOLATION LEVEL REPEATABLE READ", + Dialect::MySQL, + ); + } + + #[test] + fn set_at_at_session_variable_is_noop() { + expect_noop("SET @@session.autocommit = 1", Dialect::MySQL); + expect_noop("SET @@global.autocommit = 1", Dialect::MySQL); + expect_noop("SET @@autocommit = 1", Dialect::MySQL); + } + + #[test] + fn lock_tables_is_noop() { + expect_noop("LOCK TABLES users WRITE", Dialect::MySQL); + expect_noop("UNLOCK TABLES", Dialect::MySQL); + } } diff --git a/crates/litewire-translate/src/metadata.rs b/crates/litewire-translate/src/metadata.rs index 4c8361f..1a6ba22 100644 --- a/crates/litewire-translate/src/metadata.rs +++ b/crates/litewire-translate/src/metadata.rs @@ -21,13 +21,9 @@ pub enum MetadataQuery { /// `SELECT @@variable` queries — MySQL system variables. SystemVariables { variables: Vec }, /// `SELECT ... FROM information_schema.tables` — table listing. - InformationSchemaTables { - schema_filter: Option, - }, + InformationSchemaTables { schema_filter: Option }, /// `SELECT ... FROM information_schema.columns` — column listing. - InformationSchemaColumns { - table_filter: Option, - }, + InformationSchemaColumns { table_filter: Option }, /// `SELECT ... FROM information_schema.schemata` — schema listing. InformationSchemata, /// `SELECT ... FROM pg_catalog.pg_tables` or similar. @@ -138,8 +134,8 @@ impl MetadataQuery { /// Return a synthetic value for a MySQL system variable. fn system_variable_value(name: &str) -> &'static str { match name.to_ascii_lowercase().as_str() { - "max_allowed_packet" => "67108864", // 64 MiB - "wait_timeout" => "28800", // 8 hours + "max_allowed_packet" => "67108864", // 64 MiB + "wait_timeout" => "28800", // 8 hours "interactive_timeout" => "28800", "net_write_timeout" => "60", "net_read_timeout" => "30", @@ -314,6 +310,9 @@ pub fn detect_metadata_query(sql: &str, _dialect: Dialect) -> Option Option { let pattern = format!("{column} = "); if let Some(pos) = upper_sql.find(&pattern) { @@ -338,7 +337,7 @@ fn extract_where_value_original(original_sql: &str, column: &str) -> Option assert_eq!(schema, "mydb"), + }) => { + assert_eq!(schema, "mydb") + } other => panic!("expected InformationSchemaTables with filter, got: {other:?}"), } } @@ -695,31 +692,25 @@ mod tests { match q { Some(MetadataQuery::InformationSchemaColumns { table_filter: Some(table), - }) => assert_eq!(table, "users"), + }) => { + assert_eq!(table, "users") + } other => panic!("expected InformationSchemaColumns with filter, got: {other:?}"), } } #[test] fn detect_information_schema_columns_no_filter() { - let q = detect_metadata_query( - "SELECT * FROM INFORMATION_SCHEMA.COLUMNS", - Dialect::MySQL, - ); + let q = detect_metadata_query("SELECT * FROM INFORMATION_SCHEMA.COLUMNS", Dialect::MySQL); assert!(matches!( q, - Some(MetadataQuery::InformationSchemaColumns { - table_filter: None - }) + Some(MetadataQuery::InformationSchemaColumns { table_filter: None }) )); } #[test] fn detect_information_schema_schemata() { - let q = detect_metadata_query( - "SELECT * FROM INFORMATION_SCHEMA.SCHEMATA", - Dialect::MySQL, - ); + let q = detect_metadata_query("SELECT * FROM INFORMATION_SCHEMA.SCHEMATA", Dialect::MySQL); assert!(matches!(q, Some(MetadataQuery::InformationSchemata))); } @@ -770,10 +761,7 @@ mod tests { #[test] fn information_schema_columns_no_table_fallback() { - let sql = MetadataQuery::InformationSchemaColumns { - table_filter: None, - } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaColumns { table_filter: None }.to_sqlite_sql(); // Falls back to listing tables. assert!(sql.contains("sqlite_master"), "got: {sql}"); } @@ -827,11 +815,11 @@ mod tests { #[test] fn detect_pg_catalog_tables() { - let q = detect_metadata_query( - "SELECT * FROM pg_catalog.pg_tables", - Dialect::PostgreSQL, + let q = detect_metadata_query("SELECT * FROM pg_catalog.pg_tables", Dialect::PostgreSQL); + assert!( + matches!(q, Some(MetadataQuery::PgCatalogTables)), + "got: {q:?}" ); - assert!(matches!(q, Some(MetadataQuery::PgCatalogTables)), "got: {q:?}"); } #[test] @@ -840,7 +828,10 @@ mod tests { "SELECT * FROM pg_catalog.pg_class WHERE relkind = 'r'", Dialect::PostgreSQL, ); - assert!(matches!(q, Some(MetadataQuery::PgCatalogTables)), "got: {q:?}"); + assert!( + matches!(q, Some(MetadataQuery::PgCatalogTables)), + "got: {q:?}" + ); } #[test] diff --git a/crates/litewire-translate/src/mysql.rs b/crates/litewire-translate/src/mysql.rs index 6ff41db..bd9272b 100644 --- a/crates/litewire-translate/src/mysql.rs +++ b/crates/litewire-translate/src/mysql.rs @@ -5,8 +5,8 @@ //! MySQL-specific expressions. use sqlparser::ast::{ - DataType, DoUpdate, Expr, LimitClause, Offset, OffsetRows, OnConflict, OnConflictAction, - OnInsert, Statement, + DataType, DoUpdate, LimitClause, Offset, OffsetRows, OnConflict, OnConflictAction, OnInsert, + Statement, }; use crate::TranslateError; @@ -72,7 +72,7 @@ fn rewrite_create_table(create: &mut sqlparser::ast::CreateTable) { !matches!( &opt.option, sqlparser::ast::ColumnOption::DialectSpecific(tokens) - if tokens.iter().any(|t| t.to_string().to_ascii_uppercase() == "AUTO_INCREMENT") + if tokens.iter().any(|t| t.to_string().eq_ignore_ascii_case("AUTO_INCREMENT")) ) }); } @@ -137,7 +137,7 @@ fn rewrite_data_type(dt: &DataType) -> DataType { #[cfg(test)] mod tests { - use crate::{translate, Dialect, TranslateResult}; + use crate::{Dialect, TranslateResult, translate}; fn extract_sql(result: &TranslateResult) -> &str { match result { @@ -168,8 +168,7 @@ mod tests { #[test] fn set_sql_mode_is_noop() { - let results = - translate("SET sql_mode = 'STRICT_TRANS_TABLES'", Dialect::MySQL).unwrap(); + let results = translate("SET sql_mode = 'STRICT_TRANS_TABLES'", Dialect::MySQL).unwrap(); assert!(matches!(results[0], TranslateResult::Noop)); } @@ -239,22 +238,17 @@ mod tests { #[test] fn limit_offset_comma_rewritten() { - let results = - translate("SELECT * FROM t LIMIT 5, 10", Dialect::MySQL).unwrap(); + let results = translate("SELECT * FROM t LIMIT 5, 10", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("LIMIT 10"), "expected LIMIT 10, got: {sql}"); - assert!( - upper.contains("OFFSET 5"), - "expected OFFSET 5, got: {sql}" - ); + assert!(upper.contains("OFFSET 5"), "expected OFFSET 5, got: {sql}"); } #[test] fn standard_limit_unchanged() { // Standard LIMIT without offset should not add an OFFSET clause. - let results = - translate("SELECT * FROM t LIMIT 10", Dialect::MySQL).unwrap(); + let results = translate("SELECT * FROM t LIMIT 10", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); // sqlparser may or may not preserve LIMIT in the emitted SQL for MySQL dialect. @@ -274,10 +268,7 @@ mod tests { let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("TINYINT"), "TINYINT not rewritten: {sql}"); - assert!( - !upper.contains("SMALLINT"), - "SMALLINT not rewritten: {sql}" - ); + assert!(!upper.contains("SMALLINT"), "SMALLINT not rewritten: {sql}"); assert!( !upper.contains("MEDIUMINT"), "MEDIUMINT not rewritten: {sql}" @@ -287,8 +278,11 @@ mod tests { #[test] fn varchar_to_text() { - let results = - translate("CREATE TABLE t (name VARCHAR(255), bio TEXT)", Dialect::MySQL).unwrap(); + let results = translate( + "CREATE TABLE t (name VARCHAR(255), bio TEXT)", + Dialect::MySQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("VARCHAR"), "VARCHAR not rewritten: {sql}"); @@ -323,10 +317,7 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!( - !upper.contains("DATETIME"), - "DATETIME not rewritten: {sql}" - ); + assert!(!upper.contains("DATETIME"), "DATETIME not rewritten: {sql}"); assert!( !upper.contains("TIMESTAMP"), "TIMESTAMP not rewritten: {sql}" diff --git a/crates/litewire-translate/src/postgres.rs b/crates/litewire-translate/src/postgres.rs index 4f0514e..a25db21 100644 --- a/crates/litewire-translate/src/postgres.rs +++ b/crates/litewire-translate/src/postgres.rs @@ -35,10 +35,9 @@ fn rewrite_data_type(dt: &DataType) -> DataType { DataType::Boolean => DataType::Integer(None), - DataType::Date - | DataType::Timestamp(_, _) - | DataType::Time(_, _) - | DataType::Interval => DataType::Text, + DataType::Date | DataType::Timestamp(_, _) | DataType::Time(_, _) | DataType::Interval => { + DataType::Text + } DataType::JSON | DataType::JSONB => DataType::Text, @@ -56,7 +55,7 @@ fn rewrite_data_type(dt: &DataType) -> DataType { #[cfg(test)] mod tests { - use crate::{translate, Dialect, TranslateResult}; + use crate::{Dialect, TranslateResult, translate}; fn extract_sql(result: &TranslateResult) -> &str { match result { @@ -139,7 +138,10 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!(!upper.contains("TIMESTAMP"), "TIMESTAMP not rewritten: {sql}"); + assert!( + !upper.contains("TIMESTAMP"), + "TIMESTAMP not rewritten: {sql}" + ); } #[test] diff --git a/crates/litewire-translate/src/tds.rs b/crates/litewire-translate/src/tds.rs index f64cc46..da48c03 100644 --- a/crates/litewire-translate/src/tds.rs +++ b/crates/litewire-translate/src/tds.rs @@ -33,7 +33,7 @@ fn rewrite_top_to_limit(query: &mut sqlparser::ast::Query) { let limit_expr = match quantity { TopQuantity::Expr(e) => e, TopQuantity::Constant(n) => Expr::Value(sqlparser::ast::ValueWithSpan { - value: sqlparser::ast::Value::Number(n.to_string().into(), false), + value: sqlparser::ast::Value::Number(n.to_string(), false), span: sqlparser::tokenizer::Span::empty(), }), }; @@ -63,18 +63,18 @@ fn rewrite_data_type(dt: &DataType) -> DataType { | DataType::Decimal(..) | DataType::Numeric(..) => DataType::Real, - DataType::Varchar(_) - | DataType::Char(_) - | DataType::Nvarchar(_) - | DataType::Text => DataType::Text, + DataType::Varchar(_) | DataType::Char(_) | DataType::Nvarchar(_) | DataType::Text => { + DataType::Text + } DataType::Binary(_) | DataType::Varbinary(_) => DataType::Blob(None), DataType::Boolean | DataType::Bit(_) => DataType::Integer(None), - DataType::Date | DataType::Datetime(_) | DataType::Timestamp(_, _) | DataType::Time(_, _) => { - DataType::Text - } + DataType::Date + | DataType::Datetime(_) + | DataType::Timestamp(_, _) + | DataType::Time(_, _) => DataType::Text, DataType::Custom(name, _) => { let upper = name.to_string().to_ascii_uppercase(); @@ -92,7 +92,7 @@ fn rewrite_data_type(dt: &DataType) -> DataType { #[cfg(test)] mod tests { - use crate::{translate, Dialect, TranslateResult}; + use crate::{Dialect, TranslateResult, translate}; fn extract_sql(result: &TranslateResult) -> &str { match result { @@ -117,11 +117,7 @@ mod tests { #[test] fn nvarchar_to_text() { - let results = translate( - "CREATE TABLE t (name NVARCHAR(255))", - Dialect::TDS, - ) - .unwrap(); + let results = translate("CREATE TABLE t (name NVARCHAR(255))", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("NVARCHAR"), "NVARCHAR not rewritten: {sql}"); @@ -133,14 +129,16 @@ mod tests { let results = translate("CREATE TABLE t (data VARBINARY(MAX))", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!(!upper.contains("VARBINARY"), "VARBINARY not rewritten: {sql}"); + assert!( + !upper.contains("VARBINARY"), + "VARBINARY not rewritten: {sql}" + ); assert!(upper.contains("BLOB"), "no BLOB found: {sql}"); } #[test] fn datetime_to_text() { - let results = - translate("CREATE TABLE t (created DATETIME)", Dialect::TDS).unwrap(); + let results = translate("CREATE TABLE t (created DATETIME)", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("DATETIME"), "DATETIME not rewritten: {sql}"); @@ -149,11 +147,7 @@ mod tests { #[test] fn float_types_to_real() { - let results = translate( - "CREATE TABLE t (a FLOAT, b DECIMAL(10,2))", - Dialect::TDS, - ) - .unwrap(); + let results = translate("CREATE TABLE t (a FLOAT, b DECIMAL(10,2))", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("REAL"), "no REAL found: {sql}"); @@ -210,11 +204,7 @@ mod tests { #[test] fn uniqueidentifier_to_text() { - let results = translate( - "CREATE TABLE t (id UNIQUEIDENTIFIER)", - Dialect::TDS, - ) - .unwrap(); + let results = translate("CREATE TABLE t (id UNIQUEIDENTIFIER)", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!( diff --git a/crates/litewire/src/lib.rs b/crates/litewire/src/lib.rs index c6e5223..ac928e4 100644 --- a/crates/litewire/src/lib.rs +++ b/crates/litewire/src/lib.rs @@ -137,8 +137,7 @@ impl LiteWire { #[cfg(feature = "tds")] if let Some(addr) = self.tds_listen { let config = litewire_tds::TdsFrontendConfig { listen: addr }; - let frontend = - litewire_tds::TdsFrontend::new(config, Arc::clone(&self.backend)); + let frontend = litewire_tds::TdsFrontend::new(config, Arc::clone(&self.backend)); handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); diff --git a/crates/litewire/src/main.rs b/crates/litewire/src/main.rs index be288c8..1357205 100644 --- a/crates/litewire/src/main.rs +++ b/crates/litewire/src/main.rs @@ -4,11 +4,7 @@ use clap::Parser; use tracing::info; #[derive(Parser)] -#[command( - name = "litewire", - version, - about = "SQL protocol proxy for SQLite" -)] +#[command(name = "litewire", version, about = "SQL protocol proxy for SQLite")] struct Cli { /// SQLite database file path. #[arg(long, default_value = "litewire.db")] diff --git a/crates/litewire/tests/mysql_e2e.rs b/crates/litewire/tests/mysql_e2e.rs index 95d7d9c..35219bc 100644 --- a/crates/litewire/tests/mysql_e2e.rs +++ b/crates/litewire/tests/mysql_e2e.rs @@ -239,10 +239,7 @@ async fn multiple_connections() { // Second connection should see the data (same in-memory SQLite). let mut conn2 = connect(port).await; - let rows: Vec<(i64, String)> = conn2 - .query("SELECT id, val FROM shared") - .await - .unwrap(); + let rows: Vec<(i64, String)> = conn2.query("SELECT id, val FROM shared").await.unwrap(); assert_eq!(rows, vec![(1, "from_conn1".into())]); drop(conn2); @@ -315,10 +312,7 @@ async fn prepared_insert() { .unwrap(); assert_eq!( rows, - vec![ - (1, "Widget".into(), 10), - (2, "Gadget".into(), 20), - ] + vec![(1, "Widget".into(), 10), (2, "Gadget".into(), 20),] ); drop(conn); @@ -424,11 +418,9 @@ async fn on_duplicate_key_update() { .unwrap(); // MySQL ON DUPLICATE KEY UPDATE -> SQLite ON CONFLICT DO UPDATE - conn.query_drop( - "INSERT INTO kv (k, v) VALUES ('a', 99) ON DUPLICATE KEY UPDATE v = 99", - ) - .await - .unwrap(); + conn.query_drop("INSERT INTO kv (k, v) VALUES ('a', 99) ON DUPLICATE KEY UPDATE v = 99") + .await + .unwrap(); let rows: Vec<(String, i64)> = conn .query("SELECT k, v FROM kv WHERE k = 'a'") @@ -437,16 +429,11 @@ async fn on_duplicate_key_update() { assert_eq!(rows, vec![("a".into(), 99)]); // Insert new row (no conflict). - conn.query_drop( - "INSERT INTO kv (k, v) VALUES ('b', 2) ON DUPLICATE KEY UPDATE v = 2", - ) - .await - .unwrap(); - - let rows: Vec<(String, i64)> = conn - .query("SELECT k, v FROM kv ORDER BY k") + conn.query_drop("INSERT INTO kv (k, v) VALUES ('b', 2) ON DUPLICATE KEY UPDATE v = 2") .await .unwrap(); + + let rows: Vec<(String, i64)> = conn.query("SELECT k, v FROM kv ORDER BY k").await.unwrap(); assert_eq!(rows, vec![("a".into(), 99), ("b".into(), 2)]); drop(conn); @@ -472,10 +459,7 @@ async fn transaction_commit() { conn.query_drop("COMMIT").await.unwrap(); // Data should be visible after commit. - let rows: Vec<(i64, String)> = conn - .query("SELECT id, val FROM txn_t") - .await - .unwrap(); + let rows: Vec<(i64, String)> = conn.query("SELECT id, val FROM txn_t").await.unwrap(); assert_eq!(rows, vec![(1, "inside_txn".into())]); drop(conn); @@ -534,10 +518,7 @@ async fn transaction_atomicity() { .unwrap(); conn.query_drop("ROLLBACK").await.unwrap(); - let rows: Vec<(i64, i64)> = conn - .query("SELECT id, val FROM txn_atom") - .await - .unwrap(); + let rows: Vec<(i64, i64)> = conn.query("SELECT id, val FROM txn_atom").await.unwrap(); assert_eq!(rows, vec![(1, 100)]); drop(conn); @@ -585,15 +566,25 @@ async fn information_schema_columns() { .query("SELECT TABLE_NAME, COLUMN_NAME FROM information_schema.columns WHERE TABLE_NAME = 'users'") .await .unwrap(); - assert!(rows.len() >= 3, "expected at least 3 columns, got {}", rows.len()); + assert!( + rows.len() >= 3, + "expected at least 3 columns, got {}", + rows.len() + ); // Check that column names are present in the result. let col_names: Vec = rows .iter() .map(|r| r.get::(1).unwrap_or_default()) .collect(); assert!(col_names.contains(&"id".to_string()), "got: {col_names:?}"); - assert!(col_names.contains(&"name".to_string()), "got: {col_names:?}"); - assert!(col_names.contains(&"email".to_string()), "got: {col_names:?}"); + assert!( + col_names.contains(&"name".to_string()), + "got: {col_names:?}" + ); + assert!( + col_names.contains(&"email".to_string()), + "got: {col_names:?}" + ); drop(conn); } @@ -611,7 +602,11 @@ async fn describe_table() { // DESCRIBE should return column info via PRAGMA table_info. let rows: Vec = conn.query("DESCRIBE items").await.unwrap(); - assert!(rows.len() >= 3, "expected at least 3 columns, got {}", rows.len()); + assert!( + rows.len() >= 3, + "expected at least 3 columns, got {}", + rows.len() + ); drop(conn); } diff --git a/crates/litewire/tests/postgres_e2e.rs b/crates/litewire/tests/postgres_e2e.rs index bb93975..8dbee4a 100644 --- a/crates/litewire/tests/postgres_e2e.rs +++ b/crates/litewire/tests/postgres_e2e.rs @@ -191,10 +191,7 @@ async fn now_function_translates() { assert_eq!(rows.len(), 1); let val: &str = rows[0].get(0); // Should look like "2024-01-15 12:34:56". - assert!( - val.contains('-'), - "expected datetime string, got: {val}", - ); + assert!(val.contains('-'), "expected datetime string, got: {val}",); } #[tokio::test] @@ -205,10 +202,7 @@ async fn set_names_noop() { let client = connect(port).await; // SET NAMES should succeed as a no-op. - client - .execute("SET NAMES 'utf8mb4'", &[]) - .await - .unwrap(); + client.execute("SET NAMES 'utf8mb4'", &[]).await.unwrap(); // Connection still works after. let rows = client.query("SELECT 42", &[]).await.unwrap(); @@ -255,7 +249,10 @@ async fn empty_table_query() { let client = connect(port).await; client - .execute("CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", &[]) + .execute( + "CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", + &[], + ) .await .unwrap(); @@ -402,10 +399,7 @@ async fn large_result_set() { // Insert 100 rows. for i in 0..100 { client - .execute( - &format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), - &[], - ) + .execute(&format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), &[]) .await .unwrap(); } diff --git a/crates/litewire/tests/tds_e2e.rs b/crates/litewire/tests/tds_e2e.rs index 7bd859c..4baa32f 100644 --- a/crates/litewire/tests/tds_e2e.rs +++ b/crates/litewire/tests/tds_e2e.rs @@ -22,8 +22,7 @@ async fn start_litewire(port: u16) -> tokio::task::JoinHandle<()> { let addr: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap(); let backend = litewire::backend::Rusqlite::memory().unwrap(); let config = litewire::litewire_tds::TdsFrontendConfig { listen: addr }; - let frontend = - litewire::litewire_tds::TdsFrontend::new(config, std::sync::Arc::new(backend)); + let frontend = litewire::litewire_tds::TdsFrontend::new(config, std::sync::Arc::new(backend)); tokio::spawn(async move { frontend.serve().await.unwrap(); @@ -212,10 +211,7 @@ async fn empty_table_query() { .await .unwrap(); - let stream = client - .simple_query("SELECT * FROM empty_t") - .await - .unwrap(); + let stream = client.simple_query("SELECT * FROM empty_t").await.unwrap(); let rows = stream.into_first_result().await.unwrap(); assert!(rows.is_empty()); } @@ -537,10 +533,7 @@ async fn getdate_translates() { let row = stream.into_row().await.unwrap().unwrap(); let val: &str = row.get(0).unwrap(); // Should look like a datetime string. - assert!( - val.contains('-'), - "expected datetime string, got: {val}", - ); + assert!(val.contains('-'), "expected datetime string, got: {val}",); } #[tokio::test] diff --git a/rustfmt.toml b/rustfmt.toml new file mode 100644 index 0000000..71764a7 --- /dev/null +++ b/rustfmt.toml @@ -0,0 +1,3 @@ +# Pin rustfmt to its defaults so local runs match CI regardless of any +# rustfmt.toml in parent directories (stable rustfmt walks up the tree). +edition = "2024"