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