Files
2026-04-20 15:34:40 +08:00

317 lines
7.7 KiB
Go

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")
}
}