use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use serde::Deserialize;
use crate::err::{TemplateParameterSpecError, TemplateSubstitutionError};
use crate::tree::ast::clause::{TablePathSegment, TableSegmentPart, TemplatedTablePath};
use crate::tree::ast::dataset_identifier::{
DatasetIdentifier, QualifiedDatasetIdentifier, UnqualifiedDatasetIdentifier,
};
use crate::tree::ast::expression::{Expression, ExpressionKind, IntervalLiteral, TruncUnit};
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::ParseWithErrors;
use crate::types::{Type, INTERVAL, STRING};
mod expand;
mod substitute;
pub use expand::expand_templated_table_path;
pub use substitute::substitute_query;
#[derive(Clone, Debug, PartialEq)]
pub enum TemplateParameterKind {
Primitive(Type),
IdentifierFragment(Vec<String>),
}
#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TemplateParameterType {
Boolean,
Int,
Double,
String,
Interval,
}
impl From<TemplateParameterType> for Type {
fn from(value: TemplateParameterType) -> Self {
match value {
TemplateParameterType::Boolean => Type::Boolean,
TemplateParameterType::Int => Type::Int,
TemplateParameterType::Double => Type::Double,
TemplateParameterType::String => Type::String,
TemplateParameterType::Interval => INTERVAL,
}
}
}
fn template_value_kind_name(value: &TemplateParameterValue) -> &'static str {
match value {
TemplateParameterValue::Boolean(_) => "boolean",
TemplateParameterValue::Int(_) => "int",
TemplateParameterValue::Double(_) => "double",
TemplateParameterValue::String(_) => "string",
TemplateParameterValue::Interval(_) => "interval",
}
}
pub fn interval_from_template_value(
name: &str,
value: &TemplateParameterValue,
) -> Result<IntervalLiteral, TemplateSubstitutionError> {
let TemplateParameterValue::Interval(text) = value else {
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.to_string(),
expected: "interval".to_string(),
got: template_value_kind_name(value).to_string(),
});
};
parse_template_interval_text(text).map_err(|detail| TemplateSubstitutionError::WrongValueKind {
name: name.to_string(),
expected: "valid interval literal".to_string(),
got: detail,
})
}
pub fn trunc_unit_from_interval_template_value(
name: &str,
value: &TemplateParameterValue,
) -> Result<(TruncUnit, u32), TemplateSubstitutionError> {
let text = match value {
TemplateParameterValue::Interval(s) => s.as_str(),
_ => {
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.to_string(),
expected: "interval".to_string(),
got: template_value_kind_name(value).to_string(),
});
}
};
let lit = interval_from_template_value(name, value)?;
lit.trunc_unit_and_multiplier()
.ok_or_else(|| TemplateSubstitutionError::WrongValueKind {
name: name.to_string(),
expected: "interval with a valid @ truncation unit".to_string(),
got: format!("'{text}'"),
})
}
#[derive(Clone, Debug, PartialEq)]
pub enum TemplateParameterValue {
Boolean(bool),
Int(i64),
Double(f64),
String(String),
Interval(String),
}
pub fn parse_template_interval_text(text: &str) -> Result<IntervalLiteral, String> {
let (expr, errors) = Expression::parse_with_errors(text);
if !errors.is_empty() {
let detail = errors
.iter()
.map(|e| e.to_string())
.collect::<Vec<_>>()
.join("; ");
return Err(if detail.is_empty() {
format!("'{text}' is not a valid interval literal")
} else {
detail
});
}
match expr.kind {
ExpressionKind::IntervalLiteral(lit) => Ok(lit),
_ => Err(format!("'{text}' is not a valid interval literal")),
}
}
#[derive(Debug, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum TemplateParameterSpecSerde {
Primitive {
#[serde(rename = "type")]
typ: TemplateParameterType,
},
IdentifierFragment {
values: Vec<String>,
},
}
impl TryFrom<TemplateParameterSpecSerde> for TemplateParameterKind {
type Error = TemplateParameterSpecError;
fn try_from(value: TemplateParameterSpecSerde) -> Result<Self, Self::Error> {
match value {
TemplateParameterSpecSerde::Primitive { typ } => {
Ok(TemplateParameterKind::Primitive(typ.into()))
}
TemplateParameterSpecSerde::IdentifierFragment { values } => {
if values.is_empty() {
return Err(TemplateParameterSpecError::EmptyIdentifierFragmentValues);
}
let mut seen = HashSet::new();
for value in &values {
if !seen.insert(value) {
return Err(TemplateParameterSpecError::DuplicateIdentifierFragmentValue);
}
if let Err(e) = SimpleIdentifier::parse(value.as_str()) {
return Err(TemplateParameterSpecError::InvalidIdentifierFragmentValue {
value: value.clone(),
detail: e.to_string(),
});
}
}
Ok(TemplateParameterKind::IdentifierFragment(values))
}
}
}
}
pub fn template_context_from_specs(
spec: Option<&Arc<HashMap<String, TemplateParameterKind>>>,
) -> (
Arc<HashMap<SimpleIdentifier, Type>>,
Arc<HashMap<SimpleIdentifier, Arc<[String]>>>,
) {
let Some(spec) = spec else {
return (Arc::new(HashMap::new()), Arc::new(HashMap::new()));
};
let mut types: HashMap<SimpleIdentifier, Type> = HashMap::new();
let mut identifier_fragments: HashMap<SimpleIdentifier, Arc<[String]>> = HashMap::new();
for (name, kind) in spec.iter() {
let key = SimpleIdentifier::new(name.as_str());
match kind {
TemplateParameterKind::Primitive(t) => {
types.insert(key.clone(), t.clone());
}
TemplateParameterKind::IdentifierFragment(vals) => {
types.insert(key.clone(), STRING);
identifier_fragments.insert(key, Arc::from(vals.clone()));
}
}
}
(Arc::new(types), Arc::new(identifier_fragments))
}
const MAX_TABLE_PATH_EXPANSIONS: usize = 100;
fn cartesian_template_values(domains: &[Arc<[String]>]) -> Vec<Vec<String>> {
let mut out: Vec<Vec<String>> = vec![vec![]];
for d in domains {
let mut next = Vec::with_capacity(out.len() * d.len());
for prefix in &out {
for v in d.iter() {
let mut row = prefix.clone();
row.push(v.clone());
next.push(row);
}
}
out = next;
}
out
}
fn ordered_params_in_templated_path(
path: &TemplatedTablePath,
) -> Result<Vec<SimpleIdentifier>, TemplateSubstitutionError> {
let mut param_order = Vec::new();
let all_segments = path.space.iter().chain(path.segments.iter());
for seg in all_segments {
for part in &seg.parts {
if let TableSegmentPart::Parameter(psi) = part {
let name = psi
.clone()
.valid()
.map_err(TemplateSubstitutionError::ParameterIdentifierParseError)?;
if !param_order.iter().any(|p| p == &name) {
param_order.push(name);
}
}
}
}
Ok(param_order)
}
fn materialize_table_path_segment(
seg: &TablePathSegment,
subst: &HashMap<SimpleIdentifier, String>,
) -> Result<SimpleIdentifier, TemplateSubstitutionError> {
let mut acc = String::new();
for p in &seg.parts {
match p {
TableSegmentPart::Text(psi) => {
let si = psi
.clone()
.valid()
.map_err(TemplateSubstitutionError::ParameterIdentifierParseError)?;
acc.push_str(si.as_str());
}
TableSegmentPart::Parameter(psi) => {
let name = psi
.clone()
.valid()
.map_err(TemplateSubstitutionError::ParameterIdentifierParseError)?;
let v = subst.get(&name).ok_or_else(|| {
TemplateSubstitutionError::MissingTemplateParameter(name.as_str().to_string())
})?;
acc.push_str(v);
}
}
}
if acc.is_empty() {
return Err(TemplateSubstitutionError::EmptyTablePathSegment);
}
Ok(SimpleIdentifier::new(acc))
}
fn materialize_templated_table_path(
path: &TemplatedTablePath,
subst: &HashMap<SimpleIdentifier, String>,
) -> Result<DatasetIdentifier, TemplateSubstitutionError> {
let space = match &path.space {
Some(seg) => Some(materialize_table_path_segment(seg, subst)?),
None => None,
};
let mut segments: Vec<SimpleIdentifier> = Vec::with_capacity(path.segments.len());
for seg in &path.segments {
segments.push(materialize_table_path_segment(seg, subst)?);
}
let table = segments
.pop()
.ok_or(TemplateSubstitutionError::EmptyTablePath)?;
let id = match space {
Some(space) => DatasetIdentifier::Qualified(QualifiedDatasetIdentifier {
space,
namespace: segments,
table,
}),
None => DatasetIdentifier::Unqualified(UnqualifiedDatasetIdentifier {
namespace: segments,
table,
}),
};
Ok(id)
}
pub fn validate_template_values(
specs: &HashMap<String, TemplateParameterKind>,
values: &HashMap<String, TemplateParameterValue>,
) -> Result<(), TemplateSubstitutionError> {
for name in specs.keys() {
if !values.contains_key(name) {
return Err(TemplateSubstitutionError::MissingTemplateParameter(
name.clone(),
));
}
}
for (name, value) in values {
let spec = specs
.get(name)
.ok_or_else(|| TemplateSubstitutionError::UndeclaredTemplateParameter(name.clone()))?;
match (spec, value) {
(TemplateParameterKind::Primitive(ty), val) => {
let matches = matches!(
(ty, val),
(Type::Boolean, TemplateParameterValue::Boolean(_))
| (Type::Int, TemplateParameterValue::Int(_))
| (Type::Double, TemplateParameterValue::Double(_))
| (Type::String, TemplateParameterValue::String(_))
| (Type::Interval, TemplateParameterValue::Interval(_))
);
if !matches {
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.clone(),
expected: ty.to_string(),
got: template_value_kind_name(val).to_string(),
});
}
if matches!(val, TemplateParameterValue::Interval(_)) {
interval_from_template_value(name, val)?;
}
}
(
TemplateParameterKind::IdentifierFragment(allowed),
TemplateParameterValue::String(s),
) => {
if !allowed.iter().any(|v| v == s) {
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.clone(),
expected: format!(
"one of [{}]",
allowed
.iter()
.map(|v| format!("'{v}'"))
.collect::<Vec<_>>()
.join(", ")
),
got: format!("'{s}'"),
});
}
}
(TemplateParameterKind::IdentifierFragment(_), val) => {
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.clone(),
expected: "string (identifier_fragment)".to_string(),
got: template_value_kind_name(val).to_string(),
});
}
}
}
Ok(())
}
#[cfg(test)]
mod validate_tests {
use std::collections::HashMap;
use crate::err::TemplateSubstitutionError;
use crate::tree::template::{
trunc_unit_from_interval_template_value, validate_template_values, TemplateParameterKind,
TemplateParameterValue,
};
use crate::types::{INTERVAL, STRING};
#[test]
fn rejects_missing_value_for_declared_spec() {
let specs = HashMap::from([(
"region".to_string(),
TemplateParameterKind::Primitive(STRING),
)]);
let values = HashMap::new();
let err = validate_template_values(&specs, &values).unwrap_err();
assert_eq!(
err,
TemplateSubstitutionError::MissingTemplateParameter("region".to_string())
);
}
#[test]
fn accepts_declared_spec_with_matching_value() {
let specs = HashMap::from([(
"region".to_string(),
TemplateParameterKind::Primitive(STRING),
)]);
let values = HashMap::from([(
"region".to_string(),
TemplateParameterValue::String("us".to_string()),
)]);
validate_template_values(&specs, &values).unwrap();
}
#[test]
fn accepts_interval_primitive_value() {
let specs = HashMap::from([(
"timeslice".to_string(),
TemplateParameterKind::Primitive(INTERVAL),
)]);
let values = HashMap::from([(
"timeslice".to_string(),
TemplateParameterValue::Interval("1h".to_string()),
)]);
validate_template_values(&specs, &values).unwrap();
}
#[test]
fn rejects_trunc_only_interval_suffix() {
let specs = HashMap::from([(
"timeslice".to_string(),
TemplateParameterKind::Primitive(INTERVAL),
)]);
let values = HashMap::from([(
"timeslice".to_string(),
TemplateParameterValue::Interval("h".to_string()),
)]);
let err = validate_template_values(&specs, &values).unwrap_err();
assert!(err.to_string().contains("valid interval literal"));
}
#[test]
fn rejects_zero_interval_trunc_multiplier() {
let err = trunc_unit_from_interval_template_value(
"timeslice",
&TemplateParameterValue::Interval("0h".to_string()),
)
.unwrap_err();
assert!(err.to_string().contains("valid @ truncation unit"));
}
}