1use crate::defaults::{IntegerType, UnsignedIntegerType};
16use crate::error::{Error, Result};
17
18#[derive(Clone, Debug, Eq, PartialEq)]
34pub struct Sieve {
35 root: Node,
36}
37
38#[derive(Clone, Debug, Eq, PartialEq)]
39enum Node {
40 Residual {
42 modulus: UnsignedIntegerType,
43 shift: UnsignedIntegerType,
44 },
45 Not(Box<Node>),
46 And(Box<Node>, Box<Node>),
47 Or(Box<Node>, Box<Node>),
48 Xor(Box<Node>, Box<Node>),
49}
50
51impl Node {
52 fn contains(&self, z: IntegerType) -> bool {
53 match self {
54 Self::Residual { modulus, shift } => {
55 z.rem_euclid(*modulus as IntegerType) == *shift as IntegerType
56 }
57 Self::Not(inner) => !inner.contains(z),
58 Self::And(left, right) => left.contains(z) && right.contains(z),
59 Self::Or(left, right) => left.contains(z) || right.contains(z),
60 Self::Xor(left, right) => left.contains(z) != right.contains(z),
61 }
62 }
63
64 fn collect_moduli(&self, out: &mut Vec<UnsignedIntegerType>) {
65 match self {
66 Self::Residual { modulus, .. } => out.push(*modulus),
67 Self::Not(inner) => inner.collect_moduli(out),
68 Self::And(left, right) | Self::Or(left, right) | Self::Xor(left, right) => {
69 left.collect_moduli(out);
70 right.collect_moduli(out);
71 }
72 }
73 }
74}
75
76impl Sieve {
77 pub fn parse(expression: &str) -> Result<Self> {
84 let tokens = tokenize(expression)?;
85 let mut parser = Parser {
86 tokens: &tokens,
87 position: 0,
88 };
89 let root = parser.parse_or()?;
90 if parser.position != tokens.len() {
91 return Err(Error::Sieve(format!(
92 "trailing input in sieve {expression:?} at token {}",
93 parser.position
94 )));
95 }
96 Ok(Self { root })
97 }
98
99 pub fn contains(&self, z: IntegerType) -> bool {
101 self.root.contains(z)
102 }
103
104 pub fn period(&self) -> UnsignedIntegerType {
108 let mut moduli = Vec::new();
109 self.root.collect_moduli(&mut moduli);
110 moduli.into_iter().fold(1, lcm)
111 }
112
113 pub fn segment(&self, low: IntegerType, high: IntegerType) -> Vec<IntegerType> {
115 (low..=high).filter(|z| self.contains(*z)).collect()
116 }
117
118 pub fn interval_widths(&self) -> Result<Vec<IntegerType>> {
125 let period = self.period();
126 let members = self.segment(0, period as IntegerType);
127 if members.len() < 2 {
128 return Err(Error::Sieve(format!(
129 "sieve has {} member(s) in its period of {period}, so it defines no intervals",
130 members.len()
131 )));
132 }
133 Ok(members.windows(2).map(|pair| pair[1] - pair[0]).collect())
134 }
135}
136
137fn lcm(a: UnsignedIntegerType, b: UnsignedIntegerType) -> UnsignedIntegerType {
138 if a == 0 || b == 0 {
139 return 0;
140 }
141 a / gcd(a, b) * b
142}
143
144fn gcd(mut a: UnsignedIntegerType, mut b: UnsignedIntegerType) -> UnsignedIntegerType {
145 while b != 0 {
146 let remainder = a % b;
147 a = b;
148 b = remainder;
149 }
150 a
151}
152
153#[derive(Clone, Copy, Debug, Eq, PartialEq)]
154enum Token {
155 Number(UnsignedIntegerType),
156 At,
157 Not,
158 And,
159 Or,
160 Xor,
161 Open,
162 Close,
163}
164
165fn tokenize(expression: &str) -> Result<Vec<Token>> {
166 let mut tokens = Vec::new();
167 let mut chars = expression.chars().peekable();
168
169 while let Some(&character) = chars.peek() {
170 match character {
171 ' ' | '\t' | '\n' | '\r' => {
172 chars.next();
173 }
174 '0'..='9' => {
175 let mut value: UnsignedIntegerType = 0;
176 while let Some(&digit) = chars.peek() {
177 let Some(digit) = digit.to_digit(10) else {
178 break;
179 };
180 value = value
181 .checked_mul(10)
182 .and_then(|value| value.checked_add(digit))
183 .ok_or_else(|| {
184 Error::Sieve(format!("number overflows in sieve {expression:?}"))
185 })?;
186 chars.next();
187 }
188 tokens.push(Token::Number(value));
189 }
190 '@' => {
191 chars.next();
192 tokens.push(Token::At);
193 }
194 '-' => {
195 chars.next();
196 tokens.push(Token::Not);
197 }
198 '&' => {
199 chars.next();
200 tokens.push(Token::And);
201 }
202 '|' => {
203 chars.next();
204 tokens.push(Token::Or);
205 }
206 '^' => {
207 chars.next();
208 tokens.push(Token::Xor);
209 }
210 '{' | '(' => {
211 chars.next();
212 tokens.push(Token::Open);
213 }
214 '}' | ')' => {
215 chars.next();
216 tokens.push(Token::Close);
217 }
218 other => {
219 return Err(Error::Sieve(format!(
220 "unexpected character {other:?} in sieve {expression:?}"
221 )));
222 }
223 }
224 }
225
226 if tokens.is_empty() {
227 return Err(Error::Sieve("sieve expression is empty".to_string()));
228 }
229 Ok(tokens)
230}
231
232struct Parser<'a> {
233 tokens: &'a [Token],
234 position: usize,
235}
236
237impl Parser<'_> {
238 fn peek(&self) -> Option<Token> {
239 self.tokens.get(self.position).copied()
240 }
241
242 fn eat(&mut self, token: Token) -> bool {
243 if self.peek() == Some(token) {
244 self.position += 1;
245 return true;
246 }
247 false
248 }
249
250 fn parse_or(&mut self) -> Result<Node> {
251 let mut left = self.parse_xor()?;
252 while self.eat(Token::Or) {
253 let right = self.parse_xor()?;
254 left = Node::Or(Box::new(left), Box::new(right));
255 }
256 Ok(left)
257 }
258
259 fn parse_xor(&mut self) -> Result<Node> {
260 let mut left = self.parse_and()?;
261 while self.eat(Token::Xor) {
262 let right = self.parse_and()?;
263 left = Node::Xor(Box::new(left), Box::new(right));
264 }
265 Ok(left)
266 }
267
268 fn parse_and(&mut self) -> Result<Node> {
269 let mut left = self.parse_unary()?;
270 while self.eat(Token::And) {
271 let right = self.parse_unary()?;
272 left = Node::And(Box::new(left), Box::new(right));
273 }
274 Ok(left)
275 }
276
277 fn parse_unary(&mut self) -> Result<Node> {
278 if self.eat(Token::Not) {
279 return Ok(Node::Not(Box::new(self.parse_unary()?)));
280 }
281 self.parse_primary()
282 }
283
284 fn parse_primary(&mut self) -> Result<Node> {
285 if self.eat(Token::Open) {
286 let inner = self.parse_or()?;
287 if !self.eat(Token::Close) {
288 return Err(Error::Sieve("unclosed group in sieve".to_string()));
289 }
290 return Ok(inner);
291 }
292
293 let Some(Token::Number(modulus)) = self.peek() else {
294 return Err(Error::Sieve(format!(
295 "expected a modulus in sieve at token {}",
296 self.position
297 )));
298 };
299 self.position += 1;
300
301 if modulus == 0 {
302 return Err(Error::Sieve("sieve modulus must be non-zero".to_string()));
303 }
304
305 let shift = if self.eat(Token::At) {
307 let Some(Token::Number(shift)) = self.peek() else {
308 return Err(Error::Sieve(
309 "expected a shift after `@` in sieve".to_string(),
310 ));
311 };
312 self.position += 1;
313 shift
314 } else {
315 0
316 };
317
318 Ok(Node::Residual {
319 modulus,
320 shift: shift % modulus,
321 })
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328
329 fn widths(expression: &str) -> Vec<IntegerType> {
330 Sieve::parse(expression)
331 .expect("sieve parses")
332 .interval_widths()
333 .expect("sieve has intervals")
334 }
335
336 #[test]
337 fn a_single_residual_class_cycles_at_its_modulus() {
338 assert_eq!(widths("3@0"), [3]);
339 assert_eq!(widths("4@0"), [4]);
340 assert_eq!(widths("2@0"), [2]);
341 assert_eq!(widths("12@0"), [12]);
342 assert_eq!(widths("5"), [5]);
344 }
345
346 #[test]
347 fn the_major_scale_is_a_sieve() {
348 assert_eq!(
349 widths("(-3@2 & 4) | (-3@1 & 4@1) | (3@2 & 4@2) | (-3 & 4@3)"),
350 [2, 2, 1, 2, 2, 2, 1]
351 );
352 }
353
354 #[test]
355 fn union_intersection_and_symmetric_difference_match_music21() {
356 assert_eq!(widths("3@0|7@0"), [3, 3, 1, 2, 3, 2, 1, 3, 3]);
357 assert_eq!(widths("{3@0|4@0}"), [3, 1, 2, 2, 1, 3]);
358 assert_eq!(widths("3@0&4@0"), [12]);
359 assert_eq!(widths("3@0^4@0"), [1, 2, 2, 1]);
360 assert_eq!(widths("5@2|7@3"), [1, 4, 3, 2, 5, 5, 2, 3, 4, 1]);
361 }
362
363 #[test]
364 fn negation_applies_to_residuals_and_to_groups() {
365 assert_eq!(widths("-3@0"), [1]);
366 assert_eq!(widths("-5@2"), [1, 2, 1, 1]);
367 assert_eq!(widths("-{3@0|4@0}"), [1, 3, 2, 3, 1]);
368 }
369
370 #[test]
371 fn and_binds_tighter_than_or() {
372 assert_eq!(widths("3@0|4@0&6@0"), widths("3@0|{4@0&6@0}"));
374 assert_eq!(widths("3@0|4@0&6@0"), [3, 3, 3, 3]);
375 assert_ne!(widths("3@0|4@0&6@0"), widths("{3@0|4@0}&6@0"));
376 assert_eq!(widths("{3@0|4@0}&6@0"), [6, 6]);
377 }
378
379 #[test]
380 fn parentheses_and_braces_group_alike() {
381 assert_eq!(widths("(3@0|4@0)"), widths("{3@0|4@0}"));
382 }
383
384 #[test]
385 fn period_is_the_lcm_of_the_moduli() {
386 assert_eq!(Sieve::parse("3@0").unwrap().period(), 3);
387 assert_eq!(Sieve::parse("3@0|7@0").unwrap().period(), 21);
388 assert_eq!(Sieve::parse("5@2|7@3").unwrap().period(), 35);
389 assert_eq!(Sieve::parse("3@0|4@0").unwrap().period(), 12);
390 }
391
392 #[test]
393 fn a_sieve_too_sparse_for_intervals_errors() {
394 assert!(Sieve::parse("3@1").unwrap().interval_widths().is_err());
397 }
398
399 #[test]
400 fn malformed_expressions_error_instead_of_panicking() {
401 for bad in [
402 "", " ", "@", "3@", "|3@0", "3@0|", "(3@0", "3@0)", "0@0", "3@0 & ", "x", "3@@0",
403 ] {
404 let parsed = Sieve::parse(bad);
405 assert!(parsed.is_err(), "{bad:?} should be rejected");
406 }
407 }
408
409 #[test]
410 fn membership_wraps_for_negative_integers() {
411 let sieve = Sieve::parse("3@1").unwrap();
412 assert!(sieve.contains(1));
413 assert!(sieve.contains(4));
414 assert!(sieve.contains(-2));
415 assert!(!sieve.contains(0));
416 }
417}