1use num_traits::{One, Pow, Signed, Zero};
17use std::collections::{BTreeMap, btree_map};
18use std::fmt::Display;
19use std::iter::{Product, Sum};
20use std::ops::{Add, AddAssign, Mul, MulAssign, Neg, Sub, SubAssign};
21
22use derivative::Derivative;
23use duplicate::duplicate_item;
24
25pub trait AdditiveMonoid: Add<Output = Self> + Zero {}
27
28#[duplicate_item(T; [f32]; [f64]; [i32]; [i64]; [u32]; [u64]; [usize])]
29impl AdditiveMonoid for T {}
30
31pub trait AbGroup: AdditiveMonoid + Neg<Output = Self> {}
37
38#[duplicate_item(T; [f32]; [f64]; [i32]; [i64])]
39impl AbGroup for T {}
40
41pub trait Monoid: Mul<Output = Self> + One {}
43
44#[duplicate_item(T; [f32]; [f64]; [i32]; [i64]; [u32]; [u64]; [usize])]
45impl Monoid for T {}
46
47pub trait CommMonoid: Monoid {}
49
50#[duplicate_item(T; [f32]; [f64]; [i32]; [i64]; [u32]; [u64]; [usize])]
51impl CommMonoid for T {}
52
53pub trait Rig: Monoid + AdditiveMonoid {}
55
56#[duplicate_item(T; [f32]; [f64]; [i32]; [i64]; [u32]; [u64]; [usize])]
57impl Rig for T {}
58
59pub trait CommRig: Rig + CommMonoid {}
61
62#[duplicate_item(T; [f32]; [f64]; [i32]; [i64]; [u32]; [u64]; [usize])]
63impl CommRig for T {}
64
65pub trait Ring: Rig + AbGroup {}
67
68#[duplicate_item(T; [f32]; [f64]; [i32]; [i64])]
69impl Ring for T {}
70
71pub trait CommRing: Ring + CommRig {}
73
74#[duplicate_item(T; [f32]; [f64]; [i32]; [i64])]
75impl CommRing for T {}
76
77pub trait RigModule: AdditiveMonoid + Mul<Self::Rig, Output = Self> {
79 type Rig: CommRig;
81}
82
83pub trait Module: RigModule<Rig = Self::Ring> + AbGroup {
85 type Ring: CommRing;
87}
88
89pub trait DisplayCoef {
91 fn has_negative_sign(&self) -> bool;
98
99 fn needs_parentheses(&self) -> bool;
101}
102
103#[duplicate_item(T; [u32]; [u64]; [usize])]
104impl DisplayCoef for T {
105 fn has_negative_sign(&self) -> bool {
106 false
107 }
108 fn needs_parentheses(&self) -> bool {
109 false
110 }
111}
112
113#[duplicate_item(T; [f32]; [f64]; [i32]; [i64])]
114impl DisplayCoef for T {
115 fn has_negative_sign(&self) -> bool {
116 self.is_negative()
117 }
118 fn needs_parentheses(&self) -> bool {
119 false
120 }
121}
122
123#[derive(Clone, PartialEq, Eq, Debug, Derivative)]
137#[derivative(Default(bound = ""))]
138pub struct Combination<Var, Coef>(BTreeMap<Var, Coef>);
139
140impl<Var, Coef> Combination<Var, Coef>
141where
142 Var: Ord,
143{
144 pub fn generator(var: Var) -> Self
146 where
147 Coef: One,
148 {
149 Combination([(var, Coef::one())].into_iter().collect())
150 }
151
152 pub fn len(&self) -> usize {
154 self.0.len()
155 }
156
157 pub fn is_empty(&self) -> bool {
159 self.0.is_empty()
160 }
161
162 pub fn variables(&self) -> impl ExactSizeIterator<Item = &Var> {
164 self.0.keys()
165 }
166
167 pub fn extend_scalars<NewCoef, F>(self, mut f: F) -> Combination<Var, NewCoef>
174 where
175 F: FnMut(Coef) -> NewCoef,
176 {
177 Combination(self.0.into_iter().map(|(var, coef)| (var, f(coef))).collect())
178 }
179
180 pub fn eval<A, F>(&self, mut f: F) -> A
182 where
183 A: Mul<Coef, Output = A> + Sum,
184 F: FnMut(&Var) -> A,
185 Coef: Clone,
186 {
187 self.0.iter().map(|(var, coef)| f(var) * coef.clone()).sum()
188 }
189
190 pub fn eval_with_order<A>(&self, values: impl IntoIterator<Item = A>) -> A
196 where
197 A: Mul<Coef, Output = A> + Sum,
198 Coef: Clone,
199 {
200 let mut iter = values.into_iter();
201 let value = self.eval(|_| iter.next().expect("Should have enough values"));
202 assert!(iter.next().is_none(), "Too many values");
203 value
204 }
205
206 pub fn normalize(self) -> Self
208 where
209 Coef: Zero,
210 {
211 self.into_iter().filter(|(coef, _)| !coef.is_zero()).collect()
212 }
213}
214
215impl<Var, Coef> FromIterator<(Coef, Var)> for Combination<Var, Coef>
217where
218 Var: Ord,
219 Coef: Add<Output = Coef>,
220{
221 fn from_iter<T: IntoIterator<Item = (Coef, Var)>>(iter: T) -> Self {
222 let mut combination = Combination::default();
223 for rhs in iter {
224 combination += rhs;
225 }
226 combination
227 }
228}
229
230impl<Var, Coef> IntoIterator for Combination<Var, Coef> {
232 type Item = (Coef, Var);
233 type IntoIter = std::iter::Map<btree_map::IntoIter<Var, Coef>, fn((Var, Coef)) -> (Coef, Var)>;
234
235 fn into_iter(self) -> Self::IntoIter {
236 self.0.into_iter().map(|(var, coef)| (coef, var))
237 }
238}
239
240impl<'a, Var, Coef> IntoIterator for &'a Combination<Var, Coef> {
241 type Item = (&'a Coef, &'a Var);
242 type IntoIter = std::iter::Map<
243 btree_map::Iter<'a, Var, Coef>,
244 fn((&'a Var, &'a Coef)) -> (&'a Coef, &'a Var),
245 >;
246
247 fn into_iter(self) -> Self::IntoIter {
248 self.0.iter().map(|(var, coef)| (coef, var))
249 }
250}
251
252impl<Var, Coef> Display for Combination<Var, Coef>
253where
254 Var: Display,
255 Coef: Display + DisplayCoef + Clone + PartialEq + One + Neg<Output = Coef>,
256{
257 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
258 let fmt_scalar_mul = |f: &mut std::fmt::Formatter<'_>, coef: &Coef, var: &Var| {
259 if coef.is_one() {
260 write!(f, "{var}")
261 } else if *coef == Coef::one().neg() {
262 write!(f, "-{var}")
263 } else if coef.needs_parentheses() {
264 write!(f, "({coef}) {var}")
265 } else {
266 write!(f, "{coef} {var}")
267 }
268 };
269
270 let mut pairs = self.0.iter();
271 if let Some((var, coef)) = pairs.next() {
272 fmt_scalar_mul(f, coef, var)?;
273 } else {
274 write!(f, "0")?;
275 }
276 for (var, coef) in pairs {
277 if coef.has_negative_sign() {
278 write!(f, " - ")?;
279 fmt_scalar_mul(f, &coef.clone().neg(), var)?;
280 } else {
281 write!(f, " + ")?;
282 fmt_scalar_mul(f, coef, var)?;
283 }
284 }
285 Ok(())
286 }
287}
288
289impl<Var, Coef> AddAssign<(Coef, Var)> for Combination<Var, Coef>
290where
291 Var: Ord,
292 Coef: Add<Output = Coef>,
293{
294 fn add_assign(&mut self, rhs: (Coef, Var)) {
295 let rhs = (rhs.1, rhs.0);
296 _add_assign(&mut self.0, rhs);
297 }
298}
299
300impl<Var, Coef> AddAssign for Combination<Var, Coef>
301where
302 Var: Ord,
303 Coef: Add<Output = Coef>,
304{
305 fn add_assign(&mut self, rhs: Self) {
306 for rhs in rhs.0 {
307 _add_assign(&mut self.0, rhs);
308 }
309 }
310}
311
312fn _add_assign<K, V>(lhs: &mut BTreeMap<K, V>, rhs: (K, V))
313where
314 K: Ord,
315 V: Add<Output = V>,
316{
317 let (k, b) = rhs;
318 if let Some(a) = lhs.remove(&k) {
319 lhs.insert(k, a + b);
320 } else {
321 lhs.insert(k, b);
322 }
323}
324
325impl<Var, Coef> Add for Combination<Var, Coef>
326where
327 Var: Ord,
328 Coef: Add<Output = Coef>,
329{
330 type Output = Self;
331
332 fn add(mut self, rhs: Self) -> Self {
333 self += rhs;
334 self
335 }
336}
337
338impl<Var, Coef> Zero for Combination<Var, Coef>
339where
340 Var: Ord,
341 Coef: Add<Output = Coef> + Zero,
342{
343 fn zero() -> Self {
344 Combination(Default::default())
345 }
346
347 fn is_zero(&self) -> bool {
348 self.0.values().all(|coef| coef.is_zero())
349 }
350}
351
352impl<Var, Coef> AdditiveMonoid for Combination<Var, Coef>
353where
354 Var: Ord,
355 Coef: AdditiveMonoid,
356{
357}
358
359impl<Var, Coef> Mul<Coef> for Combination<Var, Coef>
360where
361 Var: Ord,
362 Coef: Clone + Default + Mul<Output = Coef>,
363{
364 type Output = Self;
365
366 fn mul(mut self, a: Coef) -> Self {
367 for coef in self.0.values_mut() {
368 *coef = std::mem::take(coef) * a.clone();
369 }
370 self
371 }
372}
373
374impl<Var, Coef> RigModule for Combination<Var, Coef>
375where
376 Var: Ord,
377 Coef: Clone + Default + CommRig,
378{
379 type Rig = Coef;
380}
381
382impl<Var, Coef> Neg for Combination<Var, Coef>
383where
384 Var: Ord,
385 Coef: Default + Neg<Output = Coef>,
386{
387 type Output = Self;
388
389 fn neg(mut self) -> Self {
390 for coef in self.0.values_mut() {
391 *coef = std::mem::take(coef).neg();
392 }
393 self
394 }
395}
396
397impl<Var, Coef> SubAssign for Combination<Var, Coef>
398where
399 Var: Ord,
400 Coef: Default + Add<Output = Coef> + Neg<Output = Coef> + Zero,
401{
402 fn sub_assign(&mut self, rhs: Self) {
403 *self += -rhs;
404 }
405}
406
407impl<Var, Coef> Sub for Combination<Var, Coef>
408where
409 Var: Ord,
410 Coef: Default + Add<Output = Coef> + Neg<Output = Coef> + Zero,
411{
412 type Output = Self;
413
414 fn sub(mut self, rhs: Self) -> Self::Output {
415 self -= rhs;
416 self
417 }
418}
419
420impl<Var, Coef> AbGroup for Combination<Var, Coef>
421where
422 Var: Ord,
423 Coef: Default + AbGroup,
424{
425}
426
427impl<Var, Coef> Module for Combination<Var, Coef>
428where
429 Var: Ord,
430 Coef: Clone + Default + CommRing,
431{
432 type Ring = Coef;
433}
434
435#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Debug, Derivative)]
453#[derivative(Default(bound = ""))]
454pub struct Monomial<Var, Exp>(BTreeMap<Var, Exp>);
455
456impl<Var, Exp> Monomial<Var, Exp>
457where
458 Var: Ord,
459{
460 pub fn generator(var: Var) -> Self
462 where
463 Exp: One,
464 {
465 Monomial([(var, Exp::one())].into_iter().collect())
466 }
467
468 pub fn len(&self) -> usize {
470 self.0.len()
471 }
472
473 pub fn is_empty(&self) -> bool {
475 self.0.is_empty()
476 }
477
478 pub fn variables(&self) -> impl ExactSizeIterator<Item = &Var> {
480 self.0.keys()
481 }
482
483 pub fn eval<A, F>(&self, mut f: F) -> A
485 where
486 A: Pow<Exp, Output = A> + Product,
487 F: FnMut(&Var) -> A,
488 Exp: Clone,
489 {
490 self.0.iter().map(|(var, exp)| f(var).pow(exp.clone())).product()
491 }
492
493 pub fn eval_with_order<A>(&self, values: impl IntoIterator<Item = A>) -> A
499 where
500 A: Pow<Exp, Output = A> + Product,
501 Exp: Clone,
502 {
503 let mut iter = values.into_iter();
504 let value = self.eval(|_| iter.next().expect("Should have enough values"));
505 assert!(iter.next().is_none(), "Too many values");
506 value
507 }
508
509 pub fn map_variables<NewVar, F>(&self, mut f: F) -> Monomial<NewVar, Exp>
515 where
516 Exp: Clone + Add<Output = Exp>,
517 NewVar: Ord,
518 F: FnMut(&Var) -> NewVar,
519 {
520 self.0.iter().map(|(var, exp)| (f(var), exp.clone())).collect()
521 }
522
523 pub fn normalize(self) -> Self
525 where
526 Exp: Zero,
527 {
528 self.into_iter().filter(|(_, exp)| !exp.is_zero()).collect()
529 }
530}
531
532impl<Var, Exp> FromIterator<(Var, Exp)> for Monomial<Var, Exp>
534where
535 Var: Ord,
536 Exp: Add<Output = Exp>,
537{
538 fn from_iter<T: IntoIterator<Item = (Var, Exp)>>(iter: T) -> Self {
539 let mut monomial = Monomial::default();
540 for rhs in iter {
541 monomial *= rhs;
542 }
543 monomial
544 }
545}
546
547impl<Var, Exp> IntoIterator for Monomial<Var, Exp> {
549 type Item = (Var, Exp);
550 type IntoIter = btree_map::IntoIter<Var, Exp>;
551
552 fn into_iter(self) -> Self::IntoIter {
553 self.0.into_iter()
554 }
555}
556
557impl<Var, Exp> Display for Monomial<Var, Exp>
559where
560 Var: Display,
561 Exp: Display + PartialEq + One,
562{
563 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
564 let mut pairs = self.0.iter();
565 let fmt_power = |f: &mut std::fmt::Formatter<'_>, var: &Var, exp: &Exp| {
566 write!(f, "{var}")?;
567 if !exp.is_one() {
568 let exp = exp.to_string();
569 if exp.len() == 1 {
570 write!(f, "^{exp}")?;
571 } else {
572 write!(f, "^{{{exp}}}")?;
573 }
574 }
575 Ok(())
576 };
577 let Some((var, exp)) = pairs.next() else {
578 return write!(f, "1");
579 };
580 fmt_power(f, var, exp)?;
581 for (var, exp) in pairs {
582 write!(f, " ")?;
583 fmt_power(f, var, exp)?;
584 }
585 Ok(())
586 }
587}
588
589impl<Var, Exp> Monomial<Var, Exp>
590where
591 Var: Display,
592 Exp: Display + PartialEq + One,
593{
594 pub fn to_latex(&self) -> String {
596 let fmt_power = |var: &Var, exp: &Exp| {
597 if exp.is_one() {
598 format!("{var}")
599 } else {
600 let exp = exp.to_string();
601 if exp.len() == 1 {
602 format!("{var}^{exp}")
603 } else {
604 format!("{var}^{{{exp}}}")
605 }
606 }
607 };
608 let mut pairs = self.0.iter();
609 let Some((var, exp)) = pairs.next() else {
610 return "1".to_string();
611 };
612 let mut output = fmt_power(var, exp);
613 for (var, exp) in pairs {
614 output.push_str(" \\cdot ");
615 output.push_str(&fmt_power(var, exp));
616 }
617 output
618 }
619}
620
621impl<Var, Exp> MulAssign<(Var, Exp)> for Monomial<Var, Exp>
622where
623 Var: Ord,
624 Exp: Add<Output = Exp>,
625{
626 fn mul_assign(&mut self, rhs: (Var, Exp)) {
627 _add_assign(&mut self.0, rhs);
628 }
629}
630
631impl<Var, Exp> MulAssign for Monomial<Var, Exp>
632where
633 Var: Ord,
634 Exp: Add<Output = Exp>,
635{
636 fn mul_assign(&mut self, rhs: Self) {
637 for rhs in rhs.0 {
638 *self *= rhs;
639 }
640 }
641}
642
643impl<Var, Exp> Mul for Monomial<Var, Exp>
644where
645 Var: Ord,
646 Exp: Add<Output = Exp>,
647{
648 type Output = Self;
649
650 fn mul(mut self, rhs: Self) -> Self {
651 self *= rhs;
652 self
653 }
654}
655
656impl<Var, Exp> One for Monomial<Var, Exp>
657where
658 Var: Ord,
659 Exp: Add<Output = Exp> + Zero,
660{
661 fn one() -> Self {
662 Monomial(Default::default())
663 }
664
665 fn is_one(&self) -> bool {
666 self.0.values().all(|exp| exp.is_zero())
667 }
668}
669
670impl<Var, Exp> Monoid for Monomial<Var, Exp>
671where
672 Var: Ord,
673 Exp: AdditiveMonoid,
674{
675}
676
677impl<Var, Exp> CommMonoid for Monomial<Var, Exp>
678where
679 Var: Ord,
680 Exp: AdditiveMonoid,
681{
682}
683
684impl<Var, Exp> Pow<Exp> for Monomial<Var, Exp>
685where
686 Var: Ord,
687 Exp: Clone + Default + Mul<Output = Exp>,
688{
689 type Output = Self;
690
691 fn pow(mut self, a: Exp) -> Self::Output {
692 for exp in self.0.values_mut() {
693 *exp = std::mem::take(exp) * a.clone();
694 }
695 self
696 }
697}
698
699#[cfg(test)]
700mod tests {
701 use super::*;
702
703 #[test]
704 fn combinations() {
705 let x = || Combination::<_, i32>::generator('x');
706 let y = || Combination::<_, i32>::generator('y');
707 assert_eq!(x().to_string(), "x");
708 assert_eq!((x() + y() + y() + x()).to_string(), "2 x + 2 y");
709
710 let combination = x() * 2 + y() * 3;
711 assert_eq!(combination.to_string(), "2 x + 3 y");
712 assert_eq!(combination.eval_with_order([5, 1]), 13);
713 let vars: Vec<_> = combination.variables().cloned().collect();
714 assert_eq!(vars, vec!['x', 'y']);
715
716 let combination = x() * 2 - y() * 3;
717 assert_eq!(combination.to_string(), "2 x - 3 y");
718
719 assert_eq!(Combination::<char, i32>::zero().to_string(), "0");
720
721 let x = Combination::generator('x');
722 assert_eq!((x.clone() * -1i32).to_string(), "-x");
723 assert_eq!(x.clone().neg().to_string(), "-x");
724
725 let combination = x.clone() + x.neg();
726 assert_ne!(combination, Combination::default());
727 assert_eq!(combination.normalize(), Combination::default());
728 }
729
730 #[test]
731 fn monomials() {
732 let x = || Monomial::<_, u32>::generator('x');
733 let y = || Monomial::<_, u32>::generator('y');
734 assert_eq!(x().to_string(), "x");
735 assert_eq!((x() * y() * y() * x()).to_string(), "x^2 y^2");
736
737 let monomial: Monomial<_, u32> = [('x', 1), ('y', 2)].into_iter().collect();
738 assert_eq!(monomial.to_string(), "x y^2");
739 assert_eq!(monomial.eval_with_order([10, 5]), 250);
740 let vars: Vec<_> = monomial.variables().cloned().collect();
741 assert_eq!(vars, vec!['x', 'y']);
742 assert_eq!(monomial.map_variables(|_| 'x').to_string(), "x^3");
743
744 let monomial: Monomial<char, u32> = Monomial::one();
745 assert_eq!(monomial.to_string(), "1");
746 assert_eq!(monomial.to_latex(), "1");
747
748 let monomial: Monomial<_, u32> = [('x', 1), ('y', 0), ('x', 2)].into_iter().collect();
749 assert_eq!(monomial.normalize().to_string(), "x^3");
750
751 let monomial: Monomial<_, i32> = [('x', -1), ('y', -2), ('x', 2)].into_iter().collect();
752 assert_eq!(monomial.normalize().to_string(), "x y^{-2}");
753
754 let monomial: Monomial<_, i32> = [('x', 1), ('y', 2), ('z', -1)].into_iter().collect();
755 assert_eq!(monomial.to_latex(), "x \\cdot y^2 \\cdot z^{-1}");
756 }
757}