记忆系统,上下文处理+mem0
This commit is contained in:
@@ -1,6 +1,8 @@
|
||||
"""Agent Loop 主循环:LLM 对话 -> 工具调用 -> 结果回传 -> 继续。"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import os
|
||||
from typing import AsyncGenerator, Optional
|
||||
@@ -8,9 +10,17 @@ from typing import AsyncGenerator, Optional
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from app.agent.tools import TOOL_DEFINITIONS, execute_tool
|
||||
from app.config import get_llm_model, get_llm_max_iterations, get_image_model_config
|
||||
from app.config import (
|
||||
get_llm_model,
|
||||
get_llm_max_iterations,
|
||||
get_image_model_config,
|
||||
get_max_recent_turns,
|
||||
)
|
||||
from app.memory import get_memory
|
||||
from app.services.image_gen import to_data_uri
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_client() -> AsyncOpenAI:
|
||||
base_url = os.getenv("OPENAI_BASE_URL")
|
||||
@@ -59,6 +69,7 @@ async def run_agent_loop(
|
||||
messages: list[dict],
|
||||
ref_image_url: Optional[str] = None,
|
||||
image_model: Optional[str] = None,
|
||||
session_id: Optional[str] = None,
|
||||
) -> AsyncGenerator[dict, None]:
|
||||
"""
|
||||
运行 Agent Loop,以 SSE 事件流形式 yield 结果。
|
||||
@@ -67,10 +78,42 @@ async def run_agent_loop(
|
||||
- text_delta: LLM 文字增量
|
||||
- tool_start: 开始执行工具
|
||||
- image_result: 图片生成结果
|
||||
- memory_warning: 记忆存储异常(不中断对话)
|
||||
- done: 完成
|
||||
- error: 错误
|
||||
"""
|
||||
# 构建消息列表:将当前生图模型信息注入 system prompt,避免 LLM 猜测模型名
|
||||
# ── 滑动窗口:截断过早的消息 ──
|
||||
max_turns = get_max_recent_turns()
|
||||
if len(messages) > max_turns:
|
||||
messages = messages[-max_turns:]
|
||||
|
||||
# ── 记忆检索:用最新 user 消息语义检索相关记忆 ──
|
||||
memory_block = ""
|
||||
try:
|
||||
memory = get_memory()
|
||||
last_user_content = ""
|
||||
for m in reversed(messages):
|
||||
if m["role"] == "user":
|
||||
last_user_content = m["content"]
|
||||
break
|
||||
|
||||
if last_user_content:
|
||||
search_kwargs = {"query": last_user_content, "user_id": "default_user", "limit": 10}
|
||||
if session_id:
|
||||
search_kwargs["run_id"] = session_id
|
||||
relevant = memory.search(**search_kwargs)
|
||||
|
||||
results = relevant.get("results", []) if isinstance(relevant, dict) else relevant
|
||||
if results:
|
||||
items = "\n".join(f"- {m['memory']}" for m in results if m.get("memory"))
|
||||
if items:
|
||||
memory_block = f"\n\n## 用户记忆(来自历史对话)\n{items}"
|
||||
except Exception as e:
|
||||
logger.error("Mem0 记忆检索失败: %s", e, exc_info=True)
|
||||
yield {"type": "error", "data": {"message": f"记忆系统检索失败: {e}"}}
|
||||
return
|
||||
|
||||
# ── 构建 system prompt ──
|
||||
model_config = get_image_model_config(image_model)
|
||||
current_model_name = model_config.get('name', '未知')
|
||||
current_model_id = model_config.get('id', '未知')
|
||||
@@ -81,7 +124,7 @@ async def run_agent_loop(
|
||||
f"**注意**:对话历史中可能包含之前使用其他模型的记录,忽略那些旧模型名。"
|
||||
f"本次生成使用的是 {current_model_name},在回复中只能使用这个名称。"
|
||||
)
|
||||
api_messages = [{"role": "system", "content": SYSTEM_PROMPT + model_hint}]
|
||||
api_messages = [{"role": "system", "content": SYSTEM_PROMPT + model_hint + memory_block}]
|
||||
|
||||
# 检测当前 LLM 是否支持 vision(多模态图片输入)
|
||||
llm_model = get_llm_model().lower()
|
||||
@@ -90,7 +133,6 @@ async def run_agent_loop(
|
||||
for msg in messages:
|
||||
if msg["role"] == "user" and ref_image_url and msg is messages[-1]:
|
||||
if vision_capable:
|
||||
# 支持 vision 的模型:直接发送图片
|
||||
api_messages.append({
|
||||
"role": "user",
|
||||
"content": [
|
||||
@@ -102,7 +144,6 @@ async def run_agent_loop(
|
||||
],
|
||||
})
|
||||
else:
|
||||
# 不支持 vision 的模型:以文字提示告知有参考图,风格由 IP-Adapter 处理
|
||||
hint = (
|
||||
f"{msg['content']}\n\n"
|
||||
"【系统提示:用户上传了一张参考图,已自动传递给图片生成工具的 IP-Adapter。"
|
||||
@@ -181,7 +222,9 @@ async def run_agent_loop(
|
||||
),
|
||||
})
|
||||
continue
|
||||
yield {"type": "done", "data": {}}
|
||||
# 对话正常结束,异步存储记忆
|
||||
async for evt in _store_and_done(messages, session_id):
|
||||
yield evt
|
||||
return
|
||||
|
||||
# 将 assistant 消息(含 tool_calls)加入历史
|
||||
@@ -253,6 +296,34 @@ async def run_agent_loop(
|
||||
# ref_image_url 不清除:消息构建只在循环外执行一次,
|
||||
# 后续工具调用仍需参考图(InstantStyle 等模型必须有 style_image)
|
||||
|
||||
# 迭代次数用尽,存储记忆后结束
|
||||
async for evt in _store_and_done(messages, session_id):
|
||||
yield evt
|
||||
|
||||
|
||||
# ─── 异步记忆存储 ─────────────────────────────────────────
|
||||
|
||||
|
||||
async def _store_and_done(
|
||||
messages: list[dict], session_id: Optional[str]
|
||||
) -> AsyncGenerator[dict, None]:
|
||||
"""触发 Mem0 异步存储后 yield done。存储失败时 yield warning 但不中断。"""
|
||||
try:
|
||||
memory = get_memory()
|
||||
recent = messages[-4:] if len(messages) >= 4 else messages
|
||||
add_kwargs: dict = {"user_id": "default_user"}
|
||||
if session_id:
|
||||
add_kwargs["run_id"] = session_id
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.run_in_executor(None, lambda: memory.add(recent, **add_kwargs))
|
||||
except Exception as e:
|
||||
logger.error("Mem0 记忆存储失败: %s", e, exc_info=True)
|
||||
yield {
|
||||
"type": "memory_warning",
|
||||
"data": {"message": f"记忆存储失败: {e}"},
|
||||
}
|
||||
|
||||
yield {"type": "done", "data": {}}
|
||||
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ async def chat(
|
||||
ref_image: Optional[UploadFile] = File(None),
|
||||
ref_image_url: Optional[str] = Form(None),
|
||||
image_model: Optional[str] = Form(None),
|
||||
session_id: Optional[str] = Form(None),
|
||||
):
|
||||
"""
|
||||
主对话端点。
|
||||
@@ -61,6 +62,7 @@ async def chat(
|
||||
- ref_image: 可选的参考图文件(兼容旧方式)
|
||||
- ref_image_url: 可选,已通过 /upload-ref-image 上传后的服务端路径
|
||||
- image_model: 可选,指定本次使用的生图模型短 ID
|
||||
- session_id: 可选,前端会话 ID,用于 Mem0 记忆作用域
|
||||
"""
|
||||
parsed_messages = json.loads(messages)
|
||||
|
||||
@@ -72,7 +74,10 @@ async def chat(
|
||||
|
||||
async def event_generator():
|
||||
async for event in run_agent_loop(
|
||||
parsed_messages, resolved_ref_url, image_model=image_model
|
||||
parsed_messages,
|
||||
resolved_ref_url,
|
||||
image_model=image_model,
|
||||
session_id=session_id,
|
||||
):
|
||||
yield {
|
||||
"event": event["type"],
|
||||
|
||||
@@ -148,6 +148,26 @@ def get_image_models_list() -> list[dict]:
|
||||
]
|
||||
|
||||
|
||||
# ─── 记忆系统配置 ─────────────────────────────────────
|
||||
|
||||
def get_deepseek_api_key() -> str:
|
||||
return os.getenv("DEEPSEEK_API_KEY", "")
|
||||
|
||||
|
||||
def get_ollama_base_url() -> str:
|
||||
return os.getenv("OLLAMA_BASE_URL", "http://localhost:11434")
|
||||
|
||||
|
||||
def get_mem0_embedding_model() -> str:
|
||||
return os.getenv("MEM0_EMBEDDING_MODEL", "nomic-embed-text")
|
||||
|
||||
|
||||
def get_max_recent_turns() -> int:
|
||||
return int(os.getenv("MAX_RECENT_TURNS", "20"))
|
||||
|
||||
|
||||
# ─── 图像输出配置 ─────────────────────────────────────
|
||||
|
||||
def get_image_aspect_ratio() -> str:
|
||||
return os.getenv("IMAGE_ASPECT_RATIO", "1:1")
|
||||
|
||||
|
||||
99
art-agent/backend/app/memory.py
Normal file
99
art-agent/backend/app/memory.py
Normal file
@@ -0,0 +1,99 @@
|
||||
"""
|
||||
Mem0 记忆层初始化与健康检查。
|
||||
启动时验证 Ollama 可达和模型可用,失败则抛出 RuntimeError 阻止后端启动。
|
||||
"""
|
||||
|
||||
import logging
|
||||
import httpx
|
||||
from mem0 import Memory
|
||||
|
||||
from app.config import (
|
||||
get_deepseek_api_key,
|
||||
get_ollama_base_url,
|
||||
get_mem0_embedding_model,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_memory: Memory | None = None
|
||||
|
||||
|
||||
def _check_ollama_health() -> None:
|
||||
"""检查 Ollama 服务是否可达,以及 embedding 模型是否已安装。"""
|
||||
base_url = get_ollama_base_url()
|
||||
model_name = get_mem0_embedding_model()
|
||||
|
||||
try:
|
||||
resp = httpx.get(f"{base_url}/api/tags", timeout=5)
|
||||
resp.raise_for_status()
|
||||
except (httpx.ConnectError, httpx.TimeoutException, httpx.HTTPStatusError) as e:
|
||||
raise RuntimeError(
|
||||
f"无法连接 Ollama 服务({base_url})。"
|
||||
f"请确认 Ollama 已启动:启动方式参见 https://ollama.com\n"
|
||||
f"原始错误:{e}"
|
||||
) from e
|
||||
|
||||
models = resp.json().get("models", [])
|
||||
installed = [m.get("name", "").split(":")[0] for m in models]
|
||||
if model_name not in installed:
|
||||
raise RuntimeError(
|
||||
f"Ollama 已运行,但未找到 embedding 模型 '{model_name}'。\n"
|
||||
f"已安装的模型:{installed}\n"
|
||||
f"请运行:ollama pull {model_name}"
|
||||
)
|
||||
|
||||
logger.info("Ollama 健康检查通过:%s 模型可用", model_name)
|
||||
|
||||
|
||||
def init_memory() -> Memory:
|
||||
"""初始化 Mem0 Memory 单例。首次调用时执行健康检查。"""
|
||||
global _memory
|
||||
if _memory is not None:
|
||||
return _memory
|
||||
|
||||
_check_ollama_health()
|
||||
|
||||
deepseek_key = get_deepseek_api_key()
|
||||
if not deepseek_key:
|
||||
raise RuntimeError(
|
||||
"DEEPSEEK_API_KEY 未配置。Mem0 需要该 key 进行事实提取。\n"
|
||||
"请在 .env 中设置 DEEPSEEK_API_KEY"
|
||||
)
|
||||
|
||||
config = {
|
||||
"llm": {
|
||||
"provider": "deepseek",
|
||||
"config": {
|
||||
"model": "deepseek-chat",
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 1500,
|
||||
"api_key": deepseek_key,
|
||||
},
|
||||
},
|
||||
"embedder": {
|
||||
"provider": "ollama",
|
||||
"config": {
|
||||
"model": get_mem0_embedding_model(),
|
||||
"ollama_base_url": get_ollama_base_url(),
|
||||
},
|
||||
},
|
||||
"vector_store": {
|
||||
"provider": "qdrant",
|
||||
"config": {
|
||||
"collection_name": "epeekit_memories",
|
||||
"embedding_model_dims": 768,
|
||||
"path": "./data/qdrant",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_memory = Memory.from_config(config)
|
||||
logger.info("Mem0 记忆层初始化成功")
|
||||
return _memory
|
||||
|
||||
|
||||
def get_memory() -> Memory:
|
||||
"""获取已初始化的 Memory 实例。未初始化时自动调用 init_memory()。"""
|
||||
if _memory is None:
|
||||
return init_memory()
|
||||
return _memory
|
||||
Reference in New Issue
Block a user