1use crate::arena::{ExprArena, ExprId, ExprNode, UnaryOp, VarId};
6use crate::linear::{LinearTerms, split_linear};
7
8pub(crate) const PREC_ADD: u8 = 1;
10pub(crate) const PREC_MUL: u8 = 2;
11pub(crate) const PREC_UNARY: u8 = 3;
12
13type Part = (bool, String);
15
16pub fn render_expr(arena: &ExprArena, id: ExprId, resolve: &impl Fn(VarId) -> String) -> String {
24 let (lin, residual) = split_linear(arena, id);
25 let mut parts = linear_parts(&lin, resolve);
26 for s in &residual {
27 let prec = if s.neg { PREC_UNARY } else { PREC_ADD };
28 parts.push((s.neg, render_node(arena, s.id, resolve, prec)));
29 }
30 join_parts(&parts)
31}
32
33pub fn render_linear_terms(t: &LinearTerms<'_>, resolve: &impl Fn(VarId) -> String) -> String {
41 join_parts(&linear_parts(t, resolve))
42}
43
44fn linear_parts(t: &LinearTerms<'_>, resolve: &impl Fn(VarId) -> String) -> Vec<Part> {
46 linear_parts_from_slice(&t.coeffs, t.constant, resolve)
47}
48
49fn linear_parts_from_slice(
50 coeffs: &[(VarId, f64)],
51 constant: f64,
52 resolve: &impl Fn(VarId) -> String,
53) -> Vec<Part> {
54 let mut parts = Vec::with_capacity(coeffs.len() + 1);
55 for (v, c) in coeffs {
56 if *c == 0.0 {
57 continue;
58 }
59 let mag = c.abs();
60 let text = if (mag - 1.0).abs() < f64::EPSILON {
61 resolve(*v)
62 } else {
63 format!("{} {}", fmt_num(mag), resolve(*v))
64 };
65 parts.push((*c < 0.0, text));
66 }
67 if constant != 0.0 {
68 parts.push((constant < 0.0, fmt_num(constant.abs())));
69 }
70 parts
71}
72
73fn join_parts(parts: &[Part]) -> String {
75 let Some(((first_neg, first), rest)) = parts.split_first() else {
76 return "0".to_string();
77 };
78 let mut out = String::new();
79 if *first_neg {
80 out.push('-');
81 }
82 out.push_str(first);
83 for (neg, text) in rest {
84 out.push_str(if *neg { " - " } else { " + " });
85 out.push_str(text);
86 }
87 out
88}
89
90pub(crate) fn render_node(
93 arena: &ExprArena,
94 id: ExprId,
95 resolve: &impl Fn(VarId) -> String,
96 parent_prec: u8,
97) -> String {
98 let (text, prec) = match arena.get(id) {
99 ExprNode::Const(c) => (fmt_num(*c), PREC_UNARY),
100 ExprNode::Var(v) => (resolve(*v), PREC_UNARY),
101 ExprNode::Param(p) => (fmt_num(arena.param_value(*p)), PREC_UNARY),
102 ExprNode::Unary(UnaryOp::Neg, x) => {
103 (format!("-{}", render_node(arena, *x, resolve, PREC_UNARY)), PREC_UNARY)
104 }
105 ExprNode::Add(children) => {
106 let mut parts: Vec<Part> = Vec::with_capacity(children.len());
107 for c in children.iter().copied() {
108 match arena.get(c) {
109 ExprNode::Unary(UnaryOp::Neg, inner) => {
110 parts.push((true, render_node(arena, *inner, resolve, PREC_UNARY)));
111 }
112 ExprNode::Const(v) if *v < 0.0 => parts.push((true, fmt_num(-v))),
113 ExprNode::Param(p) if arena.param_value(*p) < 0.0 => {
114 parts.push((true, fmt_num(-arena.param_value(*p))));
115 }
116 ExprNode::Linear { coeffs, constant } => {
117 parts.extend(linear_parts_from_slice(coeffs, *constant, resolve));
118 }
119 _ => parts.push((false, render_node(arena, c, resolve, PREC_ADD))),
120 }
121 }
122 (join_parts(&parts), PREC_ADD)
123 }
124 ExprNode::Mul(children) => {
125 let parts: Vec<String> =
126 children.iter().map(|c| render_node(arena, *c, resolve, PREC_MUL)).collect();
127 (parts.join(" * "), PREC_MUL)
128 }
129 ExprNode::Pow(b, e) => {
130 let base = render_node(arena, *b, resolve, PREC_UNARY);
131 let exp = render_node(arena, *e, resolve, PREC_UNARY);
132 (format!("{base}^{exp}"), PREC_UNARY)
133 }
134 ExprNode::Div(num, den) => {
135 let n = render_node(arena, *num, resolve, PREC_MUL);
136 let d = render_node(arena, *den, resolve, PREC_MUL);
137 (format!("{n} / {d}"), PREC_MUL)
138 }
139 ExprNode::Unary(op, x) => (fmt_call(op.name(), arena, *x, resolve), PREC_UNARY),
140 ExprNode::Atan2(y, x) => {
141 let y = render_node(arena, *y, resolve, PREC_ADD);
142 let x = render_node(arena, *x, resolve, PREC_ADD);
143 (format!("atan2({y}, {x})"), PREC_UNARY)
144 }
145 ExprNode::Min(children) | ExprNode::Max(children) => {
146 let name = if matches!(arena.get(id), ExprNode::Min(_)) { "min" } else { "max" };
147 let args = children
148 .iter()
149 .map(|child| render_node(arena, *child, resolve, PREC_ADD))
150 .collect::<Vec<_>>()
151 .join(", ");
152 (format!("{name}({args})"), PREC_UNARY)
153 }
154 ExprNode::Linear { coeffs, constant } => {
155 let parts = linear_parts_from_slice(coeffs, *constant, resolve);
156 let prec = match parts.as_slice() {
158 [(false, _)] => PREC_MUL,
159 [] => PREC_UNARY,
160 _ => PREC_ADD,
161 };
162 (join_parts(&parts), prec)
163 }
164 };
165 if prec < parent_prec { format!("({text})") } else { text }
166}
167
168fn fmt_call(
170 name: &str,
171 arena: &ExprArena,
172 arg: ExprId,
173 resolve: &impl Fn(VarId) -> String,
174) -> String {
175 format!("{name}({})", render_node(arena, arg, resolve, PREC_ADD))
176}
177
178pub(crate) fn fmt_num(v: f64) -> String {
180 if v == 0.0 { "0".to_string() } else { format!("{v}") }
181}
182
183#[cfg(test)]
184mod tests {
185 use super::*;
186 use crate::arena::{ExprArena, ExprNode, VarId};
187
188 fn names(v: VarId) -> String {
189 match v.0 {
190 0 => "x".to_string(),
191 1 => "y".to_string(),
192 2 => "z".to_string(),
193 n => format!("v{n}"),
194 }
195 }
196
197 fn lt(coeffs: Vec<(u32, f64)>, constant: f64) -> LinearTerms<'static> {
198 LinearTerms {
199 coeffs: std::borrow::Cow::Owned(
200 coeffs.into_iter().map(|(v, c)| (VarId(v), c)).collect(),
201 ),
202 constant,
203 }
204 }
205
206 #[test]
207 fn linear_terms_are_sign_aware() {
208 let t = lt(vec![(0, 1.0), (1, 2.0), (2, -3.0)], 0.0);
209 assert_eq!(render_linear_terms(&t, &names), "x + 2 y - 3 z");
210 }
211
212 #[test]
213 fn leading_negative_and_constant() {
214 let t = lt(vec![(0, -1.0)], 2.0);
215 assert_eq!(render_linear_terms(&t, &names), "-x + 2");
216 let t = lt(vec![(0, 1.0)], -2.5);
217 assert_eq!(render_linear_terms(&t, &names), "x - 2.5");
218 }
219
220 #[test]
221 fn zero_coeffs_skipped_and_empty_is_zero() {
222 let t = lt(vec![(0, 0.0), (1, 1.0)], 0.0);
223 assert_eq!(render_linear_terms(&t, &names), "y");
224 assert_eq!(render_linear_terms(<(vec![], 0.0), &names), "0");
225 assert_eq!(render_linear_terms(<(vec![], -0.0), &names), "0");
226 }
227
228 #[test]
229 fn constant_only() {
230 assert_eq!(render_linear_terms(<(vec![], 5.0), &names), "5");
231 assert_eq!(render_linear_terms(<(vec![], -5.0), &names), "-5");
232 }
233
234 #[test]
235 fn expr_linear_and_nonlinear_mix() {
236 let mut arena = ExprArena::new();
237 let x = arena.push(ExprNode::Var(VarId(0)));
238 let y = arena.push(ExprNode::Var(VarId(1)));
239 let z = arena.push(ExprNode::Var(VarId(2)));
240 let two = arena.constant(2.0);
241 let two_z = arena.push(ExprNode::Mul(smallvec::smallvec![two, z]));
242 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![x, y]));
243 let sum = arena.push(ExprNode::Add(smallvec::smallvec![two_z, prod]));
244 assert_eq!(render_expr(&arena, sum, &names), "2 z + x * y");
245 }
246
247 #[test]
248 fn expr_negated_residual() {
249 let mut arena = ExprArena::new();
250 let x = arena.push(ExprNode::Var(VarId(0)));
251 let s = arena.push(ExprNode::Unary(UnaryOp::Sin, x));
252 let neg = arena.push(ExprNode::Unary(UnaryOp::Neg, s));
253 assert_eq!(render_expr(&arena, neg, &names), "-sin(x)");
254
255 let y = arena.push(ExprNode::Var(VarId(1)));
256 let sum = arena.push(ExprNode::Add(smallvec::smallvec![y, neg]));
257 assert_eq!(render_expr(&arena, sum, &names), "y - sin(x)");
258 }
259
260 #[test]
261 fn expr_pure_linear_uses_split() {
262 let mut arena = ExprArena::new();
263 let e = arena.push(ExprNode::Linear {
264 coeffs: vec![(VarId(0), 3.0), (VarId(1), -1.0)],
265 constant: 1.5,
266 });
267 assert_eq!(render_expr(&arena, e, &names), "3 x - y + 1.5");
268 }
269
270 #[test]
271 fn precedence_parenthesizes_sums_in_products() {
272 let mut arena = ExprArena::new();
273 let x = arena.push(ExprNode::Var(VarId(0)));
274 let y = arena.push(ExprNode::Var(VarId(1)));
275 let one = arena.constant(1.0);
276 let sum = arena.push(ExprNode::Add(smallvec::smallvec![x, one]));
277 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![sum, y]));
278 assert_eq!(render_expr(&arena, prod, &names), "(x + 1) * y");
279 }
280
281 #[test]
282 fn nested_add_is_sign_aware() {
283 let mut arena = ExprArena::new();
284 let x = arena.push(ExprNode::Var(VarId(0)));
285 let y = arena.push(ExprNode::Var(VarId(1)));
286 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![x, y]));
287 let neg_z = arena.push(ExprNode::Linear { coeffs: vec![(VarId(2), -1.0)], constant: 0.0 });
288 let sum = arena.push(ExprNode::Add(smallvec::smallvec![prod, neg_z]));
289 let s = arena.push(ExprNode::Unary(UnaryOp::Sin, sum));
290 assert_eq!(render_expr(&arena, s, &names), "sin(x * y - z)");
291 }
292
293 #[test]
294 fn negative_param_in_nonlinear_add_is_sign_aware() {
295 let mut arena = ExprArena::new();
296 let x = arena.push(ExprNode::Var(VarId(0)));
297 let pid = arena.new_param(-3.0);
298 let p = arena.param(pid);
299 let sum = arena.push(ExprNode::Add(smallvec::smallvec![x, p]));
300 let s = arena.push(ExprNode::Unary(UnaryOp::Sin, sum));
301 assert_eq!(render_expr(&arena, s, &names), "sin(x - 3)");
302
303 arena.set_param_value(pid, 3.0);
304 assert_eq!(render_expr(&arena, s, &names), "sin(x + 3)");
305 }
306
307 #[test]
308 fn linear_node_inside_product_parenthesizes_when_needed() {
309 let mut arena = ExprArena::new();
310 let y = arena.push(ExprNode::Var(VarId(1)));
311 let two_x = arena.push(ExprNode::Linear { coeffs: vec![(VarId(0), 2.0)], constant: 0.0 });
312 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![two_x, y]));
313 assert_eq!(render_expr(&arena, prod, &names), "2 x * y");
314
315 let sum = arena.push(ExprNode::Linear { coeffs: vec![(VarId(0), 1.0)], constant: 1.0 });
316 let prod2 = arena.push(ExprNode::Mul(smallvec::smallvec![sum, y]));
317 assert_eq!(render_expr(&arena, prod2, &names), "(x + 1) * y");
318 }
319
320 #[test]
321 fn negative_zero_constant_renders_as_zero() {
322 assert_eq!(fmt_num(-0.0), "0");
323 let mut arena = ExprArena::new();
324 let c = arena.constant(-0.0);
325 assert_eq!(render_expr(&arena, c, &names), "0");
326 }
327}