Files
knowledge-graph-agent/backend/internal/service/neo4j_service.go
T
2026-04-20 15:16:51 +08:00

717 lines
17 KiB
Go

package service
import (
"context"
"fmt"
"strings"
"github.com/neo4j/neo4j-go-driver/v5/neo4j"
"knowledge-graph-backend/internal/model"
)
type Neo4jService struct {
driver neo4j.DriverWithContext
}
func NewNeo4jService(driver neo4j.DriverWithContext) *Neo4jService {
return &Neo4jService{driver: driver}
}
func (s *Neo4jService) GetGraphData() model.GraphData {
ctx := context.Background()
var nodes []model.Node
var edges []model.Edge
nodeResult, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH (n) RETURN n",
map[string]any{},
neo4j.EagerResultTransformer,
)
if err != nil {
fmt.Printf("Error querying nodes: %v\n", err)
return model.GraphData{}
}
for _, record := range nodeResult.Records {
if v, ok := record.Get("n"); ok {
if n, ok := v.(neo4j.Node); ok {
nodes = append(nodes, neo4jNodeToModel(n))
}
}
}
edgeResult, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH (a)-[r]->(b) RETURN a.id AS source, b.id AS target, r",
map[string]any{},
neo4j.EagerResultTransformer,
)
if err != nil {
fmt.Printf("Error querying edges: %v\n", err)
return model.GraphData{Nodes: nodes}
}
for _, record := range edgeResult.Records {
source, _ := record.Get("source")
target, _ := record.Get("target")
v, _ := record.Get("r")
if rel, ok := v.(neo4j.Relationship); ok {
edge := neo4jRelToModel(rel)
if s, ok := source.(string); ok {
edge.Source = s
}
if t, ok := target.(string); ok {
edge.Target = t
}
edges = append(edges, edge)
}
}
return model.GraphData{Nodes: nodes, Edges: edges}
}
func (s *Neo4jService) GetNodeByID(id string) (model.Node, bool) {
ctx := context.Background()
result, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH (n {id: $id}) RETURN n",
map[string]any{"id": id},
neo4j.EagerResultTransformer,
)
if err != nil || len(result.Records) == 0 {
return model.Node{}, false
}
v, _ := result.Records[0].Get("n")
if n, ok := v.(neo4j.Node); ok {
return neo4jNodeToModel(n), true
}
return model.Node{}, false
}
func (s *Neo4jService) SearchNodes(query string) []model.Node {
if query == "" {
return []model.Node{}
}
ctx := context.Background()
result, err := neo4j.ExecuteQuery(ctx, s.driver,
`MATCH (n)
WHERE toLower(n.label) CONTAINS toLower($query)
OR toLower(n.id) CONTAINS toLower($query)
OR (n.type IS NOT NULL AND toLower(n.type) CONTAINS toLower($query))
RETURN n`,
map[string]any{"query": query},
neo4j.EagerResultTransformer,
)
if err != nil {
return []model.Node{}
}
var nodes []model.Node
for _, record := range result.Records {
v, _ := record.Get("n")
if n, ok := v.(neo4j.Node); ok {
nodes = append(nodes, neo4jNodeToModel(n))
}
}
return nodes
}
func (s *Neo4jService) GetNeighbors(nodeID string) (model.NeighborResponse, bool) {
ctx := context.Background()
_, exists := s.GetNodeByID(nodeID)
if !exists {
return model.NeighborResponse{}, false
}
result, err := neo4j.ExecuteQuery(ctx, s.driver,
`MATCH ({id: $id})-[r]-(m)
RETURN m, r, startNode(r).id AS source, endNode(r).id AS target`,
map[string]any{"id": nodeID},
neo4j.EagerResultTransformer,
)
if err != nil {
return model.NeighborResponse{Nodes: []model.Node{}, Edges: []model.Edge{}}, true
}
neighborMap := make(map[string]model.Node)
edgeMap := make(map[string]model.Edge)
for _, record := range result.Records {
if v, ok := record.Get("m"); ok {
if n, ok := v.(neo4j.Node); ok {
modelNode := neo4jNodeToModel(n)
neighborMap[modelNode.ID] = modelNode
}
}
if v, ok := record.Get("r"); ok {
if rel, ok := v.(neo4j.Relationship); ok {
edge := neo4jRelToModel(rel)
if source, ok := record.Get("source"); ok {
if s, ok := source.(string); ok {
edge.Source = s
}
}
if target, ok := record.Get("target"); ok {
if t, ok := target.(string); ok {
edge.Target = t
}
}
edgeMap[edge.ID] = edge
}
}
}
var neighborNodes []model.Node
for _, n := range neighborMap {
neighborNodes = append(neighborNodes, n)
}
var relatedEdges []model.Edge
for _, e := range edgeMap {
relatedEdges = append(relatedEdges, e)
}
return model.NeighborResponse{
Nodes: neighborNodes,
Edges: relatedEdges,
}, true
}
func (s *Neo4jService) GetStats() map[string]int {
ctx := context.Background()
result, err := neo4j.ExecuteQuery(ctx, s.driver,
`MATCH (n)
RETURN count(n) AS totalNodes,
sum(CASE WHEN n.type = '概念' THEN 1 ELSE 0 END) AS conceptNodes,
sum(CASE WHEN n.type = '工具' THEN 1 ELSE 0 END) AS toolNodes,
sum(CASE WHEN n.type = '应用' THEN 1 ELSE 0 END) AS applicationNodes`,
map[string]any{},
neo4j.EagerResultTransformer,
)
if err != nil || len(result.Records) == 0 {
return map[string]int{}
}
record := result.Records[0]
totalNodes := getInt(record, "totalNodes")
conceptNodes := getInt(record, "conceptNodes")
toolNodes := getInt(record, "toolNodes")
applicationNodes := getInt(record, "applicationNodes")
totalEdges := 0
edgeResult, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH ()-[r]->() RETURN count(r) AS totalEdges",
map[string]any{},
neo4j.EagerResultTransformer,
)
if err == nil && len(edgeResult.Records) > 0 {
totalEdges = getInt(edgeResult.Records[0], "totalEdges")
}
return map[string]int{
"totalNodes": totalNodes,
"totalEdges": totalEdges,
"conceptNodes": conceptNodes,
"toolNodes": toolNodes,
"applicationNodes": applicationNodes,
}
}
// GetSimpleGraphData 获取简化的图数据,仅保留核心信息,适合 LLM 处理
func (s *Neo4jService) GetSimpleGraphData() model.SimpleGraphData {
ctx := context.Background()
var nodes []model.SimpleNode
var edges []model.SimpleEdge
nodeResult, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH (n) RETURN n",
map[string]any{},
neo4j.EagerResultTransformer,
)
if err != nil {
fmt.Printf("Error querying nodes: %v\n", err)
return model.SimpleGraphData{}
}
for _, record := range nodeResult.Records {
if v, ok := record.Get("n"); ok {
if n, ok := v.(neo4j.Node); ok {
props := n.Props
node := model.SimpleNode{
ID: getStr(props, "id"),
Label: getStr(props, "label"),
Type: getStr(props, "type"),
}
nodes = append(nodes, node)
}
}
}
edgeResult, err := neo4j.ExecuteQuery(ctx, s.driver,
"MATCH (a)-[r]->(b) RETURN a.id AS source, b.id AS target, r",
map[string]any{},
neo4j.EagerResultTransformer,
)
if err != nil {
fmt.Printf("Error querying edges: %v\n", err)
return model.SimpleGraphData{Nodes: nodes}
}
for _, record := range edgeResult.Records {
source, _ := record.Get("source")
target, _ := record.Get("target")
v, _ := record.Get("r")
if rel, ok := v.(neo4j.Relationship); ok {
props := rel.Props
edge := model.SimpleEdge{
Source: getStr(props, "source"),
Target: getStr(props, "target"),
Type: rel.Type,
Label: getStr(props, "label"),
}
if s, ok := source.(string); ok {
edge.Source = s
}
if t, ok := target.(string); ok {
edge.Target = t
}
edges = append(edges, edge)
}
}
return model.SimpleGraphData{Nodes: nodes, Edges: edges}
}
func neo4jNodeToModel(n neo4j.Node) model.Node {
props := n.Props
node := model.Node{
ID: getStr(props, "id"),
Label: getStr(props, "label"),
Type: getStr(props, "type"),
Properties: make(map[string]interface{}),
}
if v, ok := props["x"]; ok && v != nil {
node.X = getFloat64(v)
}
if v, ok := props["y"]; ok && v != nil {
node.Y = getFloat64(v)
}
if v, ok := props["style"]; ok && v != nil {
if s, ok := v.(map[string]interface{}); ok {
node.Style = s
}
}
for k, v := range props {
switch k {
case "id", "label", "type", "x", "y", "style":
default:
node.Properties[k] = v
}
}
return node
}
func neo4jRelToModel(r neo4j.Relationship) model.Edge {
props := r.Props
edge := model.Edge{
ID: getStr(props, "id"),
Label: getStr(props, "label"),
Type: r.Type,
Properties: make(map[string]interface{}),
}
if v, ok := props["style"]; ok && v != nil {
if s, ok := v.(map[string]interface{}); ok {
edge.Style = s
}
}
for k, v := range props {
switch k {
case "id", "label", "style":
default:
edge.Properties[k] = v
}
}
return edge
}
func getStr(props map[string]any, key string) string {
if v, ok := props[key]; ok && v != nil {
if s, ok := v.(string); ok {
return s
}
}
return ""
}
func getFloat64(v any) float64 {
switch val := v.(type) {
case float64:
return val
case int64:
return float64(val)
case int:
return float64(val)
}
return 0
}
func getInt(record *neo4j.Record, key string) int {
v, _ := record.Get(key)
switch val := v.(type) {
case int64:
return int(val)
case int:
return val
case float64:
return int(val)
}
return 0
}
// CreateNode 创建新节点
func (s *Neo4jService) CreateNode(req model.CreateNodeRequest) (model.Node, error) {
ctx := context.Background()
// 检查节点是否已存在
checkQuery := "MATCH (n {id: $id}) RETURN n"
checkResult, err := neo4j.ExecuteQuery(ctx, s.driver, checkQuery,
map[string]any{"id": req.ID}, neo4j.EagerResultTransformer)
if err != nil {
return model.Node{}, fmt.Errorf("error checking node existence: %w", err)
}
if len(checkResult.Records) > 0 {
return model.Node{}, fmt.Errorf("node with id %s already exists", req.ID)
}
// 构建创建属性
props := map[string]any{
"id": req.ID,
"label": req.Label,
}
if req.Type != "" {
props["type"] = req.Type
}
if req.X != 0 {
props["x"] = req.X
}
if req.Y != 0 {
props["y"] = req.Y
}
// 合并自定义属性(避免覆盖系统属性)
if req.Properties != nil {
for k, v := range req.Properties {
if k != "id" && k != "label" && k != "type" && k != "x" && k != "y" {
props[k] = v
}
}
}
// 执行创建
nodeType := "Node"
if req.Type != "" {
nodeType = req.Type
}
query := fmt.Sprintf(`
CREATE (n:%s)
SET n = $props
RETURN n`, nodeType)
result, err := neo4j.ExecuteQuery(ctx, s.driver, query,
map[string]any{"props": props}, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error creating node: %v\n", err)
return model.Node{}, fmt.Errorf("failed to create node: %w", err)
}
if len(result.Records) == 0 {
return model.Node{}, fmt.Errorf("failed to create node: no result returned")
}
v, _ := result.Records[0].Get("n")
if n, ok := v.(neo4j.Node); ok {
return neo4jNodeToModel(n), nil
}
return model.Node{}, fmt.Errorf("unexpected result type while creating node")
}
// CreateNodeV2 创建新节点(改进版)
func (s *Neo4jService) CreateNodeV2(req model.CreateNodeRequest) (model.Node, error) {
ctx := context.Background()
// 先检查节点是否已存在
checkQuery := "MATCH (n {id: $id}) RETURN n"
checkResult, err := neo4j.ExecuteQuery(ctx, s.driver, checkQuery,
map[string]any{"id": req.ID}, neo4j.EagerResultTransformer)
if err != nil {
return model.Node{}, err
}
if len(checkResult.Records) > 0 {
return model.Node{}, fmt.Errorf("node with id %s already exists", req.ID)
}
// 构建创建属性
props := map[string]any{
"id": req.ID,
"label": req.Label,
}
if req.Type != "" {
props["type"] = req.Type
}
if req.X != 0 {
props["x"] = req.X
}
if req.Y != 0 {
props["y"] = req.Y
}
// 合并自定义属性(避免覆盖系统属性)
if req.Properties != nil {
for k, v := range req.Properties {
if k != "id" && k != "label" && k != "type" && k != "x" && k != "y" {
props[k] = v
}
}
}
// 执行创建
nodeType := "Node"
if req.Type != "" {
nodeType = req.Type
}
query := fmt.Sprintf(`
CREATE (n:%s)
SET n = $props
RETURN n`, nodeType)
result, err := neo4j.ExecuteQuery(ctx, s.driver, query,
map[string]any{"props": props}, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error creating node: %v\n", err)
return model.Node{}, err
}
if len(result.Records) == 0 {
return model.Node{}, fmt.Errorf("failed to create node")
}
v, _ := result.Records[0].Get("n")
if n, ok := v.(neo4j.Node); ok {
return neo4jNodeToModel(n), nil
}
return model.Node{}, fmt.Errorf("unexpected result type while creating node")
}
// UpdateNode 更新节点
func (s *Neo4jService) UpdateNode(id string, req model.UpdateNodeRequest) (model.Node, error) {
ctx := context.Background()
// 检查节点是否存在
_, exists := s.GetNodeByID(id)
if !exists {
return model.Node{}, fmt.Errorf("node with id %s not found", id)
}
// 构建更新语句
setClauses := []string{}
params := map[string]any{"id": id}
if req.Label != "" {
setClauses = append(setClauses, "n.label = $label")
params["label"] = req.Label
}
if req.Type != "" {
setClauses = append(setClauses, "n.type = $type")
params["type"] = req.Type
}
if req.X != nil {
setClauses = append(setClauses, "n.x = $x")
params["x"] = *req.X
}
if req.Y != nil {
setClauses = append(setClauses, "n.y = $y")
params["y"] = *req.Y
}
query := `MATCH (n {id: $id})`
// 处理自定义属性(不管是否有标准字段更新)
if len(req.Properties) > 0 {
for key, value := range req.Properties {
// 跳过系统属性,避免冲突
if key != "id" && key != "label" && key != "type" && key != "x" && key != "y" {
setClauses = append(setClauses, fmt.Sprintf("n.%s = $%s", key, key))
params[key] = value
}
}
}
if len(setClauses) > 0 {
query += ` SET ` + fmt.Sprintf("%s", strings.Join(setClauses, ", "))
}
query += ` RETURN n`
result, err := neo4j.ExecuteQuery(ctx, s.driver, query, params, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error updating node: %v\n", err)
return model.Node{}, err
}
if len(result.Records) == 0 {
return model.Node{}, fmt.Errorf("failed to update node")
}
v, _ := result.Records[0].Get("n")
if n, ok := v.(neo4j.Node); ok {
return neo4jNodeToModel(n), nil
}
return model.Node{}, fmt.Errorf("unexpected result type while updating node")
}
// DeleteNode 删除节点
func (s *Neo4jService) DeleteNode(id string) error {
ctx := context.Background()
// 检查节点是否存在
_, exists := s.GetNodeByID(id)
if !exists {
return fmt.Errorf("node with id %s not found", id)
}
// 先删除与该节点相关的所有关系
deleteRelationsQuery := `
MATCH (n {id: $id})-[r]-(m)
DELETE r`
_, err := neo4j.ExecuteQuery(ctx, s.driver, deleteRelationsQuery,
map[string]any{"id": id}, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error deleting relations: %v\n", err)
// 继续尝试删除节点
}
// 删除节点
deleteNodeQuery := `
MATCH (n {id: $id})
DELETE n`
_, err = neo4j.ExecuteQuery(ctx, s.driver, deleteNodeQuery,
map[string]any{"id": id}, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error deleting node: %v\n", err)
return err
}
return nil
}
// CreateEdge 创建新边
func (s *Neo4jService) CreateEdge(req model.CreateEdgeRequest) (model.Edge, error) {
ctx := context.Background()
// 检查源节点和目标节点是否存在
_, sourceExists := s.GetNodeByID(req.Source)
if !sourceExists {
return model.Edge{}, fmt.Errorf("source node with id %s not found", req.Source)
}
_, targetExists := s.GetNodeByID(req.Target)
if !targetExists {
return model.Edge{}, fmt.Errorf("target node with id %s not found", req.Target)
}
// 构建创建属性
// 构建创建属性
props := map[string]any{
"id": req.ID,
"label": req.Label,
}
// 合并自定义属性
if req.Properties != nil {
for k, v := range req.Properties {
if k != "id" && k != "label" {
props[k] = v
}
}
}
// 确定边类型
relationType := "RELATED_TO"
if req.Type != "" {
relationType = strings.ToUpper(req.Type)
}
// 执行创建
query := fmt.Sprintf(`
MATCH (a {id: $source}), (b {id: $target})
CREATE (a)-[r:%s]->(b)
SET r = $props
RETURN r`, relationType)
params := map[string]any{
"source": req.Source,
"target": req.Target,
"props": props,
}
result, err := neo4j.ExecuteQuery(ctx, s.driver, query, params, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error creating edge: %v\n", err)
return model.Edge{}, err
}
if len(result.Records) == 0 {
return model.Edge{}, fmt.Errorf("failed to create edge")
}
v, _ := result.Records[0].Get("r")
if r, ok := v.(neo4j.Relationship); ok {
edge := neo4jRelToModel(r)
edge.Source = req.Source
edge.Target = req.Target
return edge, nil
}
return model.Edge{}, fmt.Errorf("unexpected result type while creating edge")
}
// DeleteEdge 删除边
func (s *Neo4jService) DeleteEdge(edgeID string) error {
ctx := context.Background()
// 检查边是否存在
checkQuery := `
MATCH ()-[r]-()
WHERE r.id = $id
RETURN r`
checkResult, err := neo4j.ExecuteQuery(ctx, s.driver, checkQuery,
map[string]any{"id": edgeID}, neo4j.EagerResultTransformer)
if err != nil {
return err
}
if len(checkResult.Records) == 0 {
return fmt.Errorf("edge with id %s not found", edgeID)
}
// 删除边
deleteQuery := `
MATCH ()-[r]-()
WHERE r.id = $id
DELETE r`
_, err = neo4j.ExecuteQuery(ctx, s.driver, deleteQuery,
map[string]any{"id": edgeID}, neo4j.EagerResultTransformer)
if err != nil {
fmt.Printf("Error deleting edge: %v\n", err)
return err
}
return nil
}