添加四则运算工具

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)
}
}
}
+2
View File
@@ -17,6 +17,7 @@ const (
defaultToolRouterTimeout = 30
defaultToolRouterMaxTokens = 512
defaultToolRouterSystemText = `你可以按需直接调用可用工具来回答用户问题。
用户询问简单数学计算、四则运算、加减乘除、括号表达式或明确要求算出数值结果时,调用 calculator。
如果用户问题包含今天、今日、明天、昨天、本周、本月、本年、最近等相对时间,且后续需要搜索或查询数据库,应先调用 time 获取绝对日期范围。
需要实时网页资料、新闻、当前版本、近期事件、网页核验或用户明确要求联网时,调用 search。
需要查询本地业务数据、日程、会议、待办、记录、统计或时间范围内数据时,调用 sql。
@@ -105,6 +106,7 @@ func DefaultToolRouterConfig() ToolRouterConfig {
MaxTokens: defaultToolRouterMaxTokens,
SystemPrompt: defaultToolRouterSystemText,
Tools: []ToolRouteConfig{
{Name: "calculator", Enabled: true, Description: ""},
{Name: "time", Enabled: true, Description: ""},
{Name: "search", Enabled: true, Description: ""},
{Name: "sql", Enabled: true, Description: ""},
+9 -8
View File
@@ -46,7 +46,7 @@ func TestNormalizeToolRouterConfigDefaults(t *testing.T) {
if strings.TrimSpace(cfg.ToolRouter.SystemPrompt) == "" {
t.Fatal("system prompt should be defaulted")
}
if len(cfg.ToolRouter.Tools) != 3 || cfg.ToolRouter.Tools[0].Name != "time" || cfg.ToolRouter.Tools[1].Name != "search" || cfg.ToolRouter.Tools[2].Name != "sql" || !cfg.ToolRouter.Tools[0].Enabled || !cfg.ToolRouter.Tools[1].Enabled || !cfg.ToolRouter.Tools[2].Enabled {
if len(cfg.ToolRouter.Tools) != 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[2].Name != "search" || cfg.ToolRouter.Tools[3].Name != "sql" || !cfg.ToolRouter.Tools[0].Enabled || !cfg.ToolRouter.Tools[1].Enabled || !cfg.ToolRouter.Tools[2].Enabled || !cfg.ToolRouter.Tools[3].Enabled {
t.Fatalf("unexpected tools: %#v", cfg.ToolRouter.Tools)
}
}
@@ -66,7 +66,7 @@ func TestNormalizeOpenAIConfigDefaultsContextWindow(t *testing.T) {
}
}
func TestNormalizeToolRouterConfigAddsTimeBeforeSQL(t *testing.T) {
func TestNormalizeToolRouterConfigAddsCalculatorAndTimeBeforeSQL(t *testing.T) {
cfg := &config.Config{ToolRouter: config.ToolRouterConfig{
Enabled: true,
Timeout: 1,
@@ -82,9 +82,9 @@ func TestNormalizeToolRouterConfigAddsTimeBeforeSQL(t *testing.T) {
t.Fatal(err)
}
if !changed {
t.Fatal("expected time tool to be added")
t.Fatal("expected calculator and time tools to be added")
}
if len(cfg.ToolRouter.Tools) < 3 || cfg.ToolRouter.Tools[0].Name != "time" || cfg.ToolRouter.Tools[2].Name != "sql" {
if len(cfg.ToolRouter.Tools) < 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[3].Name != "sql" {
t.Fatalf("unexpected tool order: %#v", cfg.ToolRouter.Tools)
}
}
@@ -112,6 +112,7 @@ func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) {
Enabled: true,
Tools: []config.ToolRouteConfig{
{Name: "search", Enabled: true},
{Name: "calculator", Enabled: true, Description: "custom calculator"},
{Name: "time", Enabled: true, Description: "custom time"},
{Name: "sql", Enabled: false},
},
@@ -121,14 +122,14 @@ func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) {
}
tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), nil, nil, nil)
if len(tools) != 1 {
if len(tools) != 2 {
t.Fatalf("tools length = %d", len(tools))
}
if tools[0].Name() != "time" {
t.Fatalf("tool name = %s", tools[0].Name())
if tools[0].Name() != "calculator" || tools[1].Name() != "time" {
t.Fatalf("unexpected tools: %#v", tools)
}
definition := tools[0].Definition()
if definition.Function == nil || definition.Function.Description != "custom time" {
if definition.Function == nil || definition.Function.Description != "custom calculator" {
t.Fatalf("unexpected definition: %#v", definition)
}
}
+13
View File
@@ -6,6 +6,7 @@ import (
"strings"
"time"
"aichat/agents/calculator"
searchagent "aichat/agents/search"
sqlquery "aichat/agents/sql"
timeagent "aichat/agents/time"
@@ -47,6 +48,18 @@ func availableAgentTools(state *State, profile *llm.Profile, searchState *search
}
description := strings.TrimSpace(item.Description)
switch item.Name {
case calculator.ToolName:
tools = append(tools, AgentTool{
name: calculator.ToolName,
definition: calculator.ToolDefinition(description),
execute: func(ctx context.Context, args string) (string, error) {
result, err := calculator.ExecuteTool(args)
if err == nil && emit != nil {
emit(stream.Frame{Type: "trace", Tool: calculator.ToolName, Stage: "calculate", Status: "success", Message: "四则运算完成"})
}
return result, err
},
})
case timeagent.ToolName:
tools = append(tools, AgentTool{
name: timeagent.ToolName,