diff --git a/Cargo.lock b/Cargo.lock index 5d25989..0f78a5c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -82,10 +82,12 @@ dependencies = [ [[package]] name = "clickhouse-query-ext" -version = "1.0.1" +version = "1.0.2" dependencies = [ "anyhow", + "bytes", "futures", + "libc", "regex", "reqwest", "secrecy", diff --git a/Cargo.toml b/Cargo.toml index 904d8b9..03b2c86 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "clickhouse-query-ext" -version = "1.0.1" +version = "1.0.2" edition = "2024" authors = ["Querya Community"] description = "High-performance ClickHouse driver for Querya Desktop" @@ -8,8 +8,9 @@ description = "High-performance ClickHouse driver for Querya Desktop" [dependencies] # Async runtime & I/O tokio = { version = "1", features = ["rt-multi-thread", "io-std", "sync", "macros", "time", "net"] } -tokio-util = { version = "0.7", features = ["codec"] } +tokio-util = { version = "0.7", features = ["codec", "io"] } futures = "0.3" +bytes = "1" # HTTP & TLS for ClickHouse API reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls"] } @@ -32,6 +33,9 @@ regex = "1" thiserror = "2.0" anyhow = "1.0" +[target.'cfg(unix)'.dependencies] +libc = "0.2" + [profile.release] lto = true strip = true diff --git a/README.md b/README.md index a4c729a..4323322 100644 --- a/README.md +++ b/README.md @@ -33,8 +33,8 @@ cargo build --release | Файл | Назначение | |------|------------| | `target/release/clickhouse-query-ext` | Бинарник драйвера | -| `dist/clickhouse-query-ext-1.0.1.qext` | Пакет для установки в Querya Desktop | -| `dist/clickhouse-query-ext-1.0.1.qext.sha256` | Контрольная сумма | +| `dist/clickhouse-query-ext-1.0.2.qext` | Пакет для установки в Querya Desktop | +| `dist/clickhouse-query-ext-1.0.2.qext.sha256` | Контрольная сумма | Установка: **Querya Desktop → Extensions → Install from file** → выбрать `.qext`. diff --git a/manifest.json b/manifest.json index 32c86eb..e6d0662 100644 --- a/manifest.json +++ b/manifest.json @@ -1,14 +1,14 @@ { "id": "queryahub.clickhouse-driver", "name": "ClickHouse Database Driver (Analyst Edition)", - "version": "1.0.1", + "version": "1.0.2", "publisher": "Querya Community", "description": "Изолированный нативный Rust-драйвер для аналитической СУБД ClickHouse с полной поддержкой MergeTree, словарей, партиций и SDUI-интроспекции.", "type": "database_driver", "main": "bin/clickhouse_rpc_server", "icon": "assets/icon.svg", "engines": { - "querya_desktop": "^2.0.0" + "querya_desktop": ">=0.4.0 <3.0.0" }, "capabilities": { "databaseDriver": true, @@ -41,6 +41,28 @@ "defaultPort": 8123, "connectionFormSchema": "assets/connection_form.json" } + ], + "commands": [ + { + "id": "clickhouse.optimizeFinal", + "title": "ClickHouse: Optimize Table (FINAL)", + "category": "ClickHouse" + }, + { + "id": "clickhouse.deduplicate", + "title": "ClickHouse: Deduplicate Table", + "category": "ClickHouse" + }, + { + "id": "clickhouse.serverStats", + "title": "ClickHouse: Show Server Statistics", + "category": "ClickHouse" + }, + { + "id": "clickhouse.dropPartition", + "title": "ClickHouse: Drop Selected Partition", + "category": "ClickHouse" + } ] } } diff --git a/scripts/package_qext.sh b/scripts/package_qext.sh index 37461d5..253bab2 100755 --- a/scripts/package_qext.sh +++ b/scripts/package_qext.sh @@ -64,11 +64,29 @@ mkdir -p "${STAGING_DIR}/bin" "${STAGING_DIR}/assets" cp manifest.json "${STAGING_DIR}/" cp -r assets/* "${STAGING_DIR}/assets/" +# For Windows targets, adjust manifest.json main entry point to include .exe extension +if [[ "$TARGET" == *"windows"* ]] || [[ "$DEST_BIN_NAME" == *".exe" ]]; then + python3 -c " +import json, sys +manifest_path = sys.argv[1] +dest_bin = sys.argv[2] +with open(manifest_path, 'r', encoding='utf-8') as f: + manifest = json.load(f) +manifest['main'] = f'bin/{dest_bin}' +with open(manifest_path, 'w', encoding='utf-8') as f: + json.dump(manifest, f, indent=2, ensure_ascii=False) +" "${STAGING_DIR}/manifest.json" "${DEST_BIN_NAME}" +fi + # 2. Copy binary into bin/ under both the manifest main entry and original name cp "${BIN_PATH}" "${STAGING_DIR}/bin/${DEST_BIN_NAME}" if [ "${DEST_BIN_NAME}" != "${BIN_NAME}" ] && [ ! -f "${STAGING_DIR}/bin/${BIN_NAME}" ]; then cp "${BIN_PATH}" "${STAGING_DIR}/bin/${BIN_NAME}" fi +# On Windows packages, also provide the extensionless binary name for backward compatibility +if [[ "$DEST_BIN_NAME" == *".exe" ]] && [ ! -f "${STAGING_DIR}/bin/clickhouse_rpc_server" ]; then + cp "${BIN_PATH}" "${STAGING_DIR}/bin/clickhouse_rpc_server" 2>/dev/null || true +fi chmod +x "${STAGING_DIR}/bin/${DEST_BIN_NAME}" if [ -f "${STAGING_DIR}/bin/${BIN_NAME}" ]; then chmod +x "${STAGING_DIR}/bin/${BIN_NAME}" diff --git a/src/driver/client.rs b/src/driver/client.rs index ce605cf..eeaa1cc 100644 --- a/src/driver/client.rs +++ b/src/driver/client.rs @@ -14,10 +14,13 @@ pub struct ConnectParams { pub connection_string: Option, pub host: Option, pub port: Option, + #[serde(alias = "username")] pub user: Option, pub database: Option, #[serde(alias = "safe_mode", alias = "safeMode")] pub readonly: Option, + #[serde(alias = "sslMode")] + pub ssl_mode: Option, } #[derive(Debug)] @@ -44,8 +47,17 @@ impl ClickHouseClient { } else { let parsed = Url::parse(&cs)?; let host = parsed.host_str().unwrap_or("localhost"); - let port = parsed.port().unwrap_or(8123); let scheme = parsed.scheme(); + let port_str = match parsed.port() { + Some(p) => format!(":{}", p), + None => { + if scheme == "http" { + ":8123".to_string() + } else { + String::new() + } + } + }; let user = if !parsed.username().is_empty() { parsed.username().to_string() } else { @@ -58,17 +70,24 @@ impl ClickHouseClient { params.database.unwrap_or_else(|| "default".to_string()) }; let readonly = params.readonly.unwrap_or(false); - let base = format!("{}://{}:{}", scheme, host, port); + let base = format!("{}://{}{}", scheme, host, port_str); (base, user, database, readonly) } } else { let host = params.host.unwrap_or_else(|| "localhost".to_string()); - let port = params.port.unwrap_or(8123); + let scheme = match params.ssl_mode.as_deref() { + Some("prefer") | Some("require") => "https", + _ => "http", + }; + let mut port = params.port.unwrap_or(8123); + if scheme == "https" && port == 8123 { + port = 8443; + } let user = params.user.unwrap_or_else(|| "default".to_string()); let database = params.database.unwrap_or_else(|| "default".to_string()); let readonly = params.readonly.unwrap_or(false); ( - format!("http://{}:{}", host, port), + format!("{}://{}:{}", scheme, host, port), user, database, readonly, @@ -146,14 +165,15 @@ impl ClickHouseClient { req } + async fn error_from_response(resp: reqwest::Response) -> DriverError { + let status = resp.status(); + let text = resp.text().await.unwrap_or_default(); + DriverError::Client(format!("ClickHouse HTTP error {}: {}", status, text)) + } + async fn read_response(resp: reqwest::Response) -> Result { if !resp.status().is_success() { - let status = resp.status(); - let text = resp.text().await.unwrap_or_default(); - return Err(DriverError::Client(format!( - "ClickHouse HTTP error {}: {}", - status, text - ))); + return Err(Self::error_from_response(resp).await); } Ok(resp.text().await?) } @@ -229,6 +249,41 @@ impl ClickHouseClient { } } + /// Like `post_sql`, but on success returns the raw streaming `reqwest::Response` + /// instead of buffering the whole body into a `String`, so a large analytical + /// result set can be parsed incrementally as it arrives over the network via + /// `crate::driver::streaming::stream_compact_output` (issue #49). Callers must + /// gate on mock/test connections themselves, same as before calling `post_sql`. + pub async fn post_sql_response( + &self, + sql: &str, + mut extra_params: impl FnMut(&mut Url), + ) -> Result { + let sql = sql.to_string(); + let mut omit = self.omit_readonly_setting(); + loop { + let mut url = Url::parse(&self.base_url)?; + url.query_pairs_mut() + .append_pair("database", &self.database); + extra_params(&mut url); + self.append_safe_mode_settings_with(&mut url, omit); + let req = self + .apply_auth(self.http_client.post(url)) + .body(sql.clone()); + let resp = req.send().await?; + if resp.status().is_success() { + return Ok(resp); + } + let err = Self::error_from_response(resp).await; + if self.readonly && !omit && Self::is_readonly_setting_conflict(&err) { + self.mark_server_readonly_enforced(); + omit = true; + continue; + } + return Err(err); + } + } + /// Check connection health by executing `SELECT version()` against ClickHouse. pub async fn ping_connection(&self) -> Result { let version = self.get_with_query("SELECT version()").await?; diff --git a/src/driver/mod.rs b/src/driver/mod.rs index f691764..a5b9f42 100644 --- a/src/driver/mod.rs +++ b/src/driver/mod.rs @@ -1,3 +1,4 @@ //! ClickHouse HTTP client and session management. pub mod client; pub mod pool; +pub mod streaming; diff --git a/src/driver/streaming.rs b/src/driver/streaming.rs new file mode 100644 index 0000000..58fe3c1 --- /dev/null +++ b/src/driver/streaming.rs @@ -0,0 +1,179 @@ +//! Incremental parsing of ClickHouse's `FORMAT JSONCompactEachRowWithNamesAndTypes` +//! HTTP response, row by row, as bytes arrive over the network. +//! +//! `db.query`'s previous implementation buffered the entire HTTP response body +//! into a single `String` (`reqwest::Response::text()`) before parsing began, +//! so a large analytical result (tens of megabytes of JSON lines) held both +//! the raw text and the parsed rows in memory simultaneously — a real OOM risk +//! under the 256 MB ClickHouse Sandbox ceiling. This module instead reads the +//! response as a line stream and stops pulling further bytes off the network +//! connection as soon as either the caller's `limit` or a safety byte cap is +//! reached (see issue #49). + +use crate::error::DriverError; +use crate::mapper::row_compact::{ + QueryResult, QueryStatistics, parse_and_normalize_row, parse_columns, +}; +use bytes::Bytes; +use futures::Stream; +use futures::StreamExt; +use tokio_util::codec::{FramedRead, LinesCodec}; +use tokio_util::io::StreamReader; + +/// Independent of any client-supplied `limit`, stop consuming further row +/// bytes once this many have been read, so an unbounded query (no `limit` +/// sent) still can't grow the result past a safe ceiling. +const MAX_ROW_BYTES: usize = 200 * 1024 * 1024; + +/// ClickHouse lines (JSON arrays of column values) are expected to be well +/// under this; it exists only to bound a single corrupt/adversarial line's +/// buffered length instead of growing unbounded. +const MAX_LINE_BYTES: usize = 64 * 1024 * 1024; + +/// Reads a byte stream (a ClickHouse HTTP response body) as a stream of +/// lines and incrementally parses it as `FORMAT JSONCompactEachRowWithNamesAndTypes` +/// output, stopping (without reading the rest of the stream) once `limit` +/// rows have been parsed or the `MAX_ROW_BYTES` safety cap is reached. +pub async fn stream_compact_output( + byte_stream: S, + limit: Option, +) -> Result +where + S: Stream> + Unpin, +{ + let stream_reader = StreamReader::new(byte_stream); + let mut lines = FramedRead::new( + stream_reader, + LinesCodec::new_with_max_length(MAX_LINE_BYTES), + ); + + let mut bytes_read: usize = 0; + + let Some(names_line) = lines.next().await.transpose().map_err(line_err)? else { + return Ok(QueryResult { + columns: vec![], + rows: vec![], + statistics: QueryStatistics { + rows_read: 0, + bytes_read: 0, + elapsed_ms: 0, + }, + query_id: None, + is_truncated: false, + }); + }; + bytes_read += names_line.len() + 1; + + let Some(types_line) = lines.next().await.transpose().map_err(line_err)? else { + return Err(DriverError::Client( + "Malformed JSONCompactEachRowWithNamesAndTypes output: missing names or types row" + .to_string(), + )); + }; + bytes_read += types_line.len() + 1; + + let columns = parse_columns(&names_line, &types_line)?; + + let mut rows = Vec::with_capacity(limit.unwrap_or(16).min(1024)); + let mut is_truncated = false; + while let Some(line) = lines.next().await.transpose().map_err(line_err)? { + bytes_read += line.len() + 1; + + if limit.is_some_and(|limit| rows.len() >= limit) || bytes_read > MAX_ROW_BYTES { + is_truncated = true; + break; + } + + rows.push(parse_and_normalize_row(&line, &columns)?); + } + // Drop the frame reader (and the underlying HTTP connection) now instead of + // reading any remaining body bytes off the network when we stopped early. + drop(lines); + + let rows_read = rows.len(); + Ok(QueryResult { + columns, + rows, + statistics: QueryStatistics { + rows_read, + bytes_read, + elapsed_ms: 0, + }, + query_id: None, + is_truncated, + }) +} + +fn line_err(e: tokio_util::codec::LinesCodecError) -> DriverError { + DriverError::Client(format!("Failed to read ClickHouse response stream: {}", e)) +} + +#[cfg(test)] +mod tests { + use super::*; + use futures::stream; + + /// Builds an in-memory byte stream out of a static string, split into + /// arbitrary chunks, to exercise `stream_compact_output` without a live + /// HTTP connection. Splitting mid-line (not just at newlines) verifies + /// the line codec correctly reassembles frames split across chunks. + fn chunked_stream( + body: &'static str, + chunk_size: usize, + ) -> impl Stream> { + let chunks: Vec> = body + .as_bytes() + .chunks(chunk_size.max(1)) + .map(|c| Ok(Bytes::copy_from_slice(c))) + .collect(); + stream::iter(chunks) + } + + #[tokio::test] + async fn test_stream_compact_output_parses_rows() { + let body = "[\"id\"]\n[\"UInt64\"]\n[1]\n[2]\n[3]\n"; + let result = stream_compact_output(chunked_stream(body, 1024), None) + .await + .unwrap(); + assert_eq!(result.rows.len(), 3); + assert!(!result.is_truncated); + } + + #[tokio::test] + async fn test_stream_compact_output_reassembles_lines_split_across_chunks() { + let body = "[\"id\"]\n[\"UInt64\"]\n[1]\n[2]\n[3]\n"; + // Force byte-at-a-time delivery so every line boundary falls mid-chunk. + let result = stream_compact_output(chunked_stream(body, 1), None) + .await + .unwrap(); + assert_eq!(result.rows.len(), 3); + } + + #[tokio::test] + async fn test_stream_compact_output_enforces_limit() { + let body = "[\"id\"]\n[\"UInt64\"]\n[1]\n[2]\n[3]\n"; + let result = stream_compact_output(chunked_stream(body, 1024), Some(2)) + .await + .unwrap(); + assert_eq!(result.rows.len(), 2); + assert!(result.is_truncated); + } + + #[tokio::test] + async fn test_stream_compact_output_empty_body() { + let result = stream_compact_output(chunked_stream("", 1024), None) + .await + .unwrap(); + assert!(result.rows.is_empty()); + assert!(result.columns.is_empty()); + } + + #[tokio::test] + async fn test_stream_compact_output_missing_types_row() { + let body = "[\"id\"]\n"; + let err = stream_compact_output(chunked_stream(body, 1024), None) + .await + .unwrap_err(); + assert!(err.to_string().contains("missing names or types row")); + } +} diff --git a/src/mapper/row_compact.rs b/src/mapper/row_compact.rs index 99012c0..c134da0 100644 --- a/src/mapper/row_compact.rs +++ b/src/mapper/row_compact.rs @@ -17,6 +17,90 @@ pub struct QueryResult { pub columns: Vec, pub rows: Vec>, pub statistics: QueryStatistics, + #[serde(skip_serializing_if = "Option::is_none")] + pub query_id: Option, + /// `true` when `limit` (or the streaming safety byte cap) cut the result + /// short of what ClickHouse actually had to offer, so the caller knows + /// `rows` isn't the complete result set. + #[serde(skip_serializing_if = "std::ops::Not::not")] + pub is_truncated: bool, +} + +/// Parses the `["name", ...]` / `["Type", ...]` header lines of +/// `FORMAT JSONCompactEachRowWithNamesAndTypes` output into `ColumnSchema`s. +/// Shared by both the buffered (`parse_compact_output`) and streaming +/// (`crate::driver::streaming::stream_compact_output`) parsers. +pub(crate) fn parse_columns( + names_line: &str, + types_line: &str, +) -> Result, DriverError> { + let names: Vec = serde_json::from_str(names_line).map_err(|e| { + DriverError::Client(format!( + "Failed to parse column names from ClickHouse output: {}", + e + )) + })?; + let types: Vec = serde_json::from_str(types_line).map_err(|e| { + DriverError::Client(format!( + "Failed to parse column types from ClickHouse output: {}", + e + )) + })?; + + if names.len() != types.len() { + return Err(DriverError::Client(format!( + "Column names count ({}) does not match types count ({})", + names.len(), + types.len() + ))); + } + + Ok(names + .into_iter() + .zip(types) + .map(|(name, ch_type)| ColumnSchema::new(name, ch_type)) + .collect()) +} + +/// Parses a single `[value, value, ...]` data row line and normalizes it +/// according to each column's `mapped_type` (e.g. converting 64-bit numbers +/// and Decimals into JSON strings to prevent 53-bit float overflow in +/// JS/Flutter). Shared by both the buffered and streaming parsers. +pub(crate) fn parse_and_normalize_row( + line: &str, + columns: &[ColumnSchema], +) -> Result, DriverError> { + let mut raw_row: Vec = serde_json::from_str(line).map_err(|e| { + DriverError::Client(format!( + "Failed to parse data row JSON array from ClickHouse: {}", + e + )) + })?; + + for (i, col) in columns.iter().enumerate() { + if let Some(val) = raw_row.get_mut(i) { + if val.is_null() { + continue; + } + match col.mapped_type { + "string" => { + // 64-bit/large integers and Decimals may arrive as JSON numbers from ClickHouse + if val.is_number() { + *val = Value::String(val.to_string()); + } + } + "integer" => { + if let Some(s) = val.as_str() + && let Ok(n) = s.parse::() + { + *val = Value::Number(serde_json::Number::from(n)); + } + } + _ => {} + } + } + } + Ok(raw_row) } /// Parses the output of ClickHouse `FORMAT JSONCompactEachRowWithNamesAndTypes`. @@ -24,17 +108,22 @@ pub struct QueryResult { /// Line 2: JSON array of ClickHouse data types `["UInt64", "Decimal(18, 4)"]` /// Lines 3+: JSON arrays of row values `[18446744073709551615, 123.4500]` /// Automatically normalizes row values according to `ColumnSchema::mapped_type` (e.g. converting 64-bit numbers and Decimals into JSON strings to prevent 53-bit float overflow in JS/Flutter). +/// +/// `limit`, when set, stops parsing (and allocating) further data rows once +/// that many have been read, instead of parsing the entire result set and +/// discarding the excess — this bounds peak memory for large result sets +/// (see issue #47). pub fn parse_compact_output( output_text: &str, elapsed_ms: u64, + limit: Option, ) -> Result { - let lines: Vec<&str> = output_text + let mut lines = output_text .lines() .map(|l| l.trim()) - .filter(|l| !l.is_empty()) - .collect(); + .filter(|l| !l.is_empty()); - if lines.is_empty() { + let Some(names_line) = lines.next() else { return Ok(QueryResult { columns: vec![], rows: vec![], @@ -43,76 +132,28 @@ pub fn parse_compact_output( bytes_read: output_text.len(), elapsed_ms, }, + query_id: None, + is_truncated: false, }); - } + }; - if lines.len() < 2 { + let Some(types_line) = lines.next() else { return Err(DriverError::Client( "Malformed JSONCompactEachRowWithNamesAndTypes output: missing names or types row" .to_string(), )); - } + }; - let names: Vec = serde_json::from_str(lines[0]).map_err(|e| { - DriverError::Client(format!( - "Failed to parse column names from ClickHouse output: {}", - e - )) - })?; - let types: Vec = serde_json::from_str(lines[1]).map_err(|e| { - DriverError::Client(format!( - "Failed to parse column types from ClickHouse output: {}", - e - )) - })?; - - if names.len() != types.len() { - return Err(DriverError::Client(format!( - "Column names count ({}) does not match types count ({})", - names.len(), - types.len() - ))); - } - - let mut columns = Vec::with_capacity(names.len()); - for (name, ch_type) in names.into_iter().zip(types) { - columns.push(ColumnSchema::new(name, ch_type)); - } + let columns = parse_columns(names_line, types_line)?; - let mut rows = Vec::with_capacity(lines.len().saturating_sub(2)); - for line in &lines[2..] { - let mut raw_row: Vec = serde_json::from_str(line).map_err(|e| { - DriverError::Client(format!( - "Failed to parse data row JSON array from ClickHouse: {}", - e - )) - })?; - - // Normalize values according to Querya schema mapped_type - for (i, col) in columns.iter().enumerate() { - if let Some(val) = raw_row.get_mut(i) { - if val.is_null() { - continue; - } - match col.mapped_type { - "string" => { - // 64-bit/large integers and Decimals may arrive as JSON numbers from ClickHouse - if val.is_number() { - *val = Value::String(val.to_string()); - } - } - "integer" => { - if let Some(s) = val.as_str() - && let Ok(n) = s.parse::() - { - *val = Value::Number(serde_json::Number::from(n)); - } - } - _ => {} - } - } + let mut rows = Vec::with_capacity(limit.unwrap_or(16).min(1024)); + let mut is_truncated = false; + for line in lines { + if limit.is_some_and(|limit| rows.len() >= limit) { + is_truncated = true; + break; } - rows.push(raw_row); + rows.push(parse_and_normalize_row(line, &columns)?); } let rows_read = rows.len(); @@ -124,6 +165,8 @@ pub fn parse_compact_output( bytes_read: output_text.len(), elapsed_ms, }, + query_id: None, + is_truncated, }) } @@ -139,7 +182,7 @@ mod tests { [18446744073709551615, "Alice", 1234567.8901, true] [102, null, 0.0000, false]"#; - let res = parse_compact_output(raw_output, 15).unwrap(); + let res = parse_compact_output(raw_output, 15, None).unwrap(); assert_eq!(res.columns.len(), 4); assert_eq!(res.columns[0].mapped_type, "string"); assert_eq!(res.columns[1].mapped_type, "string"); @@ -159,21 +202,96 @@ mod tests { assert_eq!(res.statistics.elapsed_ms, 15); } + #[test] + fn test_parse_compact_output_scalar_row() { + // Regression for issue #57: scalar introspection queries + // (SELECT version(), SELECT uptime(), SHOW CREATE TABLE) must be + // requested with FORMAT JSONCompactEachRowWithNamesAndTypes so this + // parser's fixed 2-header-line assumption holds; a bare + // JSONCompactEachRow response has no names/types lines and would + // previously fail with "missing names or types row". + let version_output = r#"["version()"] +["String"] +["24.3.1.2452"]"#; + let res = parse_compact_output(version_output, 0, None).unwrap(); + assert_eq!(res.rows.len(), 1); + assert_eq!(res.rows[0][0], json!("24.3.1.2452")); + + let uptime_output = r#"["uptime()"] +["UInt32"] +[123456]"#; + let res = parse_compact_output(uptime_output, 0, None).unwrap(); + assert_eq!(res.rows[0][0].as_u64(), Some(123456)); + } + #[test] fn test_parse_compact_output_empty() { - let res = parse_compact_output("", 5).unwrap(); + let res = parse_compact_output("", 5, None).unwrap(); assert!(res.columns.is_empty()); assert!(res.rows.is_empty()); assert_eq!(res.statistics.rows_read, 0); } + #[test] + fn test_parse_compact_output_enforces_limit() { + // Regression for issue #47: `limit` must stop row parsing early instead + // of parsing the whole result set and discarding the excess, since a + // multi-GB result would otherwise be fully materialized in memory first. + let raw_output = r#"["id"] +["UInt64"] +[1] +[2] +[3] +[4] +[5]"#; + + let unlimited = parse_compact_output(raw_output, 0, None).unwrap(); + assert_eq!(unlimited.rows.len(), 5); + assert_eq!(unlimited.statistics.rows_read, 5); + assert!(!unlimited.is_truncated); + + let limited = parse_compact_output(raw_output, 0, Some(2)).unwrap(); + assert_eq!(limited.rows.len(), 2); + // UInt64 is normalized to a JSON string to protect 53-bit JS precision. + assert_eq!(limited.rows[0][0], json!("1")); + assert_eq!(limited.rows[1][0], json!("2")); + assert_eq!(limited.statistics.rows_read, 2); + assert!(limited.is_truncated); + + // A limit larger than the actual row count is a no-op, and isn't + // reported as truncated since nothing was actually cut off. + let generous_limit = parse_compact_output(raw_output, 0, Some(100)).unwrap(); + assert_eq!(generous_limit.rows.len(), 5); + assert!(!generous_limit.is_truncated); + + // A zero limit returns no rows at all, without erroring. + let zero_limit = parse_compact_output(raw_output, 0, Some(0)).unwrap(); + assert!(zero_limit.rows.is_empty()); + } + + #[test] + fn test_parse_compact_output_sets_query_id_field() { + // Regression for issue #47: query_id lives directly on QueryResult so + // callers don't need a second serde_json::to_value pass just to splice + // a queryId key into the already-serialized response. + let raw_output = r#"["id"] +["UInt64"] +[1]"#; + let mut res = parse_compact_output(raw_output, 0, None).unwrap(); + assert_eq!(res.query_id, None); + res.query_id = Some("querya-job-1-2-3".to_string()); + + let serialized = serde_json::to_value(&res).unwrap(); + assert_eq!(serialized["queryId"], json!("querya-job-1-2-3")); + } + #[test] fn test_parse_compact_output_complex_types() { let raw_output = r#"["arr", "tup", "dt", "big_arr"] ["Array(Int32)", "Tuple(Int32, String)", "DateTime64(3)", "Array(UInt64)"] [[10, 20, 30], [100, "foo"], "2026-07-11 12:34:56.789", [18446744073709551615, 42]]"#; - let res = parse_compact_output(raw_output, 8).unwrap(); + let res = parse_compact_output(raw_output, 8, None).unwrap(); assert_eq!(res.columns.len(), 4); assert_eq!(res.columns[0].mapped_type, "array"); assert_eq!(res.columns[1].mapped_type, "json"); diff --git a/src/rpc/handlers/commands.rs b/src/rpc/handlers/commands.rs new file mode 100644 index 0000000..34f65ce --- /dev/null +++ b/src/rpc/handlers/commands.rs @@ -0,0 +1,284 @@ +use crate::error::DriverError; +use crate::rpc::handlers::{query, schema}; +use crate::utils::sql_escape::{escape_sql_string_literal, quote_identifier}; +use serde::Deserialize; +use serde_json::{Value, json}; +use tracing::info; + +/// Parameters for `commands.execute`, sent by Querya Desktop's Command Palette +/// when the user runs one of this driver's `contributions.commands` entries. +/// +/// The Command Palette does not yet forward workspace selection (selected +/// table/partition) to the driver, so `database`/`table`/`partition` are +/// optional here and validated per-command: a command that needs a target +/// returns a clear `-32602` error explaining what's missing instead of +/// silently operating on the wrong object. +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ExecuteCommandParams { + pub command_id: String, + pub connection_id: u64, + #[serde(default)] + pub database: Option, + #[serde(default)] + pub table: Option, + #[serde(default)] + pub partition: Option, +} + +/// Handler for `commands.execute`. Dispatches a Command Palette invocation +/// (`commandId`) to the matching ClickHouse maintenance operation. +pub async fn handle_execute(params: Option) -> Result { + let params_val = params.ok_or_else(|| DriverError::Rpc { + code: -32602, + message: "Invalid params: commands.execute requires commandId and connectionId".to_string(), + data: None, + })?; + + let p: ExecuteCommandParams = + serde_json::from_value(params_val).map_err(|e| DriverError::Rpc { + code: -32602, + message: format!("Malformed commands.execute parameters: {}", e), + data: None, + })?; + + info!( + "Executing command '{}' on connectionId={}", + p.command_id, p.connection_id + ); + + match p.command_id.as_str() { + "clickhouse.serverStats" => { + schema::handle_get_server_stats(Some(json!({ "connectionId": p.connection_id }))).await + } + "clickhouse.optimizeFinal" => { + let (db, tbl) = require_table_target(&p)?; + run_sql( + p.connection_id, + format!( + "OPTIMIZE TABLE {}.{} FINAL", + quote_identifier(db), + quote_identifier(tbl) + ), + ) + .await + } + "clickhouse.deduplicate" => { + let (db, tbl) = require_table_target(&p)?; + run_sql( + p.connection_id, + format!( + "OPTIMIZE TABLE {}.{} DEDUPLICATE", + quote_identifier(db), + quote_identifier(tbl) + ), + ) + .await + } + "clickhouse.dropPartition" => { + let (db, tbl) = require_table_target(&p)?; + let partition = p + .partition + .as_deref() + .filter(|s| !s.is_empty()) + .ok_or_else(|| DriverError::Rpc { + code: -32602, + message: + "commandId 'clickhouse.dropPartition' also requires a 'partition' parameter" + .to_string(), + data: None, + })?; + run_sql( + p.connection_id, + format!( + "ALTER TABLE {}.{} DROP PARTITION '{}'", + quote_identifier(db), + quote_identifier(tbl), + escape_sql_string_literal(partition) + ), + ) + .await + } + other => Err(DriverError::Rpc { + code: -32602, + message: format!("Unknown commandId: '{}'", other), + data: None, + }), + } +} + +fn require_table_target(p: &ExecuteCommandParams) -> Result<(&str, &str), DriverError> { + match ( + p.database.as_deref().filter(|s| !s.is_empty()), + p.table.as_deref().filter(|s| !s.is_empty()), + ) { + (Some(db), Some(tbl)) => Ok((db, tbl)), + _ => Err(DriverError::Rpc { + code: -32602, + message: format!( + "commandId '{}' requires 'database' and 'table' parameters (select a table first)", + p.command_id + ), + data: None, + }), + } +} + +async fn run_sql(connection_id: u64, sql: String) -> Result { + query::handle_query(Some(json!({ + "connectionId": connection_id, + "sql": sql + }))) + .await +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::driver::client::{ClickHouseClient, ConnectParams}; + use crate::driver::pool::ConnectionPool; + + #[tokio::test] + async fn test_handle_execute_server_stats() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 501, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let res = handle_execute(Some(json!({ + "commandId": "clickhouse.serverStats", + "connectionId": 501 + }))) + .await + .unwrap(); + assert_eq!(res["serverVersion"], "ClickHouse 24.3 (Mock)"); + + ConnectionPool::global().remove(501); + } + + #[tokio::test] + async fn test_handle_execute_optimize_final_and_deduplicate() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 502, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let res = handle_execute(Some(json!({ + "commandId": "clickhouse.optimizeFinal", + "connectionId": 502, + "database": "analytics", + "table": "events" + }))) + .await + .unwrap(); + assert_eq!(res["operation"], "optimize"); + + let res = handle_execute(Some(json!({ + "commandId": "clickhouse.deduplicate", + "connectionId": 502, + "database": "analytics", + "table": "events" + }))) + .await + .unwrap(); + assert_eq!(res["operation"], "optimize"); + + ConnectionPool::global().remove(502); + } + + #[tokio::test] + async fn test_handle_execute_drop_partition() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 503, + connection_string: Some("mock://localhost:8123/default".to_string()), + readonly: Some(false), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let res = handle_execute(Some(json!({ + "commandId": "clickhouse.dropPartition", + "connectionId": 503, + "database": "analytics", + "table": "events", + "partition": "202607" + }))) + .await + .unwrap(); + assert_eq!(res["operation"], "alter"); + + ConnectionPool::global().remove(503); + } + + #[tokio::test] + async fn test_handle_execute_requires_table_target() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 504, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let err = handle_execute(Some(json!({ + "commandId": "clickhouse.optimizeFinal", + "connectionId": 504 + }))) + .await + .unwrap_err(); + assert_eq!(err.to_rpc_code(), -32602); + assert!(err.to_string().contains("requires 'database' and 'table'")); + + let err = handle_execute(Some(json!({ + "commandId": "clickhouse.dropPartition", + "connectionId": 504, + "database": "analytics", + "table": "events" + }))) + .await + .unwrap_err(); + assert!(err.to_string().contains("also requires a 'partition'")); + + ConnectionPool::global().remove(504); + } + + #[tokio::test] + async fn test_handle_execute_unknown_command_id() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 505, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let err = handle_execute(Some(json!({ + "commandId": "clickhouse.doesNotExist", + "connectionId": 505 + }))) + .await + .unwrap_err(); + assert_eq!(err.to_rpc_code(), -32602); + assert!(err.to_string().contains("Unknown commandId")); + + ConnectionPool::global().remove(505); + } + + #[tokio::test] + async fn test_handle_execute_missing_params() { + let err = handle_execute(None).await.unwrap_err(); + assert_eq!(err.to_rpc_code(), -32602); + } +} diff --git a/src/rpc/handlers/mod.rs b/src/rpc/handlers/mod.rs index c5e083f..eeaf771 100644 --- a/src/rpc/handlers/mod.rs +++ b/src/rpc/handlers/mod.rs @@ -1,3 +1,4 @@ +pub mod commands; pub mod connection; pub mod query; pub mod schema; diff --git a/src/rpc/handlers/query.rs b/src/rpc/handlers/query.rs index a233daa..35a0373 100644 --- a/src/rpc/handlers/query.rs +++ b/src/rpc/handlers/query.rs @@ -2,6 +2,7 @@ use crate::driver::pool::ConnectionPool; use crate::error::DriverError; use crate::mapper::row_compact::parse_compact_output; use crate::utils::secret_guard::ConnectionSecretsPool; +use futures::StreamExt; use serde::Deserialize; use serde_json::{Value, json}; use std::sync::atomic::{AtomicU64, Ordering}; @@ -51,8 +52,182 @@ fn generate_query_id(connection_id: u64) -> String { format!("querya-job-{}-{}-{}", connection_id, now, seq) } -fn strip_sql_comments_and_trim(sql: &str) -> String { - let mut res = String::new(); +/// Validates that an identifier token (queryId or mutationId) contains only safe characters. +pub fn validate_query_or_mutation_id(id: &str, field_name: &str) -> Result<(), DriverError> { + if id.is_empty() || id.len() > 256 { + return Err(DriverError::Client(format!( + "Invalid {} length: must be between 1 and 256 characters", + field_name + ))); + } + if !id + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-' || c == '.') + { + return Err(DriverError::Client(format!( + "Invalid {} format: '{}' contains disallowed characters (allowed: [a-zA-Z0-9_.-])", + field_name, id + ))); + } + Ok(()) +} + +/// Zero-allocation, lazy SQL token scanner used by the Safe Mode precheck and +/// query classification. +/// +/// Yields whitespace-separated tokens as slices of the original text, skipping +/// `-- line` and `/* block */` comments. Quoted literals (`'...'`, `"..."`, +/// with `\\` and doubled-quote escapes) are kept inside a single token, so their +/// contents can never be mistaken for keywords or comment markers. Scanning +/// stops as soon as the caller stops pulling tokens. +struct SqlTokens<'a> { + src: &'a str, + pos: usize, +} + +impl<'a> SqlTokens<'a> { + fn new(src: &'a str) -> Self { + Self { src, pos: 0 } + } + + fn char_at(&self, i: usize) -> char { + self.src[i..].chars().next().unwrap_or('\0') + } +} + +impl<'a> Iterator for SqlTokens<'a> { + type Item = &'a str; + + fn next(&mut self) -> Option<&'a str> { + let b = self.src.as_bytes(); + let n = b.len(); + let mut i = self.pos; + + // Skip whitespace and comments before the token. + loop { + if i >= n { + self.pos = n; + return None; + } + if b[i] == b'-' && b.get(i + 1) == Some(&b'-') { + i = self.src[i..].find('\n').map_or(n, |off| i + off); + } else if b[i] == b'/' && b.get(i + 1) == Some(&b'*') { + i = self.src[i + 2..] + .find("*/") + .map_or(n, |off| i + 2 + off + 2); + } else { + let c = self.char_at(i); + if !c.is_whitespace() { + break; + } + i += c.len_utf8(); + } + } + + let start = i; + while i < n { + match b[i] { + b'-' if b.get(i + 1) == Some(&b'-') => break, + b'/' if b.get(i + 1) == Some(&b'*') => break, + q @ (b'\'' | b'"') => { + i += 1; + while i < n { + if b[i] == b'\\' { + i += 1; + if i < n { + i += self.char_at(i).len_utf8(); + } + } else if b[i] == q { + i += 1; + if b.get(i) == Some(&q) { + i += 1; + } else { + break; + } + } else { + i += 1; + } + } + } + _ => { + let c = self.char_at(i); + if c.is_whitespace() { + break; + } + i += c.len_utf8(); + } + } + } + + let i = i.min(n); + self.pos = i; + Some(&self.src[start..i]) + } +} + +/// Case-insensitive (ASCII) check that the first SQL token starts with `prefix`. +fn first_token_starts_with(sql: &str, prefix: &str) -> bool { + SqlTokens::new(sql).next().is_some_and(|t| { + t.get(..prefix.len()) + .is_some_and(|head| head.eq_ignore_ascii_case(prefix)) + }) +} + +const KW_DROP: u16 = 1 << 0; +const KW_TRUNCATE: u16 = 1 << 1; +const KW_DELETE: u16 = 1 << 2; +const KW_UPDATE: u16 = 1 << 3; +const KW_ALTER: u16 = 1 << 4; +const KW_TABLE: u16 = 1 << 5; +const KW_DATABASE: u16 = 1 << 6; +const KW_INSERT: u16 = 1 << 7; +const KW_INTO: u16 = 1 << 8; +const KW_CREATE: u16 = 1 << 9; +const KW_VALUES: u16 = 1 << 10; +const KW_SELECT: u16 = 1 << 11; +/// A mutating `ALTER TABLE` action keyword (parentheses ignored). +const KW_ALTER_ACTION: u16 = 1 << 12; + +/// Single pass over all tokens, recording which dangerous keywords appear. +fn scan_keywords(sql: &str) -> u16 { + const EXACT: [(&str, u16); 12] = [ + ("DROP", KW_DROP), + ("TRUNCATE", KW_TRUNCATE), + ("DELETE", KW_DELETE), + ("UPDATE", KW_UPDATE), + ("ALTER", KW_ALTER), + ("TABLE", KW_TABLE), + ("DATABASE", KW_DATABASE), + ("INSERT", KW_INSERT), + ("INTO", KW_INTO), + ("CREATE", KW_CREATE), + ("VALUES", KW_VALUES), + ("SELECT", KW_SELECT), + ]; + const ALTER_ACTIONS: [&str; 9] = [ + "DROP", "DELETE", "UPDATE", "MODIFY", "REPLACE", "CLEAR", "FREEZE", "ATTACH", "DETACH", + ]; + + let mut flags = 0u16; + for token in SqlTokens::new(sql) { + for (kw, bit) in EXACT { + if token.eq_ignore_ascii_case(kw) { + flags |= bit; + } + } + let clean = token.trim_matches(|c| c == '(' || c == ')'); + if ALTER_ACTIONS.iter().any(|a| clean.eq_ignore_ascii_case(a)) { + flags |= KW_ALTER_ACTION; + } + } + flags +} + +/// Splits SQL text into individual statements separated by semicolon (`;`), +/// taking care not to split inside single/double quotes, backticks, or comments. +pub fn split_sql_statements(sql: &str) -> Vec { + let mut statements = Vec::new(); + let mut current = String::new(); let mut chars = sql.chars().peekable(); let mut in_single_comment = false; let mut in_multi_comment = false; @@ -61,89 +236,133 @@ fn strip_sql_comments_and_trim(sql: &str) -> String { while let Some(c) = chars.next() { if in_single_comment { + current.push(c); if c == '\n' { in_single_comment = false; - res.push(' '); } } else if in_multi_comment { + current.push(c); if c == '*' && chars.peek() == Some(&'/') { - chars.next(); + current.push(chars.next().unwrap()); in_multi_comment = false; - res.push(' '); } } else if in_string { - res.push(c); - if c == string_quote { - in_string = false; + current.push(c); + if c == '\\' { + if let Some(next_c) = chars.next() { + current.push(next_c); + } + } else if c == string_quote { + if chars.peek() == Some(&string_quote) { + current.push(chars.next().unwrap()); + } else { + in_string = false; + } } } else if c == '-' && chars.peek() == Some(&'-') { - chars.next(); + current.push(c); + current.push(chars.next().unwrap()); in_single_comment = true; } else if c == '/' && chars.peek() == Some(&'*') { - chars.next(); + current.push(c); + current.push(chars.next().unwrap()); in_multi_comment = true; } else if c == '\'' || c == '`' || c == '"' { in_string = true; string_quote = c; - res.push(c); + current.push(c); + } else if c == ';' { + let trimmed = current.trim(); + if !trimmed.is_empty() { + statements.push(trimmed.to_string()); + } + current.clear(); } else { - res.push(c); + current.push(c); } } - res.trim().to_uppercase() + + let trimmed = current.trim(); + if !trimmed.is_empty() { + statements.push(trimmed.to_string()); + } + + statements } -/// Pre-checks AST/SQL syntax in Safe Mode (`readonly = true`) before network roundtrip. -fn enforce_safe_mode_precheck(sql: &str) -> Result<(), DriverError> { - let upper = strip_sql_comments_and_trim(sql); - let tokens: Vec<&str> = upper.split_whitespace().collect(); - if tokens.is_empty() { +/// Checks an individual SQL statement for dangerous/destructive operations in Safe Mode. +fn check_single_statement_for_safe_mode(statement_sql: &str) -> Result<(), DriverError> { + let mut tokens = SqlTokens::new(statement_sql); + let Some(first) = tokens.next() else { + return Ok(()); + }; + let first = first.trim_start_matches('('); + let is_first = |kw: &str| first.eq_ignore_ascii_case(kw); + + // Fast path: read-only leading commands need no further scanning. + const READ_ONLY: [&str; 8] = [ + "SELECT", "SHOW", "DESCRIBE", "DESC", "EXPLAIN", "EXISTS", "CHECK", "WITH", + ]; + if READ_ONLY.iter().any(|kw| is_first(kw)) { return Ok(()); } + if is_first("UPDATE") { + return Err(safe_mode_violation()); + } - let first = tokens[0]; - let second = tokens.get(1).copied().unwrap_or(""); - let third = tokens.get(2).copied().unwrap_or(""); + let second = tokens.next().unwrap_or("").trim_start_matches('('); + let third = tokens.next().unwrap_or("").trim_start_matches('('); + let second_is = |kw: &str| second.eq_ignore_ascii_case(kw); + let third_is = |kw: &str| third.eq_ignore_ascii_case(kw); + let is_object_kind = || { + second_is("DATABASE") || second_is("TABLE") || second_is("VIEW") || second_is("DICTIONARY") + }; - let is_dangerous = match first { - "DROP" => { - second == "DATABASE" || second == "TABLE" || second == "VIEW" || second == "DICTIONARY" - } - "TRUNCATE" => second == "TABLE", - "ALTER" => { - second == "TABLE" - && tokens.iter().any(|&t| { - t == "DROP" - || t == "DELETE" - || t == "UPDATE" - || t == "MODIFY" - || t == "REPLACE" - || t == "CLEAR" - || t == "FREEZE" - || t == "ATTACH" - || t == "DETACH" - }) - } - "INSERT" => second == "INTO" || third == "INTO", - "DELETE" => second == "FROM" || third == "FROM", - "UPDATE" => true, - "CREATE" => { - second == "DATABASE" || second == "TABLE" || second == "VIEW" || second == "DICTIONARY" - } - "RENAME" => second == "TABLE" || second == "DATABASE", - "ATTACH" | "DETACH" => second == "TABLE" || second == "PARTITION", - _ => { - upper.contains("DROP DATABASE") - || upper.contains("TRUNCATE TABLE") - || upper.contains("DROP TABLE") - || (upper.contains("ALTER TABLE") && upper.contains("DROP")) - } + let is_dangerous = if is_first("DROP") || is_first("CREATE") { + is_object_kind() + } else if is_first("TRUNCATE") { + second_is("TABLE") + } else if is_first("ALTER") { + second_is("TABLE") && scan_keywords(statement_sql) & KW_ALTER_ACTION != 0 + } else if is_first("INSERT") { + second_is("INTO") + || third_is("INTO") + || scan_keywords(statement_sql) & (KW_VALUES | KW_SELECT) != 0 + } else if is_first("DELETE") { + second_is("FROM") || third_is("FROM") + } else if is_first("RENAME") { + second_is("TABLE") || second_is("DATABASE") + } else if is_first("ATTACH") || is_first("DETACH") { + second_is("TABLE") || second_is("PARTITION") + } else { + let f = scan_keywords(statement_sql); + f & (KW_DROP | KW_TRUNCATE | KW_DELETE | KW_UPDATE) != 0 + || (f & KW_ALTER != 0 && f & KW_TABLE != 0) + || (f & KW_INSERT != 0 && f & KW_INTO != 0) + || (f & KW_CREATE != 0 && f & (KW_TABLE | KW_DATABASE) != 0) }; if is_dangerous { - return Err(DriverError::SafeModeViolation( - "Operation blocked by Safe Mode: write or destructive queries are forbidden in analytical read-only mode".to_string(), - )); + return Err(safe_mode_violation()); + } + Ok(()) +} + +fn safe_mode_violation() -> DriverError { + DriverError::SafeModeViolation( + "Operation blocked by Safe Mode: write or destructive queries are forbidden in analytical read-only mode".to_string(), + ) +} + +/// Pre-checks AST/SQL syntax in Safe Mode (`readonly = true`) before network roundtrip, +/// evaluating all statements in multi-statement queries. +fn enforce_safe_mode_precheck(sql: &str) -> Result<(), DriverError> { + let statements = split_sql_statements(sql); + if statements.is_empty() { + return check_single_statement_for_safe_mode(sql); + } + for stmt in &statements { + check_single_statement_for_safe_mode(stmt)?; } Ok(()) } @@ -176,22 +395,32 @@ pub async fn handle_query(params: Option) -> Result { let trimmed_sql = query_params.sql.trim(); let upper_sql = trimmed_sql.to_uppercase(); - let is_tabular_query = upper_sql.starts_with("SELECT") - || upper_sql.starts_with("SHOW") - || upper_sql.starts_with("DESCRIBE") - || upper_sql.starts_with("EXPLAIN"); - - let sql_to_run = if is_tabular_query && !upper_sql.contains("FORMAT ") { + // Classify on the comment-stripped statement so a leading `-- comment` or + // `/* comment */` doesn't hide the real starting keyword, and recognize + // `WITH ...` CTE queries as tabular too. + let is_tabular_query = ["SELECT", "SHOW", "DESCRIBE", "EXPLAIN", "WITH"] + .iter() + .any(|kw| first_token_starts_with(trimmed_sql, kw)); + let has_format_clause = SqlTokens::new(trimmed_sql).any(|t| t.eq_ignore_ascii_case("FORMAT")); + + let sql_to_run = if is_tabular_query && !has_format_clause { + // FORMAT must precede the statement-terminating `;` in ClickHouse's + // grammar, so strip any trailing semicolon before appending it. + let sql_no_trailing_semicolon = + trimmed_sql.trim_end_matches(|c: char| c == ';' || c.is_whitespace()); format!( "{}\nFORMAT JSONCompactEachRowWithNamesAndTypes", - trimmed_sql + sql_no_trailing_semicolon ) } else { trimmed_sql.to_string() }; let actual_query_id = match &query_params.query_id { - Some(qid) if !qid.is_empty() => qid.clone(), + Some(qid) if !qid.is_empty() => { + validate_query_or_mutation_id(qid, "queryId")?; + qid.clone() + } _ => generate_query_id(query_params.connection_id), }; @@ -212,14 +441,13 @@ pub async fn handle_query(params: Option) -> Result { ["UInt64", "String", "Nullable(UInt64)"] [18446744073709551615, "page_view", 42] [100, "click", null]"#; - let mut parsed_val = serde_json::to_value(parse_compact_output( + let mut result = parse_compact_output( mock_output, start_time.elapsed().as_millis() as u64, - )?)?; - if let Some(obj) = parsed_val.as_object_mut() { - obj.insert("queryId".to_string(), json!(actual_query_id)); - } - return Ok(parsed_val); + query_params.limit, + )?; + result.query_id = Some(actual_query_id); + return Ok(serde_json::to_value(result)?); } else { return Ok(build_non_tabular_result( &upper_sql, @@ -232,21 +460,33 @@ pub async fn handle_query(params: Option) -> Result { // 3. Real ClickHouse HTTP request let actual_query_id_for_url = actual_query_id.clone(); - let text = client - .post_sql(&sql_to_run, |url| { - url.query_pairs_mut() - .append_pair("query_id", &actual_query_id_for_url); - }) - .await?; - let elapsed = start_time.elapsed().as_millis() as u64; - if is_tabular_query { - let mut parsed_val = serde_json::to_value(parse_compact_output(&text, elapsed)?)?; - if let Some(obj) = parsed_val.as_object_mut() { - obj.insert("queryId".to_string(), json!(actual_query_id)); - } - Ok(parsed_val) + // Stream and parse the response row-by-row instead of buffering the + // whole body into a String first, bounding peak memory for large + // analytical result sets (issue #49). + let response = client + .post_sql_response(&sql_to_run, |url| { + url.query_pairs_mut() + .append_pair("query_id", &actual_query_id_for_url); + }) + .await?; + let byte_stream = response + .bytes_stream() + .map(|chunk| chunk.map_err(std::io::Error::other)); + let mut result = + crate::driver::streaming::stream_compact_output(byte_stream, query_params.limit) + .await?; + result.statistics.elapsed_ms = start_time.elapsed().as_millis() as u64; + result.query_id = Some(actual_query_id); + Ok(serde_json::to_value(result)?) } else { + let text = client + .post_sql(&sql_to_run, |url| { + url.query_pairs_mut() + .append_pair("query_id", &actual_query_id_for_url); + }) + .await?; + let elapsed = start_time.elapsed().as_millis() as u64; Ok(build_non_tabular_result( &upper_sql, elapsed, @@ -349,6 +589,8 @@ pub async fn handle_cancel(params: Option) -> Result .get(cancel_params.connection_id) .ok_or_else(|| DriverError::ConnectionNotFound(cancel_params.connection_id))?; + validate_query_or_mutation_id(&cancel_params.query_id, "queryId")?; + info!( "Cancelling queryId={} on connectionId={}", cancel_params.query_id, cancel_params.connection_id @@ -358,6 +600,8 @@ pub async fn handle_cancel(params: Option) -> Result return Ok(json!({ "ok": true })); } + let escaped_query_id = + crate::utils::sql_escape::escape_sql_string_literal(&cancel_params.query_id); let sync_kw = if cancel_params.sync { "SYNC" } else { "ASYNC" }; let mut url = Url::parse(&client.base_url)?; url.query_pairs_mut() @@ -366,7 +610,7 @@ pub async fn handle_cancel(params: Option) -> Result "query", &format!( "KILL QUERY WHERE query_id = '{}' {}", - cancel_params.query_id, sync_kw + escaped_query_id, sync_kw ), ); @@ -413,6 +657,8 @@ pub async fn handle_kill_mutation(params: Option) -> Result) -> Result) -> Result::new()); + } + + #[test] + fn test_split_sql_statements_handles_escaped_and_doubled_quotes() { + // Regression for issue #59: a backslash-escaped quote or a doubled quote + // inside a string literal must not be treated as the string's closing quote. + assert_eq!( + split_sql_statements("SELECT 'Customer\\'s notes'; SELECT 2"), + vec!["SELECT 'Customer\\'s notes'", "SELECT 2"] + ); + assert_eq!( + split_sql_statements("SELECT 'Don''t drop; table'; SELECT 2"), + vec!["SELECT 'Don''t drop; table'", "SELECT 2"] + ); + } + + #[test] + fn test_sql_tokens_handles_escaped_and_doubled_quotes() { + // Regression for issue #59: an escaped quote (`\'`) must not prematurely + // close a string literal and expose a following `--` as a real comment. + let escaped: Vec<&str> = + SqlTokens::new("SELECT 'Customer\\'s notes -- internal' FROM feedback").collect(); + assert_eq!(escaped.first(), Some(&"SELECT")); + assert_eq!(escaped.last(), Some(&"feedback")); + assert!(escaped.contains(&"FROM")); + + // A doubled quote (`''`, the SQL-standard escape) must not close the + // string either, so the literal's content never leaks out as keywords. + let doubled: Vec<&str> = SqlTokens::new("SELECT 'Don''t drop table' FROM logs").collect(); + assert_eq!(doubled.first(), Some(&"SELECT")); + assert_eq!(doubled.last(), Some(&"logs")); + assert!(!doubled.iter().any(|t| t.eq_ignore_ascii_case("DROP"))); + } + + #[test] + fn test_safe_mode_precheck_is_case_insensitive_and_comment_aware() { + assert!(enforce_safe_mode_precheck("-- hi\n/* x */ drop table t").is_err()); + assert!(enforce_safe_mode_precheck("Alter Table t Delete where 1").is_err()); + assert!(enforce_safe_mode_precheck("insert into t values (1)").is_err()); + assert!(enforce_safe_mode_precheck("select 'drop table t' -- drop table x").is_ok()); + assert!(enforce_safe_mode_precheck(" \n ").is_ok()); + } + + #[test] + fn test_sql_tokens_skips_comments_and_whitespace() { + let tokens: Vec<&str> = + SqlTokens::new(" -- lead\n/* block */ WITH/**/x AS (SELECT 1) -- tail").collect(); + assert_eq!(tokens, ["WITH", "x", "AS", "(SELECT", "1)"]); + assert!(SqlTokens::new("-- only comment").next().is_none()); + assert!(SqlTokens::new("/* unterminated").next().is_none()); + assert!(first_token_starts_with(" /* c */ select 1", "SELECT")); + assert!(!first_token_starts_with("(select 1)", "SELECT")); + } + + #[test] + fn test_safe_mode_multi_statement_bypass_prevention() { + // Multi-statement bypass attempts from Issue #54 + assert!( + enforce_safe_mode_precheck( + "SELECT 1; INSERT INTO telemetry VALUES ('compromised', now());" + ) + .is_err() + ); + assert!(enforce_safe_mode_precheck("SELECT 1; DROP TABLE events;").is_err()); + assert!(enforce_safe_mode_precheck("SELECT 1; TRUNCATE TABLE events;").is_err()); + assert!( + enforce_safe_mode_precheck("SELECT 1; ALTER TABLE events DROP COLUMN user_id;") + .is_err() + ); + assert!(enforce_safe_mode_precheck("SELECT 1; DELETE FROM events WHERE 1=1;").is_err()); + assert!(enforce_safe_mode_precheck("SELECT 1; UPDATE events SET id = 2;").is_err()); + assert!(enforce_safe_mode_precheck("SELECT 1; CREATE TABLE new_tbl (id Int32);").is_err()); + + // Benign multi-statement queries + assert!(enforce_safe_mode_precheck("SELECT 1; SELECT 2; SHOW TABLES;").is_ok()); + assert!(enforce_safe_mode_precheck("SELECT ';'; SELECT 'DROP TABLE in string';").is_ok()); + + // Regression for issue #59: an escaped or doubled quote inside a string + // literal must not desynchronize comment/string tracking for the rest of + // the query, which would otherwise falsely block a benign query or hide a + // dangerous statement behind a fake comment. + assert!( + enforce_safe_mode_precheck("SELECT 'Customer\\'s notes -- internal' FROM feedback") + .is_ok() + ); + assert!(enforce_safe_mode_precheck("SELECT 'Don''t drop table' FROM logs").is_ok()); + } + + #[tokio::test] + async fn test_handle_query_blocked_by_safe_mode_multi_statement() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 223, + connection_string: Some("mock://localhost:8123/default?readonly=1".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let query_params = json!({ + "connectionId": 223, + "sql": "SELECT 1; INSERT INTO telemetry VALUES ('compromised', now());" + }); + + let err = handle_query(Some(query_params)).await.unwrap_err(); + assert!(matches!(err, DriverError::SafeModeViolation(_))); + + ConnectionPool::global().remove(223); + } + #[tokio::test] async fn test_handle_query_in_mock_mode() { let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; @@ -520,6 +901,83 @@ mod tests { ConnectionPool::global().remove(111); } + #[tokio::test] + async fn test_handle_query_enforces_limit() { + // Regression for issue #47: a `limit` in the request must actually + // truncate the parsed rows instead of being silently ignored. + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 113, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let query_params = json!({ + "connectionId": 113, + "sql": "SELECT id, event_name, user_id FROM events", + "limit": 1 + }); + + let res = handle_query(Some(query_params)).await.unwrap(); + assert_eq!(res["rows"].as_array().unwrap().len(), 1); + assert_eq!(res["statistics"]["rowsRead"], 1); + + ConnectionPool::global().remove(113); + } + + #[tokio::test] + async fn test_handle_query_tabular_detection_edge_cases() { + // Regression for issue #58: leading comments, CTE `WITH` queries and a + // trailing `;` must all still be classified as tabular queries. + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 112, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + for sql in [ + "-- Top 10 users\nSELECT id, event_name, user_id FROM events", + "/* block comment */ SELECT id, event_name, user_id FROM events", + "WITH x AS (SELECT 1) SELECT id, event_name, user_id FROM events", + "SELECT id, event_name, user_id FROM events;", + "SELECT id, event_name, user_id FROM events; ", + ] { + let query_params = json!({ "connectionId": 112, "sql": sql }); + let res = handle_query(Some(query_params)).await.unwrap(); + assert_eq!( + res["columns"].as_array().unwrap().len(), + 3, + "expected tabular result for: {}", + sql + ); + assert_eq!(res["rows"].as_array().unwrap().len(), 2); + } + + ConnectionPool::global().remove(112); + } + + #[test] + fn test_tabular_query_trailing_semicolon_format_placement() { + // The FORMAT clause must be appended before any trailing `;`, never after. + let trimmed_sql = "SELECT 1;"; + assert!(first_token_starts_with(trimmed_sql, "SELECT")); + let sql_no_trailing_semicolon = + trimmed_sql.trim_end_matches(|c: char| c == ';' || c.is_whitespace()); + let sql_to_run = format!( + "{}\nFORMAT JSONCompactEachRowWithNamesAndTypes", + sql_no_trailing_semicolon + ); + assert_eq!( + sql_to_run, + "SELECT 1\nFORMAT JSONCompactEachRowWithNamesAndTypes" + ); + } + #[tokio::test] async fn test_handle_query_blocked_by_safe_mode() { let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; @@ -780,4 +1238,86 @@ mod tests { ConnectionPool::global().remove(778); } + + #[test] + fn test_validate_query_or_mutation_id() { + assert!(validate_query_or_mutation_id("q123", "queryId").is_ok()); + assert!(validate_query_or_mutation_id("mutation_123.txt", "mutationId").is_ok()); + assert!(validate_query_or_mutation_id("querya-job-1-12345-67", "queryId").is_ok()); + + // Empty + assert!(validate_query_or_mutation_id("", "queryId").is_err()); + // Too long (>256) + assert!(validate_query_or_mutation_id(&"a".repeat(257), "queryId").is_err()); + // Disallowed chars (SQL injection vectors) + assert!(validate_query_or_mutation_id("' OR 1=1 --", "queryId").is_err()); + assert!(validate_query_or_mutation_id("id; DROP TABLE x;", "queryId").is_err()); + assert!(validate_query_or_mutation_id("id`injection", "mutationId").is_err()); + assert!(validate_query_or_mutation_id("id with spaces", "queryId").is_err()); + assert!(validate_query_or_mutation_id("id\nnewline", "queryId").is_err()); + } + + #[tokio::test] + async fn test_handle_cancel_rejects_malicious_query_id() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 881, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let cancel_params = json!({ + "connectionId": 881, + "queryId": "' OR 1=1 --" + }); + let err = handle_cancel(Some(cancel_params)).await.unwrap_err(); + assert!(matches!(err, DriverError::Client(_))); + + ConnectionPool::global().remove(881); + } + + #[tokio::test] + async fn test_handle_kill_mutation_rejects_malicious_mutation_id() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 882, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let kill_params = json!({ + "connectionId": 882, + "mutationId": "' OR 1=1 --" + }); + let err = handle_kill_mutation(Some(kill_params)).await.unwrap_err(); + assert!(matches!(err, DriverError::Client(_))); + + ConnectionPool::global().remove(882); + } + + #[tokio::test] + async fn test_handle_query_rejects_malicious_custom_query_id() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 883, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let query_params = json!({ + "connectionId": 883, + "sql": "SELECT 1", + "queryId": "malicious'query" + }); + let err = handle_query(Some(query_params)).await.unwrap_err(); + assert!(matches!(err, DriverError::Client(_))); + + ConnectionPool::global().remove(883); + } } diff --git a/src/rpc/handlers/schema.rs b/src/rpc/handlers/schema.rs index 9457340..76d21e8 100644 --- a/src/rpc/handlers/schema.rs +++ b/src/rpc/handlers/schema.rs @@ -1,6 +1,7 @@ use crate::driver::pool::ConnectionPool; use crate::error::DriverError; use crate::sdui::tree::*; +use crate::utils::node_id::split_node_id; use serde::Deserialize; use serde_json::{Value, json}; use tracing::info; @@ -94,7 +95,7 @@ pub async fn handle_expand_tree_node(params: Option) -> Result = p.node_id.split('.').collect(); + let parts: Vec = split_node_id(&p.node_id); if parts.is_empty() { return Err(DriverError::Client(format!( "Invalid nodeId format: '{}'", @@ -102,19 +103,19 @@ pub async fn handle_expand_tree_node(params: Option) -> Result Groups (Tables, Views, Dictionaries) if prefix == "db" && parts.len() >= 2 { - let db_name = parts[1]; + let db_name = parts[1].as_str(); let groups = build_database_groups(db_name); return Ok(json!({ "nodes": groups })); } // 2. Expand Table or View -> Groups (Columns, Partitions) if (prefix == "table" || prefix == "view") && parts.len() >= 3 { - let db_name = parts[1]; - let table_name = parts[2]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); let groups = build_table_groups(db_name, table_name); return Ok(json!({ "nodes": groups })); } @@ -124,14 +125,14 @@ pub async fn handle_expand_tree_node(params: Option) -> Result Tables / Views / Dictionaries if prefix == "group" && parts.len() >= 3 { - let db_name = parts[1]; - let group_type = parts[2]; + let db_name = parts[1].as_str(); + let group_type = parts[2].as_str(); if group_type == "tables" || group_type == "views" { let filter_view = group_type == "views"; let sql = format!( "SELECT t.name AS name, t.engine AS engine, t.total_rows AS total_rows, formatReadableSize(t.total_bytes) AS size_readable, t.comment AS comment, multiIf(t.engine LIKE '%View%', 'view', t.engine LIKE '%Dictionary%', 'dictionary', 'table') AS object_type FROM system.tables t WHERE database = '{}' ORDER BY name FORMAT JSONCompactEachRowWithNamesAndTypes", - db_name + crate::utils::sql_escape::escape_sql_string_literal(db_name) ); let text = if is_mock { @@ -149,7 +150,7 @@ pub async fn handle_expand_tree_node(params: Option) -> Result) -> Result Columns if prefix == "group_cols" && parts.len() >= 3 { - let db_name = parts[1]; - let table_name = parts[2]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); let sql = format!( "SELECT name, type, comment FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position FORMAT JSONCompactEachRowWithNamesAndTypes", - db_name, table_name + crate::utils::sql_escape::escape_sql_string_literal(db_name), + crate::utils::sql_escape::escape_sql_string_literal(table_name) ); let text = if is_mock { @@ -191,11 +193,12 @@ pub async fn handle_expand_tree_node(params: Option) -> Result Partitions if prefix == "group_parts" && parts.len() >= 3 { - let db_name = parts[1]; - let table_name = parts[2]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); let sql = format!( "SELECT partition, sum(rows) AS total_rows, formatReadableSize(sum(data_compressed_bytes)) AS compressed_size, count() AS parts_count FROM system.parts WHERE database = '{}' AND table = '{}' AND active = 1 GROUP BY partition ORDER BY partition DESC FORMAT JSONCompactEachRowWithNamesAndTypes", - db_name, table_name + crate::utils::sql_escape::escape_sql_string_literal(db_name), + crate::utils::sql_escape::escape_sql_string_literal(table_name) ); let text = if is_mock { @@ -243,7 +246,7 @@ pub async fn handle_context_actions(params: Option) -> Result) -> Result= 4 { + let is_mock = + client.base_url.starts_with("mock://") || client.base_url.starts_with("test://"); + if is_mock { + None + } else { + let db_name = &parts[1]; + let table_name = &parts[2]; + let col_name = &parts[3]; + let sql = format!( + "SELECT type FROM system.columns WHERE database = '{}' AND table = '{}' AND name = '{}' LIMIT 1 FORMAT JSONCompactEachRowWithNamesAndTypes", + crate::utils::sql_escape::escape_sql_string_literal(db_name), + crate::utils::sql_escape::escape_sql_string_literal(table_name), + crate::utils::sql_escape::escape_sql_string_literal(col_name) + ); + let text = run_introspection_query(p.connection_id, &sql).await?; + crate::mapper::row_compact::parse_compact_output(&text, 0, None) + .ok() + .and_then(|result| result.rows.into_iter().next()) + .and_then(|row| row.into_iter().next()) + .and_then(|v| v.as_str().map(|s| s.to_string())) + } + } else { + None + } + } else { + None + }; + + let actions = crate::sdui::actions::get_context_actions_for_node( + &p.node_type, + &p.node_id, + column_type.as_deref(), + )?; Ok(json!({ "actions": actions })) } +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct GetServerStatsParams { + pub connection_id: u64, +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct GetObjectMetadataParams { + pub connection_id: u64, + pub node_id: String, + pub node_type: String, +} + +/// Handler for `db.getCapabilities`. Returns capability feature flags reported by the driver. +pub async fn handle_get_capabilities(_params: Option) -> Result { + Ok(json!({ + "supportsTransactions": false, + "supportsCancel": true, + "supportsDDLInspection": true, + "supportsPrivileges": false, + "hasServerStats": true + })) +} + +/// Handler for `db.getServerStats`. Returns server version, uptime, and database sizes. +pub async fn handle_get_server_stats(params: Option) -> Result { + let params_val = params.ok_or_else(|| DriverError::Rpc { + code: -32602, + message: "Invalid params: db.getServerStats requires connectionId".to_string(), + data: None, + })?; + + let p: GetServerStatsParams = + serde_json::from_value(params_val).map_err(|e| DriverError::Rpc { + code: -32602, + message: format!("Malformed getServerStats parameters: {}", e), + data: None, + })?; + + let client = ConnectionPool::global() + .get(p.connection_id) + .ok_or_else(|| DriverError::ConnectionNotFound(p.connection_id))?; + + if client.base_url.starts_with("mock://") || client.base_url.starts_with("test://") { + return Ok(json!({ + "serverVersion": "ClickHouse 24.3 (Mock)", + "uptimeSeconds": 3600, + "activeConnections": 5, + "activeQueries": 2, + "memoryUsageBytes": 134217728, + "databaseSizes": { + "default": 10485760, + "system": 2097152, + "analytics": 524288000 + }, + "extraMetrics": { + "queriesPerSecond": 14.5 + } + })); + } + + let version_text = client + .post_sql( + "SELECT version() FORMAT JSONCompactEachRowWithNamesAndTypes", + |_| {}, + ) + .await + .unwrap_or_else(|_| { + r#"["version()"] +["String"] +["unknown"]"# + .to_string() + }); + let uptime_text = client + .post_sql( + "SELECT uptime() FORMAT JSONCompactEachRowWithNamesAndTypes", + |_| {}, + ) + .await + .unwrap_or_else(|_| { + r#"["uptime()"] +["UInt32"] +[0]"# + .to_string() + }); + + let mut version_str = "ClickHouse".to_string(); + if let Ok(parsed) = crate::mapper::row_compact::parse_compact_output(&version_text, 0, None) + && let Some(row) = parsed.rows.first() + && let Some(v) = row.first().and_then(|x| x.as_str()) + { + version_str = format!("ClickHouse {}", v); + } + + let mut uptime_sec = 0; + if let Ok(parsed) = crate::mapper::row_compact::parse_compact_output(&uptime_text, 0, None) + && let Some(row) = parsed.rows.first() + && let Some(v) = row.first().and_then(|x| x.as_u64()) + { + uptime_sec = v; + } + + let db_sizes_text = client + .post_sql( + "SELECT database, sum(total_bytes) FROM system.tables GROUP BY database FORMAT JSONCompactEachRowWithNamesAndTypes", + |_| {}, + ) + .await + .unwrap_or_default(); + let mut db_sizes = serde_json::Map::new(); + if let Ok(parsed) = crate::mapper::row_compact::parse_compact_output(&db_sizes_text, 0, None) { + for row in parsed.rows { + if let (Some(db), Some(size)) = ( + row.first().and_then(|x| x.as_str()), + row.get(1).and_then(|x| x.as_u64()), + ) { + db_sizes.insert(db.to_string(), json!(size)); + } + } + } + + Ok(json!({ + "serverVersion": version_str, + "uptimeSeconds": uptime_sec, + "activeConnections": 1, + "activeQueries": 1, + "memoryUsageBytes": 0, + "databaseSizes": db_sizes, + "extraMetrics": {} + })) +} + +/// Handler for `db.getObjectMetadata`. Returns table/view DDL and column list. +pub async fn handle_get_object_metadata(params: Option) -> Result { + let params_val = params.ok_or_else(|| DriverError::Rpc { + code: -32602, + message: "Invalid params: db.getObjectMetadata requires connectionId, nodeId, and nodeType" + .to_string(), + data: None, + })?; + + let p: GetObjectMetadataParams = + serde_json::from_value(params_val).map_err(|e| DriverError::Rpc { + code: -32602, + message: format!("Malformed getObjectMetadata parameters: {}", e), + data: None, + })?; + + let client = ConnectionPool::global() + .get(p.connection_id) + .ok_or_else(|| DriverError::ConnectionNotFound(p.connection_id))?; + + let parts: Vec = split_node_id(&p.node_id); + let (db_name, tbl_name) = if parts.len() >= 3 && (parts[0] == "table" || parts[0] == "view") { + (parts[1].as_str(), parts[2].as_str()) + } else if parts.len() >= 2 { + (parts[0].as_str(), parts[1].as_str()) + } else { + ("default", p.node_id.as_str()) + }; + + if client.base_url.starts_with("mock://") || client.base_url.starts_with("test://") { + return Ok(json!({ + "nodeId": p.node_id, + "nodeType": p.node_type, + "ddl": format!("CREATE TABLE {}.{} (\n id UInt64,\n created_at DateTime\n) ENGINE = MergeTree ORDER BY id", db_name, tbl_name), + "columns": [ + { "name": "id", "dataType": "UInt64", "isNullable": false, "comment": "Primary ID" }, + { "name": "created_at", "dataType": "DateTime", "isNullable": false, "comment": "Creation timestamp" } + ], + "properties": { + "engine": "MergeTree" + } + })); + } + + let ddl_sql = format!( + "SHOW CREATE TABLE {}.{} FORMAT JSONCompactEachRowWithNamesAndTypes", + crate::utils::sql_escape::quote_identifier(db_name), + crate::utils::sql_escape::quote_identifier(tbl_name) + ); + let mut ddl_str = String::new(); + if let Ok(text) = client.post_sql(&ddl_sql, |_| {}).await + && let Ok(parsed) = crate::mapper::row_compact::parse_compact_output(&text, 0, None) + && let Some(row) = parsed.rows.first() + && let Some(v) = row.first().and_then(|x| x.as_str()) + { + ddl_str = v.to_string(); + } + + let cols_sql = format!( + "SELECT name, type, comment FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position FORMAT JSONCompactEachRowWithNamesAndTypes", + crate::utils::sql_escape::escape_sql_string_literal(db_name), + crate::utils::sql_escape::escape_sql_string_literal(tbl_name) + ); + let mut columns = Vec::new(); + if let Ok(text) = client.post_sql(&cols_sql, |_| {}).await + && let Ok(parsed) = crate::mapper::row_compact::parse_compact_output(&text, 0, None) + { + for row in parsed.rows { + let name = row.first().and_then(|x| x.as_str()).unwrap_or("unknown"); + let col_type = row.get(1).and_then(|x| x.as_str()).unwrap_or("String"); + let comment = row.get(2).and_then(|x| x.as_str()).unwrap_or(""); + let is_nullable = col_type.starts_with("Nullable("); + columns.push(json!({ + "name": name, + "dataType": col_type, + "isNullable": is_nullable, + "comment": comment + })); + } + } + + Ok(json!({ + "nodeId": p.node_id, + "nodeType": p.node_type, + "ddl": ddl_str, + "columns": columns, + "properties": {} + })) +} + #[cfg(test)] mod tests { use super::*; @@ -337,6 +602,47 @@ mod tests { ConnectionPool::global().remove(402); } + #[tokio::test] + async fn test_handle_expand_tree_node_dotted_database_name() { + use crate::utils::node_id::encode_id_segment; + + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 406, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + // A database name containing a literal '.' (legal in ClickHouse via + // backtick-quoted identifiers) must not be misparsed as extra path + // segments when the nodeId is split back apart (CWE-20 regression). + let enc_db = encode_id_segment("my.db"); + let db_node_id = format!("db.{}", enc_db); + assert_eq!(db_node_id, "db.my~ddb"); + + let res_db = + handle_expand_tree_node(Some(json!({ "connectionId": 406, "nodeId": db_node_id }))) + .await + .unwrap(); + let groups = res_db["nodes"].as_array().unwrap(); + assert_eq!(groups.len(), 3); + assert_eq!(groups[0]["id"], format!("group.{}.tables", enc_db)); + + let res_tbl = handle_expand_tree_node(Some( + json!({ "connectionId": 406, "nodeId": groups[0]["id"].as_str().unwrap() }), + )) + .await + .unwrap(); + assert_eq!( + res_tbl["nodes"][0]["id"], + format!("table.{}.events", enc_db) + ); + + ConnectionPool::global().remove(406); + } + #[tokio::test] async fn test_handle_get_connection_form_schema() { let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; @@ -369,4 +675,92 @@ mod tests { ConnectionPool::global().remove(403); } + + #[tokio::test] + async fn test_handle_context_actions_for_column_on_mock_connection() { + // A mock connection has no real system.columns to look up, so the type + // lookup is skipped and column_type stays None, falling back to the + // scalar-oriented profiling query (issue #60). + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 404, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let params = json!({ + "connectionId": 404, + "nodeType": "column", + "nodeId": "col.analytics.events.user_id" + }); + + let res = handle_context_actions(Some(params)).await.unwrap(); + let actions = res["actions"].as_array().unwrap(); + assert_eq!(actions.len(), 2); + assert_eq!(actions[0]["id"], "column.stats"); + assert!( + actions[0]["sql"] + .as_str() + .unwrap() + .contains("min(`user_id`)") + ); + assert_eq!(actions[1]["id"], "column.top_10"); + + ConnectionPool::global().remove(404); + } + + #[tokio::test] + async fn test_handle_get_capabilities() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let res = handle_get_capabilities(None).await.unwrap(); + assert_eq!(res["supportsCancel"], true); + assert_eq!(res["hasServerStats"], true); + } + + #[tokio::test] + async fn test_handle_get_server_stats_mock() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 404, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let res = handle_get_server_stats(Some(json!({ "connectionId": 404 }))) + .await + .unwrap(); + assert_eq!(res["serverVersion"], "ClickHouse 24.3 (Mock)"); + assert_eq!(res["uptimeSeconds"], 3600); + + ConnectionPool::global().remove(404); + } + + #[tokio::test] + async fn test_handle_get_object_metadata_mock() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 405, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let res = handle_get_object_metadata(Some(json!({ + "connectionId": 405, + "nodeId": "table.analytics.events", + "nodeType": "table" + }))) + .await + .unwrap(); + assert_eq!(res["nodeId"], "table.analytics.events"); + assert!(res["ddl"].as_str().unwrap().contains("CREATE TABLE")); + assert_eq!(res["columns"].as_array().unwrap().len(), 2); + + ConnectionPool::global().remove(405); + } } diff --git a/src/rpc/handlers/system.rs b/src/rpc/handlers/system.rs index c908fc0..f71134a 100644 --- a/src/rpc/handlers/system.rs +++ b/src/rpc/handlers/system.rs @@ -48,7 +48,10 @@ pub async fn handle_handshake(params: Option) -> Result) -> Result { @@ -17,6 +17,12 @@ pub async fn dispatch(method: &str, params: Option) -> Result schema::handle_expand_tree_node(params).await, "db.getConnectionFormSchema" => schema::handle_get_connection_form_schema(params).await, "sdui.contextActions" => schema::handle_context_actions(params).await, + "db.getCapabilities" => schema::handle_get_capabilities(params).await, + "db.getServerStats" => schema::handle_get_server_stats(params).await, + "db.getObjectMetadata" | "db.getObjectDDL" => { + schema::handle_get_object_metadata(params).await + } + "commands.execute" => commands::handle_execute(params).await, _ => Err(DriverError::Rpc { code: -32601, message: format!("Method not found: {}", method), @@ -86,6 +92,28 @@ mod tests { ConnectionPool::global().remove(889); } + #[tokio::test] + async fn test_dispatch_commands_execute() { + let _guard = crate::utils::test_lock::GLOBAL_TEST_LOCK.lock().await; + let client = ClickHouseClient::from_params(ConnectParams { + connection_id: 890, + connection_string: Some("mock://localhost:8123/default".to_string()), + ..Default::default() + }) + .unwrap(); + ConnectionPool::global().insert(client); + + let params = json!({ + "commandId": "clickhouse.serverStats", + "connectionId": 890 + }); + + let res = dispatch("commands.execute", Some(params)).await.unwrap(); + assert_eq!(res["serverVersion"], "ClickHouse 24.3 (Mock)"); + + ConnectionPool::global().remove(890); + } + #[tokio::test] async fn test_dispatch_system_handshake_and_ping() { let handshake_res = dispatch("system.handshake", None).await.unwrap(); diff --git a/src/sdui/actions.rs b/src/sdui/actions.rs index cda186b..d5337b3 100644 --- a/src/sdui/actions.rs +++ b/src/sdui/actions.rs @@ -1,4 +1,6 @@ use crate::error::DriverError; +use crate::utils::node_id::split_node_id; +use crate::utils::sql_escape::{escape_sql_string_literal, quote_identifier}; use serde::Serialize; #[derive(Debug, Clone, Serialize, PartialEq)] @@ -38,12 +40,29 @@ impl SduiContextAction { } } +/// Returns `true` if a ClickHouse column type is a complex type (`Array`, `Map`, +/// `Tuple`) that the `min`/`max`/`topK` aggregate functions cannot operate on +/// directly, so the Column Profiler must fall back to type-appropriate queries. +fn is_complex_clickhouse_type(column_type: &str) -> bool { + let mut t = column_type.trim(); + if let Some(inner) = t.strip_prefix("Nullable(") { + t = inner.trim_end_matches(')'); + } + t.starts_with("Array(") || t.starts_with("Map(") || t.starts_with("Tuple(") +} + /// Generates SDUI context menu actions based on `nodeType` and `nodeId`. +/// +/// `column_type` is the ClickHouse type of the target column (e.g. from +/// `system.columns`), used only for `nodeType == "column"` to choose a +/// profiling query that's valid for the column's type. Pass `None` when the +/// type is unknown; the scalar-oriented query is used as a safe default. pub fn get_context_actions_for_node( node_type: &str, node_id: &str, + column_type: Option<&str>, ) -> Result, DriverError> { - let parts: Vec<&str> = node_id.split('.').collect(); + let parts: Vec = split_node_id(node_id); match node_type { "server" | "root_databases" => Ok(vec![ @@ -82,8 +101,12 @@ pub fn get_context_actions_for_node( node_id ))); } - let db_name = parts[1]; - let table_name = parts[2]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); + let q_db = quote_identifier(db_name); + let q_tbl = quote_identifier(table_name); + let esc_db = escape_sql_string_literal(db_name); + let esc_tbl = escape_sql_string_literal(table_name); Ok(vec![ SduiContextAction::new( @@ -93,7 +116,7 @@ pub fn get_context_actions_for_node( "query", Some(format!( "SELECT * FROM {}.{} LIMIT 100", - db_name, table_name + q_db, q_tbl )), false, false, @@ -103,7 +126,7 @@ pub fn get_context_actions_for_node( "📜 Show DDL (SHOW CREATE TABLE)", Some("code"), "query", - Some(format!("SHOW CREATE TABLE {}.{}", db_name, table_name)), + Some(format!("SHOW CREATE TABLE {}.{}", q_db, q_tbl)), false, false, ), @@ -121,7 +144,7 @@ pub fn get_context_actions_for_node( "🔨 Optimize Table (FINAL)", Some("tool"), "execute", - Some(format!("OPTIMIZE TABLE {}.{} FINAL", db_name, table_name)), + Some(format!("OPTIMIZE TABLE {}.{} FINAL", q_db, q_tbl)), true, false, ), @@ -132,7 +155,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "OPTIMIZE TABLE {}.{} DEDUPLICATE", - db_name, table_name + q_db, q_tbl )), true, false, @@ -144,7 +167,7 @@ pub fn get_context_actions_for_node( "query", Some(format!( "SELECT mutation_id, command, create_time, parts_to_do, is_done FROM system.mutations WHERE database = '{}' AND table = '{}' AND is_done = 0", - db_name, table_name + esc_db, esc_tbl )), false, false, @@ -156,7 +179,7 @@ pub fn get_context_actions_for_node( "query", Some(format!( "SELECT query_id, user, query, elapsed, formatReadableSize(memory_usage) AS mem FROM system.processes WHERE current_database = '{}' AND query LIKE '%{}%' AND query NOT LIKE '%system.processes%'", - db_name, table_name + esc_db, esc_tbl )), false, false, @@ -168,7 +191,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "KILL MUTATION WHERE database = '{}' AND table = '{}'", - db_name, table_name + esc_db, esc_tbl )), true, true, @@ -182,8 +205,10 @@ pub fn get_context_actions_for_node( node_id ))); } - let db_name = parts[1]; - let view_name = parts[2]; + let db_name = parts[1].as_str(); + let view_name = parts[2].as_str(); + let q_db = quote_identifier(db_name); + let q_view = quote_identifier(view_name); Ok(vec![ SduiContextAction::new( @@ -191,7 +216,7 @@ pub fn get_context_actions_for_node( "⚡ Top 100 Rows", Some("eye"), "query", - Some(format!("SELECT * FROM {}.{} LIMIT 100", db_name, view_name)), + Some(format!("SELECT * FROM {}.{} LIMIT 100", q_db, q_view)), false, false, ), @@ -200,7 +225,7 @@ pub fn get_context_actions_for_node( "📜 Show DDL (SHOW CREATE TABLE)", Some("code"), "query", - Some(format!("SHOW CREATE TABLE {}.{}", db_name, view_name)), + Some(format!("SHOW CREATE TABLE {}.{}", q_db, q_view)), false, false, ), @@ -213,9 +238,12 @@ pub fn get_context_actions_for_node( node_id ))); } - let db_name = parts[1]; - let table_name = parts[2]; - let partition = parts[3]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); + let partition = parts[3].as_str(); + let q_db = quote_identifier(db_name); + let q_tbl = quote_identifier(table_name); + let esc_part = escape_sql_string_literal(partition); Ok(vec![ SduiContextAction::new( @@ -225,7 +253,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "ALTER TABLE {}.{} DROP PARTITION '{}'", - db_name, table_name, partition + q_db, q_tbl, esc_part )), true, true, @@ -237,7 +265,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "ALTER TABLE {}.{} FREEZE PARTITION '{}'", - db_name, table_name, partition + q_db, q_tbl, esc_part )), false, false, @@ -249,7 +277,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "ALTER TABLE {}.{} DETACH PARTITION '{}'", - db_name, table_name, partition + q_db, q_tbl, esc_part )), true, true, @@ -261,7 +289,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "ALTER TABLE {}.{} ATTACH PARTITION '{}'", - db_name, table_name, partition + q_db, q_tbl, esc_part )), false, false, @@ -273,7 +301,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "OPTIMIZE TABLE {}.{} PARTITION '{}' FINAL", - db_name, table_name, partition + q_db, q_tbl, esc_part )), true, false, @@ -285,7 +313,7 @@ pub fn get_context_actions_for_node( "execute", Some(format!( "OPTIMIZE TABLE {}.{} PARTITION '{}' DEDUPLICATE", - db_name, table_name, partition + q_db, q_tbl, esc_part )), true, false, @@ -299,7 +327,8 @@ pub fn get_context_actions_for_node( node_id ))); } - let db_name = parts[1]; + let db_name = parts[1].as_str(); + let esc_db = escape_sql_string_literal(db_name); Ok(vec![ SduiContextAction::new( @@ -309,7 +338,7 @@ pub fn get_context_actions_for_node( "query", Some(format!( "SELECT mutation_id, table, command, create_time, parts_to_do FROM system.mutations WHERE database = '{}' AND is_done = 0", - db_name + esc_db )), false, false, @@ -321,7 +350,7 @@ pub fn get_context_actions_for_node( "query", Some(format!( "SELECT query_id, user, query, elapsed, formatReadableSize(memory_usage) AS mem FROM system.processes WHERE current_database = '{}'", - db_name + esc_db )), false, false, @@ -331,7 +360,7 @@ pub fn get_context_actions_for_node( "🛑 Kill Mutations in Database", Some("x-circle"), "execute", - Some(format!("KILL MUTATION WHERE database = '{}'", db_name)), + Some(format!("KILL MUTATION WHERE database = '{}'", esc_db)), true, true, ), @@ -340,7 +369,7 @@ pub fn get_context_actions_for_node( "🛑 Kill Queries in Database", Some("x-circle"), "execute", - Some(format!("KILL QUERY WHERE current_database = '{}' ASYNC", db_name)), + Some(format!("KILL QUERY WHERE current_database = '{}' ASYNC", esc_db)), true, true, ), @@ -353,36 +382,56 @@ pub fn get_context_actions_for_node( node_id ))); } - let db_name = parts[1]; - let table_name = parts[2]; - let col_name = parts[3]; + let db_name = parts[1].as_str(); + let table_name = parts[2].as_str(); + let col_name = parts[3].as_str(); + let q_db = quote_identifier(db_name); + let q_tbl = quote_identifier(table_name); + let q_col = quote_identifier(col_name); - Ok(vec![ - SduiContextAction::new( - "column.stats", - "📈 Column Statistics (Быстрый профайлер)", - Some("bar-chart"), - "query", - Some(format!( + // `min`/`max`/`topK` cannot operate on Array/Map/Tuple columns in + // ClickHouse (ILLEGAL_TYPE_OF_ARGUMENT), and GROUP BY on a Map column + // is rejected outright, so complex-typed columns get a profiling + // query built from length()/uniqueness instead, and no "Top 10 + // Frequent Values" action (see issue #60). + let is_complex = column_type.is_some_and(is_complex_clickhouse_type); + + let mut actions = vec![SduiContextAction::new( + "column.stats", + "📈 Column Statistics (Быстрый профайлер)", + Some("bar-chart"), + "query", + Some(if is_complex { + format!( + "SELECT count() as total_rows, countIf(isNotNull({0})) as not_nulls, uniqExact({0}) as unique_exact, min(length({0})) as min_length, max(length({0})) as max_length FROM {1}.{2}", + q_col, q_db, q_tbl + ) + } else { + format!( "SELECT count() as total_rows, countIf(isNotNull({0})) as not_nulls, uniqExact({0}) as unique_exact, min({0}) as min_val, max({0}) as max_val, topK(5)({0}) as top_5_values FROM {1}.{2}", - col_name, db_name, table_name - )), - false, - false, - ), - SduiContextAction::new( + q_col, q_db, q_tbl + ) + }), + false, + false, + )]; + + if !is_complex { + actions.push(SduiContextAction::new( "column.top_10", "🔝 Top 10 Frequent Values", Some("list"), "query", Some(format!( "SELECT {0}, count() as cnt FROM {1}.{2} GROUP BY {0} ORDER BY cnt DESC LIMIT 10", - col_name, db_name, table_name + q_col, q_db, q_tbl )), false, false, - ), - ]) + )); + } + + Ok(actions) } _ => Ok(Vec::new()), } @@ -394,19 +443,20 @@ mod tests { #[test] fn test_table_context_actions() { - let actions = get_context_actions_for_node("table", "table.analytics.events").unwrap(); + let actions = + get_context_actions_for_node("table", "table.analytics.events", None).unwrap(); assert_eq!(actions.len(), 8); assert_eq!(actions[0].id, "table.top_100"); assert_eq!( actions[0].sql.as_deref(), - Some("SELECT * FROM analytics.events LIMIT 100") + Some("SELECT * FROM `analytics`.`events` LIMIT 100") ); assert!(!actions[0].requires_confirmation); assert_eq!(actions[3].id, "table.optimize_final"); assert_eq!( actions[3].sql.as_deref(), - Some("OPTIMIZE TABLE analytics.events FINAL") + Some("OPTIMIZE TABLE `analytics`.`events` FINAL") ); assert!(actions[3].requires_confirmation); @@ -417,12 +467,13 @@ mod tests { #[test] fn test_partition_context_actions() { let actions = - get_context_actions_for_node("partition", "part.analytics.events.202607").unwrap(); + get_context_actions_for_node("partition", "part.analytics.events.202607", None) + .unwrap(); assert_eq!(actions.len(), 6); assert_eq!(actions[0].id, "partition.drop"); assert_eq!( actions[0].sql.as_deref(), - Some("ALTER TABLE analytics.events DROP PARTITION '202607'") + Some("ALTER TABLE `analytics`.`events` DROP PARTITION '202607'") ); assert!(actions[0].requires_confirmation); assert!(actions[0].danger); @@ -430,7 +481,7 @@ mod tests { assert_eq!(actions[1].id, "partition.freeze"); assert_eq!( actions[1].sql.as_deref(), - Some("ALTER TABLE analytics.events FREEZE PARTITION '202607'") + Some("ALTER TABLE `analytics`.`events` FREEZE PARTITION '202607'") ); assert!(!actions[1].requires_confirmation); assert!(!actions[1].danger); @@ -438,7 +489,7 @@ mod tests { assert_eq!(actions[2].id, "partition.detach"); assert_eq!( actions[2].sql.as_deref(), - Some("ALTER TABLE analytics.events DETACH PARTITION '202607'") + Some("ALTER TABLE `analytics`.`events` DETACH PARTITION '202607'") ); assert!(actions[2].requires_confirmation); assert!(actions[2].danger); @@ -446,7 +497,7 @@ mod tests { assert_eq!(actions[3].id, "partition.attach"); assert_eq!( actions[3].sql.as_deref(), - Some("ALTER TABLE analytics.events ATTACH PARTITION '202607'") + Some("ALTER TABLE `analytics`.`events` ATTACH PARTITION '202607'") ); assert!(!actions[3].requires_confirmation); assert!(!actions[3].danger); @@ -454,21 +505,21 @@ mod tests { assert_eq!(actions[4].id, "partition.optimize_final"); assert_eq!( actions[4].sql.as_deref(), - Some("OPTIMIZE TABLE analytics.events PARTITION '202607' FINAL") + Some("OPTIMIZE TABLE `analytics`.`events` PARTITION '202607' FINAL") ); assert!(actions[4].requires_confirmation); assert_eq!(actions[5].id, "partition.deduplicate"); assert_eq!( actions[5].sql.as_deref(), - Some("OPTIMIZE TABLE analytics.events PARTITION '202607' DEDUPLICATE") + Some("OPTIMIZE TABLE `analytics`.`events` PARTITION '202607' DEDUPLICATE") ); assert!(actions[5].requires_confirmation); } #[test] fn test_database_and_view_actions() { - let db_actions = get_context_actions_for_node("database", "db.analytics").unwrap(); + let db_actions = get_context_actions_for_node("database", "db.analytics", None).unwrap(); assert_eq!(db_actions.len(), 4); assert_eq!(db_actions[0].id, "db.active_mutations"); assert_eq!(db_actions[1].id, "db.active_queries"); @@ -476,14 +527,19 @@ mod tests { assert_eq!(db_actions[3].id, "db.kill_queries"); let view_actions = - get_context_actions_for_node("view", "view.analytics.mv_summary").unwrap(); + get_context_actions_for_node("view", "view.analytics.mv_summary", None).unwrap(); assert_eq!(view_actions.len(), 2); assert_eq!(view_actions[0].id, "view.top_100"); + assert_eq!( + view_actions[0].sql.as_deref(), + Some("SELECT * FROM `analytics`.`mv_summary` LIMIT 100") + ); } #[test] fn test_server_and_process_monitoring_actions() { - let server_actions = get_context_actions_for_node("server", "server.cluster").unwrap(); + let server_actions = + get_context_actions_for_node("server", "server.cluster", None).unwrap(); assert_eq!(server_actions.len(), 3); assert_eq!(server_actions[0].id, "server.active_mutations"); assert_eq!(server_actions[1].id, "server.active_queries"); @@ -494,13 +550,13 @@ mod tests { #[test] fn test_column_context_actions() { let actions = - get_context_actions_for_node("column", "col.analytics.events.user_id").unwrap(); + get_context_actions_for_node("column", "col.analytics.events.user_id", None).unwrap(); assert_eq!(actions.len(), 2); assert_eq!(actions[0].id, "column.stats"); assert_eq!( actions[0].sql.as_deref(), Some( - "SELECT count() as total_rows, countIf(isNotNull(user_id)) as not_nulls, uniqExact(user_id) as unique_exact, min(user_id) as min_val, max(user_id) as max_val, topK(5)(user_id) as top_5_values FROM analytics.events" + "SELECT count() as total_rows, countIf(isNotNull(`user_id`)) as not_nulls, uniqExact(`user_id`) as unique_exact, min(`user_id`) as min_val, max(`user_id`) as max_val, topK(5)(`user_id`) as top_5_values FROM `analytics`.`events`" ) ); assert_eq!(actions[0].action_type, "query"); @@ -509,15 +565,125 @@ mod tests { assert_eq!( actions[1].sql.as_deref(), Some( - "SELECT user_id, count() as cnt FROM analytics.events GROUP BY user_id ORDER BY cnt DESC LIMIT 10" + "SELECT `user_id`, count() as cnt FROM `analytics`.`events` GROUP BY `user_id` ORDER BY cnt DESC LIMIT 10" ) ); } + #[test] + fn test_column_context_actions_for_complex_types() { + // Regression for issue #60: min()/max()/topK() abort with + // ILLEGAL_TYPE_OF_ARGUMENT on Array/Map/Tuple columns, and GROUP BY + // on a raw Map column is rejected outright, so complex-typed columns + // must get a length()-based profiling query and no top_10 action. + for complex_type in [ + "Array(String)", + "Map(String, UInt64)", + "Tuple(Int32, String)", + "Nullable(Array(String))", + ] { + let actions = get_context_actions_for_node( + "column", + "col.analytics.events.tags", + Some(complex_type), + ) + .unwrap(); + assert_eq!( + actions.len(), + 1, + "expected only column.stats for type {}", + complex_type + ); + assert_eq!(actions[0].id, "column.stats"); + let sql = actions[0].sql.as_deref().unwrap(); + assert!( + !sql.contains("min(`tags`)"), + "type {}: {}", + complex_type, + sql + ); + assert!( + !sql.contains("max(`tags`)"), + "type {}: {}", + complex_type, + sql + ); + assert!(!sql.contains("topK"), "type {}: {}", complex_type, sql); + assert!( + sql.contains("min(length(`tags`))"), + "type {}: {}", + complex_type, + sql + ); + assert!( + sql.contains("max(length(`tags`))"), + "type {}: {}", + complex_type, + sql + ); + } + + // A plain scalar type keeps the original min/max/topK query and the top_10 action. + let scalar_actions = get_context_actions_for_node( + "column", + "col.analytics.events.tags", + Some("LowCardinality(String)"), + ) + .unwrap(); + assert_eq!(scalar_actions.len(), 2); + assert!( + scalar_actions[0] + .sql + .as_deref() + .unwrap() + .contains("min(`tags`)") + ); + } + + #[test] + fn test_context_actions_sql_injection_protection() { + let malicious_part = "part.db`test.tbl'test.2026'; DROP TABLE secret; --"; + let actions = get_context_actions_for_node("partition", malicious_part, None).unwrap(); + assert_eq!( + actions[0].sql.as_deref(), + Some( + "ALTER TABLE `db``test`.`tbl'test` DROP PARTITION '2026\'\'; DROP TABLE secret; --'" + ) + ); + + let malicious_col = "col.db.tbl.user_id`; DROP TABLE users; --"; + let col_actions = get_context_actions_for_node("column", malicious_col, None).unwrap(); + assert!( + col_actions[1] + .sql + .as_ref() + .unwrap() + .contains("`user_id``; DROP TABLE users; --`") + ); + } + + #[test] + fn test_context_actions_with_dotted_table_name() { + use crate::utils::node_id::encode_id_segment; + + // A table name containing a literal '.' must survive nodeId round-trip + // unmangled instead of being split into extra bogus segments (CWE-20). + let node_id = format!( + "table.{}.{}", + encode_id_segment("analytics"), + encode_id_segment("weird.table.name") + ); + let actions = get_context_actions_for_node("table", &node_id, None).unwrap(); + assert_eq!( + actions[0].sql.as_deref(), + Some("SELECT * FROM `analytics`.`weird.table.name` LIMIT 100") + ); + } + #[test] fn test_invalid_node_id() { - assert!(get_context_actions_for_node("table", "table.only").is_err()); - assert!(get_context_actions_for_node("partition", "part.only.two").is_err()); - assert!(get_context_actions_for_node("column", "col.only.two").is_err()); + assert!(get_context_actions_for_node("table", "table.only", None).is_err()); + assert!(get_context_actions_for_node("partition", "part.only.two", None).is_err()); + assert!(get_context_actions_for_node("column", "col.only.two", None).is_err()); } } diff --git a/src/sdui/tree.rs b/src/sdui/tree.rs index fb29db1..bf99f2a 100644 --- a/src/sdui/tree.rs +++ b/src/sdui/tree.rs @@ -1,5 +1,6 @@ use crate::error::DriverError; use crate::mapper::row_compact::parse_compact_output; +use crate::utils::node_id::encode_id_segment; use serde::Serialize; use serde_json::{Value, json}; @@ -43,14 +44,14 @@ pub fn build_root_databases_nodes( let mut nodes = Vec::new(); if let Some(output) = compact_output { - let parsed = parse_compact_output(output, 0)?; + let parsed = parse_compact_output(output, 0, None)?; for row in parsed.rows { let name = row.first().and_then(|v| v.as_str()).unwrap_or("unknown"); let engine = row.get(1).and_then(|v| v.as_str()).unwrap_or(""); let comment = row.get(2).and_then(|v| v.as_str()).unwrap_or(""); nodes.push(SduiTreeNode::new( - format!("db.{}", name), + format!("db.{}", encode_id_segment(name)), name, "database", Some("database"), @@ -86,9 +87,10 @@ pub fn build_root_databases_nodes( /// Builds database child groups (`Tables`, `Views`, `Dictionaries`) when a `database` node is expanded. pub fn build_database_groups(db_name: &str) -> Vec { + let enc_db = encode_id_segment(db_name); vec![ SduiTreeNode::new( - format!("group.{}.tables", db_name), + format!("group.{}.tables", enc_db), "Таблицы (Tables)", "group", Some("folder-table"), @@ -96,7 +98,7 @@ pub fn build_database_groups(db_name: &str) -> Vec { Some(json!({ "database": db_name, "group": "tables" })), ), SduiTreeNode::new( - format!("group.{}.views", db_name), + format!("group.{}.views", enc_db), "Представления (Views)", "group", Some("folder-eye"), @@ -104,7 +106,7 @@ pub fn build_database_groups(db_name: &str) -> Vec { Some(json!({ "database": db_name, "group": "views" })), ), SduiTreeNode::new( - format!("group.{}.dictionaries", db_name), + format!("group.{}.dictionaries", enc_db), "Словари (Dictionaries)", "group", Some("folder-book"), @@ -116,9 +118,11 @@ pub fn build_database_groups(db_name: &str) -> Vec { /// Builds table/view sub-groups (`Columns`, `Partitions`) when a `table` node is expanded. pub fn build_table_groups(db_name: &str, table_name: &str) -> Vec { + let enc_db = encode_id_segment(db_name); + let enc_tbl = encode_id_segment(table_name); vec![ SduiTreeNode::new( - format!("group_cols.{}.{}", db_name, table_name), + format!("group_cols.{}.{}", enc_db, enc_tbl), "Колонки (Columns)", "group_cols", Some("folder"), @@ -126,7 +130,7 @@ pub fn build_table_groups(db_name: &str, table_name: &str) -> Vec Some(json!({ "database": db_name, "table": table_name })), ), SduiTreeNode::new( - format!("group_parts.{}.{}", db_name, table_name), + format!("group_parts.{}.{}", enc_db, enc_tbl), "Партиции (Partitions)", "group_parts", Some("folder"), @@ -142,7 +146,7 @@ pub fn parse_tables_nodes( compact_output: &str, filter_view: bool, ) -> Result, DriverError> { - let parsed = parse_compact_output(compact_output, 0)?; + let parsed = parse_compact_output(compact_output, 0, None)?; let mut nodes = Vec::new(); for row in parsed.rows { @@ -164,7 +168,12 @@ pub fn parse_tables_nodes( let icon = if obj_type == "view" { "eye" } else { "table" }; nodes.push(SduiTreeNode::new( - format!("{}.{}.{}", node_type, db_name, name), + format!( + "{}.{}.{}", + node_type, + encode_id_segment(db_name), + encode_id_segment(name) + ), name, node_type, Some(icon), @@ -186,7 +195,7 @@ pub fn parse_dictionaries_nodes( db_name: &str, compact_output: &str, ) -> Result, DriverError> { - let parsed = parse_compact_output(compact_output, 0)?; + let parsed = parse_compact_output(compact_output, 0, None)?; let mut nodes = Vec::new(); for row in parsed.rows { @@ -197,7 +206,11 @@ pub fn parse_dictionaries_nodes( let size = row.get(5).and_then(|v| v.as_str()).unwrap_or("0 B"); nodes.push(SduiTreeNode::new( - format!("dict.{}.{}", db_name, name), + format!( + "dict.{}.{}", + encode_id_segment(db_name), + encode_id_segment(name) + ), name, "dictionary", Some("book"), @@ -220,7 +233,7 @@ pub fn parse_columns_nodes( table_name: &str, compact_output: &str, ) -> Result, DriverError> { - let parsed = parse_compact_output(compact_output, 0)?; + let parsed = parse_compact_output(compact_output, 0, None)?; let mut nodes = Vec::new(); for row in parsed.rows { @@ -229,7 +242,12 @@ pub fn parse_columns_nodes( let comment = row.get(2).and_then(|v| v.as_str()).unwrap_or(""); nodes.push(SduiTreeNode::new( - format!("col.{}.{}.{}", db_name, table_name, name), + format!( + "col.{}.{}.{}", + encode_id_segment(db_name), + encode_id_segment(table_name), + encode_id_segment(name) + ), format!("{} ({})", name, col_type), "column", Some("columns"), @@ -251,7 +269,7 @@ pub fn parse_partitions_nodes( table_name: &str, compact_output: &str, ) -> Result, DriverError> { - let parsed = parse_compact_output(compact_output, 0)?; + let parsed = parse_compact_output(compact_output, 0, None)?; let mut nodes = Vec::new(); for row in parsed.rows { @@ -261,7 +279,12 @@ pub fn parse_partitions_nodes( let parts_count = row.get(3).cloned().unwrap_or(json!(1)); nodes.push(SduiTreeNode::new( - format!("part.{}.{}.{}", db_name, table_name, partition), + format!( + "part.{}.{}.{}", + encode_id_segment(db_name), + encode_id_segment(table_name), + encode_id_segment(partition) + ), format!("⚡ {}", partition), "partition", Some("archive"), diff --git a/src/transport/framing.rs b/src/transport/framing.rs index 92d588f..e852fdb 100644 --- a/src/transport/framing.rs +++ b/src/transport/framing.rs @@ -12,16 +12,21 @@ pub fn write_ndjson_stdout(payload: &str) -> io::Result<()> { /// Generic NDJSON writer that works with any `std::io::Write` sink (useful for testing). pub fn write_ndjson(writer: &mut W, payload: &str) -> io::Result<()> { // If the JSON payload contains raw '\n' or '\r' bytes (not escaped inside strings), - // replace them with spaces or strip to guarantee exact NDJSON framing. - if payload.contains('\n') || payload.contains('\r') { - let sanitized: String = payload - .chars() - .map(|c| if c == '\n' || c == '\r' { ' ' } else { c }) - .collect(); - writer.write_all(sanitized.as_bytes())?; - } else { - writer.write_all(payload.as_bytes())?; + // replace them with spaces to guarantee exact NDJSON framing. Both are single-byte + // ASCII code points, so this writes existing byte slices straight to `writer` + // between them instead of collecting a sanitized copy of the whole payload onto + // the heap first (issue #50) — a no-op payload (the common case) is written in + // one `write_all` call, same as before. + let bytes = payload.as_bytes(); + let mut start = 0; + for (i, &b) in bytes.iter().enumerate() { + if b == b'\n' || b == b'\r' { + writer.write_all(&bytes[start..i])?; + writer.write_all(b" ")?; + start = i + 1; + } } + writer.write_all(&bytes[start..])?; writer.write_all(b"\n")?; writer.flush()?; Ok(()) @@ -56,6 +61,36 @@ mod tests { ); } + #[test] + fn test_write_ndjson_sanitizes_edge_positions_and_runs() { + // Regression for issue #50: the byte-slice rewrite must handle a + // newline as the very first/last byte and consecutive newlines + // (an empty slice between them) without panicking or dropping bytes. + let mut buffer = Vec::new(); + let dirty = "\nleading\r\rmiddle\n\ntrailing\n"; + write_ndjson(&mut buffer, dirty).unwrap(); + let output = String::from_utf8(buffer).unwrap(); + assert_eq!(output, " leading middle trailing \n"); + } + + #[test] + fn test_write_ndjson_sanitizes_large_payload() { + // A payload well past any small-buffer fast path, to exercise the + // byte-slice rewrite over a realistic multi-megabyte tabular result. + let mut payload = "x".repeat(2 * 1024 * 1024); + payload.push('\n'); + payload.push_str(&"y".repeat(1024)); + + let mut buffer = Vec::new(); + write_ndjson(&mut buffer, &payload).unwrap(); + let output = String::from_utf8(buffer).unwrap(); + + assert_eq!(output.len(), payload.len() + 1); + assert_eq!(&output[..2 * 1024 * 1024], "x".repeat(2 * 1024 * 1024)); + assert_eq!(output.as_bytes()[2 * 1024 * 1024], b' '); + assert!(output.ends_with(&format!("{}\n", "y".repeat(1024)))); + } + #[test] fn test_write_ndjson_error_payload() { let mut buffer = Vec::new(); diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 906aa48..29843c1 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1,5 +1,7 @@ //! Utilities and sanitized logger. pub mod logger; +pub mod node_id; pub mod recovery; pub mod secret_guard; +pub mod sql_escape; pub mod test_lock; diff --git a/src/utils/node_id.rs b/src/utils/node_id.rs new file mode 100644 index 0000000..6952187 --- /dev/null +++ b/src/utils/node_id.rs @@ -0,0 +1,109 @@ +//! Structured encoding/decoding for SDUI tree `nodeId` path segments. +//! +//! `nodeId` values are built by joining a fixed literal prefix (e.g. `"table"`, +//! `"db"`, `"col"`, `"part"`) with dynamic ClickHouse identifiers (database, +//! table, column and partition names) using `.` as the segment separator, e.g. +//! `table.analytics.events`. ClickHouse allows `.` inside backtick-quoted +//! identifiers and inside partition expression values, so a dynamic segment +//! can itself legally contain a literal `.`. A naive `nodeId.split('.')` would +//! then misalign every subsequent segment (CWE-20: Improper Input +//! Validation). To keep parsing unambiguous, every dynamic segment is escaped +//! before being joined, and decoded back after splitting. + +/// Escapes a single dynamic `nodeId` segment so the `.` separator and the +/// escape character `~` used within it can never be confused with the path +/// separator when the full `nodeId` is later split on `.`. +pub fn encode_id_segment(raw: &str) -> String { + let mut out = String::with_capacity(raw.len()); + for ch in raw.chars() { + match ch { + '~' => out.push_str("~t"), + '.' => out.push_str("~d"), + _ => out.push(ch), + } + } + out +} + +/// Reverses `encode_id_segment`, restoring the original identifier. +/// A dangling `~` not followed by a recognized escape code is passed through +/// literally, since it cannot have been produced by `encode_id_segment`. +pub fn decode_id_segment(encoded: &str) -> String { + let mut out = String::with_capacity(encoded.len()); + let mut chars = encoded.chars().peekable(); + while let Some(ch) = chars.next() { + if ch == '~' { + match chars.peek() { + Some('d') => { + out.push('.'); + chars.next(); + } + Some('t') => { + out.push('~'); + chars.next(); + } + _ => out.push('~'), + } + } else { + out.push(ch); + } + } + out +} + +/// Splits a full `nodeId` on `.` and decodes each resulting segment. +/// The first segment (the fixed type prefix, e.g. `"table"`) is always a +/// literal constant chosen by the driver, never dynamic data, so decoding it +/// is a harmless no-op. +pub fn split_node_id(node_id: &str) -> Vec { + node_id.split('.').map(decode_id_segment).collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_roundtrip_plain_identifier() { + let encoded = encode_id_segment("analytics"); + assert_eq!(encoded, "analytics"); + assert_eq!(decode_id_segment(&encoded), "analytics"); + } + + #[test] + fn test_roundtrip_dotted_identifier() { + let encoded = encode_id_segment("my.table"); + assert_eq!(encoded, "my~dtable"); + assert_eq!(decode_id_segment(&encoded), "my.table"); + } + + #[test] + fn test_roundtrip_tilde_and_dot_mixed() { + let raw = "weird~name.with.dots~and~tildes"; + let encoded = encode_id_segment(raw); + assert_eq!(decode_id_segment(&encoded), raw); + } + + #[test] + fn test_split_node_id_preserves_dotted_segments() { + let node_id = format!( + "table.{}.{}", + encode_id_segment("my.db"), + encode_id_segment("weird.table.name") + ); + let parts = split_node_id(&node_id); + assert_eq!(parts, vec!["table", "my.db", "weird.table.name"]); + } + + #[test] + fn test_split_node_id_plain_backward_compatible() { + let parts = split_node_id("table.analytics.events"); + assert_eq!(parts, vec!["table", "analytics", "events"]); + } + + #[test] + fn test_dangling_tilde_passthrough() { + assert_eq!(decode_id_segment("foo~"), "foo~"); + assert_eq!(decode_id_segment("foo~x"), "foo~x"); + } +} diff --git a/src/utils/recovery.rs b/src/utils/recovery.rs index ec178db..933b44d 100644 --- a/src/utils/recovery.rs +++ b/src/utils/recovery.rs @@ -33,21 +33,130 @@ pub fn init_panic_hook() { })); } -/// Verifies and initializes scratch / shadow directory structure (`/tmp/clickhouse-query-ext/shadow/` or system temp). -/// Called on startup by `main.rs` to ensure temporary buffers and partition freezes have a reliable workspace. +/// Returns the secure, user-isolated base scratch directory path. +/// +/// On Unix: +/// Uses `$XDG_RUNTIME_DIR/clickhouse-query-ext/shadow` if `$XDG_RUNTIME_DIR` is set and absolute, +/// otherwise `/tmp/clickhouse-query-ext-/shadow`. +/// +/// On non-Unix platforms (e.g. Windows): +/// Uses `std::env::temp_dir().join("clickhouse-query-ext").join("shadow")`. +pub fn get_scratch_base_dir() -> PathBuf { + #[cfg(unix)] + { + if let Ok(runtime_dir) = std::env::var("XDG_RUNTIME_DIR") { + let path = PathBuf::from(runtime_dir); + if path.is_absolute() { + return path.join("clickhouse-query-ext").join("shadow"); + } + } + let uid = unsafe { libc::geteuid() }; + std::env::temp_dir() + .join(format!("clickhouse-query-ext-{}", uid)) + .join("shadow") + } + + #[cfg(not(unix))] + { + std::env::temp_dir() + .join("clickhouse-query-ext") + .join("shadow") + } +} + +/// Verifies and initializes scratch / shadow directory structure with strict permission isolation. +/// +/// On Unix systems: +/// - Isolates directory per user UID to prevent pre-creation attacks in shared `/tmp`. +/// - Rejects symbolic links at the base and parent directory levels. +/// - Validates ownership matches current effective UID. +/// - Enforces `0700` (`rwx------`) permissions so other unprivileged users cannot read buffers or partition freezes. pub fn ensure_scratch_directories() -> std::io::Result { - let base_dir = std::env::temp_dir() - .join("clickhouse-query-ext") - .join("shadow"); - if !base_dir.exists() { - std::fs::create_dir_all(&base_dir)?; - info!("Created sandbox scratch directory at {:?}", base_dir); - } else { - info!( - "Verified sandbox scratch directory integrity at {:?}", - base_dir - ); + let base_dir = get_scratch_base_dir(); + + #[cfg(unix)] + { + use std::os::unix::fs::{DirBuilderExt, MetadataExt, PermissionsExt}; + + let parent_dir = base_dir.parent().unwrap_or(&base_dir); + let euid = unsafe { libc::geteuid() }; + + // 1. Ensure parent directory (e.g. /tmp/clickhouse-query-ext-) exists with 0700 + if parent_dir.exists() { + let meta = std::fs::symlink_metadata(parent_dir)?; + if meta.file_type().is_symlink() { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + format!( + "Scratch parent directory {:?} is an insecure symlink", + parent_dir + ), + )); + } + if meta.uid() != euid { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + format!( + "Scratch parent directory {:?} is not owned by current user (owner UID {}, current UID {})", + parent_dir, + meta.uid(), + euid + ), + )); + } + std::fs::set_permissions(parent_dir, std::fs::Permissions::from_mode(0o700))?; + } else { + let mut builder = std::fs::DirBuilder::new(); + builder.recursive(true); + builder.mode(0o700); + builder.create(parent_dir)?; + } + + // 2. Ensure base_dir (shadow) exists with 0700 + if base_dir.exists() { + let meta = std::fs::symlink_metadata(&base_dir)?; + if meta.file_type().is_symlink() { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + format!("Scratch directory {:?} is an insecure symlink", base_dir), + )); + } + if meta.uid() != euid { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + format!( + "Scratch directory {:?} is not owned by current user", + base_dir + ), + )); + } + std::fs::set_permissions(&base_dir, std::fs::Permissions::from_mode(0o700))?; + info!( + "Verified secure sandbox scratch directory at {:?}", + base_dir + ); + } else { + let mut builder = std::fs::DirBuilder::new(); + builder.recursive(true); + builder.mode(0o700); + builder.create(&base_dir)?; + info!("Created secure sandbox scratch directory at {:?}", base_dir); + } + } + + #[cfg(not(unix))] + { + if !base_dir.exists() { + std::fs::create_dir_all(&base_dir)?; + info!("Created sandbox scratch directory at {:?}", base_dir); + } else { + info!( + "Verified sandbox scratch directory integrity at {:?}", + base_dir + ); + } } + Ok(base_dir) } @@ -61,6 +170,19 @@ mod tests { assert!(path.exists()); assert!(path.is_dir()); assert!(path.ends_with("shadow")); + + #[cfg(unix)] + { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + let meta = std::fs::symlink_metadata(&path).unwrap(); + assert_eq!(meta.permissions().mode() & 0o777, 0o700); + assert_eq!(meta.uid(), unsafe { libc::geteuid() }); + + let parent = path.parent().unwrap(); + let parent_meta = std::fs::symlink_metadata(parent).unwrap(); + assert_eq!(parent_meta.permissions().mode() & 0o777, 0o700); + assert_eq!(parent_meta.uid(), unsafe { libc::geteuid() }); + } } #[test] @@ -68,4 +190,24 @@ mod tests { // Calling init_panic_hook registers the hook without panic or crash init_panic_hook(); } + + #[test] + #[cfg(unix)] + fn test_symlink_rejection_in_scratch() { + let test_root = + std::env::temp_dir().join(format!("test-scratch-symlink-{}", std::process::id())); + let _ = std::fs::remove_dir_all(&test_root); + std::fs::create_dir_all(&test_root).unwrap(); + + let real_dir = test_root.join("real"); + std::fs::create_dir(&real_dir).unwrap(); + + let symlink_path = test_root.join("symlink_target"); + std::os::unix::fs::symlink(&real_dir, &symlink_path).unwrap(); + + let meta = std::fs::symlink_metadata(&symlink_path).unwrap(); + assert!(meta.file_type().is_symlink()); + + let _ = std::fs::remove_dir_all(&test_root); + } } diff --git a/src/utils/sql_escape.rs b/src/utils/sql_escape.rs new file mode 100644 index 0000000..2ef2570 --- /dev/null +++ b/src/utils/sql_escape.rs @@ -0,0 +1,58 @@ +//! ClickHouse SQL identifier and string literal escaping helpers. +//! Prevents SQL injection (CWE-89) in schema introspection queries and SDUI context actions. + +/// Escapes a string for safe inclusion inside a ClickHouse SQL string literal (`'...'`). +/// +/// Doubles all backslashes (`\`) and single quotes (`'`), ensuring malicious input cannot +/// break out of the string literal boundary. +pub fn escape_sql_string_literal(val: &str) -> String { + val.replace('\\', "\\\\").replace('\'', "''") +} + +/// Quotes and escapes a ClickHouse SQL identifier (database, table, partition, or column name). +/// +/// Wraps the identifier in backticks (`` `...` ``), escaping any embedded backticks (``` ` ```) +/// by doubling them (`` `` ``) and backslashes by doubling them (`\\`). +pub fn quote_identifier(val: &str) -> String { + let escaped = val.replace('\\', "\\\\").replace('`', "``"); + format!("`{}`", escaped) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_escape_sql_string_literal_benign() { + assert_eq!(escape_sql_string_literal("analytics"), "analytics"); + assert_eq!(escape_sql_string_literal("events_2026"), "events_2026"); + } + + #[test] + fn test_escape_sql_string_literal_injection() { + assert_eq!( + escape_sql_string_literal("test' OR 1=1 --"), + "test'' OR 1=1 --" + ); + assert_eq!( + escape_sql_string_literal(r"test\' OR 'a'='a"), + r"test\\'' OR ''a''=''a" + ); + assert_eq!(escape_sql_string_literal("O'Reilly"), "O''Reilly"); + } + + #[test] + fn test_quote_identifier_benign() { + assert_eq!(quote_identifier("default"), "`default`"); + assert_eq!(quote_identifier("system.tables"), "`system.tables`"); + } + + #[test] + fn test_quote_identifier_injection() { + assert_eq!(quote_identifier("table`name"), "`table``name`"); + assert_eq!( + quote_identifier("db`; DROP TABLE students; --"), + "`db``; DROP TABLE students; --`" + ); + } +}