From e2ecb0ccff4a1f3fb0590300fe1a248b2b74b864 Mon Sep 17 00:00:00 2001 From: Guoyu Su Date: Sat, 25 Jul 2026 09:24:08 +0800 Subject: [PATCH] fix(jdbc): use query cursors for table exports --- crates/dbx-core/src/table_export.rs | 750 +++++++++++++++--- .../main/java/app/dbx/jdbc/DbxJdbcPlugin.java | 100 ++- .../java/app/dbx/jdbc/DbxJdbcPluginTest.java | 144 ++++ 3 files changed, 886 insertions(+), 108 deletions(-) diff --git a/crates/dbx-core/src/table_export.rs b/crates/dbx-core/src/table_export.rs index bf2799ef6..097b80e24 100644 --- a/crates/dbx-core/src/table_export.rs +++ b/crates/dbx-core/src/table_export.rs @@ -15,6 +15,7 @@ use crate::database_export::{ }; use crate::db::agent_driver::AgentTableReadStartParams; use crate::models::connection::DatabaseType; +use crate::query::{close_query_session, execute_sql_statement_with_options, QueryExecutionOptions}; use crate::transfer::{ count_sql_with_where, execute_read_on_pool, execute_read_on_pool_with_max_rows, keyset_pagination_sql, pagination_sql_with_filter_order, qualified_table, quote_identifier, @@ -258,9 +259,94 @@ fn is_agent_table_read_unsupported(error: &str) -> bool { lower.contains("unknown method") || lower.contains("method not found") } -async fn pool_is_agent(state: &AppState, pool_key: &str) -> bool { +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum TableExportCursorKind { + Agent, + ExternalDriver, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +enum TableExportCursorSession { + Agent(String), + ExternalDriver(String), +} + +async fn table_export_cursor_kind(state: &AppState, pool_key: &str) -> Option { let connections = state.connections.read().await; - matches!(connections.get(pool_key), Some(PoolKind::Agent(_))) + match connections.get(pool_key) { + Some(PoolKind::Agent(_)) => Some(TableExportCursorKind::Agent), + Some(PoolKind::ExternalDriver { .. }) => Some(TableExportCursorKind::ExternalDriver), + _ => None, + } +} + +async fn table_export_query_timeout_secs(state: &AppState, pool_key: &str) -> u64 { + let configs = state.configs.read().await; + config_for_pool_key(pool_key, &configs).map(|config| config.query_timeout_secs).unwrap_or(0) +} + +async fn execute_external_driver_export_page( + state: &AppState, + pool_key: &str, + request: &TableExportRequest, + db_type: &DatabaseType, + col_names: &[String], + primary_keys: &[String], + active_batch_size: usize, + result_session_id: Option, + cancel_token: CancellationToken, +) -> Result { + let sql = table_cursor_sql(request, db_type, col_names, primary_keys); + let max_rows = request.row_limit.unwrap_or(i32::MAX as usize).min(i32::MAX as usize).max(1); + let timeout_secs = table_export_query_timeout_secs(state, pool_key).await; + execute_sql_statement_with_options( + state, + &request.connection_id, + &request.database, + &sql, + request.schema.as_deref(), + Some(cancel_token), + QueryExecutionOptions { + max_rows: Some(max_rows), + fetch_size: Some(active_batch_size), + page_size: Some(active_batch_size), + result_session_id, + client_session_id: Some(table_export_client_session_id(&request.export_id)), + timeout_secs: Some(timeout_secs), + ..Default::default() + }, + ) + .await +} + +async fn execute_table_export_count( + state: &AppState, + pool_key: &str, + request: &TableExportRequest, + sql: &str, + cancel_token: CancellationToken, +) -> Result { + if table_export_cursor_kind(state, pool_key).await != Some(TableExportCursorKind::ExternalDriver) { + return execute_read_on_pool(state, pool_key, sql).await; + } + + let timeout_secs = table_export_query_timeout_secs(state, pool_key).await; + execute_sql_statement_with_options( + state, + &request.connection_id, + &request.database, + sql, + request.schema.as_deref(), + Some(cancel_token), + QueryExecutionOptions { + max_rows: Some(1), + fetch_size: Some(1), + client_session_id: Some(table_export_client_session_id(&request.export_id)), + timeout_secs: Some(timeout_secs), + ..Default::default() + }, + ) + .await } #[allow(clippy::too_many_arguments)] @@ -275,9 +361,10 @@ async fn fetch_table_export_batch( last_pk_values: &[Value], offset: u64, active_batch_size: usize, - table_read_session_id: &mut Option, + cursor_session: &mut Option, table_read_attempted: &mut bool, table_read_completed: &mut bool, + cancel_token: CancellationToken, ) -> Result { if *table_read_completed { return Ok(QueryResult { @@ -293,84 +380,132 @@ async fn fetch_table_export_batch( }); } - if !*table_read_attempted && pool_is_agent(state, pool_key).await { - *table_read_attempted = true; - let sql = table_cursor_sql(request, db_type, col_names, primary_keys); - let max_rows = request.row_limit.unwrap_or(i32::MAX as usize); - let timeout_secs = { - let configs = state.configs.read().await; - let query_timeout = config_for_pool_key(pool_key, &configs).map(|c| c.query_timeout_secs).unwrap_or(0); - if query_timeout == 0 { - None - } else { - Some(query_timeout) + if !*table_read_attempted { + match table_export_cursor_kind(state, pool_key).await { + Some(TableExportCursorKind::Agent) => { + *table_read_attempted = true; + let sql = table_cursor_sql(request, db_type, col_names, primary_keys); + let max_rows = request.row_limit.unwrap_or(i32::MAX as usize); + let query_timeout = table_export_query_timeout_secs(state, pool_key).await; + let params = AgentTableReadStartParams { + sql, + database: Some(request.database.clone()), + schema: request.schema.clone(), + page_size: active_batch_size, + max_rows, + fetch_size: Some(active_batch_size), + timeout_secs: (query_timeout > 0).then_some(query_timeout), + }; + let connections = state.connections.read().await; + let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { + return Err("Agent table read requires an agent connection".to_string()); + }; + let client = client.clone(); + drop(connections); + let mut client = client.lock().await; + match client.start_table_read::(params).await { + Ok(result) => { + *cursor_session = result.session_id.clone().map(TableExportCursorSession::Agent); + if result.session_id.is_none() && !result.has_more { + *table_read_completed = true; + } + return Ok(result); + } + Err(error) if is_agent_table_read_unsupported(&error) => { + log::debug!("Agent table-read cursor unsupported, falling back to paginated export: {error}"); + } + Err(error) => return Err(error), + } } - }; - let params = AgentTableReadStartParams { - sql, - database: Some(request.database.clone()), - schema: request.schema.clone(), - page_size: active_batch_size, - max_rows, - fetch_size: Some(active_batch_size), - timeout_secs, - }; - let connections = state.connections.read().await; - let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { - drop(connections); - return fetch_paginated_table_export_batch( - state, - pool_key, - request, - db_type, - col_names, - primary_keys, - use_keyset, - last_pk_values, - offset, - active_batch_size, - ) - .await; - }; - let client = client.clone(); - drop(connections); - let mut client = client.lock().await; - match client.start_table_read::(params).await { - Ok(result) => { - *table_read_session_id = result.session_id.clone(); - if result.session_id.is_none() && !result.has_more { + Some(TableExportCursorKind::ExternalDriver) => { + *table_read_attempted = true; + let result = execute_external_driver_export_page( + state, + pool_key, + request, + db_type, + col_names, + primary_keys, + active_batch_size, + None, + cancel_token.clone(), + ) + .await?; + if result.has_more { + let session_id = result + .session_id + .clone() + .ok_or("JDBC export cursor did not return a session id for additional rows")?; + *cursor_session = Some(TableExportCursorSession::ExternalDriver(session_id)); + } else { *table_read_completed = true; } return Ok(result); } - Err(error) if is_agent_table_read_unsupported(&error) => { - log::debug!("Agent table-read cursor unsupported, falling back to paginated export: {error}"); - } - Err(error) => return Err(error), + None => {} } } - if let Some(session_id) = table_read_session_id.as_deref() { - let connections = state.connections.read().await; - let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { - return Err("Table read session requires an agent connection".to_string()); - }; - let client = client.clone(); - drop(connections); - let mut client = client.lock().await; - return match client.fetch_table_read_page::(session_id, active_batch_size).await { - Ok(result) => { - *table_read_session_id = result.session_id.clone().or_else(|| Some(session_id.to_string())); - if !result.has_more { - *table_read_session_id = None; - *table_read_completed = true; + if let Some(session) = cursor_session.clone() { + return match session { + TableExportCursorSession::Agent(session_id) => { + let connections = state.connections.read().await; + let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { + return Err("Table read session requires an agent connection".to_string()); + }; + let client = client.clone(); + drop(connections); + let mut client = client.lock().await; + match client.fetch_table_read_page::(&session_id, active_batch_size).await { + Ok(result) => { + *cursor_session = + result.session_id.clone().or(Some(session_id)).map(TableExportCursorSession::Agent); + if !result.has_more { + *cursor_session = None; + *table_read_completed = true; + } + Ok(result) + } + Err(error) => { + let _ = client.close_table_read_session::(&session_id).await; + *cursor_session = None; + Err(error) + } } - Ok(result) } - Err(error) => { - let _ = client.close_table_read_session::(session_id).await; - *table_read_session_id = None; - Err(error) + TableExportCursorSession::ExternalDriver(session_id) => { + match execute_external_driver_export_page( + state, + pool_key, + request, + db_type, + col_names, + primary_keys, + active_batch_size, + Some(session_id.clone()), + cancel_token.clone(), + ) + .await + { + Ok(result) => { + if result.has_more { + let next_session_id = result.session_id.clone().unwrap_or(session_id); + *cursor_session = Some(TableExportCursorSession::ExternalDriver(next_session_id)); + } else { + *cursor_session = None; + *table_read_completed = true; + } + Ok(result) + } + Err(error) => { + if cancel_token.is_cancelled() { + cursor_session.take(); + } else { + close_table_export_cursor_if_open(state, pool_key, request, cursor_session).await; + } + Err(error) + } + } } }; } @@ -416,22 +551,38 @@ async fn fetch_paginated_table_export_batch( execute_read_on_pool_with_max_rows(state, pool_key, &sql, Some(active_batch_size)).await } -async fn close_table_read_session_if_open( +async fn close_table_export_cursor_if_open( state: &AppState, pool_key: &str, - table_read_session_id: &mut Option, + request: &TableExportRequest, + cursor_session: &mut Option, ) { - let Some(session_id) = table_read_session_id.take() else { + let Some(session) = cursor_session.take() else { return; }; - let connections = state.connections.read().await; - let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { - return; - }; - let client = client.clone(); - drop(connections); - let mut client = client.lock().await; - let _ = client.close_table_read_session::(&session_id).await; + match session { + TableExportCursorSession::Agent(session_id) => { + let connections = state.connections.read().await; + let Some(PoolKind::Agent(client)) = connections.get(pool_key) else { + return; + }; + let client = client.clone(); + drop(connections); + let mut client = client.lock().await; + let _ = client.close_table_read_session::(&session_id).await; + } + TableExportCursorSession::ExternalDriver(session_id) => { + let client_session_id = table_export_client_session_id(&request.export_id); + let _ = close_query_session( + state, + &request.connection_id, + &request.database, + &session_id, + Some(&client_session_id), + ) + .await; + } + } } async fn start_export_cancel_watcher(export_id: String, cancelled: Arc, token: CancellationToken) { @@ -511,12 +662,10 @@ async fn try_export_native_table_stream( row_limit: Option, batch_size: usize, on_progress: &impl Fn(TableExportProgress), + cancelled: Arc, + cancel_token: CancellationToken, ) -> Result { let sql = table_cursor_sql(request, db_type, col_names, primary_keys); - let cancelled = Arc::new(AtomicBool::new(false)); - let cancel_token = CancellationToken::new(); - let cancel_watcher = - tokio::spawn(start_export_cancel_watcher(request.export_id.clone(), cancelled.clone(), cancel_token.clone())); let mut rows_exported = 0_u64; let progress_interval = batch_size.max(1) as u64; @@ -834,8 +983,6 @@ async fn try_export_native_table_stream( _ => Ok(false), }; - cancel_watcher.abort(); - match stream_result { Ok(false) => Ok(false), Ok(true) => { @@ -886,6 +1033,43 @@ pub async fn export_table_data_core( state: &AppState, request: &TableExportRequest, on_progress: impl Fn(TableExportProgress), +) -> Result<(), String> { + let cancelled = Arc::new(AtomicBool::new(false)); + let cancel_token = CancellationToken::new(); + let cancel_watcher = + tokio::spawn(start_export_cancel_watcher(request.export_id.clone(), cancelled.clone(), cancel_token.clone())); + let last_rows_exported = std::sync::atomic::AtomicU64::new(0); + let tracked_progress = |progress: TableExportProgress| { + last_rows_exported.store(progress.rows_exported, Ordering::SeqCst); + on_progress(progress); + }; + let result = + export_table_data_core_inner(state, request, &tracked_progress, cancelled.clone(), cancel_token.clone()).await; + cancel_watcher.abort(); + let client_session_id = table_export_client_session_id(&request.export_id); + let _ = state.close_client_session_pool(&request.connection_id, Some(&request.database), &client_session_id).await; + match result { + Err(error) if cancelled.load(Ordering::SeqCst) || error == crate::query::canceled_error() => { + on_progress(TableExportProgress { + export_id: request.export_id.clone(), + table_name: request.table_name.clone(), + rows_exported: last_rows_exported.load(Ordering::SeqCst), + total_rows: None, + status: ExportStatus::Cancelled, + error_message: Some("Export cancelled".to_string()), + }); + Ok(()) + } + result => result, + } +} + +async fn export_table_data_core_inner( + state: &AppState, + request: &TableExportRequest, + on_progress: &impl Fn(TableExportProgress), + cancelled: Arc, + cancel_token: CancellationToken, ) -> Result<(), String> { // 1. Get database type let db_type = state @@ -961,7 +1145,7 @@ pub async fn export_table_data_core( &db_type, request.where_input.as_deref(), ); - match execute_read_on_pool(state, &pool_key, &count_query).await { + match execute_table_export_count(state, &pool_key, request, &count_query, cancel_token.clone()).await { Ok(result) => result .rows .first() @@ -999,6 +1183,8 @@ pub async fn export_table_data_core( row_limit, request.batch_size.unwrap_or(DEFAULT_BATCH_SIZE).max(1), &on_progress, + cancelled, + cancel_token.clone(), ) .await? { @@ -1012,7 +1198,7 @@ pub async fn export_table_data_core( let mut rows_exported: u64 = 0; let batch_size = request.batch_size.unwrap_or(DEFAULT_BATCH_SIZE).max(1); let mut offset: u64 = 0; - let mut table_read_session_id: Option = None; + let mut cursor_session: Option = None; let mut table_read_attempted = false; let mut table_read_completed = false; @@ -1037,7 +1223,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1055,9 +1241,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1124,7 +1311,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1142,9 +1329,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1217,7 +1405,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1235,9 +1423,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1309,7 +1498,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1327,9 +1516,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1390,7 +1580,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1408,9 +1598,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1470,7 +1661,7 @@ pub async fn export_table_data_core( status: ExportStatus::Cancelled, error_message: Some("Export cancelled".to_string()), }); - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; return Ok(()); } @@ -1488,9 +1679,10 @@ pub async fn export_table_data_core( &last_pk_values, offset, active_batch_size, - &mut table_read_session_id, + &mut cursor_session, &mut table_read_attempted, &mut table_read_completed, + cancel_token.clone(), ) .await?; let row_count = result.rows.len(); @@ -1550,7 +1742,7 @@ pub async fn export_table_data_core( } } - close_table_read_session_if_open(state, &pool_key, &mut table_read_session_id).await; + close_table_export_cursor_if_open(state, &pool_key, request, &mut cursor_session).await; file.flush().map_err(|e| format!("Failed to flush export file: {e}"))?; // 8. Emit Done progress @@ -1570,9 +1762,150 @@ pub async fn export_table_data_core( mod tests { use super::*; use crate::database_export::{clear_export_cancelled, set_export_cancelled}; + #[cfg(unix)] + use crate::models::connection::ConnectionConfig; + #[cfg(unix)] + use crate::plugins::{ + InstalledPlugin, PluginDriverManifest, PluginDriverSession, PluginManifest, PluginRuntimeEnv, + }; + #[cfg(unix)] + use crate::storage::Storage; use crate::xlsx_export::{build_xlsx_workbook, XlsxWorksheetData}; use serde_json::json; use std::io::Read; + #[cfg(unix)] + use std::os::unix::fs::PermissionsExt; + + #[cfg(unix)] + struct ExternalDriverExportFixture { + state: AppState, + request: TableExportRequest, + calls: std::path::PathBuf, + output: std::path::PathBuf, + dir: std::path::PathBuf, + } + + #[cfg(unix)] + async fn external_driver_export_fixture( + rpc_body: &str, + batch_size: usize, + row_limit: Option, + skip_count: bool, + ) -> ExternalDriverExportFixture { + let dir = std::env::temp_dir().join(format!("dbx-jdbc-table-export-test-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let executable = dir.join("plugin.sh"); + let calls = dir.join("calls.log"); + let script = format!( + "#!/bin/sh\nCALLS='{}'\nwhile IFS= read -r line; do\n id=$(printf '%s' \"$line\" | sed -E 's/.*\"id\":([0-9]+).*/\\1/')\n{}\ndone\n", + calls.display(), + rpc_body + ); + std::fs::write(&executable, script).unwrap(); + let mut permissions = std::fs::metadata(&executable).unwrap().permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&executable, permissions).unwrap(); + + let plugin = InstalledPlugin { + manifest: PluginManifest { + id: "jdbc".to_string(), + name: "JDBC".to_string(), + version: "test".to_string(), + protocol_version: 1, + description: String::new(), + executable: Some("plugin.sh".to_string()), + drivers: vec![PluginDriverManifest { + id: "jdbc".to_string(), + label: "JDBC".to_string(), + kind: "external".to_string(), + database_type: Some("jdbc".to_string()), + }], + }, + path: dir.clone(), + }; + let session = Arc::new( + PluginDriverSession::start_for_test(plugin, "jdbc".to_string(), PluginRuntimeEnv::default()) + .await + .expect("test JDBC plugin should start"), + ); + let config: ConnectionConfig = serde_json::from_value(json!({ + "id": "conn-1", + "name": "JDBC", + "db_type": "jdbc", + "host": "", + "port": 0, + "username": "", + "password": "", + "database": "PUBLIC", + "query_timeout_secs": 30 + })) + .unwrap(); + let storage = Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = AppState::new(storage); + state.configs.write().await.insert(config.id.clone(), config.clone()); + let export_id = format!("export-{}", uuid::Uuid::new_v4()); + let pool_key = + format!("{}:session:{}", config.id, table_export_client_session_id(&export_id).replace(':', "_")); + state.connections.write().await.insert( + pool_key, + PoolKind::ExternalDriver { driver_id: "jdbc".to_string(), config: Arc::new(config), session }, + ); + + let output = dir.join("export.csv"); + let request = TableExportRequest { + export_id, + connection_id: "conn-1".to_string(), + database: "PUBLIC".to_string(), + schema: Some("PUBLIC".to_string()), + table_name: "EXPORT_SAMPLE".to_string(), + file_path: output.to_string_lossy().into_owned(), + format: "csv".to_string(), + columns: Some(vec!["id".to_string(), "name".to_string()]), + column_types: Some(vec![Some("INTEGER".to_string()), Some("VARCHAR".to_string())]), + primary_keys: Some(vec!["id".to_string()]), + where_input: None, + order_by: None, + skip_count, + batch_size: Some(batch_size), + row_limit, + date_time_format: None, + }; + + ExternalDriverExportFixture { state, request, calls, output, dir } + } + + #[cfg(unix)] + async fn run_external_driver_export( + fixture: &ExternalDriverExportFixture, + ) -> Result, String> { + let progress = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = progress.clone(); + let result = export_table_data_core(&fixture.state, &fixture.request, move |event| { + captured.lock().unwrap().push(event); + }) + .await; + let events = progress.lock().unwrap().clone(); + result.map(|_| events) + } + + #[cfg(unix)] + async fn wait_for_external_driver_call(calls: &std::path::Path, expected: &str) { + tokio::time::timeout(Duration::from_secs(5), async { + loop { + if std::fs::read_to_string(calls).unwrap_or_default().lines().any(|line| line == expected) { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap_or_else(|_| panic!("timed out waiting for plugin call: {expected}")); + } + + #[cfg(unix)] + fn cleanup_external_driver_export_fixture(fixture: ExternalDriverExportFixture) { + let _ = std::fs::remove_dir_all(fixture.dir); + } /// Read and decompress a single entry from an in-memory XLSX (ZIP) buffer. fn read_zip_entry(bytes: &[u8], path: &str) -> String { @@ -1782,6 +2115,217 @@ mod tests { assert!(!is_agent_table_read_unsupported("ORA-00933: SQL command not properly ended")); } + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_reads_all_cursor_pages() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[1,"Ada"],[2,"Grace"]],"affected_rows":0,"execution_time_ms":1,"session_id":"cursor-1","has_more":true}}\n' "$id" + ;; + *'"method":"fetchQueryPage"'*) + echo fetchQueryPage >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[3,"Linus"]],"affected_rows":0,"execution_time_ms":1,"session_id":null,"has_more":false}}\n' "$id" + ;; + *'"method":"executeQuery"'*) + echo executeQuery >> "$CALLS" + printf '{"id":%s,"result":{"columns":["count"],"rows":[[3]],"affected_rows":0,"execution_time_ms":1}}\n' "$id" + ;; + *'"method":"closeQuerySession"'*) + echo closeQuerySession >> "$CALLS" + printf '{"id":%s,"result":{"ok":true}}\n' "$id" + ;; + esac"#, + 2, + None, + false, + ) + .await; + + let progress = run_external_driver_export(&fixture).await.expect("multi-page JDBC export should succeed"); + let csv = std::fs::read_to_string(&fixture.output).unwrap(); + assert!(csv.contains("\"1\",\"Ada\"")); + assert!(csv.contains("\"2\",\"Grace\"")); + assert!(csv.contains("\"3\",\"Linus\"")); + assert_eq!(csv.matches("\"Ada\"").count(), 1); + assert_eq!( + std::fs::read_to_string(&fixture.calls).unwrap(), + "executeQuery\nexecuteQueryPage\nfetchQueryPage\n" + ); + assert_eq!(progress.last().and_then(|event| event.total_rows), Some(3)); + assert!(matches!(progress.last().map(|event| &event.status), Some(ExportStatus::Done))); + assert!(fixture.state.connections.read().await.is_empty()); + + cleanup_external_driver_export_fixture(fixture); + } + + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_does_not_repeat_legacy_one_shot_results() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + printf '{"id":%s,"error":{"message":"Unsupported JDBC plugin method: executeQueryPage"}}\n' "$id" + ;; + *'"method":"executeQuery"'*) + echo executeQuery >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[1,"Ada"],[2,"Grace"],[3,"Linus"]],"affected_rows":0,"execution_time_ms":1}}\n' "$id" + ;; + esac"#, + 2, + None, + true, + ) + .await; + + run_external_driver_export(&fixture).await.expect("legacy JDBC export should succeed"); + let csv = std::fs::read_to_string(&fixture.output).unwrap(); + assert_eq!(csv.matches("\"Ada\"").count(), 1); + assert_eq!(csv.matches("\"Grace\"").count(), 1); + assert_eq!(csv.matches("\"Linus\"").count(), 1); + assert_eq!(std::fs::read_to_string(&fixture.calls).unwrap(), "executeQueryPage\nexecuteQuery\n"); + + cleanup_external_driver_export_fixture(fixture); + } + + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_closes_cursor_at_row_limit() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[1,"Ada"]],"affected_rows":0,"execution_time_ms":1,"session_id":"cursor-1","has_more":true}}\n' "$id" + ;; + *'"method":"closeQuerySession"'*) + echo closeQuerySession >> "$CALLS" + printf '{"id":%s,"result":{"ok":true}}\n' "$id" + ;; + esac"#, + 1, + Some(1), + true, + ) + .await; + + run_external_driver_export(&fixture).await.expect("row-limited JDBC export should succeed"); + assert_eq!(std::fs::read_to_string(&fixture.calls).unwrap(), "executeQueryPage\ncloseQuerySession\n"); + assert!(fixture.state.connections.read().await.is_empty()); + + cleanup_external_driver_export_fixture(fixture); + } + + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_closes_cursor_after_fetch_error() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[1,"Ada"],[2,"Grace"]],"affected_rows":0,"execution_time_ms":1,"session_id":"cursor-1","has_more":true}}\n' "$id" + ;; + *'"method":"fetchQueryPage"'*) + echo fetchQueryPage >> "$CALLS" + printf '{"id":%s,"error":{"message":"simulated fetch failure"}}\n' "$id" + ;; + *'"method":"closeQuerySession"'*) + echo closeQuerySession >> "$CALLS" + printf '{"id":%s,"result":{"ok":true}}\n' "$id" + ;; + esac"#, + 2, + None, + true, + ) + .await; + + let error = run_external_driver_export(&fixture).await.expect_err("fetch errors must fail the export"); + assert!(error.starts_with("simulated fetch failure")); + assert_eq!( + std::fs::read_to_string(&fixture.calls).unwrap(), + "executeQueryPage\nfetchQueryPage\ncloseQuerySession\n" + ); + assert!(fixture.state.connections.read().await.is_empty()); + + cleanup_external_driver_export_fixture(fixture); + } + + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_cancels_blocked_execute() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + sleep 30 + ;; + esac"#, + 2, + None, + true, + ) + .await; + + let export = run_external_driver_export(&fixture); + let cancel = async { + wait_for_external_driver_call(&fixture.calls, "executeQueryPage").await; + set_export_cancelled(&fixture.request.export_id).await; + tokio::time::Instant::now() + }; + let (result, cancel_requested_at) = + tokio::time::timeout(Duration::from_secs(7), async { tokio::join!(export, cancel) }) + .await + .expect("blocked JDBC execute should be interrupted promptly"); + let progress = result.expect("cancelled JDBC export should complete without an error"); + + assert!(cancel_requested_at.elapsed() < Duration::from_secs(2)); + assert!(matches!(progress.last().map(|event| &event.status), Some(ExportStatus::Cancelled))); + assert!(fixture.state.connections.read().await.is_empty()); + clear_export_cancelled(&fixture.request.export_id).await; + cleanup_external_driver_export_fixture(fixture); + } + + #[cfg(unix)] + #[tokio::test] + async fn external_driver_table_export_cancels_blocked_fetch() { + let fixture = external_driver_export_fixture( + r#" case "$line" in + *'"method":"executeQueryPage"'*) + echo executeQueryPage >> "$CALLS" + printf '{"id":%s,"result":{"columns":["id","name"],"rows":[[1,"Ada"],[2,"Grace"]],"affected_rows":0,"execution_time_ms":1,"session_id":"cursor-1","has_more":true}}\n' "$id" + ;; + *'"method":"fetchQueryPage"'*) + echo fetchQueryPage >> "$CALLS" + sleep 30 + ;; + esac"#, + 2, + None, + true, + ) + .await; + + let export = run_external_driver_export(&fixture); + let cancel = async { + wait_for_external_driver_call(&fixture.calls, "fetchQueryPage").await; + set_export_cancelled(&fixture.request.export_id).await; + tokio::time::Instant::now() + }; + let (result, cancel_requested_at) = + tokio::time::timeout(Duration::from_secs(7), async { tokio::join!(export, cancel) }) + .await + .expect("blocked JDBC fetch should be interrupted promptly"); + let progress = result.expect("cancelled JDBC export should complete without an error"); + + assert!(cancel_requested_at.elapsed() < Duration::from_secs(2)); + assert!(matches!(progress.last().map(|event| &event.status), Some(ExportStatus::Cancelled))); + assert!(fixture.state.connections.read().await.is_empty()); + clear_export_cancelled(&fixture.request.export_id).await; + cleanup_external_driver_export_fixture(fixture); + } + #[test] fn writes_json_row_without_allocating_object_map() { let mut out = Vec::new(); diff --git a/plugins/jdbc/src/main/java/app/dbx/jdbc/DbxJdbcPlugin.java b/plugins/jdbc/src/main/java/app/dbx/jdbc/DbxJdbcPlugin.java index 43381f200..2bdd33e22 100644 --- a/plugins/jdbc/src/main/java/app/dbx/jdbc/DbxJdbcPlugin.java +++ b/plugins/jdbc/src/main/java/app/dbx/jdbc/DbxJdbcPlugin.java @@ -513,6 +513,7 @@ public final class DbxJdbcPlugin { properties.setProperty("password", password); } applyConnectTimeout(connection, properties); + applyPagedFetchProperties(connection, url, properties); applyJdbcxExtensionSecurity(connection, url, properties); if (isOracleUrl(url)) { applyOracleProperties(connection, properties); @@ -586,6 +587,33 @@ public final class DbxJdbcPlugin { normalized.equals("org.mariadb.jdbc.driver"); } + private static void applyPagedFetchProperties(JsonNode connection, String url, Properties properties) { + if (!isMysqlConnection(connection, url) || jdbcUrlHasParameter(url, "useCursorFetch")) { + return; + } + properties.putIfAbsent("useCursorFetch", "true"); + } + + private static boolean isMysqlConnection(JsonNode connection, String url) { + if (urlMatchesPrefix(url, "jdbc:mysql:")) { + return true; + } + String driverClass = optionalText(connection, "jdbc_driver_class"); + if (driverClass == null) { + return false; + } + String normalized = driverClass.toLowerCase(Locale.ROOT); + return normalized.equals("com.mysql.cj.jdbc.driver") || normalized.equals("com.mysql.jdbc.driver"); + } + + private static boolean isPostgresConnection(JsonNode connection) { + if (urlMatchesPrefix(jdbcUrl(connection), "jdbc:postgresql:")) { + return true; + } + String driverClass = optionalText(connection, "jdbc_driver_class"); + return driverClass != null && driverClass.equalsIgnoreCase("org.postgresql.Driver"); + } + private static boolean isPrestoOrTrinoConnection(JsonNode connection) { String url = jdbcUrl(connection); if (urlMatchesPrefix(url, "jdbc:presto:") || urlMatchesPrefix(url, "jdbc:trino:")) { @@ -675,6 +703,8 @@ public final class DbxJdbcPlugin { private final ArrayNode columns; private final int maxRows; private final long startNanos; + private final Connection connection; + private final boolean restoreAutoCommit; private int rowsReturned; private ArrayNode pendingRow; @@ -685,7 +715,9 @@ public final class DbxJdbcPlugin { ResultSetMetaData meta, ArrayNode columns, int maxRows, - long startNanos + long startNanos, + Connection connection, + boolean restoreAutoCommit ) { this.id = id; this.statement = statement; @@ -694,6 +726,8 @@ public final class DbxJdbcPlugin { this.columns = columns; this.maxRows = Math.max(1, maxRows); this.startNanos = startNanos; + this.connection = connection; + this.restoreAutoCommit = restoreAutoCommit; } } @@ -711,7 +745,14 @@ public final class DbxJdbcPlugin { Connection conn = openConnection(connection); applyExecutionContext(connection, conn, database, schema); JdbcDriverQuirks quirks = driverQuirks(connection); - Statement statement = conn.createStatement(); + boolean restoreAutoCommit = beginPagedQueryTransaction(connection, conn); + Statement statement; + try { + statement = createPagedQueryStatement(conn); + } catch (Exception | LinkageError error) { + restorePagedQueryTransaction(conn, restoreAutoCommit); + throw error; + } try { applyStatementOptions(statement, maxRows, fetchSize, timeoutSecs, quirks); String trimmedSql = trimStatementSql(sql); @@ -727,6 +768,7 @@ public final class DbxJdbcPlugin { result.putNull("session_id"); result.put("has_more", false); statement.close(); + restorePagedQueryTransaction(conn, restoreAutoCommit); return result; } @@ -738,18 +780,65 @@ public final class DbxJdbcPlugin { columns.add(label == null || label.isBlank() ? meta.getColumnName(i) : label); } String sessionId = UUID.randomUUID().toString(); - QuerySession session = new QuerySession(sessionId, statement, rs, meta, columns, maxRows, start); + QuerySession session = new QuerySession( + sessionId, + statement, + rs, + meta, + columns, + maxRows, + start, + conn, + restoreAutoCommit + ); QUERY_SESSIONS.put(sessionId, session); - return readQuerySessionPage(session, pageSize); - } catch (Exception error) { + try { + return readQuerySessionPage(session, pageSize); + } catch (Exception | LinkageError error) { + QUERY_SESSIONS.remove(sessionId); + throw error; + } + } catch (Exception | LinkageError error) { try { statement.close(); } catch (Exception ignored) { } + restorePagedQueryTransaction(conn, restoreAutoCommit); throw error; } } + private static Statement createPagedQueryStatement(Connection connection) throws SQLException { + try { + return connection.createStatement(ResultSet.TYPE_FORWARD_ONLY, ResultSet.CONCUR_READ_ONLY); + } catch (SQLFeatureNotSupportedException | UnsupportedOperationException | AbstractMethodError ignored) { + return connection.createStatement(); + } + } + + private static boolean beginPagedQueryTransaction(JsonNode connectionConfig, Connection connection) + throws SQLException { + if (!isPostgresConnection(connectionConfig) || !connection.getAutoCommit()) { + return false; + } + connection.setAutoCommit(false); + return true; + } + + private static void restorePagedQueryTransaction(Connection connection, boolean restoreAutoCommit) { + if (!restoreAutoCommit) { + return; + } + try { + connection.rollback(); + } catch (SQLException ignored) { + } + try { + connection.setAutoCommit(true); + } catch (SQLException ignored) { + } + } + private static JsonNode fetchQueryPage(String sessionId, int pageSize) throws SQLException { QuerySession session = QUERY_SESSIONS.get(sessionId); if (session == null) { @@ -830,6 +919,7 @@ public final class DbxJdbcPlugin { session.statement.close(); } catch (Exception ignored) { } + restorePagedQueryTransaction(session.connection, session.restoreAutoCommit); return true; } diff --git a/plugins/jdbc/src/test/java/app/dbx/jdbc/DbxJdbcPluginTest.java b/plugins/jdbc/src/test/java/app/dbx/jdbc/DbxJdbcPluginTest.java index ab9bb5b72..e2a617ecd 100644 --- a/plugins/jdbc/src/test/java/app/dbx/jdbc/DbxJdbcPluginTest.java +++ b/plugins/jdbc/src/test/java/app/dbx/jdbc/DbxJdbcPluginTest.java @@ -522,6 +522,120 @@ final class DbxJdbcPluginTest { assertFalse(properties.containsKey("connectTimeout")); } + @Test + void mysqlPagedQueriesEnableConnectorCursorFetchingByDefault() throws Exception { + Method method = DbxJdbcPlugin.class.getDeclaredMethod( + "applyPagedFetchProperties", + JsonNode.class, + String.class, + Properties.class + ); + method.setAccessible(true); + Properties properties = new Properties(); + JsonNode connection = MAPPER.readTree(""" + { + "connection_string": "jdbc:mysql://127.0.0.1:3306/app", + "jdbc_driver_class": "com.mysql.cj.jdbc.Driver" + } + """); + + method.invoke(null, connection, "jdbc:mysql://127.0.0.1:3306/app", properties); + + assertEquals("true", properties.getProperty("useCursorFetch")); + } + + @Test + void mysqlPagedQueriesPreserveExplicitCursorFetchSetting() throws Exception { + Method method = DbxJdbcPlugin.class.getDeclaredMethod( + "applyPagedFetchProperties", + JsonNode.class, + String.class, + Properties.class + ); + method.setAccessible(true); + Properties properties = new Properties(); + JsonNode connection = MAPPER.readTree(""" + { + "connection_string": "jdbc:mysql://127.0.0.1:3306/app?useCursorFetch=false" + } + """); + + method.invoke( + null, + connection, + "jdbc:mysql://127.0.0.1:3306/app?useCursorFetch=false", + properties + ); + + assertFalse(properties.containsKey("useCursorFetch")); + } + + @Test + void postgresPagedQueryUsesCursorTransactionAndRestoresAutoCommit() throws Exception { + Method begin = DbxJdbcPlugin.class.getDeclaredMethod( + "beginPagedQueryTransaction", + JsonNode.class, + Connection.class + ); + Method create = DbxJdbcPlugin.class.getDeclaredMethod("createPagedQueryStatement", Connection.class); + Method restore = DbxJdbcPlugin.class.getDeclaredMethod( + "restorePagedQueryTransaction", + Connection.class, + boolean.class + ); + begin.setAccessible(true); + create.setAccessible(true); + restore.setAccessible(true); + List calls = new ArrayList<>(); + Connection connection = pagedQueryConnection(calls, true); + JsonNode config = MAPPER.readTree(""" + { "connection_string": "jdbc:postgresql://127.0.0.1:5432/app" } + """); + + boolean restoreAutoCommit = (boolean) begin.invoke(null, config, connection); + create.invoke(null, connection); + restore.invoke(null, connection, restoreAutoCommit); + + assertEquals(true, restoreAutoCommit); + assertEquals( + List.of( + "getAutoCommit", + "setAutoCommit:false", + "createStatement:" + ResultSet.TYPE_FORWARD_ONLY + ":" + ResultSet.CONCUR_READ_ONLY, + "rollback", + "setAutoCommit:true" + ), + calls + ); + } + + @Test + void postgresPagedQueryPreservesExistingManualTransaction() throws Exception { + Method begin = DbxJdbcPlugin.class.getDeclaredMethod( + "beginPagedQueryTransaction", + JsonNode.class, + Connection.class + ); + Method restore = DbxJdbcPlugin.class.getDeclaredMethod( + "restorePagedQueryTransaction", + Connection.class, + boolean.class + ); + begin.setAccessible(true); + restore.setAccessible(true); + List calls = new ArrayList<>(); + Connection connection = pagedQueryConnection(calls, false); + JsonNode config = MAPPER.readTree(""" + { "jdbc_driver_class": "org.postgresql.Driver" } + """); + + boolean restoreAutoCommit = (boolean) begin.invoke(null, config, connection); + restore.invoke(null, connection, restoreAutoCommit); + + assertFalse(restoreAutoCommit); + assertEquals(List.of("getAutoCommit"), calls); + } + @Test void jdbcxHighPrivilegeExtensionsAreDisabledByDefault() throws Exception { Method method = DbxJdbcPlugin.class.getDeclaredMethod( @@ -1485,6 +1599,36 @@ final class DbxJdbcPluginTest { ); } + private static Connection pagedQueryConnection(List calls, boolean autoCommit) { + return (Connection) Proxy.newProxyInstance( + DbxJdbcPluginTest.class.getClassLoader(), + new Class[] { Connection.class }, + (proxy, method, args) -> switch (method.getName()) { + case "getAutoCommit" -> { + calls.add("getAutoCommit"); + yield autoCommit; + } + case "setAutoCommit" -> { + calls.add("setAutoCommit:" + args[0]); + yield null; + } + case "rollback" -> { + calls.add("rollback"); + yield null; + } + case "createStatement" -> { + if (args == null || args.length == 0) { + calls.add("createStatement"); + } else { + calls.add("createStatement:" + args[0] + ":" + args[1]); + } + yield recordingStatement(new ArrayList<>()); + } + default -> defaultValue(method.getReturnType()); + } + ); + } + private static ResultSet temporalResultSet(Object objectValue, Date dateValue) { return (ResultSet) Proxy.newProxyInstance( DbxJdbcPluginTest.class.getClassLoader(),