use std::borrow::Cow;
use mf2_model::{
Attributes, CatchAllKey, Declaration, Expression, FunctionRef, InputDeclaration, Key, Literal,
LocalDeclaration, Message, Pattern, PatternMessage, PatternPart, SelectMessage,
VariableExpression, VariableRef, Variant,
};
pub(super) const VARIANT_LIMIT: usize = 256;
pub(super) type Flat = Vec<Part>;
#[derive(Clone, Debug)]
pub(super) enum Part {
Text(String),
Expr(Expression<'static>, Option<String>),
Select(Box<Select>),
}
#[derive(Clone, Debug)]
pub(super) struct Select {
pub(super) var: String,
pub(super) function: FunctionRef<'static>,
pub(super) arms: Vec<Arm>,
}
#[derive(Clone, Debug)]
pub(super) struct Arm {
pub(super) key: Option<String>,
pub(super) pattern: Flat,
}
struct Column {
var: String,
function: FunctionRef<'static>,
keys: Vec<Option<String>>,
}
impl Column {
fn is(&self, s: &Select) -> bool {
self.var == s.var
&& self.function == s.function
&& self.keys.len() == s.arms.len()
&& self.keys.iter().zip(&s.arms).all(|(k, a)| *k == a.key)
}
}
#[derive(Clone, Default)]
struct Row {
chosen: Vec<Option<usize>>,
pattern: Pattern<'static>,
}
pub(super) struct TooMany;
pub(super) fn message(flat: &Flat) -> Result<Message<'static>, TooMany> {
let mut columns: Vec<Column> = Vec::new();
collect(flat, &mut columns);
if columns.is_empty() {
let mut pattern = Pattern::new();
for part in flat {
push(&mut pattern, part);
}
return Ok(Message::Pattern(PatternMessage {
declarations: Vec::new(),
pattern,
}));
}
let rows = extend(
vec![Row {
chosen: vec![None; columns.len()],
pattern: Pattern::new(),
}],
flat,
&columns,
)?;
let (declarations, selectors) = declare(&columns, flat);
let variants = rows
.into_iter()
.map(|row| Variant {
keys: row
.chosen
.iter()
.zip(&columns)
.map(|(chosen, column)| {
match chosen.and_then(|i| column.keys.get(i).cloned().flatten()) {
Some(k) => Key::Literal(Literal {
value: Cow::Owned(k),
}),
None => Key::CatchAll(CatchAllKey::default()),
}
})
.collect(),
value: row.pattern,
})
.collect();
Ok(Message::Select(SelectMessage {
declarations,
selectors,
variants,
}))
}
fn collect(flat: &Flat, columns: &mut Vec<Column>) {
for part in flat {
if let Part::Select(s) = part {
if !columns.iter().any(|c| c.is(s)) {
columns.push(Column {
var: s.var.clone(),
function: s.function.clone(),
keys: s.arms.iter().map(|a| a.key.clone()).collect(),
});
}
for arm in &s.arms {
collect(&arm.pattern, columns);
}
}
}
}
fn extend(rows: Vec<Row>, flat: &Flat, columns: &[Column]) -> Result<Vec<Row>, TooMany> {
let mut rows = rows;
for part in flat {
match part {
Part::Select(s) => {
let col = columns.iter().position(|c| c.is(s)).unwrap_or_default();
let mut next = Vec::new();
for row in rows {
match row.chosen.get(col).copied().flatten() {
Some(i) => {
let arm = &s.arms[i];
next.extend(extend(vec![row], &arm.pattern, columns)?);
}
None => {
for (i, arm) in s.arms.iter().enumerate() {
let mut branch = row.clone();
branch.chosen[col] = Some(i);
next.extend(extend(vec![branch], &arm.pattern, columns)?);
if next.len() > VARIANT_LIMIT {
return Err(TooMany);
}
}
}
}
}
rows = next;
}
other => {
for row in &mut rows {
push(&mut row.pattern, other);
}
}
}
}
Ok(rows)
}
fn push(pattern: &mut Pattern<'static>, part: &Part) {
match part {
Part::Text(t) => pattern.push(PatternPart::Text(Cow::Owned(t.clone()))),
Part::Expr(e, _) => pattern.push(PatternPart::Expression(e.clone())),
Part::Select(_) => {}
}
}
fn declare(
columns: &[Column],
flat: &Flat,
) -> (Vec<Declaration<'static>>, Vec<VariableRef<'static>>) {
let mut names = std::collections::BTreeSet::new();
names_in(flat, &mut names);
let mut inputs: Vec<(String, FunctionRef<'static>)> = Vec::new();
let mut locals: Vec<(String, String, FunctionRef<'static>)> = Vec::new();
let mut selectors = Vec::new();
for column in columns {
let name = if let Some((_, f)) = inputs.iter().find(|(v, _)| *v == column.var) {
if *f == column.function {
column.var.clone()
} else if let Some((n, _, _)) = locals
.iter()
.find(|(_, v, f)| *v == column.var && *f == column.function)
{
n.clone()
} else {
let mut k = 2;
let local = loop {
let candidate = format!("{}-{k}", column.var);
if !names.contains(&candidate) {
break candidate;
}
k += 1;
};
names.insert(local.clone());
locals.push((local.clone(), column.var.clone(), column.function.clone()));
local
}
} else {
inputs.push((column.var.clone(), column.function.clone()));
column.var.clone()
};
selectors.push(VariableRef {
name: Cow::Owned(name),
});
}
let mut declarations: Vec<Declaration<'static>> = inputs
.into_iter()
.map(|(var, function)| {
Declaration::Input(InputDeclaration {
name: Cow::Owned(var.clone()),
value: VariableExpression {
arg: VariableRef {
name: Cow::Owned(var),
},
function: Some(function),
attributes: Attributes::default(),
},
})
})
.collect();
declarations.extend(locals.into_iter().map(|(name, var, function)| {
Declaration::Local(LocalDeclaration {
name: Cow::Owned(name),
value: Expression::Variable(VariableExpression {
arg: VariableRef {
name: Cow::Owned(var),
},
function: Some(function),
attributes: Attributes::default(),
}),
})
}));
(declarations, selectors)
}
fn names_in(flat: &Flat, names: &mut std::collections::BTreeSet<String>) {
for part in flat {
match part {
Part::Expr(Expression::Variable(v), _) => {
names.insert(v.arg.name.to_string());
}
Part::Text(_) | Part::Expr(..) => {}
Part::Select(s) => {
names.insert(s.var.clone());
for arm in &s.arms {
names_in(&arm.pattern, names);
}
}
}
}
}