use crate::ast::{Resolver, Statement};
use crate::error::ParseResult;
use crate::interner::FrozenResolver;
use super::Dialect;
use super::engine::Parser;
pub struct Statements<'a, D: Dialect> {
parser: Parser<'a, D>,
done: bool,
}
impl<'a, D: Dialect> Statements<'a, D> {
pub fn resolver(&self) -> &impl Resolver {
self.parser.live_resolver()
}
pub fn finish(self) -> FrozenResolver {
self.parser.finish()
}
}
impl<'a, D: Dialect> Iterator for Statements<'a, D> {
type Item = ParseResult<Statement<D::Ext>>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
match self.parser.parse_next_statement() {
Ok(Some(statement)) => Some(Ok(statement)),
Ok(None) => {
self.done = true;
None
}
Err(error) => {
self.done = true;
Some(Err(error))
}
}
}
}
pub(crate) fn statements<D: Dialect>(src: &str, dialect: D) -> ParseResult<Statements<'_, D>> {
statements_with_limit(src, dialect, super::engine::DEFAULT_RECURSION_LIMIT, false)
}
pub(crate) fn statements_with_limit<D: Dialect>(
src: &str,
dialect: D,
recursion_limit: usize,
parse_float_as_decimal: bool,
) -> ParseResult<Statements<'_, D>> {
let parser = Parser::streaming(src, dialect)?
.recursion_limit(recursion_limit)
.parse_float_as_decimal(parse_float_as_decimal);
Ok(Statements {
parser,
done: false,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::{Expr, ObjectName, SelectItem, SetExpr, Symbol};
use crate::parser::{TestDialect, parse_with};
fn projection_column_symbol(statement: &Statement) -> Symbol {
let Statement::Query { query, .. } = statement else {
panic!("expected a query statement");
};
let SetExpr::Select { select, .. } = &query.body else {
panic!("expected a SELECT body");
};
match &select.projection[0] {
SelectItem::Expr {
expr:
Expr::Column {
name: ObjectName(parts),
..
},
..
} => parts[0].sym,
other => panic!("expected a single column projection, got {other:?}"),
}
}
#[test]
fn iterates_each_statement_in_order() {
let mut iter = statements("SELECT 1; SELECT 2; SELECT 3", TestDialect).expect("valid");
let collected: Vec<_> = std::iter::from_fn(|| iter.next())
.map(|result| result.expect("each statement parses"))
.collect();
assert_eq!(collected.len(), 3);
}
#[test]
fn empty_input_yields_no_statements() {
let mut iter = statements(" ; ; ", TestDialect).expect("only separators");
assert!(iter.next().is_none());
assert!(iter.next().is_none());
}
#[test]
fn fail_fast_yields_the_error_once_then_stops() {
let mut iter = statements("SELECT 1; SELECT FROM; SELECT 3", TestDialect).expect("setup");
assert!(iter.next().expect("first item").is_ok());
assert!(iter.next().expect("second item").is_err());
assert!(iter.next().is_none());
}
#[test]
fn live_resolver_resolves_each_statement_before_the_next() {
let mut iter = statements("SELECT alpha; SELECT beta", TestDialect).expect("valid");
let first = iter.next().expect("first").expect("parses");
assert_eq!(
iter.resolver().resolve(projection_column_symbol(&first)),
"alpha",
);
let second = iter.next().expect("second").expect("parses");
assert_eq!(
iter.resolver().resolve(projection_column_symbol(&first)),
"alpha",
);
assert_eq!(
iter.resolver().resolve(projection_column_symbol(&second)),
"beta",
);
assert!(iter.next().is_none());
}
#[test]
fn early_drop_after_partial_iteration_is_clean() {
let mut iter = statements("SELECT 1; SELECT 2; SELECT 3", TestDialect).expect("valid");
let _first = iter.next().expect("first").expect("parses");
drop(iter);
}
#[test]
fn parse_with_matches_the_streaming_iterator() {
let src = "SELECT a, b FROM t; SELECT 1 + 2";
let collected = parse_with(src, crate::ParseConfig::new(TestDialect)).expect("collects");
let mut iter = statements(src, TestDialect).expect("streams");
let streamed: Vec<_> = std::iter::from_fn(|| iter.next())
.map(|result| result.expect("parses"))
.collect();
assert_eq!(collected.statements(), streamed.as_slice());
}
#[test]
fn adjacent_statements_require_a_separator() {
use crate::dialect::Postgres;
for sql in [
"SELECT 1 SELECT 2",
"VALUES (1) VALUES (2)",
"TABLE t TABLE t",
"DO '' DO ''",
"Do''Do''",
"DO 'x' DO 'y'",
"DO $$a$$ DO $$b$$",
"DO '' SELECT 1",
"VALUES (1) SELECT 2",
"TABLE t SELECT 1",
] {
assert!(
parse_with(sql, crate::ParseConfig::new(Postgres)).is_err(),
"separator-less statement pair must reject: {sql:?}",
);
}
for sql in [
"SELECT 1; SELECT 2",
"VALUES (1); VALUES (2)",
"TABLE t; TABLE t",
"DO ''; DO ''",
"SELECT 1",
"SELECT 1;",
"DO ''",
"VALUES (1)",
"TABLE t",
] {
assert!(
parse_with(sql, crate::ParseConfig::new(Postgres)).is_ok(),
"`;`-delimited (or lone) statements must accept: {sql:?}",
);
}
}
#[test]
fn finish_yields_a_resolver_for_retained_statements() {
let mut iter = statements("SELECT kept", TestDialect).expect("valid");
let statement = iter.next().expect("one").expect("parses");
let sym = projection_column_symbol(&statement);
let resolver = iter.finish();
assert_eq!(resolver.resolve(sym), "kept");
}
}