Skip to main content

polydat_grammar/
pprint.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! AST → `.polydat` source pretty-printer.
5//!
6//! The printer gives a compiled `for` body its source text for
7//! diagnostics (`pp_file`), re-emits rewritten expressions in
8//! module inlining and tile lowering (`pp_expr`), and prints
9//! expressions in `polydat explain`. The runtime compiles from the
10//! AST; the printer is a faithful AST → source round-trip beside it.
11//!
12//! ## Round-trip contract
13//!
14//! For every `Statement`/`Expr` produced by the parser,
15//! `pp_statement` / `pp_expr` produces source text that re-parses
16//! into a semantically equivalent AST. "Semantically equivalent"
17//! means same node types and identical inner data (modulo
18//! `Span`s, which capture parser position and are not preserved
19//! across re-parse).
20//!
21//! ## Precedence and parens
22//!
23//! `BinOp` expressions are emitted with parens around the whole
24//! expression. This is uniformly safe — re-parsing produces the
25//! same tree structure — at the cost of extra parens. The output
26//! is the canonical spelling; the parens are uniform for round-trip
27//! safety.
28
29use crate::ast::{
30    Arg, BinOpKind, Binding, BindingModifier, CallExpr, CursorDecl, Expr, ExternPort, ForStmt,
31    ModuleDef, PolydatFile, Statement, TileBodyKind, TileDef, TileOptions, WireModifier,
32};
33
34/// Pretty-print a full file: every statement, separated by
35/// newlines.
36pub fn pp_file(file: &PolydatFile) -> String {
37    let mut out = String::new();
38    for stmt in &file.statements {
39        out.push_str(&pp_statement(stmt));
40        out.push('\n');
41    }
42    out
43}
44
45/// Pretty-print a top-level statement.
46pub fn pp_statement(stmt: &Statement) -> String {
47    match stmt {
48        Statement::InputDecl(d) => match &d.ty {
49            Some(ty) => format!("input {}: {}", d.name, ty),
50            None => format!("input {}", d.name),
51        },
52        Statement::Binding(b) => pp_binding(b),
53        Statement::ModuleDef(m) => pp_module_def(m),
54        Statement::ExternPort(p) => pp_extern_port(p),
55        Statement::Cursor(c) => pp_cursor(c),
56        Statement::Pragma { name, .. } => format!("pragma {name}"),
57        Statement::For(f) => pp_for_stmt(f, 0),
58        Statement::Tile(t) => pp_tile(t),
59    }
60}
61
62fn pp_tile(t: &TileDef) -> String {
63    let mut out = format!("tile {}", t.name);
64    if let Some(enc) = &t.encoding {
65        out.push_str(&format!(" : {enc}"));
66    }
67    let defaults = TileOptions::default();
68    let mut opts = Vec::new();
69    if t.options.open != defaults.open || t.options.close != defaults.close {
70        opts.push(format!(
71            "delims \"{}\" \"{}\"",
72            escape_string(&t.options.open),
73            escape_string(&t.options.close)
74        ));
75    }
76    if t.options.sigil != defaults.sigil {
77        opts.push(format!("sigil \"{}\"", escape_string(&t.options.sigil)));
78    }
79    if t.options.strict {
80        opts.push("strict".to_string());
81    }
82    if t.options.in_string {
83        opts.push("instring".to_string());
84    }
85    if !opts.is_empty() {
86        out.push_str(&format!(" ({})", opts.join(", ")));
87    }
88    // A tile binds a wire: `:=` precedes every body form.
89    match t.body_kind {
90        TileBodyKind::Block => {
91            out.push_str(" := ");
92            out.push_str(&t.body_text());
93        }
94        TileBodyKind::Heredoc => {
95            out.push_str(" := <<<\n");
96            out.push_str(&t.body_text());
97            out.push_str("\n>>>");
98        }
99        TileBodyKind::Literal => {
100            out.push_str(&format!(" := \"{}\"", escape_string(&t.body_text())))
101        }
102    }
103    out
104}
105
106fn pp_for_stmt(f: &ForStmt, indent: usize) -> String {
107    let pad = "    ".repeat(indent + 1);
108    let mut body = String::new();
109    for s in &f.body {
110        body.push_str(&pad);
111        body.push_str(&match s {
112            Statement::For(inner) => pp_for_stmt(inner, indent + 1),
113            other => pp_statement(other),
114        });
115        body.push('\n');
116    }
117    format!(
118        "for {} {{\n{}{}}}",
119        f.source.to_text(),
120        body,
121        "    ".repeat(indent)
122    )
123}
124
125/// Pretty-print an expression. Always emits parens around
126/// `BinOp` for round-trip safety.
127pub fn pp_expr(expr: &Expr) -> String {
128    match expr {
129        Expr::Ident(name, _) => name.clone(),
130        Expr::IntLit(v, _) => v.to_string(),
131        Expr::FloatLit(v, _) => format_float(*v),
132        Expr::StringLit(s, _) => format!("\"{}\"", escape_string(s)),
133        Expr::ArrayLit(elts, _) => {
134            let parts: Vec<String> = elts.iter().map(pp_expr).collect();
135            format!("[{}]", parts.join(", "))
136        }
137        Expr::Call(c) => pp_call(c),
138        Expr::BinOp(lhs, op, rhs) => {
139            format!("({} {} {})", pp_expr(lhs), pp_binop(*op), pp_expr(rhs))
140        }
141        Expr::UnaryNeg(e, _) => format!("(-{})", pp_expr(e)),
142        Expr::UnaryBitNot(e, _) => format!("(!{})", pp_expr(e)),
143        Expr::FieldAccess { source, field, .. } => format!("{source}.{field}"),
144        Expr::Cast(e, ty, _) => format!("({} as {})", pp_expr(e), ty.to_keyword()),
145        Expr::For(source) => format!("for {}", source.to_text()),
146    }
147}
148
149fn pp_binding(b: &Binding) -> String {
150    let mut target = if b.targets.len() == 1 {
151        b.targets[0].clone()
152    } else {
153        format!("({})", b.targets.join(", "))
154    };
155    // `shared name: type := …` — cell type annotation
156    // (scope_model.md §"Type stability").
157    if let Some(ty) = &b.type_annotation {
158        target = format!("{target}: {ty}");
159    }
160    let prefix = pp_modifier_prefix(b.modifier);
161    if prefix.is_empty() {
162        format!("{} := {}", target, pp_expr(&b.value))
163    } else {
164        format!("{} {} := {}", prefix, target, pp_expr(&b.value))
165    }
166}
167
168fn pp_extern_port(p: &ExternPort) -> String {
169    if let Some(default) = &p.default {
170        format!("extern {}: {} = {}", p.name, p.typ, pp_expr(default))
171    } else {
172        format!("extern {}: {}", p.name, p.typ)
173    }
174}
175
176fn pp_cursor(c: &CursorDecl) -> String {
177    let mut out = format!("cursor {} = {}", c.name, pp_expr(&c.constructor));
178    if let Some(over) = &c.over {
179        out.push_str(" over ");
180        out.push_str(&pp_expr(over));
181    }
182    out
183}
184
185fn pp_module_def(m: &ModuleDef) -> String {
186    let params: Vec<String> = m
187        .params
188        .iter()
189        .map(|p| format!("{}: {}", p.name, p.typ))
190        .collect();
191    let outputs: Vec<String> = m
192        .outputs
193        .iter()
194        .map(|p| format!("{}: {}", p.name, p.typ))
195        .collect();
196    let mut body = String::new();
197    for s in &m.body {
198        body.push_str("    ");
199        body.push_str(&pp_statement(s));
200        body.push('\n');
201    }
202    format!(
203        "{}({}) -> ({}) := {{\n{}}}",
204        m.name,
205        params.join(", "),
206        outputs.join(", "),
207        body
208    )
209}
210
211fn pp_call(c: &CallExpr) -> String {
212    let args: Vec<String> = c.args.iter().map(pp_arg).collect();
213    format!("{}({})", c.func, args.join(", "))
214}
215
216fn pp_arg(arg: &Arg) -> String {
217    match arg {
218        Arg::Positional(e) => pp_expr(e),
219        Arg::Named(name, e) => format!("{}: {}", name, pp_expr(e)),
220    }
221}
222
223fn pp_modifier_prefix(m: BindingModifier) -> String {
224    let mut parts: Vec<&str> = Vec::new();
225    if m.has(WireModifier::Const) {
226        parts.push("const");
227    }
228    if m.has(WireModifier::Shared) {
229        parts.push("shared");
230    }
231    if m.has(WireModifier::Volatile) {
232        parts.push("volatile");
233    }
234    parts.join(" ")
235}
236
237fn pp_binop(op: BinOpKind) -> &'static str {
238    match op {
239        BinOpKind::Add => "+",
240        BinOpKind::Sub => "-",
241        BinOpKind::Mul => "*",
242        BinOpKind::Div => "/",
243        BinOpKind::Mod => "%",
244        BinOpKind::Pow => "**",
245        BinOpKind::BitAnd => "&",
246        BinOpKind::BitOr => "|",
247        BinOpKind::BitXor => "^",
248        BinOpKind::Shl => "<<",
249        BinOpKind::Shr => ">>",
250        BinOpKind::Eq => "==",
251        BinOpKind::Ne => "!=",
252        BinOpKind::Lt => "<",
253        BinOpKind::Gt => ">",
254        BinOpKind::Le => "<=",
255        BinOpKind::Ge => ">=",
256        BinOpKind::And => "&&",
257        BinOpKind::Or => "||",
258    }
259}
260
261fn escape_string(s: &str) -> String {
262    let mut out = String::with_capacity(s.len());
263    for c in s.chars() {
264        match c {
265            '\\' => out.push_str("\\\\"),
266            '"' => out.push_str("\\\""),
267            '\n' => out.push_str("\\n"),
268            '\t' => out.push_str("\\t"),
269            '\r' => out.push_str("\\r"),
270            c => out.push(c),
271        }
272    }
273    out
274}
275
276fn format_float(v: f64) -> String {
277    if v.is_finite() && v == v.trunc() && v.abs() < 1e18 {
278        format!("{v:.1}")
279    } else {
280        format!("{v}")
281    }
282}
283
284#[cfg(test)]
285mod tests {
286    use super::*;
287    use crate::{lexer, parser};
288
289    fn parse(src: &str) -> PolydatFile {
290        let tokens = lexer::lex(src).expect("lex");
291        parser::parse(tokens).expect("parse")
292    }
293
294    fn round_trip(src: &str) {
295        let ast1 = parse(src);
296        let printed = pp_file(&ast1);
297        let ast2 = parse(&printed);
298        let printed2 = pp_file(&ast2);
299        assert_eq!(
300            printed, printed2,
301            "second-pass print should be idempotent.\n\
302             original source:\n{src}\n\n\
303             first print:\n{printed}\n\n\
304             second print:\n{printed2}"
305        );
306    }
307
308    #[test]
309    fn round_trip_simple_const() {
310        round_trip("const x := 42\n");
311    }
312
313    #[test]
314    fn round_trip_string_const() {
315        round_trip("const dataset := \"sift1m\"\n");
316    }
317
318    #[test]
319    fn round_trip_init_binding() {
320        round_trip("const prebuffer := dataset_prebuffer(\"example\")\n");
321    }
322
323    #[test]
324    fn round_trip_function_call() {
325        round_trip("ratio := mod(cycle, 100)\n");
326    }
327
328    #[test]
329    fn round_trip_named_args() {
330        round_trip("v := dist_normal(mean: 72.0, stddev: 5.0)\n");
331    }
332
333    #[test]
334    fn round_trip_binop() {
335        round_trip("y := (x + 1)\n");
336    }
337
338    #[test]
339    fn round_trip_inputs() {
340        round_trip("input (cycle: u64, thread: u64)\n");
341    }
342
343    #[test]
344    fn round_trip_extern() {
345        round_trip("extern dataset: String\n");
346    }
347
348    #[test]
349    fn round_trip_tuple_destructure() {
350        round_trip("(a, b) := unpack(cycle)\n");
351    }
352
353    #[test]
354    fn round_trip_workload_typical() {
355        // Mirrors the shape of full_cql_vector workload bindings.
356        let src = "\
357const dataset := \"sift1m\"
358const prefix := \"vec_default\"
359profiles := matching_profiles(dataset, prefix)
360table := first(profiles)
361";
362        round_trip(src);
363    }
364
365    #[test]
366    fn round_trip_string_escapes() {
367        round_trip("const s := \"hello \\\"world\\\"\"\n");
368    }
369
370    #[test]
371    fn round_trip_array_literal() {
372        round_trip("const weights := [60.0, 20.0, 15.0, 5.0]\n");
373    }
374
375    #[test]
376    fn round_trip_cursor() {
377        round_trip("cursor users = range(0, 1000000)\n");
378    }
379
380    #[test]
381    fn round_trip_cursor_with_over() {
382        // The `over <expr>` partition clause (SRD-71) must survive
383        // projection — pp_cursor emits it so over-bearing cursors
384        // round-trip faithfully.
385        round_trip("cursor q = range(0, 100) over p\n");
386    }
387
388    #[test]
389    fn pp_cursor_emits_over_clause() {
390        let ast = parse("cursor q = range(0, 100) over p\n");
391        let printed = pp_file(&ast);
392        assert!(
393            printed.contains(" over p"),
394            "projected cursor must retain its `over` clause, got:\n{printed}"
395        );
396    }
397}