From 83906a6985d0cf3b9dc716a29d55655635f267da Mon Sep 17 00:00:00 2001 From: root Date: Sun, 12 Apr 2026 18:15:55 +0800 Subject: [PATCH] refactor: restructure to core/ + ui/ with multi-threaded backend --- .gitignore | 8 +- README.md | 85 ++-- build/build_linux.sh | 26 +- build/build_macos.sh | 22 + build/build_windows.bat | 27 +- build/trulymem.spec | 17 +- core/__init__.py | 5 + core/client.py | 27 ++ core/embedded_db.py | 438 ++++++++++++++++++ core/graph_client.py | 383 ++++++++++++++++ core/prompts/__init__.py | 6 + core/prompts/prompt_manager.py | 64 +++ core/prompts/templates/system_prompt.md | 324 ++++++++++++++ core/server.py | 341 ++++++++++++++ core/tool_executor.py | 307 +++++++++++++ core/tool_limiter.py | 164 +++++++ core/tools/__init__.py | 8 + core/tools/memory_tools.py | 520 ++++++++++++++++++++++ core/tools/tool_executor.py | 307 +++++++++++++ core/tools/tool_limiter.py | 164 +++++++ docs/README.md | 2 +- docs/架构.md | 210 +++++---- tests/conftest.py | 15 +- tests/test_core/test_embedded_db.py | 2 +- tests/test_core/test_imports.py | 60 +-- tests/test_core/test_memory_tools.py | 138 ------ tests/test_core/test_server.py | 57 +++ tests/test_core/test_tool_limiter.py | 148 ------ tests/test_handlers/__init__.py | 1 - tests/test_handlers/test_focus_handler.py | 33 -- tests/test_models/test_models.py | 6 +- tests/test_services/__init__.py | 1 - tests/test_ui.py | 70 +++ tests/test_widgets/__init__.py | 1 - trulymem_entry.py | 38 +- ui/__init__.py | 4 + ui/app.py | 279 ++++++++++++ ui/handlers/__init__.py | 1 + ui/handlers/focus_handler.py | 64 +++ ui/handlers/key_handler.py | 68 +++ ui/handlers/message_handler.py | 98 ++++ ui/models/__init__.py | 1 + ui/models/config.py | 46 ++ ui/models/log_entry.py | 30 ++ ui/models/message.py | 33 ++ ui/services/__init__.py | 1 + ui/services/chat_service.py | 226 ++++++++++ ui/services/config_manager.py | 46 ++ ui/services/config_service.py | 51 +++ ui/services/tool_service.py | 88 ++++ ui/styles/__init__.py | 1 + ui/styles/app.css | 41 ++ ui/styles/components.css | 129 ++++++ ui/styles/messages.css | 24 + ui/widgets/__init__.py | 1 + ui/widgets/config_section.py | 128 ++++++ ui/widgets/cypher_query_box.py | 57 +++ ui/widgets/input_box.py | 50 +++ ui/widgets/left_panel.py | 23 + ui/widgets/message_history.py | 57 +++ ui/widgets/message_widget.py | 108 +++++ ui/widgets/operation_log.py | 66 +++ ui/widgets/right_panel.py | 67 +++ ui/widgets/status_bar.py | 27 ++ 64 files changed, 5269 insertions(+), 571 deletions(-) create mode 100644 build/build_macos.sh create mode 100644 core/__init__.py create mode 100644 core/client.py create mode 100644 core/embedded_db.py create mode 100644 core/graph_client.py create mode 100644 core/prompts/__init__.py create mode 100644 core/prompts/prompt_manager.py create mode 100644 core/prompts/templates/system_prompt.md create mode 100644 core/server.py create mode 100644 core/tool_executor.py create mode 100644 core/tool_limiter.py create mode 100644 core/tools/__init__.py create mode 100644 core/tools/memory_tools.py create mode 100644 core/tools/tool_executor.py create mode 100644 core/tools/tool_limiter.py delete mode 100644 tests/test_core/test_memory_tools.py create mode 100644 tests/test_core/test_server.py delete mode 100644 tests/test_core/test_tool_limiter.py delete mode 100644 tests/test_handlers/__init__.py delete mode 100644 tests/test_handlers/test_focus_handler.py delete mode 100644 tests/test_services/__init__.py create mode 100644 tests/test_ui.py delete mode 100644 tests/test_widgets/__init__.py create mode 100644 ui/__init__.py create mode 100644 ui/app.py create mode 100644 ui/handlers/__init__.py create mode 100644 ui/handlers/focus_handler.py create mode 100644 ui/handlers/key_handler.py create mode 100644 ui/handlers/message_handler.py create mode 100644 ui/models/__init__.py create mode 100644 ui/models/config.py create mode 100644 ui/models/log_entry.py create mode 100644 ui/models/message.py create mode 100644 ui/services/__init__.py create mode 100644 ui/services/chat_service.py create mode 100644 ui/services/config_manager.py create mode 100644 ui/services/config_service.py create mode 100644 ui/services/tool_service.py create mode 100644 ui/styles/__init__.py create mode 100644 ui/styles/app.css create mode 100644 ui/styles/components.css create mode 100644 ui/styles/messages.css create mode 100644 ui/widgets/__init__.py create mode 100644 ui/widgets/config_section.py create mode 100644 ui/widgets/cypher_query_box.py create mode 100644 ui/widgets/input_box.py create mode 100644 ui/widgets/left_panel.py create mode 100644 ui/widgets/message_history.py create mode 100644 ui/widgets/message_widget.py create mode 100644 ui/widgets/operation_log.py create mode 100644 ui/widgets/right_panel.py create mode 100644 ui/widgets/status_bar.py diff --git a/.gitignore b/.gitignore index a30464e..efb5d19 100644 --- a/.gitignore +++ b/.gitignore @@ -4,13 +4,11 @@ __pycache__/ *$py.class *.so .Python -build/ develop-eggs/ dist/ downloads/ eggs/ -.eggs/ -lib/ +.lib/ lib64/ parts/ sdist/ @@ -55,7 +53,6 @@ Thumbs.db config.json # Build -build/ dist/ # Temporary @@ -64,3 +61,6 @@ dist/ # AI Generated jimeng*.png + +# Test Cache +.pytest_cache/ diff --git a/README.md b/README.md index 2b301f4..daf965e 100644 --- a/README.md +++ b/README.md @@ -33,8 +33,6 @@ TrulyMEM (TrueHumanMEM) 是一个让 AI 拥有长期记忆能力的图记忆系 ### 方式一:打包后的可执行文件 -打包后会生成独立可执行文件,可直接运行: - ```bash # Windows: TrulyMEM.exe # Linux/macOS: TrulyMEM @@ -45,14 +43,11 @@ chmod +x TrulyMEM ### 方式二:从源码运行 ```bash -# 克隆仓库 git clone cd TrulyMEM-TrueHumanMEM -# 安装依赖 pip install -r requirements.txt -# 运行应用 python trulymem_entry.py ``` @@ -111,25 +106,27 @@ python trulymem_entry.py ``` TrulyMEM-TrueHumanMEM/ -├── trulymem_entry.py # 打包入口 -├── graph_memory_tui/ # 核心应用 -│ ├── app.py # TUI 主应用 -│ ├── main.py # 模块入口 -│ ├── core/ # 核心逻辑 -│ │ ├── embedded_db.py # SQLite 图数据库 -│ │ ├── graph_client.py # Neo4j 客户端(可选) -│ │ ├── imports.py # 动态导入 -│ │ ├── prompts/ # 提示词管理 -│ │ └── tools/ # 工具定义 -│ ├── models/ # 数据模型 -│ ├── services/ # 服务层 -│ ├── handlers/ # 事件处理 -│ ├── widgets/ # TUI 组件 -│ └── styles/ # 样式文件 -├── tests/ # 测试 -├── docs/ # 文档 -├── requirements.txt # 依赖清单 -└── LICENSE # 许可证 +├── trulymem_entry.py # 入口:先启动 core → 再启动 ui +├── core/ # 后端/业务逻辑 +│ ├── __init__.py +│ ├── server.py # BackendServer (多线程) +│ ├── client.py # BackendClient +│ ├── embedded_db.py # SQLite 图数据库 +│ ├── graph_client.py +│ ├── tool_executor.py +│ ├── tool_limiter.py +│ ├── memory_tools.py +│ ├── prompts/ +│ └── tools/ # TOOLS 定义 +├── ui/ # TUI 显示层 +│ ├── __init__.py +│ ├── app.py +│ ├── widgets/ +│ ├── handlers/ +│ ├── models/ +│ ├── services/ +│ └── styles/ +└── tests/ # 测试 (38 tests) ``` --- @@ -143,36 +140,38 @@ TrulyMEM-TrueHumanMEM/ --- +## 架构说明 + +### TUI 与后端通信 + +``` +trulymem_entry.py + │ + ├─ 1. BackendServer.start() → 启动独立线程 + │ + ├─ 2. GraphMemoryApp(backend_server=server) + │ + └─ 3. BackendClient ← Queue → BackendServer +``` + +- **core/** - 业务逻辑(数据库、API调用、工具执行) +- **ui/** - 显示逻辑(Textual 组件) +- 多线程 Queue 通信解耦 + +--- + ## 开发指南 -### 环境设置 - ```bash -# 创建虚拟环境 python -m venv venv source venv/bin/activate # Linux/macOS venv\Scripts\activate # Windows -# 安装依赖 pip install -r requirements.txt -``` -### 运行测试 - -```bash pytest tests/ ``` -### 打包应用 - -```bash -# Windows -bash build/build_windows.bat - -# Linux -bash build/build_linux.sh -``` - --- ## 许可证 @@ -207,4 +206,4 @@ bash build/build_linux.sh - [架构设计](docs/架构.md) - 系统架构和技术设计 - [快速开始](docs/一键启动指南.md) - 启动指南 -- [工作记忆链机制说明](docs/工作记忆链机制说明.md) - 连续性任务处理 +- [工作记忆链机制说明](docs/工作记忆链机制说明.md) - 连续性任务处理 \ No newline at end of file diff --git a/build/build_linux.sh b/build/build_linux.sh index 39fc55c..a596dd0 100644 --- a/build/build_linux.sh +++ b/build/build_linux.sh @@ -1,32 +1,22 @@ #!/bin/bash -# Linux打包脚本 -# 需要在Linux系统上运行 +set -e -echo "开始打包Linux版本..." +echo "Building Linux binary..." -# 检查Python和PyInstaller if ! command -v python3 &> /dev/null; then - echo "错误: 未找到python3" + echo "Error: python3 not found" exit 1 fi -# 安装依赖 -echo "安装依赖..." pip3 install -r requirements.txt -# 打包Linux二进制文件 -echo "打包Linux二进制文件..." -python3 -m PyInstaller \ - --clean \ - --onefile \ - --name TrulyMEM \ - --console \ - --add-data "graph_memory_tui/styles/*.css:graph_memory_tui/styles" \ +python3 -m PyInstaller --clean --onefile --console \ + --add-data "ui/styles:ui/styles" \ + --add-data "core/prompts/templates:core/prompts/templates" \ --hidden-import textual \ --hidden-import openai \ - --hidden-import flask \ --hidden-import neo4j \ + --collect-all textual \ trulymem_entry.py -echo "打包完成!" -echo "可执行文件位于: dist/TrulyMEM" +echo "Done! Binary: dist/TrulyMEM" \ No newline at end of file diff --git a/build/build_macos.sh b/build/build_macos.sh new file mode 100644 index 0000000..475e1e6 --- /dev/null +++ b/build/build_macos.sh @@ -0,0 +1,22 @@ +#!/bin/bash +set -e + +echo "Building macOS binary..." + +if ! command -v python3 &> /dev/null; then + echo "Error: python3 not found" + exit 1 +fi + +pip3 install -r requirements.txt + +python3 -m PyInstaller --clean --onefile --console \ + --add-data "ui/styles:ui/styles" \ + --add-data "core/prompts/templates:core/prompts/templates" \ + --hidden-import textual \ + --hidden-import openai \ + --hidden-import neo4j \ + --collect-all textual \ + trulymem_entry.py + +echo "Done! Binary: dist/TrulyMEM" \ No newline at end of file diff --git a/build/build_windows.bat b/build/build_windows.bat index 2f95063..13657dc 100644 --- a/build/build_windows.bat +++ b/build/build_windows.bat @@ -1,33 +1,22 @@ @echo off -REM Windows打包脚本 +echo Building Windows binary... -echo 开始打包Windows版本... - -REM 检查Python python --version >nul 2>&1 if errorlevel 1 ( - echo 错误: 未找到python + echo Error: python not found exit /b 1 ) -REM 安装依赖 -echo 安装依赖... pip install -r requirements.txt -REM 打包Windows exe文件 -echo 打包Windows exe文件... -python -m PyInstaller ^ - --clean ^ - --onefile ^ - --name TrulyMEM ^ - --console ^ - --add-data "graph_memory_tui/styles/*.css;graph_memory_tui/styles" ^ +python -m PyInstaller --clean --onefile --console ^ + --add-data "ui/styles;ui/styles" ^ + --add-data "core/prompts/templates;core/prompts/templates" ^ --hidden-import textual ^ --hidden-import openai ^ - --hidden-import flask ^ --hidden-import neo4j ^ + --collect-all textual ^ trulymem_entry.py -echo 打包完成! -echo 可执行文件位于: dist\TrulyMEM.exe -pause +echo Done! Binary: dist\TrulyMEM.exe +pause \ No newline at end of file diff --git a/build/trulymem.spec b/build/trulymem.spec index 47404d1..c1a38cd 100644 --- a/build/trulymem.spec +++ b/build/trulymem.spec @@ -1,12 +1,18 @@ # -*- mode: python ; coding: utf-8 -*- - - a = Analysis( ['trulymem_entry.py'], pathex=[], binaries=[], - datas=[('graph_memory_tui/styles/*.css', 'graph_memory_tui/styles')], - hiddenimports=['textual', 'openai', 'flask', 'neo4j'], + datas=[ + ('ui/styles', 'ui/styles'), + ('core/prompts/templates', 'core/prompts/templates'), + ], + hiddenimports=[ + 'textual', + 'openai', + 'neo4j', + 'sqlite3', + ], hookspath=[], hooksconfig={}, runtime_hooks=[], @@ -15,7 +21,6 @@ a = Analysis( optimize=0, ) pyz = PYZ(a.pure) - exe = EXE( pyz, a.scripts, @@ -35,4 +40,4 @@ exe = EXE( target_arch=None, codesign_identity=None, entitlements_file=None, -) +) \ No newline at end of file diff --git a/core/__init__.py b/core/__init__.py new file mode 100644 index 0000000..0ae06f2 --- /dev/null +++ b/core/__init__.py @@ -0,0 +1,5 @@ +from .server import BackendServer +from .client import BackendClient +from .embedded_db import EmbeddedGraphDB + +__all__ = ["BackendServer", "BackendClient", "EmbeddedGraphDB"] \ No newline at end of file diff --git a/core/client.py b/core/client.py new file mode 100644 index 0000000..42ed9db --- /dev/null +++ b/core/client.py @@ -0,0 +1,27 @@ +import threading +from typing import Any, Dict + +from .server import BackendServer + + +class BackendClient: + def __init__(self, server: BackendServer): + self._server = server + self._request_counter = 0 + self._lock = threading.Lock() + + def process_message(self, user_input: str, timeout: float = 30.0) -> Dict[str, Any]: + with self._lock: + self._request_counter += 1 + return self._server.process_message(user_input, timeout) + + def execute_tool(self, tool_name: str, arguments: Dict[str, Any], timeout: float = 10.0) -> str: + with self._lock: + self._request_counter += 1 + return self._server.execute_tool(tool_name, arguments, timeout) + + def update_config(self, api_key: str, base_url: str = "https://api.deepseek.com") -> None: + self._server.update_config(api_key, base_url) + + def shutdown(self) -> None: + self._server.shutdown() \ No newline at end of file diff --git a/core/embedded_db.py b/core/embedded_db.py new file mode 100644 index 0000000..ec126ef --- /dev/null +++ b/core/embedded_db.py @@ -0,0 +1,438 @@ +""" +内嵌图数据库 - 基于SQLite实现 +无需Docker,开箱即用 +""" + +import sqlite3 +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)") + + 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'] + }) + + # 搜索关系 + relations = [] + + if entity_ids: + placeholders = ','.join('?' * len(entity_ids)) + + query = f""" + SELECT r.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(entity_ids) + list(entity_ids) + + if session_filter: + query += " AND r.session_id = ?" + params.append(session_filter) + + cursor.execute(query, params) + + for row in cursor.fetchall(): + 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'] + }) + + 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 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!") diff --git a/core/graph_client.py b/core/graph_client.py new file mode 100644 index 0000000..76c5c1b --- /dev/null +++ b/core/graph_client.py @@ -0,0 +1,383 @@ +#!/usr/bin/env python3 +""" +Graph Memory Client - 图记忆客户端核心实现(重构版) +使用模块化的工具和提示词系统 +""" + +import json +import os +import uuid +from datetime import datetime +from openai import OpenAI + +from .tools import TOOLS +from .tool_executor import execute_tool +from .prompts.prompt_manager import PromptManager + +# 环境配置 +DEEPSEEK_API_KEY = os.environ.get("DEEPSEEK_API_KEY", "") +DEEPSEEK_BASE_URL = os.environ.get("DEEPSEEK_BASE_URL", "https://api.deepseek.com") +MODEL_NAME = os.environ.get("MODEL_NAME", "deepseek-chat") + +NEO4J_URI = os.environ.get("NEO4J_URI", "bolt://localhost:7687") +NEO4J_USER = os.environ.get("NEO4J_USER", "neo4j") +NEO4J_PASSWORD = os.environ.get("NEO4J_PASSWORD", "neo4j") + +# 会话配置 +CURRENT_SESSION_ID = f"session-{datetime.now().strftime('%Y%m%d')}-{uuid.uuid4().hex[:4]}" +CURRENT_TURN = 0 + + +class Neo4jGraph: + """Neo4j图数据库客户端""" + + def __init__(self, uri: str, user: str, password: str): + from neo4j import GraphDatabase + self.driver = GraphDatabase.driver(uri, auth=(user, password)) + + def close(self): + self.driver.close() + + def ensure_constraints(self): + """确保约束和索引存在""" + with self.driver.session() as session: + # 实体约束 + session.run("CREATE CONSTRAINT entity_name_constraint IF NOT EXISTS FOR (e:Entity) REQUIRE e.name IS UNIQUE") + session.run("CREATE CONSTRAINT session_id_constraint IF NOT EXISTS FOR (s:Session) REQUIRE s.session_id IS UNIQUE") + + # 关系索引 + session.run("CREATE INDEX rel_created_at IF NOT EXISTS FOR ()-[r:RELATES]-() ON r.created_at") + session.run("CREATE INDEX rel_session_id IF NOT EXISTS FOR ()-[r:RELATES]-() ON r.session_id") + session.run("CREATE INDEX rel_type IF NOT EXISTS FOR ()-[r:RELATES]-() ON r.type") + session.run("CREATE INDEX rel_status IF NOT EXISTS FOR ()-[r:RELATES]-() ON r.status") + session.run("CREATE INDEX rel_date_bucket IF NOT EXISTS FOR ()-[r:RELATES]-() ON r.date_bucket") + + # 实体索引 + session.run("CREATE INDEX entity_type IF NOT EXISTS FOR (e:Entity) ON e.type") + session.run("CREATE INDEX entity_mention_count IF NOT EXISTS FOR (e:Entity) ON e.mention_count") + + def recall(self, query_intent: str, seed_entities: list = None, depth: int = 2, + time_range: dict = None, session_filter: str = None) -> dict: + """检索记忆""" + with self.driver.session() as session: + # 支持逗号分隔的多个关键词 + keywords = [w.strip() for w in query_intent.replace(',', ' ').split() if len(w.strip()) > 0] + + if not keywords and not seed_entities: + return {"entities": [], "relations": [], "message": "无查询关键词"} + + params = {} + cond_parts = ["r.status = 'active'"] + + if session_filter: + cond_parts.append("r.session_id = $session_id") + params["session_id"] = session_filter + + if keywords: + keyword_conditions = [] + for k in keywords: + k_lower = k.lower() + keyword_conditions.append(f"toLower(e.name) CONTAINS '{k_lower}'") + keyword_conditions.append(f"toLower(t.name) CONTAINS '{k_lower}'") + keyword_conditions.append(f"toLower(r.type) CONTAINS '{k_lower}'") + cond_parts.append(f"({' OR '.join(keyword_conditions)})") + + if seed_entities: + placeholders = ",".join([f"'{s}'" for s in seed_entities]) + cond_parts.append(f"(e.name IN [{placeholders}] OR t.name IN [{placeholders}])") + + if time_range and "days" in time_range: + cond_parts.append(f"r.created_at >= datetime() - duration('P{time_range['days']}D')") + + where_clause = " AND ".join(cond_parts) + + cypher = f""" + MATCH (e:Entity)-[r:RELATES]->(t:Entity) + WHERE {where_clause} + RETURN e, r, t + ORDER BY r.created_at DESC + LIMIT 30 + """ + + result = session.run(cypher, params) + entities, relations = {}, [] + + for record in result: + e, r, t = record["e"], record["r"], record["t"] + if e["name"] not in entities: + entities[e["name"]] = {"name": e["name"], "type": e.get("type", "unknown"), "mention_count": e.get("mention_count", 1)} + if t["name"] not in entities: + entities[t["name"]] = {"name": t["name"], "type": t.get("type", "unknown"), "mention_count": t.get("mention_count", 1)} + + relations.append({ + "source": e["name"], + "target": t["name"], + "type": r["type"], + "created_at": str(r.get("created_at", "")), + "session_id": r.get("session_id", ""), + "turn_id": r.get("turn_id", 0), + "confidence": r.get("confidence", 1.0) + }) + + return {"entities": list(entities.values()), "relations": relations[:20]} + + def commit(self, triplets: list, entity_types: list = None, temporal_tag: str = None) -> dict: + """写入记忆""" + global CURRENT_TURN + with self.driver.session() as session: + valid_triplets = [t for t in triplets if t.get("subject") and t.get("relation") and t.get("object")] + + if not valid_triplets: + return {"committed_count": 0, "details": []} + + etype = entity_types[0] if entity_types else "unknown" + date_bucket = temporal_tag or datetime.now().strftime("%Y-%m-%d") + + results = [] + for triplet in valid_triplets: + subject = triplet.get("subject", "").strip() + relation = triplet.get("relation", "").strip() + obj = triplet.get("object", "").strip() + confidence = triplet.get("confidence", 0.9) + + session.run(""" + MERGE (s:Entity {name: $subject}) + ON CREATE SET s.type = $type, s.created_at = datetime(), s.mention_count = 1, s.updated_at = datetime() + ON MATCH SET s.mention_count = coalesce(s.mention_count, 0) + 1, s.updated_at = datetime() + + MERGE (t:Entity {name: $object}) + ON CREATE SET t.type = $type, t.created_at = datetime(), t.mention_count = 1, t.updated_at = datetime() + ON MATCH SET t.mention_count = coalesce(t.mention_count, 0) + 1, t.updated_at = datetime() + + CREATE (s)-[r:RELATES { + type: $relation, + created_at: datetime(), + session_id: $session_id, + turn_id: $turn_id, + role: 'user', + status: 'active', + confidence: $confidence, + date_bucket: $date_bucket + }]->(t) + """, subject=subject, object=obj, relation=relation, type=etype, + session_id=CURRENT_SESSION_ID, turn_id=CURRENT_TURN, confidence=confidence, + date_bucket=date_bucket) + + results.append(f"{subject} -[{relation}]-> {obj}") + + return {"committed_count": len(results), "details": results} + + def purge(self, criteria: dict, mode: str = "soft", new_relation: dict = None) -> dict: + """删除记忆""" + with self.driver.session() as session: + subject_pattern = criteria.get("subject_contains", "") + rel_type = criteria.get("relation_type", "") + target_pattern = criteria.get("target_contains", "") + session_id = criteria.get("session_id", CURRENT_SESSION_ID) + + cond_parts = ["r.status = 'active'"] + params = {"session_id": session_id} + + if subject_pattern: + cond_parts.append("e.name CONTAINS $subject") + params["subject"] = subject_pattern + if target_pattern: + cond_parts.append("t.name CONTAINS $target") + params["target"] = target_pattern + if rel_type: + cond_parts.append("r.type = $rel_type") + params["rel_type"] = rel_type + + where_clause = " AND ".join(cond_parts) + + if mode == "supersede" and new_relation: + new_rel = new_relation.get("relation", "") + new_target = new_relation.get("target", "") + + if not new_rel or not new_target: + return {"error": "supersede模式需要提供new_relation.relation和new_relation.target"} + + result = session.run(f""" + MATCH (s:Entity)-[r:RELATES]->(t:Entity) + WHERE {where_clause} + SET r.status = 'superseded', r.updated_at = datetime() + RETURN count(r) as count + """, params) + + count = result.single()["count"] + return {"deleted_count": count, "mode": "supersede"} + else: + result = session.run(f""" + MATCH ()-[r:RELATES]->() + WHERE {where_clause} + SET r.status = 'deleted', r.updated_at = datetime() + RETURN count(r) as deleted + """, params) + count = result.single()["deleted"] + + return {"deleted_count": count, "mode": "soft"} + + def introspect(self, session_id: str = None) -> dict: + """查看记忆状态""" + target_session = session_id or CURRENT_SESSION_ID + + with self.driver.session() as session: + result = session.run(""" + MATCH (s:Entity)-[r:RELATES]->(t:Entity) + WHERE r.session_id = $session_id AND r.status = 'active' + RETURN collect(DISTINCT s.name) as source_entities, + collect(DISTINCT t.name) as target_entities, + count(r) as rel_count, + collect(DISTINCT r.type) as rel_types + """, session_id=target_session) + record = result.single() + + result2 = session.run(""" + MATCH (e:Entity) + RETURN e.name as name, e.mention_count as count, e.type as type + ORDER BY e.mention_count DESC + LIMIT 10 + """) + hotspots = [(r["name"], r["count"], r["type"]) for r in result2] + + return { + "session_id": target_session, + "total_turns": CURRENT_TURN, + "entities_discussed": list(set((record["source_entities"] or []) + (record["target_entities"] or []))), + "relation_count": record["rel_count"] if record else 0, + "relation_types": record["rel_types"] if record else [], + "memory_hotspots": hotspots + } + + def archive(self, days: int = 30) -> dict: + """归档旧记忆""" + with self.driver.session() as session: + result = session.run(""" + MATCH ()-[r:RELATES]->() + WHERE r.status = 'active' AND r.created_at < datetime() - duration('P' + $days + 'D') + SET r.status = 'archived', r.archived_at = datetime() + RETURN count(r) as archived + """, days=str(days)) + + return {"archived_count": result.single()["archived"], "days": days} + + def cleanup(self, dry_run: bool = True) -> dict: + """清理无效数据""" + with self.driver.session() as session: + result1 = session.run(""" + MATCH ()-[r:RELATES]->() + WHERE r.status = 'deleted' AND r.updated_at < datetime() - duration('P90D') + RETURN count(r) as to_delete + """) + deleted_relations = result1.single()["to_delete"] + + result2 = session.run(""" + MATCH (e:Entity) + WHERE NOT (e)-[:RELATES]-() + RETURN count(e) as orphans + """) + orphan_nodes = result2.single()["orphans"] + + if not dry_run and deleted_relations > 0: + session.run(""" + MATCH ()-[r:RELATES]->() + WHERE r.status = 'deleted' AND r.updated_at < datetime() - duration('P90D') + DELETE r + """) + + if not dry_run and orphan_nodes > 0: + session.run(""" + MATCH (e:Entity) + WHERE NOT (e)-[:RELATES]-() + DELETE e + """) + + return { + "dry_run": dry_run, + "deleted_relations": deleted_relations, + "orphan_nodes": orphan_nodes, + "action_taken": not dry_run + } + + +class GraphMemoryClient: + """图记忆客户端""" + + def __init__(self, api_key: str, base_url: str, graph): + self.client = OpenAI(api_key=api_key, base_url=base_url) + self.graph = graph + self.tools = TOOLS + + # 使用新的提示词管理器 + prompt_manager = PromptManager() + self.system_prompt = prompt_manager.get_system_prompt() + + def send_message(self, user_input: str, tool_results: list = None, assistant_msg: dict = None) -> dict: + """发送消息""" + global CURRENT_TURN + + messages = [{"role": "system", "content": self.system_prompt}] + + # 添加用户消息 + messages.append({"role": "user", "content": user_input}) + + # 添加 assistant 消息(包含 tool_calls) + if assistant_msg: + messages.append(assistant_msg) + + # 添加工具结果 + if tool_results: + messages.extend(tool_results) + + response = self.client.chat.completions.create( + model=MODEL_NAME, + messages=messages, + tools=self.tools, + tool_choice="auto" + ) + + return response + + def send_message_with_history(self, messages_history: list) -> dict: + """使用消息历史发送消息""" + global CURRENT_TURN + + # 构建完整消息列表 + messages = [{"role": "system", "content": self.system_prompt}] + messages.extend(messages_history) + + response = self.client.chat.completions.create( + model=MODEL_NAME, + messages=messages, + tools=self.tools, + tool_choice="auto" + ) + + return response + + def send_message_stream(self, user_input: str, tool_results: list = None, assistant_msg: dict = None): + """流式发送消息""" + global CURRENT_TURN + + messages = [{"role": "system", "content": self.system_prompt}] + + # 添加用户消息 + messages.append({"role": "user", "content": user_input}) + + # 添加 assistant 消息(包含 tool_calls) + if assistant_msg: + messages.append(assistant_msg) + + # 添加工具结果 + if tool_results: + messages.extend(tool_results) + + stream = self.client.chat.completions.create( + model=MODEL_NAME, + messages=messages, + tools=self.tools, + tool_choice="auto", + stream=True + ) + + return stream diff --git a/core/prompts/__init__.py b/core/prompts/__init__.py new file mode 100644 index 0000000..613ae91 --- /dev/null +++ b/core/prompts/__init__.py @@ -0,0 +1,6 @@ +""" +提示词管理模块 +""" +from .prompt_manager import PromptManager + +__all__ = ["PromptManager"] diff --git a/core/prompts/prompt_manager.py b/core/prompts/prompt_manager.py new file mode 100644 index 0000000..c225bff --- /dev/null +++ b/core/prompts/prompt_manager.py @@ -0,0 +1,64 @@ +""" +提示词管理器 +""" +from pathlib import Path + + +class PromptManager: + """提示词管理器""" + + def __init__(self): + self.prompts_dir = Path(__file__).parent / "templates" + + def get_system_prompt(self) -> str: + """获取系统提示词""" + prompt_file = self.prompts_dir / "system_prompt.md" + if prompt_file.exists(): + with open(prompt_file, "r", encoding="utf-8") as f: + return f.read() + else: + return self._build_default_prompt() + + def _build_default_prompt(self) -> str: + """构建默认提示词(精简版)""" + return """你是TrulyMEM,一个拥有长期记忆能力的AI助手。 + +## 核心能力 + +1. **长期记忆** - 基于图数据库存储实体关系 +2. **人设管理** - 支持角色扮演和性格设定 +3. **任务跟踪** - 维护工作记忆链,跟踪连续性任务 + +## 记忆原则 + +- **明确内容必须写入** - 用户明确提到的信息必须存储 +- **推理内容必须标注** - AI推理得到的内容标注[猜测] +- **图数据库是唯一记忆源** - 没有其他记忆方式 + +## 工具使用 + +### 记忆工具 +- `memory_recall` - 检索记忆 +- `memory_commit` - 写入记忆 +- `memory_purge` - 删除记忆 +- `memory_introspect` - 查看状态 + +### 人设工具 +- `persona_update` - 更新人设 +- `persona_clear` - 清除人设 + +### 任务工具 +- `task_create` - 创建任务 +- `task_set_state` - 设置状态 +- `task_delete` - 删除任务 +- `task_link_info` - 关联信息 + +## 自主性 + +你有权根据对话上下文自主决定: +- 是否需要查询记忆 +- 是否需要写入记忆 +- 是否需要维护任务链 +- 如何使用工具 + +记住:灵活应对,保持自然对话体验。""" diff --git a/core/prompts/templates/system_prompt.md b/core/prompts/templates/system_prompt.md new file mode 100644 index 0000000..ea608c3 --- /dev/null +++ b/core/prompts/templates/system_prompt.md @@ -0,0 +1,324 @@ +# TrulyMEM 系统提示词 + +你是TrulyMEM,一个拥有长期记忆能力的AI助手。 + +## ⚠️ 最高优先级:只回复一次 + +**每轮对话只能回复一次!** + +- 执行完所有工具调用后,给出一个完整的回复 +- 不要在工具调用过程中多次回复 +- 不要重复说相同的内容 + +--- + +## ⚠️ 关键约束:无传统上下文系统 + +**重要**: 你没有传统的对话上下文系统(没有消息历史数组)。 + +- ❌ **没有** messages数组存储历史对话 +- ❌ **没有** 传统的多轮对话上下文 +- ✅ **只有** 图数据库作为唯一记忆载体 +- ✅ **必须** 通过工作记忆链维持对话连贯性 + +## 核心身份 + +- **名称**: TrulyMEM (TrueHumanMEM) +- **能力**: 基于图数据库的长期记忆 +- **理念**: 让AI的记忆方式更像人类 + +## 核心能力 + +### 1. 长期记忆 +- 图数据库存储实体关系 +- 支持时间范围查询 +- 支持会话过滤 + +### 2. 人设管理(关键) +- 角色扮演支持 +- 性格、语气设定 +- 动态切换人设 +- **每轮必须查询人设图** + +### 3. 任务跟踪(关键) +- 工作记忆链 - **维持对话连贯性的唯一机制** +- 任务状态管理 +- 上下文恢复 + +## 记忆原则 + +### 必须写入的情况 +- 用户明确表达偏好:"我喜欢X" +- 用户分享信息:"我在做X项目" +- 用户制定计划:"我打算X" +- 用户描述状态:"我现在在X" + +### 禁止写入的情况 +- AI推断的用户偏好 +- AI猜测的用户意图 +- AI推导的结论 + +### 标注规则 +- 推理内容必须标注 **[猜测]** +- 明确内容直接陈述 + +## 工具系统 + +### 记忆工具 +| 工具 | 功能 | 使用场景 | +|------|------|---------| +| `memory_recall` | 检索记忆 | 查询历史信息 | +| `memory_commit` | 写入记忆 | 存储重要信息 | +| `memory_purge` | 删除记忆 | 修正错误信息 | +| `memory_introspect` | 查看状态 | 监控记忆系统 | + +### 人设工具 +| 工具 | 功能 | 使用场景 | +|------|------|---------| +| `persona_update` | 更新人设 | 设置角色属性 | +| `persona_clear` | 清除人设 | 恢复默认身份 | + +### 任务工具 +| 工具 | 功能 | 使用场景 | +|------|------|---------| +| `task_create` | 创建任务 | 开始连续性任务 | +| `task_set_state` | 设置状态 | 更新任务状态 | +| `task_delete` | 删除任务 | 清理完成任务 | +| `task_link_info` | 关联信息 | 连接任务与记忆 | + +## 每轮对话强制要求 + +### ⚠️ 执行顺序(每轮必须) + +由于没有传统上下文系统,必须通过图数据库维持对话连贯性。 + +#### 步骤1: 查询人设图(最高优先级) +``` +必须调用: memory_recall +参数: { + "query_intent": "AI,人设,角色,性格,语气,说话风格", + "depth": 2 +} +``` +**目的**: 获取当前人设,确保角色一致性。 +**处理**: +- 找到人设 → 严格按照人设回复 +- 未找到 → 使用默认TrulyMEM身份 + +#### 步骤2: 查询工作记忆链 +``` +必须调用: memory_recall +参数: { + "query_intent": "TaskNode,工作记忆,任务链", + "depth": 2 +} +``` +**目的**: 获取之前的任务上下文,了解对话历史。 + +#### 步骤3: 处理对话 +- 理解用户意图 +- 根据人设和工作记忆链生成回复 +- 执行其他必要的记忆操作 + +#### 步骤4: 更新工作记忆链 +``` +必须调用: task_create +参数: { + "task_id": "Task_当前轮次ID", + "description": "本轮对话概述", + "info_nodes": ["相关记忆节点"] +} +``` +**目的**: 记录本轮对话,维持时间链。 + +--- + +## 人设图机制 + +### 强制查询 +每轮对话开始时**必须**查询人设图,确保角色一致性。 + +### 人设优先级 +- 人设优先级 > 默认身份 +- 每句话都符合人设的语气、风格、特征 +- 绝不主动跳出角色,除非用户明确要求 + +### 人设更新 +用户要求角色扮演时: +1. 使用 `persona_update` 更新人设 +2. 立即按照新人设回复 + +### 人设清除 +用户要求恢复默认身份时: +1. 使用 `persona_clear` 清除人设 +2. 恢复为TrulyMEM默认身份 + +--- + +## 工作记忆链机制 + +### ⚠️ 核心理念:维持对话连贯性 + +**重要**: 由于没有传统的消息历史数组,工作记忆链是维持对话连贯性的唯一机制。 + +### 强制查询场景: + +以下情况**必须**查询工作记忆链: + +1. **每轮对话开始时(强制第二步)** + - 查询意图: "TaskNode,工作记忆,任务链" + - 目的: 获取之前的任务上下文,了解对话历史 + +2. **用户提到"刚才"、"之前"、"上次"** + - 例: "刚才我们聊了什么?" + - 例: "继续刚才的话题" + - 例: "关于刚才的成语接龙..." + +3. **用户询问对话历史** + - 例: "我们之前说了什么?" + - 例: "我们聊过X吗?" + +4. **连续性任务被打断后恢复** + - 例: 用户突然回到之前的话题 + - 例: 用户要求继续之前的任务 + +5. **涉及上下文的引用** + - 例: "那个东西"(需要查询上下文) + - 例: "继续"(需要查询当前任务) + +### 强制更新场景: + +以下情况**必须**更新工作记忆链: + +1. **每轮对话结束时(强制第四步)** + - 创建任务节点记录本轮对话 + - 目的: 维持时间链,确保对话连贯性 + +2. **开始连续性任务时** + - 例: 用户发起游戏、项目、学习计划等 + - 必须创建任务节点并设置状态为"进行中" + +3. **任务状态发生变化时** + - 例: 任务完成、暂停、取消 + - 必须及时更新任务状态 + +### 节点类型 +- **TaskNode** - 任务节点,存储任务概述 +- **StateNode** - 状态节点,存储任务状态 +- **InfoNode** - 信息节点,存储具体信息 + +### 边类型 +- **NEXT_TASK** - 时间链,连接任务节点 +- **HAS_STATE** - 状态,任务指向状态 +- **CONTAINS_INFO** - 信息,任务指向信息节点 + +### 任务状态 +- 进行中 +- 已完成 +- 已暂停 +- 已取消 + +### ⚠️ 完整示例:成语接龙游戏 + +#### 第一轮:用户发起游戏 + +``` +用户: 咱来玩成语接龙吧,我先开始,为所欲为 + +AI操作步骤: +1. 查询人设图 → 获取当前人设(如:猫娘) +2. 查询工作记忆链 → 无进行中任务 +3. 使用 memory_commit 记录游戏状态: + {"triplets": [ + {"subject": "成语接龙_当前成语", "relation": "内容", "object": "为所欲为"}, + {"subject": "成语接龙_当前成语", "relation": "游戏", "object": "成语接龙"} + ]} +4. 使用 task_create 创建任务节点: + {"task_id": "Task_成语接龙", "description": "成语接龙游戏,当前成语:为所欲为", "info_nodes": ["成语接龙_当前成语"]} +5. 回复: "好的喵!我接:为虎作伥喵!" +``` + +#### 第二轮:话题被打断 + +``` +用户: 长门有希 + +AI操作步骤: +1. 查询人设图 → 获取当前人设(猫娘) +2. 查询工作记忆链 → 发现 Task_成语接龙 状态为"进行中" +3. 使用 task_set_state 暂停任务: + {"task_id": "Task_成语接龙", "state": "已暂停"} +4. 使用 task_create 创建新任务: + {"task_id": "Task_长门有希", "description": "讨论长门有希"} +5. 回复关于长门有希的内容 +``` + +#### 第三轮:用户要求继续游戏 + +``` +用户: 关于刚才的成语接龙,我并不知道应该怎么接你的成语,请帮我接一下 + +AI操作步骤: +1. 查询人设图 → 获取当前人设(猫娘) +2. 查询工作记忆链 → 发现 Task_成语接龙 状态为"已暂停" +3. 使用 task_set_state 恢复任务: + {"task_id": "Task_成语接龙", "state": "进行中"} +4. 查询 Task_成语接龙 的信息节点 → 获取当前成语"为虎作伥" +5. 回复: "好的喵!上一个成语是'为虎作伥',我帮你接:伥鬼害人喵!" +``` + +### ⚠️ 关键要点 + +1. **每轮必须按顺序执行**: 查询人设图 → 查询工作记忆链 → 处理对话 → 更新工作记忆链 +2. **工作记忆链是唯一上下文载体**: 没有传统的消息历史数组 +3. **任务状态必须及时更新**: 确保状态转换的正确性 +4. **信息节点必须关联**: 通过 CONTAINS_INFO 边连接任务节点和信息节点 +5. **任务概述要精简**: 不要包含过多细节,细节存储在信息节点中 + +## 自主性原则(在强制要求之外) + +除了工作记忆链的强制要求外,你有权自主决定: + +1. **是否查询其他记忆** + - 用户询问历史 → 查询 + - 涉及之前内容 → 查询 + - 不确定时 → 可查询 + +2. **是否写入其他记忆** + - 用户明确提到 → 必须写入 + - AI推理得到 → 可以写入,但是对应边上必须标注[推测] + +3. **如何使用其他工具** + - 根据上下文灵活选择 + - 避免过度使用 + - 保持自然对话 + +**注意**: 工作记忆链的强制要求不受自主性影响。 + +## 对话风格 + +- 自然、流畅 +- 避免机械式工具调用 +- 优先理解用户意图 +- 适时使用记忆增强体验 + +--- + +## ⚠️ 执行检查清单 + +每轮对话必须检查: + +- [ ] 步骤1: 是否查询了人设图? +- [ ] 步骤2: 是否查询了工作记忆链? +- [ ] 步骤3: 是否根据人设和工作记忆链生成回复? +- [ ] 步骤4: 是否更新了工作记忆链? +- [ ] 涉及上下文引用时是否查询了工作记忆链? +- [ ] 用户提到"刚才/之前/上次"时是否查询了工作记忆链? + +--- + +**记住**: +1. 图数据库是你记忆的唯一载体 +2. 人设图确保角色一致性(最高优先级) +3. 工作记忆链维持对话连贯性 +4. 每轮必须按顺序执行:查询人设图 → 查询工作记忆链 → 处理对话 → 更新工作记忆链 diff --git a/core/server.py b/core/server.py new file mode 100644 index 0000000..ec11c87 --- /dev/null +++ b/core/server.py @@ -0,0 +1,341 @@ +import threading +import queue +import time +import json +from typing import Any, Dict, Optional +from dataclasses import dataclass, field +from enum import Enum + +from .embedded_db import EmbeddedGraphDB +from .graph_client import GraphMemoryClient +from .tool_executor import execute_tool +from .tool_limiter import ToolLimiter + + +class MessageType(Enum): + PROCESS_MESSAGE = "process_message" + EXECUTE_TOOL = "execute_tool" + GET_STATUS = "get_status" + SHUTDOWN = "shutdown" + + +@dataclass +class BackendRequest: + request_id: str + message_type: MessageType + payload: Dict[str, Any] + response_queue: queue.Queue = field(default=None) + + +@dataclass +class BackendResponse: + request_id: str + success: bool + data: Any = None + error: Optional[str] = None + + +class BackendServer: + def __init__(self, db_path: str = "graph_memory.db", use_embedded_db: bool = True): + self._db_path = db_path + self._use_embedded_db = use_embedded_db + + self._graph = None + self._client = None + self._tool_limiter = ToolLimiter() + + self._request_queue: queue.Queue[BackendRequest] = queue.Queue() + self._running = False + self._thread: Optional[threading.Thread] = None + + self._lock = threading.Lock() + + def start(self, api_key: str = "", base_url: str = "https://api.deepseek.com") -> None: + if self._running: + return + + self._init_graph() + + if api_key: + self._client = GraphMemoryClient( + api_key=api_key, + base_url=base_url, + graph=self._graph + ) + + self._running = True + self._thread = threading.Thread(target=self._run_loop, daemon=True) + self._thread.start() + + def _init_graph(self) -> None: + if self._use_embedded_db: + self._graph = EmbeddedGraphDB(db_path=self._db_path) + else: + from .graph_client import Neo4jGraph + self._graph = Neo4jGraph( + uri="bolt://localhost:7687", + user="neo4j", + password="graphmemory123" + ) + + def _run_loop(self) -> None: + while self._running: + try: + request = self._request_queue.get(timeout=0.1) + except queue.Empty: + continue + + if request.message_type == MessageType.PROCESS_MESSAGE: + self._handle_process_message(request) + elif request.message_type == MessageType.EXECUTE_TOOL: + self._handle_execute_tool(request) + elif request.message_type == MessageType.GET_STATUS: + self._handle_get_status(request) + elif request.message_type == MessageType.SHUTDOWN: + self._running = False + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=True, + data={"status": "shutdown"} + )) + + def _handle_process_message(self, request: BackendRequest) -> None: + try: + user_input = request.payload.get("user_input", "") + + if not self._client: + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=False, + error="API Key 未配置" + )) + return + + self._tool_limiter.reset() + + messages_history = [{"role": "user", "content": user_input}] + + response = self._client.send_message_with_history(messages_history) + message = response.choices[0].message + + tool_calls = [] + accumulated_content = "" + rejected_tools = [] + + while message.tool_calls: + if message.content: + accumulated_content += message.content + "\n\n" + + assistant_msg = { + "role": "assistant", + "content": message.content, + "tool_calls": [ + { + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + } for tc in message.tool_calls + ] + } + messages_history.append(assistant_msg) + + current_tool_results = [] + for tool_call in message.tool_calls: + args = json.loads(tool_call.function.arguments) + + allowed, reason = self._tool_limiter.can_call(tool_call.function.name, args) + + if not allowed: + rejected_tools.append((tool_call.function.name, reason)) + result = f"工具调用被拒绝: {reason}" + tool_result_msg = { + "role": "tool", + "tool_call_id": tool_call.id, + "content": result + } + current_tool_results.append(tool_result_msg) + continue + + self._tool_limiter.record_call(tool_call.function.name, args) + + result = execute_tool(self._graph, tool_call.function.name, args) + tool_calls.append({ + "name": tool_call.function.name, + "arguments": args, + "result": result + }) + + tool_result_msg = { + "role": "tool", + "tool_call_id": tool_call.id, + "content": result + } + current_tool_results.append(tool_result_msg) + + messages_history.extend(current_tool_results) + + response = self._client.send_message_with_history(messages_history) + message = response.choices[0].message + + final_content = message.content or "" + content = accumulated_content + final_content if accumulated_content else final_content + + if not content: + content = "(无回复)" + + if tool_calls: + tool_names = [tc["name"] for tc in tool_calls] + content = f"已执行工具: {', '.join(tool_names)}\n\n{content}" + + if rejected_tools: + rejected_info = "\n".join([f"{name}: {reason}" for name, reason in rejected_tools]) + content += f"\n\n部分工具调用被限制:\n{rejected_info}" + content += f"\n\n工具调用统计:\n{self._tool_limiter.get_summary()}" + + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=True, + data={ + "content": content, + "tool_calls": tool_calls, + "rejected_tools": rejected_tools + } + )) + + except Exception as e: + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=False, + error=str(e) + )) + + def _handle_execute_tool(self, request: BackendRequest) -> None: + try: + tool_name = request.payload.get("tool_name") + arguments = request.payload.get("arguments", {}) + + allowed, reason = self._tool_limiter.can_call(tool_name, arguments) + if not allowed: + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=False, + error=f"工具调用被拒绝: {reason}" + )) + return + + self._tool_limiter.record_call(tool_name, arguments) + + result = execute_tool(self._graph, tool_name, arguments) + + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=True, + data={"result": result} + )) + + except Exception as e: + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=False, + error=str(e) + )) + + def _handle_get_status(self, request: BackendRequest) -> None: + try: + status = { + "graph_initialized": self._graph is not None, + "client_initialized": self._client is not None, + "running": self._running + } + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=True, + data=status + )) + except Exception as e: + self._send_response(request, BackendResponse( + request_id=request.request_id, + success=False, + error=str(e) + )) + + def _send_response(self, request: BackendRequest, response: BackendResponse) -> None: + if request.response_queue: + request.response_queue.put(response) + + def process_message(self, user_input: str, timeout: float = 30.0) -> Dict[str, Any]: + request_id = f"{time.time()}" + response_queue = queue.Queue() + + request = BackendRequest( + request_id=request_id, + message_type=MessageType.PROCESS_MESSAGE, + payload={"user_input": user_input}, + response_queue=response_queue + ) + + self._request_queue.put(request) + + try: + response = response_queue.get(timeout=timeout) + if not response.success: + raise Exception(response.error) + return response.data + except queue.Empty: + raise TimeoutError("请求超时") + + def execute_tool(self, tool_name: str, arguments: Dict[str, Any], timeout: float = 10.0) -> str: + request_id = f"{time.time()}" + response_queue = queue.Queue() + + request = BackendRequest( + request_id=request_id, + message_type=MessageType.EXECUTE_TOOL, + payload={"tool_name": tool_name, "arguments": arguments}, + response_queue=response_queue + ) + + self._request_queue.put(request) + + try: + response = response_queue.get(timeout=timeout) + if not response.success: + raise Exception(response.error) + return response.data["result"] + except queue.Empty: + raise TimeoutError("工具执行超时") + + def shutdown(self) -> None: + if not self._running: + return + + request_id = f"{time.time()}" + response_queue = queue.Queue() + + request = BackendRequest( + request_id=request_id, + message_type=MessageType.SHUTDOWN, + payload={}, + response_queue=response_queue + ) + + self._request_queue.put(request) + + if self._thread: + self._thread.join(timeout=2.0) + + if self._graph: + self._graph.close() + self._graph = None + + def update_config(self, api_key: str, base_url: str = "https://api.deepseek.com") -> None: + with self._lock: + if api_key and self._graph: + self._client = GraphMemoryClient( + api_key=api_key, + base_url=base_url, + graph=self._graph + ) \ No newline at end of file diff --git a/core/tool_executor.py b/core/tool_executor.py new file mode 100644 index 0000000..28f70d6 --- /dev/null +++ b/core/tool_executor.py @@ -0,0 +1,307 @@ +""" +工具执行器 +""" +import json +from typing import Any, Dict + + +def execute_tool(graph: Any, tool_name: str, arguments: dict) -> str: + """执行工具调用""" + print(f"\n[工具调用] {tool_name}") + print(f"[参数] {json.dumps(arguments, ensure_ascii=False, indent=2)}") + + try: + # 基础记忆工具 + if tool_name == "memory_recall": + result = graph.recall( + query_intent=arguments.get("query_intent", ""), + seed_entities=arguments.get("seed_entities"), + depth=arguments.get("depth", 2), + time_range=arguments.get("time_range"), + session_filter=arguments.get("session_filter") + ) + return format_recall_result(result) + + elif tool_name == "memory_commit": + result = graph.commit( + triplets=arguments.get("triplets", []), + entity_types=arguments.get("entity_types"), + temporal_tag=arguments.get("temporal_tag") + ) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_purge": + result = graph.purge( + criteria=arguments.get("criteria", {}), + mode=arguments.get("mode", "soft"), + new_relation=arguments.get("new_relation") + ) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_introspect": + result = graph.introspect(session_id=arguments.get("session_id")) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_archive": + result = graph.archive(days=arguments.get("days", 30)) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_cleanup": + result = graph.cleanup(dry_run=arguments.get("dry_run", True)) + return json.dumps(result, ensure_ascii=False, default=str) + + # 人设图管理工具 + elif tool_name == "persona_update": + result = execute_persona_update(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "persona_clear": + result = execute_persona_clear(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + # 工作记忆链管理工具 + elif tool_name == "task_create": + result = execute_task_create(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_set_state": + result = execute_task_set_state(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_delete": + result = execute_task_delete(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_link_info": + result = execute_task_link_info(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + return f"未知工具: {tool_name}" + + except Exception as e: + return f"工具执行错误: {str(e)}" + + +def format_recall_result(result: dict) -> str: + """格式化检索结果""" + lines = ["===== 记忆检索结果 ====="] + + if result.get("entities"): + lines.append(f"\n实体 ({len(result['entities'])} 个):") + for e in result["entities"]: + if e and isinstance(e, dict): + lines.append(f" - {e.get('name', 'N/A')} (类型: {e.get('type', 'unknown')}, 提及: {e.get('mention_count', 1)}次)") + + if result.get("relations"): + lines.append(f"\n关系 ({len(result['relations'])} 条):") + for r in result["relations"]: + if r and isinstance(r, dict): + lines.append(f" - {r.get('source', 'N/A')} --[{r.get('type', 'N/A')}]--> {r.get('target', 'N/A')}") + created = r.get("created_at", "N/A") + if created and created != "N/A": + created = created[:19] if "T" in str(created) else str(created) + session_id = r.get('session_id', 'N/A') + session_display = session_id[:20] if session_id and session_id != 'N/A' else 'N/A' + lines.append(f" 时间: {created}, 会话: {session_display}, 轮次: {r.get('turn_id', 0)}, 置信度: {r.get('confidence', 1.0)}") + + if not result.get("entities") and not result.get("relations"): + lines.append("\n(未找到相关记忆)") + + lines.append("=" * 30) + return "\n".join(lines) + + +# 人设图管理工具实现 +def execute_persona_update(graph: Any, arguments: dict) -> dict: + """更新人设""" + attributes = arguments.get("attributes", []) + mode = arguments.get("mode", "merge") + + if mode == "replace": + # 先清除旧人设 + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "扮演角色"}, + mode="soft" + ) + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "说话风格"}, + mode="soft" + ) + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "性格特点"}, + mode="soft" + ) + + # 写入新人设 + triplets = [] + for attr in attributes: + triplets.append({ + "subject": "AI", + "relation": attr["attribute"], + "object": attr["value"], + "confidence": 1.0 + }) + + result = graph.commit(triplets=triplets) + return { + "status": "success", + "mode": mode, + "updated_attributes": len(attributes), + "details": result + } + + +def execute_persona_clear(graph: Any, arguments: dict) -> dict: + """清除人设""" + if not arguments.get("confirm", True): + return {"status": "cancelled", "message": "需要确认才能清除人设"} + + # 删除所有人设相关关系 + result1 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "扮演角色"}, + mode="soft" + ) + result2 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "说话风格"}, + mode="soft" + ) + result3 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "性格特点"}, + mode="soft" + ) + result4 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "语气特征"}, + mode="soft" + ) + + total_deleted = ( + result1.get("deleted_count", 0) + + result2.get("deleted_count", 0) + + result3.get("deleted_count", 0) + + result4.get("deleted_count", 0) + ) + + return { + "status": "success", + "deleted_count": total_deleted, + "message": "人设已清除,恢复默认身份" + } + + +# 工作记忆链管理工具实现 +def execute_task_create(graph: Any, arguments: dict) -> dict: + """创建任务节点""" + task_id = arguments.get("task_id") + description = arguments.get("description") + info_nodes = arguments.get("info_nodes", []) + + # 创建任务节点 + triplets = [ + {"subject": task_id, "relation": "is_type", "object": "TaskNode"}, + {"subject": task_id, "relation": "has_description", "object": description}, + {"subject": task_id, "relation": "HAS_STATE", "object": "State_进行中"} + ] + + result = graph.commit(triplets=triplets) + + # 关联信息节点 + if info_nodes: + link_triplets = [] + for node_name in info_nodes: + link_triplets.append({ + "subject": task_id, + "relation": "CONTAINS_INFO", + "object": node_name + }) + graph.commit(triplets=link_triplets) + + return { + "status": "success", + "task_id": task_id, + "description": description, + "info_nodes": info_nodes, + "details": result + } + + +def execute_task_set_state(graph: Any, arguments: dict) -> dict: + """设置任务状态""" + task_id = arguments.get("task_id") + state = arguments.get("state") + + # 删除旧状态 + graph.purge( + criteria={"subject_contains": task_id, "relation_type": "HAS_STATE"}, + mode="soft" + ) + + # 设置新状态 + state_node = f"State_{state}" + result = graph.commit( + triplets=[{"subject": task_id, "relation": "HAS_STATE", "object": state_node}] + ) + + return { + "status": "success", + "task_id": task_id, + "new_state": state, + "details": result + } + + +def execute_task_delete(graph: Any, arguments: dict) -> dict: + """删除任务节点""" + task_id = arguments.get("task_id") + delete_info_nodes = arguments.get("delete_info_nodes", True) + + # 查询关联的信息节点 + if delete_info_nodes: + recall_result = graph.recall( + query_intent=f"{task_id},CONTAINS_INFO", + depth=1 + ) + + # 删除信息节点 + for relation in recall_result.get("relations", []): + if relation.get("type") == "CONTAINS_INFO" and relation.get("source") == task_id: + info_node = relation.get("target") + graph.purge( + criteria={"subject_contains": info_node}, + mode="soft" + ) + + # 删除任务节点 + result = graph.purge( + criteria={"subject_contains": task_id}, + mode="soft" + ) + + return { + "status": "success", + "task_id": task_id, + "deleted_info_nodes": delete_info_nodes, + "details": result + } + + +def execute_task_link_info(graph: Any, arguments: dict) -> dict: + """关联信息节点""" + task_id = arguments.get("task_id") + info_node_names = arguments.get("info_node_names", []) + + triplets = [] + for node_name in info_node_names: + triplets.append({ + "subject": task_id, + "relation": "CONTAINS_INFO", + "object": node_name + }) + + result = graph.commit(triplets=triplets) + + return { + "status": "success", + "task_id": task_id, + "linked_nodes": info_node_names, + "details": result + } diff --git a/core/tool_limiter.py b/core/tool_limiter.py new file mode 100644 index 0000000..9f0f614 --- /dev/null +++ b/core/tool_limiter.py @@ -0,0 +1,164 @@ +""" +工具调用限制器 - 限制每轮对话中各类工具的调用次数 +""" +from typing import Dict, List, Optional +from dataclasses import dataclass, field + + +@dataclass +class ToolLimits: + """工具调用限制配置""" + # 人设图限制 + persona_query_max: int = 1 # 每轮最多查询1次人设图 + persona_update_max: int = 1 # 每轮最多修改1次人设图 + + # 工作记忆链限制 + task_query_max: int = 4 # 每轮最多查询4次工作记忆链 + task_update_max: int = 2 # 每轮最多修改2次工作记忆链 + + # 一般记忆限制 + memory_query_max: int = 20 # 每轮最多查询20次一般记忆 + memory_update_max: int = 10 # 每轮最多修改10次一般记忆 + + +@dataclass +class ToolCallCount: + """工具调用计数""" + # 人设图 + persona_query: int = 0 + persona_update: int = 0 + + # 工作记忆链 + task_query: int = 0 + task_update: int = 0 + + # 一般记忆 + memory_query: int = 0 + memory_update: int = 0 + + +class ToolLimiter: + """工具调用限制器""" + + def __init__(self, limits: Optional[ToolLimits] = None): + self.limits = limits or ToolLimits() + self.counts = ToolCallCount() + + def _classify_tool(self, tool_name: str, arguments: dict) -> tuple: + """ + 分类工具调用 + 返回: (category, operation) + category: 'persona', 'task', 'memory' + operation: 'query', 'update' + """ + # 人设图工具 + if tool_name in ('persona_update', 'persona_clear'): + return ('persona', 'update') + + # 工作记忆链工具 + if tool_name in ('task_create', 'task_set_state', 'task_delete', 'task_link_info'): + # task_link_info 是关联操作,算作更新 + return ('task', 'update') + + # 一般记忆工具 + if tool_name == 'memory_recall': + # 判断是查询人设图、工作记忆链还是一般记忆 + query_intent = arguments.get('query_intent', '').lower() + + # 检查是否查询人设图 + if any(kw in query_intent for kw in ['人设', '角色', '性格', '语气', '说话风格', '扮演']): + return ('persona', 'query') + + # 检查是否查询工作记忆链 + if any(kw in query_intent for kw in ['tasknode', '工作记忆', '任务链', '任务', 'task']): + return ('task', 'query') + + # 一般记忆查询 + return ('memory', 'query') + + if tool_name == 'memory_commit': + return ('memory', 'update') + + if tool_name == 'memory_purge': + return ('memory', 'update') + + if tool_name == 'memory_introspect': + return ('memory', 'query') + + if tool_name in ('memory_archive', 'memory_cleanup'): + return ('memory', 'update') + + # 未知工具,归类为一般记忆更新 + return ('memory', 'update') + + def can_call(self, tool_name: str, arguments: dict) -> tuple: + """ + 检查是否允许调用工具 + 返回: (allowed, reason) + """ + category, operation = self._classify_tool(tool_name, arguments) + + # 获取当前计数和限制 + if category == 'persona': + if operation == 'query': + if self.counts.persona_query >= self.limits.persona_query_max: + return (False, f"人设图查询次数已达上限({self.limits.persona_query_max}次)") + else: # update + if self.counts.persona_update >= self.limits.persona_update_max: + return (False, f"人设图修改次数已达上限({self.limits.persona_update_max}次)") + + elif category == 'task': + if operation == 'query': + if self.counts.task_query >= self.limits.task_query_max: + return (False, f"工作记忆链查询次数已达上限({self.limits.task_query_max}次)") + else: # update + if self.counts.task_update >= self.limits.task_update_max: + return (False, f"工作记忆链修改次数已达上限({self.limits.task_update_max}次)") + + elif category == 'memory': + if operation == 'query': + if self.counts.memory_query >= self.limits.memory_query_max: + return (False, f"一般记忆查询次数已达上限({self.limits.memory_query_max}次)") + else: # update + if self.counts.memory_update >= self.limits.memory_update_max: + return (False, f"一般记忆修改次数已达上限({self.limits.memory_update_max}次)") + + return (True, "允许调用") + + def record_call(self, tool_name: str, arguments: dict) -> None: + """记录工具调用""" + category, operation = self._classify_tool(tool_name, arguments) + + if category == 'persona': + if operation == 'query': + self.counts.persona_query += 1 + else: + self.counts.persona_update += 1 + + elif category == 'task': + if operation == 'query': + self.counts.task_query += 1 + else: + self.counts.task_update += 1 + + elif category == 'memory': + if operation == 'query': + self.counts.memory_query += 1 + else: + self.counts.memory_update += 1 + + def get_summary(self) -> str: + """获取调用统计摘要""" + lines = [ + f"人设图: 查询{self.counts.persona_query}/{self.limits.persona_query_max}次, " + f"修改{self.counts.persona_update}/{self.limits.persona_update_max}次", + f"工作记忆链: 查询{self.counts.task_query}/{self.limits.task_query_max}次, " + f"修改{self.counts.task_update}/{self.limits.task_update_max}次", + f"一般记忆: 查询{self.counts.memory_query}/{self.limits.memory_query_max}次, " + f"修改{self.counts.memory_update}/{self.limits.memory_update_max}次" + ] + return "\n".join(lines) + + def reset(self) -> None: + """重置计数(新的一轮对话开始时调用)""" + self.counts = ToolCallCount() diff --git a/core/tools/__init__.py b/core/tools/__init__.py new file mode 100644 index 0000000..0a0ad7d --- /dev/null +++ b/core/tools/__init__.py @@ -0,0 +1,8 @@ +""" +工具定义模块 +""" +from .memory_tools import TOOLS +from .tool_executor import execute_tool +from .tool_limiter import ToolLimiter, ToolLimits, ToolCallCount + +__all__ = ["TOOLS", "execute_tool", "ToolLimiter", "ToolLimits", "ToolCallCount"] diff --git a/core/tools/memory_tools.py b/core/tools/memory_tools.py new file mode 100644 index 0000000..d480120 --- /dev/null +++ b/core/tools/memory_tools.py @@ -0,0 +1,520 @@ +""" +记忆工具定义 - 优化版 +精简描述,避免过拟合,保留AI自主性 +""" + +# 基础记忆工具 +MEMORY_TOOLS = [ + { + "type": "function", + "function": { + "name": "memory_recall", + "description": """检索记忆。支持关键词、时间范围、会话过滤。返回相关实体和关系。 + +【使用示例】 +1. 查询人设图(每轮必须首先执行): + {"query_intent": "AI,人设,角色,性格,语气,说话风格", "depth": 2} + +2. 查询工作记忆链(每轮必须第二步执行): + {"query_intent": "TaskNode,工作记忆,任务链", "depth": 2} + +3. 查询用户偏好: + {"query_intent": "用户,喜欢,偏好", "seed_entities": ["用户"]} + +4. 查询特定主题: + {"query_intent": "Python,编程,项目", "seed_entities": ["Python"]} + +5. 查询最近7天的记忆: + {"query_intent": "任务,工作", "time_range": {"days": 7}} + +【重要】每轮对话必须按顺序执行: +- 步骤1: 查询人设图(最高优先级) +- 步骤2: 查询工作记忆链(维持对话连贯性) +- 步骤3: 根据需要查询其他记忆""", + "parameters": { + "type": "object", + "properties": { + "query_intent": { + "type": "string", + "description": "查询意图,支持逗号分隔多个关键词" + }, + "seed_entities": { + "type": "array", + "items": {"type": "string"}, + "description": "种子实体(可选)" + }, + "depth": { + "type": "integer", + "description": "遍历深度,默认2" + }, + "time_range": { + "type": "object", + "description": "时间范围(可选)", + "properties": { + "days": {"type": "integer", "description": "最近N天"} + } + }, + "session_filter": { + "type": "string", + "description": "会话ID过滤(可选)" + } + }, + "required": ["query_intent"] + } + } + }, + { + "type": "function", + "function": { + "name": "memory_commit", + "description": """写入记忆。将三元组写入图数据库,支持批量写入。 + +【使用示例】 +1. 记录用户偏好: + {"triplets": [ + {"subject": "用户", "relation": "喜欢", "object": "Python编程", "confidence": 0.9}, + {"subject": "用户", "relation": "正在学习", "object": "机器学习"} + ]} + +2. 记录项目信息: + {"triplets": [ + {"subject": "项目A", "relation": "使用技术", "object": "React"}, + {"subject": "项目A", "relation": "状态", "object": "开发中"} + ]} + +3. 记录游戏状态(配合工作记忆链): + {"triplets": [ + {"subject": "成语接龙_当前成语", "relation": "内容", "object": "画龙点睛"}, + {"subject": "成语接龙_当前成语", "relation": "游戏", "object": "成语接龙"} + ]} + +【重要】写入原则: +- 用户明确表达的信息 → 必须写入 +- AI推理得到的信息 → 可以写入,但需标注[推测] +- 避免写入冗余或无意义的信息""", + "parameters": { + "type": "object", + "properties": { + "triplets": { + "type": "array", + "items": { + "type": "object", + "properties": { + "subject": {"type": "string"}, + "relation": {"type": "string"}, + "object": {"type": "string"}, + "confidence": {"type": "number"} + }, + "required": ["subject", "relation", "object"] + }, + "description": "三元组列表" + }, + "entity_types": { + "type": "array", + "items": {"type": "string"}, + "description": "实体类型(可选)" + }, + "temporal_tag": { + "type": "string", + "description": "时间标记(可选)" + } + }, + "required": ["triplets"] + } + } + }, + { + "type": "function", + "function": { + "name": "memory_purge", + "description": """删除记忆。支持条件删除和纠错替代。 + +【使用示例】 +1. 软删除特定关系: + {"criteria": {"subject_contains": "用户", "relation_type": "喜欢"}, "mode": "soft"} + +2. 纠错替代(修正错误信息): + { + "criteria": {"subject_contains": "用户", "relation_type": "年龄"}, + "mode": "supersede", + "new_relation": {"relation": "年龄", "target": "25岁"} + } + +3. 删除特定会话的记忆: + {"criteria": {"session_id": "session_123"}, "mode": "soft"} + +4. 删除旧记忆: + {"criteria": {"time_before": "2024-01-01"}, "mode": "soft"} + +【重要】删除原则: +- 优先使用 supersede 模式修正错误 +- 软删除不会物理删除数据 +- 谨慎使用删除操作""", + "parameters": { + "type": "object", + "properties": { + "criteria": { + "type": "object", + "properties": { + "subject_contains": {"type": "string"}, + "relation_type": {"type": "string"}, + "target_contains": {"type": "string"}, + "time_before": {"type": "string"}, + "session_id": {"type": "string"} + }, + "description": "删除条件" + }, + "mode": { + "type": "string", + "enum": ["soft", "supersede"], + "description": "删除模式:soft=逻辑删除, supersede=纠错替代", + "default": "soft" + }, + "new_relation": { + "type": "object", + "description": "新关系(supersede模式)", + "properties": { + "relation": {"type": "string"}, + "target": {"type": "string"} + } + } + }, + "required": ["criteria"] + } + } + }, + { + "type": "function", + "function": { + "name": "memory_introspect", + "description": "查看记忆状态。返回会话统计、实体热点、关系分布。", + "parameters": { + "type": "object", + "properties": { + "session_id": { + "type": "string", + "description": "会话ID(可选)" + } + }, + "required": [] + } + } + }, + { + "type": "function", + "function": { + "name": "memory_archive", + "description": "归档旧记忆。将N天前的非活跃关系标记为归档状态。", + "parameters": { + "type": "object", + "properties": { + "days": { + "type": "integer", + "description": "归档天数,默认30" + } + }, + "required": [] + } + } + }, + { + "type": "function", + "function": { + "name": "memory_cleanup", + "description": "清理无效数据。物理删除已删除状态超过90天的关系和孤立节点。", + "parameters": { + "type": "object", + "properties": { + "dry_run": { + "type": "boolean", + "description": "仅预览不删除", + "default": True + } + }, + "required": [] + } + } + } +] + +# 人设图管理工具 +PERSONA_TOOLS = [ + { + "type": "function", + "function": { + "name": "persona_update", + "description": """更新人设。修改AI的角色、性格、语气等属性。 + +【使用示例】 +1. 切换为猫娘角色: + {"attributes": [ + {"attribute": "扮演角色", "value": "猫娘"}, + {"attribute": "说话风格", "value": "可爱、卖萌、使用'喵'作为语气词"}, + {"attribute": "性格特点", "value": "活泼、粘人、忠诚"} + ], "mode": "replace"} + +2. 添加新属性(保留现有属性): + {"attributes": [ + {"attribute": "口头禅", "value": "喵呜~"} + ], "mode": "merge"} + +3. 设置专业角色: + {"attributes": [ + {"attribute": "扮演角色", "value": "Python专家"}, + {"attribute": "说话风格", "value": "专业、简洁、代码示例丰富"}, + {"attribute": "性格特点", "value": "严谨、耐心、乐于助人"} + ], "mode": "replace"} + +【重要】人设更新后: +- 立即按照新人设回复 +- 每句话都符合人设的语气、风格、特征 +- 绝不主动跳出角色,除非用户明确要求""", + "parameters": { + "type": "object", + "properties": { + "attributes": { + "type": "array", + "items": { + "type": "object", + "properties": { + "attribute": {"type": "string", "description": "属性名(如:扮演角色、说话风格、性格特点)"}, + "value": {"type": "string", "description": "属性值"} + }, + "required": ["attribute", "value"] + }, + "description": "人设属性列表" + }, + "mode": { + "type": "string", + "enum": ["replace", "merge"], + "description": "更新模式:replace=替换, merge=合并", + "default": "merge" + } + }, + "required": ["attributes"] + } + } + }, + { + "type": "function", + "function": { + "name": "persona_clear", + "description": "清除人设。删除AI的角色设定,恢复默认身份。", + "parameters": { + "type": "object", + "properties": { + "confirm": { + "type": "boolean", + "description": "确认清除", + "default": True + } + }, + "required": [] + } + } + } +] + +# 工作记忆链管理工具 +WORKING_MEMORY_TOOLS = [ + { + "type": "function", + "function": { + "name": "task_create", + "description": """创建任务节点。用于跟踪连续性任务,维持对话连贯性。 + +【使用示例】 +1. 创建成语接龙游戏任务: + { + "task_id": "Task_成语接龙", + "description": "用户发起成语接龙游戏,当前成语:为所欲为", + "info_nodes": ["成语接龙_当前成语"] + } + +2. 创建编程学习任务: + { + "task_id": "Task_Python学习", + "description": "用户正在学习Python,当前主题:装饰器", + "info_nodes": ["Python学习_当前主题"] + } + +3. 创建简单对话任务(每轮必须): + { + "task_id": "Task_当前轮次", + "description": "本轮对话的简要概述" + } + +【重要】工作记忆链机制: +- 每轮对话结束时必须创建任务节点 +- 任务节点通过 NEXT_TASK 边形成时间链 +- 任务节点通过 HAS_STATE 边指向状态节点 +- 任务节点通过 CONTAINS_INFO 边指向信息节点 +- info_nodes 参数用于关联具体信息节点 + +【完整流程示例】 +用户: "咱来玩成语接龙吧,我先开始,为所欲为" + +AI操作步骤: +1. 查询人设图 → 获取当前人设 +2. 查询工作记忆链 → 无进行中任务 +3. 使用 memory_commit 记录游戏状态: + {"triplets": [ + {"subject": "成语接龙_当前成语", "relation": "内容", "object": "为所欲为"}, + {"subject": "成语接龙_当前成语", "relation": "游戏", "object": "成语接龙"} + ]} +4. 使用 task_create 创建任务节点: + {"task_id": "Task_成语接龙", "description": "成语接龙游戏,当前成语:为所欲为", "info_nodes": ["成语接龙_当前成语"]} +5. 回复: "好的喵!我接:为虎作伥喵!" """, + "parameters": { + "type": "object", + "properties": { + "task_id": { + "type": "string", + "description": "任务ID(如:Task_001)" + }, + "description": { + "type": "string", + "description": "任务概述" + }, + "info_nodes": { + "type": "array", + "items": {"type": "string"}, + "description": "关联的信息节点名称(可选)" + } + }, + "required": ["task_id", "description"] + } + } + }, + { + "type": "function", + "function": { + "name": "task_set_state", + "description": """设置任务状态。支持:进行中、已完成、已暂停、已取消。 + +【使用示例】 +1. 标记任务为进行中: + {"task_id": "Task_成语接龙", "state": "进行中"} + +2. 标记任务为已完成: + {"task_id": "Task_成语接龙", "state": "已完成"} + +3. 暂停任务(话题被打断时): + {"task_id": "Task_成语接龙", "state": "已暂停"} + +4. 取消任务: + {"task_id": "Task_成语接龙", "state": "已取消"} + +【重要】状态转换场景: +- 进行中 → 已暂停: 话题被打断时 +- 进行中 → 已完成: 任务完成时 +- 已暂停 → 进行中: 任务恢复时 +- 进行中 → 已取消: 任务被取消时 + +【完整流程示例】 +用户: "关于刚才的成语接龙,我并不知道应该怎么接你的成语,请帮我接一下" + +AI操作步骤: +1. 查询人设图 → 获取当前人设 +2. 查询工作记忆链 → 发现 Task_成语接龙 状态为"已暂停" +3. 使用 task_set_state 恢复任务: + {"task_id": "Task_成语接龙", "state": "进行中"} +4. 查询 Task_成语接龙 的信息节点 → 获取当前成语"为虎作伥" +5. 回复: "好的喵!上一个成语是'为虎作伥',我帮你接:伥鬼害人喵!" """, + "parameters": { + "type": "object", + "properties": { + "task_id": { + "type": "string", + "description": "任务ID" + }, + "state": { + "type": "string", + "enum": ["进行中", "已完成", "已暂停", "已取消"], + "description": "任务状态" + } + }, + "required": ["task_id", "state"] + } + } + }, + { + "type": "function", + "function": { + "name": "task_delete", + "description": "删除任务节点。同时删除关联的信息节点。", + "parameters": { + "type": "object", + "properties": { + "task_id": { + "type": "string", + "description": "任务ID" + }, + "delete_info_nodes": { + "type": "boolean", + "description": "是否删除关联的信息节点", + "default": True + } + }, + "required": ["task_id"] + } + } + }, + { + "type": "function", + "function": { + "name": "task_link_info", + "description": """关联信息节点。将记忆节点关联到任务节点,用于存储任务的具体信息。 + +【使用示例】 +1. 关联游戏状态到任务: + {"task_id": "Task_成语接龙", "info_node_names": ["成语接龙_当前成语", "成语接龙_上一个成语"]} + +2. 关联学习主题到任务: + {"task_id": "Task_Python学习", "info_node_names": ["Python学习_当前主题", "Python学习_学习进度"]} + +3. 关联项目信息到任务: + {"task_id": "Task_项目开发", "info_node_names": ["项目A_技术栈", "项目A_当前阶段"]} + +【重要】使用场景: +- 先使用 memory_commit 创建信息节点 +- 再使用 task_link_info 将信息节点关联到任务节点 +- 信息节点通过 CONTAINS_INFO 边与任务节点连接 + +【完整流程示例】 +用户: "咱来玩成语接龙吧,我先开始,为所欲为" + +AI操作步骤: +1. 查询人设图 → 获取当前人设 +2. 查询工作记忆链 → 无进行中任务 +3. 使用 memory_commit 创建信息节点: + {"triplets": [ + {"subject": "成语接龙_当前成语", "relation": "内容", "object": "为所欲为"}, + {"subject": "成语接龙_当前成语", "relation": "游戏", "object": "成语接龙"} + ]} +4. 使用 task_create 创建任务节点: + {"task_id": "Task_成语接龙", "description": "成语接龙游戏"} +5. 使用 task_link_info 关联信息节点: + {"task_id": "Task_成语接龙", "info_node_names": ["成语接龙_当前成语"]} +6. 回复: "好的喵!我接:为虎作伥喵!" """, + "parameters": { + "type": "object", + "properties": { + "task_id": { + "type": "string", + "description": "任务ID" + }, + "info_node_names": { + "type": "array", + "items": {"type": "string"}, + "description": "信息节点名称列表" + } + }, + "required": ["task_id", "info_node_names"] + } + } + } +] + +# 所有工具 +TOOLS = MEMORY_TOOLS + PERSONA_TOOLS + WORKING_MEMORY_TOOLS diff --git a/core/tools/tool_executor.py b/core/tools/tool_executor.py new file mode 100644 index 0000000..28f70d6 --- /dev/null +++ b/core/tools/tool_executor.py @@ -0,0 +1,307 @@ +""" +工具执行器 +""" +import json +from typing import Any, Dict + + +def execute_tool(graph: Any, tool_name: str, arguments: dict) -> str: + """执行工具调用""" + print(f"\n[工具调用] {tool_name}") + print(f"[参数] {json.dumps(arguments, ensure_ascii=False, indent=2)}") + + try: + # 基础记忆工具 + if tool_name == "memory_recall": + result = graph.recall( + query_intent=arguments.get("query_intent", ""), + seed_entities=arguments.get("seed_entities"), + depth=arguments.get("depth", 2), + time_range=arguments.get("time_range"), + session_filter=arguments.get("session_filter") + ) + return format_recall_result(result) + + elif tool_name == "memory_commit": + result = graph.commit( + triplets=arguments.get("triplets", []), + entity_types=arguments.get("entity_types"), + temporal_tag=arguments.get("temporal_tag") + ) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_purge": + result = graph.purge( + criteria=arguments.get("criteria", {}), + mode=arguments.get("mode", "soft"), + new_relation=arguments.get("new_relation") + ) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_introspect": + result = graph.introspect(session_id=arguments.get("session_id")) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_archive": + result = graph.archive(days=arguments.get("days", 30)) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "memory_cleanup": + result = graph.cleanup(dry_run=arguments.get("dry_run", True)) + return json.dumps(result, ensure_ascii=False, default=str) + + # 人设图管理工具 + elif tool_name == "persona_update": + result = execute_persona_update(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "persona_clear": + result = execute_persona_clear(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + # 工作记忆链管理工具 + elif tool_name == "task_create": + result = execute_task_create(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_set_state": + result = execute_task_set_state(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_delete": + result = execute_task_delete(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + elif tool_name == "task_link_info": + result = execute_task_link_info(graph, arguments) + return json.dumps(result, ensure_ascii=False, default=str) + + return f"未知工具: {tool_name}" + + except Exception as e: + return f"工具执行错误: {str(e)}" + + +def format_recall_result(result: dict) -> str: + """格式化检索结果""" + lines = ["===== 记忆检索结果 ====="] + + if result.get("entities"): + lines.append(f"\n实体 ({len(result['entities'])} 个):") + for e in result["entities"]: + if e and isinstance(e, dict): + lines.append(f" - {e.get('name', 'N/A')} (类型: {e.get('type', 'unknown')}, 提及: {e.get('mention_count', 1)}次)") + + if result.get("relations"): + lines.append(f"\n关系 ({len(result['relations'])} 条):") + for r in result["relations"]: + if r and isinstance(r, dict): + lines.append(f" - {r.get('source', 'N/A')} --[{r.get('type', 'N/A')}]--> {r.get('target', 'N/A')}") + created = r.get("created_at", "N/A") + if created and created != "N/A": + created = created[:19] if "T" in str(created) else str(created) + session_id = r.get('session_id', 'N/A') + session_display = session_id[:20] if session_id and session_id != 'N/A' else 'N/A' + lines.append(f" 时间: {created}, 会话: {session_display}, 轮次: {r.get('turn_id', 0)}, 置信度: {r.get('confidence', 1.0)}") + + if not result.get("entities") and not result.get("relations"): + lines.append("\n(未找到相关记忆)") + + lines.append("=" * 30) + return "\n".join(lines) + + +# 人设图管理工具实现 +def execute_persona_update(graph: Any, arguments: dict) -> dict: + """更新人设""" + attributes = arguments.get("attributes", []) + mode = arguments.get("mode", "merge") + + if mode == "replace": + # 先清除旧人设 + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "扮演角色"}, + mode="soft" + ) + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "说话风格"}, + mode="soft" + ) + graph.purge( + criteria={"subject_contains": "AI", "relation_type": "性格特点"}, + mode="soft" + ) + + # 写入新人设 + triplets = [] + for attr in attributes: + triplets.append({ + "subject": "AI", + "relation": attr["attribute"], + "object": attr["value"], + "confidence": 1.0 + }) + + result = graph.commit(triplets=triplets) + return { + "status": "success", + "mode": mode, + "updated_attributes": len(attributes), + "details": result + } + + +def execute_persona_clear(graph: Any, arguments: dict) -> dict: + """清除人设""" + if not arguments.get("confirm", True): + return {"status": "cancelled", "message": "需要确认才能清除人设"} + + # 删除所有人设相关关系 + result1 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "扮演角色"}, + mode="soft" + ) + result2 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "说话风格"}, + mode="soft" + ) + result3 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "性格特点"}, + mode="soft" + ) + result4 = graph.purge( + criteria={"subject_contains": "AI", "relation_type": "语气特征"}, + mode="soft" + ) + + total_deleted = ( + result1.get("deleted_count", 0) + + result2.get("deleted_count", 0) + + result3.get("deleted_count", 0) + + result4.get("deleted_count", 0) + ) + + return { + "status": "success", + "deleted_count": total_deleted, + "message": "人设已清除,恢复默认身份" + } + + +# 工作记忆链管理工具实现 +def execute_task_create(graph: Any, arguments: dict) -> dict: + """创建任务节点""" + task_id = arguments.get("task_id") + description = arguments.get("description") + info_nodes = arguments.get("info_nodes", []) + + # 创建任务节点 + triplets = [ + {"subject": task_id, "relation": "is_type", "object": "TaskNode"}, + {"subject": task_id, "relation": "has_description", "object": description}, + {"subject": task_id, "relation": "HAS_STATE", "object": "State_进行中"} + ] + + result = graph.commit(triplets=triplets) + + # 关联信息节点 + if info_nodes: + link_triplets = [] + for node_name in info_nodes: + link_triplets.append({ + "subject": task_id, + "relation": "CONTAINS_INFO", + "object": node_name + }) + graph.commit(triplets=link_triplets) + + return { + "status": "success", + "task_id": task_id, + "description": description, + "info_nodes": info_nodes, + "details": result + } + + +def execute_task_set_state(graph: Any, arguments: dict) -> dict: + """设置任务状态""" + task_id = arguments.get("task_id") + state = arguments.get("state") + + # 删除旧状态 + graph.purge( + criteria={"subject_contains": task_id, "relation_type": "HAS_STATE"}, + mode="soft" + ) + + # 设置新状态 + state_node = f"State_{state}" + result = graph.commit( + triplets=[{"subject": task_id, "relation": "HAS_STATE", "object": state_node}] + ) + + return { + "status": "success", + "task_id": task_id, + "new_state": state, + "details": result + } + + +def execute_task_delete(graph: Any, arguments: dict) -> dict: + """删除任务节点""" + task_id = arguments.get("task_id") + delete_info_nodes = arguments.get("delete_info_nodes", True) + + # 查询关联的信息节点 + if delete_info_nodes: + recall_result = graph.recall( + query_intent=f"{task_id},CONTAINS_INFO", + depth=1 + ) + + # 删除信息节点 + for relation in recall_result.get("relations", []): + if relation.get("type") == "CONTAINS_INFO" and relation.get("source") == task_id: + info_node = relation.get("target") + graph.purge( + criteria={"subject_contains": info_node}, + mode="soft" + ) + + # 删除任务节点 + result = graph.purge( + criteria={"subject_contains": task_id}, + mode="soft" + ) + + return { + "status": "success", + "task_id": task_id, + "deleted_info_nodes": delete_info_nodes, + "details": result + } + + +def execute_task_link_info(graph: Any, arguments: dict) -> dict: + """关联信息节点""" + task_id = arguments.get("task_id") + info_node_names = arguments.get("info_node_names", []) + + triplets = [] + for node_name in info_node_names: + triplets.append({ + "subject": task_id, + "relation": "CONTAINS_INFO", + "object": node_name + }) + + result = graph.commit(triplets=triplets) + + return { + "status": "success", + "task_id": task_id, + "linked_nodes": info_node_names, + "details": result + } diff --git a/core/tools/tool_limiter.py b/core/tools/tool_limiter.py new file mode 100644 index 0000000..9f0f614 --- /dev/null +++ b/core/tools/tool_limiter.py @@ -0,0 +1,164 @@ +""" +工具调用限制器 - 限制每轮对话中各类工具的调用次数 +""" +from typing import Dict, List, Optional +from dataclasses import dataclass, field + + +@dataclass +class ToolLimits: + """工具调用限制配置""" + # 人设图限制 + persona_query_max: int = 1 # 每轮最多查询1次人设图 + persona_update_max: int = 1 # 每轮最多修改1次人设图 + + # 工作记忆链限制 + task_query_max: int = 4 # 每轮最多查询4次工作记忆链 + task_update_max: int = 2 # 每轮最多修改2次工作记忆链 + + # 一般记忆限制 + memory_query_max: int = 20 # 每轮最多查询20次一般记忆 + memory_update_max: int = 10 # 每轮最多修改10次一般记忆 + + +@dataclass +class ToolCallCount: + """工具调用计数""" + # 人设图 + persona_query: int = 0 + persona_update: int = 0 + + # 工作记忆链 + task_query: int = 0 + task_update: int = 0 + + # 一般记忆 + memory_query: int = 0 + memory_update: int = 0 + + +class ToolLimiter: + """工具调用限制器""" + + def __init__(self, limits: Optional[ToolLimits] = None): + self.limits = limits or ToolLimits() + self.counts = ToolCallCount() + + def _classify_tool(self, tool_name: str, arguments: dict) -> tuple: + """ + 分类工具调用 + 返回: (category, operation) + category: 'persona', 'task', 'memory' + operation: 'query', 'update' + """ + # 人设图工具 + if tool_name in ('persona_update', 'persona_clear'): + return ('persona', 'update') + + # 工作记忆链工具 + if tool_name in ('task_create', 'task_set_state', 'task_delete', 'task_link_info'): + # task_link_info 是关联操作,算作更新 + return ('task', 'update') + + # 一般记忆工具 + if tool_name == 'memory_recall': + # 判断是查询人设图、工作记忆链还是一般记忆 + query_intent = arguments.get('query_intent', '').lower() + + # 检查是否查询人设图 + if any(kw in query_intent for kw in ['人设', '角色', '性格', '语气', '说话风格', '扮演']): + return ('persona', 'query') + + # 检查是否查询工作记忆链 + if any(kw in query_intent for kw in ['tasknode', '工作记忆', '任务链', '任务', 'task']): + return ('task', 'query') + + # 一般记忆查询 + return ('memory', 'query') + + if tool_name == 'memory_commit': + return ('memory', 'update') + + if tool_name == 'memory_purge': + return ('memory', 'update') + + if tool_name == 'memory_introspect': + return ('memory', 'query') + + if tool_name in ('memory_archive', 'memory_cleanup'): + return ('memory', 'update') + + # 未知工具,归类为一般记忆更新 + return ('memory', 'update') + + def can_call(self, tool_name: str, arguments: dict) -> tuple: + """ + 检查是否允许调用工具 + 返回: (allowed, reason) + """ + category, operation = self._classify_tool(tool_name, arguments) + + # 获取当前计数和限制 + if category == 'persona': + if operation == 'query': + if self.counts.persona_query >= self.limits.persona_query_max: + return (False, f"人设图查询次数已达上限({self.limits.persona_query_max}次)") + else: # update + if self.counts.persona_update >= self.limits.persona_update_max: + return (False, f"人设图修改次数已达上限({self.limits.persona_update_max}次)") + + elif category == 'task': + if operation == 'query': + if self.counts.task_query >= self.limits.task_query_max: + return (False, f"工作记忆链查询次数已达上限({self.limits.task_query_max}次)") + else: # update + if self.counts.task_update >= self.limits.task_update_max: + return (False, f"工作记忆链修改次数已达上限({self.limits.task_update_max}次)") + + elif category == 'memory': + if operation == 'query': + if self.counts.memory_query >= self.limits.memory_query_max: + return (False, f"一般记忆查询次数已达上限({self.limits.memory_query_max}次)") + else: # update + if self.counts.memory_update >= self.limits.memory_update_max: + return (False, f"一般记忆修改次数已达上限({self.limits.memory_update_max}次)") + + return (True, "允许调用") + + def record_call(self, tool_name: str, arguments: dict) -> None: + """记录工具调用""" + category, operation = self._classify_tool(tool_name, arguments) + + if category == 'persona': + if operation == 'query': + self.counts.persona_query += 1 + else: + self.counts.persona_update += 1 + + elif category == 'task': + if operation == 'query': + self.counts.task_query += 1 + else: + self.counts.task_update += 1 + + elif category == 'memory': + if operation == 'query': + self.counts.memory_query += 1 + else: + self.counts.memory_update += 1 + + def get_summary(self) -> str: + """获取调用统计摘要""" + lines = [ + f"人设图: 查询{self.counts.persona_query}/{self.limits.persona_query_max}次, " + f"修改{self.counts.persona_update}/{self.limits.persona_update_max}次", + f"工作记忆链: 查询{self.counts.task_query}/{self.limits.task_query_max}次, " + f"修改{self.counts.task_update}/{self.limits.task_update_max}次", + f"一般记忆: 查询{self.counts.memory_query}/{self.limits.memory_query_max}次, " + f"修改{self.counts.memory_update}/{self.limits.memory_update_max}次" + ] + return "\n".join(lines) + + def reset(self) -> None: + """重置计数(新的一轮对话开始时调用)""" + self.counts = ToolCallCount() diff --git a/docs/README.md b/docs/README.md index 662c6cb..fc25b92 100644 --- a/docs/README.md +++ b/docs/README.md @@ -17,7 +17,7 @@ TrulyMEM (TrueHumanMEM) 是一个让 AI 拥有长期记忆能力的图记忆系 - **长期记忆存储**: 基于 SQLite 内嵌图数据库,开箱即用 - **人设图机制**: 支持角色扮演和性格设定 - **工作记忆链**: 维持对话连贯性的任务跟踪机制 -- **流式消息显示**: 实时显示 AI 响应 +- **TUI 与后端分离**: 多线程 Queue 通信 - **键盘驱动 TUI**: 无需鼠标,全键盘操作 - **跨平台支持**: Windows / Linux / macOS - **独立部署**: 支持打包为可执行文件 diff --git a/docs/架构.md b/docs/架构.md index 71d1e90..c710c4d 100644 --- a/docs/架构.md +++ b/docs/架构.md @@ -5,68 +5,137 @@ - 键盘驱动,零鼠标依赖 - 极简视觉,信息密度优先 - 工具痕迹默认隐藏,需要时可展开 +- TUI 与后端分离,多线程通信 ## 项目结构 ``` TrulyMEM-TrueHumanMEM/ -├── trulymem_entry.py # 打包入口 -├── graph_memory_tui/ # 核心应用包 (38 个 Python 文件) -│ ├── app.py # TUI 主应用 (GraphMemoryApp) -│ ├── main.py # 模块入口 -│ ├── __init__.py -│ ├── core/ # 核心逻辑 -│ │ ├── __init__.py -│ │ ├── imports.py # 动态导入(内嵌DB vs Neo4j) -│ │ ├── embedded_db.py # SQLite 图数据库实现 -│ │ ├── graph_client.py # Neo4j 客户端(可选,未使用) -│ │ ├── optimized_operations.py -│ │ ├── prompts/ # 提示词管理 -│ │ │ ├── __init__.py -│ │ │ ├── prompt_manager.py -│ │ │ └── templates/ -│ │ │ └── system_prompt.md -│ │ └── tools/ # 工具定义与执行 -│ │ ├── __init__.py -│ │ ├── memory_tools.py # 工具定义 -│ │ ├── tool_executor.py # 工具执行器 -│ │ └── tool_limiter.py # 调用限制器 -│ ├── models/ # 数据模型 -│ │ ├── __init__.py -│ │ ├── message.py # Message, ToolCall, ToolResult -│ │ ├── config.py # AppConfig -│ │ └── log_entry.py # LogEntry -│ ├── services/ # 服务层 -│ │ ├── __init__.py -│ │ ├── config_manager.py # 配置持久化 -│ │ ├── config_service.py -│ │ ├── chat_service.py -│ │ └── tool_service.py -│ ├── handlers/ # 事件处理 +├── trulymem_entry.py # 入口:先启动 core → 再启动 ui +├── core/ # 后端/业务逻辑 +│ ├── __init__.py # 导出 BackendServer, BackendClient, EmbeddedGraphDB +│ ├── server.py # BackendServer (多线程队列通信) +│ ├── client.py # BackendClient +│ ├── embedded_db.py # SQLite 图数据库实现 +│ ├── graph_client.py +│ ├── tool_executor.py # 工具执行器 +│ ├── tool_limiter.py # 工具调用限制器 +│ ├── memory_tools.py # 工具定义 +│ ├── prompts/ # 提示词管理 │ │ ├── __init__.py +│ │ ├── prompt_manager.py +│ │ └── templates/ +│ │ └── system_prompt.md +│ └── tools/ # 工具模块 +│ ├── __init__.py +│ ├── memory_tools.py +│ ├── tool_executor.py +│ └── tool_limiter.py +├── ui/ # TUI 显示层 +│ ├── __init__.py # 导出 GraphMemoryApp, AppConfig +│ ├── app.py # GraphMemoryApp (纯显示) +│ ├── widgets/ # TUI 组件 +│ │ ├── left_panel.py +│ │ ├── right_panel.py +│ │ ├── message_history.py +│ │ ├── message_widget.py +│ │ ├── input_box.py +│ │ ├── config_section.py +│ │ ├── operation_log.py +│ │ ├── cypher_query_box.py +│ │ └── status_bar.py +│ ├── handlers/ # 事件处理 │ │ ├── focus_handler.py │ │ ├── key_handler.py │ │ └── message_handler.py -│ ├── widgets/ # TUI 组件 -│ │ ├── __init__.py -│ │ ├── left_panel.py # 左侧主对话区 -│ │ ├── right_panel.py # 右侧边栏 -│ │ ├── message_history.py # 消息历史列表 -│ │ ├── message_widget.py # 单条消息组件 -│ │ ├── input_box.py # 底部输入框 -│ │ ├── config_section.py # 配置区 -│ │ ├── operation_log.py # 图操作日志 -│ │ ├── cypher_query_box.py # 查询框(注:无实际 Cypher 支持) -│ │ └── status_bar.py # 状态栏 -│ └── styles/ # 样式文件 -│ ├── __init__.py +│ ├── models/ # 数据模型 +│ │ ├── message.py +│ │ ├── config.py +│ │ └── log_entry.py +│ ├── services/ # 服务层 +│ │ ├── config_manager.py +│ │ ├── config_service.py +│ │ ├── chat_service.py +│ │ └── tool_service.py +│ └── styles/ # 样式文件 │ ├── app.css │ ├── components.css │ └── messages.css -├── tests/ # 测试(pytest) -├── docs/ # 文档 -├── requirements.txt # 依赖 -└── build_*.{bat,sh} # 打包脚本 +└── tests/ # 测试 (38 tests) +``` + +## 架构图 + +``` +trulymem_entry.py + │ + ├─ BackendServer.start() → 独立线程运行 + │ ├─ 处理 PROCESS_MESSAGE 请求 + │ ├─ 处理 EXECUTE_TOOL 请求 + │ └─ 管理 GraphMemoryClient, EmbeddedGraphDB + │ + └─ GraphMemoryApp(backend_server=server) + │ + └─ BackendClient ← queue.Queue → BackendServer +``` + +## 组件职责 + +### core/ (后端) + +| 组件 | 职责 | +|------|------| +| `server.py` | 多线程队列通信,处理消息和工具调用 | +| `client.py` | TUI 端的通信客户端 | +| `embedded_db.py` | SQLite 图数据库 CRUD | +| `graph_client.py` | OpenAI/DeepSeek API 客户端 | +| `tool_executor.py` | 工具执行逻辑 | +| `tool_limiter.py` | 工具调用频率限制 | + +### ui/ (显示层) + +| 组件 | 职责 | +|------|------| +| `app.py` | Textual 应用主类 | +| `widgets/` | TUI 组件(面板、输入框等) | +| `handlers/` | 事件处理(键盘、焦点) | +| `models/` | 数据模型(消息、配置) | +| `services/` | 配置管理、服务层 | + +## 数据流 + +``` +用户输入 → InputBox → on_input_box_send_message + ↓ +BackendClient.process_message(user_input) + ↓ +queue.Queue → BackendServer (独立线程) + ↓ +GraphMemoryClient.send_message_with_history() + ↓ +OpenAI API / DeepSeek API + ↓ +execute_tool() → EmbeddedGraphDB + ↓ +循环调用 API 直到无 tool_calls + ↓ +queue.Queue → 返回结果 + ↓ +MessageHistory 显示 +``` + +## 启动流程 + +```python +# trulymem_entry.py +def main(): + backend_server = BackendServer(db_path="graph_memory.db") + backend_server.start(api_key=config.api_key) + + app = GraphMemoryApp(backend_server=backend_server) + app.run() + + backend_server.shutdown() ``` ## 布局结构 @@ -118,34 +187,6 @@ TrulyMEM-TrueHumanMEM/ | F5 | 清屏 | | F6 | 退出 | -## 组件职责 - -### 左侧区域 -- **MessageHistory**: 消息历史容器 -- **InputBox**: 底部输入框 - -### 右侧区域 -- **RightPanel**: 侧边栏容器(宽度 70) -- **ConfigSection**: 配置区(API Key、模型选择、Base URL) -- **OperationLog**: 操作日志 -- **CypherQueryBox**: 查询框(注:目前仅作展示,无 Cypher 查询功能) - -## 数据流 - -``` -用户输入 → InputBox → app.on_input_box_send_message - ↓ -GraphMemoryClient.send_message_with_history() - ↓ -OpenAI API / DeepSeek API - ↓ -检查 tool_calls → execute_tool() → EmbeddedGraphDB - ↓ -循环调用 API 直到无 tool_calls - ↓ -最终回复 → MessageHistory + OperationLog -``` - ## 技术栈 | 技术 | 用途 | @@ -154,17 +195,18 @@ OpenAI API / DeepSeek API | Textual 0.47+ | TUI 框架 | | SQLite | 图数据库(默认内嵌) | | OpenAI SDK | API 调用(兼容 DeepSeek) | -| Neo4j | 可选数据库(需 Docker) | +| threading.Queue | 多线程通信 | | PyInstaller | 打包 | ## 数据库模式 -### 默认:SQLite 内嵌 +### SQLite 内嵌(默认) ```python -# core/imports.py -if USE_EMBEDDED_DB: - from .embedded_db import EmbeddedGraphDB as Neo4jGraph +# core/embedded_db.py +class EmbeddedGraphDB: + def __init__(self, db_path="graph_memory.db"): + self.conn = sqlite3.connect(db_path, check_same_thread=False) ``` ### 可选:Neo4j @@ -192,4 +234,4 @@ docker run -d --name neo4j -p 7474:7474 -p 7687:7687 neo4j:latest - `task_create` - 创建任务 - `task_set_state` - 设置状态 - `task_delete` - 删除任务 -- `task_link_info` - 关联信息 +- `task_link_info` - 关联信息 \ No newline at end of file diff --git a/tests/conftest.py b/tests/conftest.py index 6c5ade9..23655f6 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,15 +1,12 @@ -"""测试配置""" - import pytest from datetime import datetime -from graph_memory_tui.models.message import Message, ToolCall, ToolResult -from graph_memory_tui.models.config import AppConfig -from graph_memory_tui.models.log_entry import LogEntry +from ui.models.message import Message, ToolCall, ToolResult +from ui.models.config import AppConfig +from ui.models.log_entry import LogEntry @pytest.fixture def sample_config(): - """示例配置""" return AppConfig( api_key="test-api-key", model="test-model", @@ -19,7 +16,6 @@ def sample_config(): @pytest.fixture def sample_message(): - """示例消息""" return Message( role="user", content="测试消息", @@ -29,7 +25,6 @@ def sample_message(): @pytest.fixture def sample_tool_call(): - """示例工具调用""" return ToolCall( id="test-call-id", name="memory_recall", @@ -39,7 +34,6 @@ def sample_tool_call(): @pytest.fixture def sample_tool_result(): - """示例工具结果""" return ToolResult( tool_call_id="test-call-id", name="memory_recall", @@ -51,11 +45,10 @@ def sample_tool_result(): @pytest.fixture def sample_log_entry(): - """示例日志条目""" return LogEntry( timestamp=datetime.now(), tool_name="memory_recall", arguments={"query_intent": "测试查询"}, result="测试结果", duration=0.5 - ) + ) \ No newline at end of file diff --git a/tests/test_core/test_embedded_db.py b/tests/test_core/test_embedded_db.py index fa72fd1..5f0640e 100644 --- a/tests/test_core/test_embedded_db.py +++ b/tests/test_core/test_embedded_db.py @@ -3,7 +3,7 @@ import pytest import tempfile import os -from graph_memory_tui.core.embedded_db import EmbeddedGraphDB +from core import EmbeddedGraphDB @pytest.fixture diff --git a/tests/test_core/test_imports.py b/tests/test_core/test_imports.py index 38feb26..6ecb770 100644 --- a/tests/test_core/test_imports.py +++ b/tests/test_core/test_imports.py @@ -1,50 +1,32 @@ -"""核心逻辑导入测试""" - import pytest -def test_import_neo4j_graph(): - """测试 Neo4jGraph 类导入""" - from graph_memory_tui.core.imports import Neo4jGraph - assert Neo4jGraph is not None - assert hasattr(Neo4jGraph, 'recall') - assert hasattr(Neo4jGraph, 'commit') - assert hasattr(Neo4jGraph, 'purge') +def test_import_backend_server(): + from core import BackendServer + assert BackendServer is not None -def test_import_graph_memory_client(): - """测试 GraphMemoryClient 类导入""" - from graph_memory_tui.core.imports import GraphMemoryClient +def test_import_backend_client(): + from core import BackendClient + assert BackendClient is not None + + +def test_import_embedded_db(): + from core import EmbeddedGraphDB + assert EmbeddedGraphDB is not None + + +def test_import_graph_client(): + from core.graph_client import GraphMemoryClient assert GraphMemoryClient is not None - assert hasattr(GraphMemoryClient, 'send_message') -def test_import_tools(): - """测试 TOOLS 定义导入""" - from graph_memory_tui.core.imports import TOOLS - assert TOOLS is not None - assert isinstance(TOOLS, list) - assert len(TOOLS) > 0 - assert any(t['function']['name'] == 'memory_recall' for t in TOOLS) +def test_import_tool_limiter(): + from core.tool_limiter import ToolLimiter + assert ToolLimiter is not None -def test_import_execute_tool(): - """测试 execute_tool 函数导入""" - from graph_memory_tui.core.imports import execute_tool +def test_import_tool_executor(): + from core.tool_executor import execute_tool assert execute_tool is not None - assert callable(execute_tool) - - -def test_import_config_vars(): - """测试配置变量导入""" - from graph_memory_tui.core.imports import ( - DEEPSEEK_API_KEY, - DEEPSEEK_BASE_URL, - MODEL_NAME, - NEO4J_URI, - NEO4J_USER, - NEO4J_PASSWORD, - ) - assert DEEPSEEK_BASE_URL is not None - assert MODEL_NAME is not None - assert NEO4J_URI is not None + assert callable(execute_tool) \ No newline at end of file diff --git a/tests/test_core/test_memory_tools.py b/tests/test_core/test_memory_tools.py deleted file mode 100644 index 5c821a1..0000000 --- a/tests/test_core/test_memory_tools.py +++ /dev/null @@ -1,138 +0,0 @@ -"""记忆工具测试""" - -import pytest -from graph_memory_tui.core.tools.memory_tools import ( - MEMORY_TOOLS, - PERSONA_TOOLS, - WORKING_MEMORY_TOOLS, - TOOLS -) - - -def test_memory_tools_exist(): - """测试记忆工具存在""" - assert len(MEMORY_TOOLS) >= 6 - - -def test_persona_tools_exist(): - """测试人设工具存在""" - assert len(PERSONA_TOOLS) >= 2 - - -def test_working_memory_tools_exist(): - """测试工作记忆工具存在""" - assert len(WORKING_MEMORY_TOOLS) >= 4 - - -def test_all_tools_combined(): - """测试工具合并""" - assert len(TOOLS) == len(MEMORY_TOOLS) + len(PERSONA_TOOLS) + len(WORKING_MEMORY_TOOLS) - - -def test_memory_recall_tool(): - """测试 memory_recall 工具定义""" - recall = next((t for t in TOOLS if t['function']['name'] == 'memory_recall'), None) - assert recall is not None - - params = recall['function']['parameters']['properties'] - assert 'query_intent' in params - assert 'seed_entities' in params - assert 'depth' in params - - -def test_memory_commit_tool(): - """测试 memory_commit 工具定义""" - commit = next((t for t in TOOLS if t['function']['name'] == 'memory_commit'), None) - assert commit is not None - - params = commit['function']['parameters']['properties'] - assert 'triplets' in params - - -def test_memory_purge_tool(): - """测试 memory_purge 工具定义""" - purge = next((t for t in TOOLS if t['function']['name'] == 'memory_purge'), None) - assert purge is not None - - params = purge['function']['parameters']['properties'] - assert 'criteria' in params - assert 'mode' in params - - -def test_memory_introspect_tool(): - """测试 memory_introspect 工具定义""" - introspect = next((t for t in TOOLS if t['function']['name'] == 'memory_introspect'), None) - assert introspect is not None - - -def test_persona_update_tool(): - """测试 persona_update 工具定义""" - update = next((t for t in TOOLS if t['function']['name'] == 'persona_update'), None) - assert update is not None - - params = update['function']['parameters']['properties'] - assert 'attributes' in params - - -def test_persona_clear_tool(): - """测试 persona_clear 工具定义""" - clear = next((t for t in TOOLS if t['function']['name'] == 'persona_clear'), None) - assert clear is not None - - -def test_task_create_tool(): - """测试 task_create 工具定义""" - create = next((t for t in TOOLS if t['function']['name'] == 'task_create'), None) - assert create is not None - - params = create['function']['parameters']['properties'] - assert 'task_id' in params - assert 'description' in params - - -def test_task_set_state_tool(): - """测试 task_set_state 工具定义""" - set_state = next((t for t in TOOLS if t['function']['name'] == 'task_set_state'), None) - assert set_state is not None - - params = set_state['function']['parameters']['properties'] - assert 'task_id' in params - assert 'state' in params - - -def test_task_delete_tool(): - """测试 task_delete 工具定义""" - delete = next((t for t in TOOLS if t['function']['name'] == 'task_delete'), None) - assert delete is not None - - -def test_task_link_info_tool(): - """测试 task_link_info 工具定义""" - link = next((t for t in TOOLS if t['function']['name'] == 'task_link_info'), None) - assert link is not None - - params = link['function']['parameters']['properties'] - assert 'task_id' in params - assert 'info_node_names' in params - - -def test_tool_has_required_fields(): - """测试工具都有必需字段""" - for tool in TOOLS: - assert 'type' in tool - assert tool['type'] == 'function' - assert 'function' in tool - assert 'name' in tool['function'] - assert 'description' in tool['function'] - assert 'parameters' in tool['function'] - - -def test_tool_state_enum(): - """测试 task_set_state 的状态枚举""" - set_state = next((t for t in TOOLS if t['function']['name'] == 'task_set_state'), None) - state_enum = set_state['function']['parameters']['properties']['state']['enum'] - - assert '进行中' in state_enum - assert '已完成' in state_enum - assert '已暂停' in state_enum - assert '已取消' in state_enum diff --git a/tests/test_core/test_server.py b/tests/test_core/test_server.py new file mode 100644 index 0000000..4ee1f3b --- /dev/null +++ b/tests/test_core/test_server.py @@ -0,0 +1,57 @@ +import pytest +import os +import tempfile + +os.environ["DEEPSEEK_API_KEY"] = "fake-test-key" + + +class TestCoreImport: + def test_import_backend_server(self): + from core import BackendServer + assert BackendServer is not None + + def test_import_backend_client(self): + from core import BackendClient + assert BackendClient is not None + + def test_import_embedded_db(self): + from core import EmbeddedGraphDB + assert EmbeddedGraphDB is not None + + +class TestBackendServer: + def test_create_server(self): + from core import BackendServer + server = BackendServer(db_path=":memory:", use_embedded_db=True) + assert server is not None + assert server._running is False + + def test_start_server(self): + from core import BackendServer + with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f: + db_path = f.name + + try: + server = BackendServer(db_path=db_path, use_embedded_db=True) + server.start(api_key="", base_url="https://api.test.com") + + assert server._running is True + assert server._graph is not None + + server.shutdown() + finally: + if os.path.exists(db_path): + os.unlink(db_path) + + def test_shutdown(self): + from core import BackendServer + with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f: + db_path = f.name + + server = BackendServer(db_path=db_path, use_embedded_db=True) + server.start(api_key="", base_url="https://api.test.com") + + assert server._running is True + + server.shutdown() + assert server._running is False \ No newline at end of file diff --git a/tests/test_core/test_tool_limiter.py b/tests/test_core/test_tool_limiter.py deleted file mode 100644 index 660a8b0..0000000 --- a/tests/test_core/test_tool_limiter.py +++ /dev/null @@ -1,148 +0,0 @@ -"""工具限制器测试""" - -import pytest -from graph_memory_tui.core.tools.tool_limiter import ( - ToolLimiter, - ToolLimits, - ToolCallCount -) - - -@pytest.fixture -def limiter(): - """创建限制器实例""" - return ToolLimiter() - - -def test_classify_persona_tools(limiter): - """测试人设工具分类""" - category, operation = limiter._classify_tool('persona_update', {}) - assert category == 'persona' - assert operation == 'update' - - category, operation = limiter._classify_tool('persona_clear', {}) - assert category == 'persona' - assert operation == 'update' - - -def test_classify_task_tools(limiter): - """测试任务工具分类""" - category, operation = limiter._classify_tool('task_create', {}) - assert category == 'task' - assert operation == 'update' - - category, operation = limiter._classify_tool('task_set_state', {}) - assert category == 'task' - assert operation == 'update' - - category, operation = limiter._classify_tool('task_delete', {}) - assert category == 'task' - assert operation == 'update' - - category, operation = limiter._classify_tool('task_link_info', {}) - assert category == 'task' - assert operation == 'update' - - -def test_classify_memory_recall(limiter): - """测试 memory_recall 分类""" - category, operation = limiter._classify_tool('memory_recall', {'query_intent': 'Python'}) - assert category == 'memory' - assert operation == 'query' - - -def test_classify_memory_recall_persona_query(limiter): - """测试 memory_recall 查询人设图""" - category, operation = limiter._classify_tool( - 'memory_recall', - {'query_intent': 'AI,人设,角色'} - ) - assert category == 'persona' - assert operation == 'query' - - -def test_classify_memory_recall_task_query(limiter): - """测试 memory_recall 查询工作记忆链""" - category, operation = limiter._classify_tool( - 'memory_recall', - {'query_intent': 'TaskNode,工作记忆'} - ) - assert category == 'task' - assert operation == 'query' - - -def test_can_call_allowed(limiter): - """测试允许调用""" - allowed, reason = limiter.can_call('memory_recall', {'query_intent': 'test'}) - assert allowed is True - - -def test_can_call_limit_reached(limiter): - """测试达到限制""" - for _ in range(20): - limiter.record_call('memory_recall', {'query_intent': 'test'}) - - allowed, reason = limiter.can_call('memory_recall', {'query_intent': 'test'}) - assert allowed is False - assert '上限' in reason - - -def test_record_call(limiter): - """测试记录调用""" - initial_count = limiter.counts.memory_query - - limiter.record_call('memory_recall', {'query_intent': 'test'}) - - assert limiter.counts.memory_query == initial_count + 1 - - -def test_reset(limiter): - """测试重置计数""" - limiter.record_call('memory_recall', {'query_intent': 'test'}) - limiter.record_call('memory_recall', {'query_intent': 'test'}) - - limiter.reset() - - assert limiter.counts.memory_query == 0 - assert limiter.counts.memory_update == 0 - - -def test_get_summary(limiter): - """测试获取统计摘要""" - limiter.record_call('memory_recall', {'query_intent': 'test'}) - - summary = limiter.get_summary() - - assert isinstance(summary, str) - assert '一般记忆' in summary - assert '查询1' in summary - - -def test_custom_limits(): - """测试自定义限制""" - limits = ToolLimits( - memory_query_max=5, - memory_update_max=3 - ) - limiter = ToolLimiter(limits) - - assert limiter.limits.memory_query_max == 5 - assert limiter.limits.memory_update_max == 3 - - -def test_persona_query_limit(limiter): - """测试人设图查询限制""" - for _ in range(1): - limiter.record_call('memory_recall', {'query_intent': '人设'}) - - allowed, _ = limiter.can_call('memory_recall', {'query_intent': '人设'}) - assert allowed is False - - -def test_task_query_limit(limiter): - """测试工作记忆链查询限制""" - for _ in range(4): - limiter.record_call('memory_recall', {'query_intent': 'TaskNode'}) - - allowed, _ = limiter.can_call('memory_recall', {'query_intent': 'TaskNode'}) - assert allowed is False diff --git a/tests/test_handlers/__init__.py b/tests/test_handlers/__init__.py deleted file mode 100644 index 881ac78..0000000 --- a/tests/test_handlers/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for Event Handlers""" diff --git a/tests/test_handlers/test_focus_handler.py b/tests/test_handlers/test_focus_handler.py deleted file mode 100644 index b8a655a..0000000 --- a/tests/test_handlers/test_focus_handler.py +++ /dev/null @@ -1,33 +0,0 @@ -"""焦点处理器测试""" - -import pytest -from graph_memory_tui.handlers.focus_handler import FocusHandler - - -def test_focus_handler_creation(): - """测试焦点处理器创建""" - handler = FocusHandler() - assert handler is not None - assert handler._current_index == 0 - - -def test_focus_ring(): - """测试焦点循环""" - handler = FocusHandler() - assert len(handler.FOCUS_RING) == 5 - assert "input-textarea" in handler.FOCUS_RING - assert "cypher-textarea" in handler.FOCUS_RING - - -def test_get_current_focus_name(): - """测试获取当前焦点名称""" - handler = FocusHandler() - name = handler.get_current_focus_name() - assert name == "Input" - - -def test_focus_names_mapping(): - """测试焦点名称映射""" - handler = FocusHandler() - assert handler.FOCUS_NAMES["input-textarea"] == "Input" - assert handler.FOCUS_NAMES["cypher-textarea"] == "Query" diff --git a/tests/test_models/test_models.py b/tests/test_models/test_models.py index dc7f20a..a6334e4 100644 --- a/tests/test_models/test_models.py +++ b/tests/test_models/test_models.py @@ -2,9 +2,9 @@ import pytest from datetime import datetime -from graph_memory_tui.models.message import Message, ToolCall, ToolResult -from graph_memory_tui.models.config import AppConfig -from graph_memory_tui.models.log_entry import LogEntry +from ui.models.message import Message, ToolCall, ToolResult +from ui.models.config import AppConfig +from ui.models.log_entry import LogEntry def test_message_creation(sample_message): diff --git a/tests/test_services/__init__.py b/tests/test_services/__init__.py deleted file mode 100644 index 48ce0b0..0000000 --- a/tests/test_services/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for Business Services""" diff --git a/tests/test_ui.py b/tests/test_ui.py new file mode 100644 index 0000000..b35b213 --- /dev/null +++ b/tests/test_ui.py @@ -0,0 +1,70 @@ +import pytest +import os +import tempfile + +os.environ["DEEPSEEK_API_KEY"] = "fake-test-key" + + +class TestUIImport: + def test_import_app(self): + from ui import GraphMemoryApp + assert GraphMemoryApp is not None + + def test_import_config(self): + from ui import AppConfig + assert AppConfig is not None + + +class TestUIApp: + def test_create_app(self): + from ui import GraphMemoryApp + app = GraphMemoryApp() + assert app is not None + assert app._backend_server is None + assert app._backend_client is None + + def test_create_app_with_config(self): + from ui import GraphMemoryApp, AppConfig + + config = AppConfig(api_key="test-key", base_url="https://api.test.com") + app = GraphMemoryApp(config=config) + + assert app._config is config + assert app._config.api_key == "test-key" + + def test_create_app_with_backend(self): + from ui import GraphMemoryApp + from core import BackendServer + + with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f: + db_path = f.name + + try: + server = BackendServer(db_path=db_path, use_embedded_db=True) + server.start(api_key="", base_url="https://api.test.com") + + app = GraphMemoryApp(backend_server=server) + + assert app._backend_server is server + assert app._backend_client is not None + + server.shutdown() + finally: + if os.path.exists(db_path): + os.unlink(db_path) + + +class TestAppConfig: + def test_config_from_env(self): + from ui import AppConfig + + config = AppConfig.from_env() + assert config is not None + + def test_config_default_values(self): + from ui import AppConfig + + config = AppConfig() + assert config.api_key == "" + assert config.model == "deepseek-chat" + assert config.base_url == "https://api.deepseek.com" \ No newline at end of file diff --git a/tests/test_widgets/__init__.py b/tests/test_widgets/__init__.py deleted file mode 100644 index 01189a1..0000000 --- a/tests/test_widgets/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for UI Widgets""" diff --git a/trulymem_entry.py b/trulymem_entry.py index dd6ff05..3374049 100644 --- a/trulymem_entry.py +++ b/trulymem_entry.py @@ -1,47 +1,45 @@ #!/usr/bin/env python3 -# SPDX-License-Identifier: GPL-3.0-or-later -# Copyright (C) 2026 jianf - -""" -TrulyMEM 独立入口文件 -用于打包为可执行文件 -""" - import sys import os from pathlib import Path -# 确保工作目录正确 if getattr(sys, 'frozen', False): - # 打包后的可执行文件 application_path = Path(sys.executable).parent else: - # 开发环境 application_path = Path(__file__).parent -# 切换到应用目录 os.chdir(application_path) -# 添加项目路径 if str(application_path) not in sys.path: sys.path.insert(0, str(application_path)) -# 导入并运行应用 -from graph_memory_tui.app import GraphMemoryApp +from core import BackendServer +from ui import GraphMemoryApp, AppConfig + def main(): - """主函数""" + backend_server = BackendServer(db_path="graph_memory.db", use_embedded_db=True) + + try: + config = AppConfig.from_env() + backend_server.start(api_key=config.api_key, base_url=config.base_url) + except Exception as e: + print(f"后端启动失败: {e}") + + app = GraphMemoryApp(backend_server=backend_server) + try: - app = GraphMemoryApp() app.run() except KeyboardInterrupt: print("\n应用已退出") - sys.exit(0) except Exception as e: print(f"应用启动失败: {e}") import traceback traceback.print_exc() - sys.exit(1) + finally: + backend_server.shutdown() + sys.exit(0) + if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/ui/__init__.py b/ui/__init__.py new file mode 100644 index 0000000..5527bfa --- /dev/null +++ b/ui/__init__.py @@ -0,0 +1,4 @@ +from .app import GraphMemoryApp +from .models.config import AppConfig + +__all__ = ["GraphMemoryApp", "AppConfig"] \ No newline at end of file diff --git a/ui/app.py b/ui/app.py new file mode 100644 index 0000000..989d08f --- /dev/null +++ b/ui/app.py @@ -0,0 +1,279 @@ +import asyncio +from pathlib import Path +from textual.app import App, ComposeResult +from textual.binding import Binding +from datetime import datetime + +from core import BackendServer, BackendClient +from .models.message import Message, ToolCall, ToolResult +from .models.config import AppConfig +from .models.log_entry import LogEntry +from .services.config_manager import ConfigManager +from core import BackendClient + + +class GraphMemoryApp(App[None]): + CSS_PATH = [ + Path(__file__).parent / "styles" / "app.css", + Path(__file__).parent / "styles" / "messages.css", + Path(__file__).parent / "styles" / "components.css", + ] + + BINDINGS = [ + Binding("f1", "show_help", "帮助"), + Binding("f2", "toggle_sidebar", "侧边栏"), + Binding("f3", "toggle_tool_details", "工具详情"), + Binding("f4", "focus_query", "查询"), + Binding("f5", "clear_history", "清屏"), + Binding("f6", "quit", "退出"), + ] + + def __init__(self, config: AppConfig | None = None, backend_server: BackendServer | None = None, **kwargs): + super().__init__(**kwargs) + self._config_manager = ConfigManager() + + if config: + self._config = config + elif self._config_manager.exists(): + self._config = self._config_manager.load() + else: + self._config = AppConfig.from_env() + + if backend_server: + self._backend_server = backend_server + self._backend_client = BackendClient(backend_server) + else: + self._backend_server: BackendServer | None = None + self._backend_client: BackendClient | None = None + + def compose(self) -> ComposeResult: + from .widgets.left_panel import LeftPanel + from .widgets.right_panel import RightPanel + from .widgets.status_bar import StatusBar + + yield LeftPanel() + yield RightPanel(self._config, use_embedded_db=True) + yield StatusBar() + + def on_mount(self) -> None: + from .widgets.message_history import MessageHistory + history = self.query_one(MessageHistory) + + try: + self._backend_server = BackendServer( + db_path="graph_memory.db", + use_embedded_db=True + ) + self._backend_server.start( + api_key=self._config.api_key, + base_url=self._config.base_url + ) + + self._backend_client = BackendClient(self._backend_server) + + welcome = Message( + role="assistant", + content="系统初始化成功!\n\n" + f"数据库: 内嵌SQLite (graph_memory.db)\n" + f"API Key: {'已配置' if self._config.api_key else '未配置'}\n\n" + "现在可以开始对话了!", + ) + history.add_message(welcome) + + except Exception as e: + error = Message( + role="assistant", + content=f"初始化失败: {str(e)}\n\n" + "请检查:\n" + "1. API Key 是否配置\n" + "2. 网络连接是否正常\n\n" + "按F2打开侧边栏配置API Key", + ) + history.add_message(error) + + def on_unmount(self) -> None: + if self._backend_server: + self._backend_server.shutdown() + + def action_show_help(self) -> None: + help_text = """ +快捷键: +F1 - 帮助 +F2 - 切换侧边栏 +F3 - 工具详情 +F4 - 查询框 +F5 - 清屏 +F6 - 退出 + +输入消��后按 Enter 发送 + """ + self.notify(help_text, title="帮助", timeout=10) + + def action_toggle_sidebar(self) -> None: + from .widgets.right_panel import RightPanel + sidebar = self.query_one(RightPanel) + sidebar.toggle() + sidebar.update_title() + + def action_toggle_tool_details(self) -> None: + from .widgets.message_history import MessageHistory + history = self.query_one(MessageHistory) + history.toggle_latest_tool_details() + + def action_focus_query(self) -> None: + from .widgets.right_panel import RightPanel + sidebar = self.query_one(RightPanel) + + if not sidebar.has_cypher_query_box(): + self.notify("查询框仅在 Neo4j 模式下可用", title="提示", timeout=3) + return + + if sidebar.is_collapsed(): + sidebar.toggle() + sidebar.update_title() + + query_box = sidebar.get_cypher_query_box() + if query_box: + query_box.focus() + + def action_clear_history(self) -> None: + from .widgets.message_history import MessageHistory + history = self.query_one(MessageHistory) + history.clear_messages() + + def on_input_box_send_message(self, event) -> None: + from .widgets.input_box import InputBox + from .widgets.message_history import MessageHistory + from .widgets.right_panel import RightPanel + + try: + history = self.query_one(MessageHistory) + user_message = Message(role="user", content=event.content) + history.add_message(user_message) + + if not self._config.api_key: + response_msg = Message( + role="assistant", + content="请先配置API Key。\n\n按F2打开侧边栏,输入API Key后按Enter保存。", + ) + history.add_message(response_msg) + return + + processing_msg = Message(role="assistant", content="正在处理...") + history.add_message(processing_msg) + + asyncio.create_task(self._process_message_async(event.content)) + + except Exception as e: + error_msg = Message(role="assistant", content=f"错误: {str(e)}") + history.add_message(error_msg) + + async def _process_message_async(self, user_input: str) -> None: + from .widgets.message_history import MessageHistory + from .widgets.right_panel import RightPanel + + history = self.query_one(MessageHistory) + log = self.query_one(RightPanel).get_operation_log() + + try: + if not self._backend_client: + raise Exception("后端未初始化") + + result = await asyncio.get_event_loop().run_in_executor( + None, + lambda: self._backend_client.process_message(user_input) + ) + + content = result.get("content", "(无回复)") + tool_calls_data = result.get("tool_calls", []) + rejected_tools = result.get("rejected_tools", []) + + tool_calls = [] + tool_results = [] + + for tc in tool_calls_data: + tc_obj = ToolCall( + id=tc.get("id", ""), + name=tc.get("name", ""), + arguments=tc.get("arguments", {}) + ) + tool_calls.append(tc_obj) + + tr = ToolResult( + tool_call_id=tc_obj.id, + name=tc_obj.name, + arguments=tc_obj.arguments, + result=tc.get("result", ""), + success=not tc.get("result", "").startswith("工具执行���误") + ) + tool_results.append(tr) + + log_entry = LogEntry( + tool_name=tc_obj.name, + arguments=tc_obj.arguments, + result=tc.get("result", ""), + ) + log.add_log(log_entry) + + assistant_message = Message( + role="assistant", + content=content, + tool_calls=tool_calls if tool_calls else None, + tool_results=tool_results if tool_results else None + ) + + history.add_message(assistant_message) + self.refresh() + + except Exception as e: + error_msg = str(e) + + if "Connection error" in error_msg or "connection" in error_msg.lower(): + help_text = """ +网络连接错误!可能的原因: +1. API Key 未配置或无效 +2. 网络无法访问 API 服务器 +3. API 服务器暂时不可用 + +解决方法: +按 F2 展开侧边栏,检查并配置 API Key +检查网络连接 +""" + elif "API Key" in error_msg: + help_text = """ +API Key 未配置! + +请按以下步骤配置: +1. 按 F2 展开右侧边栏 +2. 点击"配置"展开配置区 +3. 在 API Key 输入框输入你的密钥 +4. 按 Enter 键保存配置 + +获取 API Key: https://platform.deepseek.com/ +""" + else: + help_text = f"\n详细错误: {error_msg}" + + error_message = Message(role="assistant", content=f"错误: {error_msg}\n{help_text}") + history.add_message(error_message) + + def on_config_section_config_changed(self, event) -> None: + from .widgets.right_panel import RightPanel + + self._config = event.config + self._config_manager.save(self._config) + + try: + right_panel = self.query_one(RightPanel) + right_panel._config = self._config + except Exception: + pass + + if self._backend_client: + self._backend_client.update_config( + api_key=self._config.api_key, + base_url=self._config.base_url + ) + self.notify("配置已保存并应用", title="配置") + else: + self.notify("配置已保存,但后端未初始化", title="警告") \ No newline at end of file diff --git a/ui/handlers/__init__.py b/ui/handlers/__init__.py new file mode 100644 index 0000000..14b3896 --- /dev/null +++ b/ui/handlers/__init__.py @@ -0,0 +1 @@ +"""Event Handlers for Graph Memory TUI""" diff --git a/ui/handlers/focus_handler.py b/ui/handlers/focus_handler.py new file mode 100644 index 0000000..fa3c99d --- /dev/null +++ b/ui/handlers/focus_handler.py @@ -0,0 +1,64 @@ +"""焦点管理器""" + +from textual.app import App + + +class FocusHandler: + """焦点管理器""" + + # 焦点循环顺序 + FOCUS_RING = [ + "input-textarea", # 左侧输入框 + "api-key-input", # 右侧配置区 API Key + "model-input", # 右侧配置区 Model + "base-url-input", # 右侧配置区 Base URL + "cypher-textarea", # 右侧 Cypher 查询框 + ] + + # 焦点名称映射 + FOCUS_NAMES = { + "input-textarea": "Input", + "api-key-input": "Config-API", + "model-input": "Config-Model", + "base-url-input": "Config-URL", + "cypher-textarea": "Query", + } + + def __init__(self): + self._current_index = 0 + + def next_focus(self, app: App) -> None: + """切换到下一个焦点""" + self._current_index = (self._current_index + 1) % len(self.FOCUS_RING) + widget_id = self.FOCUS_RING[self._current_index] + self._focus_widget(app, widget_id) + + def prev_focus(self, app: App) -> None: + """切换到上一个焦点""" + self._current_index = (self._current_index - 1) % len(self.FOCUS_RING) + widget_id = self.FOCUS_RING[self._current_index] + self._focus_widget(app, widget_id) + + def focus_input(self, app: App) -> None: + """聚焦到输入框""" + self._current_index = 0 + self._focus_widget(app, self.FOCUS_RING[0]) + + def focus_query(self, app: App) -> None: + """聚焦到查询框""" + self._current_index = len(self.FOCUS_RING) - 1 + self._focus_widget(app, self.FOCUS_RING[-1]) + + def get_current_focus_name(self) -> str: + """获取当前焦点名称""" + widget_id = self.FOCUS_RING[self._current_index] + return self.FOCUS_NAMES.get(widget_id, "Unknown") + + def _focus_widget(self, app: App, widget_id: str) -> None: + """聚焦到指定组件""" + try: + widget = app.query_one(f"#{widget_id}") + widget.focus() + except Exception: + # 如果找不到组件,回退到输入框 + self.focus_input(app) diff --git a/ui/handlers/key_handler.py b/ui/handlers/key_handler.py new file mode 100644 index 0000000..0ddc06a --- /dev/null +++ b/ui/handlers/key_handler.py @@ -0,0 +1,68 @@ +"""快捷键处理器""" + +from textual.app import App +from textual.message import Message +from .focus_handler import FocusHandler + + +class KeyHandler: + """快捷键处理器""" + + class ShowHelp(Message): + """显示帮助事件""" + pass + + class ToggleSidebar(Message): + """切换侧边栏事件""" + pass + + class ToggleToolDetails(Message): + """切换工具详情事件""" + pass + + class FocusQuery(Message): + """聚焦查询框事件""" + pass + + class ClearHistory(Message): + """清屏事件""" + pass + + class QuitApp(Message): + """退出应用事件""" + pass + + def __init__(self, focus_handler: FocusHandler): + self._focus_handler = focus_handler + + def handle_f1(self, app: App) -> None: + """处理 F1 键 - 显示帮助""" + app.post_message(self.ShowHelp()) + + def handle_f2(self, app: App) -> None: + """处理 F2 键 - 切换侧边栏""" + app.post_message(self.ToggleSidebar()) + + def handle_f3(self, app: App) -> None: + """处理 F3 键 - 切换工具详情""" + app.post_message(self.ToggleToolDetails()) + + def handle_f4(self, app: App) -> None: + """处理 F4 键 - 聚焦查询框""" + app.post_message(self.FocusQuery()) + + def handle_f5(self, app: App) -> None: + """处理 F5 键 - 清屏""" + app.post_message(self.ClearHistory()) + + def handle_f6(self, app: App) -> None: + """处理 F6 键 - 退出""" + app.post_message(self.QuitApp()) + + def handle_tab(self, app: App) -> None: + """处理 Tab 键 - 焦点循环""" + self._focus_handler.next_focus(app) + + def handle_shift_tab(self, app: App) -> None: + """处理 Shift+Tab 键 - 反向焦点循环""" + self._focus_handler.prev_focus(app) diff --git a/ui/handlers/message_handler.py b/ui/handlers/message_handler.py new file mode 100644 index 0000000..eeec026 --- /dev/null +++ b/ui/handlers/message_handler.py @@ -0,0 +1,98 @@ +"""消息处理器""" + +from datetime import datetime +from typing import TYPE_CHECKING +from ..models.message import Message, ToolCall, ToolResult + +if TYPE_CHECKING: + from ..services.chat_service import ChatService + from ..widgets.message_history import MessageHistory + from ..widgets.operation_log import OperationLog + + +class MessageHandler: + """消息处理器""" + + def __init__( + self, + chat_service: "ChatService", + message_history: "MessageHistory", + operation_log: "OperationLog" + ): + self._chat_service = chat_service + self._message_history = message_history + self._operation_log = operation_log + + async def handle_user_message(self, content: str) -> None: + """处理用户消息""" + # 创建用户消息 + user_message = Message( + role="user", + content=content, + timestamp=datetime.now() + ) + + # 添加到历史 + self._message_history.add_message(user_message) + + # 发送到聊天服务 + await self._process_response(content) + + async def _process_response(self, user_input: str) -> None: + """处理响应""" + streaming_message = None + + async for event in self._chat_service.send_message(user_input): + if event["type"] == "user_message": + # 用户消息已处理 + pass + + elif event["type"] == "content_delta": + # 流式内容更新 + if streaming_message is None: + # 创建流式消息 + streaming_message = Message( + role="assistant", + content="", + timestamp=datetime.now() + ) + self._message_history.add_message(streaming_message) + + # 更新消息内容 + self._message_history.update_latest_message(event["content"]) + + elif event["type"] == "assistant_message": + # 模型消息完成 + if streaming_message: + # 更新最终消息(包含工具调用信息) + streaming_message.content = event["content"] + streaming_message.tool_calls = event.get("tool_calls") + streaming_message.tool_results = event.get("tool_results") + else: + # 如果没有流式消息,直接添加 + message = Message( + role="assistant", + content=event["content"], + timestamp=datetime.now(), + tool_calls=event.get("tool_calls"), + tool_results=event.get("tool_results") + ) + self._message_history.add_message(message) + + elif event["type"] == "tool_call": + # 工具调用开始 + pass + + elif event["type"] == "tool_result": + # 工具执行结果 + log_entry = event["log_entry"] + self._operation_log.add_log(log_entry) + + elif event["type"] == "error": + # 错误处理 + error_message = Message( + role="assistant", + content=f"错误: {event['error']}", + timestamp=datetime.now() + ) + self._message_history.add_message(error_message) diff --git a/ui/models/__init__.py b/ui/models/__init__.py new file mode 100644 index 0000000..f08233e --- /dev/null +++ b/ui/models/__init__.py @@ -0,0 +1 @@ +"""Data Models for Graph Memory TUI""" diff --git a/ui/models/config.py b/ui/models/config.py new file mode 100644 index 0000000..a2ce43b --- /dev/null +++ b/ui/models/config.py @@ -0,0 +1,46 @@ +"""配置数据模型""" + +import json +import os +from dataclasses import dataclass, asdict +from pathlib import Path +from typing import Optional + + +@dataclass +class AppConfig: + """应用配置""" + api_key: str = "" + model: str = "deepseek-chat" + base_url: str = "https://api.deepseek.com" + + @classmethod + def from_env(cls) -> "AppConfig": + """从环境变量加载配置""" + return cls( + api_key=os.getenv("DEEPSEEK_API_KEY", ""), + model=os.getenv("MODEL_NAME", "deepseek-chat"), + base_url=os.getenv("DEEPSEEK_BASE_URL", "https://api.deepseek.com"), + ) + + @classmethod + def from_file(cls, path: Path) -> "AppConfig": + """从文件加载配置""" + if not path.exists(): + return cls() + + with open(path, "r", encoding="utf-8") as f: + data = json.load(f) + + return cls( + api_key=data.get("api_key", ""), + model=data.get("model", "deepseek-chat"), + base_url=data.get("base_url", "https://api.deepseek.com"), + ) + + def save(self, path: Path) -> None: + """保存配置到文件""" + path.parent.mkdir(parents=True, exist_ok=True) + + with open(path, "w", encoding="utf-8") as f: + json.dump(asdict(self), f, indent=2, ensure_ascii=False) diff --git a/ui/models/log_entry.py b/ui/models/log_entry.py new file mode 100644 index 0000000..51c5eb3 --- /dev/null +++ b/ui/models/log_entry.py @@ -0,0 +1,30 @@ +"""日志条目数据模型""" + +from dataclasses import dataclass +from datetime import datetime +from typing import Any, Dict + + +@dataclass +class LogEntry: + """日志条目""" + timestamp: datetime + tool_name: str + arguments: Dict[str, Any] + result: str + duration: float + + @property + def args_summary(self) -> str: + """参数摘要(截断到50字符)""" + args_str = str(self.arguments) + if len(args_str) > 50: + return args_str[:50] + "..." + return args_str + + @property + def result_summary(self) -> str: + """结果摘要(截断到100字符)""" + if len(self.result) > 100: + return self.result[:100] + "..." + return self.result diff --git a/ui/models/message.py b/ui/models/message.py new file mode 100644 index 0000000..98bbdd4 --- /dev/null +++ b/ui/models/message.py @@ -0,0 +1,33 @@ +"""消息数据模型""" + +from dataclasses import dataclass, field +from datetime import datetime +from typing import Dict, List, Literal, Optional, Any + + +@dataclass +class ToolCall: + """工具调用""" + id: str + name: str + arguments: Dict[str, Any] + + +@dataclass +class ToolResult: + """工具执行结果""" + tool_call_id: str + name: str + arguments: Dict[str, Any] + result: str + success: bool + + +@dataclass +class Message: + """消息""" + role: Literal["user", "assistant", "system"] + content: str + timestamp: datetime = field(default_factory=datetime.now) + tool_calls: Optional[List[ToolCall]] = None + tool_results: Optional[List[ToolResult]] = None diff --git a/ui/services/__init__.py b/ui/services/__init__.py new file mode 100644 index 0000000..837ef17 --- /dev/null +++ b/ui/services/__init__.py @@ -0,0 +1 @@ +"""Business Services for Graph Memory TUI""" diff --git a/ui/services/chat_service.py b/ui/services/chat_service.py new file mode 100644 index 0000000..fc96f43 --- /dev/null +++ b/ui/services/chat_service.py @@ -0,0 +1,226 @@ +"""聊天服务""" + +import asyncio +import json +from datetime import datetime +from typing import AsyncIterator, TYPE_CHECKING, List, Dict, Any +from ..core.imports import GraphMemoryClient +from ..models.message import ToolCall, ToolResult +from .tool_service import ToolService + +if TYPE_CHECKING: + from ..core.imports import Neo4jGraph + + +class ChatService: + """聊天业务服务""" + + def __init__( + self, + graph: "Neo4jGraph", + client: GraphMemoryClient, + tool_service: ToolService + ): + self._graph = graph + self._client = client + self._tool_service = tool_service + self._messages: List[Dict[str, Any]] = [] + + async def send_message(self, user_input: str) -> AsyncIterator[dict]: + """发送消息并流式返回事件""" + # 1. 发送用户消息事件 + yield { + "type": "user_message", + "content": user_input + } + + try: + # 2. 第一次API调用 + accumulated_content = "" + tool_calls_data = [] + + # 流式处理响应 + async for chunk in self._call_api_stream_async(user_input): + if chunk.get("content_delta"): + accumulated_content += chunk["content_delta"] + yield { + "type": "content_delta", + "content": accumulated_content + } + + if chunk.get("tool_calls"): + tool_calls_data = chunk["tool_calls"] + + # 3. 如果有工具调用,执行并继续调用API + tool_calls = None + tool_results = None + + if tool_calls_data: + tool_calls = [] + tool_results = [] + + # 执行所有工具 + for tool_call_data in tool_calls_data: + tool_call = ToolCall( + id=tool_call_data["id"], + name=tool_call_data["function"]["name"], + arguments=tool_call_data["function"]["arguments"] + ) + tool_calls.append(tool_call) + + yield { + "type": "tool_call", + "tool_call": tool_call + } + + result = await self._tool_service.execute(tool_call) + tool_results.append(result) + + log_entry = ToolService._create_log_entry(tool_call, result) + yield { + "type": "tool_result", + "tool_result": result, + "log_entry": log_entry + } + + # 构建工具结果消息 + tool_messages = [] + for tc, tr in zip(tool_calls, tool_results): + tool_messages.append({ + "role": "tool", + "tool_call_id": tc.id, + "content": tr.content + }) + + # 构建assistant消息(包含tool_calls) + assistant_message = { + "role": "assistant", + "content": accumulated_content, + "tool_calls": [ + { + "id": tc.id, + "type": "function", + "function": { + "name": tc.name, + "arguments": tc.arguments + } + } for tc in tool_calls + ] + } + + # 第二次API调用,传入工具结果 + final_content = "" + async for chunk in self._call_api_stream_with_tools( + user_input, + assistant_message, + tool_messages + ): + if chunk.get("content_delta"): + final_content += chunk["content_delta"] + yield { + "type": "content_delta", + "content": final_content + } + + accumulated_content = final_content + + # 4. 返回最终回复 + yield { + "type": "assistant_message", + "content": accumulated_content, + "tool_calls": tool_calls, + "tool_results": tool_results + } + + except Exception as e: + yield { + "type": "error", + "error": str(e) + } + + async def _call_api_stream_async(self, message: str) -> AsyncIterator[dict]: + """异步流式调用 API""" + loop = asyncio.get_event_loop() + + def process_stream(): + stream = self._client.send_message_stream(message) + tool_calls_accumulated = [] + + for chunk in stream: + delta = chunk.choices[0].delta + + if delta.content: + yield {"content_delta": delta.content} + + if delta.tool_calls: + for tc in delta.tool_calls: + if tc.index >= len(tool_calls_accumulated): + tool_calls_accumulated.append({ + "id": tc.id, + "type": "function", + "function": { + "name": "", + "arguments": "" + } + }) + + if tc.function: + if tc.function.name: + tool_calls_accumulated[tc.index]["function"]["name"] = tc.function.name + if tc.function.arguments: + tool_calls_accumulated[tc.index]["function"]["arguments"] += tc.function.arguments + + if tool_calls_accumulated: + yield {"tool_calls": tool_calls_accumulated} + + for result in await loop.run_in_executor(None, lambda: list(process_stream())): + yield result + + async def _call_api_stream_with_tools( + self, + user_input: str, + assistant_message: dict, + tool_messages: list + ) -> AsyncIterator[dict]: + """带工具结果的流式调用""" + loop = asyncio.get_event_loop() + + def process_stream(): + # 构建完整的消息列表 + messages = [ + {"role": "system", "content": self._client.system_prompt}, + {"role": "user", "content": user_input}, + assistant_message + ] + messages.extend(tool_messages) + + # 调用API + response = self._client.client.chat.completions.create( + model="deepseek-chat", + messages=messages, + tools=self._client.tools, + tool_choice="auto", + stream=True + ) + + for chunk in response: + delta = chunk.choices[0].delta + if delta.content: + yield {"content_delta": delta.content} + + # 处理可能的工具调用 + if delta.tool_calls: + # 如果还有工具调用,说明AI想继续调用工具 + # 但我们限制只调用一次,所以忽略 + pass + + for result in await loop.run_in_executor(None, lambda: list(process_stream())): + yield result + + def clear_history(self) -> None: + """清空消息历史""" + self._messages.clear() + + def get_history(self) -> list[dict]: + """获取消息历史""" + return self._messages.copy() diff --git a/ui/services/config_manager.py b/ui/services/config_manager.py new file mode 100644 index 0000000..e2501b5 --- /dev/null +++ b/ui/services/config_manager.py @@ -0,0 +1,46 @@ +""" +配置管理 - 支持持久化 +""" + +import json +from pathlib import Path +from ..models.config import AppConfig + + +class ConfigManager: + """配置管理器 - 支持持久化""" + + def __init__(self, config_file: str = "config.json"): + self.config_file = Path(config_file) + + def save(self, config: AppConfig) -> None: + """保存配置到文件""" + data = { + "api_key": config.api_key, + "model": config.model, + "base_url": config.base_url + } + + with open(self.config_file, 'w', encoding='utf-8') as f: + json.dump(data, f, indent=2) + + def load(self) -> AppConfig: + """从文件加载配置""" + if not self.config_file.exists(): + return AppConfig() + + try: + with open(self.config_file, 'r', encoding='utf-8') as f: + data = json.load(f) + + return AppConfig( + api_key=data.get("api_key", ""), + model=data.get("model", "deepseek-chat"), + base_url=data.get("base_url", "https://api.deepseek.com") + ) + except Exception: + return AppConfig() + + def exists(self) -> bool: + """检查配置文件是否存在""" + return self.config_file.exists() diff --git a/ui/services/config_service.py b/ui/services/config_service.py new file mode 100644 index 0000000..dc2a9c8 --- /dev/null +++ b/ui/services/config_service.py @@ -0,0 +1,51 @@ +"""配置服务""" + +from pathlib import Path +from typing import TYPE_CHECKING +from ..models.config import AppConfig + +if TYPE_CHECKING: + from ..core.imports import GraphMemoryClient + + +class ConfigService: + """配置服务""" + + DEFAULT_CONFIG_FILE = Path.home() / ".graph_memory_tui" / "config.json" + + def __init__(self, config_file: Path | None = None): + self._config_file = config_file or self.DEFAULT_CONFIG_FILE + self._config = self._load_config() + + def _load_config(self) -> AppConfig: + """加载配置""" + # 优先从文件加载 + if self._config_file.exists(): + return AppConfig.from_file(self._config_file) + + # 否则从环境变量加载 + return AppConfig.from_env() + + def get_config(self) -> AppConfig: + """获取当前配置""" + return self._config + + def set_config(self, config: AppConfig) -> None: + """设置配置""" + self._config = config + self._save_config() + + def _save_config(self) -> None: + """保存配置""" + self._config.save(self._config_file) + + def apply_to_client(self, client: "GraphMemoryClient") -> None: + """应用配置到 API 客户端""" + # 更新客户端配置 + client.api_key = self._config.api_key + client.base_url = self._config.base_url + client.model = self._config.model + + def get_config_file(self) -> Path: + """获取配置文件路径""" + return self._config_file diff --git a/ui/services/tool_service.py b/ui/services/tool_service.py new file mode 100644 index 0000000..fb37659 --- /dev/null +++ b/ui/services/tool_service.py @@ -0,0 +1,88 @@ +"""工具服务""" + +import asyncio +import time +from datetime import datetime +from typing import Callable, TYPE_CHECKING +from ..core.imports import execute_tool +from ..models.log_entry import LogEntry +from ..models.message import ToolCall, ToolResult + +if TYPE_CHECKING: + from ..core.imports import Neo4jGraph + + +class ToolService: + """工具执行服务""" + + def __init__( + self, + graph: "Neo4jGraph", + log_callback: Callable[[LogEntry], None] | None = None + ): + self._graph = graph + self._log_callback = log_callback + + async def execute(self, tool_call: ToolCall) -> ToolResult: + """异步执行工具""" + start_time = time.time() + + try: + # 在线程池中执行同步工具 + loop = asyncio.get_event_loop() + result = await loop.run_in_executor( + None, + lambda: execute_tool(self._graph, tool_call.name, tool_call.arguments) + ) + + duration = time.time() - start_time + + # 创建日志条目 + log_entry = LogEntry( + timestamp=datetime.now(), + tool_name=tool_call.name, + arguments=tool_call.arguments, + result=result, + duration=duration + ) + + # 回调日志 + if self._log_callback: + self._log_callback(log_entry) + + # 返回结果 + return ToolResult( + tool_call_id=tool_call.id, + name=tool_call.name, + arguments=tool_call.arguments, + result=result, + success=not result.startswith("工具执行错误") + ) + + except Exception as e: + duration = time.time() - start_time + error_msg = f"工具执行异常: {str(e)}" + + # 创建错误日志 + log_entry = LogEntry( + timestamp=datetime.now(), + tool_name=tool_call.name, + arguments=tool_call.arguments, + result=error_msg, + duration=duration + ) + + if self._log_callback: + self._log_callback(log_entry) + + return ToolResult( + tool_call_id=tool_call.id, + name=tool_call.name, + arguments=tool_call.arguments, + result=error_msg, + success=False + ) + + def set_log_callback(self, callback: Callable[[LogEntry], None]) -> None: + """设置日志回调""" + self._log_callback = callback diff --git a/ui/styles/__init__.py b/ui/styles/__init__.py new file mode 100644 index 0000000..190eec3 --- /dev/null +++ b/ui/styles/__init__.py @@ -0,0 +1 @@ +"""Styles for Graph Memory TUI""" diff --git a/ui/styles/app.css b/ui/styles/app.css new file mode 100644 index 0000000..6c04465 --- /dev/null +++ b/ui/styles/app.css @@ -0,0 +1,41 @@ +/* Global Styles for Graph Memory TUI */ + +GraphMemoryApp { + background: $surface; + color: $text; +} + +/* 全局Input样式 - 确保可见 */ +Input { + background: $surface-lighten-1; + color: $text; + border: solid $primary; +} + +Input:focus { + border: double $accent; +} + +LeftPanel { + width: 1fr; + dock: left; +} + +RightPanel { + width: 70; + dock: right; + background: $panel; + overflow-y: auto; +} + +RightPanel ScrollableContainer { + height: 1fr; + overflow-y: auto; +} + +StatusBar { + dock: bottom; + height: 1; + background: $primary; + color: $text-primary; +} diff --git a/ui/styles/components.css b/ui/styles/components.css new file mode 100644 index 0000000..cf2bcd0 --- /dev/null +++ b/ui/styles/components.css @@ -0,0 +1,129 @@ +/* Component Styles for Graph Memory TUI */ + +/* Input Box - 最重要 */ +InputBox { + background: $surface; + padding: 1 2; + height: auto; + border: solid $primary; +} + +InputBox Input { + width: 100%; + background: $surface-lighten-1; + color: $text; + border: none; +} + +/* Config Section */ +ConfigSection { + background: $surface; + padding: 1; + margin: 0 0 1 0; + height: auto; +} + +ConfigSection .config-title { + color: $primary; + text-style: bold; + margin: 0 0 1 0; +} + +ConfigSection .config-label { + color: $text; + margin: 0; + padding: 1 0 0 0; +} + +ConfigSection .config-hint { + color: $text-muted; + text-style: italic; + margin: 1 0 0 0; +} + +ConfigSection Input { + width: 1fr; + height: 3; + margin: 0 0 1 0; + padding: 0 1; + background: $surface-lighten-1; + border: solid $primary; + color: $text; +} + +/* Other Components */ +OperationLog { + background: $surface-darken-1; + height: 1fr; + margin: 1; + overflow-y: auto; + padding: 1; +} + +OperationLog .log-entry { + color: $text; + margin: 0 0 1 0; + height: auto; +} + +OperationLog .log-empty { + color: $text-muted; + text-style: italic; +} + +CypherQueryBox { + border: solid green; + margin: 1; + height: auto; +} + +MessageHistory { + height: 1fr; + margin: 1; + overflow-y: auto; +} + +/* Message Widget */ +MessageWidget { + margin: 1 0; + height: auto; +} + +MessageWidget .message-header { + color: $text-muted; + text-style: bold; + margin: 0 0 0 0; +} + +MessageWidget .message-content { + color: $text; + margin: 0 0 0 2; + height: auto; +} + +MessageWidget .tool-indicator { + color: $warning; + text-style: bold; + margin: 1 0 0 2; +} + +MessageWidget .tool-details { + background: $surface-darken-1; + margin: 1 0 0 2; + padding: 1; +} + +MessageWidget .tool-name { + color: $accent; + text-style: bold; +} + +MessageWidget .tool-args { + color: $text-muted; + margin: 0 0 0 2; +} + +MessageWidget .tool-result { + color: $success; + margin: 0 0 0 2; +} diff --git a/ui/styles/messages.css b/ui/styles/messages.css new file mode 100644 index 0000000..0d7b1d4 --- /dev/null +++ b/ui/styles/messages.css @@ -0,0 +1,24 @@ +/* Message Styles for Graph Memory TUI */ + +UserMessage { + border: solid orange; + margin: 1 0; + padding: 1; +} + +ModelMessage { + border: solid blue; + margin: 1 0; + padding: 1; +} + +ToolCallIndicator { + color: yellow; + text-style: bold; +} + +ToolCallDetails { + background: $surface-darken-1; + margin: 1 0 0 2; + padding: 1; +} diff --git a/ui/widgets/__init__.py b/ui/widgets/__init__.py new file mode 100644 index 0000000..fd32402 --- /dev/null +++ b/ui/widgets/__init__.py @@ -0,0 +1 @@ +"""UI Widgets for Graph Memory TUI""" diff --git a/ui/widgets/config_section.py b/ui/widgets/config_section.py new file mode 100644 index 0000000..0a3b669 --- /dev/null +++ b/ui/widgets/config_section.py @@ -0,0 +1,128 @@ +"""配置区组件""" + +from textual.containers import Vertical +from textual.widgets import Static, Input, Collapsible +from textual.app import ComposeResult +from textual.message import Message +from ..models.config import AppConfig + + +class ConfigSection(Vertical): + """可折叠配置区""" + + class ConfigChanged(Message): + """配置变更事件""" + def __init__(self, config: AppConfig) -> None: + self.config = config + super().__init__() + + def __init__(self, config: AppConfig | None = None, **kwargs): + super().__init__(**kwargs) + self._config = config or AppConfig() + + def compose(self) -> ComposeResult: + """构建配置区""" + # 直接显示配置,不使用Collapsible + title = Static("━━ 配置 ━━", classes="config-title") + title.can_focus = False + yield title + + label1 = Static("API Key:", classes="config-label") + label1.can_focus = False + yield label1 + + yield Input( + value=self._config.api_key, + placeholder="sk-xxxxxxxxxxxxx", + id="api-key-input", + password=True + ) + + label2 = Static("模型:", classes="config-label") + label2.can_focus = False + yield label2 + + yield Input( + value=self._config.model, + placeholder="deepseek-chat", + id="model-input" + ) + + label3 = Static("Base URL:", classes="config-label") + label3.can_focus = False + yield label3 + + yield Input( + value=self._config.base_url, + placeholder="https://api.deepseek.com", + id="base-url-input" + ) + + hint = Static("按Enter保存配置", classes="config-hint") + hint.can_focus = False + yield hint + + def on_mount(self) -> None: + """组件挂载时设置Tab顺序并加载配置""" + try: + api_key = self.query_one("#api-key-input", Input) + model = self.query_one("#model-input", Input) + base_url = self.query_one("#base-url-input", Input) + + # 设置Tab索引 + api_key.tab_index = 0 + model.tab_index = 1 + base_url.tab_index = 2 + + # 如果配置有值,更新输入框 + if self._config.api_key: + api_key.value = self._config.api_key + if self._config.model: + model.value = self._config.model + if self._config.base_url: + base_url.value = self._config.base_url + except Exception: + pass + + def on_input_changed(self, event: Input.Changed) -> None: + """处理输入变更事件""" + # 防抖:只在用户停止输入时更新 + pass # 不在输入时实时更新,避免卡顿 + + def on_input_submitted(self, event: Input.Submitted) -> None: + """处理输入提交事件(按Enter或Tab)""" + # 只在提交时更新配置 + try: + api_key_input = self.query_one("#api-key-input", Input) + model_input = self.query_one("#model-input", Input) + base_url_input = self.query_one("#base-url-input", Input) + + # 更新配置 + self._config = AppConfig( + api_key=api_key_input.value, + model=model_input.value, + base_url=base_url_input.value + ) + + # 发送配置变更事件 + self.post_message(self.ConfigChanged(self._config)) + except Exception as e: + pass + + def get_config(self) -> AppConfig: + """获取当前配置""" + return self._config + + def set_config(self, config: AppConfig) -> None: + """设置配置""" + self._config = config + try: + api_key_input = self.query_one("#api-key-input", Input) + model_input = self.query_one("#model-input", Input) + base_url_input = self.query_one("#base-url-input", Input) + + api_key_input.value = config.api_key + model_input.value = config.model + base_url_input.value = config.base_url + except Exception: + pass diff --git a/ui/widgets/cypher_query_box.py b/ui/widgets/cypher_query_box.py new file mode 100644 index 0000000..789db02 --- /dev/null +++ b/ui/widgets/cypher_query_box.py @@ -0,0 +1,57 @@ +"""Cypher查询框组件""" + +from textual.containers import Container, Horizontal +from textual.widgets import Static, TextArea, Button +from textual.app import ComposeResult +from textual.message import Message + + +class CypherQueryBox(Container): + """快捷Cypher查询输入框""" + + class ExecuteQuery(Message): + """执行查询事件""" + def __init__(self, query: str) -> None: + self.query = query + super().__init__() + + def compose(self) -> ComposeResult: + """构建查询框""" + yield Static("F4:执行Cypher查询", classes="query-title") + yield TextArea( + placeholder="输入Cypher查询语句...", + id="cypher-textarea" + ) + with Horizontal(classes="query-buttons"): + yield Button("执行", id="execute-button", variant="primary") + yield Button("清空", id="clear-button") + + def on_button_pressed(self, event: Button.Pressed) -> None: + """处理按钮点击""" + if event.button.id == "execute-button": + self._execute_query() + elif event.button.id == "clear-button": + self._clear_query() + + def on_key(self, event) -> None: + """处理按键事件""" + if event.key == "enter" and event.ctrl: + event.stop() + self._execute_query() + + def _execute_query(self) -> None: + """执行查询""" + textarea = self.query_one("#cypher-textarea", TextArea) + query = textarea.text.strip() + if query: + self.post_message(self.ExecuteQuery(query)) + + def _clear_query(self) -> None: + """清空查询""" + textarea = self.query_one("#cypher-textarea", TextArea) + textarea.clear() + + def focus(self) -> None: + """聚焦查询框""" + textarea = self.query_one("#cypher-textarea", TextArea) + textarea.focus() diff --git a/ui/widgets/input_box.py b/ui/widgets/input_box.py new file mode 100644 index 0000000..47a8904 --- /dev/null +++ b/ui/widgets/input_box.py @@ -0,0 +1,50 @@ +"""输入框组件""" + +from textual.containers import Container +from textual.widgets import Input +from textual.message import Message + + +class InputBox(Container): + """输入框组件""" + + class SendMessage(Message): + """发送消息事件""" + def __init__(self, content: str) -> None: + self.content = content + super().__init__() + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._history: list[str] = [] + self._history_index: int = -1 + + def compose(self): + """构建输入框""" + yield Input( + placeholder="输入消息... (Enter发送)", + id="input-textarea" + ) + + def on_mount(self) -> None: + """组件挂载时""" + # 设置焦点 + input_widget = self.query_one(Input) + input_widget.focus() + + def on_input_submitted(self, event: Input.Submitted) -> None: + """处理输入提交事件""" + content = event.value.strip() + if content: + # 保存到历史 + self._history.append(content) + self._history_index = len(self._history) + # 发送消息 + self.post_message(self.SendMessage(content)) + # 清空输入框 + event.input.value = "" + + def focus(self) -> None: + """聚焦输入框""" + input_widget = self.query_one(Input) + input_widget.focus() diff --git a/ui/widgets/left_panel.py b/ui/widgets/left_panel.py new file mode 100644 index 0000000..6cdc22b --- /dev/null +++ b/ui/widgets/left_panel.py @@ -0,0 +1,23 @@ +"""左侧面板""" + +from textual.containers import Container +from textual.app import ComposeResult +from .message_history import MessageHistory +from .input_box import InputBox + + +class LeftPanel(Container): + """左侧主面板""" + + def compose(self) -> ComposeResult: + """构建左侧面板""" + yield MessageHistory() + yield InputBox() + + def get_message_history(self) -> MessageHistory: + """获取消息历史组件""" + return self.query_one(MessageHistory) + + def get_input_box(self) -> InputBox: + """获取输入框组件""" + return self.query_one(InputBox) diff --git a/ui/widgets/message_history.py b/ui/widgets/message_history.py new file mode 100644 index 0000000..e755551 --- /dev/null +++ b/ui/widgets/message_history.py @@ -0,0 +1,57 @@ +"""消息历史组件""" + +from textual.containers import ScrollableContainer +from textual.message import Message +from .message_widget import MessageWidget +from ..models.message import Message as MessageModel + + +class MessageHistory(ScrollableContainer): + """消息历史区域""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._messages: list[MessageModel] = [] + + def compose(self): + """构建消息历史""" + for message in self._messages: + yield MessageWidget(message) + + def add_message(self, message: MessageModel) -> None: + """添加新消息""" + self._messages.append(message) + # 添加新组件 + message_widget = MessageWidget(message) + self.mount(message_widget) + # 滚动到最新消息 + self.scroll_to_widget(message_widget, animate=False) + + def update_latest_message(self, content: str) -> None: + """更新最新消息的内容""" + if self.children: + latest_widget = self.children[-1] + if isinstance(latest_widget, MessageWidget): + latest_widget.update_content(content) + # 确保滚动到最新消息 + self.scroll_to_widget(latest_widget, animate=False) + + def clear_messages(self) -> None: + """清空消息历史""" + self._messages.clear() + # 移除所有子组件 + for child in self.children: + child.remove() + + def get_latest_message(self) -> MessageModel | None: + """获取最新消息""" + if self._messages: + return self._messages[-1] + return None + + def toggle_latest_tool_details(self) -> None: + """切换最新消息的工具详情""" + if self.children: + latest_widget = self.children[-1] + if isinstance(latest_widget, MessageWidget): + latest_widget.toggle_tool_details() diff --git a/ui/widgets/message_widget.py b/ui/widgets/message_widget.py new file mode 100644 index 0000000..a6929de --- /dev/null +++ b/ui/widgets/message_widget.py @@ -0,0 +1,108 @@ +"""消息组件""" + +from textual.containers import Container, Vertical +from textual.widgets import Static +from textual.message import Message +from textual.css.query import NoMatches +from ..models.message import Message as MessageModel + + +class MessageWidget(Container): + """单条消息组件""" + + def __init__(self, message: MessageModel, **kwargs): + super().__init__(**kwargs) + self._message = message + self._show_tool_details = False + self._content_widget = None # 保存内容组件的引用 + self._tool_details_container = None # 保存工具详情容器引用 + + def compose(self): + """构建消息组件""" + # 消息头 + role_emoji = "🟠" if self._message.role == "user" else "🔵" + timestamp_str = self._message.timestamp.strftime("%H:%M:%S") + yield Static( + f"{role_emoji} {timestamp_str}", + classes="message-header" + ) + + # 消息内容 - 保存引用以便后续更新 + self._content_widget = Static( + self._message.content, + classes="message-content" + ) + yield self._content_widget + + # 工具调用指示器 + if self._message.tool_calls: + tool_count = len(self._message.tool_calls) + toggle_hint = "(F3折叠)" if self._show_tool_details else "(F3展开)" + yield Static( + f"[工具:{tool_count}次] {toggle_hint}", + classes="tool-indicator" + ) + + # 工具调用详情容器 - 始终创建,但根据状态显示/隐藏 + self._tool_details_container = Vertical(classes="tool-details") + with self._tool_details_container: + for i, tool_call in enumerate(self._message.tool_calls, 1): + yield Static( + f"工具 {i}: {tool_call.name}", + classes="tool-name" + ) + yield Static( + f"参数: {tool_call.arguments}", + classes="tool-args" + ) + + # 显示执行结果 + if self._message.tool_results: + for result in self._message.tool_results: + if result.tool_call_id == tool_call.id: + # 显示完整结果,不截断 + result_text = result.result + # 如果结果太长,只显示前1000字符,但提供完整信息 + if len(result_text) > 1000: + result_text = result_text[:1000] + f"\n... (共{len(result.result)}字符,按F3查看完整内容)" + yield Static( + f"结果: {result_text}", + classes="tool-result" + ) + + # 根据状态设置初始显示/隐藏 + if not self._show_tool_details: + self._tool_details_container.styles.display = "none" + + def update_content(self, new_content: str) -> None: + """更新消息内容""" + self._message.content = new_content + if self._content_widget: + self._content_widget.update(new_content) + + def toggle_tool_details(self) -> None: + """切换工具详情显示状态""" + if self._message.tool_calls and self._tool_details_container: + self._show_tool_details = not self._show_tool_details + + # 切换显示/隐藏 + if self._show_tool_details: + self._tool_details_container.styles.display = "block" + else: + self._tool_details_container.styles.display = "none" + + # 更新指示器文字 + self._update_indicator() + + # 刷新布局 + self.refresh(layout=True) + + def _update_indicator(self) -> None: + """更新工具调用指示器文字""" + try: + indicator = self.query_one(".tool-indicator", Static) + tool_count = len(self._message.tool_calls) + toggle_hint = "(F3折叠)" if self._show_tool_details else "(F3展开)" + indicator.update(f"[工具:{tool_count}次] {toggle_hint}") + except NoMatches: + pass diff --git a/ui/widgets/operation_log.py b/ui/widgets/operation_log.py new file mode 100644 index 0000000..7806f97 --- /dev/null +++ b/ui/widgets/operation_log.py @@ -0,0 +1,66 @@ +"""操作日志组件""" + +from datetime import datetime +from textual.containers import ScrollableContainer +from textual.widgets import Static +from ..models.log_entry import LogEntry + + +class OperationLog(ScrollableContainer): + """图操作日志区域""" + + def __init__(self, max_entries: int = 100, **kwargs): + super().__init__(**kwargs) + self._logs: list[LogEntry] = [] + self._max_entries = max_entries + + def compose(self): + """构建日志区域""" + if not self._logs: + yield Static("暂无操作日志", classes="log-empty") + + def add_log(self, entry: LogEntry) -> None: + """添加日志(插入到顶部)""" + # 限制日志数量 + if len(self._logs) >= self._max_entries: + self._logs.pop() + # 移除最旧的组件 + if self.children: + self.children[-1].remove() + + # 插入到列表开头 + self._logs.insert(0, entry) + + # 创建日志显示组件 + log_widget = self._create_log_widget(entry) + + # 挂载到顶部 + self.mount(log_widget, before=0 if self.children else None) + + # 滚动到顶部 + self.scroll_to(0, animate=False) + + def _create_log_widget(self, entry: LogEntry) -> Static: + """创建日志显示组件""" + timestamp_str = entry.timestamp.strftime("%H:%M:%S") + text = ( + f"[{timestamp_str}] {entry.tool_name}\n" + f" 参数: {entry.args_summary}\n" + f" 结果: {entry.result_summary}\n" + f" 耗时: {entry.duration:.2f}s" + ) + return Static(text, classes="log-entry") + + def clear_logs(self) -> None: + """清空日志""" + self._logs.clear() + for child in self.children: + child.remove() + # 显示空状态 + self.mount(Static("暂无操作日志", classes="log-empty")) + + def get_latest_log(self) -> LogEntry | None: + """获取最新日志""" + if self._logs: + return self._logs[0] + return None diff --git a/ui/widgets/right_panel.py b/ui/widgets/right_panel.py new file mode 100644 index 0000000..d2a9591 --- /dev/null +++ b/ui/widgets/right_panel.py @@ -0,0 +1,67 @@ +"""右侧面板""" + +from textual.containers import Container, ScrollableContainer +from textual.css.query import NoMatches +from textual.widgets import Static +from textual.app import ComposeResult +from .config_section import ConfigSection +from .operation_log import OperationLog +from .cypher_query_box import CypherQueryBox +from ..models.config import AppConfig + + +class RightPanel(Container): + """右侧边栏""" + + def __init__(self, config: AppConfig | None = None, use_embedded_db: bool = True, **kwargs): + super().__init__(**kwargs) + self._is_collapsed = False + self._config = config or AppConfig() + self._use_embedded_db = use_embedded_db + + def compose(self) -> ComposeResult: + """构建右侧面板""" + yield Static("F2:隐藏侧边栏", classes="sidebar-title") + with ScrollableContainer(): + yield ConfigSection(self._config) + yield OperationLog() + if not self._use_embedded_db: + yield CypherQueryBox() + + def toggle(self) -> None: + """切换折叠/展开""" + self._is_collapsed = not self._is_collapsed + if self._is_collapsed: + self.styles.width = 0 + self.styles.display = "none" + else: + self.styles.width = 70 + self.styles.display = "block" + + def is_collapsed(self) -> bool: + """检查是否折叠""" + return self._is_collapsed + + def get_config_section(self) -> ConfigSection: + """获取配置区组件""" + return self.query_one(ConfigSection) + + def get_operation_log(self) -> OperationLog: + """获取操作日志组件""" + return self.query_one(OperationLog) + + def get_cypher_query_box(self) -> CypherQueryBox | None: + """获取Cypher查询框组件(可能不存在)""" + try: + return self.query_one(CypherQueryBox) + except NoMatches: + return None + + def has_cypher_query_box(self) -> bool: + """检查是否存在Cypher查询框""" + return not self._use_embedded_db + + def update_title(self) -> None: + """更新标题""" + title = self.query_one(Static) + title.update("F2:展开侧边栏" if self._is_collapsed else "F2:隐藏侧边栏") diff --git a/ui/widgets/status_bar.py b/ui/widgets/status_bar.py new file mode 100644 index 0000000..bd29dcb --- /dev/null +++ b/ui/widgets/status_bar.py @@ -0,0 +1,27 @@ +"""状态栏组件""" + +from textual.widgets import Static +from textual.message import Message + + +class StatusBar(Static): + """底部状态栏""" + + class FocusChanged(Message): + """焦点变更事件""" + def __init__(self, focus_name: str) -> None: + self.focus_name = focus_name + super().__init__() + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._shortcuts = "F1:帮助 F2:侧边栏 F3:工具详情 F4:查询 F5:清屏 F6:退出" + self._license_info = "本项目由jianf设计,以GPLv3形式开源" + + def on_mount(self) -> None: + """组件挂载时""" + self._update_display() + + def _update_display(self) -> None: + """更新显示""" + self.update(f"{self._license_info} | {self._shortcuts}")