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 }