refactor: 重构llm包

This commit is contained in:
hhs
2026-04-20 15:34:40 +08:00
parent fe87f5ca80
commit 181edf418e
10 changed files with 1181 additions and 299 deletions
+14 -256
View File
@@ -2,253 +2,35 @@ package llm
import (
"context"
"encoding/json"
"fmt"
"knowledge-graph-backend/internal/model"
"log"
"github.com/tmc/langchaingo/llms"
)
// KnowledgeGraphToolService 定义 LLM 客户端需要的图谱操作接口
// service.Neo4jService 会隐式实现该接口,避免循环依赖
type KnowledgeGraphToolService interface {
SearchNodes(query string) []model.Node
GetNeighbors(nodeID string) (model.NeighborResponse, bool)
GetNodeByID(id string) (model.Node, bool)
CreateNode(req model.CreateNodeRequest) (model.Node, error)
CreateEdge(req model.CreateEdgeRequest) (model.Edge, error)
DeleteNode(id string) error
DeleteEdge(id string) error
}
// Agent 封装了带有工具调用和上下文管理能力的智能体
type Agent struct {
llm llms.Model
graphSvc KnowledgeGraphToolService
client *Client
executor *ToolExecutor
}
// NewAgent 创建一个新的 Agent 实例
func NewAgent(llm llms.Model, graphSvc KnowledgeGraphToolService) *Agent {
func NewAgent(client *Client, executor *ToolExecutor) *Agent {
return &Agent{
llm: llm,
graphSvc: graphSvc,
client: client,
executor: executor,
}
}
// getKnowledgeGraphTools 声明 LLM 可以调用的工具列表及其 JSON Schema
// (同原代码,此处省略重复注释)
func getKnowledgeGraphTools() []llms.Tool {
return []llms.Tool{
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "search_nodes",
Description: "当需要根据关键词在知识图谱中搜索相关节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "搜索的关键词",
},
},
"required": []string{"query"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "get_neighbors",
Description: "当需要查找某个节点的相邻节点(即与该节点有直接连线的节点)时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"node_id": map[string]any{
"type": "string",
"description": "目标节点的唯一 ID",
},
},
"required": []string{"node_id"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "create_node",
Description: "当需要在知识图谱中创建一个新的节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{"type": "string", "description": "节点的唯一ID,格式通常为 node_xxx"},
"label": map[string]any{"type": "string", "description": "节点的名称/标签,例如:人工智能"},
"type": map[string]any{"type": "string", "description": "节点的类型,例如:概念、实体"},
"x": map[string]any{"type": "number", "description": "节点在画布上的 X 坐标 (可选,默认0)"},
"y": map[string]any{"type": "number", "description": "节点在画布上的 Y 坐标 (可选,默认0)"},
"properties": map[string]any{
"type": "object",
"description": "节点的额外属性键值对 (可选)",
},
},
"required": []string{"id", "label", "type"},
},
},
},
{
Type: "Edge operation",
Function: &llms.FunctionDefinition{
Name: "create_edge",
Description: "当需要在知识图谱中创建一条连线(边)连接两个节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{"type": "string", "description": "边的唯一ID,格式通常为 edge_xxx"},
"source": map[string]any{"type": "string", "description": "起始节点的 ID"},
"target": map[string]any{"type": "string", "description": "目标节点的 ID"},
"label": map[string]any{"type": "string", "description": "连线的名称/标签,例如:包含、依赖于"},
"type": map[string]any{"type": "string", "description": "连线的类型 (可选),例如:CONTAINS"},
"properties": map[string]any{
"type": "object",
"description": "边的额外属性键值对 (可选)",
},
},
"required": []string{"id", "source", "target", "label"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "delete_node",
Description: "当需要从知识图谱中删除一个现有节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{
"type": "string",
"description": "要删除的节点的唯一 ID",
},
},
"required": []string{"id"},
},
},
},
{
Type: "Edge operation",
Function: &llms.FunctionDefinition{
Name: "delete_edge",
Description: "当需要从知识图谱中删除一条连线(边)时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{
"type": "string",
"description": "要删除的边的唯一 ID",
},
},
"required": []string{"id"},
},
},
},
}
}
// executeToolCall 执行具体的工具调用逻辑
// (同原代码,此处省略重复注释)
func executeToolCall(toolCall llms.ToolCall, graphSvc KnowledgeGraphToolService) (string, error) {
var result any
var err error
switch toolCall.FunctionCall.Name {
case "search_nodes":
var args struct {
Query string `json:"query"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for search_nodes: %w", err)
}
result = graphSvc.SearchNodes(args.Query)
case "get_neighbors":
var args struct {
NodeID string `json:"node_id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for get_neighbors: %w", err)
}
result, _ = graphSvc.GetNeighbors(args.NodeID)
case "create_node":
var args model.CreateNodeRequest
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for create_node: %w", err)
}
result, err = graphSvc.CreateNode(args)
if err != nil {
return "", err
}
case "create_edge":
var args model.CreateEdgeRequest
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for create_edge: %w", err)
}
result, err = graphSvc.CreateEdge(args)
if err != nil {
return "", err
}
case "delete_node":
var args struct {
ID string `json:"id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for delete_node: %w", err)
}
err = graphSvc.DeleteNode(args.ID)
if err != nil {
return "", err
}
result = map[string]string{"status": "success", "message": fmt.Sprintf("Node %s deleted successfully", args.ID)}
case "delete_edge":
var args struct {
ID string `json:"id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for delete_edge: %w", err)
}
err = graphSvc.DeleteEdge(args.ID)
if err != nil {
return "", err
}
result = map[string]string{"status": "success", "message": fmt.Sprintf("Edge %s deleted successfully", args.ID)}
default:
return "", fmt.Errorf("unknown tool: %s", toolCall.FunctionCall.Name)
}
resBytes, _ := json.Marshal(result)
return string(resBytes), nil
}
// ChatWithTools 带有工具调用能力的 Agent 入口
// 参考了 Client.go 的上下文管理模式,加入系统提示词及 Agent Loop 循环调用
func (a *Agent) ChatWithTools(ctx context.Context, systemPrompt string, messages []llms.MessageContent) (string, error) {
tools := getKnowledgeGraphTools()
tools := GetKnowledgeGraphTools()
// 1. 构造系统提示词消息,并插入到消息列表头部
systemMessage := llms.TextParts(llms.ChatMessageTypeSystem, systemPrompt)
allMessages := append([]llms.MessageContent{systemMessage}, messages...)
builder := NewMessageBuilder().
WithSystemPrompt(systemPrompt).
WithMessages(messages)
// 设定最大循环次数,防止死循环
maxIterations := 10
for i := 0; i < maxIterations; i++ {
// 2. 调用大模型,携带上下文和可用工具
resp, err := a.llm.GenerateContent(ctx, allMessages, llms.WithTools(tools))
resp, err := a.client.Generate(ctx, builder.Build(), llms.WithTools(tools))
if err != nil {
log.Printf("[ERROR] Agent GenerateContent failed: %v", err)
return "", fmt.Errorf("agent generate content error: %w", err)
@@ -261,51 +43,27 @@ func (a *Agent) ChatWithTools(ctx context.Context, systemPrompt string, messages
choice := resp.Choices[0]
// 3. 如果没有工具调用,说明大模型已经生成了最终回复,退出循环
if len(choice.ToolCalls) == 0 {
log.Printf("[INFO] Agent chat completed successfully, choice length: %d", len(choice.Content))
return choice.Content, nil
}
// 4. 处理工具调用:需要将 AI 的工具调用指令追加到上下文中
aiMessageParts := []llms.ContentPart{}
if choice.Content != "" {
aiMessageParts = append(aiMessageParts, llms.TextPart(choice.Content))
}
for _, tc := range choice.ToolCalls {
aiMessageParts = append(aiMessageParts, tc)
}
allMessages = append(allMessages, llms.MessageContent{
Role: llms.ChatMessageTypeAI,
Parts: aiMessageParts,
})
builder.AppendAIMessage(choice.Content, choice.ToolCalls)
// 5. 依次执行每个工具调用,并将结果作为 Tool 消息追加到上下文
for _, tc := range choice.ToolCalls {
log.Printf("[INFO] Executing tool call: %s, Args: %s", tc.FunctionCall.Name, tc.FunctionCall.Arguments)
toolResult, err := executeToolCall(tc, a.graphSvc)
toolResult, err := a.executor.Execute(tc)
if err != nil {
log.Printf("[WARN] Tool execution failed: %v", err)
// 即使工具执行失败,也将错误信息返回给 LLM,让其进行自我修正或回复用户
toolResult = fmt.Sprintf(`{"error": "%s"}`, err.Error())
}
log.Printf("[INFO] Tool call %s completed, result length: %d", tc.FunctionCall.Name, len(toolResult))
// 构造工具调用结果并追加到上下文历史
allMessages = append(allMessages, llms.MessageContent{
Role: llms.ChatMessageTypeTool,
Parts: []llms.ContentPart{
llms.ToolCallResponse{
ToolCallID: tc.ID,
Name: tc.FunctionCall.Name,
Content: toolResult,
},
},
})
builder.AppendToolResult(tc.ID, tc.FunctionCall.Name, toolResult)
}
}
return "", fmt.Errorf("agent reached maximum tool call iterations (%d)", maxIterations)
}
}
+264
View File
@@ -0,0 +1,264 @@
package llm
import (
"context"
"fmt"
"knowledge-graph-backend/internal/model"
"testing"
"github.com/tmc/langchaingo/llms"
)
type mockModel struct {
generateContent func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error)
}
func (m *mockModel) Call(ctx context.Context, prompt string, opts ...llms.CallOption) (string, error) {
return "", nil
}
func (m *mockModel) GenerateContent(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
if m.generateContent != nil {
return m.generateContent(ctx, messages, opts...)
}
return nil, nil
}
func TestAgent_ChatWithTools_DirectResponse(t *testing.T) {
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{Content: "hello from agent"},
},
}, nil
},
}
client := &Client{model: mock}
executor := NewToolExecutor(&mockGraphService{})
agent := NewAgent(client, executor)
result, err := agent.ChatWithTools(context.Background(), "system", []llms.MessageContent{
llms.TextParts(llms.ChatMessageTypeHuman, "hi"),
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "hello from agent" {
t.Errorf("expected 'hello from agent', got %q", result)
}
}
func TestAgent_ChatWithTools_ToolCallThenResponse(t *testing.T) {
callCount := 0
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
callCount++
if callCount == 1 {
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{
ToolCalls: []llms.ToolCall{
{
ID: "call_1",
FunctionCall: &llms.FunctionCall{
Name: "search_nodes",
Arguments: `{"query":"AI"}`,
},
},
},
},
},
}, nil
}
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{Content: "found AI nodes"},
},
}, nil
},
}
graphSvc := &mockGraphService{
searchNodes: func(query string) []model.Node {
return []model.Node{{ID: "1", Label: "AI"}}
},
}
client := &Client{model: mock}
executor := NewToolExecutor(graphSvc)
agent := NewAgent(client, executor)
result, err := agent.ChatWithTools(context.Background(), "system", []llms.MessageContent{
llms.TextParts(llms.ChatMessageTypeHuman, "search for AI"),
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "found AI nodes" {
t.Errorf("expected 'found AI nodes', got %q", result)
}
if callCount != 2 {
t.Errorf("expected 2 GenerateContent calls, got %d", callCount)
}
}
func TestAgent_ChatWithTools_EmptyResponse(t *testing.T) {
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return &llms.ContentResponse{Choices: []*llms.ContentChoice{}}, nil
},
}
client := &Client{model: mock}
executor := NewToolExecutor(&mockGraphService{})
agent := NewAgent(client, executor)
_, err := agent.ChatWithTools(context.Background(), "system", nil)
if err == nil {
t.Fatal("expected error for empty response, got nil")
}
}
func TestAgent_ChatWithTools_NilResponse(t *testing.T) {
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return nil, nil
},
}
client := &Client{model: mock}
executor := NewToolExecutor(&mockGraphService{})
agent := NewAgent(client, executor)
_, err := agent.ChatWithTools(context.Background(), "system", nil)
if err == nil {
t.Fatal("expected error for nil response, got nil")
}
}
func TestAgent_ChatWithTools_GenerateError(t *testing.T) {
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return nil, fmt.Errorf("network error")
},
}
client := &Client{model: mock}
executor := NewToolExecutor(&mockGraphService{})
agent := NewAgent(client, executor)
_, err := agent.ChatWithTools(context.Background(), "system", nil)
if err == nil {
t.Fatal("expected error, got nil")
}
}
func TestAgent_ChatWithTools_ToolExecutionError_ContinuesLoop(t *testing.T) {
callCount := 0
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
callCount++
if callCount == 1 {
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{
ToolCalls: []llms.ToolCall{
{
ID: "call_1",
FunctionCall: &llms.FunctionCall{
Name: "delete_node",
Arguments: `{"id":"nonexistent"}`,
},
},
},
},
},
}, nil
}
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{Content: "handled error"},
},
}, nil
},
}
graphSvc := &mockGraphService{
deleteNode: func(id string) error {
return fmt.Errorf("node not found")
},
}
client := &Client{model: mock}
executor := NewToolExecutor(graphSvc)
agent := NewAgent(client, executor)
result, err := agent.ChatWithTools(context.Background(), "system", []llms.MessageContent{
llms.TextParts(llms.ChatMessageTypeHuman, "delete nonexistent"),
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "handled error" {
t.Errorf("expected 'handled error', got %q", result)
}
}
func TestAgent_ChatWithTools_MaxIterations(t *testing.T) {
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return &llms.ContentResponse{
Choices: []*llms.ContentChoice{
{
ToolCalls: []llms.ToolCall{
{
ID: "call_loop",
FunctionCall: &llms.FunctionCall{
Name: "search_nodes",
Arguments: `{"query":"loop"}`,
},
},
},
},
},
}, nil
},
}
graphSvc := &mockGraphService{
searchNodes: func(query string) []model.Node {
return []model.Node{}
},
}
client := &Client{model: mock}
executor := NewToolExecutor(graphSvc)
agent := NewAgent(client, executor)
_, err := agent.ChatWithTools(context.Background(), "system", nil)
if err == nil {
t.Fatal("expected max iterations error, got nil")
}
}
func TestAgent_ChatWithTools_ContextCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
mock := &mockModel{
generateContent: func(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
return nil, ctx.Err()
},
}
client := &Client{model: mock}
executor := NewToolExecutor(&mockGraphService{})
agent := NewAgent(client, executor)
_, err := agent.ChatWithTools(ctx, "system", nil)
if err == nil {
t.Fatal("expected error due to cancelled context, got nil")
}
}
+33 -37
View File
@@ -11,11 +11,10 @@ import (
)
type Client struct {
llm llms.Model
model llms.Model
config config.LLMConfig
}
// 创建 LLM 实例
func NewClient(cfg config.LLMConfig) (*Client, error) {
model, err := openai.New(
openai.WithToken(cfg.APIKey),
@@ -24,47 +23,44 @@ func NewClient(cfg config.LLMConfig) (*Client, error) {
)
if err != nil {
return nil, fmt.Errorf("failed to create llm client: %w", err)
}
return nil, fmt.Errorf("failed to create llm client: %w", err)
}
return &Client{config: cfg, llm: model}, nil
return &Client{config: cfg, model: model}, nil
}
// Client简单对话调用
func (c *Client) Chat(ctx context.Context, messages []llms.MessageContent) (string, error) {
resp, err := c.llm.GenerateContent(ctx, messages)
func (c *Client) Generate(ctx context.Context, messages []llms.MessageContent, opts ...llms.CallOption) (*llms.ContentResponse, error) {
resp, err := c.model.GenerateContent(ctx, messages, opts...)
if err != nil {
log.Printf("[ERROR] GenerateContent failed: %v", err)
return "", fmt.Errorf("generate content error: %w", err)
}
if resp == nil || len(resp.Choices) == 0 {
log.Printf("[WARN] Response is empty or has no choices")
return "", fmt.Errorf("no content choices returned")
}
log.Printf("[INFO] Chat completed successfully, choice length: %d", len(resp.Choices[0].Content))
return resp.Choices[0].Content, nil
}
// Client包含系统提示词的调用
func (c *Client) ChatWithSystemPrompt(ctx context.Context, systemPrompt string, messages []llms.MessageContent) (string, error) {
// 构造系统提示词消息
systemMessage := llms.TextParts(llms.ChatMessageTypeSystem, systemPrompt)
// 将系统提示词插入到消息列表头部
allMessages := append([]llms.MessageContent{systemMessage}, messages...)
resp, err := c.llm.GenerateContent(ctx, allMessages)
if err != nil {
log.Printf("[ERROR] GenerateContent with system prompt failed: %v", err)
return "", fmt.Errorf("generate content error: %w", err)
log.Printf("[ERROR] GenerateContent failed: %v", err)
return nil, fmt.Errorf("generate content error: %w", err)
}
if resp == nil || len(resp.Choices) == 0 {
log.Printf("[WARN] Response with system prompt is empty or has no choices")
return "", fmt.Errorf("no content choices returned")
log.Printf("[WARN] Response is empty or has no choices")
return nil, fmt.Errorf("no content choices returned")
}
return resp, nil
}
func (c *Client) Chat(ctx context.Context, messages []llms.MessageContent) (string, error) {
resp, err := c.Generate(ctx, messages)
if err != nil {
return "", err
}
log.Printf("[INFO] Chat completed successfully, choice length: %d", len(resp.Choices[0].Content))
return resp.Choices[0].Content, nil
}
func (c *Client) ChatWithSystemPrompt(ctx context.Context, systemPrompt string, messages []llms.MessageContent) (string, error) {
allMessages := NewMessageBuilder().
WithSystemPrompt(systemPrompt).
WithMessages(messages).
Build()
resp, err := c.Generate(ctx, allMessages)
if err != nil {
return "", err
}
log.Printf("[INFO] Chat with system prompt completed successfully, choice length: %d", len(resp.Choices[0].Content))
return resp.Choices[0].Content, nil
}
}
+6 -6
View File
@@ -30,10 +30,10 @@ func TestNewClient(t *testing.T) {
log.Printf("[INFO] client instance is non-nil, validation passed")
// 验证 llm 实例被成功创建(不调用,只检查非 nil)
if client.llm == nil {
t.Error("client.llm is nil, expected a non-nil llms.Model")
if client.model == nil {
t.Error("client.model is nil, expected a non-nil llms.Model")
} else {
log.Printf("[INFO] client.llm is non-nil, llms.Model created successfully")
log.Printf("[INFO] client.model is non-nil, llms.Model created successfully")
}
log.Printf("[INFO] TestNewClient completed successfully")
@@ -69,7 +69,7 @@ func TestClient_Chat(t *testing.T) {
// 3. 调用被测方法
log.Printf("[INFO] calling client.Chat()...")
got, err := client.Chat(context.Background(), messages)
// 4. 校验结果
if err != nil {
log.Printf("[ERROR] Chat() returned error: %v", err)
@@ -125,7 +125,7 @@ func TestClient_ChatWithSystemPrompt(t *testing.T) {
// 4. 调用被测方法
log.Printf("[INFO] calling client.ChatWithSystemPrompt()...")
got, err := client.ChatWithSystemPrompt(context.Background(), systemPrompt, messages)
// 5. 校验结果
if err != nil {
log.Printf("[ERROR] ChatWithSystemPrompt() returned error: %v", err)
@@ -147,4 +147,4 @@ func TestClient_ChatWithSystemPrompt(t *testing.T) {
}
log.Printf("[INFO] TestClient_ChatWithSystemPrompt completed successfully")
}
}
+58
View File
@@ -0,0 +1,58 @@
package llm
import "github.com/tmc/langchaingo/llms"
type MessageBuilder struct {
messages []llms.MessageContent
}
func NewMessageBuilder() *MessageBuilder {
return &MessageBuilder{
messages: make([]llms.MessageContent, 0),
}
}
func (b *MessageBuilder) WithSystemPrompt(prompt string) *MessageBuilder {
b.messages = append(b.messages, llms.TextParts(llms.ChatMessageTypeSystem, prompt))
return b
}
func (b *MessageBuilder) WithMessages(msgs []llms.MessageContent) *MessageBuilder {
b.messages = append(b.messages, msgs...)
return b
}
func (b *MessageBuilder) AppendAIMessage(content string, toolCalls []llms.ToolCall) *MessageBuilder {
parts := []llms.ContentPart{}
if content != "" {
parts = append(parts, llms.TextPart(content))
}
for _, tc := range toolCalls {
parts = append(parts, tc)
}
b.messages = append(b.messages, llms.MessageContent{
Role: llms.ChatMessageTypeAI,
Parts: parts,
})
return b
}
func (b *MessageBuilder) AppendToolResult(toolCallID, name, result string) *MessageBuilder {
b.messages = append(b.messages, llms.MessageContent{
Role: llms.ChatMessageTypeTool,
Parts: []llms.ContentPart{
llms.ToolCallResponse{
ToolCallID: toolCallID,
Name: name,
Content: result,
},
},
})
return b
}
func (b *MessageBuilder) Build() []llms.MessageContent {
result := make([]llms.MessageContent, len(b.messages))
copy(result, b.messages)
return result
}
+155
View File
@@ -0,0 +1,155 @@
package llm
import (
"testing"
"github.com/tmc/langchaingo/llms"
)
func TestNewMessageBuilder(t *testing.T) {
b := NewMessageBuilder()
if b == nil {
t.Fatal("NewMessageBuilder() returned nil")
}
if len(b.messages) != 0 {
t.Fatalf("expected empty messages, got %d", len(b.messages))
}
}
func TestMessageBuilder_WithSystemPrompt(t *testing.T) {
b := NewMessageBuilder().WithSystemPrompt("you are a helper")
msgs := b.Build()
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
if msgs[0].Role != llms.ChatMessageTypeSystem {
t.Errorf("expected role %v, got %v", llms.ChatMessageTypeSystem, msgs[0].Role)
}
}
func TestMessageBuilder_WithMessages(t *testing.T) {
input := []llms.MessageContent{
llms.TextParts(llms.ChatMessageTypeHuman, "hello"),
llms.TextParts(llms.ChatMessageTypeAI, "hi"),
}
b := NewMessageBuilder().WithMessages(input)
msgs := b.Build()
if len(msgs) != 2 {
t.Fatalf("expected 2 messages, got %d", len(msgs))
}
if msgs[0].Role != llms.ChatMessageTypeHuman {
t.Errorf("expected role %v for msg[0], got %v", llms.ChatMessageTypeHuman, msgs[0].Role)
}
if msgs[1].Role != llms.ChatMessageTypeAI {
t.Errorf("expected role %v for msg[1], got %v", llms.ChatMessageTypeAI, msgs[1].Role)
}
}
func TestMessageBuilder_AppendAIMessage(t *testing.T) {
tc := llms.ToolCall{
ID: "call_1",
FunctionCall: &llms.FunctionCall{
Name: "search_nodes",
Arguments: `{"query":"AI"}`,
},
}
b := NewMessageBuilder().AppendAIMessage("thinking...", []llms.ToolCall{tc})
msgs := b.Build()
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
if msgs[0].Role != llms.ChatMessageTypeAI {
t.Errorf("expected role %v, got %v", llms.ChatMessageTypeAI, msgs[0].Role)
}
if len(msgs[0].Parts) != 2 {
t.Fatalf("expected 2 parts (text + toolCall), got %d", len(msgs[0].Parts))
}
}
func TestMessageBuilder_AppendAIMessage_EmptyContent(t *testing.T) {
tc := llms.ToolCall{
ID: "call_2",
FunctionCall: &llms.FunctionCall{
Name: "get_neighbors",
Arguments: `{"node_id":"1"}`,
},
}
b := NewMessageBuilder().AppendAIMessage("", []llms.ToolCall{tc})
msgs := b.Build()
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
if len(msgs[0].Parts) != 1 {
t.Fatalf("expected 1 part (only toolCall, no text), got %d", len(msgs[0].Parts))
}
}
func TestMessageBuilder_AppendToolResult(t *testing.T) {
b := NewMessageBuilder().AppendToolResult("call_1", "search_nodes", `{"nodes":[]}`)
msgs := b.Build()
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
if msgs[0].Role != llms.ChatMessageTypeTool {
t.Errorf("expected role %v, got %v", llms.ChatMessageTypeTool, msgs[0].Role)
}
}
func TestMessageBuilder_Chaining(t *testing.T) {
b := NewMessageBuilder().
WithSystemPrompt("system").
WithMessages([]llms.MessageContent{
llms.TextParts(llms.ChatMessageTypeHuman, "hello"),
}).
AppendAIMessage("hi", nil).
AppendToolResult("call_1", "search_nodes", "{}")
msgs := b.Build()
if len(msgs) != 4 {
t.Fatalf("expected 4 messages, got %d", len(msgs))
}
if msgs[0].Role != llms.ChatMessageTypeSystem {
t.Errorf("msg[0] expected system, got %v", msgs[0].Role)
}
if msgs[1].Role != llms.ChatMessageTypeHuman {
t.Errorf("msg[1] expected human, got %v", msgs[1].Role)
}
if msgs[2].Role != llms.ChatMessageTypeAI {
t.Errorf("msg[2] expected AI, got %v", msgs[2].Role)
}
if msgs[3].Role != llms.ChatMessageTypeTool {
t.Errorf("msg[3] expected tool, got %v", msgs[3].Role)
}
}
func TestMessageBuilder_Build_ReturnsCopy(t *testing.T) {
b := NewMessageBuilder().WithSystemPrompt("system")
first := b.Build()
second := b.Build()
if len(first) != len(second) {
t.Fatal("Build() should return consistent results")
}
if &first[0] == &second[0] {
t.Error("Build() should return a copy, not the same slice")
}
}
func TestMessageBuilder_Build_DoesNotMutate(t *testing.T) {
b := NewMessageBuilder().WithSystemPrompt("original")
msgs := b.Build()
msgs[0] = llms.TextParts(llms.ChatMessageTypeHuman, "mutated")
original := b.Build()
if original[0].Role == llms.ChatMessageTypeHuman {
t.Error("modifying Build() result should not affect the builder")
}
}
+105
View File
@@ -0,0 +1,105 @@
package llm
import (
"encoding/json"
"fmt"
"knowledge-graph-backend/internal/model"
"github.com/tmc/langchaingo/llms"
)
type KnowledgeGraphToolService interface {
SearchNodes(query string) []model.Node
GetNeighbors(nodeID string) (model.NeighborResponse, bool)
GetNodeByID(id string) (model.Node, bool)
CreateNode(req model.CreateNodeRequest) (model.Node, error)
CreateEdge(req model.CreateEdgeRequest) (model.Edge, error)
DeleteNode(id string) error
DeleteEdge(id string) error
}
type ToolExecutor struct {
graphSvc KnowledgeGraphToolService
}
func NewToolExecutor(graphSvc KnowledgeGraphToolService) *ToolExecutor {
return &ToolExecutor{graphSvc: graphSvc}
}
func (e *ToolExecutor) Execute(toolCall llms.ToolCall) (string, error) {
var result any
var err error
switch toolCall.FunctionCall.Name {
case "search_nodes":
var args struct {
Query string `json:"query"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for search_nodes: %w", err)
}
result = e.graphSvc.SearchNodes(args.Query)
case "get_neighbors":
var args struct {
NodeID string `json:"node_id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for get_neighbors: %w", err)
}
result, _ = e.graphSvc.GetNeighbors(args.NodeID)
case "create_node":
var args model.CreateNodeRequest
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for create_node: %w", err)
}
result, err = e.graphSvc.CreateNode(args)
if err != nil {
return "", err
}
case "create_edge":
var args model.CreateEdgeRequest
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for create_edge: %w", err)
}
result, err = e.graphSvc.CreateEdge(args)
if err != nil {
return "", err
}
case "delete_node":
var args struct {
ID string `json:"id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for delete_node: %w", err)
}
err = e.graphSvc.DeleteNode(args.ID)
if err != nil {
return "", err
}
result = map[string]string{"status": "success", "message": fmt.Sprintf("Node %s deleted successfully", args.ID)}
case "delete_edge":
var args struct {
ID string `json:"id"`
}
if err := json.Unmarshal([]byte(toolCall.FunctionCall.Arguments), &args); err != nil {
return "", fmt.Errorf("invalid arguments for delete_edge: %w", err)
}
err = e.graphSvc.DeleteEdge(args.ID)
if err != nil {
return "", err
}
result = map[string]string{"status": "success", "message": fmt.Sprintf("Edge %s deleted successfully", args.ID)}
default:
return "", fmt.Errorf("unknown tool: %s", toolCall.FunctionCall.Name)
}
resBytes, _ := json.Marshal(result)
return string(resBytes), nil
}
+316
View File
@@ -0,0 +1,316 @@
package llm
import (
"encoding/json"
"fmt"
"testing"
"knowledge-graph-backend/internal/model"
"github.com/tmc/langchaingo/llms"
)
type mockGraphService struct {
searchNodes func(query string) []model.Node
getNeighbors func(nodeID string) (model.NeighborResponse, bool)
getNodeByID func(id string) (model.Node, bool)
createNode func(req model.CreateNodeRequest) (model.Node, error)
createEdge func(req model.CreateEdgeRequest) (model.Edge, error)
deleteNode func(id string) error
deleteEdge func(id string) error
}
func (m *mockGraphService) SearchNodes(query string) []model.Node {
if m.searchNodes != nil {
return m.searchNodes(query)
}
return nil
}
func (m *mockGraphService) GetNeighbors(nodeID string) (model.NeighborResponse, bool) {
if m.getNeighbors != nil {
return m.getNeighbors(nodeID)
}
return model.NeighborResponse{}, false
}
func (m *mockGraphService) GetNodeByID(id string) (model.Node, bool) {
if m.getNodeByID != nil {
return m.getNodeByID(id)
}
return model.Node{}, false
}
func (m *mockGraphService) CreateNode(req model.CreateNodeRequest) (model.Node, error) {
if m.createNode != nil {
return m.createNode(req)
}
return model.Node{}, nil
}
func (m *mockGraphService) CreateEdge(req model.CreateEdgeRequest) (model.Edge, error) {
if m.createEdge != nil {
return m.createEdge(req)
}
return model.Edge{}, nil
}
func (m *mockGraphService) DeleteNode(id string) error {
if m.deleteNode != nil {
return m.deleteNode(id)
}
return nil
}
func (m *mockGraphService) DeleteEdge(id string) error {
if m.deleteEdge != nil {
return m.deleteEdge(id)
}
return nil
}
func TestToolExecutor_SearchNodes(t *testing.T) {
mock := &mockGraphService{
searchNodes: func(query string) []model.Node {
return []model.Node{
{ID: "1", Label: "AI", Type: "concept"},
{ID: "2", Label: "AI Agent", Type: "concept"},
}
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "search_nodes",
Arguments: `{"query":"AI"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
var nodes []model.Node
if err := json.Unmarshal([]byte(result), &nodes); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if len(nodes) != 2 {
t.Errorf("expected 2 nodes, got %d", len(nodes))
}
}
func TestToolExecutor_GetNeighbors(t *testing.T) {
mock := &mockGraphService{
getNeighbors: func(nodeID string) (model.NeighborResponse, bool) {
return model.NeighborResponse{
Nodes: []model.Node{{ID: "2", Label: "ML"}},
Edges: []model.Edge{{ID: "e1", Source: "1", Target: "2"}},
}, true
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "get_neighbors",
Arguments: `{"node_id":"1"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
var resp model.NeighborResponse
if err := json.Unmarshal([]byte(result), &resp); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if len(resp.Nodes) != 1 || len(resp.Edges) != 1 {
t.Errorf("expected 1 node and 1 edge, got %d nodes, %d edges", len(resp.Nodes), len(resp.Edges))
}
}
func TestToolExecutor_CreateNode(t *testing.T) {
mock := &mockGraphService{
createNode: func(req model.CreateNodeRequest) (model.Node, error) {
return model.Node{ID: req.ID, Label: req.Label, Type: req.Type}, nil
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "create_node",
Arguments: `{"id":"node_1","label":"Test","type":"concept"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
var node model.Node
if err := json.Unmarshal([]byte(result), &node); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if node.ID != "node_1" || node.Label != "Test" {
t.Errorf("unexpected node: %+v", node)
}
}
func TestToolExecutor_CreateNode_ServiceError(t *testing.T) {
mock := &mockGraphService{
createNode: func(req model.CreateNodeRequest) (model.Node, error) {
return model.Node{}, fmt.Errorf("duplicate id")
},
}
executor := NewToolExecutor(mock)
_, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "create_node",
Arguments: `{"id":"node_1","label":"Test","type":"concept"}`,
},
})
if err == nil {
t.Fatal("expected error, got nil")
}
}
func TestToolExecutor_CreateEdge(t *testing.T) {
mock := &mockGraphService{
createEdge: func(req model.CreateEdgeRequest) (model.Edge, error) {
return model.Edge{ID: req.ID, Source: req.Source, Target: req.Target, Label: req.Label}, nil
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "create_edge",
Arguments: `{"id":"e1","source":"1","target":"2","label":"connects"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
var edge model.Edge
if err := json.Unmarshal([]byte(result), &edge); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if edge.ID != "e1" || edge.Source != "1" || edge.Target != "2" {
t.Errorf("unexpected edge: %+v", edge)
}
}
func TestToolExecutor_DeleteNode(t *testing.T) {
deleted := false
mock := &mockGraphService{
deleteNode: func(id string) error {
deleted = true
return nil
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "delete_node",
Arguments: `{"id":"node_1"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !deleted {
t.Error("expected DeleteNode to be called")
}
var resp map[string]string
if err := json.Unmarshal([]byte(result), &resp); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if resp["status"] != "success" {
t.Errorf("expected status=success, got %s", resp["status"])
}
}
func TestToolExecutor_DeleteEdge(t *testing.T) {
deleted := false
mock := &mockGraphService{
deleteEdge: func(id string) error {
deleted = true
return nil
},
}
executor := NewToolExecutor(mock)
result, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "delete_edge",
Arguments: `{"id":"e1"}`,
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !deleted {
t.Error("expected DeleteEdge to be called")
}
var resp map[string]string
if err := json.Unmarshal([]byte(result), &resp); err != nil {
t.Fatalf("failed to unmarshal result: %v", err)
}
if resp["status"] != "success" {
t.Errorf("expected status=success, got %s", resp["status"])
}
}
func TestToolExecutor_DeleteNode_ServiceError(t *testing.T) {
mock := &mockGraphService{
deleteNode: func(id string) error {
return fmt.Errorf("not found")
},
}
executor := NewToolExecutor(mock)
_, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "delete_node",
Arguments: `{"id":"node_1"}`,
},
})
if err == nil {
t.Fatal("expected error, got nil")
}
}
func TestToolExecutor_UnknownTool(t *testing.T) {
mock := &mockGraphService{}
executor := NewToolExecutor(mock)
_, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "unknown_tool",
Arguments: `{}`,
},
})
if err == nil {
t.Fatal("expected error for unknown tool, got nil")
}
}
func TestToolExecutor_InvalidArguments(t *testing.T) {
mock := &mockGraphService{}
executor := NewToolExecutor(mock)
_, err := executor.Execute(llms.ToolCall{
FunctionCall: &llms.FunctionCall{
Name: "search_nodes",
Arguments: `{invalid json}`,
},
})
if err == nil {
t.Fatal("expected error for invalid arguments, got nil")
}
}
+117
View File
@@ -0,0 +1,117 @@
package llm
import "github.com/tmc/langchaingo/llms"
func GetKnowledgeGraphTools() []llms.Tool {
return []llms.Tool{
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "search_nodes",
Description: "当需要根据关键词在知识图谱中搜索相关节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "搜索的关键词",
},
},
"required": []string{"query"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "get_neighbors",
Description: "当需要查找某个节点的相邻节点(即与该节点有直接连线的节点)时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"node_id": map[string]any{
"type": "string",
"description": "目标节点的唯一 ID",
},
},
"required": []string{"node_id"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "create_node",
Description: "当需要在知识图谱中创建一个新的节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{"type": "string", "description": "节点的唯一ID,格式通常为 node_xxx"},
"label": map[string]any{"type": "string", "description": "节点的名称/标签,例如:人工智能"},
"type": map[string]any{"type": "string", "description": "节点的类型,例如:概念、实体"},
"x": map[string]any{"type": "number", "description": "节点在画布上的 X 坐标 (可选,默认0)"},
"y": map[string]any{"type": "number", "description": "节点在画布上的 Y 坐标 (可选,默认0)"},
"properties": map[string]any{"type": "object", "description": "节点的额外属性键值对 (可选)"},
},
"required": []string{"id", "label", "type"},
},
},
},
{
Type: "Edge operation",
Function: &llms.FunctionDefinition{
Name: "create_edge",
Description: "当需要在知识图谱中创建一条连线(边)连接两个节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{"type": "string", "description": "边的唯一ID,格式通常为 edge_xxx"},
"source": map[string]any{"type": "string", "description": "起始节点的 ID"},
"target": map[string]any{"type": "string", "description": "目标节点的 ID"},
"label": map[string]any{"type": "string", "description": "连线的名称/标签,例如:包含、依赖于"},
"type": map[string]any{"type": "string", "description": "连线的类型 (可选),例如:CONTAINS"},
"properties": map[string]any{
"type": "object",
"description": "边的额外属性键值对 (可选)",
},
},
"required": []string{"id", "source", "target", "label"},
},
},
},
{
Type: "Node operation",
Function: &llms.FunctionDefinition{
Name: "delete_node",
Description: "当需要从知识图谱中删除一个现有节点时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{
"type": "string",
"description": "要删除的节点的唯一 ID",
},
},
"required": []string{"id"},
},
},
},
{
Type: "Edge operation",
Function: &llms.FunctionDefinition{
Name: "delete_edge",
Description: "当需要从知识图谱中删除一条连线(边)时调用此工具",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"id": map[string]any{
"type": "string",
"description": "要删除的边的唯一 ID",
},
},
"required": []string{"id"},
},
},
},
}
}
+113
View File
@@ -0,0 +1,113 @@
package llm
import (
"testing"
)
func TestGetKnowledgeGraphTools_ReturnsTools(t *testing.T) {
tools := GetKnowledgeGraphTools()
if len(tools) == 0 {
t.Fatal("GetKnowledgeGraphTools() returned no tools")
}
if len(tools) != 6 {
t.Fatalf("expected 6 tools, got %d", len(tools))
}
}
func TestGetKnowledgeGraphTools_ToolNames(t *testing.T) {
tools := GetKnowledgeGraphTools()
expected := map[string]bool{
"search_nodes": true,
"get_neighbors": true,
"create_node": true,
"create_edge": true,
"delete_node": true,
"delete_edge": true,
}
for _, tool := range tools {
if tool.Function == nil {
t.Error("tool has nil Function definition")
continue
}
name := tool.Function.Name
if !expected[name] {
t.Errorf("unexpected tool name: %s", name)
}
delete(expected, name)
}
if len(expected) > 0 {
t.Errorf("missing tools: %v", expected)
}
}
func TestGetKnowledgeGraphTools_HasDescriptions(t *testing.T) {
tools := GetKnowledgeGraphTools()
for _, tool := range tools {
if tool.Function == nil {
continue
}
if tool.Function.Description == "" {
t.Errorf("tool %s has empty description", tool.Function.Name)
}
}
}
func TestGetKnowledgeGraphTools_HasParameters(t *testing.T) {
tools := GetKnowledgeGraphTools()
for _, tool := range tools {
if tool.Function == nil {
continue
}
params, ok := tool.Function.Parameters.(map[string]any)
if !ok {
t.Errorf("tool %s parameters is not map[string]any", tool.Function.Name)
continue
}
if params["type"] != "object" {
t.Errorf("tool %s parameters type is not 'object'", tool.Function.Name)
}
props, ok := params["properties"].(map[string]any)
if !ok {
t.Errorf("tool %s has no properties map", tool.Function.Name)
continue
}
if len(props) == 0 {
t.Errorf("tool %s has empty properties", tool.Function.Name)
}
}
}
func TestGetKnowledgeGraphTools_RequiredFields(t *testing.T) {
tools := GetKnowledgeGraphTools()
for _, tool := range tools {
if tool.Function == nil {
continue
}
params := tool.Function.Parameters.(map[string]any)
required, ok := params["required"].([]string)
if !ok {
t.Errorf("tool %s has no required fields", tool.Function.Name)
continue
}
if len(required) == 0 {
t.Errorf("tool %s has empty required fields", tool.Function.Name)
}
}
}
func TestGetKnowledgeGraphTools_NoDuplicates(t *testing.T) {
tools := GetKnowledgeGraphTools()
seen := map[string]bool{}
for _, tool := range tools {
if tool.Function == nil {
continue
}
name := tool.Function.Name
if seen[name] {
t.Errorf("duplicate tool name: %s", name)
}
seen[name] = true
}
}