refactor: 重构llm包
This commit is contained in:
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user