1mod lexer;
17
18use cas_expr::{Context, Expr};
19use lexer::{Tok, Token};
20use std::error::Error;
21use std::fmt;
22
23#[derive(Clone, Debug, PartialEq, Eq)]
25pub struct ParseError {
26 pub pos: usize,
27 pub msg: String,
28}
29
30impl ParseError {
31 fn new(pos: usize, msg: &str) -> Self {
32 ParseError {
33 pos,
34 msg: msg.to_string(),
35 }
36 }
37}
38
39impl fmt::Display for ParseError {
40 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
41 write!(f, "位置 {}: {}", self.pos, self.msg)
42 }
43}
44
45impl Error for ParseError {}
46
47pub fn parse(ctx: &Context, src: &str) -> Result<Expr, ParseError> {
48 let toks = lexer::lex(src)?;
49 let mut p = Parser {
50 ctx,
51 toks,
52 idx: 0,
53 src_len: src.len(),
54 };
55 let e = p.expr()?;
56 match p.peek() {
57 Some(t) => Err(ParseError::new(t.pos, "存在多余记号")),
58 None => Ok(e),
59 }
60}
61
62struct Parser<'a> {
63 ctx: &'a Context,
64 toks: Vec<Token<'a>>,
65 idx: usize,
66 src_len: usize,
67}
68
69impl<'a> Parser<'a> {
70 fn peek(&self) -> Option<Token<'a>> {
71 self.toks.get(self.idx).cloned()
72 }
73
74 fn expr(&mut self) -> Result<Expr, ParseError> {
75 let mut l = self.term()?;
76 loop {
77 let plus = match self.peek().map(|t| t.tok) {
78 Some(Tok::Plus) => true,
79 Some(Tok::Minus) => false,
80 _ => break,
81 };
82 self.idx += 1;
83 let r = self.term()?;
84 l = if plus {
85 self.ctx.add(&[l, r])
86 } else {
87 let neg = self.ctx.mul(&[r, self.ctx.int(-1)]);
88 self.ctx.add(&[l, neg])
89 };
90 }
91 Ok(l)
92 }
93
94 fn term(&mut self) -> Result<Expr, ParseError> {
95 let mut l = self.unary()?;
96 loop {
97 let mul = match self.peek().map(|t| t.tok) {
98 Some(Tok::Star) => true,
99 Some(Tok::Slash) => false,
100 _ => break,
101 };
102 self.idx += 1;
103 let r = self.unary()?;
104 l = if mul {
105 self.ctx.mul(&[l, r])
106 } else {
107 let inv = self.ctx.pow(&r, &self.ctx.int(-1));
108 self.ctx.mul(&[l, inv])
109 };
110 }
111 Ok(l)
112 }
113
114 fn unary(&mut self) -> Result<Expr, ParseError> {
115 match self.peek().map(|t| t.tok) {
116 Some(Tok::Minus) => {
117 self.idx += 1;
118 let e = self.unary()?;
119 Ok(self.ctx.mul(&[e, self.ctx.int(-1)]))
120 }
121 Some(Tok::Plus) => {
122 self.idx += 1;
123 self.unary()
124 }
125 _ => self.postfix(),
126 }
127 }
128
129 fn postfix(&mut self) -> Result<Expr, ParseError> {
130 let a = self.atom()?;
131 if matches!(self.peek().map(|t| t.tok), Some(Tok::Caret)) {
132 self.idx += 1;
133 let e = self.unary()?;
134 Ok(self.ctx.pow(&a, &e))
135 } else {
136 Ok(a)
137 }
138 }
139
140 fn atom(&mut self) -> Result<Expr, ParseError> {
141 let Some(t) = self.peek() else {
142 return Err(ParseError::new(self.src_len, "期望表达式,遇到结尾"));
143 };
144 match t.tok {
145 Tok::Int(ref v) => {
146 self.idx += 1;
147 Ok(self.ctx.integer(v))
148 }
149 Tok::Float(v) => {
150 self.idx += 1;
151 Ok(self.ctx.float(v))
152 }
153 Tok::Ident(name) => {
154 self.idx += 1;
155 if matches!(self.peek().map(|t| t.tok), Some(Tok::LParen)) {
156 self.idx += 1;
157 let mut args = vec![self.expr()?];
158 while matches!(self.peek().map(|t| t.tok), Some(Tok::Comma)) {
159 self.idx += 1;
160 args.push(self.expr()?);
161 }
162 let close = self.peek();
163 match close.as_ref().map(|t| &t.tok) {
164 Some(Tok::RParen) => {
165 self.idx += 1;
166 Ok(self.ctx.call(name, &args))
167 }
168 _ => Err(ParseError::new(
169 close.as_ref().map(|t| t.pos).unwrap_or(self.src_len),
170 "期望 ')' 结束函数参数",
171 )),
172 }
173 } else {
174 Ok(self.ctx.sym(name))
175 }
176 }
177 Tok::LParen => {
178 self.idx += 1;
179 let e = self.expr()?;
180 match self.peek().map(|t| t.tok) {
181 Some(Tok::RParen) => {
182 self.idx += 1;
183 Ok(e)
184 }
185 _ => Err(ParseError::new(self.idx_at(), "期望 ')'")),
186 }
187 }
188 _ => Err(ParseError::new(t.pos, "期望表达式")),
189 }
190 }
191
192 fn idx_at(&self) -> usize {
193 self.toks
194 .get(self.idx)
195 .map(|t| t.pos)
196 .unwrap_or(self.src_len)
197 }
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203 use cas_print::plain;
204
205 fn round(src: &str, expect: &str) {
206 let ctx = Context::new();
207 let e = parse(&ctx, src).unwrap_or_else(|err| panic!("{src} 解析失败: {err}"));
208 assert_eq!(plain(&ctx, &e), expect, "输入: {src}");
209 }
210
211 #[test]
212 fn 优先级与结合性() {
213 round("1 + 2 * 3", "7");
214 round("x^2*y", "x^2*y");
215 round("-x^2", "-x^2");
216 round("x^-2", "x^-2");
217 round("2^-2", "1/4");
218 round("x/y/z", "x*y^-1*z^-1"); round("2^3^2", "512"); round("(x + y)*x", "x*(x + y)");
221 round("7/2", "7/2");
222 }
223
224 #[test]
225 fn 函数与符号() {
226 round("sin(x)", "sin(x)");
227 round("sin(x)^2", "sin(x)^2");
228 round("sin (x)", "sin(x)");
229 round("atan2(y, x)", "atan2(y, x)");
230 round("f(f(x))", "f(f(x))");
231 round("sin", "sin"); round("b_2 + a", "a + b_2");
233 }
234
235 #[test]
236 fn 大整数字面量() {
237 round(
238 "123456789012345678901234567890 + 1",
239 "123456789012345678901234567891",
240 );
241 }
242
243 #[test]
244 fn 错误报告() {
245 let ctx = Context::new();
246 assert!(parse(&ctx, "").is_err());
247 assert!(parse(&ctx, "2x").is_err()); assert!(parse(&ctx, "1e3").is_err()); assert!(parse(&ctx, "(x").is_err());
250 assert!(parse(&ctx, "x)").is_err());
251 assert!(parse(&ctx, "f()").is_err()); assert!(parse(&ctx, "1.").is_err());
253 assert!(parse(&ctx, "1.5e+").is_err());
254 assert!(parse(&ctx, "$").is_err());
255 let err = parse(&ctx, "x + $").unwrap_err();
257 assert_eq!(err.pos, 4);
258 }
259}