diff --git a/apps/desktop/src/composables/useDataGridActions.ts b/apps/desktop/src/composables/useDataGridActions.ts index 466d8126e..f5f30e85e 100644 --- a/apps/desktop/src/composables/useDataGridActions.ts +++ b/apps/desktop/src/composables/useDataGridActions.ts @@ -13,6 +13,7 @@ import { effectiveDatabaseTypeForConnection, metadataSchemaForConnection } from import { applyMongoFindSort } from "@/lib/mongo/mongoShellCommand"; import { uuid } from "@/lib/common/utils"; import type { DataGridSortMode } from "@/lib/dataGrid/dataGridSort"; +import { queryResultBaseSql, queryResultExecutionSql } from "@/lib/tabs/tabPresentation"; const DATA_TAB_METADATA_TTL_MS = 30_000; @@ -172,16 +173,18 @@ export function useDataGridActions(activeTab: ComputedRef) const tab = activeTab.value; if (!tab) return; if (tab.mode !== "data") { - const baseSql = tab.resultSortedSql ?? tab.resultBaseSql ?? tab.lastExecutedSql ?? tab.sql; + const baseSql = queryResultExecutionSql(tab); if (!baseSql.trim()) return; const expectedNextOffset = (tab.resultPageOffset ?? 0) + (tab.resultPageLimit ?? limit); const sessionId = tab.result?.has_more && tab.result?.session_id && offset === expectedNextOffset && limit === tab.resultPageLimit ? tab.result.session_id : undefined; + const resultBaseSql = queryResultBaseSql(tab); await queryStore.executeTabSql(tab.id, baseSql, { - resultBaseSql: tab.resultBaseSql ?? tab.sql, + resultBaseSql, resultSortedSql: tab.resultSortedSql, pagination: { offset, limit, sessionId }, preserveResultDuringExecution: true, preserveTotalRowCountDuringExecution: true, + replaceActiveResultInGroup: true, }); return; } @@ -225,7 +228,7 @@ export function useDataGridActions(activeTab: ComputedRef) return; } - const baseSql = tab.resultBaseSql ?? tab.sql; + const baseSql = queryResultBaseSql(tab); if (!baseSql.trim()) return; if (!direction) { @@ -234,6 +237,7 @@ export function useDataGridActions(activeTab: ComputedRef) resultSortedSql: undefined, preserveResultDuringExecution: true, preserveTotalRowCountDuringExecution: true, + replaceActiveResultInGroup: true, }); return; } @@ -251,6 +255,7 @@ export function useDataGridActions(activeTab: ComputedRef) resultSortedSql: sortedSql, preserveResultDuringExecution: true, preserveTotalRowCountDuringExecution: true, + replaceActiveResultInGroup: true, }); return; } @@ -273,6 +278,7 @@ export function useDataGridActions(activeTab: ComputedRef) resultSortedSql: built.sql, preserveResultDuringExecution: true, preserveTotalRowCountDuringExecution: true, + replaceActiveResultInGroup: true, }); } diff --git a/apps/desktop/src/lib/__tests__/tabs/tabPresentation.spec.ts b/apps/desktop/src/lib/__tests__/tabs/tabPresentation.spec.ts new file mode 100644 index 000000000..5c98945a9 --- /dev/null +++ b/apps/desktop/src/lib/__tests__/tabs/tabPresentation.spec.ts @@ -0,0 +1,53 @@ +import { describe, expect, it } from "vitest"; +import { queryResultBaseSql, queryResultExecutionSql } from "@/lib/tabs/tabPresentation"; +import type { QueryTab } from "@/types/database"; + +function queryTab(overrides: Partial): QueryTab { + return { + id: "tab-1", + title: "SQL", + connectionId: "conn-1", + database: "db", + sql: "SELECT * FROM dbo.first;\nSELECT * FROM dbo.second;", + originalSql: "", + isExecuting: false, + isCancelling: false, + isExplaining: false, + mode: "query", + ...overrides, + } as QueryTab; +} + +describe("query result SQL selection", () => { + it("uses the active result source statement for multi-result query actions", () => { + const tab = queryTab({ + resultBaseSql: "SELECT * FROM dbo.first;\nSELECT * FROM dbo.second;", + result: { + columns: ["id"], + rows: [[1]], + affected_rows: 0, + execution_time_ms: 1, + sourceStatement: "SELECT * FROM dbo.second", + }, + }); + + expect(queryResultBaseSql(tab)).toBe("SELECT * FROM dbo.second"); + expect(queryResultExecutionSql(tab)).toBe("SELECT * FROM dbo.second"); + }); + + it("prefers the sorted SQL when the active result is sorted", () => { + const tab = queryTab({ + resultSortedSql: "SELECT * FROM dbo.second ORDER BY id DESC", + result: { + columns: ["id"], + rows: [[2]], + affected_rows: 0, + execution_time_ms: 1, + sourceStatement: "SELECT * FROM dbo.second", + }, + }); + + expect(queryResultBaseSql(tab)).toBe("SELECT * FROM dbo.second"); + expect(queryResultExecutionSql(tab)).toBe("SELECT * FROM dbo.second ORDER BY id DESC"); + }); +}); diff --git a/apps/desktop/src/lib/tabs/tabPresentation.ts b/apps/desktop/src/lib/tabs/tabPresentation.ts index f8e55c731..c6ad882ed 100644 --- a/apps/desktop/src/lib/tabs/tabPresentation.ts +++ b/apps/desktop/src/lib/tabs/tabPresentation.ts @@ -150,6 +150,14 @@ export function resultSqlForGrid(tab: Pick): string { + return resultSqlForGrid(tab); +} + +export function queryResultExecutionSql(tab: Pick): string { + return tab.resultSortedSql || resultSqlForGrid(tab); +} + export function tabularResultItems(results: QueryResult[] | undefined): { result: QueryResult; index: number; n: number; label?: string; title?: string }[] { if (!results) return []; return results diff --git a/apps/desktop/src/stores/queryStore.ts b/apps/desktop/src/stores/queryStore.ts index d55356c1e..afb0660bb 100644 --- a/apps/desktop/src/stores/queryStore.ts +++ b/apps/desktop/src/stores/queryStore.ts @@ -40,6 +40,7 @@ import { normalizeResultPageSize } from "@/lib/dataGrid/paginationPageSize"; import { splitSqlStatementRanges } from "@/lib/sql/sqlStatementRanges"; import { clearDataGridPendingSnapshotsForTab } from "@/composables/useDataGridEditor"; import { buildTabResultSnapshot, deleteTabResultSnapshot, readTabResultSnapshot, tabResultCacheKey, writeTabResultSnapshot } from "@/lib/tabs/tabResultCache"; +import { queryResultBaseSql, queryResultExecutionSql } from "@/lib/tabs/tabPresentation"; import { decodeQueryResultArchive, encodeQueryResultArchive, type DecodedQueryResultArchive } from "@/lib/query/queryResultArchive"; import * as api from "@/lib/backend/api"; import { useConnectionStore } from "@/stores/connectionStore"; @@ -2159,6 +2160,7 @@ export const useQueryStore = defineStore("query", () => { mongoSafety?: MongoAggregateSafetyOptions; preserveResultDuringExecution?: boolean; preserveTotalRowCountDuringExecution?: boolean; + replaceActiveResultInGroup?: boolean; skipRedisSafetyCheck?: boolean; sourceTraceId?: string; skipEnsureConnected?: boolean; @@ -2588,7 +2590,14 @@ export const useQueryStore = defineStore("query", () => { }); const current = tabs.value.find((t) => t.id === id); if (current?.executionId === executionId) { - if (results.length > 1) { + const activeGroupIndex = current.activeResultIndex; + const activeGroupResults = current.results; + const shouldReplaceActiveResultInGroup = options?.replaceActiveResultInGroup === true && results.length === 1 && Array.isArray(activeGroupResults) && typeof activeGroupIndex === "number" && activeGroupIndex >= 0 && activeGroupIndex < activeGroupResults.length; + if (shouldReplaceActiveResultInGroup) { + current.results = activeGroupResults.slice(); + current.results[activeGroupIndex] = results[0]; + current.result = results[0]; + } else if (results.length > 1) { const activeResultIndex = results.findIndex((result) => result.columns.length > 0); const resultIndex = activeResultIndex >= 0 ? activeResultIndex : 0; current.results = results; @@ -2667,9 +2676,19 @@ export const useQueryStore = defineStore("query", () => { } const current = tabs.value.find((t) => t.id === id); if (current?.executionId === executionId) { - current.result = toErrorResult(e); - current.results = undefined; - current.activeResultIndex = undefined; + const errorResult = toErrorResult(e); + const activeGroupIndex = current.activeResultIndex; + const activeGroupResults = current.results; + const shouldReplaceActiveResultInGroup = options?.replaceActiveResultInGroup === true && Array.isArray(activeGroupResults) && typeof activeGroupIndex === "number" && activeGroupIndex >= 0 && activeGroupIndex < activeGroupResults.length; + if (shouldReplaceActiveResultInGroup) { + current.results = activeGroupResults.slice(); + current.results[activeGroupIndex] = errorResult; + current.result = errorResult; + } else { + current.result = errorResult; + current.results = undefined; + current.activeResultIndex = undefined; + } current.queryAnalysis = undefined; current.querySourceColumns = undefined; current.queryEditabilityReason = undefined; @@ -2867,6 +2886,7 @@ export const useQueryStore = defineStore("query", () => { tab.resultSortColumnIndex = undefined; tab.resultSortDirection = undefined; tab.resultSortMode = undefined; + tab.resultSortedSql = undefined; touchResult(tab); tab.queryAnalysis = undefined; tab.querySourceColumns = undefined; @@ -3131,7 +3151,7 @@ export const useQueryStore = defineStore("query", () => { if (tab.mode !== "query") return tab.result; - const sql = tab.resultSortedSql ?? tab.resultBaseSql ?? tab.lastExecutedSql ?? tab.sql; + const sql = queryResultExecutionSql(tab); if (!sql.trim()) return tab.result; const connStore = useConnectionStore(); @@ -3140,7 +3160,7 @@ export const useQueryStore = defineStore("query", () => { const effectiveDbType = effectiveDatabaseTypeForConnection(conn); const queryTimeoutSecs = queryTimeoutSecsForConnection(conn); const useAgentCursor = usesAgentCursorForQuery(conn?.db_type); - const queryBaseSql = tab.resultBaseSql ?? sql; + const queryBaseSql = queryResultBaseSql(tab); const exportSettings = useSettingsStore().editorSettings; const exportRowLimit = exportSettings.exportRowLimitEnabled ? exportSettings.exportRowLimit : Number.POSITIVE_INFINITY; const agentExportMaxRows = exportSettings.exportRowLimitEnabled ? exportSettings.exportRowLimit : 2_147_483_647; @@ -3210,7 +3230,7 @@ export const useQueryStore = defineStore("query", () => { const tab = tabs.value.find((t) => t.id === id); if (!tab?.result || tab.mode !== "query") return undefined; - const sql = tab.resultSortedSql ?? tab.resultBaseSql ?? tab.lastExecutedSql ?? tab.sql; + const sql = queryResultExecutionSql(tab); if (!sql.trim()) return undefined; const connStore = useConnectionStore(); @@ -3220,7 +3240,7 @@ export const useQueryStore = defineStore("query", () => { const effectiveDbType = effectiveDatabaseTypeForConnection(conn); if (!effectiveDbType) return undefined; const useAgentCursor = usesAgentCursorForQuery(conn?.db_type); - const queryBaseSql = tab.resultBaseSql ?? sql; + const queryBaseSql = queryResultBaseSql(tab); const rowLimit = settings.exportRowLimitEnabled ? settings.exportRowLimit : null; const totalRows = typeof tab.resultTotalRowCount === "number" ? (rowLimit === null ? tab.resultTotalRowCount : Math.min(tab.resultTotalRowCount, rowLimit)) : null; const clientSessionId = tabClientSessionId(tab, "export"); diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index 7a7d5e429..b5d7c5cac 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -852,7 +852,12 @@ fn is_os_connection_error(lower: &str) -> bool { } pub fn timeout_error() -> String { - format!("Query timed out after {} seconds", QUERY_TIMEOUT.as_secs()) + timeout_error_for(QUERY_TIMEOUT) +} + +fn timeout_error_for(timeout_duration: Duration) -> String { + let seconds = timeout_duration.as_secs().max(1); + format!("Query timed out after {seconds} seconds") } pub fn canceled_error() -> String { @@ -882,15 +887,26 @@ pub async fn wait_for_query_with_timeout( ) -> Result where F: Future>, +{ + wait_for_result_with_timeout(cancel_token, timeout_duration, future).await +} + +async fn wait_for_result_with_timeout( + cancel_token: Option, + timeout_duration: Duration, + future: F, +) -> Result +where + F: Future>, { if let Some(token) = cancel_token { tokio::select! { biased; _ = token.cancelled() => Err(canceled_error()), - result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error())?, + result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error_for(timeout_duration))?, } } else { - timeout(timeout_duration, future).await.map_err(|_| timeout_error())? + timeout(timeout_duration, future).await.map_err(|_| timeout_error_for(timeout_duration))? } } @@ -903,9 +919,20 @@ pub async fn wait_for_query_opt( ) -> Result where F: Future>, +{ + wait_for_result_opt(cancel_token, timeout_duration, future).await +} + +async fn wait_for_result_opt( + cancel_token: Option, + timeout_duration: Option, + future: F, +) -> Result +where + F: Future>, { match timeout_duration { - Some(d) => wait_for_query_with_timeout(cancel_token, d, future).await, + Some(d) => wait_for_result_with_timeout(cancel_token, d, future).await, None => match cancel_token { Some(token) => { tokio::select! { @@ -919,6 +946,48 @@ where } } +async fn wait_for_value_opt( + cancel_token: Option, + timeout_duration: Option, + future: F, +) -> Result +where + F: Future, +{ + match timeout_duration { + Some(timeout_duration) => { + if let Some(token) = cancel_token { + tokio::select! { + biased; + _ = token.cancelled() => Err(canceled_error()), + result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error_for(timeout_duration)), + } + } else { + timeout(timeout_duration, future).await.map_err(|_| timeout_error_for(timeout_duration)) + } + } + None => match cancel_token { + Some(token) => { + tokio::select! { + biased; + _ = token.cancelled() => Err(canceled_error()), + result = future => Ok(result), + } + } + None => Ok(future.await), + }, + } +} + +async fn sqlserver_pool_is_current( + state: &AppState, + pool_key: &str, + client: &Arc>, +) -> bool { + let connections = state.connections.read().await; + matches!(connections.get(pool_key), Some(PoolKind::SqlServer(current)) if Arc::ptr_eq(current, client)) +} + fn resolve_query_timeout(timeout_secs: Option) -> Option { match timeout_secs { Some(0) => None, @@ -1867,6 +1936,7 @@ async fn execute_multi_sqlserver( check_read_only_for_connection_multi(state, pool_key, &batches).await?; let mut all_results = Vec::new(); let max_rows = options.max_rows; + let query_timeout = resolve_query_timeout(options.timeout_secs); for batch in &batches { if is_canceled(&cancel_token) { @@ -1892,17 +1962,28 @@ async fn execute_multi_sqlserver( }; drop(connections); - let mut client = match cancel_token.as_ref() { - Some(token) => tokio::select! { - biased; - _ = token.cancelled() => return Err(canceled_error()), - guard = client.lock() => guard, - }, - None => client.lock().await, + let mut client_guard = match wait_for_value_opt(cancel_token.clone(), query_timeout, client.lock()).await { + Ok(guard) => guard, + Err(err) => { + all_results.push(error_query_result(err)); + break; + } }; - let result = db::sqlserver::execute_batch_with_max_rows(&mut client, batch, max_rows).await; - drop(client); + if !sqlserver_pool_is_current(state, pool_key, &client).await { + all_results.push(error_query_result( + "SQL Server connection was reset while waiting for the query lock; please retry.".to_string(), + )); + break; + } + + let result = wait_for_result_opt( + cancel_token.clone(), + query_timeout, + db::sqlserver::execute_batch_with_max_rows(&mut client_guard, batch, max_rows), + ) + .await; + drop(client_guard); match result { Ok(results) => all_results.extend(results), @@ -2887,6 +2968,7 @@ pub async fn rollback_manual_transaction(state: &AppState, txn_session_id: &str) mod tests { use super::*; use crate::models::connection::{default_redis_key_separator, ConnectionConfig, DatabaseType}; + #[cfg(feature = "duckdb-bundled")] use crate::storage::Storage; fn test_connection_config(db_type: DatabaseType) -> ConnectionConfig { @@ -3025,7 +3107,29 @@ mod tests { }) .await; - assert_eq!(result.unwrap_err(), timeout_error()); + assert_eq!(result.unwrap_err(), timeout_error_for(Duration::from_millis(10))); + } + + #[tokio::test] + async fn wait_for_value_opt_times_out_while_waiting_for_lock() { + let lock = tokio::sync::Mutex::new(()); + let _guard = lock.lock().await; + + let result = wait_for_value_opt(None, Some(Duration::from_millis(10)), lock.lock()).await; + + assert_eq!(result.unwrap_err(), timeout_error_for(Duration::from_millis(10))); + } + + #[tokio::test] + async fn wait_for_value_opt_can_cancel_while_waiting_for_lock() { + let lock = tokio::sync::Mutex::new(()); + let _guard = lock.lock().await; + let token = CancellationToken::new(); + token.cancel(); + + let result = wait_for_value_opt(Some(token), Some(Duration::from_secs(30)), lock.lock()).await; + + assert_eq!(result.unwrap_err(), QUERY_CANCELED); } #[cfg(feature = "duckdb-bundled")]