77 lines
2.1 KiB
Go
77 lines
2.1 KiB
Go
package llm
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"ai-agent-scaffold-go/internal/model"
|
|
)
|
|
|
|
// ChatModelAdapter 适配器,将 OpenAIClient 包装为 model.ChatModel
|
|
type ChatModelAdapter struct {
|
|
client *OpenAIClient
|
|
tools []model.Tool
|
|
}
|
|
|
|
// NewChatModelAdapter 创建 ChatModel 适配器
|
|
func NewChatModelAdapter(client *OpenAIClient, tools []model.Tool) *ChatModelAdapter {
|
|
return &ChatModelAdapter{client: client, tools: tools}
|
|
}
|
|
|
|
// Generate 实现 model.ChatModel 接口
|
|
func (m *ChatModelAdapter) Generate(ctx context.Context, messages []model.ChatMessage) (model.ChatReply, error) {
|
|
return m.client.Generate(ctx, messages, m.toolDefs())
|
|
}
|
|
|
|
// Stream 实现 model.ChatModel 接口
|
|
func (m *ChatModelAdapter) Stream(ctx context.Context, messages []model.ChatMessage) (<-chan model.ChatStreamEvent, <-chan error) {
|
|
return m.client.Stream(ctx, messages, m.toolDefs())
|
|
}
|
|
|
|
// Tools 返回注册的工具列表
|
|
func (m *ChatModelAdapter) Tools() []model.Tool {
|
|
return m.tools
|
|
}
|
|
|
|
// CallTool 根据名称和参数调用对应的工具
|
|
func (m *ChatModelAdapter) CallTool(ctx context.Context, name, arguments string) (string, error) {
|
|
query := extractQuery(arguments)
|
|
for _, t := range m.tools {
|
|
if t.Name() == name {
|
|
return t.Call(ctx, query)
|
|
}
|
|
}
|
|
return "", fmt.Errorf("tool %q not found", name)
|
|
}
|
|
|
|
// toolDefs 将 model.Tool 转换为 ToolDef 列表
|
|
func (m *ChatModelAdapter) toolDefs() []ToolDef {
|
|
defs := make([]ToolDef, 0, len(m.tools))
|
|
for _, t := range m.tools {
|
|
defs = append(defs, ToolDef{Name: t.Name(), Description: t.Description()})
|
|
}
|
|
return defs
|
|
}
|
|
|
|
// extractQuery 从工具调用参数 JSON 中提取 query 字段
|
|
func extractQuery(arguments string) string {
|
|
arguments = strings.TrimSpace(arguments)
|
|
if arguments == "" {
|
|
return ""
|
|
}
|
|
if idx := strings.Index(arguments, `"query"`); idx >= 0 {
|
|
rest := arguments[idx+7:]
|
|
if colon := strings.Index(rest, `:`); colon >= 0 {
|
|
rest = strings.TrimSpace(rest[colon+1:])
|
|
if strings.HasPrefix(rest, `"`) {
|
|
rest = rest[1:]
|
|
if end := strings.Index(rest, `"`); end >= 0 {
|
|
return rest[:end]
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return arguments
|
|
}
|