添加四则运算工具

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-06-17 12:22:02 +08:00
co-authored by Claude
parent 2c4d4af070
commit 249227ef0a
5 changed files with 335 additions and 8 deletions
+259
View File
@@ -0,0 +1,259 @@
package calculator
import (
"encoding/json"
"errors"
"fmt"
"math"
"strconv"
"strings"
"unicode"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
)
const (
ToolName = "calculator"
ActivationPrompt = "执行简单、确定性的数学四则运算。当用户询问加减乘除、括号表达式、小数运算或需要准确计算表达式结果时,应直接调用此工具;不用于代数推导、方程求解、统计分析或复杂数学证明。"
)
type ToolArgs struct {
Expression string `json:"expression"`
Reason string `json:"reason"`
}
func ToolDefinition(description string) *model.Tool {
description = strings.TrimSpace(description)
if description == "" {
description = ActivationPrompt
}
return &model.Tool{
Type: model.ToolTypeFunction,
Function: &model.FunctionDefinition{
Name: ToolName,
Description: description,
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"expression": map[string]any{
"type": "string",
"description": "要计算的四则运算表达式,例如 12.5*(3+4)/2。仅支持数字、+、-、*、/ 和括号。",
},
"reason": map[string]any{
"type": "string",
"description": "调用计算器工具的原因。",
},
},
"required": []string{"expression"},
},
},
}
}
func ExecuteTool(args string) (string, error) {
var parsed ToolArgs
if err := json.Unmarshal([]byte(strings.TrimSpace(args)), &parsed); err != nil {
return "", fmt.Errorf("解析计算器工具参数失败: %w", err)
}
expression := strings.TrimSpace(parsed.Expression)
if expression == "" {
return "", errors.New("计算表达式不能为空")
}
result, err := Evaluate(expression)
if err != nil {
return "", err
}
return BuildResultContext(expression, result, parsed.Reason), nil
}
func BuildResultContext(expression string, result float64, routeReason string) string {
var b strings.Builder
b.WriteString("计算器工具结果。请优先使用这里的精确计算结果回答用户,不要重新心算。\n")
fmt.Fprintf(&b, "表达式: %s\n", strings.TrimSpace(expression))
fmt.Fprintf(&b, "结果: %s\n", FormatNumber(result))
if strings.TrimSpace(routeReason) != "" {
b.WriteString("调用原因: " + strings.TrimSpace(routeReason) + "\n")
}
return b.String()
}
func FormatNumber(value float64) string {
if math.IsInf(value, 0) || math.IsNaN(value) {
return fmt.Sprintf("%v", value)
}
return strconv.FormatFloat(value, 'f', -1, 64)
}
func Evaluate(expression string) (float64, error) {
parser := expressionParser{input: normalizeExpression(expression)}
value, err := parser.parseExpression()
if err != nil {
return 0, err
}
parser.skipSpaces()
if !parser.atEnd() {
return 0, fmt.Errorf("表达式包含不支持的字符: %q", parser.peek())
}
if math.IsInf(value, 0) || math.IsNaN(value) {
return 0, errors.New("计算结果无效")
}
return value, nil
}
func normalizeExpression(expression string) string {
replacer := strings.NewReplacer(
"×", "*",
"", "*",
"÷", "/",
"", "/",
"", "(",
"", ")",
"", "+",
"", "-",
"", ".",
)
return replacer.Replace(expression)
}
type expressionParser struct {
input string
pos int
}
func (p *expressionParser) parseExpression() (float64, error) {
value, err := p.parseTerm()
if err != nil {
return 0, err
}
for {
p.skipSpaces()
switch p.peek() {
case '+':
p.pos++
rhs, err := p.parseTerm()
if err != nil {
return 0, err
}
value += rhs
case '-':
p.pos++
rhs, err := p.parseTerm()
if err != nil {
return 0, err
}
value -= rhs
default:
return value, nil
}
}
}
func (p *expressionParser) parseTerm() (float64, error) {
value, err := p.parseFactor()
if err != nil {
return 0, err
}
for {
p.skipSpaces()
switch p.peek() {
case '*':
p.pos++
rhs, err := p.parseFactor()
if err != nil {
return 0, err
}
value *= rhs
case '/':
p.pos++
rhs, err := p.parseFactor()
if err != nil {
return 0, err
}
if rhs == 0 {
return 0, errors.New("除数不能为 0")
}
value /= rhs
default:
return value, nil
}
}
}
func (p *expressionParser) parseFactor() (float64, error) {
p.skipSpaces()
switch p.peek() {
case '+':
p.pos++
return p.parseFactor()
case '-':
p.pos++
value, err := p.parseFactor()
if err != nil {
return 0, err
}
return -value, nil
case '(':
p.pos++
value, err := p.parseExpression()
if err != nil {
return 0, err
}
p.skipSpaces()
if p.peek() != ')' {
return 0, errors.New("缺少右括号")
}
p.pos++
return value, nil
default:
return p.parseNumber()
}
}
func (p *expressionParser) parseNumber() (float64, error) {
p.skipSpaces()
start := p.pos
dotSeen := false
for !p.atEnd() {
r := p.peek()
if r == '.' {
if dotSeen {
break
}
dotSeen = true
p.pos++
continue
}
if !unicode.IsDigit(rune(r)) {
break
}
p.pos++
}
if start == p.pos {
if p.atEnd() {
return 0, errors.New("表达式不完整")
}
return 0, fmt.Errorf("期望数字,遇到 %q", p.peek())
}
value, err := strconv.ParseFloat(p.input[start:p.pos], 64)
if err != nil {
return 0, fmt.Errorf("解析数字失败: %w", err)
}
return value, nil
}
func (p *expressionParser) skipSpaces() {
for !p.atEnd() && unicode.IsSpace(rune(p.peek())) {
p.pos++
}
}
func (p *expressionParser) peek() byte {
if p.atEnd() {
return 0
}
return p.input[p.pos]
}
func (p *expressionParser) atEnd() bool {
return p.pos >= len(p.input)
}
+52
View File
@@ -0,0 +1,52 @@
package calculator
import (
"strings"
"testing"
)
func TestEvaluateBasicArithmetic(t *testing.T) {
tests := []struct {
expression string
want float64
}{
{expression: "1+2*3", want: 7},
{expression: "(1+2)*3", want: 9},
{expression: "12.5*(3+4)/2", want: 43.75},
{expression: "-2 + 3", want: 1},
{expression: "84)÷3", want: 4},
}
for _, tt := range tests {
got, err := Evaluate(tt.expression)
if err != nil {
t.Fatalf("Evaluate(%q) error: %v", tt.expression, err)
}
if got != tt.want {
t.Fatalf("Evaluate(%q) = %v, want %v", tt.expression, got, tt.want)
}
}
}
func TestEvaluateErrors(t *testing.T) {
for _, expression := range []string{"1/0", "1+", "(1+2", "2^3"} {
if _, err := Evaluate(expression); err == nil {
t.Fatalf("Evaluate(%q) expected error", expression)
}
}
}
func TestToolDefinitionAndExecuteTool(t *testing.T) {
definition := ToolDefinition("custom calculator")
if definition.Function == nil || definition.Function.Name != ToolName || definition.Function.Description != "custom calculator" {
t.Fatalf("unexpected definition: %#v", definition)
}
text, err := ExecuteTool(`{"expression":"12.5*(3+4)/2","reason":"用户询问计算结果"}`)
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"计算器工具结果", "12.5*(3+4)/2", "43.75", "用户询问计算结果"} {
if !strings.Contains(text, want) {
t.Fatalf("tool result missing %q:\n%s", want, text)
}
}
}