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