hamelin_lib 0.21.7

Core library for Hamelin query language
Documentation
//! AST-aware rewrite of catalog table references to unqualified CTE names.

use std::collections::{HashMap, HashSet};

use crate::tree::ast::clause::{TableReference, TableReferenceTarget};
use crate::tree::ast::dataset_identifier::{DatasetIdentifier, UnqualifiedDatasetIdentifier};
use crate::tree::ast::identifier::{Identifier, SimpleIdentifier};
use crate::tree::ast::query::{DefBody, Query, QueryKind, ValidQuery};

use super::catalog_key_matching::find_catalog_replacement;
use super::map_table_references::{map_table_references_in_valid_query, MapTableReferences};

#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum RewriteTableReferencesError {
    #[error("query has parse errors")]
    QueryHasParseErrors,
    #[error("invalid CTE name in replacement map: {0}")]
    InvalidCteName(String),
    #[error("ambiguous catalog replacement for {0}")]
    AmbiguousReplacement(String),
}

impl From<super::catalog_key_matching::CatalogKeyMatchingError> for RewriteTableReferencesError {
    fn from(value: super::catalog_key_matching::CatalogKeyMatchingError) -> Self {
        match value {
            super::catalog_key_matching::CatalogKeyMatchingError::AmbiguousReplacement(msg) => {
                Self::AmbiguousReplacement(msg)
            }
        }
    }
}

/// Catalog keys in `replacements` rewrite to unqualified CTE names.
///
/// Every catalog reference rewrites, including an inline's own catalog key when it appears in
/// the query body (for example recursive views that read `FROM events:my_view`). Upstream
/// queries are expected to read from underlying datasets, not their own catalog entry, except for
/// that intentional self-reference case.
pub fn rewrite_table_references(
    query: &Query,
    replacements: &HashMap<String, String>,
    default_space: Option<&SimpleIdentifier>,
) -> Result<Query, RewriteTableReferencesError> {
    let valid = match &query.kind {
        QueryKind::Valid(v) => v,
        QueryKind::Error(_) => return Err(RewriteTableReferencesError::QueryHasParseErrors),
    };

    let mapper = CatalogToCteMapper {
        replacements,
        cte_names: tabular_def_names(valid),
        default_space,
    };

    map_table_references_in_valid_query(valid, query.span.clone(), &mapper)
}

struct CatalogToCteMapper<'a> {
    replacements: &'a HashMap<String, String>,
    cte_names: HashSet<Identifier>,
    default_space: Option<&'a SimpleIdentifier>,
}

impl MapTableReferences for CatalogToCteMapper<'_> {
    type Error = RewriteTableReferencesError;

    fn map_table_reference(&self, tr: &TableReference) -> Result<TableReference, Self::Error> {
        match &tr.target {
            TableReferenceTarget::Static(_) => {
                let dataset = tr
                    .static_valid_dataset_ref()
                    .map_err(|_| RewriteTableReferencesError::QueryHasParseErrors)?;
                if dataset.is_cte_candidate(&self.cte_names) {
                    return Ok(tr.clone());
                }
                if let Some(cte_name) =
                    find_catalog_replacement(dataset, self.replacements, self.default_space)?
                {
                    let table = SimpleIdentifier::parse(&cte_name)
                        .map_err(|e| RewriteTableReferencesError::InvalidCteName(e.to_string()))?;
                    let new_id: DatasetIdentifier = UnqualifiedDatasetIdentifier {
                        namespace: vec![],
                        table,
                    }
                    .into();
                    return Ok(TableReference {
                        span: tr.span,
                        target: new_id.into(),
                    });
                }
                Ok(tr.clone())
            }
            TableReferenceTarget::Templated(_) => Ok(tr.clone()),
        }
    }
}

