@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user