diff --git a/nodedb/src/control/server/shared/ddl/neutral/query_functions/router.rs b/nodedb/src/control/server/shared/ddl/neutral/query_functions/router.rs index 6a5297c61..26b1a0844 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/query_functions/router.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/query_functions/router.rs @@ -17,6 +17,111 @@ use crate::types::DatabaseId; use super::super::super::result::{DdlError, DdlResult}; +/// Scan `sql` as code: quoted regions and comments are replaced by spaces so a +/// `contains` check cannot match a keyword inside a string literal, a quoted +/// identifier, or a comment. Everything else is preserved verbatim. +fn scan_code(sql: &str) -> String { + #[derive(PartialEq)] + enum St { + Code, + Single, + Double, + Line, + Block, + } + + let mut out = String::with_capacity(sql.len()); + let mut st = St::Code; + let mut chars = sql.chars().peekable(); + while let Some(c) = chars.next() { + match st { + St::Code => match c { + '\'' => { + st = St::Single; + out.push(' '); + } + '"' => { + st = St::Double; + out.push(' '); + } + '-' if chars.peek() == Some(&'-') => { + chars.next(); + st = St::Line; + out.push_str(" "); + } + '/' if chars.peek() == Some(&'*') => { + chars.next(); + st = St::Block; + out.push_str(" "); + } + _ => out.push(c), + }, + St::Single => { + if c == '\'' { + if chars.peek() == Some(&'\'') { + chars.next(); + out.push_str(" "); + } else { + st = St::Code; + out.push(' '); + } + } else { + out.push(' '); + } + } + St::Double => { + if c == '"' { + if chars.peek() == Some(&'"') { + chars.next(); + out.push_str(" "); + } else { + st = St::Code; + out.push(' '); + } + } else { + out.push(' '); + } + } + St::Line => { + if c == '\n' { + st = St::Code; + out.push('\n'); + } else { + out.push(' '); + } + } + St::Block => { + if c == '*' && chars.peek() == Some(&'/') { + chars.next(); + st = St::Code; + out.push_str(" "); + } else { + out.push(' '); + } + } + } + } + out +} + +/// The query function this statement routes to, if any. +/// +/// The routing decision: `scan_code` removes quoted regions and comments, then +/// the keywords are matched in the pgwire `router::function::dispatch` order. +/// Kept separate from [`try_dispatch`] so the decision itself is testable. +fn recognized_function(sql: &str) -> Option<&'static str> { + const KEYWORDS: [&str; 6] = [ + "VERIFY_AUDIT_CHAIN", + "VERIFY_HASH_CHAIN", + "BALANCE_AS_OF", + "TEMPORAL_LOOKUP", + "VERIFY_BALANCE", + "CONVERT_CURRENCY_LOOKUP", + ]; + let upper = scan_code(sql).to_uppercase(); + KEYWORDS.into_iter().find(|keyword| upper.contains(keyword)) +} + /// Try to handle `sql` as one of the temporal / audit query functions. /// /// Returns `Some(result)` when a substring matches (mirroring the pgwire @@ -28,26 +133,128 @@ pub async fn try_dispatch( database_id: DatabaseId, sql: &str, ) -> Option, DdlError>> { - let upper = sql.to_uppercase(); + match recognized_function(sql) { + Some("VERIFY_AUDIT_CHAIN") => { + Some(super::verify_audit_chain(state, identity, database_id, sql).await) + } + Some("VERIFY_HASH_CHAIN") => { + Some(super::verify_hash_chain(state, identity, database_id, sql).await) + } + Some("BALANCE_AS_OF") => { + Some(super::balance_as_of(state, identity, database_id, sql).await) + } + Some("TEMPORAL_LOOKUP") => { + Some(super::temporal_lookup(state, identity, database_id, sql).await) + } + Some("VERIFY_BALANCE") => { + Some(super::verify_balance(state, identity, database_id, sql).await) + } + Some("CONVERT_CURRENCY_LOOKUP") => { + Some(super::convert_currency_lookup(state, identity, database_id, sql).await) + } + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::{recognized_function, scan_code}; - if upper.contains("VERIFY_AUDIT_CHAIN") { - return Some(super::verify_audit_chain(state, identity, database_id, sql).await); + /// The routing decision itself: which function (if any) `try_dispatch` + /// picks for this statement. + fn hits(sql: &str, keyword: &str) -> bool { + recognized_function(sql) == Some(keyword) } - if upper.contains("VERIFY_HASH_CHAIN") { - return Some(super::verify_hash_chain(state, identity, database_id, sql).await); + + #[test] + fn scanned_text_blanks_the_literal() { + assert!( + !scan_code("SELECT 'verify_balance'") + .to_ascii_uppercase() + .contains("VERIFY_BALANCE") + ); } - if upper.contains("BALANCE_AS_OF") { - return Some(super::balance_as_of(state, identity, database_id, sql).await); + + #[test] + fn literal_contents_do_not_route() { + assert!(!hits( + "INSERT INTO kg (id, name) VALUES ('tmp_verify_balance_x', 'x')", + "VERIFY_BALANCE" + )); + assert!(!hits( + "INSERT INTO kg (id, label) VALUES ('x', 'verify_balance')", + "VERIFY_BALANCE" + )); + } + + #[test] + fn escaped_quote_keeps_the_literal_open() { + assert!(!hits( + "INSERT INTO kg (id) VALUES ('it''s verify_balance')", + "VERIFY_BALANCE" + )); } - if upper.contains("TEMPORAL_LOOKUP") { - return Some(super::temporal_lookup(state, identity, database_id, sql).await); + + #[test] + fn quoted_identifiers_do_not_route() { + assert!(!hits( + "SELECT 1 AS \"verify_balance\" FROM kg", + "VERIFY_BALANCE" + )); } - if upper.contains("VERIFY_BALANCE") { - return Some(super::verify_balance(state, identity, database_id, sql).await); + + #[test] + fn comments_do_not_route() { + assert!(!hits("SELECT 1 -- verify_balance note", "VERIFY_BALANCE")); + assert!(!hits("SELECT /* verify_balance */ 1", "VERIFY_BALANCE")); + assert!(!hits( + "SELECT 1 -- verify_audit_chain\n", + "VERIFY_AUDIT_CHAIN" + )); } - if upper.contains("CONVERT_CURRENCY_LOOKUP") { - return Some(super::convert_currency_lookup(state, identity, database_id, sql).await); + + #[test] + fn real_calls_still_route() { + assert!(hits("SELECT VERIFY_BALANCE('c', 'col')", "VERIFY_BALANCE")); + assert!(hits("select verify_balance('c','col')", "VERIFY_BALANCE")); + assert!(hits( + "SELECT VERIFY_AUDIT_CHAIN(1, 100)", + "VERIFY_AUDIT_CHAIN" + )); + assert!(hits("SELECT balance_as_of('c','k','v',1)", "BALANCE_AS_OF")); } - None + #[test] + fn blanking_does_not_fuse_neighbouring_tokens() { + // Spaces replace the quoted region, so the halves stay separate. A + // removal-style blank would glue VERIFY and _BALANCE into a match. + assert!(!hits("SELECT VERIFY'x'_BALANCE()", "VERIFY_BALANCE")); + assert!(!hits("SELECT 'verify'_balance()", "VERIFY_BALANCE")); + } + + #[test] + fn escapes_and_unterminated_regions_do_not_route() { + // `""` inside a quoted identifier stays in that identifier. + assert!(!hits( + "SELECT 1 AS \"a\"\"verify_balance\" FROM kg", + "VERIFY_BALANCE" + )); + // An unterminated block comment swallows the rest of the statement. + assert!(!hits("SELECT 1 /* verify_balance", "VERIFY_BALANCE")); + // So does an unterminated literal. + assert!(!hits("SELECT 'verify_balance", "VERIFY_BALANCE")); + } + + #[test] + fn routing_decision_picks_the_first_keyword_in_order() { + assert_eq!( + recognized_function("SELECT VERIFY_AUDIT_CHAIN(1, 2) -- VERIFY_BALANCE"), + Some("VERIFY_AUDIT_CHAIN") + ); + assert_eq!(recognized_function("SELECT 1"), None); + assert_eq!( + recognized_function("INSERT INTO kg { id: 'a', name: 'balance_as_of' }"), + None + ); + } } diff --git a/nodedb/tests/wire/cases/router_misroute_literals.rs b/nodedb/tests/wire/cases/router_misroute_literals.rs index b1188be31..eef28b35c 100644 --- a/nodedb/tests/wire/cases/router_misroute_literals.rs +++ b/nodedb/tests/wire/cases/router_misroute_literals.rs @@ -7,6 +7,11 @@ //! a doc-object UPSERT value, a string literal, a comment — belongs to its own //! handler and must never reach the function arm. These tests hold both //! directions: the literal stores verbatim, the anchored form still routes. +//! +//! The query-function family (`VERIFY_BALANCE`, `VERIFY_AUDIT_CHAIN`, +//! `VERIFY_HASH_CHAIN`, `BALANCE_AS_OF`, `TEMPORAL_LOOKUP`, +//! `CONVERT_CURRENCY_LOOKUP`) is matched by `contains` over the statement, so +//! the same rule applies to a token inside a literal there. use crate::harness::TestServer; @@ -160,3 +165,49 @@ async fn leading_whitespace_still_routes_anchored_arms() { "row must be returned despite leading whitespace" ); } + +/// The six query-function keywords are matched by `contains`, not by a prefix +/// anchor, so a value carrying one used to reach the function arm and answer +/// with its argument error. The literal must store verbatim, and the anchored +/// call must still reach the function arm. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn value_carrying_verify_balance_does_not_misroute() { + let server = TestServer::start().await; + + server + .exec( + "CREATE COLLECTION lit_qf (id STRING PRIMARY KEY, name STRING) \ + WITH (engine='kv')", + ) + .await + .expect("create collection"); + + // The value carries a query-function token. The parenthesised form is a + // non-DDL parse, so it reaches the neutral dispatcher where the pre-fix + // substring match hijacked it; the brace form never gets there. + server + .exec("INSERT INTO lit_qf (id, name) VALUES ('a', 'verify_balance')") + .await + .expect("INSERT with a query-function token in a value must not misroute"); + + let rows = server + .query_text("SELECT name FROM lit_qf WHERE id = 'a'") + .await + .expect("read back the inserted row"); + assert_eq!( + rows, + vec!["verify_balance".to_string()], + "the value must be stored verbatim" + ); + + // The anchored form still routes: the function arm reports the collection + // it was asked for, so a missing collection names it. + let err = server + .query_text("SELECT VERIFY_BALANCE('lit_qf_missing', 'name')") + .await + .expect_err("the anchored call must reach the function arm"); + assert!( + err.contains("lit_qf_missing"), + "the anchored call must reach the function arm, got: {err}" + ); +}