1use crate::Inner;
15use crate::node::Node;
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
19pub enum Predicate {
20 Real,
21 Rational,
22 Integer,
23 Even,
24 Odd,
25 Positive,
26 Negative,
27 NonZero,
28 Finite,
29}
30
31#[derive(Clone, Copy, Debug, PartialEq, Eq)]
33pub enum Trinary {
34 True,
35 False,
36 Unknown,
37}
38
39impl From<bool> for Trinary {
40 fn from(b: bool) -> Self {
41 if b { Trinary::True } else { Trinary::False }
42 }
43}
44
45#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
47pub struct Assumptions(u16);
48
49impl Predicate {
50 pub fn bit(self) -> u16 {
51 match self {
52 Predicate::Real => 1 << 0,
53 Predicate::Rational => 1 << 1,
54 Predicate::Integer => 1 << 2,
55 Predicate::Even => 1 << 3,
56 Predicate::Odd => 1 << 4,
57 Predicate::Positive => 1 << 5,
58 Predicate::Negative => 1 << 6,
59 Predicate::NonZero => 1 << 7,
60 Predicate::Finite => 1 << 8,
61 }
62 }
63}
64
65const BOUND: u16 = 1 << 15;
67
68impl Assumptions {
69 pub fn none() -> Self {
70 Assumptions(0)
71 }
72
73 pub fn with(p: Predicate) -> Self {
75 Assumptions(p.bit()).close()
76 }
77
78 pub fn union(preds: &[Predicate]) -> Self {
80 let mut bits = 0u16;
81 for &p in preds {
82 bits |= p.bit();
83 }
84 let mut a = Assumptions(bits).close();
85 a.0 |= BOUND;
86 a
87 }
88
89 pub fn bound(&self) -> bool {
91 self.0 & BOUND != 0
92 }
93
94 pub fn has(&self, p: Predicate) -> bool {
95 self.0 & p.bit() != 0
96 }
97
98 pub fn close(mut self) -> Self {
101 if self.has(Predicate::Even) || self.has(Predicate::Odd) {
103 self.0 |= Predicate::Integer.bit();
104 }
105 if self.has(Predicate::Integer) {
106 self.0 |= Predicate::Rational.bit();
107 }
108 if self.has(Predicate::Rational) {
109 self.0 |= Predicate::Real.bit();
110 }
111 for p in [Predicate::Positive, Predicate::Negative] {
113 if self.has(p) {
114 self.0 |=
115 Predicate::Real.bit() | Predicate::NonZero.bit() | Predicate::Finite.bit();
116 }
117 }
118 assert!(
121 !(self.has(Predicate::Positive) && self.has(Predicate::Negative)),
122 "假设冲突:positive ∧ negative"
123 );
124 assert!(
126 !(self.has(Predicate::Even) && self.has(Predicate::Odd)),
127 "假设冲突:even ∧ odd"
128 );
129 self
130 }
131
132 pub fn nonnegative(&self) -> bool {
134 self.has(Predicate::Positive)
135 }
136}
137
138impl Inner {
139 pub(crate) fn assumptions_of(&self, sym_id: u32) -> Assumptions {
141 *self
142 .sym_assumptions
143 .get(&sym_id)
144 .unwrap_or(&Assumptions::none())
145 }
146
147 pub(crate) fn query_at(&self, id: u32, p: Predicate, depth: u32) -> Trinary {
149 assert!(depth <= 10_000, "表达式嵌套过深");
150 match &self.nodes[id as usize] {
151 Node::Int(v) => {
152 match p {
154 Predicate::Integer
155 | Predicate::Rational
156 | Predicate::Real
157 | Predicate::Finite => Trinary::True,
158 Predicate::Positive => (!v.is_zero() && v.sign() > 0).into(),
159 Predicate::Negative => (v.sign() < 0).into(),
160 Predicate::NonZero => (!v.is_zero()).into(),
161 Predicate::Even | Predicate::Odd => match v.to_i64() {
162 Some(k) => {
163 let want_even = matches!(p, Predicate::Even);
164 (k.rem_euclid(2) == i64::from(!want_even)).into()
165 }
166 None => Trinary::Unknown, },
168 }
169 }
170 Node::Rat(_) => {
171 match p {
173 Predicate::Rational | Predicate::Real | Predicate::Finite => Trinary::True,
174 _ => Trinary::Unknown, }
176 }
177 Node::Float { .. } => match p {
178 Predicate::Real | Predicate::Finite => Trinary::True,
179 _ => Trinary::Unknown,
180 },
181 Node::Sym(s) => {
182 let a = self.assumptions_of(*s);
183 if !a.bound() {
184 return Trinary::Unknown;
185 }
186 if a.has(p) {
187 return Trinary::True;
188 }
189 match p {
192 Predicate::Negative if a.has(Predicate::Positive) => Trinary::False,
193 Predicate::Positive if a.has(Predicate::Negative) => Trinary::False,
194 Predicate::Odd if a.has(Predicate::Even) => Trinary::False,
195 Predicate::Even if a.has(Predicate::Odd) => Trinary::False,
196 _ => Trinary::Unknown,
197 }
198 }
199 Node::Add { args: sp } => {
200 self.query_seq(self.node_args(*sp), p, depth, QueryOp::All)
202 }
203 Node::Mul { args: sp } => {
204 match p {
205 Predicate::Real
207 | Predicate::Rational
208 | Predicate::Integer
209 | Predicate::Finite => {
210 self.query_seq(self.node_args(*sp), p, depth, QueryOp::All)
211 }
212 Predicate::Positive | Predicate::Negative => {
213 let mut neg = false;
215 for &a in self.node_args(*sp) {
216 match self.query_at(a, Predicate::Positive, depth + 1) {
217 Trinary::True => {}
218 Trinary::False => neg = !neg,
219 Trinary::Unknown => return Trinary::Unknown,
220 }
221 }
222 if p == Predicate::Positive {
223 (!neg).into()
224 } else {
225 neg.into()
226 }
227 }
228 _ => Trinary::Unknown,
229 }
230 }
231 Node::Pow { base, exp } => {
232 if p == Predicate::Positive {
234 match self.query_at(*base, Predicate::Positive, depth + 1) {
236 Trinary::True => {
237 match self.query_at(*exp, Predicate::Positive, depth + 1) {
238 Trinary::True => Trinary::True,
239 Trinary::Unknown => Trinary::Unknown,
240 Trinary::False => Trinary::Unknown, }
242 }
243 _ => Trinary::Unknown,
244 }
245 } else if p == Predicate::Real {
246 let b = self.query_at(*base, Predicate::Real, depth + 1);
247 let e = self.query_at(*exp, Predicate::Real, depth + 1);
248 match (b, e) {
249 (Trinary::True, Trinary::True) => {
250 match self.query_at(*exp, Predicate::Integer, depth + 1) {
252 Trinary::True => Trinary::True,
253 _ => Trinary::Unknown,
254 }
255 }
256 _ => Trinary::Unknown,
257 }
258 } else {
259 Trinary::Unknown
260 }
261 }
262 Node::Fn { head, args: sp } => {
263 let name = self.fn_names[*head as usize].as_ref();
264 let args = self.node_args(*sp);
265 match (name, p) {
266 ("exp", Predicate::Positive) => Trinary::True,
268 ("exp", Predicate::Real | Predicate::Finite | Predicate::NonZero) => {
269 Trinary::True
270 }
271 ("exp", _) => Trinary::Unknown,
272 ("sqrt", Predicate::Real) => {
274 if args.len() == 1 {
275 match self.query_at(args[0], Predicate::Positive, depth + 1) {
276 Trinary::True => Trinary::True,
277 _ => Trinary::Unknown,
278 }
279 } else {
280 Trinary::Unknown
281 }
282 }
283 ("sqrt", Predicate::Positive) => {
284 if args.len() == 1 {
285 self.query_at(args[0], Predicate::Positive, depth + 1)
286 } else {
287 Trinary::Unknown
288 }
289 }
290 ("sin" | "cos" | "tan", Predicate::Real) => {
292 if args.len() == 1 {
293 self.query_at(args[0], Predicate::Real, depth + 1)
294 } else {
295 Trinary::Unknown
296 }
297 }
298 ("log", Predicate::Real) => {
300 if args.len() == 1 {
301 match self.query_at(args[0], Predicate::Positive, depth + 1) {
302 Trinary::True => Trinary::True,
303 _ => Trinary::Unknown,
304 }
305 } else {
306 Trinary::Unknown
307 }
308 }
309 _ => Trinary::Unknown,
310 }
311 }
312 }
313 }
314
315 fn query_seq(&self, ids: &[u32], p: Predicate, depth: u32, _op: QueryOp) -> Trinary {
316 let mut all = true;
317 for &a in ids {
318 match self.query_at(a, p, depth + 1) {
319 Trinary::True => {}
320 Trinary::False => return Trinary::False,
321 Trinary::Unknown => all = false,
322 }
323 }
324 if all { Trinary::True } else { Trinary::Unknown }
325 }
326}
327
328#[derive(Clone, Copy)]
329enum QueryOp {
330 All,
331}