diff --git a/jdbc-bridge/src/main/java/sqlkit/bridge/QueryExecutor.java b/jdbc-bridge/src/main/java/sqlkit/bridge/QueryExecutor.java
index 7154995..2892853 100644
--- a/jdbc-bridge/src/main/java/sqlkit/bridge/QueryExecutor.java
+++ b/jdbc-bridge/src/main/java/sqlkit/bridge/QueryExecutor.java
@@ -4,70 +4,69 @@
import java.util.*;
/**
- * Executes SQL queries against a JDBC connection and serializes results to JSON-compatible Maps.
+ * Executes SQL statements against a JDBC connection and serializes results to JSON-compatible Maps.
*/
public class QueryExecutor {
/**
- * Execute a SQL query and return the result as a Map.
+ * Execute a SQL statement and return the result as a Map.
*
- * For SELECT queries, returns {columns: [...], rows: [[...], ...]}.
- * For UPDATE/INSERT/DELETE, returns {rows_affected: N}.
+ * Uses {@link Statement#execute(String)} so no keyword sniffing is needed:
+ * statements that return rows (SELECT, WITH ... SELECT, INSERT ... RETURNING)
+ * yield {columns, rows}; everything else (INSERT/UPDATE/DELETE/MERGE/DDL)
+ * yields {rows_affected: N}.
*/
public static Map execute(Connection conn, String sql) throws Exception {
- sql = sql.trim();
-
- boolean isQuery;
- String upper = sql.toUpperCase().trim();
- isQuery = upper.startsWith("SELECT")
- || upper.startsWith("WITH")
- || upper.startsWith("EXPLAIN")
- || upper.startsWith("SHOW")
- || upper.startsWith("DESCRIBE")
- || upper.startsWith("PRAGMA");
-
- if (isQuery) {
- return executeQuery(conn, sql);
- } else {
- return executeUpdate(conn, sql);
- }
- }
+ try (Statement stmt = conn.createStatement()) {
+ boolean isResultSet = stmt.execute(sql);
- private static Map executeQuery(Connection conn, String sql) throws Exception {
- try (Statement stmt = conn.createStatement();
- ResultSet rs = stmt.executeQuery(sql)) {
+ List columns = null;
+ List> rows = null;
+ long rowsAffected = 0;
- ResultSetMetaData meta = rs.getMetaData();
- int columnCount = meta.getColumnCount();
+ while (true) {
+ if (isResultSet) {
+ try (ResultSet rs = stmt.getResultSet()) {
+ ResultSetMetaData meta = rs.getMetaData();
+ int columnCount = meta.getColumnCount();
- List columns = new ArrayList<>();
- for (int i = 1; i <= columnCount; i++) {
- columns.add(meta.getColumnLabel(i));
- }
+ List resultColumns = new ArrayList<>();
+ for (int i = 1; i <= columnCount; i++) {
+ resultColumns.add(meta.getColumnLabel(i));
+ }
- List> rows = new ArrayList<>();
- while (rs.next()) {
- List row = new ArrayList<>();
- for (int i = 1; i <= columnCount; i++) {
- row.add(getValue(rs, i));
+ List> resultRows = new ArrayList<>();
+ while (rs.next()) {
+ List row = new ArrayList<>();
+ for (int i = 1; i <= columnCount; i++) {
+ row.add(getValue(rs, i));
+ }
+ resultRows.add(row);
+ }
+ columns = resultColumns;
+ rows = resultRows;
+ }
+ } else {
+ int count = stmt.getUpdateCount();
+ if (count >= 0) {
+ rowsAffected = count;
+ }
+ }
+ isResultSet = stmt.getMoreResults();
+ if (!isResultSet && stmt.getUpdateCount() == -1) {
+ break;
}
- rows.add(row);
}
Map result = new LinkedHashMap<>();
- result.put("columns", columns);
- result.put("rows", rows);
- return result;
- }
- }
-
- private static Map executeUpdate(Connection conn, String sql) throws Exception {
- try (Statement stmt = conn.createStatement()) {
- int affected = stmt.executeUpdate(sql);
- Map result = new LinkedHashMap<>();
- result.put("rows_affected", (long) affected);
- result.put("columns", Collections.emptyList());
- result.put("rows", Collections.emptyList());
+ if (rows != null) {
+ result.put("columns", columns);
+ result.put("rows", rows);
+ } else {
+ result.put("rows_affected", rowsAffected);
+ result.put("columns", Collections.emptyList());
+ result.put("rows", Collections.emptyList());
+ }
return result;
}
}
diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock
index 27b49b6..4b6cc47 100644
--- a/src-tauri/Cargo.lock
+++ b/src-tauri/Cargo.lock
@@ -1770,7 +1770,7 @@ dependencies = [
"libc",
"option-ext",
"redox_users",
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
@@ -2062,7 +2062,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
dependencies = [
"libc",
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
@@ -3037,7 +3037,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
- "socket2 0.5.10",
+ "socket2 0.6.4",
"system-configuration",
"tokio",
"tower-service",
@@ -3057,7 +3057,7 @@ dependencies = [
"js-sys",
"log",
"wasm-bindgen",
- "windows-core 0.61.2",
+ "windows-core 0.62.2",
]
[[package]]
@@ -5083,7 +5083,7 @@ dependencies = [
"quinn-udp",
"rustc-hash 2.1.2",
"rustls 0.23.40",
- "socket2 0.5.10",
+ "socket2 0.6.4",
"thiserror 2.0.18",
"tokio",
"tracing",
@@ -5120,9 +5120,9 @@ dependencies = [
"cfg_aliases",
"libc",
"once_cell",
- "socket2 0.5.10",
+ "socket2 0.6.4",
"tracing",
- "windows-sys 0.59.0",
+ "windows-sys 0.60.2",
]
[[package]]
@@ -5744,7 +5744,7 @@ dependencies = [
"errno",
"libc",
"linux-raw-sys",
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
@@ -5845,7 +5845,7 @@ dependencies = [
"security-framework 3.7.0",
"security-framework-sys",
"webpki-root-certs",
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
@@ -6566,9 +6566,9 @@ dependencies = [
[[package]]
name = "sqlparser"
-version = "0.55.0"
+version = "0.62.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
-checksum = "c4521174166bac1ff04fe16ef4524c70144cd29682a45978978ca3d7f4e0be11"
+checksum = "13c6d1b651dc4edf07eead2a0c6c78016ce971bc2c10da5266861b13f25e7cec"
dependencies = [
"log",
"recursive",
@@ -7271,7 +7271,7 @@ dependencies = [
"getrandom 0.4.3",
"once_cell",
"rustix",
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
@@ -8361,7 +8361,7 @@ version = "0.1.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
dependencies = [
- "windows-sys 0.59.0",
+ "windows-sys 0.61.2",
]
[[package]]
diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml
index 0c116fa..11d032a 100644
--- a/src-tauri/Cargo.toml
+++ b/src-tauri/Cargo.toml
@@ -98,7 +98,7 @@ zip = { version = "2", features = [ "deflate" ] }
# TOML parsing (drivers.toml)
toml = "0.8"
sha2 = "0.10"
-sqlparser = "0.55"
+sqlparser = "0.62"
russh = "0.60"
axum = "0.8"
portpicker = "0.1.1"
diff --git a/src-tauri/src/capabilities/sql_write.rs b/src-tauri/src/capabilities/sql_write.rs
index 9112edc..3c13f37 100644
--- a/src-tauri/src/capabilities/sql_write.rs
+++ b/src-tauri/src/capabilities/sql_write.rs
@@ -7,7 +7,7 @@
use std::sync::Arc;
use serde_json::{json, Value};
-use sqlparser::ast::Statement;
+use sqlparser::ast::{Query, SetExpr, Statement};
use sqlparser::dialect::{GenericDialect, MsSqlDialect, MySqlDialect, PostgreSqlDialect};
use sqlparser::parser::Parser;
@@ -49,10 +49,60 @@ pub(crate) fn classify_sql(db_type: &str, sql: &str) -> Result
Ok(classify_statement(&stmt))
}
+/// Merge kinds found in different parts of one statement.
+///
+/// The most restrictive kind wins so a destructive statement never rides
+/// along inside a tree otherwise classified as Read.
+fn combine_kind(a: SqlKind, b: SqlKind) -> SqlKind {
+ fn rank(kind: SqlKind) -> u8 {
+ match kind {
+ SqlKind::Read => 0,
+ SqlKind::Write => 1,
+ SqlKind::Ddl => 2,
+ SqlKind::Delete => 3,
+ SqlKind::Other => 4,
+ }
+ }
+ if rank(b) > rank(a) {
+ b
+ } else {
+ a
+ }
+}
+
+/// Classify a Query, walking into its body and every CTE.
+///
+/// sqlparser represents `WITH … INSERT/UPDATE/DELETE/MERGE` (data-modifying
+/// CTEs) as a `Query` whose body carries the DML statement, so checking only
+/// the top-level variant would let write statements through as Read.
+fn classify_query(query: &Query) -> SqlKind {
+ let mut kind = classify_set_expr(&query.body);
+ if let Some(with) = &query.with {
+ for cte in &with.cte_tables {
+ kind = combine_kind(kind, classify_query(&cte.query));
+ }
+ }
+ kind
+}
+
+fn classify_set_expr(body: &SetExpr) -> SqlKind {
+ match body {
+ SetExpr::Select(_) | SetExpr::Values(_) | SetExpr::Table(_) => SqlKind::Read,
+ SetExpr::Query(query) => classify_query(query),
+ SetExpr::SetOperation { left, right, .. } => {
+ combine_kind(classify_set_expr(left), classify_set_expr(right))
+ }
+ SetExpr::Insert(stmt)
+ | SetExpr::Update(stmt)
+ | SetExpr::Delete(stmt)
+ | SetExpr::Merge(stmt) => classify_statement(stmt),
+ }
+}
+
fn classify_statement(stmt: &Statement) -> SqlKind {
match stmt {
- Statement::Query(_)
- | Statement::Explain { .. }
+ Statement::Query(query) => classify_query(query),
+ Statement::Explain { .. }
| Statement::ExplainTable { .. }
| Statement::ShowVariable { .. }
| Statement::ShowVariables { .. }
@@ -119,13 +169,7 @@ fn classify_statement(stmt: &Statement) -> SqlKind {
| Statement::Comment { .. }
| Statement::LockTables { .. }
| Statement::UnlockTables { .. }
- | Statement::SetVariable { .. }
- | Statement::SetNames { .. }
- | Statement::SetNamesDefault { .. }
- | Statement::SetRole { .. }
- | Statement::SetSessionParam(_)
- | Statement::SetTimeZone { .. }
- | Statement::SetTransaction { .. }
+ | Statement::Set(_)
| Statement::Commit { .. }
| Statement::Rollback { .. }
| Statement::Savepoint { .. }
@@ -408,6 +452,89 @@ mod tests {
);
}
+ #[test]
+ fn classifies_with_dml_chain_as_write() {
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH u AS (SELECT id FROM users WHERE email = 'seven@wentsen.com'), \
+ o AS (INSERT INTO organizations (name, slug, created_by_email) \
+ VALUES ('wentsen', 'release-org', 'seven@wentsen.com') RETURNING id) \
+ INSERT INTO memberships (user_id, org_id, role, scope, permissions) \
+ SELECT u.id, o.id, 'owner', 'org', '{}'::jsonb FROM u, o"
+ )
+ .unwrap(),
+ SqlKind::Write
+ );
+ }
+
+ #[test]
+ fn classifies_with_dml_in_cte_as_write() {
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH i AS (INSERT INTO t (c) VALUES (1) RETURNING id) SELECT * FROM i"
+ )
+ .unwrap(),
+ SqlKind::Write
+ );
+ }
+
+ #[test]
+ fn classifies_with_update_as_write() {
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH x AS (SELECT 1 AS v) UPDATE t SET c = (SELECT v FROM x)"
+ )
+ .unwrap(),
+ SqlKind::Write
+ );
+ }
+
+ #[test]
+ fn classifies_with_delete_as_delete() {
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH x AS (SELECT 1 AS id) DELETE FROM t WHERE id IN (SELECT id FROM x)"
+ )
+ .unwrap(),
+ SqlKind::Delete
+ );
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH d AS (DELETE FROM t RETURNING id) SELECT * FROM d"
+ )
+ .unwrap(),
+ SqlKind::Delete
+ );
+ }
+
+ #[test]
+ fn classifies_with_merge_as_write() {
+ assert_eq!(
+ classify_sql(
+ "postgres",
+ "WITH x AS (SELECT 1 AS id) MERGE INTO t USING x ON t.id = x.id \
+ WHEN MATCHED THEN UPDATE SET c = 1"
+ )
+ .unwrap(),
+ SqlKind::Write
+ );
+ }
+
+ #[test]
+ fn read_guard_rejects_with_insert() {
+ let err = ensure_read_only(
+ "postgres",
+ "WITH x AS (SELECT 1) INSERT INTO t (c) SELECT 1",
+ )
+ .unwrap_err();
+ assert!(err.contains("execute_write"), "got: {}", err);
+ }
+
#[test]
fn classifies_ddl_as_ddl() {
assert_eq!(
diff --git a/src-tauri/src/database/mysql.rs b/src-tauri/src/database/mysql.rs
index c2d3672..14a9850 100644
--- a/src-tauri/src/database/mysql.rs
+++ b/src-tauri/src/database/mysql.rs
@@ -625,7 +625,8 @@ impl DatabaseAdapter for MySQLAdapter {
let is_select = query_trimmed.starts_with("SELECT")
|| query_trimmed.starts_with("SHOW")
|| query_trimmed.starts_with("DESCRIBE")
- || query_trimmed.starts_with("EXPLAIN");
+ || query_trimmed.starts_with("EXPLAIN")
+ || query_trimmed.starts_with("WITH");
let execution_time;
diff --git a/src-tauri/src/database/postgres.rs b/src-tauri/src/database/postgres.rs
index 33c9c1a..55eca29 100644
--- a/src-tauri/src/database/postgres.rs
+++ b/src-tauri/src/database/postgres.rs
@@ -1077,6 +1077,32 @@ impl DatabaseAdapter for PostgresAdapter {
.map_err(postgres_error_to_db_error)?
};
+ // A WITH-prefixed statement may be data-modifying — PostgreSQL runs
+ // `WITH … INSERT/UPDATE/DELETE` (no RETURNING) as a query that
+ // exposes zero columns, so a plain query() would silently swallow
+ // the writes. Route column-less statements through execute() so
+ // callers receive a rows_affected count instead of an empty result.
+ if statement.columns().is_empty() {
+ let affected = if let Some(timeout_duration) = timeout {
+ tokio::time::timeout(timeout_duration, client.execute(&statement, &[]))
+ .await
+ .map_err(|_| {
+ DbError::Timeout(format!(
+ "Query timed out after {:?}",
+ timeout_duration
+ ))
+ })?
+ .map_err(postgres_error_to_db_error)?
+ } else {
+ client
+ .execute(&statement, &[])
+ .await
+ .map_err(postgres_error_to_db_error)?
+ };
+ return Ok(QueryResult::affected(affected)
+ .with_execution_time(start.elapsed().as_millis() as u64));
+ }
+
let result = if let Some(timeout_duration) = timeout {
tokio::time::timeout(timeout_duration, client.query(&statement, &[]))
.await
diff --git a/src-tauri/src/database/sql_service.rs b/src-tauri/src/database/sql_service.rs
index 671e214..396a4cd 100644
--- a/src-tauri/src/database/sql_service.rs
+++ b/src-tauri/src/database/sql_service.rs
@@ -485,6 +485,10 @@ fn extract_select_items(items: &[SelectItem]) -> Vec<(String, usize)> {
let name = match item {
SelectItem::UnnamedExpr(expr) => expr_to_name(expr),
SelectItem::ExprWithAlias { alias, .. } => alias.value.clone(),
+ SelectItem::ExprWithAliases { expr, aliases } => aliases
+ .first()
+ .map(|a| a.value.clone())
+ .unwrap_or_else(|| expr_to_name(expr)),
SelectItem::QualifiedWildcard(kind, _) => match kind {
SelectItemQualifiedWildcardKind::ObjectName(obj) => obj.to_string() + ".*",
SelectItemQualifiedWildcardKind::Expr(e) => format!("({}).*", expr_to_name(e)),
diff --git a/src-tauri/src/database/sqlite.rs b/src-tauri/src/database/sqlite.rs
index 53b6617..0f3de18 100644
--- a/src-tauri/src/database/sqlite.rs
+++ b/src-tauri/src/database/sqlite.rs
@@ -380,7 +380,8 @@ impl SQLiteAdapter {
let trimmed = query.trim().to_uppercase();
let is_select = trimmed.starts_with("SELECT")
|| trimmed.starts_with("PRAGMA")
- || trimmed.starts_with("EXPLAIN");
+ || trimmed.starts_with("EXPLAIN")
+ || trimmed.starts_with("WITH");
if is_select {
let mut stmt = conn_guard