mirror of
https://gitcode.com/JianFeeeee/TrulyMEM-TrueHumanMEM.git
synced 2026-09-20 00:48:52 +00:00
715 lines
26 KiB
Python
715 lines
26 KiB
Python
"""
|
||
内嵌图数据库 - 基于SQLite实现
|
||
无需Docker,开箱即用
|
||
"""
|
||
|
||
import sqlite3
|
||
import hashlib
|
||
import json
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
from typing import List, Dict, Optional, Any
|
||
|
||
|
||
class EmbeddedGraphDB:
|
||
"""内嵌图数据库 - SQLite实现"""
|
||
|
||
def __init__(self, db_path: str = "graph_memory.db"):
|
||
"""
|
||
初始化数据库
|
||
|
||
Args:
|
||
db_path: 数据库文件路径
|
||
"""
|
||
self.db_path = Path(db_path)
|
||
self.conn = None
|
||
self._init_db()
|
||
|
||
def _init_db(self):
|
||
"""初始化数据库表"""
|
||
self.conn = sqlite3.connect(str(self.db_path), check_same_thread=False)
|
||
self.conn.row_factory = sqlite3.Row
|
||
|
||
cursor = self.conn.cursor()
|
||
|
||
# 创建实体表
|
||
cursor.execute("""
|
||
CREATE TABLE IF NOT EXISTS entities (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
name TEXT UNIQUE NOT NULL,
|
||
type TEXT,
|
||
mention_count INTEGER DEFAULT 1,
|
||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||
)
|
||
""")
|
||
|
||
# 创建关系表
|
||
cursor.execute("""
|
||
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,
|
||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||
date_bucket TEXT,
|
||
superseded_by INTEGER,
|
||
FOREIGN KEY (source_id) REFERENCES entities(id),
|
||
FOREIGN KEY (target_id) REFERENCES entities(id)
|
||
)
|
||
""")
|
||
|
||
# 创建索引
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_entity_name ON entities(name)")
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_entity_type ON entities(type)")
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_relation_source ON relations(source_id)")
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_relation_target ON relations(target_id)")
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_relation_type ON relations(relation_type)")
|
||
cursor.execute("CREATE INDEX IF NOT EXISTS idx_relation_status ON relations(status)")
|
||
|
||
cursor.execute("""
|
||
SELECT name FROM sqlite_master
|
||
WHERE type='table' AND name='chat_records'
|
||
""")
|
||
if not cursor.fetchone():
|
||
cursor.execute("""
|
||
CREATE TABLE chat_records (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
role TEXT NOT NULL,
|
||
content TEXT NOT NULL,
|
||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||
)
|
||
""")
|
||
cursor.execute("CREATE INDEX idx_chat_created ON chat_records(created_at)")
|
||
|
||
# 创建 Web 用户表(支持多用户隔离)
|
||
cursor.execute("""
|
||
CREATE TABLE IF NOT EXISTS web_users (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
username TEXT UNIQUE NOT NULL,
|
||
password_hash TEXT NOT NULL,
|
||
role TEXT NOT NULL DEFAULT 'user',
|
||
config_path TEXT,
|
||
db_path TEXT,
|
||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||
)
|
||
""")
|
||
|
||
# 检查并添加新字段(用于旧数据库迁移)
|
||
cursor.execute("PRAGMA table_info(web_users)")
|
||
columns = [row[1] for row in cursor.fetchall()]
|
||
if 'config_path' not in columns:
|
||
cursor.execute("ALTER TABLE web_users ADD COLUMN config_path TEXT")
|
||
if 'db_path' not in columns:
|
||
cursor.execute("ALTER TABLE web_users ADD COLUMN db_path TEXT")
|
||
if 'role' not in columns:
|
||
cursor.execute("ALTER TABLE web_users ADD COLUMN role TEXT NOT NULL DEFAULT 'user'")
|
||
|
||
# 确保至少有一个 admin(当 role 列刚添加时,已有用户都是 user)
|
||
cursor.execute("SELECT COUNT(*) as cnt FROM web_users WHERE role = 'admin'")
|
||
has_admin = cursor.fetchone()[0] > 0
|
||
if not has_admin:
|
||
cursor.execute("SELECT id, username FROM web_users ORDER BY created_at ASC LIMIT 1")
|
||
first_user = cursor.fetchone()
|
||
if first_user:
|
||
cursor.execute("UPDATE web_users SET role = 'admin' WHERE id = ?", (first_user[0],))
|
||
|
||
self.conn.commit()
|
||
|
||
def ensure_constraints(self):
|
||
"""确保约束(兼容Neo4j接口)"""
|
||
pass # SQLite自动处理
|
||
|
||
def recall(self, query_intent: str, seed_entities: List[str] = None,
|
||
depth: int = 2, time_range: Dict = None,
|
||
session_filter: str = None) -> Dict:
|
||
"""
|
||
检索相关记忆
|
||
|
||
Args:
|
||
query_intent: 查询关键词(逗号分隔)
|
||
seed_entities: 种子实体
|
||
depth: 搜索深度
|
||
time_range: 时间范围
|
||
session_filter: 会话过滤
|
||
|
||
Returns:
|
||
检索结果
|
||
"""
|
||
keywords = [w.strip().lower() for w in query_intent.replace(',', ' ').split() if w.strip()]
|
||
|
||
cursor = self.conn.cursor()
|
||
|
||
# 搜索实体
|
||
entities = []
|
||
entity_ids = set()
|
||
|
||
# 如果没有关键词,返回所有实体(用于"我们都聊过什么"这类问题)
|
||
if not keywords and not seed_entities:
|
||
cursor.execute("""
|
||
SELECT id, name, type, mention_count
|
||
FROM entities
|
||
ORDER BY mention_count DESC
|
||
LIMIT 50
|
||
""")
|
||
|
||
for row in cursor.fetchall():
|
||
entity_ids.add(row['id'])
|
||
entities.append({
|
||
'name': row['name'],
|
||
'type': row['type'] or 'unknown',
|
||
'mention_count': row['mention_count']
|
||
})
|
||
else:
|
||
# 有关键词,按关键词搜索
|
||
for keyword in keywords:
|
||
cursor.execute("""
|
||
SELECT id, name, type, mention_count
|
||
FROM entities
|
||
WHERE LOWER(name) LIKE ?
|
||
""", (f"%{keyword}%",))
|
||
|
||
for row in cursor.fetchall():
|
||
if row['id'] not in entity_ids:
|
||
entity_ids.add(row['id'])
|
||
entities.append({
|
||
'name': row['name'],
|
||
'type': row['type'] or 'unknown',
|
||
'mention_count': row['mention_count']
|
||
})
|
||
|
||
# 广度优先搜索(BFS)扩展实体和关系
|
||
relations = []
|
||
visited_entity_ids = set(entity_ids) # 已访问的实体
|
||
current_layer_ids = set(entity_ids) # 当前层的实体
|
||
|
||
# 记录每个实体的深度
|
||
entity_depths = {} # entity_id -> depth
|
||
for eid in entity_ids:
|
||
entity_depths[eid] = 0
|
||
|
||
for layer in range(depth):
|
||
if not current_layer_ids:
|
||
break
|
||
|
||
# 查询当前层实体的所有关系
|
||
placeholders = ','.join('?' * len(current_layer_ids))
|
||
|
||
query = f"""
|
||
SELECT r.id, r.source_id, r.target_id,
|
||
e1.name as source, e2.name as target,
|
||
r.relation_type as type, r.confidence, r.session_id,
|
||
r.turn_id, r.created_at, r.status
|
||
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 ({placeholders}) OR r.target_id IN ({placeholders}))
|
||
AND r.status = 'active'
|
||
"""
|
||
|
||
params = list(current_layer_ids) + list(current_layer_ids)
|
||
|
||
if session_filter:
|
||
query += " AND r.session_id = ?"
|
||
params.append(session_filter)
|
||
|
||
cursor.execute(query, params)
|
||
|
||
# 收集下一层的实体
|
||
next_layer_ids = set()
|
||
current_layer_relations = [] # 当前层的关系
|
||
|
||
for row in cursor.fetchall():
|
||
# 计算关系的深度(取两端实体深度的最大值+1)
|
||
source_depth = entity_depths.get(row['source_id'], layer)
|
||
target_depth = entity_depths.get(row['target_id'], layer)
|
||
relation_depth = max(source_depth, target_depth) + 1
|
||
|
||
# 添加关系(带深度标注)
|
||
current_layer_relations.append({
|
||
'source': row['source'],
|
||
'target': row['target'],
|
||
'type': row['type'],
|
||
'confidence': row['confidence'],
|
||
'session_id': row['session_id'],
|
||
'turn_id': row['turn_id'],
|
||
'created_at': row['created_at'],
|
||
'status': row['status'],
|
||
'depth': relation_depth
|
||
})
|
||
|
||
# 收集新实体(未访问过的)
|
||
source_id = row['source_id']
|
||
target_id = row['target_id']
|
||
|
||
if source_id not in visited_entity_ids:
|
||
next_layer_ids.add(source_id)
|
||
visited_entity_ids.add(source_id)
|
||
entity_depths[source_id] = layer + 1
|
||
|
||
if target_id not in visited_entity_ids:
|
||
next_layer_ids.add(target_id)
|
||
visited_entity_ids.add(target_id)
|
||
entity_depths[target_id] = layer + 1
|
||
|
||
relations.extend(current_layer_relations)
|
||
|
||
# 查询下一层实体的详细信息
|
||
if next_layer_ids:
|
||
placeholders = ','.join('?' * len(next_layer_ids))
|
||
cursor.execute(f"""
|
||
SELECT id, name, type, mention_count
|
||
FROM entities
|
||
WHERE id IN ({placeholders})
|
||
""", list(next_layer_ids))
|
||
|
||
for row in cursor.fetchall():
|
||
entities.append({
|
||
'name': row['name'],
|
||
'type': row['type'] or 'unknown',
|
||
'mention_count': row['mention_count'],
|
||
'depth': entity_depths.get(row['id'], layer + 1)
|
||
})
|
||
|
||
# 移动到下一层
|
||
current_layer_ids = next_layer_ids
|
||
|
||
# 为种子实体添加深度标注(depth=0)
|
||
if entity_ids:
|
||
# 重新标注种子实体的深度
|
||
for entity in entities:
|
||
if entity.get('depth') is None:
|
||
entity['depth'] = 0
|
||
|
||
return {
|
||
"entities": entities,
|
||
"relations": relations,
|
||
"message": f"找到 {len(entities)} 个实体, {len(relations)} 条关系"
|
||
}
|
||
|
||
def commit(self, triplets: List[Dict], entity_types: Dict = None,
|
||
temporal_tag: str = None, session_id: str = None,
|
||
turn_id: int = None) -> Dict:
|
||
"""
|
||
写入记忆
|
||
|
||
Args:
|
||
triplets: 三元组列表
|
||
entity_types: 实体类型
|
||
temporal_tag: 时间标签
|
||
session_id: 会话ID
|
||
turn_id: 轮次ID
|
||
|
||
Returns:
|
||
写入结果
|
||
"""
|
||
cursor = self.conn.cursor()
|
||
|
||
created_entities = 0
|
||
created_relations = 0
|
||
|
||
for triplet in triplets:
|
||
subject = triplet.get('subject')
|
||
relation = triplet.get('relation')
|
||
obj = triplet.get('object')
|
||
confidence = triplet.get('confidence', 1.0)
|
||
|
||
if not all([subject, relation, obj]):
|
||
continue
|
||
|
||
# 创建或更新实体
|
||
for entity_name in [subject, obj]:
|
||
entity_type = entity_types.get(entity_name) if entity_types else None
|
||
|
||
cursor.execute("""
|
||
INSERT INTO entities (name, type)
|
||
VALUES (?, ?)
|
||
ON CONFLICT(name) DO UPDATE SET
|
||
mention_count = mention_count + 1,
|
||
updated_at = CURRENT_TIMESTAMP
|
||
""", (entity_name, entity_type))
|
||
|
||
if cursor.rowcount > 0:
|
||
created_entities += 1
|
||
|
||
# 获取实体ID
|
||
cursor.execute("SELECT id FROM entities WHERE name = ?", (subject,))
|
||
source_id = cursor.fetchone()['id']
|
||
|
||
cursor.execute("SELECT id FROM entities WHERE name = ?", (obj,))
|
||
target_id = cursor.fetchone()['id']
|
||
|
||
# 创建关系
|
||
date_bucket = datetime.now().strftime('%Y-%m-%d')
|
||
|
||
cursor.execute("""
|
||
INSERT INTO relations (
|
||
source_id, target_id, relation_type, confidence,
|
||
session_id, turn_id, date_bucket
|
||
)
|
||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||
""", (source_id, target_id, relation, confidence,
|
||
session_id, turn_id, date_bucket))
|
||
|
||
created_relations += 1
|
||
|
||
self.conn.commit()
|
||
|
||
return {
|
||
"created_entities": created_entities,
|
||
"created_relations": created_relations,
|
||
"message": f"创建了 {created_entities} 个实体, {created_relations} 条关系"
|
||
}
|
||
|
||
def purge(self, criteria: Dict, mode: str = "soft",
|
||
new_relation: Dict = None) -> Dict:
|
||
"""
|
||
删除或修正记忆
|
||
|
||
Args:
|
||
criteria: 删除条件
|
||
mode: 删除模式 (soft/hard)
|
||
new_relation: 替代关系
|
||
|
||
Returns:
|
||
删除结果
|
||
"""
|
||
cursor = self.conn.cursor()
|
||
|
||
# 构建查询条件
|
||
conditions = []
|
||
params = []
|
||
|
||
if criteria.get('source'):
|
||
cursor.execute("SELECT id FROM entities WHERE name = ?", (criteria['source'],))
|
||
row = cursor.fetchone()
|
||
if row:
|
||
conditions.append("source_id = ?")
|
||
params.append(row['id'])
|
||
|
||
if criteria.get('target'):
|
||
cursor.execute("SELECT id FROM entities WHERE name = ?", (criteria['target'],))
|
||
row = cursor.fetchone()
|
||
if row:
|
||
conditions.append("target_id = ?")
|
||
params.append(row['id'])
|
||
|
||
if criteria.get('relation'):
|
||
conditions.append("relation_type = ?")
|
||
params.append(criteria['relation'])
|
||
|
||
if not conditions:
|
||
return {"deleted": 0, "message": "无删除条件"}
|
||
|
||
where_clause = " AND ".join(conditions)
|
||
|
||
if mode == "soft":
|
||
cursor.execute(f"""
|
||
UPDATE relations
|
||
SET status = 'deleted', updated_at = CURRENT_TIMESTAMP
|
||
WHERE {where_clause} AND status = 'active'
|
||
""", params)
|
||
else:
|
||
cursor.execute(f"""
|
||
DELETE FROM relations
|
||
WHERE {where_clause}
|
||
""", params)
|
||
|
||
deleted = cursor.rowcount
|
||
self.conn.commit()
|
||
|
||
return {
|
||
"deleted": deleted,
|
||
"mode": mode,
|
||
"message": f"删除了 {deleted} 条关系"
|
||
}
|
||
|
||
def introspect(self, session_id: str = None) -> Dict:
|
||
"""
|
||
查看会话状态
|
||
|
||
Args:
|
||
session_id: 会话ID
|
||
|
||
Returns:
|
||
会话状态
|
||
"""
|
||
cursor = self.conn.cursor()
|
||
|
||
# 统计实体
|
||
cursor.execute("SELECT COUNT(*) as count FROM entities")
|
||
entity_count = cursor.fetchone()['count']
|
||
|
||
# 统计关系
|
||
cursor.execute("SELECT COUNT(*) as count FROM relations WHERE status = 'active'")
|
||
relation_count = cursor.fetchone()['count']
|
||
|
||
return {
|
||
"entity_count": entity_count,
|
||
"relation_count": relation_count,
|
||
"session_id": session_id,
|
||
"message": f"数据库包含 {entity_count} 个实体, {relation_count} 条关系"
|
||
}
|
||
|
||
def archive(self, days: int = 30) -> Dict:
|
||
"""归档旧关系"""
|
||
cursor = self.conn.cursor()
|
||
|
||
cursor.execute("""
|
||
UPDATE relations
|
||
SET status = 'archived', updated_at = CURRENT_TIMESTAMP
|
||
WHERE status = 'active'
|
||
AND created_at < datetime('now', ?)
|
||
""", (f'-{days} days',))
|
||
|
||
archived = cursor.rowcount
|
||
self.conn.commit()
|
||
|
||
return {
|
||
"archived": archived,
|
||
"message": f"归档了 {archived} 条关系"
|
||
}
|
||
|
||
def cleanup(self, dry_run: bool = True) -> Dict:
|
||
"""清理已删除数据"""
|
||
cursor = self.conn.cursor()
|
||
|
||
if dry_run:
|
||
cursor.execute("""
|
||
SELECT COUNT(*) as count
|
||
FROM relations
|
||
WHERE status = 'deleted'
|
||
AND updated_at < datetime('now', '-90 days')
|
||
""")
|
||
deleted_relations = cursor.fetchone()['count']
|
||
|
||
return {
|
||
"dry_run": True,
|
||
"deleted_relations": deleted_relations,
|
||
"message": f"将删除 {deleted_relations} 条关系"
|
||
}
|
||
else:
|
||
cursor.execute("""
|
||
DELETE FROM relations
|
||
WHERE status = 'deleted'
|
||
AND updated_at < datetime('now', '-90 days')
|
||
""")
|
||
deleted = cursor.rowcount
|
||
self.conn.commit()
|
||
|
||
return {
|
||
"dry_run": False,
|
||
"deleted": deleted,
|
||
"message": f"删除了 {deleted} 条关系"
|
||
}
|
||
|
||
def save_chat_records(self, messages: list) -> Dict:
|
||
"""保存聊天记录到数据库"""
|
||
cursor = self.conn.cursor()
|
||
|
||
saved = 0
|
||
for msg in messages:
|
||
role = msg.get("role")
|
||
content = msg.get("content")
|
||
if role and content:
|
||
cursor.execute(
|
||
"INSERT INTO chat_records (role, content) VALUES (?, ?)",
|
||
(role, content)
|
||
)
|
||
saved += 1
|
||
|
||
self.conn.commit()
|
||
|
||
cursor.execute("""
|
||
DELETE FROM chat_records
|
||
WHERE id NOT IN (
|
||
SELECT id FROM chat_records
|
||
ORDER BY id DESC
|
||
LIMIT 500
|
||
)
|
||
""")
|
||
self.conn.commit()
|
||
|
||
return {"saved": saved}
|
||
|
||
def get_chat_records(self, limit: int = 500) -> list:
|
||
"""从数据库获取聊天记录"""
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("""
|
||
SELECT role, content FROM chat_records
|
||
ORDER BY id ASC LIMIT ?
|
||
""", (limit,))
|
||
return [{"role": row[0], "content": row[1]} for row in cursor.fetchall()]
|
||
|
||
def clear_chat_records(self) -> Dict:
|
||
"""清空聊天记录(保留图数据库)"""
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("DELETE FROM chat_records")
|
||
self.conn.commit()
|
||
return {"cleared": True}
|
||
|
||
def set_web_user(self, username: str, password: str, base_dir: str = None, role: str = 'user') -> Dict:
|
||
"""设置或更新 Web 登录用户。password 是明文,自动哈希存储。
|
||
自动创建用户目录并设置 config_path 和 db_path。
|
||
role: 'admin' 或 'user',默认 'user'"""
|
||
if not username or not password:
|
||
return {"success": False, "error": "用户名和密码不能为空"}
|
||
if role not in ('admin', 'user'):
|
||
return {"success": False, "error": "角色无效 (admin/user)"}
|
||
|
||
import hashlib
|
||
from pathlib import Path
|
||
|
||
password_hash = hashlib.sha256(password.encode()).hexdigest()
|
||
|
||
# 确定基础目录
|
||
if base_dir is None:
|
||
base_dir = Path.home() / ".trulymem"
|
||
else:
|
||
base_dir = Path(base_dir)
|
||
|
||
# 创建用户目录
|
||
user_dir = base_dir / username
|
||
user_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
# 设置用户文件路径
|
||
config_path = str(user_dir / "config.json")
|
||
db_path = str(user_dir / f"{username}_graph.db")
|
||
|
||
cursor = self.conn.cursor()
|
||
# 如果是第一个用户,强制设为 admin
|
||
if self.get_web_users_count() == 0:
|
||
role = 'admin'
|
||
cursor.execute("""
|
||
INSERT INTO web_users (username, password_hash, role, config_path, db_path)
|
||
VALUES (?, ?, ?, ?, ?)
|
||
ON CONFLICT(username) DO UPDATE SET
|
||
password_hash = excluded.password_hash,
|
||
role = CASE WHEN web_users.role = 'admin' THEN 'admin' ELSE excluded.role END,
|
||
config_path = COALESCE(web_users.config_path, excluded.config_path),
|
||
db_path = COALESCE(web_users.db_path, excluded.db_path),
|
||
updated_at = CURRENT_TIMESTAMP
|
||
""", (username, password_hash, role, config_path, db_path))
|
||
self.conn.commit()
|
||
return {"success": True, "username": username, "role": role, "config_path": config_path, "db_path": db_path}
|
||
|
||
def get_web_users(self) -> List[Dict]:
|
||
"""获取所有 Web 用户列表"""
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("SELECT id, username, role, config_path, db_path, created_at, updated_at FROM web_users ORDER BY created_at ASC")
|
||
users = []
|
||
for row in cursor.fetchall():
|
||
users.append({
|
||
"id": row['id'],
|
||
"username": row['username'],
|
||
"role": row['role'],
|
||
"config_path": row['config_path'],
|
||
"db_path": row['db_path'],
|
||
"created_at": row['created_at'],
|
||
"updated_at": row['updated_at']
|
||
})
|
||
return users
|
||
|
||
def get_web_user(self, username: str) -> Optional[Dict]:
|
||
"""获取单个 Web 用户信息"""
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("""
|
||
SELECT id, username, role, config_path, db_path, created_at, updated_at
|
||
FROM web_users WHERE username = ?
|
||
""", (username,))
|
||
row = cursor.fetchone()
|
||
if row:
|
||
return {
|
||
"id": row['id'],
|
||
"username": row['username'],
|
||
"role": row['role'],
|
||
"config_path": row['config_path'],
|
||
"db_path": row['db_path'],
|
||
"created_at": row['created_at'],
|
||
"updated_at": row['updated_at']
|
||
}
|
||
return None
|
||
|
||
def is_admin(self, username: str) -> bool:
|
||
"""检查用户是否为管理员"""
|
||
user = self.get_web_user(username)
|
||
return user is not None and user.get('role') == 'admin'
|
||
|
||
def delete_web_user(self, username: str) -> Dict:
|
||
"""删除 Web 用户(同时保留文件目录)"""
|
||
if not username:
|
||
return {"success": False, "error": "用户名不能为空"}
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("DELETE FROM web_users WHERE username = ?", (username,))
|
||
self.conn.commit()
|
||
if cursor.rowcount > 0:
|
||
return {"success": True, "username": username}
|
||
return {"success": False, "error": "用户不存在"}
|
||
|
||
def get_web_users_count(self) -> int:
|
||
"""获取 Web 用户数量 (用于判断是否需要首次设置)"""
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("SELECT COUNT(*) as cnt FROM web_users")
|
||
row = cursor.fetchone()
|
||
return row['cnt'] if row else 0
|
||
|
||
def verify_web_user(self, username: str, password: str) -> bool:
|
||
"""验证 Web 用户登录"""
|
||
import hashlib
|
||
password_hash = hashlib.sha256(password.encode()).hexdigest()
|
||
cursor = self.conn.cursor()
|
||
cursor.execute("""
|
||
SELECT id FROM web_users
|
||
WHERE username = ? AND password_hash = ?
|
||
""", (username, password_hash))
|
||
return cursor.fetchone() is not None
|
||
|
||
def close(self):
|
||
"""关闭数据库连接"""
|
||
if self.conn:
|
||
self.conn.close()
|
||
self.conn = None
|
||
|
||
def __enter__(self):
|
||
return self
|
||
|
||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||
self.close()
|
||
|
||
|
||
# 兼容性别名
|
||
Neo4jGraph = EmbeddedGraphDB
|
||
|
||
|
||
if __name__ == '__main__':
|
||
# 测试
|
||
print("Testing Embedded Graph Database...")
|
||
|
||
with EmbeddedGraphDB("test.db") as db:
|
||
# 写入测试
|
||
result = db.commit(
|
||
triplets=[
|
||
{"subject": "用户", "relation": "喜欢", "object": "Python"},
|
||
{"subject": "用户", "relation": "学习", "object": "AI"}
|
||
],
|
||
session_id="test-session",
|
||
turn_id=1
|
||
)
|
||
print(f"Commit: {result}")
|
||
|
||
# 检索测试
|
||
result = db.recall("Python,AI")
|
||
print(f"Recall: {result}")
|
||
|
||
# 状态测试
|
||
result = db.introspect()
|
||
print(f"Introspect: {result}")
|
||
|
||
print("\nTest completed!")
|