refactor: 统一 Packet 通信协议 + 后端配置管理 + UI 清理

- 合并 server.py 到 core/__init__.py,使用统一 Packet 协议
- 后端管理配置持久化 (~/.trulymem/config.json)
- 前端移除 ConfigService,通过 BackendClient 与后端通信
- 删除 UI 中冗余的 AI 推理逻辑 (chat_service, tool_service, message_handler)
- 删除 core/tools 重复文件 (tool_executor, tool_limiter)
- 提示词管理器支持用户自定义 (~/.trulyemem/system_prompt.md)
- 启动入口优化配置路径逻辑
- 更新测试覆盖 (42 tests)
- 更新文档
This commit is contained in:
root
2026-04-14 11:49:39 +08:00
parent 2a76fc6477
commit c742a30e1b
19 changed files with 1020 additions and 2349 deletions

View File

@ -2,17 +2,16 @@ import threading
import queue
import time
import json
import os
from pathlib import Path
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):
class PacketType(Enum):
PROCESS_MESSAGE = "process_message"
EXECUTE_TOOL = "execute_tool"
GET_STATUS = "get_status"
@ -24,46 +23,66 @@ class MessageType(Enum):
@dataclass
class BackendRequest:
request_id: str
message_type: MessageType
payload: Dict[str, Any]
response_queue: queue.Queue = field(default=None)
class Packet:
id: str
type: PacketType
body: Dict[str, Any]
response_queue: Optional[queue.Queue] = field(default=None)
created_at: float = field(default_factory=time.time)
@dataclass
class BackendResponse:
request_id: str
class PacketResponse:
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):
DEFAULT_CONFIG_PATH = Path.home() / ".trulymem" / "config.json"
def __init__(self, db_path: str = "graph_memory.db", use_embedded_db: bool = True, config_file: str = None):
self._db_path = db_path
self._use_embedded_db = use_embedded_db
self._config_file = Path(config_file) if config_file else self.DEFAULT_CONFIG_PATH
self._graph = None
self._client = None
self._tool_limiter = ToolLimiter()
self._tool_limiter = None
self._request_queue: queue.Queue[BackendRequest] = queue.Queue()
self._input_queue: queue.Queue[Packet] = queue.Queue()
self._response_queues: Dict[str, queue.Queue] = {}
self._running = False
self._thread: Optional[threading.Thread] = None
self._lock = threading.Lock()
self._config = {"api_key": "", "base_url": "https://api.deepseek.com", "model": "deepseek-chat"}
self._message_history: list = []
def start(self, api_key: str = "", base_url: str = "https://api.deepseek.com") -> None:
def start(self, api_key: str = "", base_url: str = "https://api.deepseek.com", model: str = "deepseek-chat") -> None:
if self._running:
return
self._init_graph()
self._load_config()
if api_key:
self._config["api_key"] = api_key
if base_url:
self._config["base_url"] = base_url
if model:
self._config["model"] = model
self._init_graph()
self._tool_limiter = self._create_tool_limiter()
if self._config["api_key"]:
from .graph_client import GraphMemoryClient
self._client = GraphMemoryClient(
api_key=api_key,
base_url=base_url,
api_key=self._config["api_key"],
base_url=self._config["base_url"],
model=self._config.get("model", "deepseek-chat"),
graph=self._graph
)
@ -71,6 +90,24 @@ class BackendServer:
self._thread = threading.Thread(target=self._run_loop, daemon=True)
self._thread.start()
def _load_config(self) -> None:
if self._config_file.exists():
try:
with open(self._config_file, 'r') as f:
saved = json.load(f)
self._config.update(saved)
except Exception:
pass
def _save_config(self) -> None:
self._config_file.parent.mkdir(parents=True, exist_ok=True)
with open(self._config_file, 'w') as f:
json.dump(self._config, f, indent=2)
def _create_tool_limiter(self):
from .tool_limiter import ToolLimiter
return ToolLimiter()
def _init_graph(self) -> None:
if self._use_embedded_db:
self._graph = EmbeddedGraphDB(db_path=self._db_path)
@ -85,310 +122,276 @@ class BackendServer:
def _run_loop(self) -> None:
while self._running:
try:
request = self._request_queue.get(timeout=0.1)
packet = self._input_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.GET_CONFIG:
self._handle_get_config(request)
elif request.message_type == MessageType.SET_CONFIG:
self._handle_set_config(request)
elif request.message_type == MessageType.GET_HISTORY:
self._handle_get_history(request)
elif request.message_type == MessageType.SAVE_HISTORY:
self._handle_save_history(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"}
))
self._process_packet(packet)
def _handle_process_message(self, request: BackendRequest) -> None:
def _process_packet(self, packet: Packet) -> None:
response_body = {"error": "not implemented"}
try:
user_input = request.payload.get("user_input", "")
if packet.type == PacketType.PROCESS_MESSAGE:
response_body = self._handle_process_message(packet.body)
elif packet.type == PacketType.EXECUTE_TOOL:
response_body = self._handle_execute_tool(packet.body)
elif packet.type == PacketType.GET_STATUS:
response_body = self._handle_get_status()
elif packet.type == PacketType.GET_CONFIG:
response_body = self._handle_get_config()
elif packet.type == PacketType.SET_CONFIG:
response_body = self._handle_set_config(packet.body)
elif packet.type == PacketType.GET_HISTORY:
response_body = self._handle_get_history()
elif packet.type == PacketType.SAVE_HISTORY:
response_body = self._handle_save_history(packet.body)
elif packet.type == PacketType.SHUTDOWN:
self._running = False
response_body = {"success": True, "status": "shutdown"}
if not self._client:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error="API Key 未配置"
))
return
if "success" not in response_body:
response_body["success"] = True
except Exception as e:
response_body["success"] = False
response_body["error"] = str(e)
self._send_response(packet.id, PacketResponse(
id=packet.id,
success=response_body.get("success", False),
data=response_body if response_body.get("success") else None,
error=response_body.get("error")
))
def _handle_process_message(self, body: Dict) -> Dict:
from .tool_executor import execute_tool
user_input = body.get("user_input", "")
if not self._client:
return {"success": False, "error": "API Key 未配置", "content": "请先配置 API Key"}
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"
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
assistant_msg = {
"role": "assistant",
"content": message.content,
"tool_calls": [
{
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
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
})
} 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
messages_history.extend(current_tool_results)
self._tool_limiter.record_call(tool_call.function.name, args)
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
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)
except Exception as e:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error=str(e)
))
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()}"
return {
"success": True,
"content": content,
"tool_calls": tool_calls,
"rejected_tools": rejected_tools
}
def _handle_execute_tool(self, request: BackendRequest) -> None:
"""处理直接工具调用请求(前端直接调用,不受次数限制)"""
def _handle_execute_tool(self, body: Dict) -> Dict:
from .tool_executor import execute_tool
try:
tool_name = request.payload.get("tool_name")
arguments = request.payload.get("arguments", {})
tool_name = body.get("tool_name")
arguments = body.get("arguments", {})
# 前端直接调用的工具不受次数限制,直接执行
result = execute_tool(self._graph, tool_name, arguments)
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=True,
data={"result": result}
))
return {"success": True, "result": result}
except Exception as e:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error=str(e)
))
return {"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 _handle_get_status(self) -> Dict:
return {
"running": self._running,
"config": self._config,
"graph_initialized": self._graph is not None,
"client_initialized": self._client is not None
}
def _handle_get_config(self, request: BackendRequest) -> None:
try:
config = self.get_config()
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=True,
data=config
))
except Exception as e:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error=str(e)
))
def _handle_get_config(self) -> Dict:
return self._config.copy()
def _handle_set_config(self, request: BackendRequest) -> None:
try:
api_key = request.payload.get("api_key", "")
base_url = request.payload.get("base_url", "https://api.deepseek.com")
self.update_config(api_key, base_url)
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=True,
data={"status": "config_updated"}
))
except Exception as e:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error=str(e)
))
def _handle_get_history(self, request: BackendRequest) -> None:
try:
history = self.get_message_history()
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=True,
data={"history": history}
))
except Exception as e:
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=False,
error=str(e)
))
def _handle_save_history(self, request: BackendRequest) -> None:
try:
messages = request.payload.get("messages", [])
self.save_message_history(messages)
self._send_response(request, BackendResponse(
request_id=request.request_id,
success=True,
data={"status": "history_saved"}
))
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()
def _handle_set_config(self, body: Dict) -> Dict:
api_key = body.get("api_key", "")
base_url = body.get("base_url", "https://api.deepseek.com")
model = body.get("model", "deepseek-chat")
request = BackendRequest(
request_id=request_id,
message_type=MessageType.PROCESS_MESSAGE,
payload={"user_input": user_input},
response_queue=response_queue
self.update_config(api_key, base_url, model)
self._save_config()
return {"status": "config_updated"}
def _handle_get_history(self) -> Dict:
return {"history": self._message_history}
def _handle_save_history(self, body: Dict) -> Dict:
messages = body.get("messages", [])
self._message_history = messages
return {"status": "history_saved"}
def _send_response(self, request_id: str, response: PacketResponse) -> None:
with self._lock:
q = self._response_queues.pop(request_id, None)
if q:
q.put(response)
def send(self, packet: Packet) -> Packet:
resp_q = queue.Queue()
with self._lock:
self._response_queues[packet.id] = resp_q
self._input_queue.put(packet)
try:
response = resp_q.get(timeout=30.0)
return Packet(
id=response.id,
type=packet.type,
body={
"success": response.success,
"data": response.data,
"error": response.error
}
)
except queue.Empty:
return Packet(
id=packet.id,
type=packet.type,
body={"success": False, "error": "timeout"}
)
finally:
with self._lock:
self._response_queues.pop(packet.id, None)
def process_message(self, user_input: str) -> Dict[str, Any]:
packet = Packet(
id=f"{time.time()}",
type=PacketType.PROCESS_MESSAGE,
body={"user_input": user_input}
)
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("请求超时")
response = self.send(packet)
return response.body
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
def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Dict[str, Any]:
packet = Packet(
id=f"{time.time()}",
type=PacketType.EXECUTE_TOOL,
body={"tool_name": tool_name, "arguments": arguments}
)
self._request_queue.put(request)
response = self.send(packet)
return response.body
def update_config(self, api_key: str, base_url: str = "https://api.deepseek.com", model: str = "deepseek-chat") -> None:
with self._lock:
self._config["api_key"] = api_key
self._config["base_url"] = base_url
self._config["model"] = model
try:
response = response_queue.get(timeout=timeout)
if not response.success:
raise Exception(response.error)
return response.data["result"]
except queue.Empty:
raise TimeoutError("工具执行超时")
if api_key and self._graph:
from .graph_client import GraphMemoryClient
self._client = GraphMemoryClient(
api_key=api_key,
base_url=base_url,
model=model,
graph=self._graph
)
def get_config(self) -> Dict[str, str]:
return self._config.copy()
def save_message_history(self, messages: list) -> None:
self._message_history = messages
def get_message_history(self) -> list:
return self._message_history.copy()
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
packet = Packet(
id=f"{time.time()}",
type=PacketType.SHUTDOWN,
body={}
)
self._request_queue.put(request)
self.send(packet)
if self._thread:
self._thread.join(timeout=2.0)
@ -396,22 +399,5 @@ class BackendServer:
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:
self._config = {"api_key": api_key, "base_url": base_url}
if api_key and self._graph:
self._client = GraphMemoryClient(
api_key=api_key,
base_url=base_url,
graph=self._graph
)
def get_config(self) -> Dict[str, str]:
return getattr(self, "_config", {"api_key": "", "base_url": "https://api.deepseek.com"})
def save_message_history(self, messages: list) -> None:
self._message_history = messages
def get_message_history(self) -> list:
return getattr(self, "_message_history", [])
self._running = False