fn tabular_def_names(valid: &ValidQuery) -> HashSet<Identifier> {
    let mut cte_names = HashSet::new();
    for d in &valid.defs {
        if matches!(d.body, DefBody::Pipeline(_)) {
            if let Ok(id) = d.name.valid_ref() {
                cte_names.insert(id.clone());
            }
        }
    }
    cte_names
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::tree::ast::ParseWithErrors;

    #[test]
    fn rewrites_qualified_catalog_ref_to_cte() {
        let (q, errs) = Query::parse_with_errors(
            "FROM events:okta_auth\n| SET user.name = actor.name AS string",
        );
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert(
            "events:okta_auth".to_string(),
            "events_okta_auth".to_string(),
        );
        let out = rewrite_table_references(&q, &replacements, None).unwrap();
        let rendered = out.to_string();
        assert!(rendered.contains("FROM events_okta_auth"));
        assert!(!rendered.contains("FROM events:okta_auth"));
    }

    #[test]
    fn rewrites_unqualified_ref_with_default_space() {
        let (q, errs) = Query::parse_with_errors("FROM okta_auth | SET x = 1 AS int");
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert(
            "events:okta_auth".to_string(),
            "events_okta_auth".to_string(),
        );
        let default_space = SimpleIdentifier::new("events");
        let out = rewrite_table_references(&q, &replacements, Some(&default_space)).unwrap();
        assert!(out.to_string().contains("FROM events_okta_auth"));
    }

    #[test]
    fn does_not_rewrite_local_def_name() {
        let (q, errs) = Query::parse_with_errors(
            "DEF local_cte = FROM events:sysmon | LIMIT 1;\nFROM local_cte | SET x = 1 AS int",
        );
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert("events:sysmon".to_string(), "inlined_sysmon".to_string());
        let out = rewrite_table_references(&q, &replacements, None).unwrap();
        let rendered = out.to_string();
        assert!(rendered.contains("FROM local_cte"));
        assert!(rendered.contains("FROM inlined_sysmon"));
    }

    #[test]
    fn rewrites_ref_inside_def_body() {
        let (q, errs) = Query::parse_with_errors(
            "DEF inner = FROM events:upstream | LIMIT 1;\nFROM inner | SET x = 1 AS int",
        );
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert(
            "events:upstream".to_string(),
            "inlined_upstream".to_string(),
        );
        let out = rewrite_table_references(&q, &replacements, None).unwrap();
        assert!(out.to_string().contains("FROM inlined_upstream"));
    }

    #[test]
    fn resolves_catalog_refs_without_alias_map_collision() {
        let (q, errs) = Query::parse_with_errors("FROM events:foo | SET x = 1 AS int");
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert("events:foo".to_string(), "events_foo".to_string());
        replacements.insert("signals:foo".to_string(), "signals_foo".to_string());

        let out = rewrite_table_references(&q, &replacements, None).unwrap();
        assert!(out.to_string().contains("FROM events_foo"));
    }

    #[test]
    fn does_not_rewrite_bare_table_without_default_space() {
        let (q, errs) = Query::parse_with_errors("FROM foo | SET x = 1 AS int");
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert("events:foo".to_string(), "events_foo".to_string());

        let out = rewrite_table_references(&q, &replacements, None).unwrap();
        assert!(out.to_string().contains("FROM foo"));
        assert!(!out.to_string().contains("FROM events_foo"));
    }

    #[test]
    fn resolves_unqualified_ref_with_default_space_when_table_name_collides() {
        let (q, errs) = Query::parse_with_errors("FROM foo | SET x = 1 AS int");
        assert!(errs.is_empty(), "{errs:?}");
        let mut replacements = HashMap::new();
        replacements.insert("events:foo".to_string(), "events_foo".to_string());
        replacements.insert("signals:foo".to_string(), "signals_foo".to_string());
        let default_space = SimpleIdentifier::new("events");

        let out = rewrite_table_references(&q, &replacements, Some(&default_space)).unwrap();
        assert!(out.to_string().contains("FROM events_foo"));
    }
}