package index import ( "database/sql" "encoding/json" "fmt" "os" "path/filepath" "sort" "strings" _ "github.com/mattn/go-sqlite3" "github.com/aisim/kb-cli/internal/graph" ) // Store SQLite 存储层 type Store struct { db *sql.DB } // Open 打开或创建数据库 func Open(dbPath string) (*Store, error) { // 确保目录存在 dir := filepath.Dir(dbPath) if err := os.MkdirAll(dir, 0755); err != nil { return nil, fmt.Errorf("创建目录失败: %w", err) } db, err := sql.Open("sqlite3", dbPath+"?_journal_mode=WAL") if err != nil { return nil, fmt.Errorf("打开数据库失败: %w", err) } s := &Store{db: db} if err := s.initTables(); err != nil { db.Close() return nil, err } return s, nil } // Close 关闭数据库 func (s *Store) Close() error { return s.db.Close() } // initTables 创建表结构 func (s *Store) initTables() error { schema := ` CREATE TABLE IF NOT EXISTS meta ( key TEXT PRIMARY KEY, value TEXT NOT NULL ); CREATE TABLE IF NOT EXISTS nodes ( id INTEGER PRIMARY KEY AUTOINCREMENT, path TEXT NOT NULL UNIQUE, title TEXT, section TEXT, tags TEXT, entities TEXT, wikilinks TEXT, content_fts TEXT, created_at TEXT DEFAULT (datetime('now')), updated_at TEXT DEFAULT (datetime('now')) ); CREATE TABLE IF NOT EXISTS edges ( id INTEGER PRIMARY KEY AUTOINCREMENT, from_node INTEGER NOT NULL, to_node INTEGER NOT NULL, relation TEXT NOT NULL, label TEXT, UNIQUE(from_node, to_node, relation, label) ); CREATE INDEX IF NOT EXISTS idx_nodes_section ON nodes(section); CREATE INDEX IF NOT EXISTS idx_edges_from ON edges(from_node); CREATE INDEX IF NOT EXISTS idx_edges_to ON edges(to_node); CREATE INDEX IF NOT EXISTS idx_edges_label ON edges(label); ` _, err := s.db.Exec(schema) if err != nil { return fmt.Errorf("创建表失败: %w", err) } return nil } // ClearData 清空数据(重建前调用) func (s *Store) ClearData() error { _, err := s.db.Exec("DELETE FROM edges; DELETE FROM nodes;") return err } // InsertNode 插入节点 func (s *Store) InsertNode(n *graph.Node) (int64, error) { tagsJSON, _ := json.Marshal(n.Tags) entitiesJSON, _ := json.Marshal(n.Entities) wikilinksJSON, _ := json.Marshal(n.Wikilinks) result, err := s.db.Exec(` INSERT INTO nodes (path, title, section, tags, entities, wikilinks, content_fts) VALUES (?, ?, ?, ?, ?, ?, ?) `, n.Path, n.Title, n.Section, string(tagsJSON), string(entitiesJSON), string(wikilinksJSON), n.Content) if err != nil { return 0, err } return result.LastInsertId() } // InsertEdge 插入边 func (s *Store) InsertEdge(e *graph.Edge) error { _, err := s.db.Exec(` INSERT OR IGNORE INTO edges (from_node, to_node, relation, label) VALUES (?, ?, ?, ?) `, e.FromNode, e.ToNode, e.Relation, e.Label) return err } // SetMeta 设置元信息 func (s *Store) SetMeta(key, value string) error { _, err := s.db.Exec(` INSERT OR REPLACE INTO meta (key, value) VALUES (?, ?) `, key, value) return err } // GetMeta 获取元信息 func (s *Store) GetMeta(key string) (string, error) { var value string err := s.db.QueryRow("SELECT value FROM meta WHERE key = ?", key).Scan(&value) if err == sql.ErrNoRows { return "", nil } return value, err } // NodeCount 返回节点数量 func (s *Store) NodeCount() (int, error) { var count int err := s.db.QueryRow("SELECT COUNT(*) FROM nodes").Scan(&count) return count, err } // EdgeCount 返回边数量 func (s *Store) EdgeCount() (int, error) { var count int err := s.db.QueryRow("SELECT COUNT(*) FROM edges").Scan(&count) return count, err } // GetNodeLinks 获取节点的关联链接(wikilink 目标) func (s *Store) GetNodeLinks(nodeID int64) ([]string, error) { query := ` SELECT n.path FROM edges e JOIN nodes n ON n.id = e.to_node WHERE e.from_node = ? AND e.relation = 'wikilink' ` rows, err := s.db.Query(query, nodeID) if err != nil { return nil, fmt.Errorf("查询链接失败: %w", err) } defer rows.Close() var links []string for rows.Next() { var path string if err := rows.Scan(&path); err != nil { return nil, fmt.Errorf("扫描链接失败: %w", err) } links = append(links, path) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("遍历链接失败: %w", err) } return links, nil } // NodeInfo 节点基本信息(用于 GC) type NodeInfo struct { ID int64 Path string } // GetAllNodes 获取所有节点(用于 GC 检查) func (s *Store) GetAllNodes() ([]NodeInfo, error) { rows, err := s.db.Query("SELECT id, path FROM nodes") if err != nil { return nil, fmt.Errorf("查询节点失败: %w", err) } defer rows.Close() var nodes []NodeInfo for rows.Next() { var n NodeInfo if err := rows.Scan(&n.ID, &n.Path); err != nil { return nil, fmt.Errorf("扫描节点失败: %w", err) } nodes = append(nodes, n) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("遍历节点失败: %w", err) } return nodes, nil } // DeleteNodesByPaths 删除指定路径的节点及其关联边 func (s *Store) DeleteNodesByPaths(paths []string) (int, error) { if len(paths) == 0 { return 0, nil } // 构建 IN 子句 placeholders := make([]string, len(paths)) args := make([]interface{}, len(paths)) for i, path := range paths { placeholders[i] = "?" args[i] = path } tx, err := s.db.Begin() if err != nil { return 0, fmt.Errorf("开始事务失败: %w", err) } // 先删除关联边 query := fmt.Sprintf(` DELETE FROM edges WHERE from_node IN (SELECT id FROM nodes WHERE path IN (%s)) OR to_node IN (SELECT id FROM nodes WHERE path IN (%s)) `, strings.Join(placeholders, ","), strings.Join(placeholders, ",")) // 参数需要重复两次 allArgs := append(args, args...) _, err = tx.Exec(query, allArgs...) if err != nil { tx.Rollback() return 0, fmt.Errorf("删除边失败: %w", err) } // 删除节点 query = fmt.Sprintf("DELETE FROM nodes WHERE path IN (%s)", strings.Join(placeholders, ",")) result, err := tx.Exec(query, args...) if err != nil { tx.Rollback() return 0, fmt.Errorf("删除节点失败: %w", err) } if err := tx.Commit(); err != nil { return 0, fmt.Errorf("提交事务失败: %w", err) } affected, _ := result.RowsAffected() return int(affected), nil } // GetRelationStats 获取关系类型分布统计 func (s *Store) GetRelationStats() (map[string]int, error) { rows, err := s.db.Query(` SELECT relation, COUNT(*) as count FROM edges GROUP BY relation ORDER BY count DESC `) if err != nil { return nil, fmt.Errorf("查询关系统计失败: %w", err) } defer rows.Close() stats := make(map[string]int) for rows.Next() { var relation string var count int if err := rows.Scan(&relation, &count); err != nil { return nil, fmt.Errorf("扫描关系统计失败: %w", err) } stats[relation] = count } if err := rows.Err(); err != nil { return nil, fmt.Errorf("遍历关系统计失败: %w", err) } return stats, nil } // FindNodesByKeyword 根据关键词查找节点(模糊匹配标题、路径、标签、实体) func (s *Store) FindNodesByKeyword(keyword string) ([]*graph.Node, error) { keyword = "%" + keyword + "%" query := ` SELECT id, path, title, section, tags, entities, wikilinks, content_fts FROM nodes WHERE title LIKE ? OR path LIKE ? OR tags LIKE ? OR entities LIKE ? LIMIT 20 ` rows, err := s.db.Query(query, keyword, keyword, keyword, keyword) if err != nil { return nil, fmt.Errorf("查询节点失败: %w", err) } defer rows.Close() var nodes []*graph.Node for rows.Next() { var n graph.Node var tagsJSON, entitiesJSON, wikilinksJSON string var content sql.NullString err := rows.Scan(&n.ID, &n.Path, &n.Title, &n.Section, &tagsJSON, &entitiesJSON, &wikilinksJSON, &content) if err != nil { return nil, fmt.Errorf("扫描节点失败: %w", err) } json.Unmarshal([]byte(tagsJSON), &n.Tags) json.Unmarshal([]byte(entitiesJSON), &n.Entities) json.Unmarshal([]byte(wikilinksJSON), &n.Wikilinks) if content.Valid { n.Content = content.String } nodes = append(nodes, &n) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("遍历节点失败: %w", err) } return nodes, nil } // GetNodeEdges 获取节点的关联边(支持关系类型过滤和深度查询) func (s *Store) GetNodeEdges(nodeID int64, relation string, depth int) ([]*graph.Edge, error) { if depth < 1 { depth = 1 } if depth > 3 { depth = 3 } var edges []*graph.Edge visited := make(map[int64]bool) currentLevel := []int64{nodeID} for d := 0; d < depth; d++ { var nextLevel []int64 for _, id := range currentLevel { if visited[id] && d > 0 { continue } visited[id] = true query := ` SELECT id, from_node, to_node, relation, label FROM edges WHERE from_node = ? ` args := []interface{}{id} if relation != "" { query += " AND relation = ?" args = append(args, relation) } rows, err := s.db.Query(query, args...) if err != nil { return nil, fmt.Errorf("查询边失败: %w", err) } for rows.Next() { var e graph.Edge var edgeID int64 if err := rows.Scan(&edgeID, &e.FromNode, &e.ToNode, &e.Relation, &e.Label); err != nil { rows.Close() return nil, fmt.Errorf("扫描边失败: %w", err) } edges = append(edges, &e) nextLevel = append(nextLevel, e.ToNode) } rows.Close() } currentLevel = nextLevel } return edges, nil } // FindRelatedNodes 查找与关键词相关的所有节点(通过共同标签、实体、wikilink) func (s *Store) FindRelatedNodes(keyword string, topN int) ([]RelatedNode, error) { // 先找到匹配关键词的节点 keyword = "%" + keyword + "%" query := ` SELECT id, path, title, section, tags, entities FROM nodes WHERE title LIKE ? OR path LIKE ? OR tags LIKE ? OR entities LIKE ? LIMIT 5 ` rows, err := s.db.Query(query, keyword, keyword, keyword, keyword) if err != nil { return nil, fmt.Errorf("查询节点失败: %w", err) } defer rows.Close() var seedNodes []*graph.Node for rows.Next() { var n graph.Node var tagsJSON, entitiesJSON string err := rows.Scan(&n.ID, &n.Path, &n.Title, &n.Section, &tagsJSON, &entitiesJSON) if err != nil { return nil, fmt.Errorf("扫描节点失败: %w", err) } json.Unmarshal([]byte(tagsJSON), &n.Tags) json.Unmarshal([]byte(entitiesJSON), &n.Entities) seedNodes = append(seedNodes, &n) } if len(seedNodes) == 0 { return []RelatedNode{}, nil } // 收集种子节点的标签和实体 tagSet := make(map[string]bool) entitySet := make(map[string]bool) for _, n := range seedNodes { for _, tag := range n.Tags { tagSet[tag] = true } for _, entity := range n.Entities { entitySet[entity] = true } } // 查找共享标签或实体的节点 relevanceMap := make(map[int64]int) pathMap := make(map[int64]string) titleMap := make(map[int64]string) sectionMap := make(map[int64]string) tagsMap := make(map[int64][]string) // 通过标签查找 for tag := range tagSet { tagPattern := "%" + tag + "%" rows, err := s.db.Query(` SELECT id, path, title, section, tags FROM nodes WHERE tags LIKE ? `, tagPattern) if err != nil { continue } for rows.Next() { var id int64 var path, title, section, tagsJSON string if err := rows.Scan(&id, &path, &title, §ion, &tagsJSON); err != nil { rows.Close() continue } relevanceMap[id]++ pathMap[id] = path titleMap[id] = title sectionMap[id] = section var tags []string json.Unmarshal([]byte(tagsJSON), &tags) tagsMap[id] = tags } rows.Close() } // 通过实体查找 for entity := range entitySet { entityPattern := "%" + entity + "%" rows, err := s.db.Query(` SELECT id, path, title, section, tags FROM nodes WHERE entities LIKE ? `, entityPattern) if err != nil { continue } for rows.Next() { var id int64 var path, title, section, tagsJSON string if err := rows.Scan(&id, &path, &title, §ion, &tagsJSON); err != nil { rows.Close() continue } relevanceMap[id]++ pathMap[id] = path titleMap[id] = title sectionMap[id] = section var tags []string json.Unmarshal([]byte(tagsJSON), &tags) tagsMap[id] = tags } rows.Close() } // 转换为结果列表 var results []RelatedNode for id, relevance := range relevanceMap { results = append(results, RelatedNode{ ID: id, Path: pathMap[id], Title: titleMap[id], Section: sectionMap[id], Tags: tagsMap[id], Relevance: relevance, }) } // 按相关度排序 sort.Slice(results, func(i, j int) bool { return results[i].Relevance > results[j].Relevance }) // 限制返回数量 if topN > 0 && len(results) > topN { results = results[:topN] } return results, nil } // RelatedNode 相关节点 type RelatedNode struct { ID int64 `json:"id"` Path string `json:"path"` Title string `json:"title"` Section string `json:"section"` Tags []string `json:"tags"` Relevance int `json:"relevance"` }