Skip to main content

cas_expr/
assume.rs

1//! 假设系统(D5,对标 MATLAB `assume`,P1 后期)。
2//!
3//! - 闭谓词集(9 个):real / rational / integer / even / odd /
4//!   positive / negative / nonzero / finite,位集表示;
5//! - 构造时闭包补全(`even⇒integer⇒rational⇒real`、`positive⇒
6//!   real∧nonzero∧¬negative` 等)并检冲突(`positive∧negative` 为
7//!   编程错误);
8//! - 绑定:同一 Context 内同名符号**一次性设定**,重设不同值报错;
9//! - 查询:`query(e, P) → True | False | Unknown`,自底向上按头函数
10//!   事实表传播(`exp(x)>0`、`x^2≥0`(x real)等),遇 Unknown 截断;
11//! - 与化简的耦合只有一条路:规则 guard 调查询,仅 `True` 触发
12//!   (宁可少化简,不可错化简)。不做 SAT / 一阶逻辑引擎。
13
14use crate::Inner;
15use crate::node::Node;
16
17/// 闭谓词集(D5:P0 固定 9 个)。
18#[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/// 三值查询结果。
32#[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/// 假设位集(u16;位序与 `Predicate::bit` 一致)。恒存闭包后的形态。
46#[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
65/// "已绑定"标记位(区分未绑定与绑定空集:前者查询返回 Unknown)。
66const BOUND: u16 = 1 << 15;
67
68impl Assumptions {
69    pub fn none() -> Self {
70        Assumptions(0)
71    }
72
73    /// 单谓词构造(经闭包补全)。
74    pub fn with(p: Predicate) -> Self {
75        Assumptions(p.bit()).close()
76    }
77
78    /// 多谓词并集构造(经闭包补全 + 绑定标记)。
79    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    /// 是否有绑定(无绑定 ⇒ 一切查询 Unknown)。
90    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    /// 闭包补全 + 冲突检测。返回补全后的位集;冲突(positive∧negative)
99    /// 属编程错误,直接 panic(与符号名校验同级的契约)。
100    pub fn close(mut self) -> Self {
101        // even/odd ⇒ integer ⇒ rational ⇒ real
102        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        // positive/negative ⇒ real ∧ nonzero ∧ finite
112        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        // nonzero ∧ real:与零假设无冲突(零不是 nonzero)
119        // 冲突检测
120        assert!(
121            !(self.has(Predicate::Positive) && self.has(Predicate::Negative)),
122            "假设冲突:positive ∧ negative"
123        );
124        // even ∧ odd 冲突
125        assert!(
126            !(self.has(Predicate::Even) && self.has(Predicate::Odd)),
127            "假设冲突:even ∧ odd"
128        );
129        self
130    }
131
132    /// 非负(positive ∨ zero)——sqrt(x^2) 规则的判据。
133    pub fn nonnegative(&self) -> bool {
134        self.has(Predicate::Positive)
135    }
136}
137
138impl Inner {
139    /// 符号当前假设(无绑定返回空)。
140    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    /// 三值查询:自底向上传播,遇 Unknown 截断。
148    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                // 字面量谓词精确判定(符号/奇偶/整数性)
153                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, // 大数奇偶可判但 P1 保守
167                    },
168                }
169            }
170            Node::Rat(_) => {
171                // 有理数:rational/real/finite 恒真;符号看分子
172                match p {
173                    Predicate::Rational | Predicate::Real | Predicate::Finite => Trinary::True,
174                    _ => Trinary::Unknown, // 精确符号判定 P1 不做(避免大数分解)
175                }
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                // 闭包蕴含的否定:positive⇒¬negative、negative⇒¬positive、
190                // even⇒¬odd、odd⇒¬even
191                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                // 正性:全部 positive ⇒ positive;实性:全部 real ⇒ real
201                self.query_seq(self.node_args(*sp), p, depth, QueryOp::All)
202            }
203            Node::Mul { args: sp } => {
204                match p {
205                    // 实性:全部 real ⇒ real
206                    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                        // 符号 = 各因子符号的乘积:奇数个 negative 翻转
214                        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                // x^2(x real)非负;x^k(k 偶、x real)非负——sqrt(x^2)→x 判据
233                if p == Predicate::Positive {
234                    // exp(x) 型此处不进;x^k:k 正且 base positive ⇒ positive
235                    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, // x^0 = 1 仍 positive——保守 Unknown
241                            }
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                            // base real 且 exp 整数 ⇒ real;exp 实数 ⇒ 需 base>0
251                            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(x) > 0 恒成立;exp(x) real 当 x real
267                    ("exp", Predicate::Positive) => Trinary::True,
268                    ("exp", Predicate::Real | Predicate::Finite | Predicate::NonZero) => {
269                        Trinary::True
270                    }
271                    ("exp", _) => Trinary::Unknown,
272                    // sqrt(x):x≥0 ⇒ real;x>0 ⇒ positive
273                    ("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:real ⇒ real
291                    ("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(x):x>0 ⇒ real
299                    ("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}