1use 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#[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
60fn 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}