317 lines
7.7 KiB
Go
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")
|
|
}
|
|
}
|