Files
aichat/agents/calculator/calculator.go
T
2026-06-17 12:22:02 +08:00

260 lines
5.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}