use std::borrow::Cow;
use std::fmt;
use std::sync::{Arc, OnceLock};
use crate::ast::dialect::StringLiteralSyntax;
use crate::ast::render::{Render, RenderConfig, RenderCtx, RenderExt as _};
use crate::ast::{
Extension, LineIndex, Literal, LiteralValueError, NoExt, SourceStore, Span, Statement,
};
use crate::interner::FrozenResolver;
use crate::tokenizer::{TriviaIndex, TriviaRange};
use super::clause_marks::{ClauseMark, ClauseMarkIndex};
#[derive(Debug)]
pub struct Parsed<S: SourceStore = Arc<str>, X: Extension = NoExt> {
source: S,
resolver: FrozenResolver,
statements: Vec<Statement<X>>,
string_literals: StringLiteralSyntax,
line_index: OnceLock<LineIndex>,
trivia: TriviaIndex,
clause_marks: ClauseMarkIndex,
}
pub type StockParsed<S = Arc<str>> = Parsed<S, NoExt>;
impl<S: SourceStore, X: Extension> Parsed<S, X> {
pub(crate) fn new(
source: S,
resolver: FrozenResolver,
statements: Vec<Statement<X>>,
string_literals: StringLiteralSyntax,
) -> Self {
Self {
source,
resolver,
statements,
string_literals,
line_index: OnceLock::new(),
trivia: TriviaIndex::default(),
clause_marks: ClauseMarkIndex::default(),
}
}
pub(crate) fn with_trivia(mut self, trivia: TriviaIndex) -> Self {
self.trivia = trivia;
self
}
pub(crate) fn with_clause_marks(mut self, clause_marks: ClauseMarkIndex) -> Self {
self.clause_marks = clause_marks;
self
}
pub fn statements(&self) -> &[Statement<X>] {
&self.statements
}
pub fn into_statements(self) -> Vec<Statement<X>> {
self.statements
}
pub fn source(&self) -> &str {
&self.source
}
pub fn resolver(&self) -> &FrozenResolver {
&self.resolver
}
pub fn string_literal_syntax(&self) -> StringLiteralSyntax {
self.string_literals
}
pub fn literal_str(&self, literal: &Literal) -> Result<Cow<'_, str>, LiteralValueError> {
literal.as_str_in(self.source(), self.string_literals)
}
pub fn line_index(&self) -> &LineIndex {
self.line_index
.get_or_init(|| LineIndex::from_str(&self.source))
}
pub fn line_col(&self, offset: u32) -> (u32, u32) {
self.line_index().lookup(offset)
}
pub fn span_line_col(&self, span: Span) -> Option<((u32, u32), (u32, u32))> {
if span.is_synthetic() {
return None;
}
let index = self.line_index();
Some((index.lookup(span.start()), index.lookup(span.end())))
}
pub fn trivia(&self) -> &[TriviaRange] {
self.trivia.all()
}
pub fn trivia_in(&self, span: Span) -> &[TriviaRange] {
self.trivia.in_span(span)
}
pub fn trivia_before(&self, offset: u32) -> &[TriviaRange] {
self.trivia.before(offset)
}
pub fn clause_marks(&self) -> &[ClauseMark] {
self.clause_marks.all()
}
pub fn clause_marks_in(&self, span: Span) -> &[ClauseMark] {
self.clause_marks.in_span(span)
}
}
impl<S: SourceStore, X: Extension + Render> fmt::Display for Parsed<S, X> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let config = RenderConfig::default();
let ctx = RenderCtx::new(self.resolver(), self.source(), &config);
self.render_statements_into(&ctx, f)
}
}
impl<S: SourceStore, X: Extension + Render> Parsed<S, X> {
fn render_statements_into<'r, W: fmt::Write>(
&self,
ctx: &'r RenderCtx<'r>,
out: &mut W,
) -> fmt::Result {
for (i, statement) in self.statements().iter().enumerate() {
if i > 0 {
out.write_str("; ")?;
}
write!(out, "{}", statement.displayed(ctx))?;
}
Ok(())
}
pub fn render_into(&self, out: &mut String) -> fmt::Result {
let config = RenderConfig::default();
let ctx = RenderCtx::new(self.resolver(), self.source(), &config);
self.render_statements_into(&ctx, out)
}
pub fn to_sql(&self) -> String {
let mut out = String::with_capacity(self.source().len());
self.render_into(&mut out)
.expect("rendering to a String cannot fail");
out
}
}
#[cfg(any(feature = "serde-serialize", feature = "serde-deserialize"))]
mod serde_impls {
use super::*;
#[cfg(feature = "serde-deserialize")]
use crate::ast::generated::visit::{self, Visit};
#[cfg(feature = "serde-deserialize")]
use crate::ast::serde_depth::{DEFAULT_DESERIALIZE_DEPTH, DepthLimited};
#[cfg(feature = "serde-deserialize")]
use crate::ast::*;
#[cfg(feature = "serde-deserialize")]
use crate::interner::Interner;
#[cfg(feature = "serde-deserialize")]
use serde::Deserialize;
#[cfg(feature = "serde-serialize")]
use serde::Serialize;
#[cfg(feature = "serde-deserialize")]
use serde::de::{Deserializer, Error as _};
#[cfg_attr(feature = "serde-serialize", derive(Serialize))]
#[cfg_attr(feature = "serde-deserialize", derive(Deserialize))]
#[serde(remote = "StringLiteralSyntax")]
struct StringLiteralSyntaxDef {
escape_strings: bool,
dollar_quoted_strings: bool,
national_strings: bool,
double_quoted_strings: bool,
backslash_escapes: bool,
unicode_strings: bool,
bit_string_literals: bool,
blob_literals: bool,
charset_introducers: bool,
same_line_adjacent_concat: bool,
}
#[cfg(feature = "serde-serialize")]
impl<S: SourceStore, X: Extension + Serialize> Serialize for Parsed<S, X> {
fn serialize<Sr>(&self, serializer: Sr) -> Result<Sr::Ok, Sr::Error>
where
Sr: serde::Serializer,
{
#[derive(Serialize)]
#[serde(bound(serialize = "X: Serialize"))]
struct ParsedRef<'a, X: Extension> {
source: &'a str,
symbols: &'a [Box<str>],
#[serde(with = "StringLiteralSyntaxDef")]
string_literals: StringLiteralSyntax,
statements: &'a [Statement<X>],
}
ParsedRef {
source: self.source(),
symbols: self.resolver.dynamic_strings(),
string_literals: self.string_literals,
statements: &self.statements,
}
.serialize(serializer)
}
}
#[cfg(feature = "serde-deserialize")]
impl<'de, S, X> Deserialize<'de> for Parsed<S, X>
where
S: SourceStore + From<String>,
X: Extension + Deserialize<'de>,
{
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
Self::deserialize_with_depth(deserializer, DEFAULT_DESERIALIZE_DEPTH)
}
}
#[cfg(feature = "serde-deserialize")]
impl<S: SourceStore + From<String>, X: Extension> Parsed<S, X> {
pub fn deserialize_with_depth<'de, D>(
deserializer: D,
max_depth: usize,
) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
X: Deserialize<'de>,
{
#[derive(Deserialize)]
#[serde(bound(deserialize = "X: Deserialize<'de>"))]
struct ParsedData<X: Extension> {
source: String,
symbols: Vec<String>,
#[serde(with = "StringLiteralSyntaxDef")]
string_literals: StringLiteralSyntax,
statements: Vec<Statement<X>>,
}
let data = ParsedData::<X>::deserialize(DepthLimited::new(deserializer, max_depth))?;
let mut interner = Interner::new();
for text in &data.symbols {
interner.intern_nonkeyword(text);
}
let resolver = interner.freeze();
if resolver.dynamic_strings().len() != data.symbols.len() {
return Err(D::Error::custom(format!(
"symbol table carries duplicate entries: {} declared, {} distinct \
after interning (duplicates would silently misresolve symbols)",
data.symbols.len(),
resolver.dynamic_strings().len(),
)));
}
validate(&data.source, &resolver, &data.statements).map_err(D::Error::custom)?;
Ok(Parsed::new(
S::from(data.source),
resolver,
data.statements,
data.string_literals,
))
}
}
#[cfg(feature = "serde-deserialize")]
fn validate<X: Extension>(
source: &str,
resolver: &FrozenResolver,
statements: &[Statement<X>],
) -> Result<(), String> {
let source_len = u32::try_from(source.len())
.map_err(|_| "source length exceeds the u32 span range".to_string())?;
let mut validator = DocumentValidator {
source_len,
resolver,
visited: 0,
error: None,
};
for statement in statements {
if validator.error.is_some() {
break;
}
validator.visit_statement(statement);
}
match validator.error {
Some(message) => Err(message),
None => Ok(()),
}
}
#[cfg(feature = "serde-deserialize")]
struct DocumentValidator<'a> {
source_len: u32,
resolver: &'a FrozenResolver,
#[cfg_attr(not(test), allow(dead_code))]
visited: usize,
error: Option<String>,
}
#[cfg(feature = "serde-deserialize")]
impl DocumentValidator<'_> {
fn check_span(&mut self, node: &'static str, span: Span) {
self.visited += 1;
if span.is_synthetic() {
return;
}
let start = span.start();
let end = span.end();
if start > end || end > self.source_len {
self.error.get_or_insert_with(|| {
format!(
"node {node} carries span {start}..{end}, out of bounds for the \
{}-byte source (requires start <= end <= source length)",
self.source_len,
)
});
}
}
fn check_symbol(&mut self, sym: Symbol) {
if self.resolver.try_resolve(sym).is_none() {
self.error.get_or_insert_with(|| {
format!(
"symbol {} is absent from the deserialized symbol table (would \
panic on canonical render)",
sym.as_u32(),
)
});
}
}
}
#[cfg(feature = "serde-deserialize")]
macro_rules! check_spans {
($lt:lifetime, $(($method:ident, $walk:ident, $ty:ty)),+ $(,)?) => {
$(
fn $method(&mut self, node: &$lt $ty) {
if self.error.is_some() {
return;
}
self.check_span(stringify!($method), node.span());
visit::$walk::<Self, X>(self, node);
}
)+
};
}
#[cfg(feature = "serde-deserialize")]
impl<'ast, X: Extension> Visit<'ast, X> for DocumentValidator<'_> {
check_spans!('ast,
(visit_session_statement, walk_session_statement, SessionStatement<X>),
(visit_set_value, walk_set_value, SetValue),
(visit_set_parameter_value, walk_set_parameter_value, SetParameterValue),
(visit_special_set_value, walk_special_set_value, SpecialSetValue),
(visit_constraints_target, walk_constraints_target, ConstraintsTarget),
(visit_set_names_value, walk_set_names_value, SetNamesValue),
(visit_config_parameter, walk_config_parameter, ConfigParameter),
(visit_access_control_statement, walk_access_control_statement, AccessControlStatement<X>),
(visit_privileges, walk_privileges, Privileges),
(visit_privilege, walk_privilege, Privilege),
(visit_grant_object, walk_grant_object, GrantObject<X>),
(visit_routine_signature, walk_routine_signature, RoutineSignature<X>),
(visit_grantee, walk_grantee, Grantee),
(visit_role_spec, walk_role_spec, RoleSpec),
(visit_create_table, walk_create_table, CreateTable<X>),
(visit_create_table_body, walk_create_table_body, CreateTableBody<X>),
(visit_table_element, walk_table_element, TableElement<X>),
(visit_column_def, walk_column_def, ColumnDef<X>),
(visit_column_constraint, walk_column_constraint, ColumnConstraint<X>),
(visit_constraint_characteristics, walk_constraint_characteristics, ConstraintCharacteristics),
(visit_column_option, walk_column_option, ColumnOption<X>),
(visit_foreign_key_ref, walk_foreign_key_ref, ForeignKeyRef),
(visit_referential_action, walk_referential_action, ReferentialAction),
(visit_generated_column, walk_generated_column, GeneratedColumn<X>),
(visit_identity_column, walk_identity_column, IdentityColumn<X>),
(visit_identity_option, walk_identity_option, IdentityOption<X>),
(visit_table_constraint_def, walk_table_constraint_def, TableConstraintDef<X>),
(visit_table_constraint, walk_table_constraint, TableConstraint<X>),
(visit_create_table_option, walk_create_table_option, CreateTableOption<X>),
(visit_create_table_option_kind, walk_create_table_option_kind, CreateTableOptionKind<X>),
(visit_table_option, walk_table_option, TableOption),
(visit_table_option_value, walk_table_option_value, TableOptionValue),
(visit_table_storage_parameter, walk_table_storage_parameter, TableStorageParameter<X>),
(visit_alter_table, walk_alter_table, AlterTable<X>),
(visit_alter_table_action, walk_alter_table_action, AlterTableAction<X>),
(visit_alter_column_action, walk_alter_column_action, AlterColumnAction<X>),
(visit_drop_statement, walk_drop_statement, DropStatement),
(visit_comment_on_statement, walk_comment_on_statement, CommentOnStatement<X>),
(visit_create_schema, walk_create_schema, CreateSchema<X>),
(visit_create_view, walk_create_view, CreateView<X>),
(visit_create_index, walk_create_index, CreateIndex<X>),
(visit_index_column, walk_index_column, IndexColumn<X>),
(visit_create_trigger, walk_create_trigger, CreateTrigger<X>),
(visit_trigger_event, walk_trigger_event, TriggerEvent),
(visit_create_database, walk_create_database, CreateDatabase),
(visit_create_function, walk_create_function, CreateFunction<X>),
(visit_function_param, walk_function_param, FunctionParam<X>),
(visit_function_option, walk_function_option, FunctionOption<X>),
(visit_insert, walk_insert, Insert<X>),
(visit_insert_target, walk_insert_target, InsertTarget),
(visit_insert_source, walk_insert_source, InsertSource<X>),
(visit_insert_values, walk_insert_values, InsertValues<X>),
(visit_insert_value, walk_insert_value, InsertValue<X>),
(visit_dml_target, walk_dml_target, DmlTarget),
(visit_update, walk_update, Update<X>),
(visit_update_assignment, walk_update_assignment, UpdateAssignment<X>),
(visit_update_value, walk_update_value, UpdateValue<X>),
(visit_update_tuple_source, walk_update_tuple_source, UpdateTupleSource<X>),
(visit_delete, walk_delete, Delete<X>),
(visit_dml_selection, walk_dml_selection, DmlSelection<X>),
(visit_default_value, walk_default_value, DefaultValue),
(visit_returning, walk_returning, Returning<X>),
(visit_upsert, walk_upsert, Upsert<X>),
(visit_on_conflict, walk_on_conflict, OnConflict<X>),
(visit_conflict_target, walk_conflict_target, ConflictTarget<X>),
(visit_conflict_action, walk_conflict_action, ConflictAction<X>),
(visit_merge, walk_merge, Merge<X>),
(visit_merge_when_clause, walk_merge_when_clause, MergeWhenClause<X>),
(visit_merge_action, walk_merge_action, MergeAction<X>),
(visit_subscript_expr, walk_subscript_expr, SubscriptExpr<X>),
(visit_collate_expr, walk_collate_expr, CollateExpr<X>),
(visit_at_time_zone_expr, walk_at_time_zone_expr, AtTimeZoneExpr<X>),
(visit_array_expr, walk_array_expr, ArrayExpr<X>),
(visit_row_expr, walk_row_expr, RowExpr<X>),
(visit_field_selection_expr, walk_field_selection_expr, FieldSelectionExpr<X>),
(visit_function_call, walk_function_call, FunctionCall<X>),
(visit_case_expr, walk_case_expr, CaseExpr<X>),
(visit_when_clause, walk_when_clause, WhenClause<X>),
(visit_extract_expr, walk_extract_expr, ExtractExpr<X>),
(visit_literal, walk_literal, Literal),
(visit_query, walk_query, Query<X>),
(visit_set_expr, walk_set_expr, SetExpr<X>),
(visit_with, walk_with, With<X>),
(visit_cte, walk_cte, Cte<X>),
(visit_values, walk_values, Values<X>),
(visit_values_item, walk_values_item, ValuesItem<X>),
(visit_select, walk_select, Select<X>),
(visit_into_target, walk_into_target, IntoTarget),
(visit_group_by_item, walk_group_by_item, GroupByItem<X>),
(visit_select_item, walk_select_item, SelectItem<X>),
(visit_select_distinct, walk_select_distinct, SelectDistinct<X>),
(visit_table_with_joins, walk_table_with_joins, TableWithJoins<X>),
(visit_table_alias, walk_table_alias, TableAlias),
(visit_table_sample, walk_table_sample, TableSample<X>),
(visit_table_function_column, walk_table_function_column, TableFunctionColumn<X>),
(visit_rows_from_item, walk_rows_from_item, RowsFromItem<X>),
(visit_table_factor, walk_table_factor, TableFactor<X>),
(visit_join, walk_join, Join<X>),
(visit_join_operator, walk_join_operator, JoinOperator<X>),
(visit_join_constraint, walk_join_constraint, JoinConstraint<X>),
(visit_order_by_expr, walk_order_by_expr, OrderByExpr<X>),
(visit_order_by_using, walk_order_by_using, OrderByUsing),
(visit_limit, walk_limit, Limit<X>),
(visit_statement, walk_statement, Statement<X>),
(visit_transaction_statement, walk_transaction_statement, TransactionStatement),
(visit_transaction_mode, walk_transaction_mode, TransactionMode),
(visit_data_type, walk_data_type, DataType<X>),
(visit_copy_statement, walk_copy_statement, CopyStatement<X>),
(visit_copy_source, walk_copy_source, CopySource<X>),
(visit_copy_target, walk_copy_target, CopyTarget),
(visit_copy_option, walk_copy_option, CopyOption),
(visit_copy_option_value, walk_copy_option_value, CopyOptionValue),
(visit_explain_statement, walk_explain_statement, ExplainStatement<X>),
(visit_explain_option, walk_explain_option, ExplainOption),
(visit_pragma_statement, walk_pragma_statement, PragmaStatement),
(visit_attach_statement, walk_attach_statement, AttachStatement<X>),
(visit_detach_statement, walk_detach_statement, DetachStatement),
(visit_vacuum_statement, walk_vacuum_statement, VacuumStatement<X>),
(visit_reindex_statement, walk_reindex_statement, ReindexStatement),
(visit_analyze_statement, walk_analyze_statement, AnalyzeStatement),
(visit_use_statement, walk_use_statement, UseStatement),
(visit_pivot, walk_pivot, Pivot<X>),
(visit_unpivot, walk_unpivot, Unpivot<X>),
(visit_pivot_expr, walk_pivot_expr, PivotExpr<X>),
(visit_pivot_column, walk_pivot_column, PivotColumn<X>),
(visit_unpivot_column, walk_unpivot_column, UnpivotColumn<X>),
(visit_window_spec, walk_window_spec, WindowSpec<X>),
(visit_window_definition, walk_window_definition, WindowDefinition<X>),
(visit_window_frame, walk_window_frame, WindowFrame<X>),
(visit_window_frame_bound, walk_window_frame_bound, WindowFrameBound<X>),
(visit_named_window, walk_named_window, NamedWindow<X>),
);
fn visit_ident(&mut self, node: &'ast Ident) {
if self.error.is_some() {
return;
}
self.check_span("visit_ident", node.span());
self.check_symbol(node.sym);
visit::walk_ident::<Self, X>(self, node);
}
fn visit_expr(&mut self, node: &'ast Expr<X>) {
if self.error.is_some() {
return;
}
self.check_span("visit_expr", node.span());
if let Expr::SessionVariable { name, .. } = node {
self.check_symbol(*name);
}
visit::walk_expr(self, node);
}
fn visit_function_arg(&mut self, node: &'ast FunctionArg<X>) {
if self.error.is_some() {
return;
}
self.check_span("visit_function_arg", node.span());
if let Some(name) = node.name {
self.check_symbol(name);
}
visit::walk_function_arg(self, node);
}
fn visit_named_operator_expr(&mut self, node: &'ast NamedOperatorExpr<X>) {
if self.error.is_some() {
return;
}
self.check_span("visit_named_operator_expr", node.span());
self.check_symbol(node.op);
visit::walk_named_operator_expr(self, node);
}
fn visit_parameter_kind(&mut self, node: &'ast ParameterKind) {
if self.error.is_some() {
return;
}
match node {
ParameterKind::Named { name, .. } => self.check_symbol(*name),
ParameterKind::PositionalLarge { digits } => self.check_symbol(*digits),
_ => {}
}
visit::walk_parameter_kind::<Self, X>(self, node);
}
}
#[cfg(all(test, feature = "serde-deserialize"))]
mod tests {
use super::*;
use crate::ast::generated::NodeIdWalk;
use crate::parser::{TestDialect, parse_with};
#[test]
fn span_walk_covers_every_node() {
let sql = "CREATE TABLE t (id INT PRIMARY KEY, name TEXT); \
INSERT INTO t (id, name) VALUES (1, 'a'); \
UPDATE t SET name = 'b' WHERE id = 1; \
DELETE FROM t WHERE id > 0; \
SELECT a + 1 AS n, count(DISTINCT b), \
CASE a WHEN 1 THEN b ELSE c END, \
avg(a) OVER (ORDER BY b ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW) \
FROM s.t AS t JOIN u ON t.a = u.b \
WHERE a IN (SELECT x FROM w) ORDER BY n";
let parsed =
parse_with(sql, crate::ParseConfig::new(TestDialect)).expect("rich corpus parses");
let mut walk = NodeIdWalk::default();
let mut validator = DocumentValidator {
source_len: parsed.source().len() as u32,
resolver: parsed.resolver(),
visited: 0,
error: None,
};
for statement in parsed.statements() {
walk.visit_statement(statement);
validator.visit_statement(statement);
}
assert!(validator.error.is_none(), "valid tree must not fault");
assert_eq!(
validator.visited,
walk.metas.len(),
"the check_spans! list drifted from the generated node set: the \
validator visited {} span-bearing nodes but NodeIdWalk recorded {}",
validator.visited,
walk.metas.len(),
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::interner::Interner;
use crate::parser::{ParseConfig, TestDialect, parse_rc_with, parse_with};
use std::rc::Rc;
fn parsed_for(source: &str) -> Parsed {
Parsed::new(
Arc::from(source),
Interner::new().freeze(),
Vec::new(),
StringLiteralSyntax::ANSI,
)
}
#[test]
fn line_col_maps_offsets_across_lines() {
let parsed = parsed_for("ab\ncde\nz");
assert_eq!(parsed.line_col(0), (0, 0));
assert_eq!(parsed.line_col(2), (0, 2));
assert_eq!(parsed.line_col(3), (1, 0));
assert_eq!(parsed.line_col(7), (2, 0));
assert_eq!(parsed.line_col(8), (2, 1));
}
#[test]
fn empty_input_has_a_single_origin_line() {
let parsed = parsed_for("");
assert_eq!(parsed.line_col(0), (0, 0));
}
#[test]
fn line_index_is_built_once_and_reused() {
let parsed = parsed_for("a\nb");
assert!(std::ptr::eq(parsed.line_index(), parsed.line_index()));
}
#[test]
fn span_line_col_spans_lines_and_skips_synthetic() {
let parsed = parsed_for("ab\ncde");
assert_eq!(
parsed.span_line_col(Span::new(1, 5)),
Some(((0, 1), (1, 2)))
);
assert_eq!(parsed.span_line_col(Span::SYNTHETIC), None);
}
#[test]
fn parsed_arc_root_stays_send_and_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Parsed>();
}
#[test]
fn public_roots_are_static_and_never_borrow_the_input() {
fn assert_static<T: 'static>() {}
assert_static::<Parsed>(); assert_static::<Parsed<Rc<str>>>();
}
#[test]
fn arc_root_parses_and_renders_canonical_sql() {
let parsed = parse_with("select 1, a", crate::ParseConfig::new(TestDialect))
.expect("Arc root parses");
assert_eq!(parsed.source(), "select 1, a");
assert_eq!(format!("{parsed}"), "SELECT 1, a");
}
#[test]
fn rc_root_parses_and_renders_canonical_sql() {
let parsed =
parse_rc_with("select 1, a", ParseConfig::new(TestDialect)).expect("Rc root parses");
assert_eq!(parsed.source(), "select 1, a");
assert_eq!(format!("{parsed}"), "SELECT 1, a");
}
#[test]
fn into_statements_yields_a_structure_only_send_vec() {
fn assert_send<T: Send>() {}
assert_send::<Vec<Statement>>();
let parsed =
parse_with("SELECT 1; SELECT 2", crate::ParseConfig::new(TestDialect)).expect("parses");
let statements = parsed.into_statements();
assert_eq!(statements.len(), 2);
assert!(matches!(statements[0], Statement::Query { .. }));
}
#[test]
fn display_joins_statements_and_renders_empty_input() {
let parsed =
parse_with("select a; select b", crate::ParseConfig::new(TestDialect)).expect("parses");
assert_eq!(format!("{parsed}"), "SELECT a; SELECT b");
assert_eq!(parsed.to_string(), "SELECT a; SELECT b");
let empty =
parse_with(" ; ; ", crate::ParseConfig::new(TestDialect)).expect("only separators");
assert!(empty.statements().is_empty());
assert_eq!(format!("{empty}"), "");
}
#[test]
fn default_parse_homes_no_trivia_on_the_root() {
let parsed = parse_with(
"SELECT /* c */ 1 -- note\n",
crate::ParseConfig::new(TestDialect),
)
.expect("parses");
assert!(
parsed.trivia().is_empty(),
"default parse captures no trivia"
);
assert!(parsed.trivia_in(Span::new(0, 100)).is_empty());
assert!(parsed.trivia_before(15).is_empty());
}
#[test]
fn trivia_config_homes_recoverable_trivia_on_the_root() {
use crate::tokenizer::TriviaKind::{BlockComment, LineComment, Whitespace};
let src = "SELECT /* c */ 1 -- note\n";
let parsed = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("parses with trivia");
let kinds: Vec<_> = parsed.trivia().iter().map(|r| r.kind()).collect();
assert_eq!(
kinds,
[
Whitespace, BlockComment, Whitespace, Whitespace, LineComment, Whitespace, ],
);
let block = parsed.trivia()[1];
assert_eq!(block.kind(), BlockComment);
let span = block.span();
assert_eq!(&src[span.start() as usize..span.end() as usize], "/* c */");
}
#[test]
fn root_trivia_queries_recover_runs_by_offset() {
use crate::tokenizer::TriviaKind::{BlockComment, LineComment};
let src = "SELECT /* c */ 1 -- note\n";
let parsed = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("parses with trivia");
let inner = parsed.trivia_in(Span::new(7, 14));
assert_eq!(inner.len(), 1);
assert_eq!(inner[0].kind(), BlockComment);
let leading = parsed.trivia_before(15);
assert_eq!(
leading.first().map(|r| r.span().start()),
Some(6),
"the leading run reaches back to just after SELECT",
);
assert_eq!(
leading.last().map(|r| r.span().end()),
Some(15),
"and ends exactly at the token",
);
let before_newline = parsed.trivia_before(24);
assert!(before_newline.iter().any(|r| r.kind() == LineComment));
}
#[test]
fn trivia_capture_does_not_change_statement_structure() {
let src = "SELECT a, /* x */ b FROM t -- trailing";
let plain = parse_with(src, crate::ParseConfig::new(TestDialect)).expect("plain");
let with_trivia = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("with trivia");
assert_eq!(plain.statements(), with_trivia.statements());
assert!(plain.trivia().is_empty());
assert!(!with_trivia.trivia().is_empty());
}
#[test]
fn capture_trivia_config_is_deterministic() {
use crate::parser::{ParseConfig, parse_with};
let src = "SELECT /* c */ a -- note\nFROM t";
let first = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("parses with trivia");
let second = parse_with(src, ParseConfig::new(TestDialect).capture_trivia(true))
.expect("parses with trivia");
assert!(!second.trivia().is_empty());
assert_eq!(first.trivia(), second.trivia());
let off = parse_with(src, ParseConfig::new(TestDialect)).expect("parses without trivia");
assert!(off.trivia().is_empty());
}
#[test]
fn clause_marks_capture_kinds_owners_and_offsets() {
use crate::ast::SetExpr;
use crate::parser::ClauseKw;
let src = "SELECT a FROM t WHERE b GROUP BY c HAVING d ORDER BY e LIMIT 1";
let parsed = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("parses with clause marks");
let Statement::Query { query, .. } = &parsed.statements()[0] else {
panic!("expected a query statement");
};
let query_id = query.meta.node_id;
let SetExpr::Select { select, .. } = &query.body else {
panic!("expected a SELECT body");
};
let select_id = select.meta.node_id;
let marks = parsed.clause_marks();
assert_eq!(
marks.iter().map(ClauseMark::kind).collect::<Vec<_>>(),
[
ClauseKw::From,
ClauseKw::Where,
ClauseKw::GroupBy,
ClauseKw::Having,
ClauseKw::OrderBy,
ClauseKw::Limit,
],
);
let at = |kw: &str| src.find(kw).expect("keyword present") as u32;
assert_eq!(marks[0].offset(), at("FROM"));
assert_eq!(marks[1].offset(), at("WHERE"));
assert_eq!(marks[2].offset(), at("GROUP BY"));
assert_eq!(marks[3].offset(), at("HAVING"));
assert_eq!(marks[4].offset(), at("ORDER BY"));
assert_eq!(marks[5].offset(), at("LIMIT"));
assert_eq!(marks[0].owner(), select_id);
assert_eq!(marks[1].owner(), select_id);
assert_eq!(marks[2].owner(), select_id);
assert_eq!(marks[3].owner(), select_id);
assert_eq!(marks[4].owner(), query_id);
assert_eq!(marks[5].owner(), query_id);
assert!(marks.windows(2).all(|w| w[0].offset() <= w[1].offset()));
let in_select = parsed.clause_marks_in(select.meta.span);
assert_eq!(
in_select.iter().map(ClauseMark::kind).collect::<Vec<_>>(),
[
ClauseKw::From,
ClauseKw::Where,
ClauseKw::GroupBy,
ClauseKw::Having,
],
);
assert_eq!(parsed.clause_marks_in(query.meta.span).len(), 6);
}
#[test]
fn default_parse_homes_no_clause_marks() {
let parsed = parse_with(
"SELECT a FROM t WHERE b GROUP BY c ORDER BY d LIMIT 1",
crate::ParseConfig::new(TestDialect),
)
.expect("parses");
assert!(
parsed.clause_marks().is_empty(),
"default parse captures no clause marks"
);
assert!(parsed.clause_marks_in(Span::new(0, 100)).is_empty());
}
#[test]
fn nested_query_clause_marks_are_owned_by_the_inner_nodes() {
use crate::ast::SetExpr;
let src = "SELECT a FROM t WHERE a IN (SELECT x FROM u WHERE y) ORDER BY a";
let parsed = parse_with(
src,
crate::ParseConfig::new(TestDialect).capture_trivia(true),
)
.expect("parses");
let Statement::Query { query, .. } = &parsed.statements()[0] else {
panic!("expected a query statement");
};
let outer_query_id = query.meta.node_id;
let SetExpr::Select { select, .. } = &query.body else {
panic!("expected a SELECT body");
};
let outer_select_id = select.meta.node_id;
let marks = parsed.clause_marks();
let owner_at = |offset: u32| {
marks
.iter()
.find(|m| m.offset() == offset)
.unwrap_or_else(|| panic!("a mark at offset {offset}"))
.owner()
};
let outer_where = src.find("WHERE").expect("outer WHERE") as u32;
let inner_where = src.rfind("WHERE").expect("inner WHERE") as u32;
let inner_from = src.rfind("FROM").expect("inner FROM") as u32;
let order_by = src.find("ORDER BY").expect("ORDER BY") as u32;
assert_eq!(owner_at(outer_where), outer_select_id);
assert_eq!(owner_at(order_by), outer_query_id);
let inner_owner = owner_at(inner_where);
assert_ne!(inner_owner, outer_select_id);
assert_ne!(inner_owner, outer_query_id);
assert_eq!(owner_at(inner_from), inner_owner);
}
}