Skip to main content

machi_tools/
calc.rs

1//! Simple arithmetic evaluation tool for demos and tests.
2
3use async_trait::async_trait;
4use serde_json::{Value, json};
5
6use crate::context::ToolCallContext;
7use crate::error::{ToolError, codes};
8use crate::metadata::ToolMetadata;
9use crate::tool::{DynTool, ToolResult};
10
11/// Evaluates a restricted arithmetic expression (`+ - * / ( )` and numbers).
12#[derive(Debug, Default, Clone, Copy)]
13pub struct CalcTool;
14
15#[async_trait]
16impl DynTool for CalcTool {
17    fn name(&self) -> &'static str {
18        "calc"
19    }
20
21    fn description(&self) -> &'static str {
22        "Evaluate a basic arithmetic expression with + - * / and parentheses. \
23         Example: expr=\"(2+3)*4\""
24    }
25
26    fn parameters(&self) -> Value {
27        json!({
28            "type": "object",
29            "properties": {
30                "expr": {
31                    "type": "string",
32                    "description": "Arithmetic expression to evaluate"
33                }
34            },
35            "required": ["expr"],
36            "additionalProperties": false
37        })
38    }
39
40    fn metadata(&self) -> ToolMetadata {
41        ToolMetadata::read_only()
42    }
43
44    async fn call(&self, _ctx: ToolCallContext, arguments: Value) -> Result<ToolResult, ToolError> {
45        let expr = arguments
46            .get("expr")
47            .and_then(Value::as_str)
48            .map(str::trim)
49            .filter(|s| !s.is_empty())
50            .ok_or_else(|| codes::invalid_args("calc requires non-empty expr"))?;
51        let value = eval_expr(expr).map_err(codes::execution)?;
52        Ok(ToolResult {
53            content: value.to_string(),
54            structured: Some(json!({ "expr": expr, "value": value })),
55            is_error: false,
56        })
57    }
58}
59
60/// Recursive-descent evaluator for numbers and + - * / ( ).
61fn eval_expr(input: &str) -> Result<f64, String> {
62    let tokens = tokenize(input)?;
63    let mut p = Parser { tokens, i: 0 };
64    let v = p.parse_expr()?;
65    if p.i != p.tokens.len() {
66        return Err("unexpected trailing tokens".into());
67    }
68    Ok(v)
69}
70
71#[derive(Debug, Clone, PartialEq)]
72enum Tok {
73    Num(f64),
74    Op(char),
75    LParen,
76    RParen,
77}
78
79fn tokenize(s: &str) -> Result<Vec<Tok>, String> {
80    let mut out = Vec::new();
81    let chars: Vec<char> = s.chars().collect();
82    let mut i = 0usize;
83    while i < chars.len() {
84        let Some(&c) = chars.get(i) else {
85            break;
86        };
87        if c.is_whitespace() {
88            i = i.saturating_add(1);
89            continue;
90        }
91        if c.is_ascii_digit() || c == '.' {
92            let start = i;
93            i = i.saturating_add(1);
94            while i < chars.len()
95                && chars
96                    .get(i)
97                    .is_some_and(|ch| ch.is_ascii_digit() || *ch == '.')
98            {
99                i = i.saturating_add(1);
100            }
101            let slice: String = chars
102                .get(start..i)
103                .ok_or_else(|| "bad number slice".to_owned())?
104                .iter()
105                .collect();
106            let n: f64 = slice
107                .parse()
108                .map_err(|_| format!("invalid number: {slice}"))?;
109            out.push(Tok::Num(n));
110            continue;
111        }
112        match c {
113            '+' | '-' | '*' | '/' => {
114                out.push(Tok::Op(c));
115                i = i.saturating_add(1);
116            }
117            '(' => {
118                out.push(Tok::LParen);
119                i = i.saturating_add(1);
120            }
121            ')' => {
122                out.push(Tok::RParen);
123                i = i.saturating_add(1);
124            }
125            other => return Err(format!("invalid character: {other}")),
126        }
127    }
128    Ok(out)
129}
130
131struct Parser {
132    tokens: Vec<Tok>,
133    i: usize,
134}
135
136impl Parser {
137    fn peek(&self) -> Option<&Tok> {
138        self.tokens.get(self.i)
139    }
140
141    fn bump(&mut self) -> Option<Tok> {
142        let t = self.tokens.get(self.i).cloned();
143        if t.is_some() {
144            self.i = self.i.saturating_add(1);
145        }
146        t
147    }
148
149    fn parse_expr(&mut self) -> Result<f64, String> {
150        let mut v = self.parse_term()?;
151        loop {
152            match self.peek() {
153                Some(Tok::Op('+')) => {
154                    self.bump();
155                    v += self.parse_term()?;
156                }
157                Some(Tok::Op('-')) => {
158                    self.bump();
159                    v -= self.parse_term()?;
160                }
161                _ => break,
162            }
163        }
164        Ok(v)
165    }
166
167    fn parse_term(&mut self) -> Result<f64, String> {
168        let mut v = self.parse_factor()?;
169        loop {
170            match self.peek() {
171                Some(Tok::Op('*')) => {
172                    self.bump();
173                    v *= self.parse_factor()?;
174                }
175                Some(Tok::Op('/')) => {
176                    self.bump();
177                    let d = self.parse_factor()?;
178                    if d == 0.0 {
179                        return Err("division by zero".into());
180                    }
181                    v /= d;
182                }
183                _ => break,
184            }
185        }
186        Ok(v)
187    }
188
189    fn parse_factor(&mut self) -> Result<f64, String> {
190        match self.bump() {
191            Some(Tok::Num(n)) => Ok(n),
192            Some(Tok::Op('-')) => Ok(-self.parse_factor()?),
193            Some(Tok::Op('+')) => self.parse_factor(),
194            Some(Tok::LParen) => {
195                let v = self.parse_expr()?;
196                match self.bump() {
197                    Some(Tok::RParen) => Ok(v),
198                    _ => Err("expected ')'".into()),
199                }
200            }
201            other => Err(format!("unexpected token: {other:?}")),
202        }
203    }
204}
205
206#[cfg(test)]
207mod tests {
208    use super::*;
209
210    #[tokio::test]
211    async fn evaluates_expression() {
212        let tool = CalcTool;
213        let result = tool
214            .call(ToolCallContext::default(), json!({"expr": "(2+3)*4"}))
215            .await
216            .expect("calc");
217        assert_eq!(result.content, "20");
218        assert_eq!(
219            result.structured.as_ref().and_then(|v| v.get("value")),
220            Some(&json!(20.0))
221        );
222    }
223
224    #[test]
225    fn rejects_letters() {
226        assert!(eval_expr("1+foo").is_err());
227    }
228}