use crate::Inner;
use crate::node::Node;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Predicate {
Real,
Rational,
Integer,
Even,
Odd,
Positive,
Negative,
NonZero,
Finite,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Trinary {
True,
False,
Unknown,
}
impl From<bool> for Trinary {
fn from(b: bool) -> Self {
if b { Trinary::True } else { Trinary::False }
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Assumptions(u16);
impl Predicate {
pub fn bit(self) -> u16 {
match self {
Predicate::Real => 1 << 0,
Predicate::Rational => 1 << 1,
Predicate::Integer => 1 << 2,
Predicate::Even => 1 << 3,
Predicate::Odd => 1 << 4,
Predicate::Positive => 1 << 5,
Predicate::Negative => 1 << 6,
Predicate::NonZero => 1 << 7,
Predicate::Finite => 1 << 8,
}
}
}
const BOUND: u16 = 1 << 15;
impl Assumptions {
pub fn none() -> Self {
Assumptions(0)
}
pub fn with(p: Predicate) -> Self {
Assumptions(p.bit()).close()
}
pub fn union(preds: &[Predicate]) -> Self {
let mut bits = 0u16;
for &p in preds {
bits |= p.bit();
}
let mut a = Assumptions(bits).close();
a.0 |= BOUND;
a
}
pub fn bound(&self) -> bool {
self.0 & BOUND != 0
}
pub fn has(&self, p: Predicate) -> bool {
self.0 & p.bit() != 0
}
pub fn close(mut self) -> Self {
if self.has(Predicate::Even) || self.has(Predicate::Odd) {
self.0 |= Predicate::Integer.bit();
}
if self.has(Predicate::Integer) {
self.0 |= Predicate::Rational.bit();
}
if self.has(Predicate::Rational) {
self.0 |= Predicate::Real.bit();
}
for p in [Predicate::Positive, Predicate::Negative] {
if self.has(p) {
self.0 |=
Predicate::Real.bit() | Predicate::NonZero.bit() | Predicate::Finite.bit();
}
}
assert!(
!(self.has(Predicate::Positive) && self.has(Predicate::Negative)),
"假设冲突:positive ∧ negative"
);
assert!(
!(self.has(Predicate::Even) && self.has(Predicate::Odd)),
"假设冲突:even ∧ odd"
);
self
}
pub fn nonnegative(&self) -> bool {
self.has(Predicate::Positive)
}
}
impl Inner {
pub(crate) fn assumptions_of(&self, sym_id: u32) -> Assumptions {
*self
.sym_assumptions
.get(&sym_id)
.unwrap_or(&Assumptions::none())
}
pub(crate) fn query_at(&self, id: u32, p: Predicate, depth: u32) -> Trinary {
assert!(depth <= 10_000, "表达式嵌套过深");
match &self.nodes[id as usize] {
Node::Int(v) => {
match p {
Predicate::Integer
| Predicate::Rational
| Predicate::Real
| Predicate::Finite => Trinary::True,
Predicate::Positive => (!v.is_zero() && v.sign() > 0).into(),
Predicate::Negative => (v.sign() < 0).into(),
Predicate::NonZero => (!v.is_zero()).into(),
Predicate::Even | Predicate::Odd => match v.to_i64() {
Some(k) => {
let want_even = matches!(p, Predicate::Even);
(k.rem_euclid(2) == i64::from(!want_even)).into()
}
None => Trinary::Unknown, },
}
}
Node::Rat(_) => {
match p {
Predicate::Rational | Predicate::Real | Predicate::Finite => Trinary::True,
_ => Trinary::Unknown, }
}
Node::Float { .. } => match p {
Predicate::Real | Predicate::Finite => Trinary::True,
_ => Trinary::Unknown,
},
Node::Sym(s) => {
let a = self.assumptions_of(*s);
if !a.bound() {
return Trinary::Unknown;
}
if a.has(p) {
return Trinary::True;
}
match p {
Predicate::Negative if a.has(Predicate::Positive) => Trinary::False,
Predicate::Positive if a.has(Predicate::Negative) => Trinary::False,
Predicate::Odd if a.has(Predicate::Even) => Trinary::False,
Predicate::Even if a.has(Predicate::Odd) => Trinary::False,
_ => Trinary::Unknown,
}
}
Node::Add { args: sp } => {
self.query_seq(self.node_args(*sp), p, depth, QueryOp::All)
}
Node::Mul { args: sp } => {
match p {
Predicate::Real
| Predicate::Rational
| Predicate::Integer
| Predicate::Finite => {
self.query_seq(self.node_args(*sp), p, depth, QueryOp::All)
}
Predicate::Positive | Predicate::Negative => {
let mut neg = false;
for &a in self.node_args(*sp) {
match self.query_at(a, Predicate::Positive, depth + 1) {
Trinary::True => {}
Trinary::False => neg = !neg,
Trinary::Unknown => return Trinary::Unknown,
}
}
if p == Predicate::Positive {
(!neg).into()
} else {
neg.into()
}
}
_ => Trinary::Unknown,
}
}
Node::Pow { base, exp } => {
if p == Predicate::Positive {
match self.query_at(*base, Predicate::Positive, depth + 1) {
Trinary::True => {
match self.query_at(*exp, Predicate::Positive, depth + 1) {
Trinary::True => Trinary::True,
Trinary::Unknown => Trinary::Unknown,
Trinary::False => Trinary::Unknown, }
}
_ => Trinary::Unknown,
}
} else if p == Predicate::Real {
let b = self.query_at(*base, Predicate::Real, depth + 1);
let e = self.query_at(*exp, Predicate::Real, depth + 1);
match (b, e) {
(Trinary::True, Trinary::True) => {
match self.query_at(*exp, Predicate::Integer, depth + 1) {
Trinary::True => Trinary::True,
_ => Trinary::Unknown,
}
}
_ => Trinary::Unknown,
}
} else {
Trinary::Unknown
}
}
Node::Fn { head, args: sp } => {
let name = self.fn_names[*head as usize].as_ref();
let args = self.node_args(*sp);
match (name, p) {
("exp", Predicate::Positive) => Trinary::True,
("exp", Predicate::Real | Predicate::Finite | Predicate::NonZero) => {
Trinary::True
}
("exp", _) => Trinary::Unknown,
("sqrt", Predicate::Real) => {
if args.len() == 1 {
match self.query_at(args[0], Predicate::Positive, depth + 1) {
Trinary::True => Trinary::True,
_ => Trinary::Unknown,
}
} else {
Trinary::Unknown
}
}
("sqrt", Predicate::Positive) => {
if args.len() == 1 {
self.query_at(args[0], Predicate::Positive, depth + 1)
} else {
Trinary::Unknown
}
}
("sin" | "cos" | "tan", Predicate::Real) => {
if args.len() == 1 {
self.query_at(args[0], Predicate::Real, depth + 1)
} else {
Trinary::Unknown
}
}
("log", Predicate::Real) => {
if args.len() == 1 {
match self.query_at(args[0], Predicate::Positive, depth + 1) {
Trinary::True => Trinary::True,
_ => Trinary::Unknown,
}
} else {
Trinary::Unknown
}
}
_ => Trinary::Unknown,
}
}
}
}
fn query_seq(&self, ids: &[u32], p: Predicate, depth: u32, _op: QueryOp) -> Trinary {
let mut all = true;
for &a in ids {
match self.query_at(a, p, depth + 1) {
Trinary::True => {}
Trinary::False => return Trinary::False,
Trinary::Unknown => all = false,
}
}
if all { Trinary::True } else { Trinary::Unknown }
}
}
#[derive(Clone, Copy)]
enum QueryOp {
All,
}