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";
const READ_KEYWORDS: [&str; 2] = ["SELECT", "WITH"];
#[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 values: Vec<Value>,
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
)));
}
let (rewritten, values) = resolve(sql, &scanned.params, ¶ms, stream, &sessions)?;
Ok(QueryPlan {
sql: rewritten,
values,
stream: stream.to_string(),
sessions,
params,
timeout: Duration::from_millis(opts.timeout_ms),
limit: opts.limit,
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum Param {
Stream,
Sessions,
Positional,
Numbered,
Named(String),
}
struct Token {
at: Range<usize>,
param: Param,
}
struct Scanned {
keyword: String,
params: Vec<Token>,
}
fn resolve(
sql: &str,
tokens: &[Token],
params: &QueryParams,
stream: &str,
sessions: &[String],
) -> KnlResult<(String, Vec<Value>)> {
const NO_VALUES: &[Value] = &[];
let given: &[Value] = match params {
QueryParams::Positional(values) => values,
_ => NO_VALUES,
};
let mut out = String::with_capacity(sql.len());
let mut values: Vec<Value> = Vec::with_capacity(tokens.len());
let mut cursor = 0;
let mut taken = 0;
for token in tokens {
out.push_str(&sql[cursor..token.at.start]);
cursor = token.at.end;
match &token.param {
Param::Stream => {
out.push('?');
values.push(Value::from(stream));
}
Param::Sessions => {
out.push('(');
for (index, id) in sessions.iter().enumerate() {
if index > 0 {
out.push_str(", ");
}
out.push('?');
values.push(Value::from(id.as_str()));
}
out.push(')');
}
Param::Positional => {
let value = given.get(taken).ok_or_else(|| {
KnlError::Validation(format!(
"the query has more `?` parameters than the {} value(s) given",
given.len()
))
})?;
taken += 1;
out.push('?');
values.push(scalar(value)?.clone());
}
Param::Numbered => {
return Err(KnlError::Validation(format!(
"{:?} is a numbered parameter, and the kernel assigns the positions: \
number your parameters by position (a bare `?`) or name them",
&sql[token.at.clone()]
)));
}
Param::Named(name) => {
let QueryParams::Named(named) = params else {
return Err(KnlError::Validation(format!(
"the query names the parameter {name:?}, so params must be a table of \
names to values"
)));
};
let value = named
.get(&name[1..])
.or_else(|| named.get(name.as_str()))
.ok_or_else(|| {
KnlError::Validation(format!(
"no value was given for the parameter {name:?}"
))
})?;
out.push('?');
values.push(scalar(value)?.clone());
}
}
}
out.push_str(&sql[cursor..]);
if given.len() > taken {
return Err(KnlError::Validation(format!(
"{} value(s) were given for {taken} `?` parameter(s)",
given.len()
)));
}
Ok((out, values))
}
fn scalar(value: &Value) -> KnlResult<&Value> {
match value {
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => Ok(value),
other => Err(KnlError::Validation(format!(
"a {} is not a SQLite value",
type_name_of(other)
))),
}
}
fn type_name_of(value: &Value) -> &'static str {
match value {
Value::Null => "null",
Value::Bool(_) => "boolean",
Value::Number(_) => "number",
Value::String(_) => "string",
Value::Array(_) => "list",
Value::Object(_) => "table",
}
}
fn scan(sql: &str) -> KnlResult<Scanned> {
let bytes = sql.as_bytes();
let mut at = 0;
let mut keyword: Option<String> = None;
let mut params: Vec<Token> = 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 mut end = at + 1;
while end < bytes.len() && bytes[end].is_ascii_digit() {
end += 1;
}
let param = if end == at + 1 {
Param::Positional
} else {
Param::Numbered
};
params.push(Token { at: at..end, param });
at = end;
}
b'$' | b':' | b'@' => {
let end = ident_end(bytes, at + 1);
if end > at + 1 {
let name = &sql[at..end];
let param = match name {
STREAM_PARAM => Param::Stream,
SESSIONS_TOKEN => Param::Sessions,
other => Param::Named(other.to_string()),
};
params.push(Token { at: at..end, param });
}
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, params })
}
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
}
#[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_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 (?, ?) ORDER BY seq"
);
assert_eq!(plan.sessions, ["stream-one", "stream-two"]);
assert_eq!(plan.values, [json!("stream-one"), json!("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 (?)");
assert_eq!(plan.sessions, ["s-1"]);
assert_eq!(plan.stream, "s-1");
assert_eq!(plan.values, [json!("s-1")]);
}
#[test]
fn only_a_token_outside_quotes_and_comments_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",
r#"SELECT 'a ? b' AS literal"#,
r#"SELECT ':kind' AS literal"#,
"SELECT 1 -- :kind ? @who\n",
] {
let plan = plan_over(sql, &["a", "b"]).expect("plan");
assert_eq!(plan.sql, sql, "the text must be untouched: {sql:?}");
assert!(plan.values.is_empty(), "{sql:?}: {:?}", plan.values);
}
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 (?) UNION \
SELECT * FROM events WHERE stream IN (?)"
);
assert_eq!(plan.values, [json!("a"), json!("a")]);
}
#[test]
fn a_name_that_starts_like_a_reserved_one_is_the_callers() {
let opts = QueryOpts::default();
let mut named = Map::new();
named.insert("sessions2".to_string(), json!("x"));
let plan =
plan("SELECT $sessions2", QueryParams::Named(named), &opts, "s-1").expect("plan");
assert_eq!(plan.sql, "SELECT ?");
assert_eq!(plan.values, [json!("x")]);
}
#[test]
fn stream_resolves_to_the_sessions_own_stream() {
let plan =
plan_of("SELECT * FROM events WHERE stream = $stream ORDER BY seq").expect("plan");
assert_eq!(
plan.sql,
"SELECT * FROM events WHERE stream = ? ORDER BY seq"
);
assert_eq!(plan.values, [json!("s-1")]);
}
#[test]
fn the_callers_parameters_are_resolved_in_order() {
let opts = QueryOpts::default();
let positional = plan(
"SELECT * FROM events WHERE stream = $stream AND kind = ? AND seq > ?",
QueryParams::Positional(vec![json!("note"), json!(3)]),
&opts,
"s-1",
)
.expect("plan");
assert_eq!(
positional.sql,
"SELECT * FROM events WHERE stream = ? AND kind = ? AND seq > ?"
);
assert_eq!(positional.values, [json!("s-1"), json!("note"), json!(3)]);
let mut named = Map::new();
named.insert("kind".to_string(), json!("note"));
named.insert("@who".to_string(), json!("me"));
let by_name = plan(
"SELECT * FROM events WHERE kind = :kind AND stream = @who AND kind = $kind",
QueryParams::Named(named),
&opts,
"s-1",
)
.expect("plan");
assert_eq!(
by_name.sql,
"SELECT * FROM events WHERE kind = ? AND stream = ? AND kind = ?"
);
assert_eq!(by_name.values, [json!("note"), json!("me"), json!("note")]);
}
#[test]
fn every_parameter_is_answered_and_every_value_is_used() {
let opts = QueryOpts::default();
let err = plan_of("SELECT * FROM events WHERE kind = :kind")
.expect_err("an unanswered parameter must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{err}");
assert!(err.reason().contains(":kind"), "{}", err.reason());
let err = plan(
"SELECT * FROM events WHERE kind = :kind",
QueryParams::Named(Map::new()),
&opts,
"s-1",
)
.expect_err("a name with no value must be refused");
assert!(err.reason().contains(":kind"), "{}", err.reason());
let err = plan(
"SELECT * FROM events WHERE kind = ?",
QueryParams::Positional(vec![json!("a"), json!("b")]),
&opts,
"s-1",
)
.expect_err("a value with no parameter must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{err}");
let err = plan(
"SELECT * FROM events WHERE kind = ? AND seq = ?",
QueryParams::Positional(vec![json!("a")]),
&opts,
"s-1",
)
.expect_err("a parameter with no value must be refused");
assert!(err.reason().contains("more `?`"), "{}", err.reason());
}
#[test]
fn a_numbered_parameter_is_refused() {
let opts = QueryOpts::default();
let err = plan(
"SELECT * FROM events WHERE kind = ?1",
QueryParams::Positional(vec![json!("a")]),
&opts,
"s-1",
)
.expect_err("a numbered parameter must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{err}");
assert!(err.reason().contains("?1"), "{}", err.reason());
}
#[test]
fn a_list_or_a_table_is_not_a_value() {
let opts = QueryOpts::default();
for value in [json!([1, 2]), json!({ "a": 1 })] {
let err = plan(
"SELECT ?",
QueryParams::Positional(vec![value.clone()]),
&opts,
"s-1",
)
.expect_err("a composite must be refused");
assert_eq!(err.kind(), KnlError::VALIDATION, "{value}: {err}");
assert!(
err.reason().contains("not a SQLite value"),
"{}",
err.reason()
);
}
let plan = plan(
"SELECT ?, ?, ?, ?",
QueryParams::Positional(vec![Value::Null, json!(true), json!(1.5), json!("text")]),
&opts,
"s-1",
)
.expect("plan");
assert_eq!(plan.values.len(), 4);
}
#[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);
}
}