hamelin_lib 0.21.4

Core library for Hamelin query language
Documentation
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, ParsedIdentifier, SimpleIdentifier};
use crate::tree::ast::query::{DefBody, DefStatement, Query, ValidQuery};
use crate::tree::transform::{map_table_references_in_valid_query, MapTableReferences};

use super::def_name_registry::pick_available_def_name;
use super::CombineQueriesError;

/// Rename `DEF`s whose names collide with `used_def_names`, updating table and field references.
pub fn rename_conflicting_defs(
    query: &Query,
    used_def_names: &HashSet<Identifier>,
) -> Result<Query, CombineQueriesError> {
    let valid = query
        .valid_ref()
        .map_err(|_| CombineQueriesError::QueryHasParseErrors)?;

    let def_renames = plan_def_renames(valid, used_def_names)?;
    if def_renames.is_empty() {
        return Ok(query.clone());
    }

    apply_def_renames(query, valid, &def_renames)
}

fn plan_def_renames(
    valid: &ValidQuery,
    used_def_names: &HashSet<Identifier>,
) -> Result<HashMap<String, String>, CombineQueriesError> {
    let local_def_names: HashSet<Identifier> = valid
        .defs
        .iter()
        .filter_map(|d| d.name.valid_ref().ok().cloned())
        .collect();

    let mut def_renames: HashMap<String, String> = HashMap::new();
    let mut rename_targets: HashSet<Identifier> = HashSet::new();
    for d in &valid.defs {
        let Ok(id) = d.name.valid_ref() else {
            continue;
        };
        if used_def_names.contains(id) {
            if matches!(d.body, DefBody::Expression(_)) {
                return Err(CombineQueriesError::Message(format!(
                    "scalar DEF name collision: {id}"
                )));
            }

            let mut blocked = used_def_names.clone();
            blocked.extend(local_def_names.iter().cloned());
            blocked.extend(rename_targets.iter().cloned());

            let allocated = pick_available_def_name(&id.to_string(), &blocked)
                .map_err(CombineQueriesError::InvalidCteName)?;
            rename_targets.insert(allocated.clone().into());
            def_renames.insert(id.to_string(), allocated.as_str().to_string());
        }
    }

    Ok(def_renames)
}

fn apply_def_renames(
    query: &Query,
    valid: &ValidQuery,
    def_renames: &HashMap<String, String>,
) -> Result<Query, CombineQueriesError> {
    let mut defs = Vec::with_capacity(valid.defs.len());
    for d in &valid.defs {
        defs.push(DefStatement {
            span: d.span,
            name: renamed_def_name(&d.name, def_renames)?,
            body: d.body.clone(),
        });
    }

    let interim = ValidQuery {
        span: valid.span,
        defs,
        main_pipeline: valid.main_pipeline.clone(),
    };

    let table_mapper = DefRenameTableMapper { def_renames };
    map_table_references_in_valid_query(&interim, query.span.clone(), &table_mapper)
}

fn renamed_def_name(
    name: &ParsedIdentifier,
    def_renames: &HashMap<String, String>,
) -> Result<ParsedIdentifier, CombineQueriesError> {
    let Ok(id) = name.valid_ref() else {
        return Ok(name.clone());
    };
    if let Some(new_name) = def_renames.get(&id.to_string()) {
        let table = SimpleIdentifier::parse(new_name)
            .map_err(|e| CombineQueriesError::InvalidCteName(e.to_string()))?;
        return Ok(ParsedIdentifier::from(table));
    }
    Ok(name.clone())
}

struct DefRenameTableMapper<'a> {
    def_renames: &'a HashMap<String, String>,
}

impl MapTableReferences for DefRenameTableMapper<'_> {
    type Error = CombineQueriesError;

    fn map_table_reference(&self, tr: &TableReference) -> Result<TableReference, Self::Error> {
        match &tr.target {
            TableReferenceTarget::Static(id) => {
                let dataset = id
                    .valid_ref()
                    .map_err(|_| CombineQueriesError::QueryHasParseErrors)?;
                if let DatasetIdentifier::Unqualified(unqual) = dataset {
                    if unqual.namespace.is_empty() {
                        if let Some(new_name) = self.def_renames.get(unqual.table.as_str()) {
                            let table = SimpleIdentifier::parse(new_name)
                                .map_err(|e| CombineQueriesError::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()),
        }
    }
}

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

    #[test]
    fn renames_conflicting_internal_def() {
        let (q, errs) = Query::parse_with_errors(
            "DEF network_events = FROM events:sysmon | LIMIT 1;\nFROM network_events | SET x = 1 AS int",
        );
        assert!(errs.is_empty(), "{errs:?}");

        let mut registry = DefNameRegistry::new();
        let (seed, errs) = Query::parse_with_errors(
            "DEF network_events = FROM events:a | LIMIT 1;\nFROM network_events",
        );
        assert!(errs.is_empty(), "{errs:?}");
        registry.extend_from_query(&seed);

        let out = rename_conflicting_defs(&q, registry.used_names()).unwrap();
        let rendered = out.to_string();
        assert!(rendered.contains("DEF network_events_2 = FROM events:sysmon"));
        assert!(rendered.contains("FROM network_events_2"));
        assert!(!rendered.contains("FROM network_events\n"));
    }

    #[test]
    fn does_not_rename_sibling_def_matching_allocated_suffix() {
        let (q, errs) = Query::parse_with_errors(
            "DEF foo = FROM events:a | LIMIT 1;\nDEF foo_2 = FROM events:b | LIMIT 1;\nFROM foo_2 | SET x = 1 AS int",
        );
        assert!(errs.is_empty(), "{errs:?}");

        let mut registry = DefNameRegistry::new();
        let (seed, errs) = Query::parse_with_errors("DEF foo = FROM events:a | LIMIT 1;\nFROM foo");
        assert!(errs.is_empty(), "{errs:?}");
        registry.extend_from_query(&seed);

        let out = rename_conflicting_defs(&q, registry.used_names()).unwrap();
        let rendered = out.to_string();
        assert!(rendered.contains("DEF foo_3 = FROM events:a"));
        assert!(rendered.contains("DEF foo_2 = FROM events:b"));
        assert!(rendered.contains("FROM foo_2"));
    }

    #[test]
    fn rejects_conflicting_scalar_def() {
        let (q, errs) = Query::parse_with_errors(
            "DEF my_limit = 20 AS int;\nFROM events:source | LIMIT my_limit",
        );
        assert!(errs.is_empty(), "{errs:?}");

        let mut registry = DefNameRegistry::new();
        let (seed, errs) = Query::parse_with_errors("DEF my_limit = 10 AS int;\nFROM events:base");
        assert!(errs.is_empty(), "{errs:?}");
        registry.extend_from_query(&seed);

        let err = rename_conflicting_defs(&q, registry.used_names()).unwrap_err();
        assert!(
            matches!(err, CombineQueriesError::Message(msg) if msg.contains("scalar DEF name collision"))
        );
    }
}