diff --git a/crates/dbx-core/src/query_result_sql.rs b/crates/dbx-core/src/query_result_sql.rs index 50a290364..455035bd9 100644 --- a/crates/dbx-core/src/query_result_sql.rs +++ b/crates/dbx-core/src/query_result_sql.rs @@ -224,6 +224,19 @@ pub fn build_count_query_sql(options: CountQuerySqlOptions) -> QuerySqlBuildResu if unsupported_pagination_type(options.database_type) { return err("unsupported"); } + // A locking clause does not affect cardinality and cannot appear inside + // every dialect's derived-table count query. PostgreSQL permits pagination + // after the lock clause; decline counting that uncommon order rather than + // accidentally dropping the user's explicit LIMIT/OFFSET. + let tokens = top_level_sql_tokens(&statement); + let statement = if let Some(index) = locking_clause_index(&tokens) { + if has_pagination_clause_after(&tokens, index) { + return err("locking"); + } + statement[..index].trim_end().to_string() + } else { + statement + }; // ES SQL can't wrap a SELECT in `SELECT COUNT(*) FROM (...)` — the // driver already reports the true match count via affected_rows. if matches!(options.database_type, Some(DatabaseType::Elasticsearch | DatabaseType::Easysearch)) { @@ -952,12 +965,13 @@ fn add_questdb_limit(statement: &str, limit: usize, offset: usize) -> String { } return format!("{statement};"); } - if offset > 0 { + let limit_sql = if offset > 0 { let upper_bound = offset + limit; - append_sql_suffix(statement, &format!("LIMIT {offset}, {upper_bound};")) + format!("LIMIT {offset}, {upper_bound}") } else { - append_sql_suffix(statement, &format!("LIMIT {limit};")) - } + format!("LIMIT {limit}") + }; + append_or_insert_before_locking(statement, &limit_sql) } fn has_top_level_limit(sql: &str) -> bool { @@ -1057,7 +1071,7 @@ fn add_fetch_first_limit(statement: &str, limit: usize, offset: usize) -> String return format!("{statement};"); } let offset_sql = if offset > 0 { format!(" OFFSET {offset} ROWS") } else { String::new() }; - append_sql_suffix(statement, &format!("{offset_sql} FETCH FIRST {limit} ROWS ONLY;")) + append_or_insert_before_locking(statement, &format!("{offset_sql} FETCH FIRST {limit} ROWS ONLY")) } fn add_firebird_rows_limit(statement: &str, limit: usize, offset: usize) -> String { @@ -1114,7 +1128,7 @@ fn add_standard_limit( if database_type == Some(DatabaseType::ClickHouse) { return add_clickhouse_limit(statement, &limit_sql); } - append_sql_suffix(statement, &format!("{limit_sql};")) + append_or_insert_before_locking(statement, &limit_sql) } fn add_outer_standard_limit( @@ -1138,7 +1152,45 @@ fn add_clickhouse_limit(statement: &str, limit_sql: &str) -> String { return format!("{statement_before_settings} {limit_sql} {settings_clause};"); } - append_sql_suffix(statement, &format!("{limit_sql};")) + append_or_insert_before_locking(statement, limit_sql) +} + +/// Insert pagination before a top-level locking clause; SQL dialects require +/// LIMIT/FETCH to precede FOR UPDATE, FOR SHARE, or LOCK IN SHARE MODE. +fn append_or_insert_before_locking(statement: &str, clause: &str) -> String { + let clause = clause.trim(); + if let Some(index) = locking_clause_index(&top_level_sql_tokens(statement)) { + let before = statement[..index].trim_end(); + let after = statement[index..].trim_start(); + let separator = if sql_suffix_needs_newline(before) { "\n" } else { " " }; + return format!("{before}{separator}{clause} {after};"); + } + append_sql_suffix(statement, &format!("{clause};")) +} + +const LOCKING_CLAUSE_PATTERNS: &[&[&str]] = &[ + &["FOR", "UPDATE"], + &["FOR", "SHARE"], + &["FOR", "KEY", "SHARE"], + &["FOR", "NO", "KEY", "UPDATE"], + &["LOCK", "IN", "SHARE", "MODE"], +]; + +fn locking_clause_index(tokens: &[SqlToken]) -> Option { + tokens.iter().enumerate().find_map(|(index, token)| { + LOCKING_CLAUSE_PATTERNS + .iter() + .any(|pattern| token_sequence_matches(&tokens[index..], pattern)) + .then_some(token.start) + }) +} + +fn token_sequence_matches(tokens: &[SqlToken], expected: &[&str]) -> bool { + tokens.len() >= expected.len() && tokens.iter().zip(expected).all(|(token, expected)| token.text == *expected) +} + +fn has_pagination_clause_after(tokens: &[SqlToken], index: usize) -> bool { + tokens.iter().any(|token| token.start > index && matches!(token.text.as_str(), "LIMIT" | "OFFSET" | "FETCH")) } fn clickhouse_settings_clause_index(statement: &str) -> Option { @@ -1340,6 +1392,14 @@ fn top_level_sql_tokens(sql: &str) -> Vec { continue; } + if ch == '#' { + i += 1; + while i < sql.len() && next_char(sql, i) != '\n' { + i += next_char(sql, i).len_utf8(); + } + continue; + } + if ch == '/' && next == Some('*') { i += 2; while i < sql.len() { @@ -1354,6 +1414,15 @@ fn top_level_sql_tokens(sql: &str) -> Vec { continue; } + // PostgreSQL dollar-quoted bodies may contain arbitrary SQL keywords. + // Skip them before scanning for top-level clauses such as FOR UPDATE. + if ch == '$' { + if let Some(end) = skip_sql_dollar_quoted(sql, i) { + i = end; + continue; + } + } + if matches!(ch, '\'' | '"' | '`') { i = skip_sql_quoted(sql, i, ch); continue; @@ -1392,6 +1461,22 @@ fn top_level_sql_tokens(sql: &str) -> Vec { tokens } +fn skip_sql_dollar_quoted(sql: &str, pos: usize) -> Option { + let tag_end_offset = sql.get(pos + 1..)?.find('$')?; + let tag_end = pos + 1 + tag_end_offset; + let tag = &sql[pos + 1..tag_end]; + let valid_tag = tag.is_empty() + || (tag.chars().next().is_some_and(|ch| ch.is_ascii_alphabetic() || ch == '_') + && tag.chars().all(|ch| ch.is_ascii_alphanumeric() || ch == '_')); + if !valid_tag { + return None; + } + + let delimiter = &sql[pos..=tag_end]; + let content_start = tag_end + 1; + sql.get(content_start..)?.find(delimiter).map(|closing_offset| content_start + closing_offset + delimiter.len()) +} + fn skip_sql_quoted(sql: &str, pos: usize, quote: char) -> usize { let mut i = pos + quote.len_utf8(); while i < sql.len() { @@ -2147,6 +2232,185 @@ WHERE u.id = picked.id; assert_eq!(plan.page_offset, Some(0)); } + #[test] + fn mysql_for_update_places_limit_before_locking_clause() { + let result = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: "SELECT * FROM `test`\nwhere id=1 for update".to_string(), + database_type: Some(DatabaseType::Mysql), + limit: 100, + offset: 0, + }); + + assert!(result.ok); + assert_eq!(result.sql.unwrap(), "SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;"); + } + + #[test] + fn locking_query_plan_keeps_server_pagination_and_count() { + let sql = "SELECT * FROM `test`\nwhere id=1 for update".to_string(); + let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions { + sql: sql.clone(), + query_base_sql: sql.clone(), + database_type: Some(DatabaseType::Mysql), + pagination: QueryPagination { limit: 100, offset: 0, session_id: None }, + use_agent_cursor: false, + first_page_uses_actual_sql: false, + }); + + assert_eq!(plan.sql_to_execute, "SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;"); + assert_eq!(plan.page_sql, Some("SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;".to_string())); + assert_eq!(plan.page_limit, Some(100)); + assert_eq!(plan.page_offset, Some(0)); + assert_eq!( + plan.count_sql, + Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM `test`\nwhere id=1) `dbx_count`;".to_string()) + ); + assert!(!plan.use_agent_result_session); + } + + #[test] + fn locking_query_later_page_places_offset_before_locking_clause() { + let sql = "SELECT * FROM t WHERE deleted = 0 FOR UPDATE SKIP LOCKED".to_string(); + let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions { + sql: sql.clone(), + query_base_sql: sql.clone(), + database_type: Some(DatabaseType::Mysql), + pagination: QueryPagination { limit: 100, offset: 100, session_id: None }, + use_agent_cursor: false, + first_page_uses_actual_sql: false, + }); + + assert_eq!( + plan.sql_to_execute, + "SELECT * FROM t WHERE deleted = 0 LIMIT 100 OFFSET 100 FOR UPDATE SKIP LOCKED;" + ); + assert_eq!(plan.page_limit, Some(100)); + assert_eq!(plan.page_offset, Some(100)); + } + + #[test] + fn locking_query_count_removes_locking_clause() { + let result = build_count_query_sql(CountQuerySqlOptions { + original_sql: "SELECT * FROM t WHERE deleted = 0 FOR UPDATE".to_string(), + database_type: Some(DatabaseType::Mysql), + }); + + assert_eq!( + result.sql, + Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM t WHERE deleted = 0) `dbx_count`;".to_string()) + ); + } + + #[test] + fn locking_query_count_preserves_postgres_limit_after_lock_by_declining_rewrite() { + let result = build_count_query_sql(CountQuerySqlOptions { + original_sql: "SELECT * FROM t FOR UPDATE LIMIT 10".to_string(), + database_type: Some(DatabaseType::Postgres), + }); + + assert_eq!(result, err("locking")); + } + + #[test] + fn places_limit_before_supported_top_level_locking_clause_variants() { + for (sql, expected) in [ + ("SELECT * FROM t FOR UPDATE", "SELECT * FROM t LIMIT 100 FOR UPDATE;"), + ("SELECT * FROM t FOR SHARE", "SELECT * FROM t LIMIT 100 FOR SHARE;"), + ("SELECT * FROM t FOR KEY SHARE", "SELECT * FROM t LIMIT 100 FOR KEY SHARE;"), + ("SELECT * FROM t FOR NO KEY UPDATE", "SELECT * FROM t LIMIT 100 FOR NO KEY UPDATE;"), + ("SELECT * FROM t LOCK IN SHARE MODE", "SELECT * FROM t LIMIT 100 LOCK IN SHARE MODE;"), + ] { + let result = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: sql.to_string(), + database_type: Some(DatabaseType::Postgres), + limit: 100, + offset: 0, + }); + assert_eq!(result.sql.as_deref(), Some(expected), "{sql}"); + } + } + + #[test] + fn nested_for_update_does_not_block_outer_limit_append() { + let result = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: "SELECT * FROM (SELECT id FROM t FOR UPDATE) locked".to_string(), + database_type: Some(DatabaseType::Mysql), + limit: 100, + offset: 0, + }); + + assert!(result.ok); + assert_eq!(result.sql.unwrap(), "SELECT * FROM (SELECT id FROM t FOR UPDATE) locked LIMIT 100;"); + } + + #[test] + fn for_xml_is_not_treated_as_locking_clause() { + let result = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: "SELECT id, name FROM users FOR XML PATH('row')".to_string(), + database_type: Some(DatabaseType::Mysql), + limit: 100, + offset: 0, + }); + + assert!(result.ok); + assert_eq!(result.sql.unwrap(), "SELECT id, name FROM users FOR XML PATH('row') LIMIT 100;"); + } + + #[test] + fn ordinary_select_still_appends_limit() { + let result = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: "SELECT * FROM t WHERE deleted = 0".to_string(), + database_type: Some(DatabaseType::Mysql), + limit: 100, + offset: 0, + }); + + assert!(result.ok); + assert_eq!(result.sql.unwrap(), "SELECT * FROM t WHERE deleted = 0 LIMIT 100;"); + } + + #[test] + fn locking_keywords_inside_postgres_dollar_quote_are_not_rewritten() { + let sql = "SELECT $$FOR UPDATE$$ AS message"; + let paginated = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: sql.to_string(), + database_type: Some(DatabaseType::Postgres), + limit: 100, + offset: 0, + }); + let counted = build_count_query_sql(CountQuerySqlOptions { + original_sql: sql.to_string(), + database_type: Some(DatabaseType::Postgres), + }); + + assert_eq!(paginated.sql.as_deref(), Some("SELECT $$FOR UPDATE$$ AS message LIMIT 100;")); + assert_eq!( + counted.sql.as_deref(), + Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT $$FOR UPDATE$$ AS message) \"dbx_count\";") + ); + } + + #[test] + fn locking_keywords_inside_mysql_hash_comment_are_not_rewritten() { + let sql = "SELECT * FROM t\n# FOR UPDATE LIMIT 1"; + let paginated = build_paginated_query_sql(PaginatedQuerySqlOptions { + original_sql: sql.to_string(), + database_type: Some(DatabaseType::Mysql), + limit: 100, + offset: 0, + }); + let counted = build_count_query_sql(CountQuerySqlOptions { + original_sql: sql.to_string(), + database_type: Some(DatabaseType::Mysql), + }); + + assert_eq!(paginated.sql.as_deref(), Some("SELECT * FROM t\n# FOR UPDATE LIMIT 1\nLIMIT 100;")); + assert_eq!( + counted.sql.as_deref(), + Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM t\n# FOR UPDATE LIMIT 1\n) `dbx_count`;") + ); + } + #[test] fn clickhouse_scalar_with_select_can_be_counted() { let result = build_count_query_sql(CountQuerySqlOptions {