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
|
}
|