Files
HomeAgent/internal/memory/graph.go

744 lines
18 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package memory
import (
"database/sql"
"fmt"
"sync"
"time"
_ "modernc.org/sqlite"
)
type Entity struct {
ID int64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
MentionCount int `json:"mention_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type Relation struct {
ID int64 `json:"id"`
SourceID int64 `json:"source_id"`
TargetID int64 `json:"target_id"`
SourceName string `json:"source_name"`
TargetName string `json:"target_name"`
RelationType string `json:"relation_type"`
Confidence float64 `json:"confidence"`
Status string `json:"status"`
SessionID string `json:"session_id"`
TurnID int `json:"turn_id"`
CreatedAt time.Time `json:"created_at"`
DateBucket string `json:"date_bucket"`
}
type Triple struct {
Subject string `json:"subject"`
Relation string `json:"relation"`
Object string `json:"object"`
Confidence float64 `json:"confidence,omitempty"`
SubjectType string `json:"subject_type,omitempty"`
ObjectType string `json:"object_type,omitempty"`
}
type GraphDB struct {
db *sql.DB
mu sync.RWMutex
dbPath string
}
func NewGraphDB(dbPath string) (*GraphDB, error) {
db, err := sql.Open("sqlite", dbPath+"?_journal_mode=WAL&_foreign_keys=on")
if err != nil {
return nil, fmt.Errorf("open graph db: %w", err)
}
g := &GraphDB{db: db, dbPath: dbPath}
if err := g.initSchema(); err != nil {
return nil, fmt.Errorf("init schema: %w", err)
}
return g, nil
}
func (g *GraphDB) initSchema() error {
g.mu.Lock()
defer g.mu.Unlock()
tx, err := g.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
schemas := []string{
`CREATE TABLE IF NOT EXISTS entities (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT UNIQUE NOT NULL,
type TEXT DEFAULT 'Concept',
mention_count INTEGER DEFAULT 1,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`,
`CREATE TABLE IF NOT EXISTS relations (
id INTEGER PRIMARY KEY AUTOINCREMENT,
source_id INTEGER NOT NULL,
target_id INTEGER NOT NULL,
relation_type TEXT NOT NULL,
confidence REAL DEFAULT 1.0,
status TEXT DEFAULT 'active',
session_id TEXT,
turn_id INTEGER DEFAULT 0,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
date_bucket TEXT,
FOREIGN KEY (source_id) REFERENCES entities(id),
FOREIGN KEY (target_id) REFERENCES entities(id)
)`,
`CREATE INDEX IF NOT EXISTS idx_entity_name ON entities(name)`,
`CREATE INDEX IF NOT EXISTS idx_entity_type ON entities(type)`,
`CREATE INDEX IF NOT EXISTS idx_relation_source ON relations(source_id)`,
`CREATE INDEX IF NOT EXISTS idx_relation_target ON relations(target_id)`,
`CREATE INDEX IF NOT EXISTS idx_relation_type ON relations(relation_type)`,
`CREATE INDEX IF NOT EXISTS idx_relation_status ON relations(status)`,
`CREATE INDEX IF NOT EXISTS idx_relation_session ON relations(session_id)`,
}
for _, s := range schemas {
if _, err := tx.Exec(s); err != nil {
return fmt.Errorf("schema exec: %w", err)
}
}
return tx.Commit()
}
func (g *GraphDB) Commit(triples []Triple, sessionID string, turnID int) (int, int, error) {
g.mu.Lock()
defer g.mu.Unlock()
tx, err := g.db.Begin()
if err != nil {
return 0, 0, err
}
defer tx.Rollback()
entitiesCreated := 0
relationsCreated := 0
dateBucket := time.Now().Format("2006-01-02")
for _, t := range triples {
if t.Subject == "" || t.Relation == "" || t.Object == "" {
continue
}
subjType := t.SubjectType
if subjType == "" {
subjType = "Concept"
}
objType := t.ObjectType
if objType == "" {
objType = "Concept"
}
confidence := t.Confidence
if confidence == 0 {
confidence = 1.0
}
ec, err := g.upsertEntity(tx, t.Subject, subjType)
if err != nil {
return 0, 0, err
}
entitiesCreated += ec
ec, err = g.upsertEntity(tx, t.Object, objType)
if err != nil {
return 0, 0, err
}
entitiesCreated += ec
var sourceID, targetID int64
err = tx.QueryRow("SELECT id FROM entities WHERE name = ?", t.Subject).Scan(&sourceID)
if err != nil {
return 0, 0, err
}
err = tx.QueryRow("SELECT id FROM entities WHERE name = ?", t.Object).Scan(&targetID)
if err != nil {
return 0, 0, err
}
_, err = tx.Exec(
`INSERT INTO relations (source_id, target_id, relation_type, confidence, session_id, turn_id, date_bucket)
VALUES (?, ?, ?, ?, ?, ?, ?)`,
sourceID, targetID, t.Relation, confidence, sessionID, turnID, dateBucket,
)
if err != nil {
return 0, 0, err
}
relationsCreated++
}
if err := tx.Commit(); err != nil {
return 0, 0, err
}
return entitiesCreated, relationsCreated, nil
}
func (g *GraphDB) upsertEntity(tx *sql.Tx, name string, entityType string) (int, error) {
result, err := tx.Exec(
`INSERT INTO entities (name, type) VALUES (?, ?)
ON CONFLICT(name) DO UPDATE SET
mention_count = mention_count + 1,
updated_at = CURRENT_TIMESTAMP`,
name, entityType,
)
if err != nil {
return 0, err
}
rows, _ := result.RowsAffected()
if rows > 0 {
return 1, nil
}
return 0, nil
}
type RecallResult struct {
Entities []Entity `json:"entities"`
Relations []Relation `json:"relations"`
}
func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, sessionFilter string) (*RecallResult, error) {
g.mu.RLock()
defer g.mu.RUnlock()
result := &RecallResult{}
if len(keywords) == 0 && len(seedEntities) == 0 {
rows, err := g.db.Query(
`SELECT id, name, type, mention_count, created_at, updated_at
FROM entities ORDER BY mention_count DESC LIMIT 50`,
)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var e Entity
if err := rows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil {
return nil, err
}
result.Entities = append(result.Entities, e)
}
relRows, err := g.db.Query(
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
r.relation_type, r.confidence, r.status, r.session_id,
r.turn_id, r.created_at, COALESCE(r.date_bucket, '')
FROM relations r
JOIN entities e1 ON r.source_id = e1.id
JOIN entities e2 ON r.target_id = e2.id
WHERE r.status = 'active'
ORDER BY r.created_at DESC LIMIT 30`,
)
if err != nil {
return nil, err
}
defer relRows.Close()
for relRows.Next() {
var rel Relation
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
&rel.SourceName, &rel.TargetName, &rel.RelationType,
&rel.Confidence, &rel.Status, &rel.SessionID,
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
return nil, err
}
result.Relations = append(result.Relations, rel)
}
return result, nil
}
entityIDs := make(map[int64]bool)
for _, kw := range keywords {
rows, err := g.db.Query(
`SELECT id, name, type, mention_count, created_at, updated_at
FROM entities WHERE LOWER(name) LIKE ?`,
"%"+kw+"%",
)
if err != nil {
return nil, err
}
for rows.Next() {
var e Entity
if err := rows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil {
rows.Close()
return nil, err
}
if !entityIDs[e.ID] {
entityIDs[e.ID] = true
result.Entities = append(result.Entities, e)
}
}
rows.Close()
}
for _, se := range seedEntities {
row := g.db.QueryRow(
`SELECT id, name, type, mention_count, created_at, updated_at
FROM entities WHERE name = ?`, se)
var e Entity
if err := row.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err == nil {
if !entityIDs[e.ID] {
entityIDs[e.ID] = true
result.Entities = append(result.Entities, e)
}
}
}
if len(entityIDs) == 0 {
return result, nil
}
for depthLevel := 0; depthLevel < depth; depthLevel++ {
ids := make([]interface{}, 0, len(entityIDs))
for id := range entityIDs {
ids = append(ids, id)
}
if len(ids) == 0 {
break
}
query := fmt.Sprintf(
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
r.relation_type, r.confidence, r.status, r.session_id,
r.turn_id, r.created_at, COALESCE(r.date_bucket, '')
FROM relations r
JOIN entities e1 ON r.source_id = e1.id
JOIN entities e2 ON r.target_id = e2.id
WHERE (r.source_id IN (%s) OR r.target_id IN (%s))
AND r.status = 'active'`,
placeholders(len(ids)),
placeholders(len(ids)),
)
allIDs := append(ids, ids...)
if sessionFilter != "" {
query += " AND r.session_id = ?"
allIDs = append(allIDs, sessionFilter)
}
relRows, err := g.db.Query(query, allIDs...)
if err != nil {
return nil, err
}
newIDs := make(map[int64]bool)
for relRows.Next() {
var rel Relation
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
&rel.SourceName, &rel.TargetName, &rel.RelationType,
&rel.Confidence, &rel.Status, &rel.SessionID,
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
relRows.Close()
return nil, err
}
result.Relations = append(result.Relations, rel)
if !entityIDs[rel.SourceID] {
newIDs[rel.SourceID] = true
}
if !entityIDs[rel.TargetID] {
newIDs[rel.TargetID] = true
}
}
relRows.Close()
if len(newIDs) == 0 {
break
}
ids2 := make([]interface{}, 0, len(newIDs))
for id := range newIDs {
ids2 = append(ids2, id)
}
eRows, err := g.db.Query(
fmt.Sprintf(
`SELECT id, name, type, mention_count, created_at, updated_at
FROM entities WHERE id IN (%s)`, placeholders(len(ids2))),
ids2...,
)
if err != nil {
return nil, err
}
for eRows.Next() {
var e Entity
if err := eRows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil {
eRows.Close()
return nil, err
}
if !entityIDs[e.ID] {
entityIDs[e.ID] = true
result.Entities = append(result.Entities, e)
}
}
eRows.Close()
for id := range newIDs {
entityIDs[id] = true
}
}
return result, nil
}
func (g *GraphDB) Purge(criteria map[string]string, mode string) (int, error) {
g.mu.Lock()
defer g.mu.Unlock()
conds := []string{"status = 'active'"}
args := []interface{}{}
if v, ok := criteria["subject_contains"]; ok {
rows, err := g.db.Query("SELECT id FROM entities WHERE name LIKE ?", "%"+v+"%")
if err != nil {
return 0, err
}
var ids []interface{}
for rows.Next() {
var id int64
rows.Scan(&id)
ids = append(ids, id)
}
rows.Close()
if len(ids) > 0 {
conds = append(conds, fmt.Sprintf("source_id IN (%s)", placeholders(len(ids))))
args = append(args, ids...)
}
}
if v, ok := criteria["target_contains"]; ok {
rows, err := g.db.Query("SELECT id FROM entities WHERE name LIKE ?", "%"+v+"%")
if err != nil {
return 0, err
}
var ids []interface{}
for rows.Next() {
var id int64
rows.Scan(&id)
ids = append(ids, id)
}
rows.Close()
if len(ids) > 0 {
conds = append(conds, fmt.Sprintf("target_id IN (%s)", placeholders(len(ids))))
args = append(args, ids...)
}
}
if v, ok := criteria["relation_type"]; ok {
conds = append(conds, "relation_type = ?")
args = append(args, v)
}
if v, ok := criteria["session_id"]; ok {
conds = append(conds, "session_id = ?")
args = append(args, v)
}
if len(conds) == 1 {
return 0, fmt.Errorf("no criteria provided")
}
where := ""
for i, c := range conds {
if i == 0 {
where = c
} else {
where += " AND " + c
}
}
if mode == "hard" {
result, err := g.db.Exec(
fmt.Sprintf(`DELETE FROM relations WHERE %s`, where), args...)
if err != nil {
return 0, err
}
n, _ := result.RowsAffected()
g.db.Exec(`DELETE FROM entities WHERE id NOT IN (
SELECT DISTINCT source_id FROM relations
UNION SELECT DISTINCT target_id FROM relations)`)
return int(n), nil
}
result, err := g.db.Exec(
fmt.Sprintf(`UPDATE relations SET status = 'deleted', updated_at = CURRENT_TIMESTAMP WHERE %s`, where),
args...,
)
if err != nil {
return 0, err
}
n, _ := result.RowsAffected()
return int(n), nil
}
func (g *GraphDB) GraphData() (map[string]interface{}, error) {
g.mu.RLock()
defer g.mu.RUnlock()
rows, err := g.db.Query(`SELECT id, name, type, mention_count, created_at, updated_at FROM entities ORDER BY mention_count DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
type graphEntity struct {
ID int64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
MentionCount int `json:"mention_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
var entities []graphEntity
for rows.Next() {
var e graphEntity
if err := rows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil {
return nil, err
}
entities = append(entities, e)
}
if err := rows.Err(); err != nil {
return nil, err
}
rrows, err := g.db.Query(`SELECT id, source_id, target_id, relation_type, confidence, status, created_at FROM relations WHERE status = 'active' ORDER BY created_at DESC`)
if err != nil {
return nil, err
}
defer rrows.Close()
type graphRelation struct {
ID int64 `json:"id"`
SourceID int64 `json:"source_id"`
TargetID int64 `json:"target_id"`
RelationType string `json:"relation_type"`
Confidence float64 `json:"confidence"`
Status string `json:"status"`
CreatedAt time.Time `json:"created_at"`
}
var relations []graphRelation
for rrows.Next() {
var r graphRelation
if err := rrows.Scan(&r.ID, &r.SourceID, &r.TargetID, &r.RelationType, &r.Confidence, &r.Status, &r.CreatedAt); err != nil {
return nil, err
}
relations = append(relations, r)
}
if err := rrows.Err(); err != nil {
return nil, err
}
return map[string]interface{}{
"nodes": entities,
"edges": relations,
}, nil
}
func (g *GraphDB) Introspect() (map[string]interface{}, error) {
g.mu.RLock()
defer g.mu.RUnlock()
var entityCount, relationCount int
g.db.QueryRow("SELECT COUNT(*) FROM entities").Scan(&entityCount)
g.db.QueryRow("SELECT COUNT(*) FROM relations WHERE status = 'active'").Scan(&relationCount)
hotspots := []map[string]interface{}{}
rows, err := g.db.Query(
`SELECT name, mention_count, type FROM entities ORDER BY mention_count DESC LIMIT 10`,
)
if err == nil {
defer rows.Close()
for rows.Next() {
var name, etype string
var count int
if err := rows.Scan(&name, &count, &etype); err == nil {
hotspots = append(hotspots, map[string]interface{}{
"name": name, "count": count, "type": etype,
})
}
}
}
return map[string]interface{}{
"entity_count": entityCount,
"relation_count": relationCount,
"memory_hotspots": hotspots,
}, nil
}
// MergeEntities 合并两个实体:将 sourceName 的所有信息合并到 targetName
// 1. sourceName 的所有关系重新指向 targetName
// 2. targetName 的 mention_count 增加 sourceName 的计数
// 3. sourceName 彻底删除(不再残留 @merged_ 实体)
// 返回 (关系的重定向数, error)
func (g *GraphDB) MergeEntities(sourceName, targetName string) (int, error) {
g.mu.Lock()
defer g.mu.Unlock()
tx, err := g.db.Begin()
if err != nil {
return 0, err
}
defer tx.Rollback()
var sourceID, targetID int64
var sourceCount, targetCount int
err = tx.QueryRow("SELECT id, mention_count FROM entities WHERE name = ?", sourceName).Scan(&sourceID, &sourceCount)
if err != nil {
return 0, fmt.Errorf("source entity '%s' not found: %w", sourceName, err)
}
err = tx.QueryRow("SELECT id, mention_count FROM entities WHERE name = ?", targetName).Scan(&targetID, &targetCount)
if err != nil {
return 0, fmt.Errorf("target entity '%s' not found: %w", targetName, err)
}
if sourceID == targetID {
return 0, fmt.Errorf("cannot merge entity with itself")
}
// 重定向 source → target 的活跃关系(作为 source
res, err := tx.Exec(
`UPDATE relations SET source_id = ?, updated_at = CURRENT_TIMESTAMP
WHERE source_id = ? AND status = 'active'`,
targetID, sourceID,
)
if err != nil {
return 0, err
}
redirectedSource, _ := res.RowsAffected()
// 重定向 source → target 的活跃关系(作为 target
res, err = tx.Exec(
`UPDATE relations SET target_id = ?, updated_at = CURRENT_TIMESTAMP
WHERE target_id = ? AND status = 'active'`,
targetID, sourceID,
)
if err != nil {
return 0, err
}
redirectedTarget, _ := res.RowsAffected()
// 删除可能产生的自引用关系
_, err = tx.Exec(
`DELETE FROM relations
WHERE source_id = target_id AND source_id = ?`,
targetID,
)
if err != nil {
return 0, err
}
// 清理 source 残留的非活跃关系archived/deleted否则外键约束阻止删除实体
_, err = tx.Exec(`DELETE FROM relations WHERE source_id = ? OR target_id = ?`, sourceID, sourceID)
if err != nil {
return 0, err
}
// 更新 target 的 mention_count
_, err = tx.Exec(
`UPDATE entities SET mention_count = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
targetCount+sourceCount, targetID,
)
if err != nil {
return 0, err
}
// 彻底删除 source 实体(所有关系已重定向,自引用已删除)
_, err = tx.Exec(`DELETE FROM entities WHERE id = ?`, sourceID)
if err != nil {
return 0, err
}
if err := tx.Commit(); err != nil {
return 0, err
}
total := int(redirectedSource + redirectedTarget)
return total, nil
}
// DeleteEntity 彻底删除一个实体及其所有关联关系。
func (g *GraphDB) DeleteEntity(name string) error {
g.mu.Lock()
defer g.mu.Unlock()
tx, err := g.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var id int64
err = tx.QueryRow("SELECT id FROM entities WHERE name = ?", name).Scan(&id)
if err != nil {
return fmt.Errorf("entity '%s' not found: %w", name, err)
}
_, err = tx.Exec(`DELETE FROM relations WHERE source_id = ? OR target_id = ?`, id, id)
if err != nil {
return err
}
_, err = tx.Exec(`DELETE FROM entities WHERE id = ?`, id)
if err != nil {
return err
}
return tx.Commit()
}
func (g *GraphDB) Archive(days int) (int, error) {
g.mu.Lock()
defer g.mu.Unlock()
result, err := g.db.Exec(
`UPDATE relations SET status = 'archived', updated_at = CURRENT_TIMESTAMP
WHERE status = 'active' AND created_at < datetime('now', ?)`,
fmt.Sprintf("-%d days", days),
)
if err != nil {
return 0, err
}
n, _ := result.RowsAffected()
return int(n), nil
}
func (g *GraphDB) Close() error {
return g.db.Close()
}
func placeholders(n int) string {
if n <= 0 {
return "NULL"
}
b := make([]byte, 0, n*2-1)
for i := 0; i < n; i++ {
if i > 0 {
b = append(b, ',')
}
b = append(b, '?')
}
return string(b)
}