196 lines
5.0 KiB
Go
196 lines
5.0 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"github.com/neo4j/neo4j-go-driver/v5/neo4j"
|
|
|
|
"knowledge-graph-backend/internal/model"
|
|
)
|
|
|
|
// CreateNode 创建新节点
|
|
func (s *Neo4jService) CreateNode(req model.CreateNodeRequest) (model.Node, error) {
|
|
ctx := context.Background()
|
|
|
|
if req.ID == "" {
|
|
return model.Node{}, fmt.Errorf("node id is required")
|
|
}
|
|
if req.Label == "" {
|
|
return model.Node{}, fmt.Errorf("node label is required")
|
|
}
|
|
|
|
if _, found, err := s.fetchNodeByID(req.ID); err != nil {
|
|
return model.Node{}, err
|
|
} else if found {
|
|
return model.Node{}, fmt.Errorf("node with id %s already exists", req.ID)
|
|
}
|
|
|
|
props := buildNodeProps(req)
|
|
nodeType := sanitizeLabel(req.Type)
|
|
if nodeType == "" {
|
|
nodeType = "Node"
|
|
}
|
|
|
|
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 {
|
|
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")
|
|
}
|
|
|
|
// UpdateNode 更新节点
|
|
func (s *Neo4jService) UpdateNode(id string, req model.UpdateNodeRequest) (model.Node, error) {
|
|
ctx := context.Background()
|
|
|
|
if _, found, err := s.fetchNodeByID(id); err != nil {
|
|
return model.Node{}, err
|
|
} else if !found {
|
|
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
|
|
}
|
|
|
|
for key, value := range req.Properties {
|
|
if isReservedNodeProperty(key) {
|
|
continue
|
|
}
|
|
safeKey := sanitizeLabel(key)
|
|
setClauses = append(setClauses, fmt.Sprintf("n.%s = $%s", safeKey, safeKey))
|
|
params[safeKey] = value
|
|
}
|
|
|
|
query := `MATCH (n {id: $id})`
|
|
if len(setClauses) > 0 {
|
|
query += ` SET ` + strings.Join(setClauses, ", ")
|
|
}
|
|
query += ` RETURN n`
|
|
|
|
result, err := neo4j.ExecuteQuery(ctx, s.driver, query, params, neo4j.EagerResultTransformer)
|
|
if err != nil {
|
|
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()
|
|
|
|
if _, found, err := s.fetchNodeByID(id); err != nil {
|
|
return err
|
|
} else if !found {
|
|
return fmt.Errorf("node with id %s not found", id)
|
|
}
|
|
|
|
_, err := neo4j.ExecuteQuery(ctx, s.driver, `
|
|
MATCH (n {id: $id})
|
|
DETACH DELETE n`, map[string]any{"id": id}, neo4j.EagerResultTransformer)
|
|
return err
|
|
}
|
|
|
|
// CreateEdge 创建新边
|
|
func (s *Neo4jService) CreateEdge(req model.CreateEdgeRequest) (model.Edge, error) {
|
|
ctx := context.Background()
|
|
|
|
if req.Source == "" || req.Target == "" {
|
|
return model.Edge{}, fmt.Errorf("source and target are required")
|
|
}
|
|
if _, found, err := s.fetchNodeByID(req.Source); err != nil {
|
|
return model.Edge{}, err
|
|
} else if !found {
|
|
return model.Edge{}, fmt.Errorf("source node with id %s not found", req.Source)
|
|
}
|
|
if _, found, err := s.fetchNodeByID(req.Target); err != nil {
|
|
return model.Edge{}, err
|
|
} else if !found {
|
|
return model.Edge{}, fmt.Errorf("target node with id %s not found", req.Target)
|
|
}
|
|
|
|
props := buildEdgeProps(req)
|
|
relationType := sanitizeRelationshipType(req.Type)
|
|
if relationType == "" {
|
|
relationType = "RELATED_TO"
|
|
}
|
|
|
|
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 {
|
|
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()
|
|
|
|
_, err := neo4j.ExecuteQuery(ctx, s.driver, `
|
|
MATCH ()-[r]-()
|
|
WHERE r.id = $id
|
|
DELETE r`, map[string]any{"id": edgeID}, neo4j.EagerResultTransformer)
|
|
return err
|
|
}
|