Skip to main content

prima_core/
render.rs

1use num_bigint::BigInt;
2use num_rational::BigRational;
3use num_traits::Signed;
4
5use crate::expr_pool::{ExprData, ExprId, ExprPool};
6use crate::number::{Number, Real};
7use crate::symbol::SymbolTable;
8
9/// LaTeX view of a number (spec §8.3 default `print_format := latex`): rationals render as `\frac{n}{d}`.
10pub fn render_number(n: &Number) -> String {
11    match n {
12        Number::Integer(i) => i.to_string(),
13        Number::Rational(r) => format!("\\frac{{{}}}{{{}}}", r.numer(), r.denom()),
14        Number::Real(Real::F64(f)) => f.to_string(),
15        Number::Real(Real::F32(f)) => f.to_string(),
16        Number::Complex { re, im } => format!("{} + {}i", render_number(re), render_number(im)),
17        Number::I8(v) => v.to_string(),
18        Number::I16(v) => v.to_string(),
19        Number::I32(v) => v.to_string(),
20        Number::I64(v) => v.to_string(),
21        Number::I128(v) => v.to_string(),
22        Number::U8(v) => v.to_string(),
23        Number::U16(v) => v.to_string(),
24        Number::U32(v) => v.to_string(),
25        Number::U64(v) => v.to_string(),
26        Number::U128(v) => v.to_string(),
27        Number::Isize(v) => v.to_string(),
28        Number::Usize(v) => v.to_string(),
29        Number::BigFloat(f) => f.to_string(),
30    }
31}
32
33/// ExprDAG → LaTeX view (spec §8.3 level 0, preserving the original form). TeX names come from the symbol table (spec §7);
34/// rendering is just a view conversion, decoupled from the reverse parsing of `tex"..."` (spec §4.9).
35pub fn render_latex(pool: &ExprPool, symbols: &SymbolTable, id: ExprId) -> String {
36    match pool.get(id) {
37        Some(ExprData::Symbol(s)) => symbols.name(s).unwrap_or_else(|| format!("?{}", s.0)),
38        Some(ExprData::Integer(i)) => i.to_string(),
39        Some(ExprData::Rational(r)) => render_number(&Number::Rational(*r)),
40        Some(ExprData::Real(Real::F64(f))) => f.to_string(),
41        Some(ExprData::Real(Real::F32(f))) => f.to_string(),
42        Some(ExprData::Add(items)) => render_add(pool, symbols, &items),
43        Some(ExprData::Mul(items)) => render_mul(pool, symbols, &items),
44        Some(ExprData::Pow { base, exp }) => render_pow(pool, symbols, base, exp),
45        Some(ExprData::Apply { f, args }) => render_apply(pool, symbols, f, &args),
46        Some(ExprData::Indeterminate(_)) => "\\text{indeterminate}".into(),
47        None => "?".into(),
48    }
49}
50
51fn render_add(pool: &ExprPool, symbols: &SymbolTable, items: &[ExprId]) -> String {
52    let mut ordered: Vec<ExprId> = items.to_vec();
53    ordered.sort_by_key(|&id| match pool.get(id) {
54        Some(ExprData::Integer(_) | ExprData::Rational(_) | ExprData::Real(_)) => (1u8, 0u8),
55        Some(ExprData::Symbol(_)) => (0u8, 1u8),
56        _ => (0u8, 0u8),
57    });
58    let mut parts = Vec::new();
59    for (i, &item) in ordered.iter().enumerate() {
60        let s = render_signed(pool, symbols, item);
61        if i == 0 {
62            let trimmed = s.strip_prefix("+ ").unwrap_or(&s);
63            parts.push(trimmed.to_string());
64        } else {
65            parts.push(s);
66        }
67    }
68    parts.join(" ")
69}
70
71fn render_signed(pool: &ExprPool, symbols: &SymbolTable, id: ExprId) -> String {
72    match pool.get(id) {
73        Some(ExprData::Integer(i)) if *i < BigInt::from(0) => format!("- {}", -(*i)),
74        Some(ExprData::Rational(r)) if *r < BigRational::new(BigInt::from(0), BigInt::from(1)) => {
75            format!("- {}", render_number(&Number::Rational(r.abs())))
76        }
77        Some(ExprData::Real(Real::F64(f))) if f < 0.0 => format!("- {}", -f),
78        Some(ExprData::Real(Real::F32(f))) if f < 0.0 => format!("- {}", -f),
79        _ => {
80            let s = render_latex(pool, symbols, id);
81            if s.starts_with('-') {
82                s
83            } else {
84                format!("+ {s}")
85            }
86        }
87    }
88}
89
90fn render_mul(pool: &ExprPool, symbols: &SymbolTable, items: &[ExprId]) -> String {
91    let mut neg = false;
92    let mut parts = Vec::new();
93    for &item in items.iter() {
94        match pool.get(item) {
95            Some(ExprData::Integer(i)) if *i == BigInt::from(-1) && items.len() > 1 => neg = !neg,
96            _ => {
97                let s = render_latex(pool, symbols, item);
98                if is_atomic(pool, item) {
99                    parts.push(s);
100                } else {
101                    parts.push(format!("\\left({s}\\right)"));
102                }
103            }
104        }
105    }
106    let body = if parts.is_empty() {
107        "1".to_string()
108    } else {
109        parts.join(" ")
110    };
111    if neg { format!("-{body}") } else { body }
112}
113
114fn is_atomic(pool: &ExprPool, id: ExprId) -> bool {
115    matches!(
116        pool.get(id),
117        Some(
118            ExprData::Symbol(_) | ExprData::Integer(_) | ExprData::Rational(_) | ExprData::Real(_)
119        )
120    )
121}
122
123fn render_pow(pool: &ExprPool, symbols: &SymbolTable, base: ExprId, exp: ExprId) -> String {
124    let half = BigRational::new(BigInt::from(1), BigInt::from(2));
125    if matches!(pool.const_number(exp), Some(Number::Rational(r)) if r == half) {
126        return format!("\\sqrt{{{}}}", render_latex(pool, symbols, base));
127    }
128    let base_s = render_latex(pool, symbols, base);
129    let base_s = if is_atomic(pool, base) {
130        base_s
131    } else {
132        format!("\\left({base_s}\\right)")
133    };
134    format!("{base_s}^{{{}}}", render_latex(pool, symbols, exp))
135}
136
137fn render_apply(pool: &ExprPool, symbols: &SymbolTable, f: ExprId, args: &[ExprId]) -> String {
138    let name = match pool.get(f) {
139        Some(ExprData::Symbol(s)) => symbols.name(s).unwrap_or_else(|| "f".to_string()),
140        _ => "f".to_string(),
141    };
142    let arg_strs: Vec<String> = args
143        .iter()
144        .map(|&a| render_latex(pool, symbols, a))
145        .collect();
146    if name == "\\sqrt" && args.len() == 1 {
147        return format!("\\sqrt{{{}}}", arg_strs[0]);
148    }
149    format!("{name}\\left({}\\right)", arg_strs.join(", "))
150}