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;
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"))
);
}
}