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