use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use serde::Deserialize;
use crate::err::{TemplateParameterSpecError, TemplateSubstitutionError};
use crate::tree::ast::clause::{TableSegmentPart, TemplatedTablePath};
use crate::tree::ast::identifier::{CompoundIdentifier, Identifier, SimpleIdentifier};
use crate::types::{Type, 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,
}
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,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum TemplateParameterValue {
Boolean(bool),
Int(i64),
Double(f64),
String(String),
}
#[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 enums: 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);
enums.insert(key, Arc::from(vals.clone()));
}
}
}
(Arc::new(types), Arc::new(enums))
}
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();
for seg in &path.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_templated_table_path(
path: &TemplatedTablePath,
subst: &HashMap<SimpleIdentifier, String>,
) -> Result<Identifier, TemplateSubstitutionError> {
let mut parts: Vec<SimpleIdentifier> = Vec::with_capacity(path.segments.len());
for seg in &path.segments {
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);
}
parts.push(SimpleIdentifier::new(&acc));
}
match parts.as_slice() {
[] => Err(TemplateSubstitutionError::EmptyTablePath),
[one] => Ok(one.clone().into()),
[f, s, rest @ ..] => {
Ok(CompoundIdentifier::new(f.clone(), s.clone(), rest.to_vec()).into())
}
}
}
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(_))
);
if !matches {
let got = match val {
TemplateParameterValue::Boolean(_) => "boolean",
TemplateParameterValue::Int(_) => "int",
TemplateParameterValue::Double(_) => "double",
TemplateParameterValue::String(_) => "string",
};
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.clone(),
expected: ty.to_string(),
got: got.to_string(),
});
}
}
(
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) => {
let got = match val {
TemplateParameterValue::Boolean(_) => "boolean",
TemplateParameterValue::Int(_) => "int",
TemplateParameterValue::Double(_) => "double",
TemplateParameterValue::String(_) => "string",
};
return Err(TemplateSubstitutionError::WrongValueKind {
name: name.clone(),
expected: "string (identifier_fragment)".to_string(),
got: got.to_string(),
});
}
}
}
Ok(())
}
#[cfg(test)]
mod validate_tests {
use std::collections::HashMap;
use crate::err::TemplateSubstitutionError;
use crate::tree::template::{
validate_template_values, TemplateParameterKind, TemplateParameterValue,
};
use crate::types::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();
}
}