106 lines
3.0 KiB
Go
106 lines
3.0 KiB
Go
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
|
|
}
|