package rag import ( "context" "encoding/json" "fmt" "log" "strings" "github.com/neo4j/neo4j-go-driver/v5/neo4j" "github.com/neo4j/neo4j-go-driver/v5/neo4j/auth" "github.com/sundynix/sundynix-shared/prompts" ) // Triple 是一条知识三元组(主体-关系-客体)。 type Triple struct { S string `json:"s"` P string `json:"p"` O string `json:"o"` } // graphStore 是 GraphRAG 的图路:实体/关系存 Neo4j。 type graphStore struct { driver neo4j.DriverWithContext } func openGraph(ctx context.Context, uri, user, pass string) *graphStore { if uri == "" { return &graphStore{} } drv, err := neo4j.NewDriverWithContext(uri, auth.BasicTokenManager(func(context.Context) (neo4j.AuthToken, error) { return neo4j.BasicAuth(user, pass, ""), nil })) if err != nil { log.Printf("[rag] Neo4j 连接失败,图谱路降级: %v", err) return &graphStore{} } if err := drv.VerifyConnectivity(ctx); err != nil { log.Printf("[rag] Neo4j 不可用,图谱路降级: %v", err) return &graphStore{} } // 实体唯一约束(kb+name)。 _, _ = neo4j.ExecuteQuery(ctx, drv, "CREATE CONSTRAINT entity_key IF NOT EXISTS FOR (e:Entity) REQUIRE (e.kb, e.name) IS UNIQUE", nil, neo4j.EagerResultTransformer) log.Printf("[rag] Neo4j connected %s", uri) return &graphStore{driver: drv} } func (g *graphStore) ready() bool { return g != nil && g.driver != nil } func (g *graphStore) close(ctx context.Context) { if g.ready() { _ = g.driver.Close(ctx) } } // store 把三元组 MERGE 进 Neo4j(实体 + 关系,按 kb 隔离)。 // 关系带 file_id(来源文档稳定 ID)→ 同一三元组来自不同文档各成一条边,供按文档级联删; // 实体仍按 (kb,name) 共享去重(一个实体可被多篇文档提及)。 func (g *graphStore) store(ctx context.Context, kb, fileID string, triples []Triple) (int, error) { if !g.ready() { return 0, nil } n := 0 for _, t := range triples { if t.S == "" || t.O == "" || t.P == "" { continue } _, err := neo4j.ExecuteQuery(ctx, g.driver, `MERGE (a:Entity {kb:$kb, name:$s}) MERGE (b:Entity {kb:$kb, name:$o}) MERGE (a)-[r:REL {type:$p, file_id:$fid}]->(b)`, map[string]any{"kb": kb, "s": t.S, "o": t.O, "p": t.P, "fid": fileID}, neo4j.EagerResultTransformer, neo4j.ExecuteQueryWithDatabase("neo4j")) if err != nil { return n, err } n++ } return n, nil } // deleteByFile 删除某文档(file_id)贡献的全部关系,再清掉因此变孤儿的实体(无任何关系)。 // 共享实体(仍被其它文档的关系连着)保留。kb 作用域隔离。 func (g *graphStore) deleteByFile(ctx context.Context, kb, fileID string) error { if !g.ready() || fileID == "" { return nil } if _, err := neo4j.ExecuteQuery(ctx, g.driver, `MATCH (:Entity {kb:$kb})-[r:REL {file_id:$fid}]->(:Entity {kb:$kb}) DELETE r`, map[string]any{"kb": kb, "fid": fileID}, neo4j.EagerResultTransformer, neo4j.ExecuteQueryWithDatabase("neo4j")); err != nil { return err } // 孤儿实体清理:本 kb 内不再连任何关系的实体删除。 _, err := neo4j.ExecuteQuery(ctx, g.driver, `MATCH (e:Entity {kb:$kb}) WHERE NOT (e)--() DELETE e`, map[string]any{"kb": kb}, neo4j.EagerResultTransformer, neo4j.ExecuteQueryWithDatabase("neo4j")) return err } // search 图谱召回:找查询里提到的实体,返回其相连三元组(文本化)。 // 错误如实返回:Neo4j 不可用/查询失败以前一律回 nil,与"图谱里没有相关实体"无法区分。 func (g *graphStore) search(ctx context.Context, kb, query string, limit int) ([]Hit, error) { if !g.ready() || query == "" { return nil, nil } // 匹配两路:① 查询整体含实体名($q CONTAINS);② 查询的字符 n-gram 与实体名互为子串 // —— 解决"查询说'星云一号'、实体名抽成'星云一号卫星'"这类后缀错配(纯 $q CONTAINS 会漏)。 res, err := neo4j.ExecuteQuery(ctx, g.driver, `MATCH (a:Entity {kb:$kb})-[r:REL]->(b:Entity {kb:$kb}) WHERE $q CONTAINS a.name OR $q CONTAINS b.name OR any(ng IN $ngrams WHERE a.name CONTAINS ng OR b.name CONTAINS ng) RETURN a.name AS s, r.type AS p, b.name AS o LIMIT $k`, map[string]any{"kb": kb, "q": query, "ngrams": queryNgrams(query), "k": limit}, neo4j.EagerResultTransformer, neo4j.ExecuteQueryWithDatabase("neo4j")) if err != nil { return nil, err } var hits []Hit for _, rec := range res.Records { s, _ := rec.Get("s") p, _ := rec.Get("p") o, _ := rec.Get("o") hits = append(hits, Hit{Text: fmt.Sprintf("%v —%v→ %v", s, p, o), Score: 1}) } return hits, nil } // queryNgrams 生成查询的字符 n-gram(长度 2..8,去重并截断)——用于和图谱实体名做双向子串匹配, // 解决实体名比查询里提到的更长/更短的错配(中文无词边界,按字符 n-gram 是务实做法)。 func queryNgrams(q string) []string { r := []rune(q) const minLen, maxLen, cap = 2, 8, 80 seen := make(map[string]bool) out := make([]string, 0, cap) for i := 0; i < len(r); i++ { for l := minLen; l <= maxLen && i+l <= len(r); l++ { s := string(r[i : i+l]) if !seen[s] { seen[s] = true out = append(out, s) if len(out) >= cap { return out } } } } return out } // triples 返回某 kb 的全部三元组(供 UI 图谱可视化)。 func (g *graphStore) triples(ctx context.Context, kb string, limit int) []Triple { if !g.ready() { return nil } res, err := neo4j.ExecuteQuery(ctx, g.driver, `MATCH (a:Entity {kb:$kb})-[r:REL]->(b:Entity {kb:$kb}) RETURN a.name AS s, r.type AS p, b.name AS o LIMIT $k`, map[string]any{"kb": kb, "k": limit}, neo4j.EagerResultTransformer, neo4j.ExecuteQueryWithDatabase("neo4j")) if err != nil { return nil } var out []Triple for _, rec := range res.Records { s, _ := rec.Get("s") p, _ := rec.Get("p") o, _ := rec.Get("o") out = append(out, Triple{S: fmt.Sprint(s), P: fmt.Sprint(p), O: fmt.Sprint(o)}) } return out } // extractTriples 用 LLM 从文本抽取知识三元组。 func extractTriples(ctx context.Context, chat *chatClient, text string) ([]Triple, error) { if !chat.ready() { return nil, nil } out, err := chat.complete(ctx, prompts.Get(prompts.GraphExtract), text) if err != nil { return nil, err } return parseTriples(out), nil } // parseTriples 容忍代码块/前后噪声地解析三元组 JSON。 func parseTriples(s string) []Triple { s = strings.TrimSpace(s) s = strings.TrimPrefix(s, "```json") s = strings.TrimPrefix(s, "```") s = strings.TrimSuffix(s, "```") if i := strings.Index(s, "["); i >= 0 { if j := strings.LastIndex(s, "]"); j > i { s = s[i : j+1] } } var triples []Triple if json.Unmarshal([]byte(s), &triples) != nil { return nil } return triples }