package index import ( "database/sql" "encoding/json" "fmt" "os" "path/filepath" "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 }