use std::collections::HashSet;
use std::fmt::{self, Display, Formatter};
use std::ops::RangeInclusive;
use std::rc::Rc;
use std::sync::Arc;
use derive_more::From;
use ordermap::OrderSet;
use super::command::CommandKind;
use super::expression::Expression;
use super::identifier::{Identifier, ParsedIdentifier, ParsedSimpleIdentifier};
use super::node::{Span, Spannable};
use super::pipeline::Pipeline;
use crate::antlr::hamelinparser::{
ExpressionQueryContextAttrs, QueryContextAll, QueryEOFContextAttrs,
StandaloneQueryContextAttrs, WithQueryContextAttrs,
};
use crate::antlr::interval;
use crate::antlr::resilient_parse_query;
use crate::err::{TranslationError, TranslationErrors};
use crate::tree::ast::context::{FromCst, ParseContext};
use crate::tree::ast::{ParseWithContext, ParseWithErrors, TypeCheck};
use crate::tree::typed_ast::context::StatementTranslationContext;
use crate::tree::typed_ast::query::TypedStatement;
use crate::write_utils::pad;
#[derive(Debug, Clone)]
pub struct Query {
pub span: Span,
pub kind: QueryKind,
}
impl Query {
pub fn valid_ref(&self) -> Result<&ValidQuery, Arc<TranslationError>> {
match &self.kind {
QueryKind::Valid(valid_query) => Ok(valid_query),
QueryKind::Error(translation_error) => Err(translation_error.clone()),
}
}
pub fn datasets(&self) -> Result<Vec<Identifier>, Arc<TranslationError>> {
let valid = self.valid_ref()?;
let mut cte_names: HashSet<&Identifier> = HashSet::new();
for wc in &valid.with_clauses {
cte_names.insert(wc.name.valid_ref()?);
}
let mut datasets = OrderSet::new();
let mut pipelines: Vec<&Pipeline> = vec![&valid.main_pipeline];
for wc in &valid.with_clauses {
pipelines.push(&wc.pipeline);
}
for pipeline in pipelines {
for cmd in &pipeline.commands {
match &cmd.kind {
CommandKind::From(from_cmd) => {
for clause in &from_cmd.clauses {
let table_ref = clause.table_reference();
let identifier = table_ref.identifier.valid_ref()?;
if !cte_names.contains(identifier) {
datasets.insert(identifier.clone());
}
}
}
CommandKind::Union(union_cmd) => {
for clause in &union_cmd.clauses {
let table_ref = clause.table_reference();
let identifier = table_ref.identifier.valid_ref()?;
if !cte_names.contains(identifier) {
datasets.insert(identifier.clone());
}
}
}
CommandKind::Join(join_cmd) => {
let table_ref = join_cmd.other.table_reference();
let identifier = table_ref.identifier.valid_ref()?;
if !cte_names.contains(identifier) {
datasets.insert(identifier.clone());
}
}
CommandKind::Lookup(lookup_cmd) => {
let table_ref = lookup_cmd.other.table_reference();
let identifier = table_ref.identifier.valid_ref()?;
if !cte_names.contains(identifier) {
datasets.insert(identifier.clone());
}
}
CommandKind::Append(append_cmd) => {
let identifier = append_cmd.table.identifier.valid_ref()?;
if !cte_names.contains(identifier) {
datasets.insert(identifier.clone());
}
}
CommandKind::Match(match_cmd) => {
for pattern in &match_cmd.pattern {
pattern.extract_datasets(&cte_names, &mut datasets)?;
}
}
_ => {}
}
}
}
Ok(datasets.into_iter().collect())
}
}
impl PartialEq for Query {
fn eq(&self, other: &Self) -> bool {
self.kind == other.kind
}
}
impl ParseWithContext for Query {
fn parse_with_context(input: String, mut ctx: ParseContext) -> (Self, TranslationErrors) {
match resilient_parse_query(input) {
Ok((csteof, parse_errors)) => match csteof.query() {
Some(cst) => {
let ast = Self::from_cst_with_context(cst, &mut ctx);
let mut errors = ctx.take_errors();
errors.extend(parse_errors);
(ast, errors)
}
None => {
let err = ctx.error("parse did not resolve a query tree").emit();
let mut errors = ctx.take_errors();
errors.extend(parse_errors);
(
Self {
span: Span::NONE,
kind: QueryKind::Error(err),
},
errors,
)
}
},
Err(e) => {
let error = ctx
.error("fatal error initializing parser")
.with_source_boxed(e.into())
.emit();
(
Self {
span: Span::NONE,
kind: QueryKind::Error(error.clone()),
},
(*error).clone().single(),
)
}
}
}
}
impl ParseWithErrors for Query {
fn parse_with_errors(input: impl Into<String>) -> (Self, TranslationErrors) {
Self::parse_with_context(input.into(), ParseContext::new())
}
}
impl Spannable for Query {
fn span(&self) -> Option<RangeInclusive<usize>> {
self.span.to_range()
}
}
#[derive(Debug, Clone, PartialEq, From)]
pub enum QueryKind {
Valid(ValidQuery),
Error(Arc<TranslationError>),
}
#[derive(Debug, Clone)]
pub struct ValidQuery {
pub span: Span,
pub with_clauses: Vec<WithClause>,
pub main_pipeline: Arc<Pipeline>,
}
impl PartialEq for ValidQuery {
fn eq(&self, other: &Self) -> bool {
self.with_clauses == other.with_clauses && self.main_pipeline == other.main_pipeline
}
}
impl Spannable for ValidQuery {
fn span(&self) -> Option<RangeInclusive<usize>> {
self.span.to_range()
}
}
#[derive(Debug, Clone)]
pub struct WithClause {
pub span: Span,
pub name: ParsedIdentifier,
pub pipeline: Arc<Pipeline>,
}
impl PartialEq for WithClause {
fn eq(&self, other: &Self) -> bool {
self.name == other.name && self.pipeline == other.pipeline
}
}
impl Spannable for WithClause {
fn span(&self) -> Option<RangeInclusive<usize>> {
self.span.to_range()
}
}
impl Query {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
self.kind.fmt_indented(f, indentation)
}
}
impl Display for Query {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
self.fmt_indented(f, 0)
}
}
impl QueryKind {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
match self {
QueryKind::Valid(v) => v.fmt_indented(f, indentation),
QueryKind::Error(_) => write!(f, "?!"),
}
}
}
impl ValidQuery {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
for with_clause in &self.with_clauses {
with_clause.fmt_indented(f, indentation)?;
writeln!(f)?;
writeln!(f)?;
pad(f, indentation)?;
}
self.main_pipeline.fmt_indented(f, indentation)
}
}
impl Display for ValidQuery {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
self.fmt_indented(f, 0)
}
}
impl WithClause {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: usize) -> fmt::Result {
write!(f, "WITH {} = ", self.name)?;
self.pipeline.fmt_indented(f, indentation)
}
}
impl Display for WithClause {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
self.fmt_indented(f, 0)
}
}
impl FromCst<Rc<QueryContextAll<'static>>> for Query {
fn from_cst_with_context(cst: Rc<QueryContextAll<'static>>, ctx: &mut ParseContext) -> Self {
let query_span = Some(interval(cst.as_ref()));
let kind = match cst.as_ref() {
QueryContextAll::WithQueryContext(wctx) => {
let simple_idents = wctx.simpleIdentifier_all();
let pipelines = wctx.pipeline_all();
if pipelines.is_empty() {
QueryKind::Error(ctx.error("WITH query missing pipelines").at(wctx).emit())
} else {
let with_clauses: Vec<WithClause> = simple_idents
.into_iter()
.zip(pipelines.iter())
.map(|(ident_cst, pipeline_cst)| {
let ident_span = interval(ident_cst.as_ref());
let pipeline_span = interval(pipeline_cst.as_ref());
let clause_span = *ident_span.start()..=*pipeline_span.end();
WithClause {
span: clause_span.into(),
name: ParsedSimpleIdentifier::from_cst_with_context(
ident_cst.as_ref(),
ctx,
)
.into(),
pipeline: Arc::new(Pipeline::from_cst_with_context(
pipeline_cst.clone(),
ctx,
)),
}
})
.collect();
let main_pipeline = match pipelines.last() {
Some(pipeline_cst) => {
Arc::new(Pipeline::from_cst_with_context(pipeline_cst.clone(), ctx))
}
None => {
return Query {
span: query_span.into(),
kind: QueryKind::Error(
ctx.error("WITH query missing main pipeline")
.at(wctx)
.emit(),
),
};
}
};
QueryKind::Valid(ValidQuery {
span: query_span.clone().into(),
with_clauses,
main_pipeline,
})
}
}
QueryContextAll::StandaloneQueryContext(sctx) => {
match sctx.pipeline() {
Some(pipeline_cst) => QueryKind::Valid(ValidQuery {
span: query_span.clone().into(),
with_clauses: vec![],
main_pipeline: Arc::new(Pipeline::from_cst_with_context(
pipeline_cst.clone(),
ctx,
)),
}),
None => QueryKind::Error(ctx.error("Query missing its main pipeline").emit()),
}
}
QueryContextAll::ExpressionQueryContext(ectx) => {
match ectx.expression() {
Some(expr_cst) => {
let expr = Expression::from_cst_with_context(expr_cst.clone(), ctx);
let pipeline = Pipeline::from(expr);
QueryKind::Valid(ValidQuery {
span: query_span.clone().into(),
with_clauses: vec![],
main_pipeline: Arc::new(pipeline),
})
}
None => QueryKind::Error(ctx.error("Expression did not parse").emit()),
}
}
QueryContextAll::Error(_) => ctx
.error("query parse error")
.at(cst.as_ref())
.emit()
.into(),
};
Query {
span: query_span.into(),
kind,
}
}
}
impl TypeCheck for Query {
type Output = TypedStatement;
fn type_check_with_context(
ast: Arc<Self>,
ctx: &mut StatementTranslationContext,
) -> Self::Output {
TypedStatement::from_ast_with_context(ast, ctx)
}
}