use std::fmt;
use std::ops::BitOr;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Assumption {
Real,
Complex,
Integer,
Rational,
Positive,
Negative,
NonNegative,
NonPositive,
NonZero,
Finite,
Even,
Odd,
Prime,
}
impl fmt::Display for Assumption {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Assumption::Real => f.write_str("real"),
Assumption::Complex => f.write_str("complex"),
Assumption::Integer => f.write_str("integer"),
Assumption::Rational => f.write_str("rational"),
Assumption::Positive => f.write_str("positive"),
Assumption::Negative => f.write_str("negative"),
Assumption::NonNegative => f.write_str("non-negative"),
Assumption::NonPositive => f.write_str("non-positive"),
Assumption::NonZero => f.write_str("non-zero"),
Assumption::Finite => f.write_str("finite"),
Assumption::Even => f.write_str("even"),
Assumption::Odd => f.write_str("odd"),
Assumption::Prime => f.write_str("prime"),
}
}
}
impl Assumption {
pub fn implied(&self) -> &'static [Assumption] {
match self {
Assumption::Real => &[],
Assumption::Complex => &[Assumption::Real],
Assumption::Integer => &[Assumption::Rational, Assumption::Real],
Assumption::Rational => &[Assumption::Real],
Assumption::Positive => &[
Assumption::NonNegative,
Assumption::NonZero,
Assumption::Real,
],
Assumption::Negative => &[
Assumption::NonPositive,
Assumption::NonZero,
Assumption::Real,
],
Assumption::NonNegative => &[Assumption::Real],
Assumption::NonPositive => &[Assumption::Real],
Assumption::NonZero => &[],
Assumption::Finite => &[],
Assumption::Even => &[Assumption::Integer],
Assumption::Odd => &[Assumption::Integer],
Assumption::Prime => &[Assumption::Integer, Assumption::Positive],
}
}
pub fn conflicts(&self) -> &'static [Assumption] {
match self {
Assumption::Real => &[],
Assumption::Complex => &[],
Assumption::Integer => &[],
Assumption::Rational => &[],
Assumption::Positive => &[Assumption::Negative, Assumption::NonPositive],
Assumption::Negative => &[Assumption::Positive, Assumption::NonNegative],
Assumption::NonNegative => &[Assumption::Negative],
Assumption::NonPositive => &[Assumption::Positive],
Assumption::NonZero => &[],
Assumption::Finite => &[],
Assumption::Even => &[Assumption::Odd],
Assumption::Odd => &[Assumption::Even],
Assumption::Prime => &[],
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct Assumptions {
inner: Vec<Assumption>,
}
impl Assumptions {
pub fn new() -> Self {
Self { inner: Vec::new() }
}
pub fn single(a: Assumption) -> Self {
let mut s = Self::new();
s.insert(a);
s
}
pub fn len(&self) -> usize {
self.inner.len()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
pub fn contains(&self, a: Assumption) -> bool {
self.inner.contains(&a)
}
pub fn implies(&self, other: Assumption) -> bool {
if self.contains(other) {
return true;
}
for &a in &self.inner {
if a.implied().contains(&other) {
return true;
}
}
false
}
pub fn insert(&mut self, a: Assumption) -> bool {
if self.inner.contains(&a) {
return !self.conflicts_with(a);
}
self.inner.push(a);
let implied: Vec<Assumption> = a.implied().to_vec();
for imp in implied {
if !self.inner.contains(&imp) {
self.inner.push(imp);
}
}
self.inner.sort_unstable_by_key(|a| *a as u8);
self.inner.dedup();
!self.conflicts_with(a)
}
pub fn remove(&mut self, a: Assumption) {
self.inner.retain(|&x| x != a);
}
pub fn is_consistent(&self) -> bool {
for &a in &self.inner {
if self.conflicts_with(a) {
return false;
}
}
true
}
pub fn iter(&self) -> impl Iterator<Item = Assumption> + '_ {
self.inner.iter().copied()
}
fn conflicts_with(&self, a: Assumption) -> bool {
for &conflict in a.conflicts() {
if self.inner.contains(&conflict) {
return true;
}
}
false
}
}
impl BitOr<Assumption> for Assumption {
type Output = Assumptions;
fn bitor(self, rhs: Assumption) -> Assumptions {
let mut s = Assumptions::single(self);
s.insert(rhs);
s
}
}
impl BitOr<Assumption> for Assumptions {
type Output = Assumptions;
fn bitor(mut self, rhs: Assumption) -> Assumptions {
self.insert(rhs);
self
}
}
impl BitOr<Assumptions> for Assumptions {
type Output = Assumptions;
fn bitor(mut self, rhs: Assumptions) -> Assumptions {
for a in rhs.inner {
self.insert(a);
}
self
}
}
impl FromIterator<Assumption> for Assumptions {
fn from_iter<I: IntoIterator<Item = Assumption>>(iter: I) -> Self {
let mut s = Self::new();
for a in iter {
s.insert(a);
}
s
}
}
impl fmt::Display for Assumptions {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.is_empty() {
return f.write_str("(none)");
}
let parts: Vec<String> = self.inner.iter().map(|a| a.to_string()).collect();
f.write_str(&parts.join(", "))
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct SymbolAssumptions {
entries: Vec<(String, Assumptions)>,
}
impl SymbolAssumptions {
pub fn new() -> Self {
Self {
entries: Vec::new(),
}
}
pub fn set(&mut self, symbol: &str, assumptions: Assumptions) {
match self
.entries
.binary_search_by(|(s, _)| s.as_str().cmp(symbol))
{
Ok(idx) => {
self.entries[idx].1 = assumptions;
}
Err(idx) => {
self.entries.insert(idx, (symbol.to_owned(), assumptions));
}
}
}
pub fn get(&self, symbol: &str) -> Option<&Assumptions> {
match self
.entries
.binary_search_by(|(s, _)| s.as_str().cmp(symbol))
{
Ok(idx) => Some(&self.entries[idx].1),
Err(_) => None,
}
}
pub fn remove(&mut self, symbol: &str) {
if let Ok(idx) = self
.entries
.binary_search_by(|(s, _)| s.as_str().cmp(symbol))
{
self.entries.remove(idx);
}
}
pub fn check(&self, symbol: &str, assumption: Assumption) -> bool {
self.get(symbol)
.map(|a| a.implies(assumption))
.unwrap_or(false)
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = &(String, Assumptions)> {
self.entries.iter()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_assumptions() {
let a = Assumptions::new();
assert!(a.is_empty());
assert!(a.is_consistent());
assert!(!a.implies(Assumption::Real));
}
#[test]
fn positive_implies_real_and_nonzero() {
let a = Assumptions::single(Assumption::Positive);
assert!(a.implies(Assumption::Real));
assert!(a.implies(Assumption::NonNegative));
assert!(a.implies(Assumption::NonZero));
assert!(!a.implies(Assumption::Integer));
}
#[test]
fn integer_implies_rational_and_real() {
let a = Assumptions::single(Assumption::Integer);
assert!(a.implies(Assumption::Rational));
assert!(a.implies(Assumption::Real));
}
#[test]
fn complex_implies_real() {
let a = Assumptions::single(Assumption::Complex);
assert!(a.implies(Assumption::Real));
}
#[test]
fn positive_and_integer() {
let mut a = Assumptions::new();
a.insert(Assumption::Positive);
a.insert(Assumption::Integer);
assert!(a.implies(Assumption::Real));
assert!(a.implies(Assumption::NonZero));
assert!(a.implies(Assumption::Rational));
}
#[test]
fn conflict_positive_negative() {
let mut a = Assumptions::new();
a.insert(Assumption::Positive);
a.insert(Assumption::Negative);
assert!(!a.is_consistent());
}
#[test]
fn conflict_even_odd() {
let mut a = Assumptions::new();
a.insert(Assumption::Even);
a.insert(Assumption::Odd);
assert!(!a.is_consistent());
}
#[test]
fn no_conflict_real_integer() {
let mut a = Assumptions::new();
a.insert(Assumption::Real);
a.insert(Assumption::Integer);
assert!(a.is_consistent());
}
#[test]
fn bitor_operator() {
let a = Assumption::Positive | Assumption::Integer;
assert!(a.implies(Assumption::Real));
assert!(a.implies(Assumption::Rational));
assert!(a.implies(Assumption::NonZero));
}
#[test]
fn symbol_assumptions_basics() {
let mut sa = SymbolAssumptions::new();
sa.set("x", Assumptions::single(Assumption::Positive));
assert!(sa.check("x", Assumption::Real));
assert!(sa.check("x", Assumption::NonNegative));
assert!(!sa.check("x", Assumption::Integer));
assert!(!sa.check("y", Assumption::Real));
}
#[test]
fn symbol_assumptions_override() {
let mut sa = SymbolAssumptions::new();
sa.set("x", Assumptions::single(Assumption::Positive));
sa.set("x", Assumptions::single(Assumption::Integer));
assert!(sa.check("x", Assumption::Integer));
assert!(!sa.check("x", Assumption::Positive));
}
#[test]
fn prime_implies_integer_and_positive() {
let a = Assumptions::single(Assumption::Prime);
assert!(a.implies(Assumption::Integer));
assert!(a.implies(Assumption::Positive));
assert!(a.implies(Assumption::Real));
}
}