use std::ops::Range;
use std::time::Duration;
use serde_json::{Map, Value};
use super::{KnlError, KnlResult};
pub const DEFAULT_TIMEOUT_MS: u64 = 5_000;
pub const DEFAULT_LIMIT: usize = 1_000;
pub const STREAM_PARAM: &str = "$stream";
pub const SESSIONS_TOKEN: &str = "$sessions";
pub const SESSION_SLOT_PREFIX: &str = ":knl_sessions_";
const READ_KEYWORDS: [&str; 2] = ["SELECT", "WITH"];
pub fn session_slot(index: usize) -> String {
format!("{SESSION_SLOT_PREFIX}{index}")
}
#[derive(Debug, Clone, PartialEq, Default)]
pub enum QueryParams {
#[default]
None,
Positional(Vec<Value>),
Named(Map<String, Value>),
}
#[derive(Debug, Clone, PartialEq)]
pub struct QueryOpts {
pub sessions: Option<Vec<String>>,
pub timeout_ms: u64,
pub limit: usize,
}
impl Default for QueryOpts {
fn default() -> Self {
Self {
sessions: None,
timeout_ms: DEFAULT_TIMEOUT_MS,
limit: DEFAULT_LIMIT,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct QueryPlan {
pub sql: String,
pub stream: String,
pub sessions: Vec<String>,
pub params: QueryParams,
pub timeout: Duration,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct QueryRows {
pub rows: Vec<Map<String, Value>>,
pub truncated: bool,
}
pub fn plan(
sql: &str,
params: QueryParams,
opts: &QueryOpts,
stream: &str,
) -> KnlResult<QueryPlan> {
let sessions = match opts.sessions.as_ref() {
None => vec![stream.to_string()],
Some(list) if list.is_empty() => {
return Err(KnlError::Validation(
"opts.sessions is empty; a set that selects no stream is not a request \
(omit it to read this session's own)"
.to_string(),
));
}
Some(list) => list.clone(),
};
if opts.timeout_ms == 0 {
return Err(KnlError::Validation(
"opts.timeout_ms must be a positive whole number of milliseconds".to_string(),
));
}
let scanned = scan(sql)?;
if !READ_KEYWORDS.contains(&scanned.keyword.as_str()) {
return Err(KnlError::Validation(format!(
"a query reads: it must start with SELECT or WITH, got {:?}",
scanned.keyword
)));
}
Ok(QueryPlan {
sql: expand_sessions(sql, &scanned.sessions, sessions.len()),
stream: stream.to_string(),
sessions,
params,
timeout: Duration::from_millis(opts.timeout_ms),
limit: opts.limit,
})
}
struct Scanned {
keyword: String,
sessions: Vec<Range<usize>>,
}
fn scan(sql: &str) -> KnlResult<Scanned> {
let bytes = sql.as_bytes();
let mut at = 0;
let mut keyword: Option<String> = None;
let mut sessions: Vec<Range<usize>> = Vec::new();
let mut ended = false;
while at < bytes.len() {
let byte = bytes[at];
if byte.is_ascii_whitespace() {
at += 1;
continue;
}
if byte == b'-' && bytes.get(at + 1) == Some(&b'-') {
at += 2;
while at < bytes.len() && bytes[at] != b'\n' {
at += 1;
}
continue;
}
if byte == b'/' && bytes.get(at + 1) == Some(&b'*') {
at += 2;
while at < bytes.len() && !(bytes[at] == b'*' && bytes.get(at + 1) == Some(&b'/')) {
at += 1;
}
at = usize::min(at + 2, bytes.len());
continue;
}
if ended {
return Err(KnlError::Validation(
"a query is one statement: there is SQL after the `;`".to_string(),
));
}
match byte {
b';' => {
ended = true;
at += 1;
}
b'\'' | b'"' | b'`' | b'[' => at = skip_quoted(bytes, at),
b'$' => {
let end = ident_end(bytes, at + 1);
if &sql[at..end] == SESSIONS_TOKEN {
sessions.push(at..end);
}
at = end;
}
_ if byte.is_ascii_alphabetic() || byte == b'_' => {
let end = ident_end(bytes, at);
if keyword.is_none() {
keyword = Some(sql[at..end].to_ascii_uppercase());
}
at = end;
}
_ => at += 1,
}
}
let keyword = keyword.ok_or_else(|| {
KnlError::Validation("a query needs a statement; the SQL is empty".to_string())
})?;
Ok(Scanned { keyword, sessions })
}
fn skip_quoted(bytes: &[u8], start: usize) -> usize {
let open = bytes[start];
let close = if open == b'[' { b']' } else { open };
let doubles = open != b'[';
let mut at = start + 1;
while at < bytes.len() {
if bytes[at] == close {
if doubles && bytes.get(at + 1) == Some(&close) {
at += 2;
continue;
}
return at + 1;
}
at += 1;
}
at
}
fn ident_end(bytes: &[u8], from: usize) -> usize {
let mut at = from;
while at < bytes.len()
&& (bytes[at].is_ascii_alphanumeric() || matches!(bytes[at], b'_' | b'$'))
{
at += 1;
}
at
}
fn expand_sessions(sql: &str, tokens: &[Range<usize>], count: usize) -> String {
if tokens.is_empty() {
return sql.to_string();
}
let slots = (0..count).map(session_slot).collect::<Vec<_>>().join(", ");
let slots = format!("({slots})");
let mut out = String::with_capacity(sql.len() + tokens.len() * slots.len());
let mut cursor = 0;
for token in tokens {
out.push_str(&sql[cursor..token.start]);
out.push_str(&slots);
cursor = token.end;
}
out.push_str(&sql[cursor..]);
out
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn plan_of(sql: &str) -> KnlResult<QueryPlan> {
plan(sql, QueryParams::None, &QueryOpts::default(), "s-1")
}
fn plan_over(sql: &str, sessions: &[&str]) -> KnlResult<QueryPlan> {
let opts = QueryOpts {
sessions: Some(sessions.iter().map(|s| (*s).to_string()).collect()),
..QueryOpts::default()
};
plan(sql, QueryParams::None, &opts, "s-1")
}
#[test]
fn a_select_or_with_is_a_query_whatever_it_is_dressed_in() {
for sql in [
"SELECT 1",
"select 1",
" \n\t select 1",
"-- a comment first\nSELECT 1",
"/* and a block one */ WITH x AS (SELECT 1) SELECT * FROM x",
"SELECT 1;",
"SELECT 1; -- trailing comment\n",
] {
plan_of(sql).unwrap_or_else(|e| panic!("{sql:?} must be a query: {e}"));
}
}
#[test]
fn anything_that_is_not_a_read_is_refused_as_validation() {
for sql in [
"INSERT INTO events (stream) VALUES ('x')",
"UPDATE events SET kind = 'x'",
"DELETE FROM events",
"DROP TABLE events",
"PRAGMA table_info(events)",
"ATTACH DATABASE '/tmp/other.db' AS other",
"BEGIN",
"VACUUM",
] {
let err = plan_of(sql).expect_err("a statement that is not a read must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{sql:?}: {err}");
assert!(
err.reason().contains("SELECT or WITH"),
"{sql:?}: {}",
err.reason()
);
}
let err = plan_of(" ").expect_err("empty SQL must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION);
}
#[test]
fn a_second_statement_is_refused() {
for sql in [
"SELECT 1; DELETE FROM events",
"SELECT 1;SELECT 2",
"SELECT 1; -- a comment\n DROP TABLE events",
] {
let err = plan_of(sql).expect_err("a second statement must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{sql:?}: {err}");
assert!(err.reason().contains("one statement"), "{}", err.reason());
}
}
#[test]
fn a_semicolon_in_a_literal_or_a_comment_is_not_a_statement_boundary() {
for sql in [
r#"SELECT * FROM events WHERE kind = 'a;b'"#,
r#"SELECT * FROM events WHERE kind = 'it''s;fine'"#,
r#"SELECT "we;ird" FROM events"#,
"SELECT 1 -- ; not a boundary\n",
"SELECT /* ; */ 1",
] {
plan_of(sql).unwrap_or_else(|e| panic!("{sql:?} must be one statement: {e}"));
}
}
#[test]
fn sessions_expands_to_one_named_placeholder_per_id() {
let plan = plan_over(
"SELECT * FROM events WHERE stream IN $sessions ORDER BY seq",
&["stream-one", "stream-two"],
)
.expect("plan");
assert_eq!(
plan.sql,
"SELECT * FROM events WHERE stream IN (:knl_sessions_0, :knl_sessions_1) ORDER BY seq"
);
assert_eq!(plan.sessions, ["stream-one", "stream-two"]);
for id in &plan.sessions {
assert!(
!plan.sql.contains(id.as_str()),
"an id is bound, never written into the SQL: {}",
plan.sql
);
}
let plan = plan_of("SELECT * FROM events WHERE stream IN $sessions").expect("plan");
assert_eq!(
plan.sql,
"SELECT * FROM events WHERE stream IN (:knl_sessions_0)"
);
assert_eq!(plan.sessions, ["s-1"]);
assert_eq!(plan.stream, "s-1");
}
#[test]
fn only_the_token_itself_is_rewritten() {
for sql in [
r#"SELECT '$sessions' AS literal"#,
r#"SELECT "$sessions" FROM events"#,
"SELECT 1 -- $sessions in a comment\n",
"SELECT /* $sessions */ 1",
"SELECT $sessions2",
] {
let plan = plan_over(sql, &["a", "b"]).expect("plan");
assert_eq!(plan.sql, sql, "the text must be untouched: {sql:?}");
}
let plan = plan_over(
"SELECT * FROM events WHERE stream IN $sessions UNION \
SELECT * FROM events WHERE stream IN $sessions",
&["a"],
)
.expect("plan");
assert_eq!(
plan.sql,
"SELECT * FROM events WHERE stream IN (:knl_sessions_0) UNION \
SELECT * FROM events WHERE stream IN (:knl_sessions_0)"
);
}
#[test]
fn an_empty_session_set_is_refused() {
let err = plan_over("SELECT 1", &[]).expect_err("an empty set must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION);
assert!(err.reason().contains("sessions"), "{}", err.reason());
}
#[test]
fn a_zero_timeout_is_refused() {
let opts = QueryOpts {
timeout_ms: 0,
..QueryOpts::default()
};
let err = plan("SELECT 1", QueryParams::None, &opts, "s-1")
.expect_err("a zero timeout must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION);
assert!(err.reason().contains("timeout_ms"), "{}", err.reason());
}
#[test]
fn the_plan_carries_the_callers_values_for_binding() {
let opts = QueryOpts::default();
let plan = plan(
"SELECT * FROM events WHERE kind = ?",
QueryParams::Positional(vec![json!("note")]),
&opts,
"s-1",
)
.expect("plan");
assert_eq!(plan.params, QueryParams::Positional(vec![json!("note")]));
assert_eq!(plan.timeout, Duration::from_millis(DEFAULT_TIMEOUT_MS));
assert_eq!(plan.limit, DEFAULT_LIMIT);
}
}