diff --git a/apps/desktop/src/components/editor/QueryEditor.vue b/apps/desktop/src/components/editor/QueryEditor.vue index 361d28dae..d51bcdfb5 100644 --- a/apps/desktop/src/components/editor/QueryEditor.vue +++ b/apps/desktop/src/components/editor/QueryEditor.vue @@ -29,6 +29,7 @@ import { getSqlFunctionSignatureHelp, getSqlCompletionContext, getSqlCompletionResultValidFor, + isSqlLikeCompletionStatement, shouldAutoOpenSqlCompletion, extractCteDefinitions, } from "@/lib/sqlCompletion"; @@ -739,7 +740,13 @@ async function refreshSemanticDiagnostics() { setSemanticDiagnostics([]); return; } - if (!shouldRunSqlSemanticDiagnostics(sql, currentView.state.selection.main.head)) { + if (props.databaseType === "elasticsearch") { + setSemanticDiagnostics([]); + return; + } + if ( + !shouldRunSqlSemanticDiagnostics(sql, currentView.state.selection.main.head, { databaseType: props.databaseType }) + ) { scheduleSemanticDiagnostics(1200); return; } @@ -933,15 +940,17 @@ async function provideSqlCompletions( explicit: boolean, ) { if (!props.connectionId) return null; + const fullDoc = currentState.doc.toString(); if (props.databaseType === "elasticsearch") { - return provideElasticsearchCompletions(currentState, position, explicit); + if (!isSqlLikeCompletionStatement(fullDoc, position)) { + return provideElasticsearchCompletions(currentState, position, explicit); + } } const hasDatabase = props.database != null; const epoch = ++completionEpoch; try { - const fullDoc = currentState.doc.toString(); if (!explicit && !shouldAutoOpenSqlCompletion(fullDoc, position)) return null; const completionContext = getSqlCompletionContext(fullDoc, position); diff --git a/apps/desktop/src/components/grid/DataGrid.vue b/apps/desktop/src/components/grid/DataGrid.vue index f4cc95648..4a9286615 100644 --- a/apps/desktop/src/components/grid/DataGrid.vue +++ b/apps/desktop/src/components/grid/DataGrid.vue @@ -1857,10 +1857,30 @@ const showTruncationWarning = computed( () => props.result.truncated === true && typeof props.pageLimit !== "number" && props.result.has_more !== true, ); const isResultsContext = computed(() => props.context === "results"); -const displayedTotalRowCount = computed(() => props.totalRowCount ?? manualTotalRowCount.value); +// affected_rows reported by the backend can be larger than the rows we +// actually have in memory — e.g. ES auto-pages SELECT * on a big index and +// reports the index's true match count. Surface that in the status bar so +// the user sees the real total, but do NOT use it to unlock pagination: +// we don't have those rows, so letting the user page into them would just +// show blank screens. +const inferredBackendTotalRowCount = computed(() => { + const affected = props.result.affected_rows; + if (typeof affected !== "number" || !Number.isFinite(affected)) return undefined; + if (affected <= props.result.rows.length) return undefined; + return affected; +}); +const serverKnownTotalRowCount = computed(() => props.totalRowCount ?? manualTotalRowCount.value); +const displayedTotalRowCount = computed(() => serverKnownTotalRowCount.value ?? inferredBackendTotalRowCount.value); +// Only a server-confirmed total drives pagination — an inferred total means +// rows exist that we never fetched, so navigation must stay inside rows.length. const hasKnownTotalRowCount = computed( - () => typeof displayedTotalRowCount.value === "number" && displayedTotalRowCount.value >= 0, + () => typeof serverKnownTotalRowCount.value === "number" && serverKnownTotalRowCount.value >= 0, ); +// When context=results and the caller hasn't configured server-side +// pagination (no pageLimit), the backend handed us every row up-front and +// rowCount IS the total. Without this hint, the "page is full → assume more" +// fallback in canGoNextDataGridPage lets the user keep clicking next forever. +const allRowsLoaded = computed(() => isResultsContext.value && props.pageLimit === undefined); const canGoNextPage = computed(() => { return canGoNextDataGridPage({ hasMore: props.result.has_more, @@ -1869,10 +1889,13 @@ const canGoNextPage = computed(() => { pageOffset: props.pageOffset, currentPage: currentPage.value, totalRowCount: hasKnownTotalRowCount.value ? displayedTotalRowCount.value : undefined, + allRowsLoaded: allRowsLoaded.value, }); }); const canJumpLastPage = computed( - () => canGoNextPage.value && (hasKnownTotalRowCount.value || !!props.tableMeta || !!props.countSql), + () => + canGoNextPage.value && + (hasKnownTotalRowCount.value || allRowsLoaded.value || !!props.tableMeta || !!props.countSql), ); const totalRowCountBusy = computed(() => props.totalRowCountLoading === true || manualTotalRowCountLoading.value); const canCalculateTotalRowCount = computed( @@ -2015,6 +2038,15 @@ async function lastPage() { emit("paginate", (lastPageNum - 1) * pageSize.value, pageSize.value, currentWhereInput(), currentOrderBy()); return; } + if (allRowsLoaded.value) { + const total = props.result.rows.length; + if (total <= 0) return; + const lastPageNum = Math.ceil(total / pageSize.value); + if (lastPageNum <= currentPage.value) return; + currentPage.value = lastPageNum; + resetGridVerticalScroll(true); + return; + } if (!props.connectionId) return; const countTarget = await buildCurrentCountTarget(); const sql = countTarget?.sql; diff --git a/apps/desktop/src/lib/dataGridPagination.ts b/apps/desktop/src/lib/dataGridPagination.ts index 255c2371a..c0f3edf22 100644 --- a/apps/desktop/src/lib/dataGridPagination.ts +++ b/apps/desktop/src/lib/dataGridPagination.ts @@ -5,20 +5,29 @@ export interface CanGoNextDataGridPageOptions { pageOffset?: number; currentPage?: number; totalRowCount?: number; + // True when every result row is already in memory (SQL editor result with no + // server-side pagination). rowCount IS the authoritative total in that case, + // so a full final page must not appear as "more available". + allRowsLoaded?: boolean; } export function canGoNextDataGridPage(options: CanGoNextDataGridPageOptions): boolean { if (options.hasMore === true) return true; const pageSize = Math.max(1, options.pageSize); + const currentOffset = + typeof options.pageOffset === "number" && options.pageOffset >= 0 + ? options.pageOffset + : Math.max(0, (options.currentPage ?? 1) - 1) * pageSize; + const totalRowCount = options.totalRowCount; if (typeof totalRowCount === "number" && Number.isFinite(totalRowCount) && totalRowCount >= 0) { - const currentOffset = - typeof options.pageOffset === "number" && options.pageOffset >= 0 - ? options.pageOffset - : Math.max(0, (options.currentPage ?? 1) - 1) * pageSize; return currentOffset + pageSize < totalRowCount; } + if (options.allRowsLoaded === true) { + return currentOffset + pageSize < options.rowCount; + } + return options.rowCount >= pageSize; } diff --git a/apps/desktop/src/lib/sqlCompletion.ts b/apps/desktop/src/lib/sqlCompletion.ts index abb97e53b..2f1ca4a25 100644 --- a/apps/desktop/src/lib/sqlCompletion.ts +++ b/apps/desktop/src/lib/sqlCompletion.ts @@ -1109,7 +1109,71 @@ export function shouldAutoOpenSqlCompletion(sql: string, cursor: number): boolea ) { return true; } - return /[\w$.]/.test(previousChar); + return /[\w$.@]/.test(previousChar); +} + +export function isSqlLikeCompletionStatement(sql: string, cursor: number): boolean { + const statement = extractStatementAt(sql, cursor).trimStart(); + if (/^(select|with)\b/i.test(statement)) return true; + return currentLineBlockStartsSql(sql, cursor); +} + +function currentLineBlockStartsSql(sql: string, cursor: number): boolean { + return currentSqlLikeLineBlockSpan(sql, cursor) != null; +} + +function currentSqlLikeLineBlockSpan(sql: string, cursor: number): { start: number; end: number } | null { + const safeCursor = Math.max(0, Math.min(cursor, sql.length)); + const beforeCursor = sql.slice(0, safeCursor); + const lines = beforeCursor.split(/\r?\n/); + let start: number | null = null; + let offset = 0; + + for (const line of lines) { + const trimmed = line.trimStart(); + if (trimmed) { + const indentation = line.length - trimmed.length; + if (/^(select|with)\b/i.test(trimmed)) start = offset + indentation; + if (/^(get|post|put|delete|patch|head)\s+\//i.test(trimmed)) start = null; + } + offset += line.length + 1; + } + + if (start == null) return null; + + let end = sql.length; + let inSingleQuote = false; + let inDoubleQuote = false; + for (let index = start; index < sql.length; index += 1) { + const ch = sql[index]; + if (ch === "'" && !inDoubleQuote) inSingleQuote = !inSingleQuote; + else if (ch === '"' && !inSingleQuote) inDoubleQuote = !inDoubleQuote; + else if (ch === ";" && !inSingleQuote && !inDoubleQuote && index >= safeCursor) { + end = index; + break; + } + } + + const blockEnd = currentLineBlockEnd(sql, safeCursor, start); + if (blockEnd != null) end = Math.min(end, blockEnd); + + return { start, end }; +} + +function currentLineBlockEnd(sql: string, cursor: number, start: number): number | null { + let lineStart = sql.lastIndexOf("\n", cursor - 1) + 1; + while (lineStart < sql.length) { + const lineEnd = sql.indexOf("\n", lineStart); + const boundedLineEnd = lineEnd >= 0 ? lineEnd : sql.length; + const line = sql.slice(lineStart, boundedLineEnd); + const trimmed = line.trimStart(); + if (lineStart > start && (!trimmed || /^(get|post|put|delete|patch|head)\s+\//i.test(trimmed))) { + return lineStart; + } + if (lineEnd < 0) break; + lineStart = lineEnd + 1; + } + return null; } export function getSqlCompletionResultValidFor(sql: string, cursor: number): RegExp | undefined { @@ -1144,6 +1208,9 @@ export function getSqlFunctionSignatureHelp(sql: string, cursor: number): SqlFun * Respects semicolons and string literals. */ function extractStatementStart(sql: string, cursor: number): number { + const lineBlock = currentSqlLikeLineBlockSpan(sql, cursor); + if (lineBlock) return lineBlock.start; + let start = 0; let inSingleQuote = false; let inDoubleQuote = false; @@ -1166,6 +1233,9 @@ function extractStatementStart(sql: string, cursor: number): number { * Respects semicolons and string literals. */ function extractStatementAt(sql: string, cursor: number): string { + const lineBlock = currentSqlLikeLineBlockSpan(sql, cursor); + if (lineBlock) return sql.slice(lineBlock.start, lineBlock.end).trim(); + const start = extractStatementStart(sql, cursor); let end = sql.length; let inSingleQuote = false; @@ -1364,11 +1434,11 @@ function parseTrailingIdentifierPart(input: string, endExclusive: number): { sta } let start = end; - while (start >= 0 && /[A-Za-z0-9_$]/.test(input[start] ?? "")) start -= 1; + while (start >= 0 && /[A-Za-z0-9_$@]/.test(input[start] ?? "")) start -= 1; start += 1; if (start >= endExclusive) return null; const raw = input.slice(start, endExclusive); - if (!/^[A-Za-z_][\w$]*$/.test(raw)) return null; + if (!/^[@A-Za-z_][\w$@]*$/.test(raw)) return null; return { start, raw }; } @@ -1599,21 +1669,31 @@ function extractReferencedTables(sql: string): SqlCompletionReferencedTable[] { ]); const pattern = - /\b(?:from|join|update|into|apply)\s+((?:"[^"]+"|`[^`]+`|[A-Za-z_][\w$]*)(?:\.(?:"[^"]+"|`[^`]+`|[A-Za-z_][\w$]*))?)(?:\s+(?:as\s+)?([A-Za-z_][\w$]*))?/gi; + /\b(?:from|join|update|into|apply)\s+((?:"[^"]+"|`[^`]+`|[^\s,;()]+)(?:\.(?:"[^"]+"|`[^`]+`|[^\s,;()]+))?)(?:\s+(?:as\s+)?([A-Za-z_][\w$]*))?/gi; const referenced: SqlCompletionReferencedTable[] = []; for (const match of sql.matchAll(pattern)) { const rawName = match[1]; const alias = match[2]; - const [first, second] = splitQualifiedName(rawName); - if (!first) continue; // Filter out SQL keywords that accidentally matched as aliases const cleanAlias = alias && !ALIAS_BLACKLIST.has(alias.toLowerCase()) ? alias : undefined; + if (isElasticsearchStyleIndexName(rawName)) { + referenced.push({ name: unquoteIdentifier(rawName), alias: cleanAlias }); + continue; + } + const [first, second] = splitQualifiedName(rawName); + if (!first) continue; const table = second ? { schema: first, name: second, alias: cleanAlias } : { name: first, alias: cleanAlias }; referenced.push(table); } return referenced; } +function isElasticsearchStyleIndexName(name: string | undefined): name is string { + if (!name) return false; + if ((name.startsWith('"') && name.endsWith('"')) || (name.startsWith("`") && name.endsWith("`"))) return false; + return /[-*]/.test(name); +} + function extractSelectAliases(sql: string): string[] { const selectList = extractSelectList(sql); if (!selectList) return []; @@ -2292,7 +2372,7 @@ function buildColumnApply( context: SqlCompletionContext, dialect?: "mysql" | "postgres" | "sqlserver", ): string { - if (context.qualifier || !column.displayLabel.includes(".")) { + if (context.qualifier || column.displayLabel === column.name || !column.displayLabel.includes(".")) { return quoteSqlIdentifier(column.name, dialect); } return `${quoteSqlIdentifier(column.table, dialect)}.${quoteSqlIdentifier(column.name, dialect)}`; diff --git a/apps/desktop/src/lib/sqlSemanticDiagnostics.ts b/apps/desktop/src/lib/sqlSemanticDiagnostics.ts index ba116ac29..c0dfa0064 100644 --- a/apps/desktop/src/lib/sqlSemanticDiagnostics.ts +++ b/apps/desktop/src/lib/sqlSemanticDiagnostics.ts @@ -1,6 +1,12 @@ import type { SqlCompletionColumn, SqlCompletionTable } from "@/lib/sqlCompletion"; import { getSqlCompletionContext } from "@/lib/sqlCompletion"; -import type { SqlColumnReference, SqlReferenceAnalysis, SqlTableReference, SqlTextSpan } from "@/types/database"; +import type { + DatabaseType, + SqlColumnReference, + SqlReferenceAnalysis, + SqlTableReference, + SqlTextSpan, +} from "@/types/database"; export interface SqlSemanticDiagnostic { span: SqlTextSpan; @@ -93,7 +99,12 @@ export function areSqlSemanticDiagnosticsEqual( }); } -export function shouldRunSqlSemanticDiagnostics(sql: string, cursor: number): boolean { +export function shouldRunSqlSemanticDiagnostics( + sql: string, + cursor: number, + options: { databaseType?: DatabaseType } = {}, +): boolean { + if (options.databaseType === "elasticsearch") return false; const context = getSqlCompletionContext(sql, cursor); if (context.suggestTables || context.exclusiveTableSuggestions || context.exclusiveColumnSuggestions) return false; if (context.qualifier) return false; diff --git a/apps/desktop/src/stores/queryStore.ts b/apps/desktop/src/stores/queryStore.ts index 98c91501e..1784f34eb 100644 --- a/apps/desktop/src/stores/queryStore.ts +++ b/apps/desktop/src/stores/queryStore.ts @@ -1270,6 +1270,20 @@ export const useQueryStore = defineStore("query", () => { current.resultTotalRowCount = undefined; } current.resultTotalRowCountLoading = current.mode === "query" && !!current.result && !!countSql; + // Server-side pagination without a countSql: the backend (currently + // the Elasticsearch driver) already reports the true match total via + // affected_rows. Use it directly so the result-grid can compute the + // page count without issuing a separate COUNT query. + if ( + current.result && + current.mode === "query" && + typeof pageLimit === "number" && + !countSql && + typeof current.result.affected_rows === "number" + ) { + current.resultTotalRowCount = current.result.affected_rows; + current.resultTotalRowCountLoading = false; + } touchResult(current); if (current.mode === "query" && current.result) { countQueryTotalRowsInBackground({ diff --git a/crates/dbx-core/src/db/elasticsearch_driver.rs b/crates/dbx-core/src/db/elasticsearch_driver.rs index d1eb53f55..80ccf492b 100644 --- a/crates/dbx-core/src/db/elasticsearch_driver.rs +++ b/crates/dbx-core/src/db/elasticsearch_driver.rs @@ -1,5 +1,7 @@ use reqwest::Client as HttpClient; use serde::Deserialize; +use serde_json::Value; +use std::collections::HashSet; use std::error::Error; use std::time::Duration; @@ -222,6 +224,90 @@ pub async fn list_indices(client: &EsClient) -> Result, String> { Ok(names) } +pub async fn get_columns(client: &EsClient, index: &str) -> Result, String> { + let path = format!("/{index}/_mapping"); + let resp = client.get(&path).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?; + if !resp.status().is_success() { + let body = resp.text().await.unwrap_or_default(); + return Err(format!("Elasticsearch error: {body}")); + } + + let body: Value = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?; + let mut seen = HashSet::new(); + let mut columns = Vec::new(); + + if let Some(indices) = body.as_object() { + for index_mapping in indices.values() { + if let Some(properties) = mapping_properties(index_mapping) { + collect_mapping_columns("", properties, &mut seen, &mut columns); + } + } + } + + columns.sort_by(|left, right| left.name.cmp(&right.name)); + Ok(columns) +} + +fn mapping_properties(mapping: &Value) -> Option<&serde_json::Map> { + if let Some(properties) = mapping.pointer("/mappings/properties").and_then(Value::as_object) { + return Some(properties); + } + + mapping + .get("mappings") + .and_then(Value::as_object)? + .values() + .find_map(|typed_mapping| typed_mapping.get("properties").and_then(Value::as_object)) +} + +fn collect_mapping_columns( + prefix: &str, + properties: &serde_json::Map, + seen: &mut HashSet, + columns: &mut Vec, +) { + for (name, definition) in properties { + let field_name = if prefix.is_empty() { name.clone() } else { format!("{prefix}.{name}") }; + let field_type = definition.get("type").and_then(Value::as_str); + + if let Some(data_type) = field_type { + push_mapping_column(&field_name, data_type, seen, columns); + } + + if let Some(fields) = definition.get("fields").and_then(Value::as_object) { + collect_mapping_columns(&field_name, fields, seen, columns); + } + + if let Some(children) = definition.get("properties").and_then(Value::as_object) { + collect_mapping_columns(&field_name, children, seen, columns); + } + } +} + +fn push_mapping_column( + name: &str, + data_type: &str, + seen: &mut HashSet, + columns: &mut Vec, +) { + if !seen.insert(name.to_string()) { + return; + } + + columns.push(crate::db::ColumnInfo { + name: name.to_string(), + data_type: data_type.to_string(), + is_nullable: true, + column_default: None, + is_primary_key: false, + extra: None, + comment: None, + numeric_precision: None, + numeric_scale: None, + character_maximum_length: None, + }); +} + #[derive(Deserialize)] struct SearchResponse { hits: SearchHits, @@ -330,6 +416,14 @@ pub async fn execute_rest_query(client: &EsClient, input: &str) -> Result Result Result { + let report_index_total = query.from_plan_pagination; + let path = format!("/{}/_search", query.index); + let resp = + client.post(&path).json(&query.body).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?; + let status = resp.status().as_u16(); + let body: serde_json::Value = resp.json().await.unwrap_or_else(|_| serde_json::Value::Null); + // Capture the index's true match total before the body is consumed by the + // parser — needed below when we report total instead of rows.len(). + let index_total = body.pointer("/hits/total/value").and_then(|v| v.as_u64()); + + let mut result = parse_elasticsearch_response(status, body, start)?; + if report_index_total { + if let Some(total) = index_total { + result.affected_rows = total; + } + } + Ok(result) +} + +fn parse_elasticsearch_response( + status: u16, + body: serde_json::Value, + start: std::time::Instant, +) -> Result { + if let Some(result) = parse_sql_response(&body, start) { + Ok(result) + } else if let Some(hits) = body.pointer("/hits/hits").and_then(|v| v.as_array()).filter(|h| !h.is_empty()) { let mut all_keys = Vec::::new(); let docs: Vec> = hits .iter() @@ -417,13 +562,13 @@ pub async fn execute_rest_query(client: &EsClient, input: &str) -> Result Result Option { + let mut cursor = skip_sql_whitespace(input, 0); + cursor = consume_sql_keyword(input, cursor, "select")?; + cursor = skip_sql_whitespace(input, cursor); + if next_char_at(input, cursor)? != '*' { + return None; + } + cursor += '*'.len_utf8(); + cursor = skip_sql_whitespace(input, cursor); + cursor = consume_sql_keyword(input, cursor, "from")?; + cursor = skip_sql_whitespace(input, cursor); + + let (index, next_cursor) = read_sql_token(input, cursor)?; + cursor = next_cursor; + + let mut sort_field = None; + let mut sort_order = "asc"; + let mut limit = None; + let mut offset: Option = None; + + loop { + cursor = skip_sql_whitespace(input, cursor); + if cursor >= input.len() { + break; + } + if next_char_at(input, cursor) == Some(';') { + cursor += ';'.len_utf8(); + cursor = skip_sql_whitespace(input, cursor); + if cursor == input.len() { + break; + } + return None; + } + + if is_keyword_at(input, cursor, "order") { + cursor = consume_sql_keyword(input, cursor, "order")?; + cursor = skip_sql_whitespace(input, cursor); + cursor = consume_sql_keyword(input, cursor, "by")?; + cursor = skip_sql_whitespace(input, cursor); + let (field, next_cursor) = read_sql_token(input, cursor)?; + sort_field = Some(field); + cursor = skip_sql_whitespace(input, next_cursor); + if is_keyword_at(input, cursor, "asc") { + sort_order = "asc"; + cursor = consume_sql_keyword(input, cursor, "asc")?; + } else if is_keyword_at(input, cursor, "desc") { + sort_order = "desc"; + cursor = consume_sql_keyword(input, cursor, "desc")?; + } + } else if is_keyword_at(input, cursor, "limit") { + cursor = consume_sql_keyword(input, cursor, "limit")?; + cursor = skip_sql_whitespace(input, cursor); + let (value, next_cursor) = read_while(input, cursor, |ch| ch.is_ascii_digit()); + limit = value.parse::().ok(); + cursor = next_cursor; + } else if is_keyword_at(input, cursor, "offset") { + cursor = consume_sql_keyword(input, cursor, "offset")?; + cursor = skip_sql_whitespace(input, cursor); + let (value, next_cursor) = read_while(input, cursor, |ch| ch.is_ascii_digit()); + offset = value.parse::().ok(); + cursor = next_cursor; + } else { + return None; + } + } + + // The pagination plan emits `LIMIT N OFFSET M` (always with OFFSET, even + // when 0) for ES; a user-written SQL that only has `LIMIT N` leaves + // OFFSET absent. We use that as the signal for whether the front-end is + // driving server-side pagination — in that case affected_rows must reflect + // the index's true total so the grid can compute the total page count. + let from_plan_pagination = offset.is_some(); + let effective_size = limit.unwrap_or(AUTO_PAGED_SELECT_STAR_SIZE); + let effective_from = offset.unwrap_or(0); + let mut body = serde_json::Map::new(); + body.insert("size".to_string(), serde_json::json!(effective_size)); + if effective_from > 0 { + body.insert("from".to_string(), serde_json::json!(effective_from)); + } + + if let Some(field) = sort_field { + let mut sort_item = serde_json::Map::new(); + sort_item.insert(field, serde_json::json!({ "order": sort_order })); + body.insert("sort".to_string(), serde_json::Value::Array(vec![serde_json::Value::Object(sort_item)])); + } + + Some(ElasticsearchSearchQuery { index, body: serde_json::Value::Object(body), from_plan_pagination }) +} + +fn is_elasticsearch_sql_query(input: &str) -> bool { + input + .trim_start() + .split_once(char::is_whitespace) + .map(|(keyword, _)| keyword.eq_ignore_ascii_case("select")) + .unwrap_or_else(|| input.trim_start().eq_ignore_ascii_case("select")) +} + +async fn execute_sql_query( + client: &EsClient, + query: &str, + start: std::time::Instant, +) -> Result { + let query = adapt_elasticsearch_sql_query(query); + let body = serde_json::json!({ "query": query }); + let resp = + client.post("/_sql").json(&body).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?; + let status = resp.status(); + let response_body: serde_json::Value = resp.json().await.unwrap_or_else(|_| serde_json::Value::Null); + + if !status.is_success() { + return Err(format_sql_error(status, &response_body)); + } + + parse_sql_response(&response_body, start).ok_or_else(|| { + let pretty = serde_json::to_string_pretty(&response_body).unwrap_or_else(|_| response_body.to_string()); + format!("Unexpected Elasticsearch SQL response: {pretty}") + }) +} + +fn adapt_elasticsearch_sql_query(query: &str) -> String { + let mut output = String::with_capacity(query.len()); + let mut index = 0; + let mut state = SqlScanState::Normal; + + while let Some(ch) = next_char_at(query, index) { + match state { + SqlScanState::Normal => match ch { + '\'' => { + output.push(ch); + index += ch.len_utf8(); + state = SqlScanState::SingleQuoted; + } + '"' => { + output.push(ch); + index += ch.len_utf8(); + state = SqlScanState::DoubleQuoted; + } + '`' => { + output.push(ch); + index += ch.len_utf8(); + state = SqlScanState::BacktickQuoted; + } + '-' if query[index..].starts_with("--") => { + output.push_str("--"); + index += 2; + state = SqlScanState::LineComment; + } + '/' if query[index..].starts_with("/*") => { + output.push_str("/*"); + index += 2; + state = SqlScanState::BlockComment; + } + '@' if is_at_identifier_boundary(&output) => { + let (identifier, next_index) = read_while(query, index, is_elasticsearch_identifier_part); + output.push('"'); + output.push_str(identifier); + output.push('"'); + index = next_index; + } + _ => { + if let Some(keyword) = relation_keyword_at(query, index) { + index = quote_relation_after_keyword(query, index, keyword, &mut output); + } else { + output.push(ch); + index += ch.len_utf8(); + } + } + }, + SqlScanState::SingleQuoted => { + if copy_quoted_char(query, &mut index, ch, '\'', &mut output) { + state = SqlScanState::Normal; + } + } + SqlScanState::DoubleQuoted => { + if copy_quoted_char(query, &mut index, ch, '"', &mut output) { + state = SqlScanState::Normal; + } + } + SqlScanState::BacktickQuoted => { + if copy_quoted_char(query, &mut index, ch, '`', &mut output) { + state = SqlScanState::Normal; + } + } + SqlScanState::LineComment => { + output.push(ch); + index += ch.len_utf8(); + if ch == '\n' { + state = SqlScanState::Normal; + } + } + SqlScanState::BlockComment => { + if query[index..].starts_with("*/") { + output.push_str("*/"); + index += 2; + state = SqlScanState::Normal; + } else { + output.push(ch); + index += ch.len_utf8(); + } + } + } + } + + output +} + +fn quote_relation_after_keyword(query: &str, index: usize, keyword: &str, output: &mut String) -> usize { + let mut cursor = index + keyword.len(); + output.push_str(&query[index..cursor]); + + while let Some(ch) = next_char_at(query, cursor) { + if !ch.is_whitespace() { + break; + } + output.push(ch); + cursor += ch.len_utf8(); + } + + if matches!(next_char_at(query, cursor), Some('"' | '`' | '\'' | '(')) { + return cursor; + } + + let relation_start = cursor; + while let Some(ch) = next_char_at(query, cursor) { + if !is_relation_name_char(ch) { + break; + } + cursor += ch.len_utf8(); + } + + let relation = &query[relation_start..cursor]; + if relation_name_needs_quotes(relation) { + output.push('"'); + output.push_str(relation); + output.push('"'); + } else { + output.push_str(relation); + } + + cursor +} + +fn copy_quoted_char(query: &str, index: &mut usize, ch: char, quote: char, output: &mut String) -> bool { + output.push(ch); + *index += ch.len_utf8(); + + if ch != quote { + return false; + } + + if next_char_at(query, *index).is_some_and(|next| next == quote) { + output.push(quote); + *index += quote.len_utf8(); + false + } else { + true + } +} + +fn read_while(query: &str, start: usize, predicate: fn(char) -> bool) -> (&str, usize) { + let mut cursor = start; + while let Some(ch) = next_char_at(query, cursor) { + if !predicate(ch) { + break; + } + cursor += ch.len_utf8(); + } + + (&query[start..cursor], cursor) +} + +fn skip_sql_whitespace(query: &str, mut cursor: usize) -> usize { + while let Some(ch) = next_char_at(query, cursor) { + if !ch.is_whitespace() { + break; + } + cursor += ch.len_utf8(); + } + + cursor +} + +fn consume_sql_keyword(query: &str, cursor: usize, keyword: &str) -> Option { + is_keyword_at(query, cursor, keyword).then_some(cursor + keyword.len()) +} + +fn read_sql_token(query: &str, cursor: usize) -> Option<(String, usize)> { + let quote = match next_char_at(query, cursor)? { + '"' => Some('"'), + '`' => Some('`'), + _ => None, + }; + + if let Some(quote) = quote { + let mut output = String::new(); + let mut next_cursor = cursor + quote.len_utf8(); + while let Some(ch) = next_char_at(query, next_cursor) { + next_cursor += ch.len_utf8(); + if ch == quote { + if next_char_at(query, next_cursor).is_some_and(|next| next == quote) { + output.push(quote); + next_cursor += quote.len_utf8(); + } else { + return Some((output, next_cursor)); + } + } else { + output.push(ch); + } + } + return None; + } + + let (token, next_cursor) = read_while(query, cursor, is_relation_name_char); + (!token.is_empty()).then(|| (token.to_string(), next_cursor)) +} + +fn relation_keyword_at(query: &str, index: usize) -> Option<&'static str> { + ["from", "join"].into_iter().find(|keyword| is_keyword_at(query, index, keyword)) +} + +#[derive(Clone, Copy)] +enum SqlScanState { + Normal, + SingleQuoted, + DoubleQuoted, + BacktickQuoted, + LineComment, + BlockComment, +} + +fn is_at_identifier_boundary(output: &str) -> bool { + output.chars().next_back().is_none_or(|ch| !is_sql_identifier_part(ch)) +} + +fn is_sql_identifier_part(ch: char) -> bool { + ch.is_ascii_alphanumeric() || matches!(ch, '_' | '.') +} + +fn is_elasticsearch_identifier_part(ch: char) -> bool { + ch.is_ascii_alphanumeric() || matches!(ch, '_' | '.' | '-' | '@') +} + +fn is_relation_name_char(ch: char) -> bool { + !ch.is_whitespace() && !matches!(ch, ',' | ';' | '(' | ')') +} + +fn relation_name_needs_quotes(relation: &str) -> bool { + relation.chars().any(|ch| matches!(ch, '-' | '*' | '@')) +} + +fn is_keyword_at(query: &str, index: usize, keyword: &str) -> bool { + query.get(index..index + keyword.len()).is_some_and(|candidate| candidate.eq_ignore_ascii_case(keyword)) + && query[..index].chars().next_back().is_none_or(|ch| !is_keyword_boundary_char(ch)) + && query[index + keyword.len()..].chars().next().is_none_or(|ch| !is_keyword_boundary_char(ch)) +} + +fn is_keyword_boundary_char(ch: char) -> bool { + ch.is_ascii_alphanumeric() || ch == '_' +} + +fn next_char_at(query: &str, index: usize) -> Option { + query.get(index..)?.chars().next() +} + +fn parse_sql_response(body: &serde_json::Value, start: std::time::Instant) -> Option { + let columns = body.get("columns")?.as_array()?; + let rows = body.get("rows")?.as_array()?; + let column_names: Vec = columns + .iter() + .filter_map(|column| column.get("name").and_then(|name| name.as_str()).map(str::to_string)) + .collect(); + + if column_names.is_empty() && !columns.is_empty() { + return None; + } + + let result_rows: Vec> = + rows.iter().filter_map(|row| row.as_array().map(|values| values.to_vec())).collect(); + + Some(crate::types::QueryResult { + columns: column_names, + column_types: Vec::new(), + rows: result_rows, + affected_rows: rows.len() as u64, + execution_time_ms: start.elapsed().as_millis(), + truncated: false, + session_id: body.get("cursor").and_then(|cursor| cursor.as_str()).map(str::to_string), + has_more: body.get("cursor").and_then(|cursor| cursor.as_str()).is_some(), + }) +} + +fn format_sql_error(status: reqwest::StatusCode, body: &serde_json::Value) -> String { + let detail = body + .pointer("/error/reason") + .and_then(|reason| reason.as_str()) + .map(str::to_string) + .unwrap_or_else(|| serde_json::to_string_pretty(body).unwrap_or_else(|_| body.to_string())); + + if status == reqwest::StatusCode::NOT_FOUND { + format!("Elasticsearch SQL API is not available ({status}): {detail}") + } else { + format!("Elasticsearch SQL error ({status}): {detail}") + } +} + fn parse_aggregations(aggs: &serde_json::Map) -> (Vec, Vec>) { for (_name, agg_value) in aggs { if let Some(buckets) = agg_value.get("buckets").and_then(|b| b.as_array()) { diff --git a/crates/dbx-core/src/query_result_sql.rs b/crates/dbx-core/src/query_result_sql.rs index fd777fb8e..1666be7ce 100644 --- a/crates/dbx-core/src/query_result_sql.rs +++ b/crates/dbx-core/src/query_result_sql.rs @@ -178,6 +178,19 @@ pub fn build_paginated_query_sql(options: PaginatedQuerySqlOptions) -> QuerySqlB return ok(add_mysql_limit(&statement, safe_limit, safe_offset)); } + if options.database_type == Some(DatabaseType::Elasticsearch) { + // If the user wrote their own LIMIT, leave the SQL alone — they + // explicitly bounded the result set and the front-end will paginate + // client-side. Otherwise wrap with an explicit OFFSET (even when + // 0) so the ES driver can tell a plan-wrapped query from one the + // user wrote, which decides whether affected_rows should reflect + // the index total or the row count we actually returned. + if has_top_level_limit(&statement) { + return err("unsupported"); + } + return ok(format!("{statement} LIMIT {safe_limit} OFFSET {safe_offset};")); + } + if options.database_type.is_some_and(uses_fetch_first) { return ok(add_fetch_first_limit(&statement, safe_limit, safe_offset)); } @@ -192,6 +205,11 @@ pub fn build_count_query_sql(options: CountQuerySqlOptions) -> QuerySqlBuildResu if unsupported_pagination_type(options.database_type) { return err("unsupported"); } + // ES SQL can't wrap a SELECT in `SELECT COUNT(*) FROM (...)` — the + // driver already reports the true match count via affected_rows. + if options.database_type == Some(DatabaseType::Elasticsearch) { + return err("unsupported"); + } let statement = LIMIT_OFFSET_STRIP_RE.replace(&statement, "").to_string(); @@ -274,10 +292,7 @@ fn err(reason: &str) -> QuerySqlBuildResult { } fn unsupported_pagination_type(database_type: Option) -> bool { - matches!( - database_type, - Some(DatabaseType::Neo4j | DatabaseType::MongoDb | DatabaseType::Redis | DatabaseType::Elasticsearch) - ) + matches!(database_type, Some(DatabaseType::Neo4j | DatabaseType::MongoDb | DatabaseType::Redis)) } fn single_selectable_statement(original_sql: &str) -> Result { diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index 88c9009b6..f7ed6125c 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -1297,6 +1297,9 @@ pub async fn get_columns_core( PoolKind::Rqlite(client) => { db::rqlite_driver::get_columns(client, schema, table).await.map(deduplicate_column_infos) } + PoolKind::Elasticsearch(client) => { + db::elasticsearch_driver::get_columns(client, table).await.map(deduplicate_column_infos) + } _ => Ok(vec![]), } } diff --git a/docs/screenshot-es-sql-completion.png b/docs/screenshot-es-sql-completion.png new file mode 100644 index 000000000..9ca465ef9 Binary files /dev/null and b/docs/screenshot-es-sql-completion.png differ diff --git a/docs/screenshot-es-sql-pagination.png b/docs/screenshot-es-sql-pagination.png new file mode 100644 index 000000000..e4acc206a Binary files /dev/null and b/docs/screenshot-es-sql-pagination.png differ