diff --git a/agents/calculator/calculator.go b/agents/calculator/calculator.go new file mode 100644 index 0000000..125564b --- /dev/null +++ b/agents/calculator/calculator.go @@ -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) +} diff --git a/agents/calculator/calculator_test.go b/agents/calculator/calculator_test.go new file mode 100644 index 0000000..f5f5e8b --- /dev/null +++ b/agents/calculator/calculator_test.go @@ -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: "(8+4)÷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) + } + } +} diff --git a/config/config.go b/config/config.go index 9356a35..66bb234 100644 --- a/config/config.go +++ b/config/config.go @@ -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: ""}, diff --git a/main_test.go b/main_test.go index 9cb7879..1a69333 100644 --- a/main_test.go +++ b/main_test.go @@ -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) } } diff --git a/toolrouter/tools.go b/toolrouter/tools.go index 2784dff..0c16733 100644 --- a/toolrouter/tools.go +++ b/toolrouter/tools.go @@ -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,