1use super::*;
2use nom::branch::alt;
3use nom::bytes::complete::tag;
4use nom::character::complete::{alpha1, alphanumeric1, digit1, one_of};
5use nom::combinator::{all_consuming, map, map_res, recognize};
6use nom::multi::{fold, many0, separated_list0};
7use nom::sequence::{delimited, pair, preceded, separated_pair};
8use nom::{IResult, Parser};
9use nom_language::error::VerboseError;
10
11type R<'i, O> = IResult<&'i str, O, VerboseError<&'i str>>;
12
13pub fn parse_tdim(symbol_table: &SymbolScope, input: &str) -> TractResult<TDim> {
14 match all_consuming(|i| expr(symbol_table, i)).parse(input) {
15 Ok(pair) => Ok(pair.1),
16 Err(e) => bail!("Failed to parse {:?}, {:?}", input, e),
17 }
18}
19
20pub fn parse_assertion(symbol_table: &SymbolScope, input: &str) -> TractResult<Assertion> {
21 match all_consuming(|i| assertion(symbol_table, i)).parse(input) {
22 Ok(pair) => Ok(pair.1),
23 Err(e) => bail!("Failed to parse {:?}, {:?}", input, e),
24 }
25}
26
27fn assertion<'i>(s: &SymbolScope, i: &'i str) -> R<'i, Assertion> {
28 delimited(
29 spaces,
30 alt((
31 map(separated_pair(|i| expr(s, i), stag("=="), |i| expr(s, i)), |(a, b)| {
32 Assertion::Eq(a, b)
33 }),
34 map(separated_pair(|i| expr(s, i), stag("<="), |i| expr(s, i)), |(a, b)| {
35 Assertion::LTE(a, b)
36 }),
37 map(separated_pair(|i| expr(s, i), stag(">="), |i| expr(s, i)), |(a, b)| {
38 Assertion::GTE(a, b)
39 }),
40 map(separated_pair(|i| expr(s, i), stag("<"), |i| expr(s, i)), |(a, b)| {
41 Assertion::LT(a, b)
42 }),
43 map(separated_pair(|i| expr(s, i), stag(">"), |i| expr(s, i)), |(a, b)| {
44 Assertion::GT(a, b)
45 }),
46 )),
47 spaces,
48 )
49 .parse(i)
50}
51
52fn expr<'i>(symbol_table: &SymbolScope, i: &'i str) -> R<'i, TDim> {
53 broadcast(symbol_table, i)
54}
55
56fn broadcast<'i>(symbol_table: &SymbolScope, input: &'i str) -> R<'i, TDim> {
57 let s = symbol_table;
58 let (mut input, mut result) = add(s, input)?;
59 while let Ok((i, _)) = stag("#").parse(input) {
60 let (i, next) = map_res(|i| add(s, i), |v| result.clone().broadcast(v)).parse(i)?;
61 (input, result) = (i, next);
62 }
63 Ok((input, result))
64}
65
66macro_rules! bin {
67 ($name: ident, $left: expr, $right: expr, $op: expr, $builder: expr) => {
68 fn $name<'i>(symbol_table: &SymbolScope, input: &'i str) -> R<'i, TDim> {
69 let s = symbol_table;
70 let (input, result) = $left(s, input)?;
71 fold(0.., preceded(stag($op), |i| $right(s, i)), move || result.clone(), $builder)
72 .parse(input)
73 }
74 };
75}
76
77bin!(add, sub, sub, "+", |a, b| a + b);
78bin!(sub, mul, mul, "-", |a, b| a - b);
79bin!(mul, div, div, "*", |a, b| a * b);
80bin!(div, atom, |_s, i| numeric(i), "/", |a, b| a / b);
81
82fn atom<'i>(symbol_table: &SymbolScope, i: &'i str) -> R<'i, TDim> {
83 alt((
84 map(numeric, TDim::Val),
85 map(|i| func(symbol_table, "min", i), TDim::Min),
86 map(|i| func(symbol_table, "max", i), TDim::Max),
87 map(|i| func(symbol_table, "broadcast", i), TDim::Broadcast),
88 map(|i| func(symbol_table, "floor", i), |xs| xs[0].clone()),
89 map(|i| identifier(symbol_table, i), TDim::Sym),
90 map(pair(recognize(stag("-")), |i| atom(symbol_table, i)), |(_, dim)| dim * -1),
91 delimited(stag("("), |i| expr(symbol_table, i), stag(")")),
92 ))
93 .parse(i)
94}
95
96fn func<'i>(symbol_table: &SymbolScope, name: &'static str, i: &'i str) -> R<'i, Vec<TDim>> {
97 preceded(
98 stag(name),
99 delimited(stag("("), separated_list0(stag(","), |i| expr(symbol_table, i)), stag(")")),
100 )
101 .parse(i)
102}
103
104fn identifier<'i>(symbol_table: &SymbolScope, i: &'i str) -> R<'i, Symbol> {
105 map(
106 recognize(pair(
107 alt((alpha1, tag("_"))),
108 many0(alt((alphanumeric1, tag("_"), tag("."), recognize(pair(tag("/"), alpha1))))),
109 )),
110 |s| symbol_table.sym(s),
111 )
112 .parse(i)
113}
114
115fn numeric(i: &str) -> R<'_, i64> {
116 map_res(digit1, std::str::FromStr::from_str).parse(i)
117}
118
119fn spaces(i: &str) -> R<'_, ()> {
120 map(many0(one_of(" \t\n\r")), |_| ()).parse(i)
121}
122
123fn spaced<'s, O, P>(it: P) -> impl Parser<&'s str, Output = O, Error = VerboseError<&'s str>>
124where
125 P: Parser<&'s str, Output = O, Error = VerboseError<&'s str>>,
126{
127 delimited(spaces, it, spaces)
128}
129
130pub(super) fn stag<'s>(
131 t: &'static str,
132) -> impl Parser<&'s str, Output = &'s str, Error = VerboseError<&'s str>> {
133 spaced(tag(t))
134}
135
136#[cfg(test)]
137mod test {
138 use super::*;
139
140 #[test]
141 fn parse_int() {
142 let table = SymbolScope::default();
143 assert_eq!(parse_tdim(&table, "12").unwrap(), TDim::Val(12));
144 assert_eq!(parse_tdim(&table, "-12").unwrap(), TDim::Val(-12));
145 }
146
147 #[test]
148 fn parse_sym() {
149 let table = SymbolScope::default();
150 assert_eq!(parse_tdim(&table, "x").unwrap(), TDim::Sym(table.sym("x")));
151 assert_eq!(
152 parse_tdim(&table, "-y").unwrap(),
153 TDim::MulInt(-1, Box::new(table.sym("y").into()))
154 );
155 }
156
157 #[test]
158 fn parse_bin() {
159 let table = SymbolScope::default();
160 assert_eq!(parse_tdim(&table, "1+2").unwrap(), 3.into());
161 assert_eq!(parse_tdim(&table, "1-2").unwrap(), (-1).into());
162 assert_eq!(parse_tdim(&table, "1*2").unwrap(), 2.into());
163 assert_eq!(parse_tdim(&table, "1/2").unwrap(), 0.into());
164 }
165
166 #[test]
167 fn parse_prio() {
168 let table = SymbolScope::default();
169 assert_eq!(parse_tdim(&table, "1+2*3").unwrap(), 7.into());
170 assert_eq!(parse_tdim(&table, "1*2+3").unwrap(), 5.into());
171 }
172
173 #[test]
174 fn parse_min() {
175 let table = SymbolScope::default();
176 assert_eq!(
177 parse_tdim(&table, "min(P,S)").unwrap(),
178 TDim::Min(vec!(table.sym("P").into(), table.sym("S").into()))
179 );
180 }
181
182 #[test]
183 fn parse_broadcast_func() {
184 let table = SymbolScope::default();
185 assert_eq!(
186 parse_tdim(&table, "broadcast(P,S)").unwrap(),
187 TDim::Broadcast(vec!(table.sym("P").into(), table.sym("S").into()))
188 );
189 }
190
191 #[test]
192 fn parse_broadcast_display_roundtrip() {
193 let table = SymbolScope::default();
194 let original = TDim::Broadcast(vec![table.sym("P").into(), table.sym("S").into()]);
195 let printed = format!("{original}");
196 let reparsed = parse_tdim(&table, &printed).unwrap();
197 assert_eq!(reparsed, original);
198 }
199
200 #[test]
201 fn parse_inequality_0() {
202 let table = SymbolScope::default();
203 assert_eq!(
204 parse_assertion(&table, "P+S<4096").unwrap(),
205 Assertion::LT(parse_tdim(&table, "P+S").unwrap(), 4096.to_dim())
206 );
207 }
208
209 #[test]
210 fn parse_dot_ids() {
211 let table = SymbolScope::default();
212 assert_eq!(parse_tdim(&table, "dot.0").unwrap(), table.sym("dot.0").into());
213 }
214
215 #[test]
216 fn parse_dot_ids_arith() {
217 let table = SymbolScope::default();
218 assert_eq!(parse_tdim(&table, "dot.0/2").unwrap(), table.sym("dot.0").to_dim() / 2);
219 }
220
221 #[test]
222 fn parse_floors() {
223 let table = SymbolScope::default();
224 assert_eq!(parse_tdim(&table, "floor(a)").unwrap(), table.sym("a").to_dim());
225 }
226
227 #[test]
228 fn parse_slash_ids() {
229 let table = SymbolScope::default();
230 assert_eq!(parse_tdim(&table, "foo/bar").unwrap(), table.sym("foo/bar").into());
231 assert_eq!(parse_tdim(&table, "foo/bar/baz").unwrap(), table.sym("foo/bar/baz").into());
232 }
233
234 #[test]
235 fn parse_slash_ids_arith() {
236 let table = SymbolScope::default();
237 assert_eq!(parse_tdim(&table, "foo/bar/2").unwrap(), table.sym("foo/bar").to_dim() / 2);
238 }
239
240 #[test]
241 fn parse_slash_display_roundtrip() {
242 let table = SymbolScope::default();
243 let original: TDim = table.sym("foo/bar").into();
244 let reparsed = parse_tdim(&table, &format!("{original}")).unwrap();
245 assert_eq!(reparsed, original);
246 }
247}