diff --git a/backend/internal/handler/crud_handler.go b/backend/internal/handler/crud_handler.go index 740c189..a429260 100644 --- a/backend/internal/handler/crud_handler.go +++ b/backend/internal/handler/crud_handler.go @@ -46,23 +46,17 @@ func NewEdgeCRUDHandler(svc service.GraphService) *EdgeCRUDHandler { func (h *NodeCRUDHandler) CreateNode(c *gin.Context) { var req model.CreateNodeRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, model.ErrorResponse{ - Error: "invalid_request", - Message: err.Error(), - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "invalid_request", Message: err.Error()}) return } node, err := h.service.CreateNode(req) if err != nil { - c.JSON(http.StatusConflict, model.ErrorResponse{ - Error: "create_failed", - Message: err.Error(), - }) + c.JSON(http.StatusConflict, model.ErrorResponse{Error: "create_failed", Message: err.Error()}) return } - c.JSON(http.StatusOK, node) + c.JSON(http.StatusCreated, node) } // UpdateNode 更新节点 @@ -82,19 +76,13 @@ func (h *NodeCRUDHandler) UpdateNode(c *gin.Context) { var req model.UpdateNodeRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, model.ErrorResponse{ - Error: "invalid_request", - Message: err.Error(), - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "invalid_request", Message: err.Error()}) return } node, err := h.service.UpdateNode(id, req) if err != nil { - c.JSON(http.StatusNotFound, model.ErrorResponse{ - Error: "update_failed", - Message: err.Error(), - }) + c.JSON(http.StatusNotFound, model.ErrorResponse{Error: "update_failed", Message: err.Error()}) return } @@ -113,20 +101,12 @@ func (h *NodeCRUDHandler) UpdateNode(c *gin.Context) { func (h *NodeCRUDHandler) DeleteNode(c *gin.Context) { id := c.Param("id") - err := h.service.DeleteNode(id) - if err != nil { - c.JSON(http.StatusNotFound, model.ErrorResponse{ - Error: "delete_failed", - Message: err.Error(), - }) + if err := h.service.DeleteNode(id); err != nil { + c.JSON(http.StatusNotFound, model.ErrorResponse{Error: "delete_failed", Message: err.Error()}) return } - c.JSON(http.StatusOK, model.DeleteResponse{ - Success: true, - Message: "节点已成功删除", - ID: id, - }) + c.JSON(http.StatusOK, model.DeleteResponse{Success: true, Message: "节点已成功删除", ID: id}) } // CreateEdge 创建新边 @@ -144,33 +124,17 @@ func (h *NodeCRUDHandler) DeleteNode(c *gin.Context) { func (h *EdgeCRUDHandler) CreateEdge(c *gin.Context) { var req model.CreateEdgeRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, model.ErrorResponse{ - Error: "invalid_request", - Message: err.Error(), - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "invalid_request", Message: err.Error()}) return } edge, err := h.service.CreateEdge(req) if err != nil { - c.JSON(http.StatusConflict, model.ErrorResponse{ - Error: "create_failed", - Message: err.Error(), - }) + c.JSON(http.StatusConflict, model.ErrorResponse{Error: "create_failed", Message: err.Error()}) return } - // 转换为响应格式 - response := model.EdgeResponse{ - ID: edge.ID, - Source: edge.Source, - Target: edge.Target, - Label: edge.Label, - Type: edge.Type, - Properties: edge.Properties, - } - - c.JSON(http.StatusOK, response) + c.JSON(http.StatusCreated, model.EdgeResponse{ID: edge.ID, Source: edge.Source, Target: edge.Target, Label: edge.Label, Type: edge.Type, Properties: edge.Properties}) } // DeleteEdge 删除边 @@ -185,18 +149,10 @@ func (h *EdgeCRUDHandler) CreateEdge(c *gin.Context) { func (h *EdgeCRUDHandler) DeleteEdge(c *gin.Context) { edgeID := c.Param("id") - err := h.service.DeleteEdge(edgeID) - if err != nil { - c.JSON(http.StatusNotFound, model.ErrorResponse{ - Error: "delete_failed", - Message: err.Error(), - }) + if err := h.service.DeleteEdge(edgeID); err != nil { + c.JSON(http.StatusNotFound, model.ErrorResponse{Error: "delete_failed", Message: err.Error()}) return } - c.JSON(http.StatusOK, model.DeleteResponse{ - Success: true, - Message: "边已成功删除", - ID: edgeID, - }) + c.JSON(http.StatusOK, model.DeleteResponse{Success: true, Message: "边已成功删除", ID: edgeID}) } diff --git a/backend/internal/handler/node_handler.go b/backend/internal/handler/node_handler.go index bb3ad25..cc6148c 100644 --- a/backend/internal/handler/node_handler.go +++ b/backend/internal/handler/node_handler.go @@ -4,6 +4,7 @@ import ( "net/http" "github.com/gin-gonic/gin" + "knowledge-graph-backend/internal/model" "knowledge-graph-backend/internal/service" ) @@ -29,19 +30,13 @@ func NewNodeHandler(svc service.GraphService) *NodeHandler { func (h *NodeHandler) GetNodeByID(c *gin.Context) { nodeID := c.Param("id") if nodeID == "" { - c.JSON(http.StatusBadRequest, gin.H{ - "error": "Bad Request", - "message": "Node ID is required", - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "bad_request", Message: "Node ID is required"}) return } node, found := h.service.GetNodeByID(nodeID) if !found { - c.JSON(http.StatusNotFound, gin.H{ - "error": "Not Found", - "message": "Node with ID '" + nodeID + "' not found", - }) + c.JSON(http.StatusNotFound, model.ErrorResponse{Error: "not_found", Message: "Node with ID '" + nodeID + "' not found"}) return } @@ -60,19 +55,13 @@ func (h *NodeHandler) GetNodeByID(c *gin.Context) { func (h *NodeHandler) GetNeighbors(c *gin.Context) { nodeID := c.Param("id") if nodeID == "" { - c.JSON(http.StatusBadRequest, gin.H{ - "error": "Bad Request", - "message": "Node ID is required", - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "bad_request", Message: "Node ID is required"}) return } neighbors, found := h.service.GetNeighbors(nodeID) if !found { - c.JSON(http.StatusNotFound, gin.H{ - "error": "Not Found", - "message": "Node with ID '" + nodeID + "' not found", - }) + c.JSON(http.StatusNotFound, model.ErrorResponse{Error: "not_found", Message: "Node with ID '" + nodeID + "' not found"}) return } diff --git a/backend/internal/handler/search_handler.go b/backend/internal/handler/search_handler.go index 32cef9c..e898e55 100644 --- a/backend/internal/handler/search_handler.go +++ b/backend/internal/handler/search_handler.go @@ -4,6 +4,7 @@ import ( "net/http" "github.com/gin-gonic/gin" + "knowledge-graph-backend/internal/model" "knowledge-graph-backend/internal/service" ) @@ -28,10 +29,7 @@ func NewSearchHandler(svc service.GraphService) *SearchHandler { func (h *SearchHandler) SearchNodes(c *gin.Context) { query := c.Query("q") if query == "" { - c.JSON(http.StatusBadRequest, gin.H{ - "error": "Bad Request", - "message": "Query parameter 'q' is required", - }) + c.JSON(http.StatusBadRequest, model.ErrorResponse{Error: "bad_request", Message: "Query parameter 'q' is required"}) return } diff --git a/backend/internal/model/response.go b/backend/internal/model/response.go index 2aeb15b..eeda49d 100644 --- a/backend/internal/model/response.go +++ b/backend/internal/model/response.go @@ -1,5 +1,11 @@ package model +// ErrorResponse represents a standardized API error response. +type ErrorResponse struct { + Error string `json:"error" example:"bad_request"` + Message string `json:"message" example:"invalid request"` +} + // CreateNodeRequest creates a new node type CreateNodeRequest struct { ID string `json:"id" binding:"required" example:"node_123"` diff --git a/backend/internal/service/graph_crud.go b/backend/internal/service/graph_crud.go new file mode 100644 index 0000000..c900b70 --- /dev/null +++ b/backend/internal/service/graph_crud.go @@ -0,0 +1,195 @@ +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 +} diff --git a/backend/internal/service/graph_queries.go b/backend/internal/service/graph_queries.go new file mode 100644 index 0000000..f5fef61 --- /dev/null +++ b/backend/internal/service/graph_queries.go @@ -0,0 +1,194 @@ +package service + +import ( + "context" + "fmt" + + "github.com/neo4j/neo4j-go-driver/v5/neo4j" + + "knowledge-graph-backend/internal/model" +) + +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) { + node, found, err := s.fetchNodeByID(id) + if err != nil || !found { + return model.Node{}, false + } + return node, true +} + +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() + + if _, exists := s.GetNodeByID(nodeID); !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, + } +} diff --git a/backend/internal/service/graph_simple.go b/backend/internal/service/graph_simple.go new file mode 100644 index 0000000..16f10a1 --- /dev/null +++ b/backend/internal/service/graph_simple.go @@ -0,0 +1,64 @@ +package service + +import ( + "context" + "fmt" + + "github.com/neo4j/neo4j-go-driver/v5/neo4j" + + "knowledge-graph-backend/internal/model" +) + +// 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} +} diff --git a/backend/internal/service/interface.go b/backend/internal/service/interface.go index 062eb75..9a3266e 100644 --- a/backend/internal/service/interface.go +++ b/backend/internal/service/interface.go @@ -10,7 +10,6 @@ type GraphService interface { GetNeighbors(nodeID string) (model.NeighborResponse, bool) GetStats() map[string]int - // 新增增删改接口 CreateNode(req model.CreateNodeRequest) (model.Node, error) UpdateNode(id string, req model.UpdateNodeRequest) (model.Node, error) DeleteNode(id string) error diff --git a/backend/internal/service/neo4j_service.go b/backend/internal/service/neo4j_service.go deleted file mode 100644 index 9aa6ac5..0000000 --- a/backend/internal/service/neo4j_service.go +++ /dev/null @@ -1,716 +0,0 @@ -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 -} diff --git a/backend/internal/service/neo4j_utils.go b/backend/internal/service/neo4j_utils.go new file mode 100644 index 0000000..e8a6757 --- /dev/null +++ b/backend/internal/service/neo4j_utils.go @@ -0,0 +1,174 @@ +package service + +import ( + "fmt" + "strings" + + "github.com/neo4j/neo4j-go-driver/v5/neo4j" + + "knowledge-graph-backend/internal/model" +) + +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 +} + +func buildNodeProps(req model.CreateNodeRequest) map[string]any { + 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 + } + for k, v := range req.Properties { + if isReservedNodeProperty(k) { + continue + } + props[k] = v + } + return props +} + +func buildEdgeProps(req model.CreateEdgeRequest) map[string]any { + props := map[string]any{"id": req.ID, "label": req.Label} + for k, v := range req.Properties { + if isReservedEdgeProperty(k) { + continue + } + props[k] = v + } + return props +} + +func isReservedNodeProperty(key string) bool { + switch key { + case "id", "label", "type", "x", "y", "style": + return true + default: + return false + } +} + +func isReservedEdgeProperty(key string) bool { + switch key { + case "id", "label", "type", "style", "source", "target": + return true + default: + return false + } +} + +func sanitizeLabel(label string) string { + if label == "" { + return "" + } + var b strings.Builder + for i, r := range label { + if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '_' { + b.WriteRune(r) + continue + } + if i == 0 { + continue + } + b.WriteRune('_') + } + return strings.Trim(b.String(), "_") +} + +func sanitizeRelationshipType(relType string) string { + if relType == "" { + return "" + } + return sanitizeLabel(strings.ToUpper(relType)) +} + +func wrapNotFound(entity string, id string) error { + return fmt.Errorf("%s with id %s not found", entity, id) +}