Files
knowledge-graph-agent/backend/internal/llm/executor.go
T

106 lines
3.0 KiB
Go
Raw Normal View History

2026-04-20 15:34:40 +08:00
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
}