From 35692d33fa95379e0595b44226d00186097b3cf5 Mon Sep 17 00:00:00 2001 From: Luther Monson Date: Thu, 9 Jul 2026 19:03:36 -0700 Subject: [PATCH 1/2] compat: session/transaction rewrites + real MySQL/PG error codes + CI Adds a first CI workflow (fmt / clippy -D warnings / test on ubuntu-latest) and the fixes that were needed to make it green today. Compiler + clippy hygiene - Drop unused sqlparser::Expr import in translate/mysql.rs - Drop dead DONE_MORE constant in tds/token.rs - Remove unread PreparedStmt.param_count in mysql/handler.rs - Collapse redundant closures on pgwire portal parameter extraction - Fix manual case-insensitive cmp, manual strip, useless conversion, single-match, and too-many-args lints across the workspace - extract_where_value / trim_offset in metadata.rs marked as WIP scaffolding (underscored / #[allow(dead_code)]) rather than removed -- they are part of an in-progress WHERE-clause extraction rework Session / transaction compatibility (litewire-translate) - START TRANSACTION [READ ONLY|READ WRITE|WITH CONSISTENT SNAPSHOT] -> BEGIN - BEGIN [WORK|TRANSACTION [name]] -> BEGIN (fixes PDO::beginTransaction) - COMMIT / ROLLBACK [WORK|TRANSACTION [name]] normalized - ROLLBACK TRANSACTION TO [SAVEPOINT] name rewritten to ROLLBACK TO SAVEPOINT - SAVEPOINT / RELEASE SAVEPOINT / ROLLBACK TO SAVEPOINT verified passthrough - SET autocommit = 1|ON|TRUE|0|OFF|FALSE -> noop (with WARN on 0/OFF/FALSE since litewire does not emulate MySQL implicit transactions) - SET [SESSION|GLOBAL] TRANSACTION ISOLATION LEVEL ... -> noop with debug log - SET @@[session|global.]var forms routed to their plain-name handling - LOCK TABLES / UNLOCK TABLES -> noop with warn - LAST_INSERT_ID() -> last_insert_rowid(); ROW_COUNT() -> changes() - DATABASE() / VERSION() / USER() / CURRENT_USER() / SESSION_USER() / SYSTEM_USER() / CONNECTION_ID() -> constant projections that match the values used by the @@variable metadata fast path - SQL_CALC_FOUND_ROWS / FOUND_ROWS() intentionally NOT implemented -- would require session state; documented as a known gap. Error-code fidelity - New litewire-mysql::error_map: classifies backend error strings to real MySQL codes: SQLITE_BUSY -> 1205, unique/PK -> 1062 (23000), FK -> 1452 (23000), readonly -> 1290, else 1105. - New litewire-postgres::error_map: unique -> 23505, FK -> 23503, not-null -> 23502, check -> 23514, busy -> 55P03, readonly -> 25006, else XX000. - Both classifiers unit-tested. Prepared-statement cap - Per-connection cap of 1024 prepared statements; overflow returns MySQL error 1461 (ER_MAX_PREPARED_STMT_COUNT_REACHED). Protects the server against clients that never send COM_STMT_CLOSE. Cast safety - Route affected_rows / last_insert_id through .max(0) + TryFrom rather than reinterpret-casting negative i64 into u64. README truth pass - Replace unverified "Tested with WordPress/Laravel/..." with an accurate Compatibility section: mysql_e2e covers the wire protocol end-to-end; postgres/tds are wire-compatible for basic CRUD; TDS marked experimental (auth simplified, no SSL). Note that postgres+tds require build-time features. All flag names verified against crates/litewire/src/main.rs. --- .github/workflows/ci.yml | 42 ++ README.md | 36 +- crates/litewire-backend/src/hrana_client.rs | 97 +---- crates/litewire-backend/src/lib.rs | 20 +- .../litewire-backend/src/rusqlite_backend.rs | 103 ++--- .../tests/hrana_client_integration.rs | 104 +---- crates/litewire-hrana/src/http.rs | 90 +---- crates/litewire-hrana/src/lib.rs | 2 +- crates/litewire-hrana/src/types.rs | 45 +-- crates/litewire-mysql/src/error_map.rs | 149 +++++++ crates/litewire-mysql/src/handler.rs | 121 +++--- crates/litewire-mysql/src/lib.rs | 1 + crates/litewire-mysql/src/types.rs | 70 +--- crates/litewire-postgres/src/error_map.rs | 125 ++++++ crates/litewire-postgres/src/handler.rs | 124 ++---- crates/litewire-postgres/src/lib.rs | 5 +- crates/litewire-tds/src/handler.rs | 28 +- crates/litewire-tds/src/packet.rs | 12 +- crates/litewire-tds/src/token.rs | 55 +-- crates/litewire-translate/src/common.rs | 170 +++++--- crates/litewire-translate/src/lib.rs | 381 ++++++++++++++++-- crates/litewire-translate/src/metadata.rs | 189 +++------ crates/litewire-translate/src/mysql.rs | 101 ++--- crates/litewire-translate/src/postgres.rs | 52 +-- crates/litewire-translate/src/tds.rs | 59 +-- crates/litewire/src/lib.rs | 19 +- crates/litewire/src/main.rs | 6 +- crates/litewire/tests/mysql_e2e.rs | 258 +++--------- crates/litewire/tests/postgres_e2e.rs | 268 +++--------- crates/litewire/tests/tds_e2e.rs | 130 ++---- 30 files changed, 1299 insertions(+), 1563 deletions(-) create mode 100644 .github/workflows/ci.yml create mode 100644 crates/litewire-mysql/src/error_map.rs create mode 100644 crates/litewire-postgres/src/error_map.rs 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..8e7a91f 100644 --- a/crates/litewire-backend/src/hrana_client.rs +++ b/crates/litewire-backend/src/hrana_client.rs @@ -143,10 +143,7 @@ impl HranaClient { if resp.status().is_success() { Ok(()) } else { - Err(BackendError::Other(format!( - "health check returned {}", - resp.status() - ))) + Err(BackendError::Other(format!("health check returned {}", resp.status()))) } } @@ -161,10 +158,7 @@ impl HranaClient { let request = PipelineRequest { baton: None, requests: vec![StreamRequest::Execute(ExecuteRequest { - stmt: StmtRequest { - sql: sql.to_string(), - args, - }, + stmt: StmtRequest { sql: sql.to_string(), args }, })], }; @@ -179,9 +173,7 @@ impl HranaClient { if !resp.status().is_success() { let status = resp.status(); let body = resp.text().await.unwrap_or_default(); - return Err(BackendError::Other(format!( - "sqld returned {status}: {body}" - ))); + return Err(BackendError::Other(format!("sqld returned {status}: {body}"))); } let pipeline: PipelineResponse = resp @@ -218,10 +210,7 @@ impl Backend for HranaClient { let exec = self.execute_pipeline(sql, params).await?; Ok(ExecuteResult { affected_rows: exec.affected_row_count, - last_insert_rowid: exec - .last_insert_rowid - .as_deref() - .and_then(|s| s.parse().ok()), + last_insert_rowid: exec.last_insert_rowid.as_deref().and_then(|s| s.parse().ok()), }) } } @@ -232,18 +221,12 @@ impl Backend for HranaClient { fn value_to_hrana(val: &Value) -> HranaValue { match val { Value::Null => HranaValue::Null, - Value::Integer(i) => HranaValue::Integer { - value: i.to_string(), - }, + Value::Integer(i) => HranaValue::Integer { 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 { - base64: base64::engine::general_purpose::STANDARD.encode(b), - } + HranaValue::Blob { base64: base64::engine::general_purpose::STANDARD.encode(b) } } } } @@ -257,9 +240,7 @@ fn response_value_to_backend(rv: &ResponseValue) -> Value { ResponseValue::Text { value } => Value::Text(value.clone()), ResponseValue::Blob { base64: b64 } => { use base64::Engine; - let bytes = base64::engine::general_purpose::STANDARD - .decode(b64) - .unwrap_or_default(); + let bytes = base64::engine::general_purpose::STANDARD.decode(b64).unwrap_or_default(); Value::Blob(bytes) } } @@ -270,17 +251,11 @@ fn execute_response_to_result_set(exec: ExecuteResponse) -> ResultSet { let columns = exec .cols .iter() - .map(|c| Column { - name: c.name.clone(), - decltype: c.decltype.clone(), - }) + .map(|c| Column { name: c.name.clone(), decltype: c.decltype.clone() }) .collect(); - let rows = exec - .rows - .iter() - .map(|row| row.iter().map(response_value_to_backend).collect()) - .collect(); + let rows = + exec.rows.iter().map(|row| row.iter().map(response_value_to_backend).collect()).collect(); ResultSet { columns, rows } } @@ -326,9 +301,7 @@ mod tests { match value_to_hrana(&Value::Blob(data.clone())) { HranaValue::Blob { base64: b64 } => { use base64::Engine; - let decoded = base64::engine::general_purpose::STANDARD - .decode(&b64) - .unwrap(); + let decoded = base64::engine::general_purpose::STANDARD.decode(&b64).unwrap(); assert_eq!(decoded, data); } other => panic!("expected Blob, got: {other:?}"), @@ -339,19 +312,9 @@ mod tests { fn response_value_roundtrip() { let cases = vec![ (Value::Null, ResponseValue::Null), - ( - Value::Integer(123), - ResponseValue::Integer { - value: "123".into(), - }, - ), + (Value::Integer(123), ResponseValue::Integer { value: "123".into() }), (Value::Float(3.14), ResponseValue::Float { value: 3.14 }), - ( - Value::Text("test".into()), - ResponseValue::Text { - value: "test".into(), - }, - ), + (Value::Text("test".into()), ResponseValue::Text { value: "test".into() }), ]; for (val, rv) in cases { @@ -364,31 +327,17 @@ mod tests { fn execute_response_to_result_set_maps_correctly() { let exec = ExecuteResponse { cols: vec![ - ColResponse { - name: "id".into(), - decltype: Some("INTEGER".into()), - }, - ColResponse { - name: "name".into(), - decltype: Some("TEXT".into()), - }, + ColResponse { name: "id".into(), decltype: Some("INTEGER".into()) }, + ColResponse { name: "name".into(), decltype: Some("TEXT".into()) }, ], rows: vec![ vec![ - ResponseValue::Integer { - value: "1".into(), - }, - ResponseValue::Text { - value: "alice".into(), - }, + ResponseValue::Integer { value: "1".into() }, + ResponseValue::Text { value: "alice".into() }, ], vec![ - ResponseValue::Integer { - value: "2".into(), - }, - ResponseValue::Text { - value: "bob".into(), - }, + ResponseValue::Integer { value: "2".into() }, + ResponseValue::Text { value: "bob".into() }, ], ], affected_row_count: 0, @@ -419,10 +368,8 @@ mod tests { #[test] fn hrana_generic_error() { - let err = hrana_error_to_backend(ErrorResponse { - message: "something broke".into(), - code: None, - }); + let err = + hrana_error_to_backend(ErrorResponse { message: "something broke".into(), code: None }); match err { BackendError::Other(msg) => assert_eq!(msg, "something broke"), other => panic!("expected Other error, got: {other:?}"), diff --git a/crates/litewire-backend/src/lib.rs b/crates/litewire-backend/src/lib.rs index d1e617c..fb12cbe 100644 --- a/crates/litewire-backend/src/lib.rs +++ b/crates/litewire-backend/src/lib.rs @@ -120,10 +120,7 @@ mod tests { #[test] fn display_blob() { - assert_eq!( - format!("{}", Value::Blob(vec![0xDE, 0xAD])), - "" - ); + assert_eq!(format!("{}", Value::Blob(vec![0xDE, 0xAD])), ""); assert_eq!(format!("{}", Value::Blob(vec![])), ""); } @@ -156,20 +153,14 @@ mod tests { #[test] fn column_with_decltype() { - let c = Column { - name: "id".into(), - decltype: Some("INTEGER".into()), - }; + let c = Column { name: "id".into(), decltype: Some("INTEGER".into()) }; assert_eq!(c.name, "id"); assert_eq!(c.decltype.as_deref(), Some("INTEGER")); } #[test] fn column_without_decltype() { - let c = Column { - name: "expr".into(), - decltype: None, - }; + let c = Column { name: "expr".into(), decltype: None }; assert!(c.decltype.is_none()); } @@ -177,10 +168,7 @@ mod tests { #[test] fn execute_result_no_insert() { - let r = ExecuteResult { - affected_rows: 3, - last_insert_rowid: None, - }; + let r = ExecuteResult { affected_rows: 3, last_insert_rowid: None }; assert_eq!(r.affected_rows, 3); assert!(r.last_insert_rowid.is_none()); } diff --git a/crates/litewire-backend/src/rusqlite_backend.rs b/crates/litewire-backend/src/rusqlite_backend.rs index de086d2..1f8d878 100644 --- a/crates/litewire-backend/src/rusqlite_backend.rs +++ b/crates/litewire-backend/src/rusqlite_backend.rs @@ -31,9 +31,7 @@ impl Rusqlite { conn.execute_batch("PRAGMA journal_mode=WAL; PRAGMA busy_timeout=5000;") .map_err(|e| BackendError::Sqlite(e.to_string()))?; - Ok(Self { - conn: Arc::new(Mutex::new(conn)), - }) + Ok(Self { conn: Arc::new(Mutex::new(conn)) }) } /// Open an in-memory SQLite database. @@ -42,11 +40,8 @@ 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()))?; - Ok(Self { - conn: Arc::new(Mutex::new(conn)), - }) + let conn = Connection::open_in_memory().map_err(|e| BackendError::Sqlite(e.to_string()))?; + Ok(Self { conn: Arc::new(Mutex::new(conn)) }) } } @@ -73,9 +68,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())), } } @@ -89,9 +82,7 @@ impl Backend for Rusqlite { task::spawn_blocking(move || { let conn = conn.lock(); - let mut stmt = conn - .prepare(&sql) - .map_err(|e| BackendError::Sqlite(e.to_string()))?; + let mut stmt = conn.prepare(&sql).map_err(|e| BackendError::Sqlite(e.to_string()))?; let col_count = stmt.column_count(); let columns: Vec = (0..col_count) @@ -114,17 +105,13 @@ impl Backend for Rusqlite { 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); } - Ok(ResultSet { - columns, - rows: result_rows, - }) + Ok(ResultSet { columns, rows: result_rows }) }) .await .map_err(|e| BackendError::Other(format!("spawn_blocking join error: {e}")))? @@ -175,20 +162,14 @@ mod tests { .unwrap(); let result = backend - .execute( - "INSERT INTO users (name) VALUES (?1)", - &[Value::Text("Alice".into())], - ) + .execute("INSERT INTO users (name) VALUES (?1)", &[Value::Text("Alice".into())]) .await .unwrap(); assert_eq!(result.affected_rows, 1); assert_eq!(result.last_insert_rowid, Some(1)); let result = backend - .execute( - "INSERT INTO users (name) VALUES (?1)", - &[Value::Text("Bob".into())], - ) + .execute("INSERT INTO users (name) VALUES (?1)", &[Value::Text("Bob".into())]) .await .unwrap(); assert_eq!(result.last_insert_rowid, Some(2)); @@ -201,10 +182,8 @@ mod tests { assert_eq!(rs.rows[0][1], Value::Text("Alice".into())); assert_eq!(rs.rows[1][1], Value::Text("Bob".into())); - let rs = backend - .query("SELECT * FROM users WHERE id = ?1", &[Value::Integer(1)]) - .await - .unwrap(); + let rs = + backend.query("SELECT * FROM users WHERE id = ?1", &[Value::Integer(1)]).await.unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][1], Value::Text("Alice".into())); } @@ -213,10 +192,7 @@ mod tests { async fn types_roundtrip() { let backend = Rusqlite::memory().unwrap(); backend - .execute( - "CREATE TABLE typed (i INTEGER, r REAL, t TEXT, b BLOB)", - &[], - ) + .execute("CREATE TABLE typed (i INTEGER, r REAL, t TEXT, b BLOB)", &[]) .await .unwrap(); @@ -244,10 +220,7 @@ mod tests { async fn null_handling() { let backend = Rusqlite::memory().unwrap(); backend.execute("CREATE TABLE t (v TEXT)", &[]).await.unwrap(); - backend - .execute("INSERT INTO t VALUES (?1)", &[Value::Null]) - .await - .unwrap(); + backend.execute("INSERT INTO t VALUES (?1)", &[Value::Null]).await.unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.rows[0][0], Value::Null); @@ -256,10 +229,7 @@ mod tests { #[tokio::test] async fn empty_table_query() { let backend = Rusqlite::memory().unwrap(); - backend - .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.columns.len(), 2); @@ -269,27 +239,20 @@ mod tests { #[tokio::test] async fn multiple_params() { let backend = Rusqlite::memory().unwrap(); - backend - .execute("CREATE TABLE t (a INTEGER, b TEXT, c REAL)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (a INTEGER, b TEXT, c REAL)", &[]).await.unwrap(); backend .execute( "INSERT INTO t VALUES (?1, ?2, ?3)", - &[ - Value::Integer(1), - Value::Text("hello".into()), - Value::Float(9.99), - ], + &[Value::Integer(1), Value::Text("hello".into()), Value::Float(9.99)], ) .await .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); @@ -299,10 +262,7 @@ mod tests { #[tokio::test] async fn affected_rows_count() { let backend = Rusqlite::memory().unwrap(); - backend - .execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); for i in 0..5 { backend @@ -314,10 +274,8 @@ mod tests { .unwrap(); } - let result = backend - .execute("DELETE FROM t WHERE id >= ?1", &[Value::Integer(3)]) - .await - .unwrap(); + let result = + backend.execute("DELETE FROM t WHERE id >= ?1", &[Value::Integer(3)]).await.unwrap(); assert_eq!(result.affected_rows, 2); } @@ -338,16 +296,10 @@ mod tests { #[tokio::test] async fn blob_roundtrip() { let backend = Rusqlite::memory().unwrap(); - backend - .execute("CREATE TABLE t (data BLOB)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (data BLOB)", &[]).await.unwrap(); let data = vec![0x00, 0xFF, 0xDE, 0xAD, 0xBE, 0xEF]; - backend - .execute("INSERT INTO t VALUES (?1)", &[Value::Blob(data.clone())]) - .await - .unwrap(); + backend.execute("INSERT INTO t VALUES (?1)", &[Value::Blob(data.clone())]).await.unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.rows[0][0], Value::Blob(data)); @@ -357,10 +309,7 @@ mod tests { async fn last_insert_rowid_increments() { let backend = Rusqlite::memory().unwrap(); backend - .execute( - "CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, v TEXT)", - &[], - ) + .execute("CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, v TEXT)", &[]) .await .unwrap(); diff --git a/crates/litewire-backend/tests/hrana_client_integration.rs b/crates/litewire-backend/tests/hrana_client_integration.rs index 4061844..4120c7e 100644 --- a/crates/litewire-backend/tests/hrana_client_integration.rs +++ b/crates/litewire-backend/tests/hrana_client_integration.rs @@ -16,9 +16,7 @@ use litewire_hrana::HranaFrontendConfig; /// Start a Hrana server on a random port, return the client and server task handle. async fn start_server() -> (HranaClient, tokio::task::JoinHandle<()>) { // Bind to port 0 to get a random available port. - let listener = tokio::net::TcpListener::bind("127.0.0.1:0") - .await - .expect("failed to bind"); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("failed to bind"); let addr: SocketAddr = listener.local_addr().unwrap(); let backend = Rusqlite::memory().expect("failed to create in-memory SQLite"); @@ -63,10 +61,7 @@ async fn create_table_and_insert() { // Insert let result = client - .execute( - "INSERT INTO users (name) VALUES (?)", - &[Value::Text("alice".into())], - ) + .execute("INSERT INTO users (name) VALUES (?)", &[Value::Text("alice".into())]) .await .expect("INSERT failed"); assert_eq!(result.affected_rows, 1); @@ -74,10 +69,7 @@ async fn create_table_and_insert() { // Insert another let result = client - .execute( - "INSERT INTO users (name) VALUES (?)", - &[Value::Text("bob".into())], - ) + .execute("INSERT INTO users (name) VALUES (?)", &[Value::Text("bob".into())]) .await .expect("INSERT failed"); assert_eq!(result.affected_rows, 1); @@ -129,10 +121,7 @@ async fn query_rows() { async fn query_with_params() { let (client, _server) = start_server().await; - client - .execute("CREATE TABLE kv (key TEXT, val TEXT)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE kv (key TEXT, val TEXT)", &[]).await.unwrap(); client .execute( "INSERT INTO kv (key, val) VALUES (?, ?)", @@ -148,13 +137,8 @@ async fn query_with_params() { .await .unwrap(); - let rs = client - .query( - "SELECT val FROM kv WHERE key = ?", - &[Value::Text("b".into())], - ) - .await - .unwrap(); + let rs = + client.query("SELECT val FROM kv WHERE key = ?", &[Value::Text("b".into())]).await.unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Text("2".into())); @@ -164,15 +148,9 @@ async fn query_with_params() { async fn query_empty_result() { let (client, _server) = start_server().await; - client - .execute("CREATE TABLE empty_test (id INTEGER)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE empty_test (id INTEGER)", &[]).await.unwrap(); - let rs = client - .query("SELECT id FROM empty_test", &[]) - .await - .unwrap(); + let rs = client.query("SELECT id FROM empty_test", &[]).await.unwrap(); assert_eq!(rs.columns.len(), 1); assert!(rs.rows.is_empty()); @@ -182,24 +160,15 @@ async fn query_empty_result() { async fn blob_roundtrip() { let (client, _server) = start_server().await; - client - .execute("CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", &[]).await.unwrap(); let blob_data = vec![0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0xFF]; client - .execute( - "INSERT INTO blobs (data) VALUES (?)", - &[Value::Blob(blob_data.clone())], - ) + .execute("INSERT INTO blobs (data) VALUES (?)", &[Value::Blob(blob_data.clone())]) .await .unwrap(); - let rs = client - .query("SELECT data FROM blobs WHERE id = 1", &[]) - .await - .unwrap(); + let rs = client.query("SELECT data FROM blobs WHERE id = 1", &[]).await.unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Blob(blob_data)); @@ -209,22 +178,13 @@ async fn blob_roundtrip() { async fn null_values() { let (client, _server) = start_server().await; + client.execute("CREATE TABLE nullable (id INTEGER, val TEXT)", &[]).await.unwrap(); client - .execute("CREATE TABLE nullable (id INTEGER, val TEXT)", &[]) - .await - .unwrap(); - client - .execute( - "INSERT INTO nullable (id, val) VALUES (?, ?)", - &[Value::Integer(1), Value::Null], - ) + .execute("INSERT INTO nullable (id, val) VALUES (?, ?)", &[Value::Integer(1), Value::Null]) .await .unwrap(); - let rs = client - .query("SELECT val FROM nullable WHERE id = 1", &[]) - .await - .unwrap(); + let rs = client.query("SELECT val FROM nullable WHERE id = 1", &[]).await.unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Null); @@ -234,9 +194,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(); @@ -250,36 +208,16 @@ async fn sql_error_returns_backend_error() { async fn update_and_delete() { let (client, _server) = start_server().await; - client - .execute("CREATE TABLE counters (name TEXT, val INTEGER)", &[]) - .await - .unwrap(); - client - .execute( - "INSERT INTO counters VALUES ('hits', 0)", - &[], - ) - .await - .unwrap(); + client.execute("CREATE TABLE counters (name TEXT, val INTEGER)", &[]).await.unwrap(); + client.execute("INSERT INTO counters VALUES ('hits', 0)", &[]).await.unwrap(); - let result = client - .execute( - "UPDATE counters SET val = val + 1 WHERE name = 'hits'", - &[], - ) - .await - .unwrap(); + let result = + client.execute("UPDATE counters SET val = val + 1 WHERE name = 'hits'", &[]).await.unwrap(); assert_eq!(result.affected_rows, 1); - let result = client - .execute("DELETE FROM counters WHERE name = 'hits'", &[]) - .await - .unwrap(); + let result = client.execute("DELETE FROM counters WHERE name = 'hits'", &[]).await.unwrap(); assert_eq!(result.affected_rows, 1); - let rs = client - .query("SELECT COUNT(*) FROM counters", &[]) - .await - .unwrap(); + let rs = client.query("SELECT COUNT(*) FROM counters", &[]).await.unwrap(); assert_eq!(rs.rows[0][0], Value::Integer(0)); } diff --git a/crates/litewire-hrana/src/http.rs b/crates/litewire-hrana/src/http.rs index e71bad7..dbc0172 100644 --- a/crates/litewire-hrana/src/http.rs +++ b/crates/litewire-hrana/src/http.rs @@ -44,19 +44,14 @@ async fn pipeline_handler( for stream_req in &req.requests { let result = match stream_req { StreamRequest::Execute(exec) => execute_stmt(&state.backend, &exec.stmt).await, - StreamRequest::Close => Ok(StreamResult::Ok { - response: StreamResponse::Close, - }), + StreamRequest::Close => Ok(StreamResult::Ok { response: StreamResponse::Close }), }; match result { Ok(r) => results.push(r), Err(e) => { results.push(StreamResult::Error { - error: ErrorResponse { - message: e.to_string(), - code: None, - }, + error: ErrorResponse { message: e.to_string(), code: None }, }); } } @@ -74,11 +69,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(); @@ -92,10 +84,7 @@ async fn execute_stmt( let cols: Vec = rs .columns .iter() - .map(|c| ColResponse { - name: c.name.clone(), - decltype: c.decltype.clone(), - }) + .map(|c| ColResponse { name: c.name.clone(), decltype: c.decltype.clone() }) .collect(); let rows: Vec> = rs @@ -157,12 +146,7 @@ mod tests { async fn health_check() { let app = build_router(test_backend()); let resp = app - .oneshot( - Request::builder() - .uri("/health") - .body(Body::empty()) - .unwrap(), - ) + .oneshot(Request::builder().uri("/health").body(Body::empty()).unwrap()) .await .unwrap(); @@ -175,12 +159,7 @@ mod tests { async fn version_endpoint() { let app = build_router(test_backend()); let resp = app - .oneshot( - Request::builder() - .uri("/version") - .body(Body::empty()) - .unwrap(), - ) + .oneshot(Request::builder().uri("/version").body(Body::empty()).unwrap()) .await .unwrap(); @@ -252,17 +231,8 @@ mod tests { let backend = test_backend(); // Pre-create table. - backend - .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) - .await - .unwrap(); - backend - .execute( - "INSERT INTO t VALUES (1, 'Alice')", - &[], - ) - .await - .unwrap(); + backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); + backend.execute("INSERT INTO t VALUES (1, 'Alice')", &[]).await.unwrap(); let app = build_router(backend); @@ -356,31 +326,18 @@ 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] async fn pipeline_mutation_returns_affected_rows() { let backend = test_backend(); - backend - .execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]) - .await - .unwrap(); - backend - .execute("INSERT INTO t VALUES (1, 'a')", &[]) - .await - .unwrap(); - backend - .execute("INSERT INTO t VALUES (2, 'b')", &[]) - .await - .unwrap(); - backend - .execute("INSERT INTO t VALUES (3, 'c')", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + backend.execute("INSERT INTO t VALUES (1, 'a')", &[]).await.unwrap(); + backend.execute("INSERT INTO t VALUES (2, 'b')", &[]).await.unwrap(); + backend.execute("INSERT INTO t VALUES (3, 'c')", &[]).await.unwrap(); let app = build_router(backend); @@ -417,10 +374,7 @@ mod tests { async fn pipeline_insert_returns_last_insert_rowid() { let backend = test_backend(); backend - .execute( - "CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, val TEXT)", - &[], - ) + .execute("CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, val TEXT)", &[]) .await .unwrap(); @@ -486,10 +440,7 @@ mod tests { #[tokio::test] async fn pipeline_pragma_is_query() { let backend = test_backend(); - backend - .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); let app = build_router(backend); @@ -528,10 +479,7 @@ mod tests { #[tokio::test] async fn pipeline_with_blob_param() { let backend = test_backend(); - backend - .execute("CREATE TABLE t (data BLOB)", &[]) - .await - .unwrap(); + backend.execute("CREATE TABLE t (data BLOB)", &[]).await.unwrap(); let app = build_router(backend); 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..1db16eb 100644 --- a/crates/litewire-hrana/src/types.rs +++ b/crates/litewire-hrana/src/types.rs @@ -47,16 +47,13 @@ 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 } => { use base64::Engine; - let bytes = base64::engine::general_purpose::STANDARD - .decode(base64) - .unwrap_or_default(); + let bytes = + base64::engine::general_purpose::STANDARD.decode(base64).unwrap_or_default(); Value::Blob(bytes) } } @@ -115,18 +112,12 @@ impl ResponseValue { pub fn from_backend_value(val: &Value) -> Self { match val { Value::Null => Self::Null, - Value::Integer(i) => Self::Integer { - value: i.to_string(), - }, + Value::Integer(i) => Self::Integer { 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 { - base64: base64::engine::general_purpose::STANDARD.encode(b), - } + Self::Blob { base64: base64::engine::general_purpose::STANDARD.encode(b) } } } } @@ -153,17 +144,13 @@ 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))); } #[test] fn integer_invalid_to_backend() { - let v = HranaValue::Integer { - value: "not_a_number".into(), - }; + let v = HranaValue::Integer { value: "not_a_number".into() }; assert!(matches!(v.to_backend_value(), Value::Integer(0))); } @@ -175,9 +162,7 @@ mod tests { #[test] fn text_to_backend() { - let v = HranaValue::Text { - value: "hello".into(), - }; + let v = HranaValue::Text { value: "hello".into() }; assert!(matches!(v.to_backend_value(), Value::Text(s) if s == "hello")); } @@ -192,9 +177,7 @@ mod tests { #[test] fn blob_invalid_base64_to_backend() { - let v = HranaValue::Blob { - base64: "!!!invalid!!!".into(), - }; + let v = HranaValue::Blob { base64: "!!!invalid!!!".into() }; // Invalid base64 should return empty blob. assert!(matches!(v.to_backend_value(), Value::Blob(b) if b.is_empty())); } @@ -241,9 +224,7 @@ mod tests { match rv { ResponseValue::Blob { base64: encoded } => { use base64::Engine; - let decoded = base64::engine::general_purpose::STANDARD - .decode(&encoded) - .unwrap(); + let decoded = base64::engine::general_purpose::STANDARD.decode(&encoded).unwrap(); assert_eq!(decoded, data); } other => panic!("expected Blob, got: {other:?}"), @@ -297,9 +278,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..d4bc742 100644 --- a/crates/litewire-mysql/src/handler.rs +++ b/crates/litewire-mysql/src/handler.rs @@ -11,19 +11,21 @@ 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 { - StatusFlags::SERVER_STATUS_IN_TRANS - } else { - StatusFlags::empty() - }; - OkResponse { - affected_rows, - last_insert_id, - status_flags, - ..OkResponse::default() - } + let status_flags = + if in_transaction { StatusFlags::SERVER_STATUS_IN_TRANS } else { StatusFlags::empty() }; + OkResponse { affected_rows, last_insert_id, status_flags, ..OkResponse::default() } } use crate::types::sqlite_to_mysql_column_type; @@ -34,8 +36,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. @@ -51,12 +51,7 @@ pub struct LiteWireHandler { impl LiteWireHandler { pub fn new(backend: SharedBackend) -> Self { - Self { - backend, - stmts: HashMap::new(), - next_stmt_id: 1, - in_transaction: false, - } + Self { backend, stmts: HashMap::new(), next_stmt_id: 1, in_transaction: false } } /// Execute a query and write result set. @@ -96,11 +91,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 +104,15 @@ 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 +133,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 +160,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 +202,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 +238,31 @@ 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 } @@ -287,9 +292,7 @@ impl AsyncMysqlShim for LiteWireHandler { // Noop statements (empty SQL from SET NAMES etc.) if sql.is_empty() { - return results - .completed(ok_response(0, 0, self.in_transaction)) - .await; + return results.completed(ok_response(0, 0, self.in_transaction)).await; } match kind { @@ -315,9 +318,7 @@ impl AsyncMysqlShim for LiteWireHandler { Ok(r) => r, Err(e) => { warn!("SQL translation error: {e}"); - return results - .error(ErrorKind::ER_PARSE_ERROR, e.to_string().as_bytes()) - .await; + return results.error(ErrorKind::ER_PARSE_ERROR, e.to_string().as_bytes()).await; } }; @@ -335,9 +336,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-mysql/src/types.rs b/crates/litewire-mysql/src/types.rs index 4451980..61ddd36 100644 --- a/crates/litewire-mysql/src/types.rs +++ b/crates/litewire-mysql/src/types.rs @@ -38,55 +38,34 @@ mod tests { #[test] fn type_mapping_integer() { - assert_eq!( - sqlite_to_mysql_column_type(Some("INTEGER")), - ColumnType::MYSQL_TYPE_LONGLONG - ); + assert_eq!(sqlite_to_mysql_column_type(Some("INTEGER")), ColumnType::MYSQL_TYPE_LONGLONG); } #[test] fn type_mapping_int_substring() { // "BIGINT", "TINYINT", etc. all contain "INT" - assert_eq!( - sqlite_to_mysql_column_type(Some("BIGINT")), - ColumnType::MYSQL_TYPE_LONGLONG - ); - assert_eq!( - sqlite_to_mysql_column_type(Some("TINYINT")), - ColumnType::MYSQL_TYPE_LONGLONG - ); + assert_eq!(sqlite_to_mysql_column_type(Some("BIGINT")), ColumnType::MYSQL_TYPE_LONGLONG); + assert_eq!(sqlite_to_mysql_column_type(Some("TINYINT")), ColumnType::MYSQL_TYPE_LONGLONG); } #[test] fn type_mapping_real() { - assert_eq!( - sqlite_to_mysql_column_type(Some("REAL")), - ColumnType::MYSQL_TYPE_DOUBLE - ); + assert_eq!(sqlite_to_mysql_column_type(Some("REAL")), ColumnType::MYSQL_TYPE_DOUBLE); } #[test] fn type_mapping_float() { - assert_eq!( - sqlite_to_mysql_column_type(Some("FLOAT")), - ColumnType::MYSQL_TYPE_DOUBLE - ); + assert_eq!(sqlite_to_mysql_column_type(Some("FLOAT")), ColumnType::MYSQL_TYPE_DOUBLE); } #[test] fn type_mapping_double() { - assert_eq!( - sqlite_to_mysql_column_type(Some("DOUBLE")), - ColumnType::MYSQL_TYPE_DOUBLE - ); + assert_eq!(sqlite_to_mysql_column_type(Some("DOUBLE")), ColumnType::MYSQL_TYPE_DOUBLE); } #[test] fn type_mapping_text() { - assert_eq!( - sqlite_to_mysql_column_type(Some("TEXT")), - ColumnType::MYSQL_TYPE_VAR_STRING - ); + assert_eq!(sqlite_to_mysql_column_type(Some("TEXT")), ColumnType::MYSQL_TYPE_VAR_STRING); } #[test] @@ -107,26 +86,17 @@ mod tests { #[test] fn type_mapping_blob() { - assert_eq!( - sqlite_to_mysql_column_type(Some("BLOB")), - ColumnType::MYSQL_TYPE_BLOB - ); + assert_eq!(sqlite_to_mysql_column_type(Some("BLOB")), ColumnType::MYSQL_TYPE_BLOB); } #[test] fn type_mapping_bytea() { - assert_eq!( - sqlite_to_mysql_column_type(Some("BYTEA")), - ColumnType::MYSQL_TYPE_BLOB - ); + assert_eq!(sqlite_to_mysql_column_type(Some("BYTEA")), ColumnType::MYSQL_TYPE_BLOB); } #[test] fn type_mapping_none_defaults_to_string() { - assert_eq!( - sqlite_to_mysql_column_type(None), - ColumnType::MYSQL_TYPE_VAR_STRING - ); + assert_eq!(sqlite_to_mysql_column_type(None), ColumnType::MYSQL_TYPE_VAR_STRING); } #[test] @@ -140,21 +110,9 @@ mod tests { #[test] fn type_mapping_case_insensitive() { // The function uppercases, so lowercase should work too. - assert_eq!( - sqlite_to_mysql_column_type(Some("integer")), - ColumnType::MYSQL_TYPE_LONGLONG - ); - assert_eq!( - sqlite_to_mysql_column_type(Some("real")), - ColumnType::MYSQL_TYPE_DOUBLE - ); - assert_eq!( - sqlite_to_mysql_column_type(Some("text")), - ColumnType::MYSQL_TYPE_VAR_STRING - ); - assert_eq!( - sqlite_to_mysql_column_type(Some("blob")), - ColumnType::MYSQL_TYPE_BLOB - ); + assert_eq!(sqlite_to_mysql_column_type(Some("integer")), ColumnType::MYSQL_TYPE_LONGLONG); + assert_eq!(sqlite_to_mysql_column_type(Some("real")), ColumnType::MYSQL_TYPE_DOUBLE); + assert_eq!(sqlite_to_mysql_column_type(Some("text")), ColumnType::MYSQL_TYPE_VAR_STRING); + assert_eq!(sqlite_to_mysql_column_type(Some("blob")), ColumnType::MYSQL_TYPE_BLOB); } } diff --git a/crates/litewire-postgres/src/error_map.rs b/crates/litewire-postgres/src/error_map.rs new file mode 100644 index 0000000..139639f --- /dev/null +++ b/crates/litewire-postgres/src/error_map.rs @@ -0,0 +1,125 @@ +//! 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..45162df 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. @@ -33,16 +34,13 @@ pub struct PostgresHandler { impl PostgresHandler { pub fn new(backend: SharedBackend) -> Self { - Self { - backend, - query_parser: Arc::new(NoopQueryParser::new()), - } + Self { backend, query_parser: Arc::new(NoopQueryParser::new()) } } /// 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)); @@ -82,11 +80,7 @@ impl PostgresHandler { params: &[Value], format: &Format, ) -> PgWireResult> { - let rs = self - .backend - .query(sql, params) - .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + let rs = self.backend.query(sql, params).await.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 @@ -106,13 +100,7 @@ impl PostgresHandler { .map(value_to_pg_type) .unwrap_or(Type::TEXT) }; - FieldInfo::new( - col.name.clone(), - None, - None, - pg_type, - format.format_for(idx), - ) + FieldInfo::new(col.name.clone(), None, None, pg_type, format.format_for(idx)) }) .collect(); @@ -127,10 +115,7 @@ impl PostgresHandler { rows.push(encoder.finish()); } - Ok(Response::Query(QueryResponse::new( - schema, - stream::iter(rows), - ))) + Ok(Response::Query(QueryResponse::new(schema, stream::iter(rows)))) } /// Execute a mutation (INSERT/UPDATE/DELETE/DDL) and return an execution response. @@ -143,10 +128,7 @@ impl PostgresHandler { // Transaction commands need special Response variants, handle before // the generic execute path to avoid double-execution. if *kind == StatementKind::Transaction { - self.backend - .execute(sql, params) - .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + self.backend.execute(sql, params).await.map_err(|e| pg_backend_error(&e))?; let upper = sql.trim().to_ascii_uppercase(); if upper.starts_with("BEGIN") || upper.starts_with("START") { @@ -161,11 +143,7 @@ impl PostgresHandler { return Ok(Response::Execution(Tag::new("OK"))); } - let result = self - .backend - .execute(sql, params) - .await - .map_err(|e| PgWireError::ApiError(Box::new(e)))?; + let result = self.backend.execute(sql, params).await.map_err(|e| pg_backend_error(&e))?; let tag_name = match kind { StatementKind::Mutation => { @@ -207,11 +185,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 @@ -228,13 +202,7 @@ impl PostgresHandler { .map(value_to_pg_type) .unwrap_or(Type::TEXT) }; - FieldInfo::new( - col.name.clone(), - None, - None, - pg_type, - format.format_for(idx), - ) + FieldInfo::new(col.name.clone(), None, None, pg_type, format.format_for(idx)) }) .collect()), Err(_) => Ok(vec![]), @@ -256,11 +224,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) => { @@ -282,12 +246,7 @@ fn encode_value( fn extract_params(portal: &Portal) -> Vec { let mut values = Vec::with_capacity(portal.parameter_len()); for i in 0..portal.parameter_len() { - let param_type = portal - .statement - .parameter_types - .get(i) - .cloned() - .unwrap_or(Type::TEXT); + let param_type = portal.statement.parameter_types.get(i).cloned().unwrap_or(Type::TEXT); let val = match ¶m_type { t if *t == Type::BOOL => portal @@ -312,7 +271,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 +283,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 +303,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 +314,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 +337,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,20 +347,16 @@ 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) - .await? + self.exec_query(&sqlite_sql, &[], &Format::UnifiedText).await? } TranslateResult::Sql(sqlite_sql) => { if sqlite_sql.is_empty() { Response::Execution(Tag::new("OK")) } else { - self.exec_query(&sqlite_sql, &[], &Format::UnifiedText) - .await? + self.exec_query(&sqlite_sql, &[], &Format::UnifiedText).await? } } }; @@ -431,8 +399,7 @@ impl ExtendedQueryHandler for PostgresHandler { } let params = extract_params(portal); - self.exec_query(&sqlite_sql, ¶ms, &portal.result_column_format) - .await + self.exec_query(&sqlite_sql, ¶ms, &portal.result_column_format).await } async fn do_describe_statement( @@ -443,16 +410,12 @@ impl ExtendedQueryHandler for PostgresHandler { where C: ClientInfo + Unpin + Send + Sync, { - let (sqlite_sql, kind) = self - .translate_sql(&stmt.statement) - .map_err(|e| pg_error(&e))?; + let (sqlite_sql, kind) = self.translate_sql(&stmt.statement).map_err(|e| pg_error(&e))?; let param_types = stmt.parameter_types.clone(); if kind == StatementKind::Query && !sqlite_sql.is_empty() { - let fields = self - .probe_columns(&sqlite_sql, &Format::UnifiedBinary) - .await?; + let fields = self.probe_columns(&sqlite_sql, &Format::UnifiedBinary).await?; Ok(DescribeStatementResponse::new(param_types, fields)) } else { Ok(DescribeStatementResponse::new(param_types, vec![])) @@ -467,14 +430,11 @@ impl ExtendedQueryHandler for PostgresHandler { where C: ClientInfo + Unpin + Send + Sync, { - let (sqlite_sql, kind) = self - .translate_sql(&portal.statement.statement) - .map_err(|e| pg_error(&e))?; + let (sqlite_sql, kind) = + self.translate_sql(&portal.statement.statement).map_err(|e| pg_error(&e))?; if kind == StatementKind::Query && !sqlite_sql.is_empty() { - let fields = self - .probe_columns(&sqlite_sql, &portal.result_column_format) - .await?; + let fields = self.probe_columns(&sqlite_sql, &portal.result_column_format).await?; Ok(DescribePortalResponse::new(fields)) } else { Ok(DescribePortalResponse::new(vec![])) 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..47620ee 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. @@ -25,11 +25,7 @@ struct TdsSession { impl TdsSession { fn new() -> Self { - Self { - in_transaction: false, - next_tran_id: 1, - current_tran_id: 0, - } + Self { in_transaction: false, next_tran_id: 1, current_tran_id: 0 } } fn begin(&mut self) -> u64 { @@ -207,10 +203,7 @@ fn decode_utf16le(data: &[u8]) -> Option { if data.len() % 2 != 0 { return None; } - let chars: Vec = data - .chunks_exact(2) - .map(|c| u16::from_le_bytes([c[0], c[1]])) - .collect(); + let chars: Vec = data.chunks_exact(2).map(|c| u16::from_le_bytes([c[0], c[1]])).collect(); String::from_utf16(&chars).ok() } @@ -245,13 +238,8 @@ 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; - if total_len >= 4 && total_len <= payload.len() { - total_len - } else { - 0 - } + 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 { 0 } } /// Handle an RPC request. @@ -408,10 +396,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/packet.rs b/crates/litewire-tds/src/packet.rs index 8eac8db..f12c605 100644 --- a/crates/litewire-tds/src/packet.rs +++ b/crates/litewire-tds/src/packet.rs @@ -93,14 +93,10 @@ pub async fn read_message( } match msg_type { - Some(pt) => Ok(Some(TdsMessage { - packet_type: pt, - payload, - })), - None => Err(std::io::Error::new( - std::io::ErrorKind::InvalidData, - "unknown TDS packet type", - )), + Some(pt) => Ok(Some(TdsMessage { packet_type: pt, payload })), + None => { + Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "unknown TDS packet type")) + } } } diff --git a/crates/litewire-tds/src/token.rs b/crates/litewire-tds/src/token.rs index f3b5b73..4c2b717 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 ─────────────────────────────────────────── @@ -100,10 +99,7 @@ pub fn build_columns(columns: &[Column], first_row: Option<&[Value]>) -> Vec) -> Vec = server_name - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let name_utf16: Vec = server_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); // Token + length(u16) + interface(u8) + tds_version(u32) + name_len(u8) + name + version(u32) let body_len = 1 + 4 + 1 + name_utf16.len() + 4; @@ -133,10 +126,7 @@ pub fn write_loginack(buf: &mut BytesMut, server_name: &str) { /// Write an ENVCHANGE token for database change. pub fn write_envchange_database(buf: &mut BytesMut, db_name: &str) { - let name_utf16: Vec = db_name - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let name_utf16: Vec = db_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); let char_len = db_name.chars().count() as u8; // type(1) + new_len(1) + new_value + old_len(1) + old_value @@ -154,15 +144,9 @@ pub fn write_envchange_database(buf: &mut BytesMut, db_name: &str) { /// Write an ENVCHANGE token for packet size. pub fn write_envchange_packet_size(buf: &mut BytesMut, size: u32) { let new_str = size.to_string(); - let new_utf16: Vec = new_str - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let new_utf16: Vec = new_str.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); let old_str = "4096"; - let old_utf16: Vec = old_str - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let old_utf16: Vec = old_str.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); let body_len = 1 + 1 + new_utf16.len() + 1 + old_utf16.len(); @@ -185,7 +169,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, @@ -196,18 +181,9 @@ fn write_info_or_error( proc_name: &str, line: u32, ) { - let msg_utf16: Vec = message - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); - let srv_utf16: Vec = server_name - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); - let proc_utf16: Vec = proc_name - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let msg_utf16: Vec = message.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let srv_utf16: Vec = server_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let proc_utf16: Vec = proc_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); // number(4) + state(1) + class(1) + msg_len(2) + msg + srv_len(1) + srv + proc_len(1) + proc + line(4) let body_len = 4 + 1 + 1 + 2 + msg_utf16.len() + 1 + srv_utf16.len() + 1 + proc_utf16.len() + 4; @@ -257,11 +233,7 @@ pub fn write_colmetadata(buf: &mut BytesMut, columns: &[TdsColumn]) { } // Column name (B_VARCHAR: length in chars as u8, then UTF-16LE). - let name_utf16: Vec = col - .name - .encode_utf16() - .flat_map(|c| c.to_le_bytes()) - .collect(); + let name_utf16: Vec = col.name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); buf.put_u8(col.name.chars().count() as u8); buf.put_slice(&name_utf16); } @@ -354,10 +326,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..ccacc2c 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; @@ -25,11 +25,7 @@ fn rewrite_statement_exprs(stmt: &mut Statement) { rewrite_query_exprs(source); } } - Statement::Update { - assignments, - selection, - .. - } => { + Statement::Update { assignments, selection, .. } => { for assign in assignments { rewrite_expr(&mut assign.value); } @@ -104,30 +100,18 @@ fn rewrite_expr(expr: &mut Expr) { Expr::IsNull(inner) | Expr::IsNotNull(inner) => { rewrite_expr(inner); } - Expr::InList { - expr: inner, list, .. - } => { + Expr::InList { expr: inner, list, .. } => { rewrite_expr(inner); for e in list { rewrite_expr(e); } } - Expr::Between { - expr: inner, - low, - high, - .. - } => { + Expr::Between { expr: inner, low, high, .. } => { rewrite_expr(inner); rewrite_expr(low); rewrite_expr(high); } - Expr::Case { - operand, - conditions, - else_result, - .. - } => { + Expr::Case { operand, conditions, else_result, .. } => { if let Some(op) = operand { rewrite_expr(op); } @@ -143,11 +127,8 @@ fn rewrite_expr(expr: &mut Expr) { rewrite_query_exprs(q); } Expr::CompoundIdentifier(parts) => { - let joined = parts - .iter() - .map(|p| p.value.to_ascii_uppercase()) - .collect::>() - .join("."); + let joined = + parts.iter().map(|p| p.value.to_ascii_uppercase()).collect::>().join("."); match joined.as_str() { "@@IDENTITY" => { *expr = Expr::Function(Function { @@ -182,26 +163,18 @@ fn rewrite_expr(expr: &mut Expr) { /// Helper to create a `ValueWithSpan` from a `Value`. fn value_expr(val: Value) -> Expr { - Expr::Value(ValueWithSpan { - value: val, - span: sqlparser::tokenizer::Span::empty(), - }) + Expr::Value(ValueWithSpan { value: val, span: sqlparser::tokenizer::Span::empty() }) } /// Helper to build a function name `ObjectName`. fn func_name(name: &str) -> ObjectName { - ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier( - Ident::new(name), - )]) + ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier(Ident::new(name))]) } /// Helper to build function args list. fn func_args(args: Vec) -> FunctionArguments { FunctionArguments::List(FunctionArgumentList { - args: args - .into_iter() - .map(|e| FunctionArg::Unnamed(FunctionArgExpr::Expr(e))) - .collect(), + args: args.into_iter().map(|e| FunctionArg::Unnamed(FunctionArgExpr::Expr(e))).collect(), duplicate_treatment: None, clauses: vec![], }) @@ -230,6 +203,44 @@ 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 +288,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,46 +352,84 @@ 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}"); } #[test] fn boolean_in_case_expression() { - let results = translate( - "SELECT CASE WHEN x = TRUE THEN 'yes' ELSE 'no' END FROM t", - Dialect::MySQL, - ) - .unwrap(); + let results = + translate("SELECT CASE WHEN x = TRUE THEN 'yes' ELSE 'no' END FROM t", Dialect::MySQL) + .unwrap(); let sql = extract_sql(&results[0]); 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() { @@ -404,11 +453,8 @@ mod tests { #[test] fn multiple_dollar_placeholders() { - let results = translate( - "SELECT * FROM t WHERE a = $1 AND b = $2", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("SELECT * FROM t WHERE a = $1 AND b = $2", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains("?1"), "got: {sql}"); assert!(sql.contains("?2"), "got: {sql}"); diff --git a/crates/litewire-translate/src/lib.rs b/crates/litewire-translate/src/lib.rs index 0939ea2..9fde609 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)?; @@ -185,42 +350,27 @@ mod tests { #[test] fn classify_insert() { - assert_eq!( - classify("INSERT INTO users VALUES (1)"), - StatementKind::Mutation - ); + assert_eq!(classify("INSERT INTO users VALUES (1)"), StatementKind::Mutation); } #[test] fn classify_update() { - assert_eq!( - classify("UPDATE users SET name = 'x'"), - StatementKind::Mutation - ); + assert_eq!(classify("UPDATE users SET name = 'x'"), StatementKind::Mutation); } #[test] fn classify_delete() { - assert_eq!( - classify("DELETE FROM users WHERE id = 1"), - StatementKind::Mutation - ); + assert_eq!(classify("DELETE FROM users WHERE id = 1"), StatementKind::Mutation); } #[test] fn classify_replace() { - assert_eq!( - classify("REPLACE INTO users VALUES (1, 'x')"), - StatementKind::Mutation - ); + assert_eq!(classify("REPLACE INTO users VALUES (1, 'x')"), StatementKind::Mutation); } #[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] @@ -230,10 +380,7 @@ mod tests { #[test] fn classify_alter() { - assert_eq!( - classify("ALTER TABLE users ADD col TEXT"), - StatementKind::Ddl - ); + assert_eq!(classify("ALTER TABLE users ADD col TEXT"), StatementKind::Ddl); } #[test] @@ -276,11 +423,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 +435,157 @@ 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..d67b4eb 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", @@ -180,19 +176,15 @@ pub fn detect_metadata_query(sql: &str, _dialect: Dialect) -> Option / SHOW FIELDS FROM - if let Some(rest) = upper - .strip_prefix("SHOW COLUMNS FROM ") - .or_else(|| upper.strip_prefix("SHOW FIELDS FROM ")) + if let Some(rest) = + upper.strip_prefix("SHOW COLUMNS FROM ").or_else(|| upper.strip_prefix("SHOW FIELDS FROM ")) { let table = extract_table_name(rest, trimmed); return Some(MetadataQuery::ShowColumns { table }); } // DESCRIBE
/ DESC
- if let Some(rest) = upper - .strip_prefix("DESCRIBE ") - .or_else(|| upper.strip_prefix("DESC ")) - { + if let Some(rest) = upper.strip_prefix("DESCRIBE ").or_else(|| upper.strip_prefix("DESC ")) { let table = extract_table_name(rest, trimmed); return Some(MetadataQuery::ShowColumns { table }); } @@ -273,9 +265,7 @@ pub fn detect_metadata_query(sql: &str, _dialect: Dialect) -> Option Option Option Option Option { let pattern = format!("{column} = "); if let Some(pos) = upper_sql.find(&pattern) { @@ -338,7 +324,7 @@ fn extract_where_value_original(original_sql: &str, column: &str) -> Option assert_eq!(schema, "mydb"), + Some(MetadataQuery::InformationSchemaTables { schema_filter: Some(schema) }) => { + assert_eq!(schema, "mydb") + } other => panic!("expected InformationSchemaTables with filter, got: {other:?}"), } } @@ -693,33 +654,22 @@ mod tests { Dialect::MySQL, ); match q { - Some(MetadataQuery::InformationSchemaColumns { - table_filter: Some(table), - }) => assert_eq!(table, "users"), + Some(MetadataQuery::InformationSchemaColumns { table_filter: Some(table) }) => { + 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, - ); - assert!(matches!( - q, - Some(MetadataQuery::InformationSchemaColumns { - table_filter: None - }) - )); + let q = detect_metadata_query("SELECT * FROM INFORMATION_SCHEMA.COLUMNS", Dialect::MySQL); + assert!(matches!(q, 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))); } @@ -727,10 +677,7 @@ mod tests { #[test] fn information_schema_tables_sql() { - let sql = MetadataQuery::InformationSchemaTables { - schema_filter: None, - } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaTables { schema_filter: None }.to_sqlite_sql(); assert!(sql.contains("sqlite_master"), "got: {sql}"); assert!(sql.contains("TABLE_NAME"), "got: {sql}"); assert!(sql.contains("TABLE_TYPE"), "got: {sql}"); @@ -738,10 +685,8 @@ mod tests { #[test] fn information_schema_tables_with_main_filter() { - let sql = MetadataQuery::InformationSchemaTables { - schema_filter: Some("main".into()), - } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaTables { schema_filter: Some("main".into()) } + .to_sqlite_sql(); assert!(sql.contains("sqlite_master"), "got: {sql}"); // Should NOT contain "AND 0" (main is a valid schema). assert!(!sql.contains("AND 0"), "got: {sql}"); @@ -749,20 +694,17 @@ mod tests { #[test] fn information_schema_tables_with_unknown_schema() { - let sql = MetadataQuery::InformationSchemaTables { - schema_filter: Some("nonexistent".into()), - } - .to_sqlite_sql(); + let sql = + MetadataQuery::InformationSchemaTables { schema_filter: Some("nonexistent".into()) } + .to_sqlite_sql(); // Should return empty (AND 0). assert!(sql.contains("AND 0"), "got: {sql}"); } #[test] fn information_schema_columns_with_table_sql() { - let sql = MetadataQuery::InformationSchemaColumns { - table_filter: Some("users".into()), - } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaColumns { table_filter: Some("users".into()) } + .to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("users"), "got: {sql}"); assert!(sql.contains("COLUMN_NAME"), "got: {sql}"); @@ -770,10 +712,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}"); } @@ -789,48 +728,35 @@ mod tests { #[test] fn identity_system_variable() { - let sql = MetadataQuery::SystemVariables { - variables: vec!["identity".into()], - } - .to_sqlite_sql(); + let sql = + MetadataQuery::SystemVariables { variables: vec!["identity".into()] }.to_sqlite_sql(); assert!(sql.contains("last_insert_rowid()"), "got: {sql}"); } #[test] fn rowcount_system_variable() { - let sql = MetadataQuery::SystemVariables { - variables: vec!["rowcount".into()], - } - .to_sqlite_sql(); + let sql = + MetadataQuery::SystemVariables { variables: vec!["rowcount".into()] }.to_sqlite_sql(); assert!(sql.contains("changes()"), "got: {sql}"); } #[test] fn detect_select_at_identity() { let q = detect_metadata_query("SELECT @@IDENTITY", Dialect::TDS); - assert!( - matches!(q, Some(MetadataQuery::SystemVariables { .. })), - "got: {q:?}" - ); + assert!(matches!(q, Some(MetadataQuery::SystemVariables { .. })), "got: {q:?}"); } #[test] fn detect_select_at_rowcount() { let q = detect_metadata_query("SELECT @@ROWCOUNT", Dialect::TDS); - assert!( - matches!(q, Some(MetadataQuery::SystemVariables { .. })), - "got: {q:?}" - ); + assert!(matches!(q, Some(MetadataQuery::SystemVariables { .. })), "got: {q:?}"); } // ── pg_catalog detection ─────────────────────────────────────────────── #[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:?}"); } @@ -852,10 +778,7 @@ mod tests { #[test] fn pg_catalog_columns_sql() { - let sql = MetadataQuery::PgCatalogColumns { - table: "users".into(), - } - .to_sqlite_sql(); + let sql = MetadataQuery::PgCatalogColumns { table: "users".into() }.to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("users"), "got: {sql}"); } @@ -888,10 +811,7 @@ mod tests { #[test] fn sys_columns_sql() { - let sql = MetadataQuery::SysColumns { - table: "orders".into(), - } - .to_sqlite_sql(); + let sql = MetadataQuery::SysColumns { table: "orders".into() }.to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("orders"), "got: {sql}"); } @@ -907,9 +827,6 @@ mod tests { #[test] fn detect_sp_columns() { let q = detect_metadata_query("EXEC sp_columns 'users'", Dialect::TDS); - assert!( - matches!(q, Some(MetadataQuery::SysColumns { .. })), - "got: {q:?}" - ); + assert!(matches!(q, Some(MetadataQuery::SysColumns { .. })), "got: {q:?}"); } } diff --git a/crates/litewire-translate/src/mysql.rs b/crates/litewire-translate/src/mysql.rs index 6ff41db..9d49e4a 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; @@ -35,10 +35,7 @@ fn rewrite_insert_on_duplicate(insert: &mut sqlparser::ast::Insert) { if let Some(OnInsert::DuplicateKeyUpdate(assignments)) = insert.on.take() { insert.on = Some(OnInsert::OnConflict(OnConflict { conflict_target: None, - action: OnConflictAction::DoUpdate(DoUpdate { - assignments, - selection: None, - }), + action: OnConflictAction::DoUpdate(DoUpdate { assignments, selection: None }), })); } } @@ -50,10 +47,7 @@ fn rewrite_limit_clause(query: &mut sqlparser::ast::Query) { if let Some(LimitClause::OffsetCommaLimit { offset, limit }) = query.limit_clause.take() { query.limit_clause = Some(LimitClause::LimitOffset { limit: Some(limit), - offset: Some(Offset { - value: offset, - rows: OffsetRows::None, - }), + offset: Some(Offset { value: offset, rows: OffsetRows::None }), limit_by: vec![], }); } @@ -72,7 +66,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 +131,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 +162,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)); } @@ -193,10 +186,7 @@ mod tests { fn boolean_translated() { let results = translate("SELECT TRUE, FALSE", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); - assert!( - sql.contains('1') && sql.contains('0'), - "expected 1 and 0, got: {sql}" - ); + assert!(sql.contains('1') && sql.contains('0'), "expected 1 and 0, got: {sql}"); } // ── ON DUPLICATE KEY UPDATE ───────────────────────────────────────────── @@ -210,18 +200,9 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!( - upper.contains("ON CONFLICT"), - "expected ON CONFLICT, got: {sql}" - ); - assert!( - upper.contains("DO UPDATE"), - "expected DO UPDATE, got: {sql}" - ); - assert!( - !upper.contains("DUPLICATE KEY"), - "DUPLICATE KEY should be removed: {sql}" - ); + assert!(upper.contains("ON CONFLICT"), "expected ON CONFLICT, got: {sql}"); + assert!(upper.contains("DO UPDATE"), "expected DO UPDATE, got: {sql}"); + assert!(!upper.contains("DUPLICATE KEY"), "DUPLICATE KEY should be removed: {sql}"); } #[test] @@ -239,22 +220,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,14 +250,8 @@ 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("MEDIUMINT"), - "MEDIUMINT not rewritten: {sql}" - ); + assert!(!upper.contains("SMALLINT"), "SMALLINT not rewritten: {sql}"); + assert!(!upper.contains("MEDIUMINT"), "MEDIUMINT not rewritten: {sql}"); assert!(!upper.contains("BIGINT"), "BIGINT not rewritten: {sql}"); } @@ -296,11 +266,9 @@ mod tests { #[test] fn float_types_to_real() { - let results = translate( - "CREATE TABLE t (a FLOAT, b DOUBLE, c DECIMAL(10,2))", - Dialect::MySQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (a FLOAT, b DOUBLE, c DECIMAL(10,2))", Dialect::MySQL) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("REAL"), "no REAL found: {sql}"); @@ -316,21 +284,13 @@ mod tests { #[test] fn datetime_to_text() { - let results = translate( - "CREATE TABLE t (created DATETIME, updated TIMESTAMP)", - Dialect::MySQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (created DATETIME, updated TIMESTAMP)", Dialect::MySQL) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!( - !upper.contains("DATETIME"), - "DATETIME not rewritten: {sql}" - ); - assert!( - !upper.contains("TIMESTAMP"), - "TIMESTAMP not rewritten: {sql}" - ); + assert!(!upper.contains("DATETIME"), "DATETIME not rewritten: {sql}"); + assert!(!upper.contains("TIMESTAMP"), "TIMESTAMP not rewritten: {sql}"); } #[test] @@ -342,10 +302,7 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!( - !upper.contains("AUTO_INCREMENT"), - "AUTO_INCREMENT not removed: {sql}" - ); + assert!(!upper.contains("AUTO_INCREMENT"), "AUTO_INCREMENT not removed: {sql}"); } #[test] @@ -377,11 +334,9 @@ mod tests { #[test] fn insert_passthrough() { - let results = translate( - "INSERT INTO users (name, age) VALUES ('Alice', 30)", - Dialect::MySQL, - ) - .unwrap(); + let results = + translate("INSERT INTO users (name, age) VALUES ('Alice', 30)", Dialect::MySQL) + .unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains("Alice"), "got: {sql}"); } diff --git a/crates/litewire-translate/src/postgres.rs b/crates/litewire-translate/src/postgres.rs index 4f0514e..7cc616d 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 { @@ -67,11 +66,8 @@ mod tests { #[test] fn int_types_to_integer() { - let results = translate( - "CREATE TABLE t (a SMALLINT, b INT, c BIGINT)", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (a SMALLINT, b INT, c BIGINT)", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("SMALLINT"), "SMALLINT not rewritten: {sql}"); @@ -80,11 +76,8 @@ mod tests { #[test] fn float_to_real() { - let results = translate( - "CREATE TABLE t (a FLOAT(8), b NUMERIC(10,2))", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (a FLOAT(8), b NUMERIC(10,2))", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("REAL"), "no REAL found: {sql}"); @@ -93,11 +86,8 @@ mod tests { #[test] fn varchar_to_text() { - let results = translate( - "CREATE TABLE t (name VARCHAR(255), bio TEXT)", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (name VARCHAR(255), bio TEXT)", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("VARCHAR"), "VARCHAR not rewritten: {sql}"); @@ -132,11 +122,9 @@ mod tests { #[test] fn timestamp_to_text() { - let results = translate( - "CREATE TABLE t (created TIMESTAMP, updated DATE)", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (created TIMESTAMP, updated DATE)", Dialect::PostgreSQL) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("TIMESTAMP"), "TIMESTAMP not rewritten: {sql}"); @@ -160,11 +148,8 @@ mod tests { #[test] fn serial_to_integer() { - let results = translate( - "CREATE TABLE t (id SERIAL PRIMARY KEY)", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (id SERIAL PRIMARY KEY)", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("SERIAL"), "SERIAL not rewritten: {sql}"); @@ -173,11 +158,8 @@ mod tests { #[test] fn bigserial_to_integer() { - let results = translate( - "CREATE TABLE t (id BIGSERIAL PRIMARY KEY)", - Dialect::PostgreSQL, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (id BIGSERIAL PRIMARY KEY)", Dialect::PostgreSQL).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("INTEGER"), "no INTEGER found: {sql}"); diff --git a/crates/litewire-translate/src/tds.rs b/crates/litewire-translate/src/tds.rs index f64cc46..bc192b4 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 { @@ -103,11 +103,9 @@ mod tests { #[test] fn int_types_to_integer() { - let results = translate( - "CREATE TABLE t (a TINYINT, b SMALLINT, c INT, d BIGINT)", - Dialect::TDS, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (a TINYINT, b SMALLINT, c INT, d BIGINT)", Dialect::TDS) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("TINYINT"), "TINYINT not rewritten: {sql}"); @@ -117,11 +115,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}"); @@ -139,8 +133,7 @@ mod tests { #[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 +142,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,17 +199,10 @@ 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!( - !upper.contains("UNIQUEIDENTIFIER"), - "UNIQUEIDENTIFIER not rewritten: {sql}" - ); + assert!(!upper.contains("UNIQUEIDENTIFIER"), "UNIQUEIDENTIFIER not rewritten: {sql}"); assert!(upper.contains("TEXT"), "no TEXT found: {sql}"); } @@ -235,11 +217,8 @@ mod tests { #[test] fn identity_column_option_removed() { - let results = translate( - "CREATE TABLE t (id INT IDENTITY(1,1) PRIMARY KEY)", - Dialect::TDS, - ) - .unwrap(); + let results = + translate("CREATE TABLE t (id INT IDENTITY(1,1) PRIMARY KEY)", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("IDENTITY"), "IDENTITY not removed: {sql}"); diff --git a/crates/litewire/src/lib.rs b/crates/litewire/src/lib.rs index c6e5223..cd6234b 100644 --- a/crates/litewire/src/lib.rs +++ b/crates/litewire/src/lib.rs @@ -110,18 +110,14 @@ impl LiteWire { if let Some(addr) = self.mysql_listen { let config = litewire_mysql::MysqlFrontendConfig { listen: addr }; let frontend = litewire_mysql::MysqlFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { - frontend.serve().await.map_err(Into::into) - })); + handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); } #[cfg(feature = "hrana")] if let Some(addr) = self.hrana_listen { let config = litewire_hrana::HranaFrontendConfig { listen: addr }; let frontend = litewire_hrana::HranaFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { - frontend.serve().await.map_err(Into::into) - })); + handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); } #[cfg(feature = "postgres")] @@ -129,19 +125,14 @@ impl LiteWire { let config = litewire_postgres::PostgresFrontendConfig { listen: addr }; let frontend = litewire_postgres::PostgresFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { - frontend.serve().await.map_err(Into::into) - })); + handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); } #[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)); - handles.push(tokio::spawn(async move { - frontend.serve().await.map_err(Into::into) - })); + let frontend = litewire_tds::TdsFrontend::new(config, Arc::clone(&self.backend)); + handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); } if handles.is_empty() { 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..9963e99 100644 --- a/crates/litewire/tests/mysql_e2e.rs +++ b/crates/litewire/tests/mysql_e2e.rs @@ -50,10 +50,7 @@ async fn connect(port: u16) -> Conn { } fn init_tracing() { - let _ = tracing_subscriber::fmt() - .with_env_filter("debug") - .with_test_writer() - .try_init(); + let _ = tracing_subscriber::fmt().with_env_filter("debug").with_test_writer().try_init(); } #[tokio::test] @@ -80,17 +77,11 @@ async fn create_table_insert_select() { .await .unwrap(); - conn.query_drop("INSERT INTO users (id, name) VALUES (1, 'Alice')") - .await - .unwrap(); - conn.query_drop("INSERT INTO users (id, name) VALUES (2, 'Bob')") - .await - .unwrap(); + conn.query_drop("INSERT INTO users (id, name) VALUES (1, 'Alice')").await.unwrap(); + conn.query_drop("INSERT INTO users (id, name) VALUES (2, 'Bob')").await.unwrap(); - let rows: Vec<(i64, String)> = conn - .query("SELECT id, name FROM users ORDER BY id") - .await - .unwrap(); + let rows: Vec<(i64, String)> = + conn.query("SELECT id, name FROM users ORDER BY id").await.unwrap(); assert_eq!(rows, vec![(1, "Alice".into()), (2, "Bob".into())]); drop(conn); @@ -107,11 +98,7 @@ async fn now_function_translates() { let result: Vec<(String,)> = conn.query("SELECT NOW()").await.unwrap(); assert_eq!(result.len(), 1); // Should look like "2024-01-15 12:34:56". - assert!( - result[0].0.contains('-'), - "expected datetime string, got: {}", - result[0].0 - ); + assert!(result[0].0.contains('-'), "expected datetime string, got: {}", result[0].0); drop(conn); } @@ -136,12 +123,8 @@ async fn show_tables() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE alpha (id INTEGER PRIMARY KEY)") - .await - .unwrap(); - conn.query_drop("CREATE TABLE beta (id INTEGER PRIMARY KEY)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE alpha (id INTEGER PRIMARY KEY)").await.unwrap(); + conn.query_drop("CREATE TABLE beta (id INTEGER PRIMARY KEY)").await.unwrap(); let result: Vec<(String,)> = conn.query("SHOW TABLES").await.unwrap(); let names: Vec<&str> = result.iter().map(|r| r.0.as_str()).collect(); @@ -190,29 +173,16 @@ async fn update_and_delete() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)") - .await - .unwrap(); - conn.query_drop("INSERT INTO items VALUES (1, 10)") - .await - .unwrap(); - conn.query_drop("INSERT INTO items VALUES (2, 20)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)").await.unwrap(); + conn.query_drop("INSERT INTO items VALUES (1, 10)").await.unwrap(); + conn.query_drop("INSERT INTO items VALUES (2, 20)").await.unwrap(); - conn.query_drop("UPDATE items SET qty = 15 WHERE id = 1") - .await - .unwrap(); + conn.query_drop("UPDATE items SET qty = 15 WHERE id = 1").await.unwrap(); - let rows: Vec<(i64, i64)> = conn - .query("SELECT id, qty FROM items ORDER BY id") - .await - .unwrap(); + let rows: Vec<(i64, i64)> = conn.query("SELECT id, qty FROM items ORDER BY id").await.unwrap(); assert_eq!(rows, vec![(1, 15), (2, 20)]); - conn.query_drop("DELETE FROM items WHERE id = 2") - .await - .unwrap(); + conn.query_drop("DELETE FROM items WHERE id = 2").await.unwrap(); let rows: Vec<(i64, i64)> = conn.query("SELECT id, qty FROM items").await.unwrap(); assert_eq!(rows, vec![(1, 15)]); @@ -227,22 +197,13 @@ async fn multiple_connections() { let _server = start_litewire(port).await; let mut conn1 = connect(port).await; - conn1 - .query_drop("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); - conn1 - .query_drop("INSERT INTO shared VALUES (1, 'from_conn1')") - .await - .unwrap(); + conn1.query_drop("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + conn1.query_drop("INSERT INTO shared VALUES (1, 'from_conn1')").await.unwrap(); drop(conn1); // 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); @@ -257,27 +218,17 @@ async fn prepared_select_with_param() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)") - .await - .unwrap(); - conn.query_drop("INSERT INTO users VALUES (1, 'Alice')") - .await - .unwrap(); - conn.query_drop("INSERT INTO users VALUES (2, 'Bob')") - .await - .unwrap(); + conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)").await.unwrap(); + conn.query_drop("INSERT INTO users VALUES (1, 'Alice')").await.unwrap(); + conn.query_drop("INSERT INTO users VALUES (2, 'Bob')").await.unwrap(); // Prepared SELECT with parameter binding. - let rows: Vec<(i64, String)> = conn - .exec("SELECT id, name FROM users WHERE id = ?", (1_i64,)) - .await - .unwrap(); + let rows: Vec<(i64, String)> = + conn.exec("SELECT id, name FROM users WHERE id = ?", (1_i64,)).await.unwrap(); assert_eq!(rows, vec![(1, "Alice".into())]); - let rows: Vec<(i64, String)> = conn - .exec("SELECT id, name FROM users WHERE id = ?", (2_i64,)) - .await - .unwrap(); + let rows: Vec<(i64, String)> = + conn.exec("SELECT id, name FROM users WHERE id = ?", (2_i64,)).await.unwrap(); assert_eq!(rows, vec![(2, "Bob".into())]); drop(conn); @@ -295,31 +246,17 @@ async fn prepared_insert() { .unwrap(); // Prepared INSERT with parameters. - conn.exec_drop( - "INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", - (1_i64, "Widget", 10_i64), - ) - .await - .unwrap(); - - conn.exec_drop( - "INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", - (2_i64, "Gadget", 20_i64), - ) - .await - .unwrap(); - - let rows: Vec<(i64, String, i64)> = conn - .query("SELECT id, name, qty FROM items ORDER BY id") + conn.exec_drop("INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", (1_i64, "Widget", 10_i64)) .await .unwrap(); - assert_eq!( - rows, - vec![ - (1, "Widget".into(), 10), - (2, "Gadget".into(), 20), - ] - ); + + conn.exec_drop("INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", (2_i64, "Gadget", 20_i64)) + .await + .unwrap(); + + let rows: Vec<(i64, String, i64)> = + conn.query("SELECT id, name, qty FROM items ORDER BY id").await.unwrap(); + assert_eq!(rows, vec![(1, "Widget".into(), 10), (2, "Gadget".into(), 20),]); drop(conn); } @@ -331,23 +268,15 @@ async fn prepared_update_and_delete() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); - conn.query_drop("INSERT INTO t VALUES (1, 'old')") - .await - .unwrap(); + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + conn.query_drop("INSERT INTO t VALUES (1, 'old')").await.unwrap(); - conn.exec_drop("UPDATE t SET val = ? WHERE id = ?", ("new", 1_i64)) - .await - .unwrap(); + conn.exec_drop("UPDATE t SET val = ? WHERE id = ?", ("new", 1_i64)).await.unwrap(); let rows: Vec<(i64, String)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert_eq!(rows, vec![(1, "new".into())]); - conn.exec_drop("DELETE FROM t WHERE id = ?", (1_i64,)) - .await - .unwrap(); + conn.exec_drop("DELETE FROM t WHERE id = ?", (1_i64,)).await.unwrap(); let rows: Vec<(i64, String)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert!(rows.is_empty()); @@ -362,17 +291,12 @@ async fn prepared_with_null_param() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)") + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + + conn.exec_drop("INSERT INTO t (id, val) VALUES (?, ?)", (1_i64, Option::::None)) .await .unwrap(); - conn.exec_drop( - "INSERT INTO t (id, val) VALUES (?, ?)", - (1_i64, Option::::None), - ) - .await - .unwrap(); - let rows: Vec<(i64, Option)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert_eq!(rows, vec![(1, None)]); @@ -386,9 +310,7 @@ async fn prepared_reuse_same_statement() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY)").await.unwrap(); // Execute the same prepared statement multiple times. let stmt = conn.prep("INSERT INTO t (id) VALUES (?)").await.unwrap(); @@ -415,38 +337,24 @@ async fn on_duplicate_key_update() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE kv (k TEXT PRIMARY KEY, v INTEGER)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE kv (k TEXT PRIMARY KEY, v INTEGER)").await.unwrap(); - conn.query_drop("INSERT INTO kv (k, v) VALUES ('a', 1)") - .await - .unwrap(); + conn.query_drop("INSERT INTO kv (k, v) VALUES ('a', 1)").await.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(); - - let rows: Vec<(String, i64)> = conn - .query("SELECT k, v FROM kv WHERE k = 'a'") + 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'").await.unwrap(); 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); @@ -461,21 +369,14 @@ async fn transaction_commit() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("INSERT INTO txn_t VALUES (1, 'inside_txn')") - .await - .unwrap(); + conn.query_drop("INSERT INTO txn_t VALUES (1, 'inside_txn')").await.unwrap(); 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); @@ -488,25 +389,17 @@ async fn transaction_rollback() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); - conn.query_drop("INSERT INTO txn_rb VALUES (1, 'before')") - .await - .unwrap(); + conn.query_drop("INSERT INTO txn_rb VALUES (1, 'before')").await.unwrap(); conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("INSERT INTO txn_rb VALUES (2, 'rolled_back')") - .await - .unwrap(); + conn.query_drop("INSERT INTO txn_rb VALUES (2, 'rolled_back')").await.unwrap(); conn.query_drop("ROLLBACK").await.unwrap(); // Only the row inserted before the transaction should exist. - let rows: Vec<(i64, String)> = conn - .query("SELECT id, val FROM txn_rb ORDER BY id") - .await - .unwrap(); + let rows: Vec<(i64, String)> = + conn.query("SELECT id, val FROM txn_rb ORDER BY id").await.unwrap(); assert_eq!(rows, vec![(1, "before".into())]); drop(conn); @@ -519,25 +412,16 @@ async fn transaction_atomicity() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_atom (id INTEGER PRIMARY KEY, val INTEGER)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE txn_atom (id INTEGER PRIMARY KEY, val INTEGER)").await.unwrap(); - conn.query_drop("INSERT INTO txn_atom VALUES (1, 100)") - .await - .unwrap(); + conn.query_drop("INSERT INTO txn_atom VALUES (1, 100)").await.unwrap(); // Begin a transaction, update, then rollback — value should remain 100. conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("UPDATE txn_atom SET val = 200 WHERE id = 1") - .await - .unwrap(); + conn.query_drop("UPDATE txn_atom SET val = 200 WHERE id = 1").await.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); @@ -550,18 +434,12 @@ async fn information_schema_tables() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY)") - .await - .unwrap(); - conn.query_drop("CREATE TABLE posts (id INTEGER PRIMARY KEY)") - .await - .unwrap(); + conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY)").await.unwrap(); + conn.query_drop("CREATE TABLE posts (id INTEGER PRIMARY KEY)").await.unwrap(); // Returns TABLE_NAME, TABLE_TYPE, TABLE_SCHEMA columns. - let rows: Vec<(String, String, String)> = conn - .query("SELECT TABLE_NAME FROM information_schema.tables") - .await - .unwrap(); + let rows: Vec<(String, String, String)> = + conn.query("SELECT TABLE_NAME FROM information_schema.tables").await.unwrap(); let names: Vec<&str> = rows.iter().map(|r| r.0.as_str()).collect(); assert!(names.contains(&"users"), "got: {names:?}"); assert!(names.contains(&"posts"), "got: {names:?}"); @@ -587,10 +465,8 @@ async fn information_schema_columns() { .unwrap(); 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(); + 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:?}"); diff --git a/crates/litewire/tests/postgres_e2e.rs b/crates/litewire/tests/postgres_e2e.rs index bb93975..f9ff829 100644 --- a/crates/litewire/tests/postgres_e2e.rs +++ b/crates/litewire/tests/postgres_e2e.rs @@ -52,10 +52,7 @@ async fn connect(port: u16) -> Client { } fn init_tracing() { - let _ = tracing_subscriber::fmt() - .with_env_filter("debug") - .with_test_writer() - .try_init(); + let _ = tracing_subscriber::fmt().with_env_filter("debug").with_test_writer().try_init(); } // ── Simple query tests ───────────────────────────────────────────────────── @@ -81,26 +78,14 @@ async fn create_table_insert_select() { let client = connect(port).await; client - .execute( - "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)", - &[], - ) + .execute("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)", &[]) .await .unwrap(); - client - .execute("INSERT INTO users (id, name) VALUES (1, 'Alice')", &[]) - .await - .unwrap(); - client - .execute("INSERT INTO users (id, name) VALUES (2, 'Bob')", &[]) - .await - .unwrap(); + client.execute("INSERT INTO users (id, name) VALUES (1, 'Alice')", &[]).await.unwrap(); + client.execute("INSERT INTO users (id, name) VALUES (2, 'Bob')", &[]).await.unwrap(); - let rows = client - .query("SELECT id, name FROM users ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, name FROM users ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 2); let id: i64 = rows[0].get(0); @@ -121,44 +106,20 @@ async fn update_and_delete() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute( - "CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)", - &[], - ) - .await - .unwrap(); - client - .execute("INSERT INTO items VALUES (1, 10)", &[]) - .await - .unwrap(); - client - .execute("INSERT INTO items VALUES (2, 20)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)", &[]).await.unwrap(); + client.execute("INSERT INTO items VALUES (1, 10)", &[]).await.unwrap(); + client.execute("INSERT INTO items VALUES (2, 20)", &[]).await.unwrap(); - client - .execute("UPDATE items SET qty = 15 WHERE id = 1", &[]) - .await - .unwrap(); + client.execute("UPDATE items SET qty = 15 WHERE id = 1", &[]).await.unwrap(); - let rows = client - .query("SELECT id, qty FROM items ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, qty FROM items ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 2); let qty: i64 = rows[0].get(1); assert_eq!(qty, 15); - client - .execute("DELETE FROM items WHERE id = 2", &[]) - .await - .unwrap(); + client.execute("DELETE FROM items WHERE id = 2", &[]).await.unwrap(); - let rows = client - .query("SELECT id, qty FROM items", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, qty FROM items", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let id: i64 = rows[0].get(0); assert_eq!(id, 1); @@ -191,10 +152,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 +163,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(); @@ -223,25 +178,13 @@ async fn multiple_connections() { let _server = start_litewire(port).await; let client1 = connect(port).await; - client1 - .execute( - "CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)", - &[], - ) - .await - .unwrap(); - client1 - .execute("INSERT INTO shared VALUES (1, 'from_conn1')", &[]) - .await - .unwrap(); + client1.execute("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + client1.execute("INSERT INTO shared VALUES (1, 'from_conn1')", &[]).await.unwrap(); drop(client1); // Second connection should see the data. let client2 = connect(port).await; - let rows = client2 - .query("SELECT id, val FROM shared", &[]) - .await - .unwrap(); + let rows = client2.query("SELECT id, val FROM shared", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "from_conn1"); @@ -254,10 +197,7 @@ async fn empty_table_query() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute("CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); let rows = client.query("SELECT * FROM empty_t", &[]).await.unwrap(); assert!(rows.is_empty()); @@ -270,10 +210,7 @@ async fn drop_table() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute("CREATE TABLE temp_t (id INTEGER PRIMARY KEY)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE temp_t (id INTEGER PRIMARY KEY)", &[]).await.unwrap(); client.execute("DROP TABLE temp_t", &[]).await.unwrap(); // Selecting from dropped table should fail. @@ -291,27 +228,12 @@ async fn pg_serial_to_integer() { let client = connect(port).await; // SERIAL should be translated to INTEGER. - client - .execute( - "CREATE TABLE auto_t (id SERIAL PRIMARY KEY, name TEXT)", - &[], - ) - .await - .unwrap(); + client.execute("CREATE TABLE auto_t (id SERIAL PRIMARY KEY, name TEXT)", &[]).await.unwrap(); - client - .execute("INSERT INTO auto_t (name) VALUES ('first')", &[]) - .await - .unwrap(); - client - .execute("INSERT INTO auto_t (name) VALUES ('second')", &[]) - .await - .unwrap(); + client.execute("INSERT INTO auto_t (name) VALUES ('first')", &[]).await.unwrap(); + client.execute("INSERT INTO auto_t (name) VALUES ('second')", &[]).await.unwrap(); - let rows = client - .query("SELECT id, name FROM auto_t ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, name FROM auto_t ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 2); let name: &str = rows[0].get(1); assert_eq!(name, "first"); @@ -325,25 +247,16 @@ async fn pg_varchar_to_text() { let client = connect(port).await; client - .execute( - "CREATE TABLE typed_t (name VARCHAR(255), bio TEXT, data BYTEA)", - &[], - ) + .execute("CREATE TABLE typed_t (name VARCHAR(255), bio TEXT, data BYTEA)", &[]) .await .unwrap(); client - .execute( - "INSERT INTO typed_t (name, bio, data) VALUES ('test', 'a bio', 'raw')", - &[], - ) + .execute("INSERT INTO typed_t (name, bio, data) VALUES ('test', 'a bio', 'raw')", &[]) .await .unwrap(); - let rows = client - .query("SELECT name, bio FROM typed_t", &[]) - .await - .unwrap(); + let rows = client.query("SELECT name, bio FROM typed_t", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let name: &str = rows[0].get(0); assert_eq!(name, "test"); @@ -358,26 +271,14 @@ async fn pg_boolean_column() { // BOOLEAN should be translated to INTEGER in SQLite. client - .execute( - "CREATE TABLE flags (id INTEGER PRIMARY KEY, active BOOLEAN)", - &[], - ) + .execute("CREATE TABLE flags (id INTEGER PRIMARY KEY, active BOOLEAN)", &[]) .await .unwrap(); - client - .execute("INSERT INTO flags VALUES (1, TRUE)", &[]) - .await - .unwrap(); - client - .execute("INSERT INTO flags VALUES (2, FALSE)", &[]) - .await - .unwrap(); + client.execute("INSERT INTO flags VALUES (1, TRUE)", &[]).await.unwrap(); + client.execute("INSERT INTO flags VALUES (2, FALSE)", &[]).await.unwrap(); - let rows = client - .query("SELECT id, active FROM flags ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, active FROM flags ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 2); let active: i64 = rows[0].get(1); assert_eq!(active, 1); @@ -394,26 +295,14 @@ async fn large_result_set() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute("CREATE TABLE big_t (id INTEGER PRIMARY KEY, val TEXT)", &[]) - .await - .unwrap(); + client.execute("CREATE TABLE big_t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); // Insert 100 rows. for i in 0..100 { - client - .execute( - &format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), - &[], - ) - .await - .unwrap(); + client.execute(&format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), &[]).await.unwrap(); } - let rows = client - .query("SELECT id, val FROM big_t ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, val FROM big_t ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 100); let first_id: i64 = rows[0].get(0); @@ -429,27 +318,12 @@ async fn null_handling() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute( - "CREATE TABLE nullable (id INTEGER PRIMARY KEY, val TEXT)", - &[], - ) - .await - .unwrap(); + client.execute("CREATE TABLE nullable (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); - client - .execute("INSERT INTO nullable VALUES (1, NULL)", &[]) - .await - .unwrap(); - client - .execute("INSERT INTO nullable VALUES (2, 'present')", &[]) - .await - .unwrap(); + client.execute("INSERT INTO nullable VALUES (1, NULL)", &[]).await.unwrap(); + client.execute("INSERT INTO nullable VALUES (2, 'present')", &[]).await.unwrap(); - let rows = client - .query("SELECT id, val FROM nullable ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, val FROM nullable ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 2); let val: Option<&str> = rows[0].get(1); @@ -469,23 +343,14 @@ async fn transaction_commit() { let client = connect(port).await; // Use batch_execute (simple query protocol) for all transaction-related ops. - client - .batch_execute("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); + client.batch_execute("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); client.batch_execute("BEGIN").await.unwrap(); - client - .batch_execute("INSERT INTO txn_t VALUES (1, 'inside_txn')") - .await - .unwrap(); + client.batch_execute("INSERT INTO txn_t VALUES (1, 'inside_txn')").await.unwrap(); client.batch_execute("COMMIT").await.unwrap(); // Data should be visible after commit. - let rows = client - .query("SELECT id, val FROM txn_t", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, val FROM txn_t", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "inside_txn"); @@ -498,28 +363,16 @@ async fn transaction_rollback() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .batch_execute("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)") - .await - .unwrap(); + client.batch_execute("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); - client - .batch_execute("INSERT INTO txn_rb VALUES (1, 'before')") - .await - .unwrap(); + client.batch_execute("INSERT INTO txn_rb VALUES (1, 'before')").await.unwrap(); client.batch_execute("BEGIN").await.unwrap(); - client - .batch_execute("INSERT INTO txn_rb VALUES (2, 'rolled_back')") - .await - .unwrap(); + client.batch_execute("INSERT INTO txn_rb VALUES (2, 'rolled_back')").await.unwrap(); client.batch_execute("ROLLBACK").await.unwrap(); // Only the row inserted before the transaction should exist. - let rows = client - .query("SELECT id, val FROM txn_rb ORDER BY id", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, val FROM txn_rb ORDER BY id", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "before"); @@ -537,23 +390,14 @@ async fn transaction_atomicity() { .await .unwrap(); - client - .batch_execute("INSERT INTO txn_atom VALUES (1, 100)") - .await - .unwrap(); + client.batch_execute("INSERT INTO txn_atom VALUES (1, 100)").await.unwrap(); // Begin a transaction, update, then rollback — value should remain 100. client.batch_execute("BEGIN").await.unwrap(); - client - .batch_execute("UPDATE txn_atom SET val = 200 WHERE id = 1") - .await - .unwrap(); + client.batch_execute("UPDATE txn_atom SET val = 200 WHERE id = 1").await.unwrap(); client.batch_execute("ROLLBACK").await.unwrap(); - let rows = client - .query("SELECT id, val FROM txn_atom", &[]) - .await - .unwrap(); + let rows = client.query("SELECT id, val FROM txn_atom", &[]).await.unwrap(); assert_eq!(rows.len(), 1); let val: i64 = rows[0].get(1); assert_eq!(val, 100); @@ -566,23 +410,11 @@ async fn float_values() { let _server = start_litewire(port).await; let client = connect(port).await; - client - .execute( - "CREATE TABLE floats (id INTEGER PRIMARY KEY, val REAL)", - &[], - ) - .await - .unwrap(); + client.execute("CREATE TABLE floats (id INTEGER PRIMARY KEY, val REAL)", &[]).await.unwrap(); - client - .execute("INSERT INTO floats VALUES (1, 3.14)", &[]) - .await - .unwrap(); + client.execute("INSERT INTO floats VALUES (1, 3.14)", &[]).await.unwrap(); - let rows = client - .query("SELECT val FROM floats WHERE id = 1", &[]) - .await - .unwrap(); + let rows = client.query("SELECT val FROM floats WHERE id = 1", &[]).await.unwrap(); let val: f64 = rows[0].get(0); assert!((val - 3.14).abs() < 0.001); } diff --git a/crates/litewire/tests/tds_e2e.rs b/crates/litewire/tests/tds_e2e.rs index 7bd859c..61469b8 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(); @@ -68,10 +67,7 @@ async fn connect(port: u16) -> Client Date: Thu, 9 Jul 2026 19:37:51 -0700 Subject: [PATCH 2/2] style: add repo rustfmt.toml, reformat to stable-default rustfmt CI formats this repo standalone with stable rustfmt defaults; local checkouts nested under other projects were inheriting a parent rustfmt.toml and disagreeing with CI. Pin the config in-repo so both always match. --- crates/litewire-backend/src/hrana_client.rs | 85 ++++-- crates/litewire-backend/src/lib.rs | 30 +- .../litewire-backend/src/rusqlite_backend.rs | 115 ++++++-- .../tests/hrana_client_integration.rs | 107 ++++++-- crates/litewire-hrana/src/http.rs | 78 +++++- crates/litewire-hrana/src/types.rs | 29 +- crates/litewire-mysql/src/handler.rs | 38 ++- crates/litewire-mysql/src/types.rs | 70 ++++- crates/litewire-postgres/src/error_map.rs | 35 ++- crates/litewire-postgres/src/handler.rs | 76 ++++-- crates/litewire-tds/src/handler.rs | 23 +- crates/litewire-tds/src/packet.rs | 12 +- crates/litewire-tds/src/token.rs | 46 +++- crates/litewire-translate/src/common.rs | 90 ++++-- crates/litewire-translate/src/lib.rs | 60 +++- crates/litewire-translate/src/metadata.rs | 148 +++++++--- crates/litewire-translate/src/mysql.rs | 76 ++++-- crates/litewire-translate/src/postgres.rs | 48 +++- crates/litewire-translate/src/tds.rs | 25 +- crates/litewire/src/lib.rs | 16 +- crates/litewire/tests/mysql_e2e.rs | 245 ++++++++++++----- crates/litewire/tests/postgres_e2e.rs | 258 ++++++++++++++---- crates/litewire/tests/tds_e2e.rs | 117 ++++++-- rustfmt.toml | 3 + 24 files changed, 1425 insertions(+), 405 deletions(-) create mode 100644 rustfmt.toml diff --git a/crates/litewire-backend/src/hrana_client.rs b/crates/litewire-backend/src/hrana_client.rs index 8e7a91f..922cb80 100644 --- a/crates/litewire-backend/src/hrana_client.rs +++ b/crates/litewire-backend/src/hrana_client.rs @@ -143,7 +143,10 @@ impl HranaClient { if resp.status().is_success() { Ok(()) } else { - Err(BackendError::Other(format!("health check returned {}", resp.status()))) + Err(BackendError::Other(format!( + "health check returned {}", + resp.status() + ))) } } @@ -158,7 +161,10 @@ impl HranaClient { let request = PipelineRequest { baton: None, requests: vec![StreamRequest::Execute(ExecuteRequest { - stmt: StmtRequest { sql: sql.to_string(), args }, + stmt: StmtRequest { + sql: sql.to_string(), + args, + }, })], }; @@ -173,7 +179,9 @@ impl HranaClient { if !resp.status().is_success() { let status = resp.status(); let body = resp.text().await.unwrap_or_default(); - return Err(BackendError::Other(format!("sqld returned {status}: {body}"))); + return Err(BackendError::Other(format!( + "sqld returned {status}: {body}" + ))); } let pipeline: PipelineResponse = resp @@ -210,7 +218,10 @@ impl Backend for HranaClient { let exec = self.execute_pipeline(sql, params).await?; Ok(ExecuteResult { affected_rows: exec.affected_row_count, - last_insert_rowid: exec.last_insert_rowid.as_deref().and_then(|s| s.parse().ok()), + last_insert_rowid: exec + .last_insert_rowid + .as_deref() + .and_then(|s| s.parse().ok()), }) } } @@ -221,12 +232,16 @@ impl Backend for HranaClient { fn value_to_hrana(val: &Value) -> HranaValue { match val { Value::Null => HranaValue::Null, - Value::Integer(i) => HranaValue::Integer { value: i.to_string() }, + Value::Integer(i) => HranaValue::Integer { + value: i.to_string(), + }, Value::Float(f) => HranaValue::Float { value: *f }, Value::Text(s) => HranaValue::Text { value: s.clone() }, Value::Blob(b) => { use base64::Engine; - HranaValue::Blob { base64: base64::engine::general_purpose::STANDARD.encode(b) } + HranaValue::Blob { + base64: base64::engine::general_purpose::STANDARD.encode(b), + } } } } @@ -240,7 +255,9 @@ fn response_value_to_backend(rv: &ResponseValue) -> Value { ResponseValue::Text { value } => Value::Text(value.clone()), ResponseValue::Blob { base64: b64 } => { use base64::Engine; - let bytes = base64::engine::general_purpose::STANDARD.decode(b64).unwrap_or_default(); + let bytes = base64::engine::general_purpose::STANDARD + .decode(b64) + .unwrap_or_default(); Value::Blob(bytes) } } @@ -251,11 +268,17 @@ fn execute_response_to_result_set(exec: ExecuteResponse) -> ResultSet { let columns = exec .cols .iter() - .map(|c| Column { name: c.name.clone(), decltype: c.decltype.clone() }) + .map(|c| Column { + name: c.name.clone(), + decltype: c.decltype.clone(), + }) .collect(); - let rows = - exec.rows.iter().map(|row| row.iter().map(response_value_to_backend).collect()).collect(); + let rows = exec + .rows + .iter() + .map(|row| row.iter().map(response_value_to_backend).collect()) + .collect(); ResultSet { columns, rows } } @@ -301,7 +324,9 @@ mod tests { match value_to_hrana(&Value::Blob(data.clone())) { HranaValue::Blob { base64: b64 } => { use base64::Engine; - let decoded = base64::engine::general_purpose::STANDARD.decode(&b64).unwrap(); + let decoded = base64::engine::general_purpose::STANDARD + .decode(&b64) + .unwrap(); assert_eq!(decoded, data); } other => panic!("expected Blob, got: {other:?}"), @@ -312,9 +337,19 @@ mod tests { fn response_value_roundtrip() { let cases = vec![ (Value::Null, ResponseValue::Null), - (Value::Integer(123), ResponseValue::Integer { value: "123".into() }), + ( + Value::Integer(123), + ResponseValue::Integer { + value: "123".into(), + }, + ), (Value::Float(3.14), ResponseValue::Float { value: 3.14 }), - (Value::Text("test".into()), ResponseValue::Text { value: "test".into() }), + ( + Value::Text("test".into()), + ResponseValue::Text { + value: "test".into(), + }, + ), ]; for (val, rv) in cases { @@ -327,17 +362,27 @@ mod tests { fn execute_response_to_result_set_maps_correctly() { let exec = ExecuteResponse { cols: vec![ - ColResponse { name: "id".into(), decltype: Some("INTEGER".into()) }, - ColResponse { name: "name".into(), decltype: Some("TEXT".into()) }, + ColResponse { + name: "id".into(), + decltype: Some("INTEGER".into()), + }, + ColResponse { + name: "name".into(), + decltype: Some("TEXT".into()), + }, ], rows: vec![ vec![ ResponseValue::Integer { value: "1".into() }, - ResponseValue::Text { value: "alice".into() }, + ResponseValue::Text { + value: "alice".into(), + }, ], vec![ ResponseValue::Integer { value: "2".into() }, - ResponseValue::Text { value: "bob".into() }, + ResponseValue::Text { + value: "bob".into(), + }, ], ], affected_row_count: 0, @@ -368,8 +413,10 @@ mod tests { #[test] fn hrana_generic_error() { - let err = - hrana_error_to_backend(ErrorResponse { message: "something broke".into(), code: None }); + let err = hrana_error_to_backend(ErrorResponse { + message: "something broke".into(), + code: None, + }); match err { BackendError::Other(msg) => assert_eq!(msg, "something broke"), other => panic!("expected Other error, got: {other:?}"), diff --git a/crates/litewire-backend/src/lib.rs b/crates/litewire-backend/src/lib.rs index fb12cbe..c391fd3 100644 --- a/crates/litewire-backend/src/lib.rs +++ b/crates/litewire-backend/src/lib.rs @@ -120,7 +120,10 @@ mod tests { #[test] fn display_blob() { - assert_eq!(format!("{}", Value::Blob(vec![0xDE, 0xAD])), ""); + assert_eq!( + format!("{}", Value::Blob(vec![0xDE, 0xAD])), + "" + ); assert_eq!(format!("{}", Value::Blob(vec![])), ""); } @@ -153,14 +156,20 @@ mod tests { #[test] fn column_with_decltype() { - let c = Column { name: "id".into(), decltype: Some("INTEGER".into()) }; + let c = Column { + name: "id".into(), + decltype: Some("INTEGER".into()), + }; assert_eq!(c.name, "id"); assert_eq!(c.decltype.as_deref(), Some("INTEGER")); } #[test] fn column_without_decltype() { - let c = Column { name: "expr".into(), decltype: None }; + let c = Column { + name: "expr".into(), + decltype: None, + }; assert!(c.decltype.is_none()); } @@ -168,7 +177,10 @@ mod tests { #[test] fn execute_result_no_insert() { - let r = ExecuteResult { affected_rows: 3, last_insert_rowid: None }; + let r = ExecuteResult { + affected_rows: 3, + last_insert_rowid: None, + }; assert_eq!(r.affected_rows, 3); assert!(r.last_insert_rowid.is_none()); } @@ -179,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 1f8d878..ea9be27 100644 --- a/crates/litewire-backend/src/rusqlite_backend.rs +++ b/crates/litewire-backend/src/rusqlite_backend.rs @@ -31,7 +31,9 @@ impl Rusqlite { conn.execute_batch("PRAGMA journal_mode=WAL; PRAGMA busy_timeout=5000;") .map_err(|e| BackendError::Sqlite(e.to_string()))?; - Ok(Self { conn: Arc::new(Mutex::new(conn)) }) + Ok(Self { + conn: Arc::new(Mutex::new(conn)), + }) } /// Open an in-memory SQLite database. @@ -41,7 +43,9 @@ 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()))?; - Ok(Self { conn: Arc::new(Mutex::new(conn)) }) + Ok(Self { + conn: Arc::new(Mutex::new(conn)), + }) } } @@ -82,7 +86,9 @@ impl Backend for Rusqlite { task::spawn_blocking(move || { let conn = conn.lock(); - let mut stmt = conn.prepare(&sql).map_err(|e| BackendError::Sqlite(e.to_string()))?; + let mut stmt = conn + .prepare(&sql) + .map_err(|e| BackendError::Sqlite(e.to_string()))?; let col_count = stmt.column_count(); let columns: Vec = (0..col_count) @@ -101,7 +107,10 @@ 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( @@ -111,7 +120,10 @@ impl Backend for Rusqlite { result_rows.push(values); } - Ok(ResultSet { columns, rows: result_rows }) + Ok(ResultSet { + columns, + rows: result_rows, + }) }) .await .map_err(|e| BackendError::Other(format!("spawn_blocking join error: {e}")))? @@ -162,19 +174,28 @@ mod tests { .unwrap(); let result = backend - .execute("INSERT INTO users (name) VALUES (?1)", &[Value::Text("Alice".into())]) + .execute( + "INSERT INTO users (name) VALUES (?1)", + &[Value::Text("Alice".into())], + ) .await .unwrap(); assert_eq!(result.affected_rows, 1); assert_eq!(result.last_insert_rowid, Some(1)); let result = backend - .execute("INSERT INTO users (name) VALUES (?1)", &[Value::Text("Bob".into())]) + .execute( + "INSERT INTO users (name) VALUES (?1)", + &[Value::Text("Bob".into())], + ) .await .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"); @@ -182,8 +203,10 @@ mod tests { assert_eq!(rs.rows[0][1], Value::Text("Alice".into())); assert_eq!(rs.rows[1][1], Value::Text("Bob".into())); - let rs = - backend.query("SELECT * FROM users WHERE id = ?1", &[Value::Integer(1)]).await.unwrap(); + let rs = backend + .query("SELECT * FROM users WHERE id = ?1", &[Value::Integer(1)]) + .await + .unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][1], Value::Text("Alice".into())); } @@ -192,7 +215,10 @@ mod tests { async fn types_roundtrip() { let backend = Rusqlite::memory().unwrap(); backend - .execute("CREATE TABLE typed (i INTEGER, r REAL, t TEXT, b BLOB)", &[]) + .execute( + "CREATE TABLE typed (i INTEGER, r REAL, t TEXT, b BLOB)", + &[], + ) .await .unwrap(); @@ -219,8 +245,14 @@ mod tests { #[tokio::test] async fn null_handling() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (v TEXT)", &[]).await.unwrap(); - backend.execute("INSERT INTO t VALUES (?1)", &[Value::Null]).await.unwrap(); + backend + .execute("CREATE TABLE t (v TEXT)", &[]) + .await + .unwrap(); + backend + .execute("INSERT INTO t VALUES (?1)", &[Value::Null]) + .await + .unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.rows[0][0], Value::Null); @@ -229,7 +261,10 @@ mod tests { #[tokio::test] async fn empty_table_query() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) + .await + .unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.columns.len(), 2); @@ -239,11 +274,18 @@ mod tests { #[tokio::test] async fn multiple_params() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (a INTEGER, b TEXT, c REAL)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (a INTEGER, b TEXT, c REAL)", &[]) + .await + .unwrap(); backend .execute( "INSERT INTO t VALUES (?1, ?2, ?3)", - &[Value::Integer(1), Value::Text("hello".into()), Value::Float(9.99)], + &[ + Value::Integer(1), + Value::Text("hello".into()), + Value::Float(9.99), + ], ) .await .unwrap(); @@ -262,7 +304,10 @@ mod tests { #[tokio::test] async fn affected_rows_count() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]) + .await + .unwrap(); for i in 0..5 { backend @@ -274,8 +319,10 @@ mod tests { .unwrap(); } - let result = - backend.execute("DELETE FROM t WHERE id >= ?1", &[Value::Integer(3)]).await.unwrap(); + let result = backend + .execute("DELETE FROM t WHERE id >= ?1", &[Value::Integer(3)]) + .await + .unwrap(); assert_eq!(result.affected_rows, 2); } @@ -296,10 +343,16 @@ mod tests { #[tokio::test] async fn blob_roundtrip() { let backend = Rusqlite::memory().unwrap(); - backend.execute("CREATE TABLE t (data BLOB)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (data BLOB)", &[]) + .await + .unwrap(); let data = vec![0x00, 0xFF, 0xDE, 0xAD, 0xBE, 0xEF]; - backend.execute("INSERT INTO t VALUES (?1)", &[Value::Blob(data.clone())]).await.unwrap(); + backend + .execute("INSERT INTO t VALUES (?1)", &[Value::Blob(data.clone())]) + .await + .unwrap(); let rs = backend.query("SELECT * FROM t", &[]).await.unwrap(); assert_eq!(rs.rows[0][0], Value::Blob(data)); @@ -309,7 +362,10 @@ mod tests { async fn last_insert_rowid_increments() { let backend = Rusqlite::memory().unwrap(); backend - .execute("CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, v TEXT)", &[]) + .execute( + "CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, v TEXT)", + &[], + ) .await .unwrap(); @@ -335,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"); @@ -348,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 4120c7e..176762a 100644 --- a/crates/litewire-backend/tests/hrana_client_integration.rs +++ b/crates/litewire-backend/tests/hrana_client_integration.rs @@ -16,7 +16,9 @@ use litewire_hrana::HranaFrontendConfig; /// Start a Hrana server on a random port, return the client and server task handle. async fn start_server() -> (HranaClient, tokio::task::JoinHandle<()>) { // Bind to port 0 to get a random available port. - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("failed to bind"); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("failed to bind"); let addr: SocketAddr = listener.local_addr().unwrap(); let backend = Rusqlite::memory().expect("failed to create in-memory SQLite"); @@ -54,14 +56,20 @@ 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); // Insert let result = client - .execute("INSERT INTO users (name) VALUES (?)", &[Value::Text("alice".into())]) + .execute( + "INSERT INTO users (name) VALUES (?)", + &[Value::Text("alice".into())], + ) .await .expect("INSERT failed"); assert_eq!(result.affected_rows, 1); @@ -69,7 +77,10 @@ async fn create_table_and_insert() { // Insert another let result = client - .execute("INSERT INTO users (name) VALUES (?)", &[Value::Text("bob".into())]) + .execute( + "INSERT INTO users (name) VALUES (?)", + &[Value::Text("bob".into())], + ) .await .expect("INSERT failed"); assert_eq!(result.affected_rows, 1); @@ -81,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 @@ -121,7 +135,10 @@ async fn query_rows() { async fn query_with_params() { let (client, _server) = start_server().await; - client.execute("CREATE TABLE kv (key TEXT, val TEXT)", &[]).await.unwrap(); + client + .execute("CREATE TABLE kv (key TEXT, val TEXT)", &[]) + .await + .unwrap(); client .execute( "INSERT INTO kv (key, val) VALUES (?, ?)", @@ -137,8 +154,13 @@ async fn query_with_params() { .await .unwrap(); - let rs = - client.query("SELECT val FROM kv WHERE key = ?", &[Value::Text("b".into())]).await.unwrap(); + let rs = client + .query( + "SELECT val FROM kv WHERE key = ?", + &[Value::Text("b".into())], + ) + .await + .unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Text("2".into())); @@ -148,9 +170,15 @@ async fn query_with_params() { async fn query_empty_result() { let (client, _server) = start_server().await; - client.execute("CREATE TABLE empty_test (id INTEGER)", &[]).await.unwrap(); + client + .execute("CREATE TABLE empty_test (id INTEGER)", &[]) + .await + .unwrap(); - let rs = client.query("SELECT id FROM empty_test", &[]).await.unwrap(); + let rs = client + .query("SELECT id FROM empty_test", &[]) + .await + .unwrap(); assert_eq!(rs.columns.len(), 1); assert!(rs.rows.is_empty()); @@ -160,15 +188,27 @@ async fn query_empty_result() { async fn blob_roundtrip() { let (client, _server) = start_server().await; - client.execute("CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE blobs (id INTEGER PRIMARY KEY, data BLOB)", + &[], + ) + .await + .unwrap(); let blob_data = vec![0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0xFF]; client - .execute("INSERT INTO blobs (data) VALUES (?)", &[Value::Blob(blob_data.clone())]) + .execute( + "INSERT INTO blobs (data) VALUES (?)", + &[Value::Blob(blob_data.clone())], + ) .await .unwrap(); - let rs = client.query("SELECT data FROM blobs WHERE id = 1", &[]).await.unwrap(); + let rs = client + .query("SELECT data FROM blobs WHERE id = 1", &[]) + .await + .unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Blob(blob_data)); @@ -178,13 +218,22 @@ async fn blob_roundtrip() { async fn null_values() { let (client, _server) = start_server().await; - client.execute("CREATE TABLE nullable (id INTEGER, val TEXT)", &[]).await.unwrap(); client - .execute("INSERT INTO nullable (id, val) VALUES (?, ?)", &[Value::Integer(1), Value::Null]) + .execute("CREATE TABLE nullable (id INTEGER, val TEXT)", &[]) + .await + .unwrap(); + client + .execute( + "INSERT INTO nullable (id, val) VALUES (?, ?)", + &[Value::Integer(1), Value::Null], + ) .await .unwrap(); - let rs = client.query("SELECT val FROM nullable WHERE id = 1", &[]).await.unwrap(); + let rs = client + .query("SELECT val FROM nullable WHERE id = 1", &[]) + .await + .unwrap(); assert_eq!(rs.rows.len(), 1); assert_eq!(rs.rows[0][0], Value::Null); @@ -208,16 +257,30 @@ async fn sql_error_returns_backend_error() { async fn update_and_delete() { let (client, _server) = start_server().await; - client.execute("CREATE TABLE counters (name TEXT, val INTEGER)", &[]).await.unwrap(); - client.execute("INSERT INTO counters VALUES ('hits', 0)", &[]).await.unwrap(); + client + .execute("CREATE TABLE counters (name TEXT, val INTEGER)", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO counters VALUES ('hits', 0)", &[]) + .await + .unwrap(); - let result = - client.execute("UPDATE counters SET val = val + 1 WHERE name = 'hits'", &[]).await.unwrap(); + let result = client + .execute("UPDATE counters SET val = val + 1 WHERE name = 'hits'", &[]) + .await + .unwrap(); assert_eq!(result.affected_rows, 1); - let result = client.execute("DELETE FROM counters WHERE name = 'hits'", &[]).await.unwrap(); + let result = client + .execute("DELETE FROM counters WHERE name = 'hits'", &[]) + .await + .unwrap(); assert_eq!(result.affected_rows, 1); - let rs = client.query("SELECT COUNT(*) FROM counters", &[]).await.unwrap(); + let rs = client + .query("SELECT COUNT(*) FROM counters", &[]) + .await + .unwrap(); assert_eq!(rs.rows[0][0], Value::Integer(0)); } diff --git a/crates/litewire-hrana/src/http.rs b/crates/litewire-hrana/src/http.rs index dbc0172..b164170 100644 --- a/crates/litewire-hrana/src/http.rs +++ b/crates/litewire-hrana/src/http.rs @@ -44,14 +44,19 @@ async fn pipeline_handler( for stream_req in &req.requests { let result = match stream_req { StreamRequest::Execute(exec) => execute_stmt(&state.backend, &exec.stmt).await, - StreamRequest::Close => Ok(StreamResult::Ok { response: StreamResponse::Close }), + StreamRequest::Close => Ok(StreamResult::Ok { + response: StreamResponse::Close, + }), }; match result { Ok(r) => results.push(r), Err(e) => { results.push(StreamResult::Error { - error: ErrorResponse { message: e.to_string(), code: None }, + error: ErrorResponse { + message: e.to_string(), + code: None, + }, }); } } @@ -84,7 +89,10 @@ async fn execute_stmt( let cols: Vec = rs .columns .iter() - .map(|c| ColResponse { name: c.name.clone(), decltype: c.decltype.clone() }) + .map(|c| ColResponse { + name: c.name.clone(), + decltype: c.decltype.clone(), + }) .collect(); let rows: Vec> = rs @@ -146,7 +154,12 @@ mod tests { async fn health_check() { let app = build_router(test_backend()); let resp = app - .oneshot(Request::builder().uri("/health").body(Body::empty()).unwrap()) + .oneshot( + Request::builder() + .uri("/health") + .body(Body::empty()) + .unwrap(), + ) .await .unwrap(); @@ -159,7 +172,12 @@ mod tests { async fn version_endpoint() { let app = build_router(test_backend()); let resp = app - .oneshot(Request::builder().uri("/version").body(Body::empty()).unwrap()) + .oneshot( + Request::builder() + .uri("/version") + .body(Body::empty()) + .unwrap(), + ) .await .unwrap(); @@ -231,8 +249,14 @@ mod tests { let backend = test_backend(); // Pre-create table. - backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); - backend.execute("INSERT INTO t VALUES (1, 'Alice')", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) + .await + .unwrap(); + backend + .execute("INSERT INTO t VALUES (1, 'Alice')", &[]) + .await + .unwrap(); let app = build_router(backend); @@ -327,17 +351,32 @@ mod tests { 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") + resp["results"][0]["error"]["message"] + .as_str() + .unwrap() + .contains("nonexistent_table") ); } #[tokio::test] async fn pipeline_mutation_returns_affected_rows() { let backend = test_backend(); - backend.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); - backend.execute("INSERT INTO t VALUES (1, 'a')", &[]).await.unwrap(); - backend.execute("INSERT INTO t VALUES (2, 'b')", &[]).await.unwrap(); - backend.execute("INSERT INTO t VALUES (3, 'c')", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)", &[]) + .await + .unwrap(); + backend + .execute("INSERT INTO t VALUES (1, 'a')", &[]) + .await + .unwrap(); + backend + .execute("INSERT INTO t VALUES (2, 'b')", &[]) + .await + .unwrap(); + backend + .execute("INSERT INTO t VALUES (3, 'c')", &[]) + .await + .unwrap(); let app = build_router(backend); @@ -374,7 +413,10 @@ mod tests { async fn pipeline_insert_returns_last_insert_rowid() { let backend = test_backend(); backend - .execute("CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, val TEXT)", &[]) + .execute( + "CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, val TEXT)", + &[], + ) .await .unwrap(); @@ -440,7 +482,10 @@ mod tests { #[tokio::test] async fn pipeline_pragma_is_query() { let backend = test_backend(); - backend.execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (id INTEGER, name TEXT)", &[]) + .await + .unwrap(); let app = build_router(backend); @@ -479,7 +524,10 @@ mod tests { #[tokio::test] async fn pipeline_with_blob_param() { let backend = test_backend(); - backend.execute("CREATE TABLE t (data BLOB)", &[]).await.unwrap(); + backend + .execute("CREATE TABLE t (data BLOB)", &[]) + .await + .unwrap(); let app = build_router(backend); diff --git a/crates/litewire-hrana/src/types.rs b/crates/litewire-hrana/src/types.rs index 1db16eb..b887954 100644 --- a/crates/litewire-hrana/src/types.rs +++ b/crates/litewire-hrana/src/types.rs @@ -52,8 +52,9 @@ impl HranaValue { Self::Text { value } => Value::Text(value.clone()), Self::Blob { base64 } => { use base64::Engine; - let bytes = - base64::engine::general_purpose::STANDARD.decode(base64).unwrap_or_default(); + let bytes = base64::engine::general_purpose::STANDARD + .decode(base64) + .unwrap_or_default(); Value::Blob(bytes) } } @@ -112,12 +113,16 @@ impl ResponseValue { pub fn from_backend_value(val: &Value) -> Self { match val { Value::Null => Self::Null, - Value::Integer(i) => Self::Integer { value: i.to_string() }, + Value::Integer(i) => Self::Integer { + value: i.to_string(), + }, Value::Float(f) => Self::Float { value: *f }, Value::Text(s) => Self::Text { value: s.clone() }, Value::Blob(b) => { use base64::Engine; - Self::Blob { base64: base64::engine::general_purpose::STANDARD.encode(b) } + Self::Blob { + base64: base64::engine::general_purpose::STANDARD.encode(b), + } } } } @@ -150,7 +155,9 @@ mod tests { #[test] fn integer_invalid_to_backend() { - let v = HranaValue::Integer { value: "not_a_number".into() }; + let v = HranaValue::Integer { + value: "not_a_number".into(), + }; assert!(matches!(v.to_backend_value(), Value::Integer(0))); } @@ -162,7 +169,9 @@ mod tests { #[test] fn text_to_backend() { - let v = HranaValue::Text { value: "hello".into() }; + let v = HranaValue::Text { + value: "hello".into(), + }; assert!(matches!(v.to_backend_value(), Value::Text(s) if s == "hello")); } @@ -177,7 +186,9 @@ mod tests { #[test] fn blob_invalid_base64_to_backend() { - let v = HranaValue::Blob { base64: "!!!invalid!!!".into() }; + let v = HranaValue::Blob { + base64: "!!!invalid!!!".into(), + }; // Invalid base64 should return empty blob. assert!(matches!(v.to_backend_value(), Value::Blob(b) if b.is_empty())); } @@ -224,7 +235,9 @@ mod tests { match rv { ResponseValue::Blob { base64: encoded } => { use base64::Engine; - let decoded = base64::engine::general_purpose::STANDARD.decode(&encoded).unwrap(); + let decoded = base64::engine::general_purpose::STANDARD + .decode(&encoded) + .unwrap(); assert_eq!(decoded, data); } other => panic!("expected Blob, got: {other:?}"), diff --git a/crates/litewire-mysql/src/handler.rs b/crates/litewire-mysql/src/handler.rs index d4bc742..06ed08c 100644 --- a/crates/litewire-mysql/src/handler.rs +++ b/crates/litewire-mysql/src/handler.rs @@ -23,9 +23,17 @@ 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 { StatusFlags::SERVER_STATUS_IN_TRANS } else { StatusFlags::empty() }; - OkResponse { affected_rows, last_insert_id, status_flags, ..OkResponse::default() } + let status_flags = if in_transaction { + StatusFlags::SERVER_STATUS_IN_TRANS + } else { + StatusFlags::empty() + }; + OkResponse { + affected_rows, + last_insert_id, + status_flags, + ..OkResponse::default() + } } use crate::types::sqlite_to_mysql_column_type; @@ -51,7 +59,12 @@ pub struct LiteWireHandler { impl LiteWireHandler { pub fn new(backend: SharedBackend) -> Self { - Self { backend, stmts: HashMap::new(), next_stmt_id: 1, in_transaction: false } + Self { + backend, + stmts: HashMap::new(), + next_stmt_id: 1, + in_transaction: false, + } } /// Execute a query and write result set. @@ -107,8 +120,10 @@ impl LiteWireHandler { // 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 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 } @@ -262,7 +277,8 @@ impl AsyncMysqlShim for LiteWireHandler { let stmt_id = self.next_stmt_id; self.next_stmt_id += 1; - self.stmts.insert(stmt_id, PreparedStmt { sqlite_sql, kind }); + self.stmts + .insert(stmt_id, PreparedStmt { sqlite_sql, kind }); info.reply(stmt_id, ¶ms, &columns).await } @@ -292,7 +308,9 @@ impl AsyncMysqlShim for LiteWireHandler { // Noop statements (empty SQL from SET NAMES etc.) if sql.is_empty() { - return results.completed(ok_response(0, 0, self.in_transaction)).await; + return results + .completed(ok_response(0, 0, self.in_transaction)) + .await; } match kind { @@ -318,7 +336,9 @@ impl AsyncMysqlShim for LiteWireHandler { Ok(r) => r, Err(e) => { warn!("SQL translation error: {e}"); - return results.error(ErrorKind::ER_PARSE_ERROR, e.to_string().as_bytes()).await; + return results + .error(ErrorKind::ER_PARSE_ERROR, e.to_string().as_bytes()) + .await; } }; diff --git a/crates/litewire-mysql/src/types.rs b/crates/litewire-mysql/src/types.rs index 61ddd36..4451980 100644 --- a/crates/litewire-mysql/src/types.rs +++ b/crates/litewire-mysql/src/types.rs @@ -38,34 +38,55 @@ mod tests { #[test] fn type_mapping_integer() { - assert_eq!(sqlite_to_mysql_column_type(Some("INTEGER")), ColumnType::MYSQL_TYPE_LONGLONG); + assert_eq!( + sqlite_to_mysql_column_type(Some("INTEGER")), + ColumnType::MYSQL_TYPE_LONGLONG + ); } #[test] fn type_mapping_int_substring() { // "BIGINT", "TINYINT", etc. all contain "INT" - assert_eq!(sqlite_to_mysql_column_type(Some("BIGINT")), ColumnType::MYSQL_TYPE_LONGLONG); - assert_eq!(sqlite_to_mysql_column_type(Some("TINYINT")), ColumnType::MYSQL_TYPE_LONGLONG); + assert_eq!( + sqlite_to_mysql_column_type(Some("BIGINT")), + ColumnType::MYSQL_TYPE_LONGLONG + ); + assert_eq!( + sqlite_to_mysql_column_type(Some("TINYINT")), + ColumnType::MYSQL_TYPE_LONGLONG + ); } #[test] fn type_mapping_real() { - assert_eq!(sqlite_to_mysql_column_type(Some("REAL")), ColumnType::MYSQL_TYPE_DOUBLE); + assert_eq!( + sqlite_to_mysql_column_type(Some("REAL")), + ColumnType::MYSQL_TYPE_DOUBLE + ); } #[test] fn type_mapping_float() { - assert_eq!(sqlite_to_mysql_column_type(Some("FLOAT")), ColumnType::MYSQL_TYPE_DOUBLE); + assert_eq!( + sqlite_to_mysql_column_type(Some("FLOAT")), + ColumnType::MYSQL_TYPE_DOUBLE + ); } #[test] fn type_mapping_double() { - assert_eq!(sqlite_to_mysql_column_type(Some("DOUBLE")), ColumnType::MYSQL_TYPE_DOUBLE); + assert_eq!( + sqlite_to_mysql_column_type(Some("DOUBLE")), + ColumnType::MYSQL_TYPE_DOUBLE + ); } #[test] fn type_mapping_text() { - assert_eq!(sqlite_to_mysql_column_type(Some("TEXT")), ColumnType::MYSQL_TYPE_VAR_STRING); + assert_eq!( + sqlite_to_mysql_column_type(Some("TEXT")), + ColumnType::MYSQL_TYPE_VAR_STRING + ); } #[test] @@ -86,17 +107,26 @@ mod tests { #[test] fn type_mapping_blob() { - assert_eq!(sqlite_to_mysql_column_type(Some("BLOB")), ColumnType::MYSQL_TYPE_BLOB); + assert_eq!( + sqlite_to_mysql_column_type(Some("BLOB")), + ColumnType::MYSQL_TYPE_BLOB + ); } #[test] fn type_mapping_bytea() { - assert_eq!(sqlite_to_mysql_column_type(Some("BYTEA")), ColumnType::MYSQL_TYPE_BLOB); + assert_eq!( + sqlite_to_mysql_column_type(Some("BYTEA")), + ColumnType::MYSQL_TYPE_BLOB + ); } #[test] fn type_mapping_none_defaults_to_string() { - assert_eq!(sqlite_to_mysql_column_type(None), ColumnType::MYSQL_TYPE_VAR_STRING); + assert_eq!( + sqlite_to_mysql_column_type(None), + ColumnType::MYSQL_TYPE_VAR_STRING + ); } #[test] @@ -110,9 +140,21 @@ mod tests { #[test] fn type_mapping_case_insensitive() { // The function uppercases, so lowercase should work too. - assert_eq!(sqlite_to_mysql_column_type(Some("integer")), ColumnType::MYSQL_TYPE_LONGLONG); - assert_eq!(sqlite_to_mysql_column_type(Some("real")), ColumnType::MYSQL_TYPE_DOUBLE); - assert_eq!(sqlite_to_mysql_column_type(Some("text")), ColumnType::MYSQL_TYPE_VAR_STRING); - assert_eq!(sqlite_to_mysql_column_type(Some("blob")), ColumnType::MYSQL_TYPE_BLOB); + assert_eq!( + sqlite_to_mysql_column_type(Some("integer")), + ColumnType::MYSQL_TYPE_LONGLONG + ); + assert_eq!( + sqlite_to_mysql_column_type(Some("real")), + ColumnType::MYSQL_TYPE_DOUBLE + ); + assert_eq!( + sqlite_to_mysql_column_type(Some("text")), + ColumnType::MYSQL_TYPE_VAR_STRING + ); + assert_eq!( + sqlite_to_mysql_column_type(Some("blob")), + ColumnType::MYSQL_TYPE_BLOB + ); } } diff --git a/crates/litewire-postgres/src/error_map.rs b/crates/litewire-postgres/src/error_map.rs index 139639f..1709255 100644 --- a/crates/litewire-postgres/src/error_map.rs +++ b/crates/litewire-postgres/src/error_map.rs @@ -25,22 +25,34 @@ pub fn classify(err_msg: &str) -> PgError { // 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() }; + 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() }; + 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() }; + 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() }; + return PgError { + sqlstate: "23514", + message: err_msg.to_string(), + }; } // SQLITE_BUSY / SQLITE_LOCKED -> 55P03 lock_not_available @@ -49,7 +61,10 @@ pub fn classify(err_msg: &str) -> PgError { || lower.contains("sqlite_busy") || lower.contains("sqlite_locked") { - return PgError { sqlstate: "55P03", message: err_msg.to_string() }; + return PgError { + sqlstate: "55P03", + message: err_msg.to_string(), + }; } // SQLITE_READONLY -> 25006 read_only_sql_transaction @@ -57,11 +72,17 @@ pub fn classify(err_msg: &str) -> PgError { || lower.contains("readonly database") || lower.contains("sqlite_readonly") { - return PgError { sqlstate: "25006", message: err_msg.to_string() }; + return PgError { + sqlstate: "25006", + message: err_msg.to_string(), + }; } // Fallback: internal_error - PgError { sqlstate: "XX000", message: err_msg.to_string() } + PgError { + sqlstate: "XX000", + message: err_msg.to_string(), + } } #[cfg(test)] diff --git a/crates/litewire-postgres/src/handler.rs b/crates/litewire-postgres/src/handler.rs index 45162df..db9bc89 100644 --- a/crates/litewire-postgres/src/handler.rs +++ b/crates/litewire-postgres/src/handler.rs @@ -34,7 +34,10 @@ pub struct PostgresHandler { impl PostgresHandler { pub fn new(backend: SharedBackend) -> Self { - Self { backend, query_parser: Arc::new(NoopQueryParser::new()) } + Self { + backend, + query_parser: Arc::new(NoopQueryParser::new()), + } } /// Translate SQL from PostgreSQL dialect to SQLite and classify it. @@ -80,7 +83,11 @@ impl PostgresHandler { params: &[Value], format: &Format, ) -> PgWireResult> { - let rs = self.backend.query(sql, params).await.map_err(|e| pg_backend_error(&e))?; + let rs = self + .backend + .query(sql, params) + .await + .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 @@ -100,7 +107,13 @@ impl PostgresHandler { .map(value_to_pg_type) .unwrap_or(Type::TEXT) }; - FieldInfo::new(col.name.clone(), None, None, pg_type, format.format_for(idx)) + FieldInfo::new( + col.name.clone(), + None, + None, + pg_type, + format.format_for(idx), + ) }) .collect(); @@ -115,7 +128,10 @@ impl PostgresHandler { rows.push(encoder.finish()); } - Ok(Response::Query(QueryResponse::new(schema, stream::iter(rows)))) + Ok(Response::Query(QueryResponse::new( + schema, + stream::iter(rows), + ))) } /// Execute a mutation (INSERT/UPDATE/DELETE/DDL) and return an execution response. @@ -128,7 +144,10 @@ impl PostgresHandler { // Transaction commands need special Response variants, handle before // the generic execute path to avoid double-execution. if *kind == StatementKind::Transaction { - self.backend.execute(sql, params).await.map_err(|e| pg_backend_error(&e))?; + self.backend + .execute(sql, params) + .await + .map_err(|e| pg_backend_error(&e))?; let upper = sql.trim().to_ascii_uppercase(); if upper.starts_with("BEGIN") || upper.starts_with("START") { @@ -143,7 +162,11 @@ impl PostgresHandler { return Ok(Response::Execution(Tag::new("OK"))); } - let result = self.backend.execute(sql, params).await.map_err(|e| pg_backend_error(&e))?; + let result = self + .backend + .execute(sql, params) + .await + .map_err(|e| pg_backend_error(&e))?; let tag_name = match kind { StatementKind::Mutation => { @@ -202,7 +225,13 @@ impl PostgresHandler { .map(value_to_pg_type) .unwrap_or(Type::TEXT) }; - FieldInfo::new(col.name.clone(), None, None, pg_type, format.format_for(idx)) + FieldInfo::new( + col.name.clone(), + None, + None, + pg_type, + format.format_for(idx), + ) }) .collect()), Err(_) => Ok(vec![]), @@ -246,7 +275,12 @@ fn encode_value(encoder: &mut DataRowEncoder, val: &Value, field: &FieldInfo) -> fn extract_params(portal: &Portal) -> Vec { let mut values = Vec::with_capacity(portal.parameter_len()); for i in 0..portal.parameter_len() { - let param_type = portal.statement.parameter_types.get(i).cloned().unwrap_or(Type::TEXT); + let param_type = portal + .statement + .parameter_types + .get(i) + .cloned() + .unwrap_or(Type::TEXT); let val = match ¶m_type { t if *t == Type::BOOL => portal @@ -350,13 +384,15 @@ impl SimpleQueryHandler for PostgresHandler { TranslateResult::Noop => Response::Execution(Tag::new("SET")), TranslateResult::Metadata(meta) => { let sqlite_sql = meta.to_sqlite_sql(); - self.exec_query(&sqlite_sql, &[], &Format::UnifiedText).await? + self.exec_query(&sqlite_sql, &[], &Format::UnifiedText) + .await? } TranslateResult::Sql(sqlite_sql) => { if sqlite_sql.is_empty() { Response::Execution(Tag::new("OK")) } else { - self.exec_query(&sqlite_sql, &[], &Format::UnifiedText).await? + self.exec_query(&sqlite_sql, &[], &Format::UnifiedText) + .await? } } }; @@ -399,7 +435,8 @@ impl ExtendedQueryHandler for PostgresHandler { } let params = extract_params(portal); - self.exec_query(&sqlite_sql, ¶ms, &portal.result_column_format).await + self.exec_query(&sqlite_sql, ¶ms, &portal.result_column_format) + .await } async fn do_describe_statement( @@ -410,12 +447,16 @@ impl ExtendedQueryHandler for PostgresHandler { where C: ClientInfo + Unpin + Send + Sync, { - let (sqlite_sql, kind) = self.translate_sql(&stmt.statement).map_err(|e| pg_error(&e))?; + let (sqlite_sql, kind) = self + .translate_sql(&stmt.statement) + .map_err(|e| pg_error(&e))?; let param_types = stmt.parameter_types.clone(); if kind == StatementKind::Query && !sqlite_sql.is_empty() { - let fields = self.probe_columns(&sqlite_sql, &Format::UnifiedBinary).await?; + let fields = self + .probe_columns(&sqlite_sql, &Format::UnifiedBinary) + .await?; Ok(DescribeStatementResponse::new(param_types, fields)) } else { Ok(DescribeStatementResponse::new(param_types, vec![])) @@ -430,11 +471,14 @@ impl ExtendedQueryHandler for PostgresHandler { where C: ClientInfo + Unpin + Send + Sync, { - let (sqlite_sql, kind) = - self.translate_sql(&portal.statement.statement).map_err(|e| pg_error(&e))?; + let (sqlite_sql, kind) = self + .translate_sql(&portal.statement.statement) + .map_err(|e| pg_error(&e))?; if kind == StatementKind::Query && !sqlite_sql.is_empty() { - let fields = self.probe_columns(&sqlite_sql, &portal.result_column_format).await?; + let fields = self + .probe_columns(&sqlite_sql, &portal.result_column_format) + .await?; Ok(DescribePortalResponse::new(fields)) } else { Ok(DescribePortalResponse::new(vec![])) diff --git a/crates/litewire-tds/src/handler.rs b/crates/litewire-tds/src/handler.rs index 47620ee..d6aec33 100644 --- a/crates/litewire-tds/src/handler.rs +++ b/crates/litewire-tds/src/handler.rs @@ -25,7 +25,11 @@ struct TdsSession { impl TdsSession { fn new() -> Self { - Self { in_transaction: false, next_tran_id: 1, current_tran_id: 0 } + Self { + in_transaction: false, + next_tran_id: 1, + current_tran_id: 0, + } } fn begin(&mut self) -> u64 { @@ -163,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?; @@ -203,7 +211,10 @@ fn decode_utf16le(data: &[u8]) -> Option { if data.len() % 2 != 0 { return None; } - let chars: Vec = data.chunks_exact(2).map(|c| u16::from_le_bytes([c[0], c[1]])).collect(); + let chars: Vec = data + .chunks_exact(2) + .map(|c| u16::from_le_bytes([c[0], c[1]])) + .collect(); String::from_utf16(&chars).ok() } @@ -239,7 +250,11 @@ fn skip_all_headers(payload: &[u8]) -> usize { return 0; } 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 { 0 } + if total_len >= 4 && total_len <= payload.len() { + total_len + } else { + 0 + } } /// Handle an RPC request. diff --git a/crates/litewire-tds/src/packet.rs b/crates/litewire-tds/src/packet.rs index f12c605..8eac8db 100644 --- a/crates/litewire-tds/src/packet.rs +++ b/crates/litewire-tds/src/packet.rs @@ -93,10 +93,14 @@ pub async fn read_message( } match msg_type { - Some(pt) => Ok(Some(TdsMessage { packet_type: pt, payload })), - None => { - Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "unknown TDS packet type")) - } + Some(pt) => Ok(Some(TdsMessage { + packet_type: pt, + payload, + })), + None => Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "unknown TDS packet type", + )), } } diff --git a/crates/litewire-tds/src/token.rs b/crates/litewire-tds/src/token.rs index 4c2b717..9498dba 100644 --- a/crates/litewire-tds/src/token.rs +++ b/crates/litewire-tds/src/token.rs @@ -99,7 +99,10 @@ pub fn build_columns(columns: &[Column], first_row: Option<&[Value]>) -> Vec) -> Vec = server_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let name_utf16: Vec = server_name + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); // Token + length(u16) + interface(u8) + tds_version(u32) + name_len(u8) + name + version(u32) let body_len = 1 + 4 + 1 + name_utf16.len() + 4; @@ -126,7 +132,10 @@ pub fn write_loginack(buf: &mut BytesMut, server_name: &str) { /// Write an ENVCHANGE token for database change. pub fn write_envchange_database(buf: &mut BytesMut, db_name: &str) { - let name_utf16: Vec = db_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let name_utf16: Vec = db_name + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); let char_len = db_name.chars().count() as u8; // type(1) + new_len(1) + new_value + old_len(1) + old_value @@ -144,9 +153,15 @@ pub fn write_envchange_database(buf: &mut BytesMut, db_name: &str) { /// Write an ENVCHANGE token for packet size. pub fn write_envchange_packet_size(buf: &mut BytesMut, size: u32) { let new_str = size.to_string(); - let new_utf16: Vec = new_str.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let new_utf16: Vec = new_str + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); let old_str = "4096"; - let old_utf16: Vec = old_str.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let old_utf16: Vec = old_str + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); let body_len = 1 + 1 + new_utf16.len() + 1 + old_utf16.len(); @@ -181,9 +196,18 @@ fn write_info_or_error( proc_name: &str, line: u32, ) { - let msg_utf16: Vec = message.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); - let srv_utf16: Vec = server_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); - let proc_utf16: Vec = proc_name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let msg_utf16: Vec = message + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); + let srv_utf16: Vec = server_name + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); + let proc_utf16: Vec = proc_name + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); // number(4) + state(1) + class(1) + msg_len(2) + msg + srv_len(1) + srv + proc_len(1) + proc + line(4) let body_len = 4 + 1 + 1 + 2 + msg_utf16.len() + 1 + srv_utf16.len() + 1 + proc_utf16.len() + 4; @@ -233,7 +257,11 @@ pub fn write_colmetadata(buf: &mut BytesMut, columns: &[TdsColumn]) { } // Column name (B_VARCHAR: length in chars as u8, then UTF-16LE). - let name_utf16: Vec = col.name.encode_utf16().flat_map(|c| c.to_le_bytes()).collect(); + let name_utf16: Vec = col + .name + .encode_utf16() + .flat_map(|c| c.to_le_bytes()) + .collect(); buf.put_u8(col.name.chars().count() as u8); buf.put_slice(&name_utf16); } diff --git a/crates/litewire-translate/src/common.rs b/crates/litewire-translate/src/common.rs index ccacc2c..88d2c77 100644 --- a/crates/litewire-translate/src/common.rs +++ b/crates/litewire-translate/src/common.rs @@ -25,7 +25,11 @@ fn rewrite_statement_exprs(stmt: &mut Statement) { rewrite_query_exprs(source); } } - Statement::Update { assignments, selection, .. } => { + Statement::Update { + assignments, + selection, + .. + } => { for assign in assignments { rewrite_expr(&mut assign.value); } @@ -100,18 +104,30 @@ fn rewrite_expr(expr: &mut Expr) { Expr::IsNull(inner) | Expr::IsNotNull(inner) => { rewrite_expr(inner); } - Expr::InList { expr: inner, list, .. } => { + Expr::InList { + expr: inner, list, .. + } => { rewrite_expr(inner); for e in list { rewrite_expr(e); } } - Expr::Between { expr: inner, low, high, .. } => { + Expr::Between { + expr: inner, + low, + high, + .. + } => { rewrite_expr(inner); rewrite_expr(low); rewrite_expr(high); } - Expr::Case { operand, conditions, else_result, .. } => { + Expr::Case { + operand, + conditions, + else_result, + .. + } => { if let Some(op) = operand { rewrite_expr(op); } @@ -127,8 +143,11 @@ fn rewrite_expr(expr: &mut Expr) { rewrite_query_exprs(q); } Expr::CompoundIdentifier(parts) => { - let joined = - parts.iter().map(|p| p.value.to_ascii_uppercase()).collect::>().join("."); + let joined = parts + .iter() + .map(|p| p.value.to_ascii_uppercase()) + .collect::>() + .join("."); match joined.as_str() { "@@IDENTITY" => { *expr = Expr::Function(Function { @@ -163,18 +182,26 @@ fn rewrite_expr(expr: &mut Expr) { /// Helper to create a `ValueWithSpan` from a `Value`. fn value_expr(val: Value) -> Expr { - Expr::Value(ValueWithSpan { value: val, span: sqlparser::tokenizer::Span::empty() }) + Expr::Value(ValueWithSpan { + value: val, + span: sqlparser::tokenizer::Span::empty(), + }) } /// Helper to build a function name `ObjectName`. fn func_name(name: &str) -> ObjectName { - ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier(Ident::new(name))]) + ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier( + Ident::new(name), + )]) } /// Helper to build function args list. fn func_args(args: Vec) -> FunctionArguments { FunctionArguments::List(FunctionArgumentList { - args: args.into_iter().map(|e| FunctionArg::Unnamed(FunctionArgExpr::Expr(e))).collect(), + args: args + .into_iter() + .map(|e| FunctionArg::Unnamed(FunctionArgExpr::Expr(e))) + .collect(), duplicate_treatment: None, clauses: vec![], }) @@ -229,13 +256,15 @@ fn rewrite_function(func: &mut Function) { } "VERSION" => { func.name = func_name("coalesce"); - func.args = - func_args(vec![value_expr(Value::SingleQuotedString("8.0.0-litewire".into()))]); + 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()))]); + func.args = func_args(vec![value_expr(Value::SingleQuotedString( + "root@localhost".into(), + ))]); } "CONNECTION_ID" => { func.name = func_name("coalesce"); @@ -374,9 +403,11 @@ mod tests { #[test] fn boolean_in_case_expression() { - let results = - translate("SELECT CASE WHEN x = TRUE THEN 'yes' ELSE 'no' END FROM t", Dialect::MySQL) - .unwrap(); + let results = translate( + "SELECT CASE WHEN x = TRUE THEN 'yes' ELSE 'no' END FROM t", + Dialect::MySQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains('1'), "got: {sql}"); } @@ -387,8 +418,14 @@ mod tests { 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}"); + assert!( + sql.to_ascii_lowercase().contains("last_insert_rowid"), + "got: {sql}" + ); + assert!( + !sql.to_ascii_uppercase().contains("LAST_INSERT_ID("), + "got: {sql}" + ); } #[test] @@ -396,7 +433,10 @@ mod tests { 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}"); + assert!( + !sql.to_ascii_uppercase().contains("ROW_COUNT("), + "got: {sql}" + ); } #[test] @@ -426,7 +466,10 @@ mod tests { 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}"); + assert!( + !sql.to_ascii_uppercase().contains("CONNECTION_ID("), + "got: {sql}" + ); } // -- NEWID rewrite -------------------------------------------------------- @@ -453,8 +496,11 @@ mod tests { #[test] fn multiple_dollar_placeholders() { - let results = - translate("SELECT * FROM t WHERE a = $1 AND b = $2", Dialect::PostgreSQL).unwrap(); + let results = translate( + "SELECT * FROM t WHERE a = $1 AND b = $2", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains("?1"), "got: {sql}"); assert!(sql.contains("?2"), "got: {sql}"); diff --git a/crates/litewire-translate/src/lib.rs b/crates/litewire-translate/src/lib.rs index 9fde609..4e971bd 100644 --- a/crates/litewire-translate/src/lib.rs +++ b/crates/litewire-translate/src/lib.rs @@ -350,22 +350,34 @@ mod tests { #[test] fn classify_insert() { - assert_eq!(classify("INSERT INTO users VALUES (1)"), StatementKind::Mutation); + assert_eq!( + classify("INSERT INTO users VALUES (1)"), + StatementKind::Mutation + ); } #[test] fn classify_update() { - assert_eq!(classify("UPDATE users SET name = 'x'"), StatementKind::Mutation); + assert_eq!( + classify("UPDATE users SET name = 'x'"), + StatementKind::Mutation + ); } #[test] fn classify_delete() { - assert_eq!(classify("DELETE FROM users WHERE id = 1"), StatementKind::Mutation); + assert_eq!( + classify("DELETE FROM users WHERE id = 1"), + StatementKind::Mutation + ); } #[test] fn classify_replace() { - assert_eq!(classify("REPLACE INTO users VALUES (1, 'x')"), StatementKind::Mutation); + assert_eq!( + classify("REPLACE INTO users VALUES (1, 'x')"), + StatementKind::Mutation + ); } #[test] @@ -380,7 +392,10 @@ mod tests { #[test] fn classify_alter() { - assert_eq!(classify("ALTER TABLE users ADD col TEXT"), StatementKind::Ddl); + assert_eq!( + classify("ALTER TABLE users ADD col TEXT"), + StatementKind::Ddl + ); } #[test] @@ -464,7 +479,10 @@ mod tests { #[test] fn start_transaction_read_only_becomes_begin() { - assert_eq!(expect_sql("START TRANSACTION READ ONLY", Dialect::MySQL), "BEGIN"); + assert_eq!( + expect_sql("START TRANSACTION READ ONLY", Dialect::MySQL), + "BEGIN" + ); } #[test] @@ -482,7 +500,10 @@ mod tests { #[test] fn tsql_named_begin_transaction_strips_name() { - assert_eq!(expect_sql("BEGIN TRANSACTION my_txn", Dialect::TDS), "BEGIN"); + assert_eq!( + expect_sql("BEGIN TRANSACTION my_txn", Dialect::TDS), + "BEGIN" + ); } #[test] @@ -493,13 +514,19 @@ mod tests { #[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"); + 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"); + assert_eq!( + expect_sql("ROLLBACK TRANSACTION my_txn", Dialect::TDS), + "ROLLBACK" + ); } #[test] @@ -571,9 +598,18 @@ mod tests { #[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); + 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] diff --git a/crates/litewire-translate/src/metadata.rs b/crates/litewire-translate/src/metadata.rs index d67b4eb..1a6ba22 100644 --- a/crates/litewire-translate/src/metadata.rs +++ b/crates/litewire-translate/src/metadata.rs @@ -176,15 +176,19 @@ pub fn detect_metadata_query(sql: &str, _dialect: Dialect) -> Option / SHOW FIELDS FROM
- if let Some(rest) = - upper.strip_prefix("SHOW COLUMNS FROM ").or_else(|| upper.strip_prefix("SHOW FIELDS FROM ")) + if let Some(rest) = upper + .strip_prefix("SHOW COLUMNS FROM ") + .or_else(|| upper.strip_prefix("SHOW FIELDS FROM ")) { let table = extract_table_name(rest, trimmed); return Some(MetadataQuery::ShowColumns { table }); } // DESCRIBE
/ DESC
- if let Some(rest) = upper.strip_prefix("DESCRIBE ").or_else(|| upper.strip_prefix("DESC ")) { + if let Some(rest) = upper + .strip_prefix("DESCRIBE ") + .or_else(|| upper.strip_prefix("DESC ")) + { let table = extract_table_name(rest, trimmed); return Some(MetadataQuery::ShowColumns { table }); } @@ -265,7 +269,9 @@ pub fn detect_metadata_query(sql: &str, _dialect: Dialect) -> Option Option Option { + Some(MetadataQuery::InformationSchemaTables { + schema_filter: Some(schema), + }) => { assert_eq!(schema, "mydb") } other => panic!("expected InformationSchemaTables with filter, got: {other:?}"), @@ -654,7 +690,9 @@ mod tests { Dialect::MySQL, ); match q { - Some(MetadataQuery::InformationSchemaColumns { table_filter: Some(table) }) => { + Some(MetadataQuery::InformationSchemaColumns { + table_filter: Some(table), + }) => { assert_eq!(table, "users") } other => panic!("expected InformationSchemaColumns with filter, got: {other:?}"), @@ -664,7 +702,10 @@ mod tests { #[test] fn detect_information_schema_columns_no_filter() { let q = detect_metadata_query("SELECT * FROM INFORMATION_SCHEMA.COLUMNS", Dialect::MySQL); - assert!(matches!(q, Some(MetadataQuery::InformationSchemaColumns { table_filter: None }))); + assert!(matches!( + q, + Some(MetadataQuery::InformationSchemaColumns { table_filter: None }) + )); } #[test] @@ -677,7 +718,10 @@ mod tests { #[test] fn information_schema_tables_sql() { - let sql = MetadataQuery::InformationSchemaTables { schema_filter: None }.to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaTables { + schema_filter: None, + } + .to_sqlite_sql(); assert!(sql.contains("sqlite_master"), "got: {sql}"); assert!(sql.contains("TABLE_NAME"), "got: {sql}"); assert!(sql.contains("TABLE_TYPE"), "got: {sql}"); @@ -685,8 +729,10 @@ mod tests { #[test] fn information_schema_tables_with_main_filter() { - let sql = MetadataQuery::InformationSchemaTables { schema_filter: Some("main".into()) } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaTables { + schema_filter: Some("main".into()), + } + .to_sqlite_sql(); assert!(sql.contains("sqlite_master"), "got: {sql}"); // Should NOT contain "AND 0" (main is a valid schema). assert!(!sql.contains("AND 0"), "got: {sql}"); @@ -694,17 +740,20 @@ mod tests { #[test] fn information_schema_tables_with_unknown_schema() { - let sql = - MetadataQuery::InformationSchemaTables { schema_filter: Some("nonexistent".into()) } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaTables { + schema_filter: Some("nonexistent".into()), + } + .to_sqlite_sql(); // Should return empty (AND 0). assert!(sql.contains("AND 0"), "got: {sql}"); } #[test] fn information_schema_columns_with_table_sql() { - let sql = MetadataQuery::InformationSchemaColumns { table_filter: Some("users".into()) } - .to_sqlite_sql(); + let sql = MetadataQuery::InformationSchemaColumns { + table_filter: Some("users".into()), + } + .to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("users"), "got: {sql}"); assert!(sql.contains("COLUMN_NAME"), "got: {sql}"); @@ -728,28 +777,38 @@ mod tests { #[test] fn identity_system_variable() { - let sql = - MetadataQuery::SystemVariables { variables: vec!["identity".into()] }.to_sqlite_sql(); + let sql = MetadataQuery::SystemVariables { + variables: vec!["identity".into()], + } + .to_sqlite_sql(); assert!(sql.contains("last_insert_rowid()"), "got: {sql}"); } #[test] fn rowcount_system_variable() { - let sql = - MetadataQuery::SystemVariables { variables: vec!["rowcount".into()] }.to_sqlite_sql(); + let sql = MetadataQuery::SystemVariables { + variables: vec!["rowcount".into()], + } + .to_sqlite_sql(); assert!(sql.contains("changes()"), "got: {sql}"); } #[test] fn detect_select_at_identity() { let q = detect_metadata_query("SELECT @@IDENTITY", Dialect::TDS); - assert!(matches!(q, Some(MetadataQuery::SystemVariables { .. })), "got: {q:?}"); + assert!( + matches!(q, Some(MetadataQuery::SystemVariables { .. })), + "got: {q:?}" + ); } #[test] fn detect_select_at_rowcount() { let q = detect_metadata_query("SELECT @@ROWCOUNT", Dialect::TDS); - assert!(matches!(q, Some(MetadataQuery::SystemVariables { .. })), "got: {q:?}"); + assert!( + matches!(q, Some(MetadataQuery::SystemVariables { .. })), + "got: {q:?}" + ); } // ── pg_catalog detection ─────────────────────────────────────────────── @@ -757,7 +816,10 @@ mod tests { #[test] fn detect_pg_catalog_tables() { 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] @@ -766,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] @@ -778,7 +843,10 @@ mod tests { #[test] fn pg_catalog_columns_sql() { - let sql = MetadataQuery::PgCatalogColumns { table: "users".into() }.to_sqlite_sql(); + let sql = MetadataQuery::PgCatalogColumns { + table: "users".into(), + } + .to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("users"), "got: {sql}"); } @@ -811,7 +879,10 @@ mod tests { #[test] fn sys_columns_sql() { - let sql = MetadataQuery::SysColumns { table: "orders".into() }.to_sqlite_sql(); + let sql = MetadataQuery::SysColumns { + table: "orders".into(), + } + .to_sqlite_sql(); assert!(sql.contains("pragma_table_info"), "got: {sql}"); assert!(sql.contains("orders"), "got: {sql}"); } @@ -827,6 +898,9 @@ mod tests { #[test] fn detect_sp_columns() { let q = detect_metadata_query("EXEC sp_columns 'users'", Dialect::TDS); - assert!(matches!(q, Some(MetadataQuery::SysColumns { .. })), "got: {q:?}"); + assert!( + matches!(q, Some(MetadataQuery::SysColumns { .. })), + "got: {q:?}" + ); } } diff --git a/crates/litewire-translate/src/mysql.rs b/crates/litewire-translate/src/mysql.rs index 9d49e4a..bd9272b 100644 --- a/crates/litewire-translate/src/mysql.rs +++ b/crates/litewire-translate/src/mysql.rs @@ -35,7 +35,10 @@ fn rewrite_insert_on_duplicate(insert: &mut sqlparser::ast::Insert) { if let Some(OnInsert::DuplicateKeyUpdate(assignments)) = insert.on.take() { insert.on = Some(OnInsert::OnConflict(OnConflict { conflict_target: None, - action: OnConflictAction::DoUpdate(DoUpdate { assignments, selection: None }), + action: OnConflictAction::DoUpdate(DoUpdate { + assignments, + selection: None, + }), })); } } @@ -47,7 +50,10 @@ fn rewrite_limit_clause(query: &mut sqlparser::ast::Query) { if let Some(LimitClause::OffsetCommaLimit { offset, limit }) = query.limit_clause.take() { query.limit_clause = Some(LimitClause::LimitOffset { limit: Some(limit), - offset: Some(Offset { value: offset, rows: OffsetRows::None }), + offset: Some(Offset { + value: offset, + rows: OffsetRows::None, + }), limit_by: vec![], }); } @@ -186,7 +192,10 @@ mod tests { fn boolean_translated() { let results = translate("SELECT TRUE, FALSE", Dialect::MySQL).unwrap(); let sql = extract_sql(&results[0]); - assert!(sql.contains('1') && sql.contains('0'), "expected 1 and 0, got: {sql}"); + assert!( + sql.contains('1') && sql.contains('0'), + "expected 1 and 0, got: {sql}" + ); } // ── ON DUPLICATE KEY UPDATE ───────────────────────────────────────────── @@ -200,9 +209,18 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!(upper.contains("ON CONFLICT"), "expected ON CONFLICT, got: {sql}"); - assert!(upper.contains("DO UPDATE"), "expected DO UPDATE, got: {sql}"); - assert!(!upper.contains("DUPLICATE KEY"), "DUPLICATE KEY should be removed: {sql}"); + assert!( + upper.contains("ON CONFLICT"), + "expected ON CONFLICT, got: {sql}" + ); + assert!( + upper.contains("DO UPDATE"), + "expected DO UPDATE, got: {sql}" + ); + assert!( + !upper.contains("DUPLICATE KEY"), + "DUPLICATE KEY should be removed: {sql}" + ); } #[test] @@ -251,14 +269,20 @@ mod tests { 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("MEDIUMINT"), "MEDIUMINT not rewritten: {sql}"); + assert!( + !upper.contains("MEDIUMINT"), + "MEDIUMINT not rewritten: {sql}" + ); assert!(!upper.contains("BIGINT"), "BIGINT not rewritten: {sql}"); } #[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}"); @@ -266,9 +290,11 @@ mod tests { #[test] fn float_types_to_real() { - let results = - translate("CREATE TABLE t (a FLOAT, b DOUBLE, c DECIMAL(10,2))", Dialect::MySQL) - .unwrap(); + let results = translate( + "CREATE TABLE t (a FLOAT, b DOUBLE, c DECIMAL(10,2))", + Dialect::MySQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("REAL"), "no REAL found: {sql}"); @@ -284,13 +310,18 @@ mod tests { #[test] fn datetime_to_text() { - let results = - translate("CREATE TABLE t (created DATETIME, updated TIMESTAMP)", Dialect::MySQL) - .unwrap(); + let results = translate( + "CREATE TABLE t (created DATETIME, updated TIMESTAMP)", + Dialect::MySQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("DATETIME"), "DATETIME not rewritten: {sql}"); - assert!(!upper.contains("TIMESTAMP"), "TIMESTAMP not rewritten: {sql}"); + assert!( + !upper.contains("TIMESTAMP"), + "TIMESTAMP not rewritten: {sql}" + ); } #[test] @@ -302,7 +333,10 @@ mod tests { .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!(!upper.contains("AUTO_INCREMENT"), "AUTO_INCREMENT not removed: {sql}"); + assert!( + !upper.contains("AUTO_INCREMENT"), + "AUTO_INCREMENT not removed: {sql}" + ); } #[test] @@ -334,9 +368,11 @@ mod tests { #[test] fn insert_passthrough() { - let results = - translate("INSERT INTO users (name, age) VALUES ('Alice', 30)", Dialect::MySQL) - .unwrap(); + let results = translate( + "INSERT INTO users (name, age) VALUES ('Alice', 30)", + Dialect::MySQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); assert!(sql.contains("Alice"), "got: {sql}"); } diff --git a/crates/litewire-translate/src/postgres.rs b/crates/litewire-translate/src/postgres.rs index 7cc616d..a25db21 100644 --- a/crates/litewire-translate/src/postgres.rs +++ b/crates/litewire-translate/src/postgres.rs @@ -66,8 +66,11 @@ mod tests { #[test] fn int_types_to_integer() { - let results = - translate("CREATE TABLE t (a SMALLINT, b INT, c BIGINT)", Dialect::PostgreSQL).unwrap(); + let results = translate( + "CREATE TABLE t (a SMALLINT, b INT, c BIGINT)", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("SMALLINT"), "SMALLINT not rewritten: {sql}"); @@ -76,8 +79,11 @@ mod tests { #[test] fn float_to_real() { - let results = - translate("CREATE TABLE t (a FLOAT(8), b NUMERIC(10,2))", Dialect::PostgreSQL).unwrap(); + let results = translate( + "CREATE TABLE t (a FLOAT(8), b NUMERIC(10,2))", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("REAL"), "no REAL found: {sql}"); @@ -86,8 +92,11 @@ mod tests { #[test] fn varchar_to_text() { - let results = - translate("CREATE TABLE t (name VARCHAR(255), bio TEXT)", Dialect::PostgreSQL).unwrap(); + let results = translate( + "CREATE TABLE t (name VARCHAR(255), bio TEXT)", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("VARCHAR"), "VARCHAR not rewritten: {sql}"); @@ -122,12 +131,17 @@ mod tests { #[test] fn timestamp_to_text() { - let results = - translate("CREATE TABLE t (created TIMESTAMP, updated DATE)", Dialect::PostgreSQL) - .unwrap(); + let results = translate( + "CREATE TABLE t (created TIMESTAMP, updated DATE)", + Dialect::PostgreSQL, + ) + .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] @@ -148,8 +162,11 @@ mod tests { #[test] fn serial_to_integer() { - let results = - translate("CREATE TABLE t (id SERIAL PRIMARY KEY)", Dialect::PostgreSQL).unwrap(); + let results = translate( + "CREATE TABLE t (id SERIAL PRIMARY KEY)", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("SERIAL"), "SERIAL not rewritten: {sql}"); @@ -158,8 +175,11 @@ mod tests { #[test] fn bigserial_to_integer() { - let results = - translate("CREATE TABLE t (id BIGSERIAL PRIMARY KEY)", Dialect::PostgreSQL).unwrap(); + let results = translate( + "CREATE TABLE t (id BIGSERIAL PRIMARY KEY)", + Dialect::PostgreSQL, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(upper.contains("INTEGER"), "no INTEGER found: {sql}"); diff --git a/crates/litewire-translate/src/tds.rs b/crates/litewire-translate/src/tds.rs index bc192b4..da48c03 100644 --- a/crates/litewire-translate/src/tds.rs +++ b/crates/litewire-translate/src/tds.rs @@ -103,9 +103,11 @@ mod tests { #[test] fn int_types_to_integer() { - let results = - translate("CREATE TABLE t (a TINYINT, b SMALLINT, c INT, d BIGINT)", Dialect::TDS) - .unwrap(); + let results = translate( + "CREATE TABLE t (a TINYINT, b SMALLINT, c INT, d BIGINT)", + Dialect::TDS, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("TINYINT"), "TINYINT not rewritten: {sql}"); @@ -127,7 +129,10 @@ 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}"); } @@ -202,7 +207,10 @@ mod tests { let results = translate("CREATE TABLE t (id UNIQUEIDENTIFIER)", Dialect::TDS).unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); - assert!(!upper.contains("UNIQUEIDENTIFIER"), "UNIQUEIDENTIFIER not rewritten: {sql}"); + assert!( + !upper.contains("UNIQUEIDENTIFIER"), + "UNIQUEIDENTIFIER not rewritten: {sql}" + ); assert!(upper.contains("TEXT"), "no TEXT found: {sql}"); } @@ -217,8 +225,11 @@ mod tests { #[test] fn identity_column_option_removed() { - let results = - translate("CREATE TABLE t (id INT IDENTITY(1,1) PRIMARY KEY)", Dialect::TDS).unwrap(); + let results = translate( + "CREATE TABLE t (id INT IDENTITY(1,1) PRIMARY KEY)", + Dialect::TDS, + ) + .unwrap(); let sql = extract_sql(&results[0]); let upper = sql.to_ascii_uppercase(); assert!(!upper.contains("IDENTITY"), "IDENTITY not removed: {sql}"); diff --git a/crates/litewire/src/lib.rs b/crates/litewire/src/lib.rs index cd6234b..ac928e4 100644 --- a/crates/litewire/src/lib.rs +++ b/crates/litewire/src/lib.rs @@ -110,14 +110,18 @@ impl LiteWire { if let Some(addr) = self.mysql_listen { let config = litewire_mysql::MysqlFrontendConfig { listen: addr }; let frontend = litewire_mysql::MysqlFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); + handles.push(tokio::spawn(async move { + frontend.serve().await.map_err(Into::into) + })); } #[cfg(feature = "hrana")] if let Some(addr) = self.hrana_listen { let config = litewire_hrana::HranaFrontendConfig { listen: addr }; let frontend = litewire_hrana::HranaFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); + handles.push(tokio::spawn(async move { + frontend.serve().await.map_err(Into::into) + })); } #[cfg(feature = "postgres")] @@ -125,14 +129,18 @@ impl LiteWire { let config = litewire_postgres::PostgresFrontendConfig { listen: addr }; let frontend = litewire_postgres::PostgresFrontend::new(config, Arc::clone(&self.backend)); - handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); + handles.push(tokio::spawn(async move { + frontend.serve().await.map_err(Into::into) + })); } #[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)); - handles.push(tokio::spawn(async move { frontend.serve().await.map_err(Into::into) })); + handles.push(tokio::spawn(async move { + frontend.serve().await.map_err(Into::into) + })); } if handles.is_empty() { diff --git a/crates/litewire/tests/mysql_e2e.rs b/crates/litewire/tests/mysql_e2e.rs index 9963e99..35219bc 100644 --- a/crates/litewire/tests/mysql_e2e.rs +++ b/crates/litewire/tests/mysql_e2e.rs @@ -50,7 +50,10 @@ async fn connect(port: u16) -> Conn { } fn init_tracing() { - let _ = tracing_subscriber::fmt().with_env_filter("debug").with_test_writer().try_init(); + let _ = tracing_subscriber::fmt() + .with_env_filter("debug") + .with_test_writer() + .try_init(); } #[tokio::test] @@ -77,11 +80,17 @@ async fn create_table_insert_select() { .await .unwrap(); - conn.query_drop("INSERT INTO users (id, name) VALUES (1, 'Alice')").await.unwrap(); - conn.query_drop("INSERT INTO users (id, name) VALUES (2, 'Bob')").await.unwrap(); + conn.query_drop("INSERT INTO users (id, name) VALUES (1, 'Alice')") + .await + .unwrap(); + conn.query_drop("INSERT INTO users (id, name) VALUES (2, 'Bob')") + .await + .unwrap(); - let rows: Vec<(i64, String)> = - conn.query("SELECT id, name FROM users ORDER BY id").await.unwrap(); + let rows: Vec<(i64, String)> = conn + .query("SELECT id, name FROM users ORDER BY id") + .await + .unwrap(); assert_eq!(rows, vec![(1, "Alice".into()), (2, "Bob".into())]); drop(conn); @@ -98,7 +107,11 @@ async fn now_function_translates() { let result: Vec<(String,)> = conn.query("SELECT NOW()").await.unwrap(); assert_eq!(result.len(), 1); // Should look like "2024-01-15 12:34:56". - assert!(result[0].0.contains('-'), "expected datetime string, got: {}", result[0].0); + assert!( + result[0].0.contains('-'), + "expected datetime string, got: {}", + result[0].0 + ); drop(conn); } @@ -123,8 +136,12 @@ async fn show_tables() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE alpha (id INTEGER PRIMARY KEY)").await.unwrap(); - conn.query_drop("CREATE TABLE beta (id INTEGER PRIMARY KEY)").await.unwrap(); + conn.query_drop("CREATE TABLE alpha (id INTEGER PRIMARY KEY)") + .await + .unwrap(); + conn.query_drop("CREATE TABLE beta (id INTEGER PRIMARY KEY)") + .await + .unwrap(); let result: Vec<(String,)> = conn.query("SHOW TABLES").await.unwrap(); let names: Vec<&str> = result.iter().map(|r| r.0.as_str()).collect(); @@ -173,16 +190,29 @@ async fn update_and_delete() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)").await.unwrap(); - conn.query_drop("INSERT INTO items VALUES (1, 10)").await.unwrap(); - conn.query_drop("INSERT INTO items VALUES (2, 20)").await.unwrap(); + conn.query_drop("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)") + .await + .unwrap(); + conn.query_drop("INSERT INTO items VALUES (1, 10)") + .await + .unwrap(); + conn.query_drop("INSERT INTO items VALUES (2, 20)") + .await + .unwrap(); - conn.query_drop("UPDATE items SET qty = 15 WHERE id = 1").await.unwrap(); + conn.query_drop("UPDATE items SET qty = 15 WHERE id = 1") + .await + .unwrap(); - let rows: Vec<(i64, i64)> = conn.query("SELECT id, qty FROM items ORDER BY id").await.unwrap(); + let rows: Vec<(i64, i64)> = conn + .query("SELECT id, qty FROM items ORDER BY id") + .await + .unwrap(); assert_eq!(rows, vec![(1, 15), (2, 20)]); - conn.query_drop("DELETE FROM items WHERE id = 2").await.unwrap(); + conn.query_drop("DELETE FROM items WHERE id = 2") + .await + .unwrap(); let rows: Vec<(i64, i64)> = conn.query("SELECT id, qty FROM items").await.unwrap(); assert_eq!(rows, vec![(1, 15)]); @@ -197,8 +227,14 @@ async fn multiple_connections() { let _server = start_litewire(port).await; let mut conn1 = connect(port).await; - conn1.query_drop("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); - conn1.query_drop("INSERT INTO shared VALUES (1, 'from_conn1')").await.unwrap(); + conn1 + .query_drop("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); + conn1 + .query_drop("INSERT INTO shared VALUES (1, 'from_conn1')") + .await + .unwrap(); drop(conn1); // Second connection should see the data (same in-memory SQLite). @@ -218,17 +254,27 @@ async fn prepared_select_with_param() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)").await.unwrap(); - conn.query_drop("INSERT INTO users VALUES (1, 'Alice')").await.unwrap(); - conn.query_drop("INSERT INTO users VALUES (2, 'Bob')").await.unwrap(); + conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)") + .await + .unwrap(); + conn.query_drop("INSERT INTO users VALUES (1, 'Alice')") + .await + .unwrap(); + conn.query_drop("INSERT INTO users VALUES (2, 'Bob')") + .await + .unwrap(); // Prepared SELECT with parameter binding. - let rows: Vec<(i64, String)> = - conn.exec("SELECT id, name FROM users WHERE id = ?", (1_i64,)).await.unwrap(); + let rows: Vec<(i64, String)> = conn + .exec("SELECT id, name FROM users WHERE id = ?", (1_i64,)) + .await + .unwrap(); assert_eq!(rows, vec![(1, "Alice".into())]); - let rows: Vec<(i64, String)> = - conn.exec("SELECT id, name FROM users WHERE id = ?", (2_i64,)).await.unwrap(); + let rows: Vec<(i64, String)> = conn + .exec("SELECT id, name FROM users WHERE id = ?", (2_i64,)) + .await + .unwrap(); assert_eq!(rows, vec![(2, "Bob".into())]); drop(conn); @@ -246,17 +292,28 @@ async fn prepared_insert() { .unwrap(); // Prepared INSERT with parameters. - conn.exec_drop("INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", (1_i64, "Widget", 10_i64)) + conn.exec_drop( + "INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", + (1_i64, "Widget", 10_i64), + ) + .await + .unwrap(); + + conn.exec_drop( + "INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", + (2_i64, "Gadget", 20_i64), + ) + .await + .unwrap(); + + let rows: Vec<(i64, String, i64)> = conn + .query("SELECT id, name, qty FROM items ORDER BY id") .await .unwrap(); - - conn.exec_drop("INSERT INTO items (id, name, qty) VALUES (?, ?, ?)", (2_i64, "Gadget", 20_i64)) - .await - .unwrap(); - - let rows: Vec<(i64, String, i64)> = - conn.query("SELECT id, name, qty FROM items ORDER BY id").await.unwrap(); - assert_eq!(rows, vec![(1, "Widget".into(), 10), (2, "Gadget".into(), 20),]); + assert_eq!( + rows, + vec![(1, "Widget".into(), 10), (2, "Gadget".into(), 20),] + ); drop(conn); } @@ -268,15 +325,23 @@ async fn prepared_update_and_delete() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); - conn.query_drop("INSERT INTO t VALUES (1, 'old')").await.unwrap(); + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); + conn.query_drop("INSERT INTO t VALUES (1, 'old')") + .await + .unwrap(); - conn.exec_drop("UPDATE t SET val = ? WHERE id = ?", ("new", 1_i64)).await.unwrap(); + conn.exec_drop("UPDATE t SET val = ? WHERE id = ?", ("new", 1_i64)) + .await + .unwrap(); let rows: Vec<(i64, String)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert_eq!(rows, vec![(1, "new".into())]); - conn.exec_drop("DELETE FROM t WHERE id = ?", (1_i64,)).await.unwrap(); + conn.exec_drop("DELETE FROM t WHERE id = ?", (1_i64,)) + .await + .unwrap(); let rows: Vec<(i64, String)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert!(rows.is_empty()); @@ -291,12 +356,17 @@ async fn prepared_with_null_param() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); - - conn.exec_drop("INSERT INTO t (id, val) VALUES (?, ?)", (1_i64, Option::::None)) + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)") .await .unwrap(); + conn.exec_drop( + "INSERT INTO t (id, val) VALUES (?, ?)", + (1_i64, Option::::None), + ) + .await + .unwrap(); + let rows: Vec<(i64, Option)> = conn.query("SELECT id, val FROM t").await.unwrap(); assert_eq!(rows, vec![(1, None)]); @@ -310,7 +380,9 @@ async fn prepared_reuse_same_statement() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY)").await.unwrap(); + conn.query_drop("CREATE TABLE t (id INTEGER PRIMARY KEY)") + .await + .unwrap(); // Execute the same prepared statement multiple times. let stmt = conn.prep("INSERT INTO t (id) VALUES (?)").await.unwrap(); @@ -337,16 +409,23 @@ async fn on_duplicate_key_update() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE kv (k TEXT PRIMARY KEY, v INTEGER)").await.unwrap(); + conn.query_drop("CREATE TABLE kv (k TEXT PRIMARY KEY, v INTEGER)") + .await + .unwrap(); - conn.query_drop("INSERT INTO kv (k, v) VALUES ('a', 1)").await.unwrap(); + conn.query_drop("INSERT INTO kv (k, v) VALUES ('a', 1)") + .await + .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(); - let rows: Vec<(String, i64)> = conn.query("SELECT k, v FROM kv WHERE k = 'a'").await.unwrap(); + let rows: Vec<(String, i64)> = conn + .query("SELECT k, v FROM kv WHERE k = 'a'") + .await + .unwrap(); assert_eq!(rows, vec![("a".into(), 99)]); // Insert new row (no conflict). @@ -369,10 +448,14 @@ async fn transaction_commit() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + conn.query_drop("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("INSERT INTO txn_t VALUES (1, 'inside_txn')").await.unwrap(); + conn.query_drop("INSERT INTO txn_t VALUES (1, 'inside_txn')") + .await + .unwrap(); conn.query_drop("COMMIT").await.unwrap(); // Data should be visible after commit. @@ -389,17 +472,25 @@ async fn transaction_rollback() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + conn.query_drop("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); - conn.query_drop("INSERT INTO txn_rb VALUES (1, 'before')").await.unwrap(); + conn.query_drop("INSERT INTO txn_rb VALUES (1, 'before')") + .await + .unwrap(); conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("INSERT INTO txn_rb VALUES (2, 'rolled_back')").await.unwrap(); + conn.query_drop("INSERT INTO txn_rb VALUES (2, 'rolled_back')") + .await + .unwrap(); conn.query_drop("ROLLBACK").await.unwrap(); // Only the row inserted before the transaction should exist. - let rows: Vec<(i64, String)> = - conn.query("SELECT id, val FROM txn_rb ORDER BY id").await.unwrap(); + let rows: Vec<(i64, String)> = conn + .query("SELECT id, val FROM txn_rb ORDER BY id") + .await + .unwrap(); assert_eq!(rows, vec![(1, "before".into())]); drop(conn); @@ -412,13 +503,19 @@ async fn transaction_atomicity() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE txn_atom (id INTEGER PRIMARY KEY, val INTEGER)").await.unwrap(); + conn.query_drop("CREATE TABLE txn_atom (id INTEGER PRIMARY KEY, val INTEGER)") + .await + .unwrap(); - conn.query_drop("INSERT INTO txn_atom VALUES (1, 100)").await.unwrap(); + conn.query_drop("INSERT INTO txn_atom VALUES (1, 100)") + .await + .unwrap(); // Begin a transaction, update, then rollback — value should remain 100. conn.query_drop("BEGIN").await.unwrap(); - conn.query_drop("UPDATE txn_atom SET val = 200 WHERE id = 1").await.unwrap(); + conn.query_drop("UPDATE txn_atom SET val = 200 WHERE id = 1") + .await + .unwrap(); conn.query_drop("ROLLBACK").await.unwrap(); let rows: Vec<(i64, i64)> = conn.query("SELECT id, val FROM txn_atom").await.unwrap(); @@ -434,12 +531,18 @@ async fn information_schema_tables() { let _server = start_litewire(port).await; let mut conn = connect(port).await; - conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY)").await.unwrap(); - conn.query_drop("CREATE TABLE posts (id INTEGER PRIMARY KEY)").await.unwrap(); + conn.query_drop("CREATE TABLE users (id INTEGER PRIMARY KEY)") + .await + .unwrap(); + conn.query_drop("CREATE TABLE posts (id INTEGER PRIMARY KEY)") + .await + .unwrap(); // Returns TABLE_NAME, TABLE_TYPE, TABLE_SCHEMA columns. - let rows: Vec<(String, String, String)> = - conn.query("SELECT TABLE_NAME FROM information_schema.tables").await.unwrap(); + let rows: Vec<(String, String, String)> = conn + .query("SELECT TABLE_NAME FROM information_schema.tables") + .await + .unwrap(); let names: Vec<&str> = rows.iter().map(|r| r.0.as_str()).collect(); assert!(names.contains(&"users"), "got: {names:?}"); assert!(names.contains(&"posts"), "got: {names:?}"); @@ -463,13 +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(); + 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); } @@ -487,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 f9ff829..8dbee4a 100644 --- a/crates/litewire/tests/postgres_e2e.rs +++ b/crates/litewire/tests/postgres_e2e.rs @@ -52,7 +52,10 @@ async fn connect(port: u16) -> Client { } fn init_tracing() { - let _ = tracing_subscriber::fmt().with_env_filter("debug").with_test_writer().try_init(); + let _ = tracing_subscriber::fmt() + .with_env_filter("debug") + .with_test_writer() + .try_init(); } // ── Simple query tests ───────────────────────────────────────────────────── @@ -78,14 +81,26 @@ async fn create_table_insert_select() { let client = connect(port).await; client - .execute("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)", &[]) + .execute( + "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL)", + &[], + ) .await .unwrap(); - client.execute("INSERT INTO users (id, name) VALUES (1, 'Alice')", &[]).await.unwrap(); - client.execute("INSERT INTO users (id, name) VALUES (2, 'Bob')", &[]).await.unwrap(); + client + .execute("INSERT INTO users (id, name) VALUES (1, 'Alice')", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO users (id, name) VALUES (2, 'Bob')", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, name FROM users ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, name FROM users ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 2); let id: i64 = rows[0].get(0); @@ -106,20 +121,44 @@ async fn update_and_delete() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)", &[]).await.unwrap(); - client.execute("INSERT INTO items VALUES (1, 10)", &[]).await.unwrap(); - client.execute("INSERT INTO items VALUES (2, 20)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE items (id INTEGER PRIMARY KEY, qty INTEGER)", + &[], + ) + .await + .unwrap(); + client + .execute("INSERT INTO items VALUES (1, 10)", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO items VALUES (2, 20)", &[]) + .await + .unwrap(); - client.execute("UPDATE items SET qty = 15 WHERE id = 1", &[]).await.unwrap(); + client + .execute("UPDATE items SET qty = 15 WHERE id = 1", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, qty FROM items ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, qty FROM items ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 2); let qty: i64 = rows[0].get(1); assert_eq!(qty, 15); - client.execute("DELETE FROM items WHERE id = 2", &[]).await.unwrap(); + client + .execute("DELETE FROM items WHERE id = 2", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, qty FROM items", &[]).await.unwrap(); + let rows = client + .query("SELECT id, qty FROM items", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let id: i64 = rows[0].get(0); assert_eq!(id, 1); @@ -178,13 +217,25 @@ async fn multiple_connections() { let _server = start_litewire(port).await; let client1 = connect(port).await; - client1.execute("CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); - client1.execute("INSERT INTO shared VALUES (1, 'from_conn1')", &[]).await.unwrap(); + client1 + .execute( + "CREATE TABLE shared (id INTEGER PRIMARY KEY, val TEXT)", + &[], + ) + .await + .unwrap(); + client1 + .execute("INSERT INTO shared VALUES (1, 'from_conn1')", &[]) + .await + .unwrap(); drop(client1); // Second connection should see the data. let client2 = connect(port).await; - let rows = client2.query("SELECT id, val FROM shared", &[]).await.unwrap(); + let rows = client2 + .query("SELECT id, val FROM shared", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "from_conn1"); @@ -197,7 +248,13 @@ async fn empty_table_query() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE empty_t (id INTEGER PRIMARY KEY, val TEXT)", + &[], + ) + .await + .unwrap(); let rows = client.query("SELECT * FROM empty_t", &[]).await.unwrap(); assert!(rows.is_empty()); @@ -210,7 +267,10 @@ async fn drop_table() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE temp_t (id INTEGER PRIMARY KEY)", &[]).await.unwrap(); + client + .execute("CREATE TABLE temp_t (id INTEGER PRIMARY KEY)", &[]) + .await + .unwrap(); client.execute("DROP TABLE temp_t", &[]).await.unwrap(); // Selecting from dropped table should fail. @@ -228,12 +288,27 @@ async fn pg_serial_to_integer() { let client = connect(port).await; // SERIAL should be translated to INTEGER. - client.execute("CREATE TABLE auto_t (id SERIAL PRIMARY KEY, name TEXT)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE auto_t (id SERIAL PRIMARY KEY, name TEXT)", + &[], + ) + .await + .unwrap(); - client.execute("INSERT INTO auto_t (name) VALUES ('first')", &[]).await.unwrap(); - client.execute("INSERT INTO auto_t (name) VALUES ('second')", &[]).await.unwrap(); + client + .execute("INSERT INTO auto_t (name) VALUES ('first')", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO auto_t (name) VALUES ('second')", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, name FROM auto_t ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, name FROM auto_t ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 2); let name: &str = rows[0].get(1); assert_eq!(name, "first"); @@ -247,16 +322,25 @@ async fn pg_varchar_to_text() { let client = connect(port).await; client - .execute("CREATE TABLE typed_t (name VARCHAR(255), bio TEXT, data BYTEA)", &[]) + .execute( + "CREATE TABLE typed_t (name VARCHAR(255), bio TEXT, data BYTEA)", + &[], + ) .await .unwrap(); client - .execute("INSERT INTO typed_t (name, bio, data) VALUES ('test', 'a bio', 'raw')", &[]) + .execute( + "INSERT INTO typed_t (name, bio, data) VALUES ('test', 'a bio', 'raw')", + &[], + ) .await .unwrap(); - let rows = client.query("SELECT name, bio FROM typed_t", &[]).await.unwrap(); + let rows = client + .query("SELECT name, bio FROM typed_t", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let name: &str = rows[0].get(0); assert_eq!(name, "test"); @@ -271,14 +355,26 @@ async fn pg_boolean_column() { // BOOLEAN should be translated to INTEGER in SQLite. client - .execute("CREATE TABLE flags (id INTEGER PRIMARY KEY, active BOOLEAN)", &[]) + .execute( + "CREATE TABLE flags (id INTEGER PRIMARY KEY, active BOOLEAN)", + &[], + ) .await .unwrap(); - client.execute("INSERT INTO flags VALUES (1, TRUE)", &[]).await.unwrap(); - client.execute("INSERT INTO flags VALUES (2, FALSE)", &[]).await.unwrap(); + client + .execute("INSERT INTO flags VALUES (1, TRUE)", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO flags VALUES (2, FALSE)", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, active FROM flags ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, active FROM flags ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 2); let active: i64 = rows[0].get(1); assert_eq!(active, 1); @@ -295,14 +391,23 @@ async fn large_result_set() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE big_t (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + client + .execute("CREATE TABLE big_t (id INTEGER PRIMARY KEY, val TEXT)", &[]) + .await + .unwrap(); // Insert 100 rows. for i in 0..100 { - client.execute(&format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), &[]).await.unwrap(); + client + .execute(&format!("INSERT INTO big_t VALUES ({i}, 'row_{i}')"), &[]) + .await + .unwrap(); } - let rows = client.query("SELECT id, val FROM big_t ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, val FROM big_t ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 100); let first_id: i64 = rows[0].get(0); @@ -318,12 +423,27 @@ async fn null_handling() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE nullable (id INTEGER PRIMARY KEY, val TEXT)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE nullable (id INTEGER PRIMARY KEY, val TEXT)", + &[], + ) + .await + .unwrap(); - client.execute("INSERT INTO nullable VALUES (1, NULL)", &[]).await.unwrap(); - client.execute("INSERT INTO nullable VALUES (2, 'present')", &[]).await.unwrap(); + client + .execute("INSERT INTO nullable VALUES (1, NULL)", &[]) + .await + .unwrap(); + client + .execute("INSERT INTO nullable VALUES (2, 'present')", &[]) + .await + .unwrap(); - let rows = client.query("SELECT id, val FROM nullable ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, val FROM nullable ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 2); let val: Option<&str> = rows[0].get(1); @@ -343,14 +463,23 @@ async fn transaction_commit() { let client = connect(port).await; // Use batch_execute (simple query protocol) for all transaction-related ops. - client.batch_execute("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + client + .batch_execute("CREATE TABLE txn_t (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); client.batch_execute("BEGIN").await.unwrap(); - client.batch_execute("INSERT INTO txn_t VALUES (1, 'inside_txn')").await.unwrap(); + client + .batch_execute("INSERT INTO txn_t VALUES (1, 'inside_txn')") + .await + .unwrap(); client.batch_execute("COMMIT").await.unwrap(); // Data should be visible after commit. - let rows = client.query("SELECT id, val FROM txn_t", &[]).await.unwrap(); + let rows = client + .query("SELECT id, val FROM txn_t", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "inside_txn"); @@ -363,16 +492,28 @@ async fn transaction_rollback() { let _server = start_litewire(port).await; let client = connect(port).await; - client.batch_execute("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)").await.unwrap(); + client + .batch_execute("CREATE TABLE txn_rb (id INTEGER PRIMARY KEY, val TEXT)") + .await + .unwrap(); - client.batch_execute("INSERT INTO txn_rb VALUES (1, 'before')").await.unwrap(); + client + .batch_execute("INSERT INTO txn_rb VALUES (1, 'before')") + .await + .unwrap(); client.batch_execute("BEGIN").await.unwrap(); - client.batch_execute("INSERT INTO txn_rb VALUES (2, 'rolled_back')").await.unwrap(); + client + .batch_execute("INSERT INTO txn_rb VALUES (2, 'rolled_back')") + .await + .unwrap(); client.batch_execute("ROLLBACK").await.unwrap(); // Only the row inserted before the transaction should exist. - let rows = client.query("SELECT id, val FROM txn_rb ORDER BY id", &[]).await.unwrap(); + let rows = client + .query("SELECT id, val FROM txn_rb ORDER BY id", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let val: &str = rows[0].get(1); assert_eq!(val, "before"); @@ -390,14 +531,23 @@ async fn transaction_atomicity() { .await .unwrap(); - client.batch_execute("INSERT INTO txn_atom VALUES (1, 100)").await.unwrap(); + client + .batch_execute("INSERT INTO txn_atom VALUES (1, 100)") + .await + .unwrap(); // Begin a transaction, update, then rollback — value should remain 100. client.batch_execute("BEGIN").await.unwrap(); - client.batch_execute("UPDATE txn_atom SET val = 200 WHERE id = 1").await.unwrap(); + client + .batch_execute("UPDATE txn_atom SET val = 200 WHERE id = 1") + .await + .unwrap(); client.batch_execute("ROLLBACK").await.unwrap(); - let rows = client.query("SELECT id, val FROM txn_atom", &[]).await.unwrap(); + let rows = client + .query("SELECT id, val FROM txn_atom", &[]) + .await + .unwrap(); assert_eq!(rows.len(), 1); let val: i64 = rows[0].get(1); assert_eq!(val, 100); @@ -410,11 +560,23 @@ async fn float_values() { let _server = start_litewire(port).await; let client = connect(port).await; - client.execute("CREATE TABLE floats (id INTEGER PRIMARY KEY, val REAL)", &[]).await.unwrap(); + client + .execute( + "CREATE TABLE floats (id INTEGER PRIMARY KEY, val REAL)", + &[], + ) + .await + .unwrap(); - client.execute("INSERT INTO floats VALUES (1, 3.14)", &[]).await.unwrap(); + client + .execute("INSERT INTO floats VALUES (1, 3.14)", &[]) + .await + .unwrap(); - let rows = client.query("SELECT val FROM floats WHERE id = 1", &[]).await.unwrap(); + let rows = client + .query("SELECT val FROM floats WHERE id = 1", &[]) + .await + .unwrap(); let val: f64 = rows[0].get(0); assert!((val - 3.14).abs() < 0.001); } diff --git a/crates/litewire/tests/tds_e2e.rs b/crates/litewire/tests/tds_e2e.rs index 61469b8..4baa32f 100644 --- a/crates/litewire/tests/tds_e2e.rs +++ b/crates/litewire/tests/tds_e2e.rs @@ -67,7 +67,10 @@ async fn connect(port: u16) -> Client