Skip to main content

xlog_logic/
list_normalize.rs

1//! Finite list normalization for the language-completeness surface.
2
3use std::collections::{HashMap, HashSet};
4
5use xlog_core::{Result, ScalarType, XlogError};
6
7use crate::ast::{
8    Atom, BodyLiteral, Comparison, DomainDecl, PredColumn, PredDecl, Program, Query, Rule, Term,
9    TypeRef,
10};
11
12const LIST_ID_TYPE: ScalarType = ScalarType::U64;
13
14/// Normalize finite list syntax into scalar list identifiers plus helper relations.
15///
16/// The normalized program uses ordinary XLOG relations:
17///
18/// - `__xlog_list_len(list_id, len)`
19/// - `__xlog_list_item_<type>(list_id, index, value)`
20/// - `__xlog_list_cons_<type>(list_id, head, tail_id)`
21/// - selected built-in helper relations such as `append`, `sort`, `msort`, and `list_to_set`
22///
23/// This keeps accepted list programs on the existing relational lowering/runtime path.
24pub fn normalize_list_builtins(program: &Program) -> Result<Program> {
25    normalize_list_builtins_owned(program.clone())
26}
27
28/// [`normalize_list_builtins`] taking the program by value.
29///
30/// Avoids cloning the whole AST: facts without list literals are moved through
31/// untouched (the pass is the identity on them, see
32/// `ListNormalizer::fact_is_identity`), everything else is rebuilt as before.
33pub fn normalize_list_builtins_owned(program: Program) -> Result<Program> {
34    let mut normalizer = ListNormalizer::new(&program)?;
35    normalizer.normalize_program(program)
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
39enum ListOp {
40    Append,
41    Sort,
42    MSort,
43    ListToSet,
44}
45
46#[derive(Debug, Clone, PartialEq, Eq, Hash)]
47enum LiteralKey {
48    U32(u32),
49    U64(u64),
50    I32(i32),
51    I64(i64),
52    F32(u32),
53    F64(u64),
54    Bool(bool),
55    Symbol(String),
56}
57
58#[derive(Debug, Clone)]
59struct ListDef {
60    id: u64,
61    elem_type: ScalarType,
62    items: Vec<Term>,
63    keys: Vec<LiteralKey>,
64}
65
66#[derive(Debug, Clone)]
67struct ListArg {
68    term: Term,
69    id: Option<u64>,
70    elem_type: ScalarType,
71    is_literal: bool,
72}
73
74#[derive(Debug, Clone, Copy)]
75struct ListColumn {
76    elem_type: ScalarType,
77}
78
79struct ListNormalizer {
80    /// Declared `list<T>` columns per predicate: `(column index, column)`.
81    list_columns: HashMap<String, Vec<(usize, ListColumn)>>,
82    lists_by_key: HashMap<(ScalarType, Vec<LiteralKey>), u64>,
83    lists: Vec<ListDef>,
84    next_list_id: u64,
85    helper_preds: HashSet<String>,
86    operation_rows: HashSet<(ListOp, ScalarType, Vec<u64>)>,
87    fresh_counter: usize,
88}
89
90impl ListNormalizer {
91    fn new(program: &Program) -> Result<Self> {
92        let domains: HashMap<String, ScalarType> = program
93            .domains
94            .iter()
95            .map(|DomainDecl { name, typ }| (name.clone(), *typ))
96            .collect();
97
98        let mut list_columns: HashMap<String, Vec<(usize, ListColumn)>> = HashMap::new();
99        for pred in &program.predicates {
100            for (idx, col) in pred.schema_columns().iter().enumerate() {
101                if let TypeRef::List(inner) = &col.typ {
102                    let elem_type = Self::storage_type_for_type_ref(inner, &domains)?;
103                    list_columns
104                        .entry(pred.name.clone())
105                        .or_default()
106                        .push((idx, ListColumn { elem_type }));
107                }
108            }
109        }
110
111        Ok(Self {
112            list_columns,
113            lists_by_key: HashMap::new(),
114            lists: Vec::new(),
115            next_list_id: 1,
116            helper_preds: HashSet::new(),
117            operation_rows: HashSet::new(),
118            fresh_counter: 0,
119        })
120    }
121
122    /// True when list normalization is the identity on this fact: with an
123    /// empty body only the head is rewritten, and `normalize_value_term` only
124    /// touches `Term::List`/`Term::Cons` (everything else is cloned as is).
125    /// Such facts are moved through without being rebuilt.
126    fn fact_is_identity(rule: &Rule) -> bool {
127        rule.body.is_empty()
128            && !rule
129                .head
130                .terms
131                .iter()
132                .any(|term| matches!(term, Term::List(_) | Term::Cons { .. }))
133    }
134
135    fn normalize_program(&mut self, program: Program) -> Result<Program> {
136        let mut out = program;
137        let mut rules = Vec::with_capacity(out.rules.len());
138        for rule in std::mem::take(&mut out.rules) {
139            if Self::fact_is_identity(&rule) {
140                rules.push(rule);
141            } else {
142                rules.extend(self.normalize_rule(&rule)?);
143            }
144        }
145        out.rules = rules;
146        out.constraints = std::mem::take(&mut out.constraints)
147            .into_iter()
148            .map(|constraint| {
149                let body = self.normalize_body(&constraint.body)?;
150                Ok(crate::ast::Constraint {
151                    authored_index: constraint.authored_index,
152                    body,
153                })
154            })
155            .collect::<Result<Vec<_>>>()?;
156        out.queries = std::mem::take(&mut out.queries)
157            .into_iter()
158            .map(|query| self.normalize_query(&query))
159            .collect::<Result<Vec<_>>>()?;
160
161        self.append_helper_declarations_and_facts(&mut out)?;
162        Ok(out)
163    }
164
165    fn normalize_query(&mut self, query: &Query) -> Result<Query> {
166        let mut var_list_types = HashMap::new();
167        let atoms = self.expand_atom(&query.atom, false, &mut var_list_types)?;
168        let atom = atoms
169            .into_iter()
170            .next()
171            .unwrap_or_else(|| query.atom.clone());
172        Ok(Query { atom })
173    }
174
175    fn normalize_rule(&mut self, rule: &Rule) -> Result<Vec<Rule>> {
176        let mut var_list_types = HashMap::new();
177        let head = self.normalize_head_atom(&rule.head)?;
178        let body = self.normalize_body_with_env(&rule.body, &mut var_list_types)?;
179        Ok(vec![Rule { head, body }])
180    }
181
182    fn normalize_body(&mut self, body: &[BodyLiteral]) -> Result<Vec<BodyLiteral>> {
183        let mut var_list_types = HashMap::new();
184        self.normalize_body_with_env(body, &mut var_list_types)
185    }
186
187    fn normalize_body_with_env(
188        &mut self,
189        body: &[BodyLiteral],
190        var_list_types: &mut HashMap<String, ScalarType>,
191    ) -> Result<Vec<BodyLiteral>> {
192        let mut out = Vec::new();
193        for lit in body {
194            match lit {
195                BodyLiteral::Positive(atom) => {
196                    for atom in self.expand_atom(atom, false, var_list_types)? {
197                        out.push(BodyLiteral::Positive(atom));
198                    }
199                }
200                BodyLiteral::Negated(atom) => {
201                    let atoms = self.expand_atom(atom, true, var_list_types)?;
202                    if atoms.len() != 1 {
203                        return Err(list_error(
204                            "negated list pattern expansion would require multiple atoms",
205                        ));
206                    }
207                    out.push(BodyLiteral::Negated(atoms.into_iter().next().unwrap()));
208                }
209                BodyLiteral::Comparison(cmp) => {
210                    out.push(BodyLiteral::Comparison(Comparison {
211                        op: cmp.op,
212                        left: self.normalize_value_term(&cmp.left, None)?,
213                        right: self.normalize_value_term(&cmp.right, None)?,
214                    }));
215                }
216                BodyLiteral::Epistemic(lit) => out.push(BodyLiteral::Epistemic(lit.clone())),
217                BodyLiteral::IsExpr(is_expr) => out.push(BodyLiteral::IsExpr(is_expr.clone())),
218                BodyLiteral::Univ(_) => {
219                    return Err(list_error(
220                        "univ literals must be normalized by the meta-term pass before list normalization",
221                    ));
222                }
223            }
224        }
225        Ok(out)
226    }
227
228    /// Element type of the declared `list<T>` column `idx` of `pred`, if any.
229    fn list_column_elem(&self, pred: &str, idx: usize) -> Option<ScalarType> {
230        self.list_columns
231            .get(pred)
232            .and_then(|cols| cols.iter().find(|(i, _)| *i == idx))
233            .map(|(_, col)| col.elem_type)
234    }
235
236    fn normalize_head_atom(&mut self, atom: &Atom) -> Result<Atom> {
237        let mut terms = Vec::with_capacity(atom.terms.len());
238        for (idx, term) in atom.terms.iter().enumerate() {
239            let expected = self.list_column_elem(&atom.predicate, idx);
240            terms.push(self.normalize_value_term(term, expected)?);
241        }
242        Ok(Atom {
243            predicate: atom.predicate.clone(),
244            terms,
245        })
246    }
247
248    fn expand_atom(
249        &mut self,
250        atom: &Atom,
251        negated: bool,
252        var_list_types: &mut HashMap<String, ScalarType>,
253    ) -> Result<Vec<Atom>> {
254        if let Some(atoms) = self.expand_list_builtin(atom, var_list_types)? {
255            return Ok(atoms);
256        }
257        if is_reserved_pair_helper(atom) {
258            return Err(list_error(
259                "pair helpers are reserved for meta-term normalization and are not implemented in list normalization",
260            ));
261        }
262
263        let mut terms = Vec::with_capacity(atom.terms.len());
264        let mut extra_atoms = Vec::new();
265        for (idx, term) in atom.terms.iter().enumerate() {
266            let expected = self.list_column_elem(&atom.predicate, idx);
267            match term {
268                Term::Variable(name) => {
269                    if let Some(elem_type) = expected {
270                        var_list_types.insert(name.clone(), elem_type);
271                    }
272                    terms.push(term.clone());
273                }
274                Term::Cons { head, tail } => {
275                    if negated {
276                        return Err(list_error(
277                            "cons patterns in negated atoms are unsupported before a positive binder",
278                        ));
279                    }
280                    let elem_type = expected.ok_or_else(|| {
281                        list_error("cons pattern requires a declared list<T> column")
282                    })?;
283                    let list_var = self.fresh_var("list_cons");
284                    let head_term = self.normalize_scalar_term(head, Some(elem_type))?.1;
285                    let tail_arg =
286                        self.normalize_list_arg(tail, Some(elem_type), var_list_types)?;
287                    if let Term::Variable(name) = &tail_arg.term {
288                        var_list_types.insert(name.clone(), elem_type);
289                    }
290                    terms.push(Term::Variable(list_var.clone()));
291                    extra_atoms.push(Atom {
292                        predicate: cons_pred(elem_type),
293                        terms: vec![Term::Variable(list_var), head_term, tail_arg.term],
294                    });
295                }
296                _ => {
297                    let normalized = self.normalize_value_term(term, expected)?;
298                    terms.push(normalized);
299                }
300            }
301        }
302
303        let mut atoms = vec![Atom {
304            predicate: atom.predicate.clone(),
305            terms,
306        }];
307        atoms.extend(extra_atoms);
308        Ok(atoms)
309    }
310
311    fn expand_list_builtin(
312        &mut self,
313        atom: &Atom,
314        var_list_types: &mut HashMap<String, ScalarType>,
315    ) -> Result<Option<Vec<Atom>>> {
316        let pred = atom.predicate.as_str();
317        match pred {
318            "is_list" => {
319                require_arity(atom, 1)?;
320                let arg = self.normalize_list_arg(&atom.terms[0], None, var_list_types)?;
321                Ok(Some(vec![Atom {
322                    predicate: len_pred().to_string(),
323                    terms: vec![arg.term, Term::Anonymous],
324                }]))
325            }
326            "member" | "memberchk" => {
327                require_arity(atom, 2)?;
328                let list = self.normalize_list_arg(&atom.terms[1], None, var_list_types)?;
329                let value = self
330                    .normalize_scalar_term(&atom.terms[0], Some(list.elem_type))?
331                    .1;
332                Ok(Some(vec![Atom {
333                    predicate: item_pred(list.elem_type),
334                    terms: vec![list.term, Term::Anonymous, value],
335                }]))
336            }
337            "length" => {
338                require_arity(atom, 2)?;
339                let list = self.normalize_list_arg(&atom.terms[0], None, var_list_types)?;
340                let len = self
341                    .normalize_scalar_term(&atom.terms[1], Some(ScalarType::U32))?
342                    .1;
343                Ok(Some(vec![Atom {
344                    predicate: len_pred().to_string(),
345                    terms: vec![list.term, len],
346                }]))
347            }
348            "nth" => {
349                require_arity(atom, 3)?;
350                let idx = self
351                    .normalize_scalar_term(&atom.terms[0], Some(ScalarType::U32))?
352                    .1;
353                let list = self.normalize_list_arg(&atom.terms[1], None, var_list_types)?;
354                let value = self
355                    .normalize_scalar_term(&atom.terms[2], Some(list.elem_type))?
356                    .1;
357                Ok(Some(vec![Atom {
358                    predicate: item_pred(list.elem_type),
359                    terms: vec![list.term, idx, value],
360                }]))
361            }
362            "append" => {
363                require_arity(atom, 3)?;
364                if matches!(atom.terms[2], Term::List(_) | Term::Cons { .. })
365                    && !matches!(atom.terms[0], Term::List(_) | Term::Cons { .. })
366                    && !matches!(atom.terms[1], Term::List(_) | Term::Cons { .. })
367                {
368                    return Err(list_error(
369                        "unbounded append/3 generation is unsupported during list normalization; bind the first two lists to finite literals",
370                    ));
371                }
372                let left = self.normalize_list_arg(&atom.terms[0], None, var_list_types)?;
373                let right =
374                    self.normalize_list_arg(&atom.terms[1], Some(left.elem_type), var_list_types)?;
375                if left.elem_type != right.elem_type {
376                    return Err(list_error(
377                        "append/3 list arguments must have the same element type",
378                    ));
379                }
380                let out =
381                    self.normalize_list_arg(&atom.terms[2], Some(left.elem_type), var_list_types)?;
382                if !left.is_literal || !right.is_literal {
383                    return Err(list_error(
384                        "unbounded append/3 generation is unsupported during list normalization; bind the first two lists to finite literals",
385                    ));
386                }
387                let left_id = left.id.expect("literal list has id");
388                let right_id = right.id.expect("literal list has id");
389                let out_id = match out.id {
390                    Some(id) => id,
391                    None => self.concat_lists(left_id, right_id)?,
392                };
393                self.operation_rows.insert((
394                    ListOp::Append,
395                    left.elem_type,
396                    vec![left_id, right_id, out_id],
397                ));
398                if let Term::Variable(name) = &out.term {
399                    var_list_types.insert(name.clone(), left.elem_type);
400                }
401                Ok(Some(vec![Atom {
402                    predicate: append_pred(left.elem_type),
403                    terms: vec![left.term, right.term, out.term],
404                }]))
405            }
406            "sort" | "msort" | "list_to_set" => {
407                require_arity(atom, 2)?;
408                let input = self.normalize_list_arg(&atom.terms[0], None, var_list_types)?;
409                if !input.is_literal {
410                    return Err(list_error(
411                        "sort/msort/list_to_set require a finite list literal during list normalization",
412                    ));
413                }
414                let op = match pred {
415                    "sort" => ListOp::Sort,
416                    "msort" => ListOp::MSort,
417                    "list_to_set" => ListOp::ListToSet,
418                    _ => unreachable!(),
419                };
420                let out =
421                    self.normalize_list_arg(&atom.terms[1], Some(input.elem_type), var_list_types)?;
422                let input_id = input.id.expect("literal list has id");
423                let out_id = match out.id {
424                    Some(id) => id,
425                    None => self.derived_ordered_list(input_id, op)?,
426                };
427                self.operation_rows
428                    .insert((op, input.elem_type, vec![input_id, out_id]));
429                if let Term::Variable(name) = &out.term {
430                    var_list_types.insert(name.clone(), input.elem_type);
431                }
432                let helper = match op {
433                    ListOp::Sort => sort_pred(input.elem_type),
434                    ListOp::MSort => msort_pred(input.elem_type),
435                    ListOp::ListToSet => list_to_set_pred(input.elem_type),
436                    ListOp::Append => unreachable!(),
437                };
438                Ok(Some(vec![Atom {
439                    predicate: helper,
440                    terms: vec![input.term, out.term],
441                }]))
442            }
443            _ => Ok(None),
444        }
445    }
446
447    fn normalize_value_term(
448        &mut self,
449        term: &Term,
450        expected_list_elem: Option<ScalarType>,
451    ) -> Result<Term> {
452        match term {
453            Term::List(_) => Ok(self
454                .normalize_list_arg(term, expected_list_elem, &mut HashMap::new())?
455                .term),
456            Term::Cons { head, tail } => {
457                let elem_type = expected_list_elem.ok_or_else(|| {
458                    list_error("cons literal requires a declared list<T> context")
459                })?;
460                let mut items = vec![self.normalize_scalar_term(head, Some(elem_type))?.1];
461                let tail_arg =
462                    self.normalize_list_arg(tail, Some(elem_type), &mut HashMap::new())?;
463                let tail_id = tail_arg
464                    .id
465                    .ok_or_else(|| list_error("cons literal tail must be a finite list literal"))?;
466                let tail_items = self.list_by_id(tail_id)?.items.clone();
467                items.extend(tail_items);
468                let id = self.register_list_terms(items, Some(elem_type))?;
469                Ok(Term::Integer(id as i64))
470            }
471            _ => Ok(term.clone()),
472        }
473    }
474
475    fn normalize_list_arg(
476        &mut self,
477        term: &Term,
478        expected_elem_type: Option<ScalarType>,
479        var_list_types: &mut HashMap<String, ScalarType>,
480    ) -> Result<ListArg> {
481        match term {
482            Term::List(items) => {
483                let id = self.register_list_terms(items.clone(), expected_elem_type)?;
484                let elem_type = self.list_by_id(id)?.elem_type;
485                Ok(ListArg {
486                    term: Term::Integer(id as i64),
487                    id: Some(id),
488                    elem_type,
489                    is_literal: true,
490                })
491            }
492            Term::Cons { .. } => {
493                let value = self.normalize_value_term(term, expected_elem_type)?;
494                let id = match value {
495                    Term::Integer(id) if id >= 0 => id as u64,
496                    _ => {
497                        return Err(list_error(
498                            "cons list argument must normalize to a finite list id",
499                        ));
500                    }
501                };
502                let elem_type = self.list_by_id(id)?.elem_type;
503                Ok(ListArg {
504                    term: Term::Integer(id as i64),
505                    id: Some(id),
506                    elem_type,
507                    is_literal: true,
508                })
509            }
510            Term::Integer(id) if *id > 0 => {
511                let elem_type = expected_elem_type
512                    .or_else(|| self.list_by_id(*id as u64).ok().map(|list| list.elem_type))
513                    .ok_or_else(|| list_error("list id requires a known list<T> type"))?;
514                Ok(ListArg {
515                    term: term.clone(),
516                    id: Some(*id as u64),
517                    elem_type,
518                    is_literal: false,
519                })
520            }
521            Term::Variable(name) => {
522                let elem_type = expected_elem_type
523                    .or_else(|| var_list_types.get(name).copied())
524                    .ok_or_else(|| {
525                        list_error(
526                            "list built-in requires a finite literal or known list<T> variable",
527                        )
528                    })?;
529                var_list_types.insert(name.clone(), elem_type);
530                Ok(ListArg {
531                    term: term.clone(),
532                    id: None,
533                    elem_type,
534                    is_literal: false,
535                })
536            }
537            Term::Anonymous => {
538                let elem_type = expected_elem_type.ok_or_else(|| {
539                    list_error("anonymous list argument requires a known list<T> type")
540                })?;
541                Ok(ListArg {
542                    term: term.clone(),
543                    id: None,
544                    elem_type,
545                    is_literal: false,
546                })
547            }
548            _ => Err(list_error(
549                "expected finite list literal or known list<T> variable",
550            )),
551        }
552    }
553
554    fn normalize_scalar_term(
555        &mut self,
556        term: &Term,
557        expected: Option<ScalarType>,
558    ) -> Result<(ScalarType, Term)> {
559        match term {
560            Term::Variable(_) | Term::Anonymous => {
561                let typ = expected.ok_or_else(|| {
562                    list_error("list element variable requires a known list<T> element type")
563                })?;
564                Ok((typ, term.clone()))
565            }
566            Term::Integer(i) => match expected {
567                Some(ScalarType::U32) => {
568                    let value = u32::try_from(*i)
569                        .map_err(|_| list_error("integer list element is out of range for u32"))?;
570                    Ok((ScalarType::U32, Term::Integer(value as i64)))
571                }
572                Some(ScalarType::U64) => {
573                    let value = u64::try_from(*i)
574                        .map_err(|_| list_error("integer list element is out of range for u64"))?;
575                    Ok((ScalarType::U64, Term::Integer(value as i64)))
576                }
577                Some(ScalarType::I32) => {
578                    let value = i32::try_from(*i)
579                        .map_err(|_| list_error("integer list element is out of range for i32"))?;
580                    Ok((ScalarType::I32, Term::Integer(value as i64)))
581                }
582                Some(ScalarType::I64) | None if *i < 0 || *i > u32::MAX as i64 => {
583                    Ok((ScalarType::I64, Term::Integer(*i)))
584                }
585                Some(ScalarType::I64) => Ok((ScalarType::I64, Term::Integer(*i))),
586                None => Ok((ScalarType::U32, Term::Integer(*i))),
587                Some(ScalarType::F32) => Ok((ScalarType::F32, Term::Float(*i as f64))),
588                Some(ScalarType::F64) => Ok((ScalarType::F64, Term::Float(*i as f64))),
589                Some(ScalarType::Bool) if *i == 0 || *i == 1 => {
590                    Ok((ScalarType::Bool, Term::Integer(*i)))
591                }
592                Some(ScalarType::Bool) => Err(list_error("bool list elements must be 0 or 1")),
593                Some(ScalarType::Symbol) => {
594                    Err(list_error("integer list element is not valid for symbol"))
595                }
596            },
597            Term::Float(f) => match expected {
598                Some(ScalarType::F32) => Ok((ScalarType::F32, Term::Float(*f))),
599                Some(ScalarType::F64) | None => Ok((ScalarType::F64, Term::Float(*f))),
600                Some(_) => Err(list_error(
601                    "float list element is not valid for expected type",
602                )),
603            },
604            Term::String(s) => {
605                if expected.is_none() || expected == Some(ScalarType::Symbol) {
606                    Ok((ScalarType::Symbol, Term::String(s.clone())))
607                } else {
608                    Err(list_error(
609                        "string list element is only valid for symbol lists",
610                    ))
611                }
612            }
613            Term::Symbol(id) => {
614                if expected.is_none() || expected == Some(ScalarType::Symbol) {
615                    Ok((ScalarType::Symbol, Term::Symbol(*id)))
616                } else {
617                    Err(list_error(
618                        "symbol list element is not valid for expected type",
619                    ))
620                }
621            }
622            Term::List(_) | Term::Cons { .. } => {
623                let id = self
624                    .normalize_list_arg(term, None, &mut HashMap::new())?
625                    .id
626                    .expect("literal list has id");
627                let typ = expected.unwrap_or(LIST_ID_TYPE);
628                if typ != LIST_ID_TYPE {
629                    return Err(list_error(
630                        "nested list element requires list<list<T>> or list<u64> context",
631                    ));
632                }
633                Ok((LIST_ID_TYPE, Term::Integer(id as i64)))
634            }
635            Term::Compound { .. } | Term::PredRef(_) | Term::Aggregate(_) => Err(list_error(
636                "compound, predref, and aggregate list elements require meta-term normalization support",
637            )),
638        }
639    }
640
641    fn register_list_terms(
642        &mut self,
643        items: Vec<Term>,
644        expected_elem_type: Option<ScalarType>,
645    ) -> Result<u64> {
646        let mut elem_type = expected_elem_type;
647        let mut normalized = Vec::with_capacity(items.len());
648        let mut keys = Vec::with_capacity(items.len());
649
650        for item in &items {
651            let (item_type, term) = match self.normalize_scalar_term(item, elem_type) {
652                Ok(value) => value,
653                Err(err) if elem_type.is_some() => {
654                    return Err(list_error(format!("heterogeneous list literal: {}", err)));
655                }
656                Err(err) => return Err(err),
657            };
658            if let Some(existing) = elem_type {
659                if existing != item_type {
660                    return Err(list_error(
661                        "heterogeneous list literals require a declared finite term type",
662                    ));
663                }
664            } else {
665                elem_type = Some(item_type);
666            }
667            keys.push(literal_key(&term, item_type)?);
668            normalized.push(term);
669        }
670
671        let elem_type = elem_type.unwrap_or(LIST_ID_TYPE);
672        let key = (elem_type, keys.clone());
673        if let Some(id) = self.lists_by_key.get(&key) {
674            return Ok(*id);
675        }
676
677        let id = self.next_list_id;
678        self.next_list_id += 1;
679        self.lists_by_key.insert(key, id);
680        self.lists.push(ListDef {
681            id,
682            elem_type,
683            items: normalized,
684            keys,
685        });
686        Ok(id)
687    }
688
689    fn concat_lists(&mut self, left_id: u64, right_id: u64) -> Result<u64> {
690        let left = self.list_by_id(left_id)?.clone();
691        let right = self.list_by_id(right_id)?.clone();
692        if left.elem_type != right.elem_type {
693            return Err(list_error(
694                "append/3 list arguments must have the same element type",
695            ));
696        }
697        let mut items = left.items;
698        items.extend(right.items);
699        self.register_list_terms(items, Some(right.elem_type))
700    }
701
702    fn derived_ordered_list(&mut self, input_id: u64, op: ListOp) -> Result<u64> {
703        let input = self.list_by_id(input_id)?.clone();
704        let mut pairs: Vec<(LiteralKey, Term)> = input.keys.into_iter().zip(input.items).collect();
705        pairs.sort_by(|(a, _), (b, _)| format!("{a:?}").cmp(&format!("{b:?}")));
706        if matches!(op, ListOp::Sort | ListOp::ListToSet) {
707            pairs.dedup_by(|(a, _), (b, _)| a == b);
708        }
709        let items = pairs.into_iter().map(|(_, term)| term).collect();
710        self.register_list_terms(items, Some(input.elem_type))
711    }
712
713    fn append_helper_declarations_and_facts(&mut self, program: &mut Program) -> Result<()> {
714        self.ensure_tail_lists()?;
715
716        let mut helper_rules = Vec::new();
717        for list in self.lists.clone() {
718            self.helper_preds.insert(len_pred().to_string());
719            helper_rules.push(fact(
720                len_pred(),
721                vec![
722                    Term::Integer(list.id as i64),
723                    Term::Integer(list.items.len() as i64),
724                ],
725            ));
726
727            let item_name = item_pred(list.elem_type);
728            self.helper_preds.insert(item_name.clone());
729            for (idx, item) in list.items.iter().enumerate() {
730                helper_rules.push(fact(
731                    &item_name,
732                    vec![
733                        Term::Integer(list.id as i64),
734                        Term::Integer(idx as i64),
735                        item.clone(),
736                    ],
737                ));
738            }
739
740            if let Some((head, tail_id)) = self.cons_fact_parts(list.id)? {
741                let cons_name = cons_pred(list.elem_type);
742                self.helper_preds.insert(cons_name.clone());
743                helper_rules.push(fact(
744                    &cons_name,
745                    vec![
746                        Term::Integer(list.id as i64),
747                        head,
748                        Term::Integer(tail_id as i64),
749                    ],
750                ));
751            }
752        }
753
754        // Deterministic emission order: `operation_rows` is a HashSet, so
755        // iterating it directly would make helper-fact order (and with it
756        // relation numbering in the plan) vary between processes.
757        let mut operation_rows: Vec<(ListOp, ScalarType, Vec<u64>)> =
758            self.operation_rows.iter().cloned().collect();
759        operation_rows.sort_unstable_by(|a, b| {
760            (a.0 as u8, a.1 as u8, &a.2).cmp(&(b.0 as u8, b.1 as u8, &b.2))
761        });
762        for (op, elem_type, ids) in operation_rows {
763            let pred = match op {
764                ListOp::Append => append_pred(elem_type),
765                ListOp::Sort => sort_pred(elem_type),
766                ListOp::MSort => msort_pred(elem_type),
767                ListOp::ListToSet => list_to_set_pred(elem_type),
768            };
769            self.helper_preds.insert(pred.clone());
770            helper_rules.push(fact(
771                &pred,
772                ids.into_iter().map(|id| Term::Integer(id as i64)).collect(),
773            ));
774        }
775
776        self.append_helper_pred_decls(program);
777        program.rules.extend(helper_rules);
778        Ok(())
779    }
780
781    fn ensure_tail_lists(&mut self) -> Result<()> {
782        loop {
783            let before = self.lists.len();
784            for list in self.lists.clone() {
785                if list.items.is_empty() {
786                    continue;
787                }
788                self.register_list_terms(list.items[1..].to_vec(), Some(list.elem_type))?;
789            }
790            if self.lists.len() == before {
791                break;
792            }
793        }
794        Ok(())
795    }
796
797    fn cons_fact_parts(&self, list_id: u64) -> Result<Option<(Term, u64)>> {
798        let list = self.list_by_id(list_id)?;
799        let Some(head) = list.items.first() else {
800            return Ok(None);
801        };
802        let tail_key = (list.elem_type, list.keys[1..].to_vec());
803        let tail_id = self
804            .lists_by_key
805            .get(&tail_key)
806            .copied()
807            .ok_or_else(|| list_error("missing normalized list tail"))?;
808        Ok(Some((head.clone(), tail_id)))
809    }
810
811    fn append_helper_pred_decls(&self, program: &mut Program) {
812        let mut existing: HashSet<String> = program
813            .predicates
814            .iter()
815            .map(|pred| pred.name.clone())
816            .collect();
817
818        // Sorted for a process-independent declaration order (see
819        // `append_helper_declarations_and_facts`).
820        let mut helper_preds: Vec<&String> = self.helper_preds.iter().collect();
821        helper_preds.sort();
822        for pred in helper_preds {
823            if existing.contains(pred) {
824                continue;
825            }
826            let columns = helper_columns(pred);
827            program.predicates.push(PredDecl {
828                name: pred.clone(),
829                types: columns.iter().map(|col| col.typ.clone()).collect(),
830                columns,
831                is_private: true,
832            });
833            existing.insert(pred.clone());
834        }
835    }
836
837    fn list_by_id(&self, id: u64) -> Result<&ListDef> {
838        self.lists
839            .iter()
840            .find(|list| list.id == id)
841            .ok_or_else(|| list_error("unknown normalized list id"))
842    }
843
844    fn fresh_var(&mut self, prefix: &str) -> String {
845        let var = format!(
846            "__XLOG_{}_{}",
847            prefix.to_ascii_uppercase(),
848            self.fresh_counter
849        );
850        self.fresh_counter += 1;
851        var
852    }
853
854    fn storage_type_for_type_ref(
855        typ: &TypeRef,
856        domains: &HashMap<String, ScalarType>,
857    ) -> Result<ScalarType> {
858        match typ {
859            TypeRef::Scalar(ty) => Ok(*ty),
860            TypeRef::Domain(name) => domains
861                .get(name)
862                .copied()
863                .ok_or_else(|| list_error(format!("unknown domain alias '{}' in list<T>", name))),
864            TypeRef::List(_) | TypeRef::Term | TypeRef::Compound | TypeRef::PredRef => {
865                Ok(LIST_ID_TYPE)
866            }
867        }
868    }
869}
870
871fn fact(pred: &str, terms: Vec<Term>) -> Rule {
872    Rule {
873        head: Atom {
874            predicate: pred.to_string(),
875            terms,
876        },
877        body: vec![],
878    }
879}
880
881fn require_arity(atom: &Atom, expected: usize) -> Result<()> {
882    if atom.terms.len() == expected {
883        Ok(())
884    } else {
885        Err(list_error(format!(
886            "{} expects {} arguments, got {}",
887            atom.predicate,
888            expected,
889            atom.terms.len()
890        )))
891    }
892}
893
894fn literal_key(term: &Term, typ: ScalarType) -> Result<LiteralKey> {
895    match (typ, term) {
896        (ScalarType::U32, Term::Integer(v)) => Ok(LiteralKey::U32(*v as u32)),
897        (ScalarType::U64, Term::Integer(v)) => Ok(LiteralKey::U64(*v as u64)),
898        (ScalarType::I32, Term::Integer(v)) => Ok(LiteralKey::I32(*v as i32)),
899        (ScalarType::I64, Term::Integer(v)) => Ok(LiteralKey::I64(*v)),
900        (ScalarType::F32, Term::Float(v)) => Ok(LiteralKey::F32((*v as f32).to_bits())),
901        (ScalarType::F64, Term::Float(v)) => Ok(LiteralKey::F64(v.to_bits())),
902        (ScalarType::Bool, Term::Integer(v)) => Ok(LiteralKey::Bool(*v != 0)),
903        (ScalarType::Symbol, Term::String(s)) => Ok(LiteralKey::Symbol(s.clone())),
904        (ScalarType::Symbol, Term::Symbol(id)) => {
905            Ok(LiteralKey::Symbol(xlog_core::symbol::resolve(*id)))
906        }
907        _ => Err(list_error(
908            "list literal element did not normalize to expected scalar type",
909        )),
910    }
911}
912
913fn helper_columns(pred: &str) -> Vec<PredColumn> {
914    let scalar = |name: &str, typ| PredColumn {
915        name: Some(name.to_string()),
916        typ: TypeRef::Scalar(typ),
917    };
918    if pred == len_pred() {
919        return vec![
920            scalar("list_id", LIST_ID_TYPE),
921            scalar("len", ScalarType::U32),
922        ];
923    }
924    if let Some(elem_type) = helper_elem_type(pred, "__xlog_list_item_") {
925        return vec![
926            scalar("list_id", LIST_ID_TYPE),
927            scalar("idx", ScalarType::U32),
928            scalar("value", elem_type),
929        ];
930    }
931    if let Some(elem_type) = helper_elem_type(pred, "__xlog_list_cons_") {
932        return vec![
933            scalar("list_id", LIST_ID_TYPE),
934            scalar("head", elem_type),
935            scalar("tail_id", LIST_ID_TYPE),
936        ];
937    }
938    if helper_elem_type(pred, "__xlog_list_append_").is_some() {
939        return vec![
940            scalar("left_id", LIST_ID_TYPE),
941            scalar("right_id", LIST_ID_TYPE),
942            scalar("out_id", LIST_ID_TYPE),
943        ];
944    }
945    vec![
946        scalar("input_id", LIST_ID_TYPE),
947        scalar("out_id", LIST_ID_TYPE),
948    ]
949}
950
951fn helper_elem_type(pred: &str, prefix: &str) -> Option<ScalarType> {
952    pred.strip_prefix(prefix).and_then(scalar_type_from_suffix)
953}
954
955fn scalar_type_from_suffix(s: &str) -> Option<ScalarType> {
956    match s {
957        "u32" => Some(ScalarType::U32),
958        "u64" => Some(ScalarType::U64),
959        "i32" => Some(ScalarType::I32),
960        "i64" => Some(ScalarType::I64),
961        "f32" => Some(ScalarType::F32),
962        "f64" => Some(ScalarType::F64),
963        "bool" => Some(ScalarType::Bool),
964        "symbol" => Some(ScalarType::Symbol),
965        _ => None,
966    }
967}
968
969fn scalar_suffix(typ: ScalarType) -> &'static str {
970    match typ {
971        ScalarType::U32 => "u32",
972        ScalarType::U64 => "u64",
973        ScalarType::I32 => "i32",
974        ScalarType::I64 => "i64",
975        ScalarType::F32 => "f32",
976        ScalarType::F64 => "f64",
977        ScalarType::Bool => "bool",
978        ScalarType::Symbol => "symbol",
979    }
980}
981
982fn len_pred() -> &'static str {
983    "__xlog_list_len"
984}
985
986fn item_pred(typ: ScalarType) -> String {
987    format!("__xlog_list_item_{}", scalar_suffix(typ))
988}
989
990fn cons_pred(typ: ScalarType) -> String {
991    format!("__xlog_list_cons_{}", scalar_suffix(typ))
992}
993
994fn append_pred(typ: ScalarType) -> String {
995    format!("__xlog_list_append_{}", scalar_suffix(typ))
996}
997
998fn sort_pred(typ: ScalarType) -> String {
999    format!("__xlog_list_sort_{}", scalar_suffix(typ))
1000}
1001
1002fn msort_pred(typ: ScalarType) -> String {
1003    format!("__xlog_list_msort_{}", scalar_suffix(typ))
1004}
1005
1006fn list_to_set_pred(typ: ScalarType) -> String {
1007    format!("__xlog_list_to_set_{}", scalar_suffix(typ))
1008}
1009
1010fn is_reserved_pair_helper(atom: &Atom) -> bool {
1011    matches!(
1012        (atom.predicate.as_str(), atom.terms.len()),
1013        ("pair", 3) | ("pairs", 2) | ("zip", 3) | ("zip_with_index", 2) | ("enumerate", 2)
1014    )
1015}
1016
1017fn list_error(message: impl Into<String>) -> XlogError {
1018    XlogError::Compilation(format!("list normalization error: {}", message.into()))
1019}