From da05cd73e6f29d0fb1a22cc26d54461841139e7b Mon Sep 17 00:00:00 2001 From: guo zebin Date: Tue, 11 Nov 2025 14:45:49 +0800 Subject: [PATCH] =?UTF-8?q?=E8=87=AA=E7=84=B6=E8=BF=9B=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1.加入仿生记忆功能。 --- bionicmemory/__init__.py | 17 + bionicmemory/algorithms/__init__.py | 7 + .../algorithms/clustering_suppression.py | 147 ++ .../algorithms/newton_cooling_helper.py | 48 + bionicmemory/api/__init__.py | 6 + bionicmemory/api/proxy_server.py | 511 ++++++ bionicmemory/core/__init__.py | 7 + bionicmemory/core/chroma_service.py | 552 ++++++ bionicmemory/core/memory_system.py | 1488 ++++++++++++++++ bionicmemory/services/__init__.py | 10 + .../services/api_embedding_service.py | 218 +++ bionicmemory/services/chat_helper.py | 109 ++ .../services/local_embedding_service.py | 199 +++ .../services/memory_cleanup_scheduler.py | 312 ++++ bionicmemory/services/summary_service.py | 169 ++ bionicmemory/utils/__init__.py | 7 + bionicmemory/utils/logging_config.py | 88 + core/fay_core.py | 10 +- fay_booter.py | 34 +- faymcp/data/mcp_servers.json | 8 +- genagents/genagents_flask.py | 126 +- gui/flask_server.py | 129 +- gui/static/js/setting.js | 35 +- gui/templates/setting.html | 8 +- llm/nlp_bionicmemory_stream.py | 1502 +++++++++++++++++ mcp_servers/schedule_manager/server.py | 1 - requirements.txt | 4 +- test/test_fay_gpt_nonstream.py | 6 +- test/test_fay_gpt_stream.py | 6 +- utils/config_util.py | 27 +- 30 files changed, 5663 insertions(+), 128 deletions(-) create mode 100644 bionicmemory/__init__.py create mode 100644 bionicmemory/algorithms/__init__.py create mode 100644 bionicmemory/algorithms/clustering_suppression.py create mode 100644 bionicmemory/algorithms/newton_cooling_helper.py create mode 100644 bionicmemory/api/__init__.py create mode 100644 bionicmemory/api/proxy_server.py create mode 100644 bionicmemory/core/__init__.py create mode 100644 bionicmemory/core/chroma_service.py create mode 100644 bionicmemory/core/memory_system.py create mode 100644 bionicmemory/services/__init__.py create mode 100644 bionicmemory/services/api_embedding_service.py create mode 100644 bionicmemory/services/chat_helper.py create mode 100644 bionicmemory/services/local_embedding_service.py create mode 100644 bionicmemory/services/memory_cleanup_scheduler.py create mode 100644 bionicmemory/services/summary_service.py create mode 100644 bionicmemory/utils/__init__.py create mode 100644 bionicmemory/utils/logging_config.py create mode 100644 llm/nlp_bionicmemory_stream.py diff --git a/bionicmemory/__init__.py b/bionicmemory/__init__.py new file mode 100644 index 0000000..2d317d2 --- /dev/null +++ b/bionicmemory/__init__.py @@ -0,0 +1,17 @@ +""" +BionicMemory - 仿生记忆系统 + +基于仿生学原理的AI记忆管理系统,模拟人类大脑的长短期记忆机制, +通过科学的遗忘算法和智能的记忆管理,为AI应用提供真正个性化的记忆体验。 + +主要特性: +- 长短期记忆分层管理 +- 牛顿冷却遗忘算法 +- 聚类抑制机制 +- 上下文增强技术 +- 多租户安全隔离 +""" + +__version__ = "2.0.0" +__author__ = "BionicMemory Team" +__email__ = "contact@bionicmemory.ai" diff --git a/bionicmemory/algorithms/__init__.py b/bionicmemory/algorithms/__init__.py new file mode 100644 index 0000000..b22445e --- /dev/null +++ b/bionicmemory/algorithms/__init__.py @@ -0,0 +1,7 @@ +""" +算法模块 + +包含仿生记忆系统的核心算法: +- 牛顿冷却遗忘算法 +- 聚类抑制机制 +""" diff --git a/bionicmemory/algorithms/clustering_suppression.py b/bionicmemory/algorithms/clustering_suppression.py new file mode 100644 index 0000000..4f0a905 --- /dev/null +++ b/bionicmemory/algorithms/clustering_suppression.py @@ -0,0 +1,147 @@ +""" +基于聚类的记忆抑制机制 +实现从短期记忆中加载数倍目标条数,进行k-means聚类,每簇取最相似的代表 + +主要思路: +1. 从短期记忆中加载数倍(t:聚类平均条数)目标所需条数(k*n:从n倍的检索结果中取topk)的相关记录(含embedding),总检索条数=t*k*n +2. 对结果根据embedding进行k-means聚类,簇数为k*n(同条数/t) +3. 每簇取与检索最相似的代表当前簇,返回k个簇代表作为最终结果 +""" + +import numpy as np +from sklearn.cluster import KMeans +from typing import List, Dict, Tuple +import logging + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +class ClusteringSuppression: + """ + 聚类抑制机制 + 通过k-means聚类对相似记忆进行分组,从每组中选择最相关的代表 + """ + + def __init__(self, + cluster_multiplier: int = 3, + retrieval_multiplier: int = 2): + """ + 初始化聚类抑制机制 + + Args: + cluster_multiplier: 每个簇期望包含的记录数量,默认3条 + retrieval_multiplier: 检索结果倍数,默认2倍 + """ + self.cluster_multiplier = cluster_multiplier + self.retrieval_multiplier = retrieval_multiplier + logger.info(f"聚类抑制机制初始化: 每簇期望记录数={cluster_multiplier}, 检索倍数={retrieval_multiplier}") + + + + + + + + def calculate_retrieval_parameters(self, target_k: int) -> Tuple[int, int]: + """ + 计算检索参数 + + Args: + target_k: 目标返回条数 + + Returns: + (总检索条数, 聚类数) + """ + # 聚类数 = 目标条数 * 检索倍数 + cluster_count = target_k * self.retrieval_multiplier + + # 总检索条数 = 聚类数 * 每簇期望记录数 + total_retrieval = cluster_count * self.cluster_multiplier + + return total_retrieval, cluster_count + + def cluster_by_query_similarity_and_aggregate(self, + records: List[Dict], + embeddings_array: np.ndarray, + distances: List[float], + cluster_count: int, + target_k: int) -> List[Dict]: + """ + 基于查询相似度的聚类: + - 簇内选与查询distance最小的记录为代表; + - 代表记录的valid_access_count = 簇内所有记录的valid_access_count之和; + - 最终结果 = 分别按相关度与valid_access_count各取target_k条,按doc_id去重后返回合集。 + Args: + records: 与embeddings_array、distances一一对齐的记录列表(每条含embedding、distance、valid_access_count) + embeddings_array: 形如 (N, D) 的向量数组 + distances: 长度为 N 的距离列表(越小越相似) + cluster_count: 聚类簇数 + target_k: 返回前k条代表 + """ + import numpy as np + from sklearn.cluster import KMeans + + if not isinstance(cluster_count, int) or cluster_count < 1: + cluster_count = 1 + + n = len(records) + if n == 0: + return [] + + # 样本数 <= 聚类数:不聚类,直接在原集合上做双路topK并去重 + if n <= cluster_count: + base = [] + for i in range(n): + rep = dict(records[i]) + rep["cluster_size"] = 1 + base.append(rep) + # 分别取topK + by_rel = sorted(base, key=lambda x: float(x.get("distance", float("inf"))))[:target_k] + by_cnt = sorted(base, key=lambda x: float(x.get("valid_access_count", 0.0)), reverse=True)[:target_k] + # 合并去重(按doc_id) + seen = set() + merged = [] + for r in by_rel + by_cnt: + rid = r.get("doc_id") + if rid not in seen: + seen.add(rid) + merged.append(r) + return merged + + # KMeans 聚类 + kmeans = KMeans(n_clusters=cluster_count, random_state=42, n_init=10) + labels = kmeans.fit_predict(embeddings_array) + + # 簇代表选择与累计 + representatives = [] + for cid in np.unique(labels): + idx = np.where(labels == cid)[0] + if len(idx) == 0: + continue + + # 代表:簇内与查询distance最小 + local_dist = [(i, float(distances[i]) if distances[i] is not None else float("inf")) for i in idx] + rep_idx, _ = min(local_dist, key=lambda t: t[1]) + + # 累计簇内valid_access_count + sum_valid = float(sum(float(records[i].get("valid_access_count", 0.0)) for i in idx)) + + rep = dict(records[rep_idx]) + rep["valid_access_count"] = sum_valid + rep["cluster_size"] = len(idx) + representatives.append(rep) + + # 分别按相关度与valid_access_count取topK,然后合并去重 + top_by_relevance = sorted(representatives, key=lambda x: float(x.get("distance", float("inf"))))[:target_k] + top_by_count = sorted(representatives, key=lambda x: float(x.get("valid_access_count", 0.0)), reverse=True)[:target_k] + + seen_ids = set() + final_selection = [] + for r in top_by_relevance + top_by_count: + rid = r.get("doc_id") + if rid not in seen_ids: + seen_ids.add(rid) + final_selection.append(r) + + return final_selection \ No newline at end of file diff --git a/bionicmemory/algorithms/newton_cooling_helper.py b/bionicmemory/algorithms/newton_cooling_helper.py new file mode 100644 index 0000000..32e1c17 --- /dev/null +++ b/bionicmemory/algorithms/newton_cooling_helper.py @@ -0,0 +1,48 @@ +import math +from enum import Enum +from datetime import datetime + +class CoolingRate(Enum): + MINUTES_20 = (0.582, 20 * 60) + HOURS_1 = (0.442, 1 * 60 * 60) + HOURS_9 = (0.358, 9 * 60 * 60) + DAYS_1 = (0.337, 1 * 24 * 60 * 60) + DAYS_2 = (0.278, 2 * 24 * 60 * 60) + DAYS_6 = (0.254, 6 * 24 * 60 * 60) + DAYS_31 = (0.211, 31 * 24 * 60 * 60) + +class NewtonCoolingHelper: + @staticmethod + def calculate_cooling_rate(enum_value: CoolingRate) -> float: + """ + 根据枚举值计算冷却速率系数(alpha)。 + """ + final_temperature_ratio, time_interval = enum_value.value + return -math.log(final_temperature_ratio) / time_interval + + @staticmethod + def calculate_newton_cooling_effect(initial_temperature: float, time_interval: float, cooling_rate: float = None) -> float: + """ + 根据牛顿冷却定律计算当前时间的温度。 + """ + if cooling_rate is None: + cooling_rate = NewtonCoolingHelper.calculate_cooling_rate(CoolingRate.DAYS_31) + return initial_temperature * math.exp(-cooling_rate * time_interval) + + @staticmethod + def calculate_time_difference(update_time: datetime, current_time: datetime) -> float: + """ + 计算上次更新时间与当前时间之间的时间差。 + """ + if isinstance(update_time, str): + update_time = datetime.fromisoformat(update_time) + if isinstance(current_time, str): + current_time = datetime.fromisoformat(current_time) + time_delta = current_time - update_time + return time_delta.total_seconds() + + @staticmethod + def get_threshold(cooling_rate: CoolingRate=None) -> float: + if cooling_rate is None: + cooling_rate=CoolingRate.DAYS_31 + return cooling_rate.value[0] diff --git a/bionicmemory/api/__init__.py b/bionicmemory/api/__init__.py new file mode 100644 index 0000000..1baab55 --- /dev/null +++ b/bionicmemory/api/__init__.py @@ -0,0 +1,6 @@ +""" +API模块 + +包含仿生记忆系统的API接口: +- FastAPI代理服务器 +""" diff --git a/bionicmemory/api/proxy_server.py b/bionicmemory/api/proxy_server.py new file mode 100644 index 0000000..7e441b7 --- /dev/null +++ b/bionicmemory/api/proxy_server.py @@ -0,0 +1,511 @@ +""" +基于OpenAI官方库的代理服务器 +使用OpenAI官方客户端处理所有请求,确保完全兼容 +""" + +from contextlib import asynccontextmanager +import os +import json +import logging +import asyncio +from datetime import datetime +from typing import Dict, Any, List, Optional, Tuple +from fastapi import FastAPI, Request, Response, Depends +from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import JSONResponse, StreamingResponse +import uvicorn + +# OpenAI官方库 +from openai import OpenAI, AsyncOpenAI +from openai.types.chat import ChatCompletion, ChatCompletionChunk +from openai.types.embedding import Embedding + +# BionicMemory核心组件 +from bionicmemory.core.memory_system import LongShortTermMemorySystem, SourceType +from bionicmemory.services.memory_cleanup_scheduler import MemoryCleanupScheduler +from bionicmemory.core.chroma_service import ChromaService +from bionicmemory.algorithms.newton_cooling_helper import CoolingRate +from bionicmemory.services.local_embedding_service import get_embedding_service + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +# ========== 环境变量配置 ========== +# 禁用ChromaDB遥测 +os.environ["ANONYMIZED_TELEMETRY"] = "False" + +API_HOST = os.getenv("API_HOST", "0.0.0.0") +API_PORT = int(os.getenv("API_PORT", "8000")) +CHROMA_HOST = os.getenv("CHROMA_HOST", "localhost") +CHROMA_PORT = int(os.getenv("CHROMA_PORT", "8001")) +CHROMA_CLIENT_TYPE = os.getenv("CHROMA_CLIENT_TYPE", "persistent") + +# ========== OpenAI配置 ========== +OPENAI_API_KEY = os.getenv("OPENAI_API_KEY", "") +OPENAI_API_BASE = os.getenv("OPENAI_API_BASE", "https://api.deepseek.com") +OPENAI_MODEL_NAME = os.getenv("OPENAI_MODEL_NAME", "deepseek-chat") + +# 记忆系统配置 +SUMMARY_MAX_LENGTH = int(os.getenv('SUMMARY_MAX_LENGTH', '500')) +MAX_RETRIEVAL_RESULTS = int(os.getenv('MAX_RETRIEVAL_RESULTS', '7')) +CLUSTER_MULTIPLIER = int(os.getenv('CLUSTER_MULTIPLIER', '3')) +RETRIEVAL_MULTIPLIER = int(os.getenv('RETRIEVAL_MULTIPLIER', '2')) + +# ========== 工具函数 ========== + +def extract_user_message(messages: List[Dict]) -> Optional[str]: + """从消息列表中提取用户消息""" + for message in reversed(messages): # 从最新消息开始查找 + if message.get("role") == "user": + return message.get("content", "") + return None + +def extract_user_id_from_request(body_data: Dict) -> str: + """从OpenAI请求中提取用户ID""" + try: + logger.info("🔍 开始提取用户ID...") + + # 1. 优先从对话协议中的user字段提取 + if "user" in body_data: + raw_user = body_data["user"] + if isinstance(raw_user, str) and raw_user.strip(): + user_id = raw_user.strip() + logger.info(f"✅ 使用对话协议user字段: {user_id}") + return user_id + + # 2. 默认值:default_user + user_id = "default_user" + logger.info(f"✅ 使用默认用户ID: {user_id}") + return user_id + + except Exception as e: + logger.error(f"❌ 提取用户ID失败: {e}") + return "default_user" + +def enhance_chat_with_memory(body_data: Dict, user_id: str) -> Tuple[Dict, List[float]]: + """ + 使用记忆系统增强聊天请求 + + Args: + body_data: 请求体数据 + user_id: 用户ID + + Returns: + (增强后的body_data, enhanced_query_embedding) + """ + global memory_system + + if not memory_system: + logger.warning("⚠️ 记忆系统未初始化,跳过记忆增强") + return body_data, None + + try: + messages = body_data.get("messages", []) + if not messages: + return body_data, None + + # 提取用户消息 + user_message = extract_user_message(messages) + if not user_message: + return body_data, None + + # 使用记忆系统处理用户消息 + short_term_records, system_prompt, query_embedding = memory_system.process_user_message( + user_message, user_id + ) + + if short_term_records: + logger.info(f"🧠 找到 {len(short_term_records)} 条相关记忆") + logger.info(f"🧠 生成的系统提示语长度: {len(system_prompt)}") + + # 直接使用memory_system生成的系统提示语作为系统消息 + system_message = { + "role": "system", + "content": system_prompt + } + + # 在用户消息前插入系统消息 + enhanced_messages = [system_message] + (messages[-3:] if len(messages) > 3 else messages) + body_data["messages"] = enhanced_messages + + logger.info(f"🧠 记忆增强完成,消息数量: {len(messages)} -> {len(enhanced_messages)}") + logger.info(f"🧠 记忆增强完成,消息内容: {enhanced_messages}") + + return body_data, query_embedding + + except Exception as e: + logger.error(f"❌ 记忆增强失败: {e}") + return body_data, None + +async def process_ai_reply_async(response_content: str, user_id: str, current_user_content: str = None): + """异步处理AI回复(不阻塞响应性能)""" + global memory_system + + if not memory_system: + return + + try: + # 执行记忆系统处理(正确的业务逻辑顺序) + await memory_system.process_agent_reply_async(response_content, user_id, current_user_content) + + except Exception as e: + logger.error(f"❌ 异步处理AI回复失败: {e}") + +# ========== 全局变量 ========== +memory_system = None +memory_cleanup_scheduler = None +chroma_service = None + +# OpenAI客户端 +openai_client = None +async_openai_client = None + +# ========== 初始化函数 ========== + +def initialize_memory_system(): + """初始化记忆系统""" + global memory_system, memory_cleanup_scheduler, chroma_service + + try: + logger.info("正在初始化记忆系统...") + + # 初始化ChromaDB服务(只使用本地embedding) + chroma_service = ChromaService() + logger.info("ChromaDB服务初始化完成(本地embedding模式)") + + # 初始化记忆系统 + memory_system = LongShortTermMemorySystem( + chroma_service=chroma_service, + summary_threshold=SUMMARY_MAX_LENGTH, + max_retrieval_results=MAX_RETRIEVAL_RESULTS, + cluster_multiplier=CLUSTER_MULTIPLIER, + retrieval_multiplier=RETRIEVAL_MULTIPLIER, + ) + + # 启动时清空短期记忆库 + try: + # 清空短期记忆库 + short_term_deleted_ids = chroma_service.delete_documents( + memory_system.short_term_collection_name + ) + logger.info(f"启动清空短期记忆库,删除 {len(short_term_deleted_ids)} 条记录") + + except Exception as _e: + logger.warning("启动清空短期记忆库失败", exc_info=True) + + # 初始化清理调度器 + memory_cleanup_scheduler = MemoryCleanupScheduler(memory_system=memory_system) + memory_cleanup_scheduler.start() + + logger.info("记忆系统初始化完成") + return True + except Exception as e: + logger.error(f"记忆系统初始化失败: {str(e)}", exc_info=True) + return False + +def initialize_openai_clients(): + """初始化OpenAI客户端""" + global openai_client, async_openai_client + + try: + logger.info("正在初始化OpenAI客户端...") + + # 同步客户端 + openai_client = OpenAI( + api_key=OPENAI_API_KEY, + base_url=OPENAI_API_BASE + ) + + # 异步客户端 + async_openai_client = AsyncOpenAI( + api_key=OPENAI_API_KEY, + base_url=OPENAI_API_BASE + ) + + logger.info("OpenAI客户端初始化完成") + return True + except Exception as e: + logger.error(f"OpenAI客户端初始化失败: {e}") + return False + +# ========== 生命周期事件处理器 ========== +@asynccontextmanager +async def lifespan(app: FastAPI): + # 启动时初始化 + initialize_memory_system() + initialize_openai_clients() + yield + # 关闭时清理 + if memory_cleanup_scheduler: + memory_cleanup_scheduler.stop() + logger.info("记忆清理调度器已停止") + +# ========== FastAPI应用初始化 ========== +app = FastAPI(title="BionicMemory OpenAI Proxy", version="2.0.0", lifespan=lifespan) + +# 添加CORS中间件 +app.add_middleware( + CORSMiddleware, + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + +# ========== 健康检查端点 ========== +@app.get("/health") +async def health_check(): + return { + "status": "healthy", + "service": "BionicMemory OpenAI Proxy", + "timestamp": datetime.now().isoformat(), + "memory_system_initialized": memory_system is not None, + "openai_client_initialized": openai_client is not None, + "cleanup_scheduler_running": memory_cleanup_scheduler is not None if memory_cleanup_scheduler else False + } + +# ========== 主要路由处理 ========== +@app.api_route("/v1/{path:path}", methods=["POST", "GET"]) +async def proxy(request: Request, path: str): + """ + 代理所有 /v1/* 请求 + 使用OpenAI官方库处理,确保完全兼容 + """ + body = await request.body() + + # 记录基本请求信息 + logger.info(f"📥 收到请求: {request.method} /v1/{path}") + + # ========== 路由处理 ========== + if path.startswith("embeddings"): + # Embedding API - 使用本地embedding服务 + return await handle_embedding_request(request, path, body) + + elif path == "chat/completions": + # Chat Completions API - 使用OpenAI客户端 + 记忆增强 + return await handle_chat_request(request, path, body) + + else: + # 其他 API - 使用OpenAI客户端透传 + return await handle_other_request(request, path, body) + +# ========== 处理函数 ========== + +async def handle_embedding_request(request: Request, path: str, body: bytes): + """处理embedding请求 - 使用本地embedding服务""" + try: + # 解析请求体 + if body: + body_data = json.loads(body) + input_text = body_data.get("input", "") + model = body_data.get("model", "") + + # 使用本地embedding服务 + logger.info("使用本地embedding服务") + embedding_service = get_embedding_service() + embeddings = embedding_service.get_embeddings([input_text]) + + # 构造OpenAI兼容的响应 + response_data = { + "object": "list", + "data": [{ + "object": "embedding", + "index": 0, + "embedding": embeddings[0] + }], + "model": model, + "usage": { + "prompt_tokens": len(input_text.split()), + "total_tokens": len(input_text.split()) + } + } + + return JSONResponse(content=response_data) + + except Exception as e: + logger.error(f"❌ 处理embedding请求失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"处理embedding请求失败: {str(e)}"} + ) + +async def handle_chat_request(request: Request, path: str, body: bytes): + """处理对话请求 - 使用OpenAI客户端 + 记忆增强""" + try: + # 解析请求体 + body_data = None + user_id = None + enhanced_query_embedding = None + current_user_content = None + + if body: + body_data = json.loads(body) + # 提取用户ID + user_id = extract_user_id_from_request(body_data) + + # 替换模型名称 + if "model" in body_data: + body_data["model"] = OPENAI_MODEL_NAME + + # 记忆增强处理 + enhanced_body_data, query_embedding = enhance_chat_with_memory(body_data, user_id) + current_user_content = body_data.get("messages", [])[-1].get("content", "") + body_data = enhanced_body_data + + # 检查是否为流式响应 + is_stream = body_data and body_data.get("stream", False) if body_data else False + + if is_stream: + # 流式响应 - 使用异步OpenAI客户端 + logger.info("🌊 处理流式响应(使用OpenAI客户端)") + + try: + # 使用OpenAI客户端创建流式响应 + stream = await async_openai_client.chat.completions.create( + model=body_data.get("model", OPENAI_MODEL_NAME), + messages=body_data.get("messages", []), + stream=True, + **{k: v for k, v in body_data.items() + if k not in ["model", "messages", "stream"]} + ) + + async def openai_stream_wrapper(): + full_content = "" + async for chunk in stream: + # 使用OpenAI原生格式 + chunk_data = chunk.model_dump() + content = chunk_data.get('choices', [{}])[0].get('delta', {}).get('content', '') + if content: + full_content += content + + # 转换为SSE格式 + yield f"data: {json.dumps(chunk_data)}\n\n" + + # 流式结束后异步存储记忆 + if full_content and body_data: + asyncio.create_task(process_ai_reply_async( + full_content, user_id, current_user_content + )) + + yield "data: [DONE]\n\n" + + return StreamingResponse( + openai_stream_wrapper(), + status_code=200, + headers={ + "Content-Type": "text/plain; charset=utf-8", + "Cache-Control": "no-cache", + "Connection": "keep-alive" + } + ) + + except Exception as e: + logger.error(f"❌ OpenAI流式处理失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"流式处理失败: {str(e)}"} + ) + else: + # 非流式响应 - 使用同步OpenAI客户端 + logger.info("📝 处理非流式响应(使用OpenAI客户端)") + + try: + response = openai_client.chat.completions.create( + model=body_data.get("model", OPENAI_MODEL_NAME), + messages=body_data.get("messages", []), + **{k: v for k, v in body_data.items() + if k not in ["model", "messages"]} + ) + + # 异步存储记忆 + if response.choices[0].message.content and body_data: + asyncio.create_task(process_ai_reply_async( + response.choices[0].message.content, + user_id, + current_user_content + )) + + # 返回OpenAI原生响应 + return JSONResponse(content=response.model_dump()) + + except Exception as e: + logger.error(f"❌ OpenAI非流式处理失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"非流式处理失败: {str(e)}"} + ) + + except Exception as e: + logger.error(f"❌ 处理对话请求失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"处理对话请求失败: {str(e)}"} + ) + +async def handle_other_request(request: Request, path: str, body: bytes): + """处理其他API - 使用OpenAI客户端透传""" + try: + # 解析请求体 + body_data = json.loads(body) if body else {} + + # 使用OpenAI客户端处理其他请求 + logger.info(f"🔄 处理其他请求: {path}") + + # 根据路径选择处理方法 + if path == "models": + # 模型列表请求 + models_response = { + "object": "list", + "data": [ + { + "id": OPENAI_MODEL_NAME, + "object": "model", + "created": int(datetime.now().timestamp()), + "owned_by": "bionicmemory" + } + ] + } + return JSONResponse(content=models_response) + + else: + # 其他请求透传 + try: + # 使用OpenAI客户端处理 + if request.method == "GET": + # GET请求处理 + response = openai_client._client.get(f"/v1/{path}") + return JSONResponse(content=response.json()) + else: + # POST请求处理 + response = openai_client._client.post( + f"/v1/{path}", + json=body_data, + headers={"Authorization": f"Bearer {OPENAI_API_KEY}"} + ) + return JSONResponse(content=response.json()) + + except Exception as e: + logger.error(f"❌ OpenAI客户端处理其他请求失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"处理请求失败: {str(e)}"} + ) + + except Exception as e: + logger.error(f"❌ 处理其他请求失败: {e}") + return JSONResponse( + status_code=500, + content={"error": f"处理其他请求失败: {str(e)}"} + ) + +# ========== 启动配置 ========== +if __name__ == "__main__": + uvicorn.run( + "bionicmemory.api.proxy_server_openai:app", + host=API_HOST, + port=API_PORT, + log_level="info", + access_log=True, + reload=False + ) diff --git a/bionicmemory/core/__init__.py b/bionicmemory/core/__init__.py new file mode 100644 index 0000000..54fd39c --- /dev/null +++ b/bionicmemory/core/__init__.py @@ -0,0 +1,7 @@ +""" +核心模块 + +包含仿生记忆系统的核心功能: +- 长短期记忆系统 +- ChromaDB服务封装 +""" diff --git a/bionicmemory/core/chroma_service.py b/bionicmemory/core/chroma_service.py new file mode 100644 index 0000000..a49e03a --- /dev/null +++ b/bionicmemory/core/chroma_service.py @@ -0,0 +1,552 @@ +import chromadb +from chromadb import Documents, EmbeddingFunction, Embeddings +from typing import Optional, List, Dict, Any, Union, Callable +import json +import logging +import os +from dotenv import load_dotenv +from bionicmemory.services.chat_helper import ChatHelper + +# 加载.env文件 +load_dotenv() + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +# 在文件顶部添加导入 +from bionicmemory.services.local_embedding_service import get_embedding_service + + +class ChromaService: + """ + ChromaDB向量数据库操作服务 + """ + + def __init__(self, + client_type: str = None, + path: Optional[str] = None, + host: str = None, + port: int = None, + chat_api_key: str = None, + chat_base_url: str = None): + """ + 初始化ChromaDB服务 + + Args: + client_type (str): 客户端类型,支持 'persistent', 'ephemeral', 'http' + path (str): 持久化存储路径(仅persistent模式) + host (str): 服务器地址(仅http模式) + port (int): 服务器端口(仅http模式) + chat_api_key (str): 聊天API密钥 + chat_base_url (str): 聊天API基础URL + """ + try: + # 从环境变量读取配置 + from dotenv import load_dotenv + import os + + # 加载.env文件 + load_dotenv() + + # 设置默认值 + client_type = client_type or os.getenv('CHROMA_CLIENT_TYPE', 'persistent') + path = path or os.getenv('CHROMA_PATH', './memory/chroma_db') + path = os.path.abspath(path) # 转换为绝对路径 + host = host or os.getenv('CHROMA_HOST', 'localhost') + port = int(port or os.getenv('CHROMA_PORT', '8001')) + chat_api_key = chat_api_key or os.getenv('OPENAI_API_KEY') + chat_base_url = chat_base_url or os.getenv('OPENAI_API_BASE') + + # 初始化ChromaDB客户端 + if client_type == "persistent": + self.client = chromadb.PersistentClient(path=path) + elif client_type == "ephemeral": + self.client = chromadb.EphemeralClient() + elif client_type == "http": + self.client = chromadb.HttpClient(host=host, port=port) + else: + raise ValueError(f"不支持的客户端类型: {client_type}") + + # 初始化聊天助手(如果需要) + if chat_api_key and chat_base_url: + self.chat_helper = ChatHelper(chat_api_key, chat_base_url) + logger.info("聊天助手初始化完成") + else: + self.chat_helper = None + logger.info("未配置聊天API,聊天功能不可用") + + # 初始化本地embedding服务 + self.local_embedding_service = get_embedding_service() + logger.info("使用本地embedding服务") + + # 初始化自定义embedding函数相关变量 + self._custom_embedding_func = None + self._embedding_function = None # 本地模式不需要embedding函数 + + except Exception as e: + raise Exception(f"初始化ChromaDB客户端失败: {str(e)}") + + def create_collection(self, name: str, metadata: Optional[Dict[str, Any]] = None): + """ + 创建新的集合 + + Args: + name (str): 集合名称 + metadata (Dict[str, Any], optional): 集合元数据 + + Returns: + Collection: 集合对象 + """ + try: + # 本地embedding模式,不使用ChromaDB的embedding函数 + embedding_function = None + + collection = self.client.create_collection( + name=name, + metadata=metadata, + embedding_function=embedding_function + ) + logger.info(f"成功创建集合: {name}") + return collection + except Exception as e: + logger.error(f"创建集合失败: {name}, 错误: {e}") + raise + + def get_or_create_collection(self, name: str, metadata: Optional[Dict[str, Any]] = None): + """ + 获取或创建集合 + + Args: + name (str): 集合名称 + metadata (Dict[str, Any], optional): 集合元数据 + + Returns: + Collection: 集合对象 + """ + try: + embedding_function = None + if self._custom_embedding_func is not None: + self._embedding_function.custom_func = self._custom_embedding_func + embedding_function = self._embedding_function + + collection = self.client.get_or_create_collection( + name=name, + metadata=metadata, + embedding_function=embedding_function + ) + logger.info(f"成功获取或创建集合: {name}") + return collection + except Exception as e: + logger.error(f"获取或创建集合失败: {name}, 错误: {e}") + raise + + def list_collections(self): + """ + 列出所有集合 + + Returns: + List[Collection]: 集合对象列表 + """ + try: + collections = self.client.list_collections() + logger.info(f"找到 {len(collections)} 个集合") + return collections + except Exception as e: + logger.error(f"获取集合列表失败: {e}") + raise + + def delete_collection(self, name: str): + """ + 删除集合 + + Args: + name (str): 集合名称 + + Returns: + None + """ + try: + self.client.delete_collection(name=name) + logger.info(f"成功删除集合: {name}") + except Exception as e: + logger.error(f"删除集合失败: {name}, 错误: {e}") + raise + + def add_documents(self, + collection_name: str, + documents: List[str], + embeddings: List[List[float]] = None, + ids: Optional[List[str]] = None, + metadatas: Optional[List[Dict[str, Any]]] = None) -> List[str]: + """ + 向集合添加文档 + + Args: + collection_name (str): 集合名称 + documents (List[str]): 文档内容列表 + embeddings (List[List[float]], optional): 预计算的embedding向量列表 + ids (List[str], optional): 文档ID列表 + metadatas (List[Dict[str, Any]], optional): 文档元数据列表 + + Returns: + List[str]: 添加的文档ID列表 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + + # 如果没有提供ID,自动生成 + if ids is None: + ids = [f"doc_{i}" for i in range(len(documents))] + + # 如果提供了预计算的embedding,使用它们 + if embeddings is not None: + # 验证参数长度一致性 + if len(documents) != len(embeddings): + raise ValueError(f"文档数量({len(documents)})与embedding数量({len(embeddings)})不匹配") + + collection.add( + documents=documents, + embeddings=embeddings, + ids=ids, + metadatas=metadatas + ) + else: + # 让ChromaDB自动生成embedding + collection.add( + documents=documents, + ids=ids, + metadatas=metadatas + ) + + return ids # ✅ 返回实际数据 + except Exception as e: + logger.error(f"添加文档失败: {e}") + raise # ✅ 抛出异常 + + def query_documents(self, + collection_name: str, + query_texts: List[str] = None, + query_embeddings: List[List[float]] = None, + n_results: int = 10, + where: Optional[Dict[str, Any]] = None, + include: Optional[List[str]] = None) -> Dict: + """ + 查询文档 + + Args: + collection_name (str): 集合名称 + query_texts (List[str], optional): 查询文本列表 + query_embeddings (List[List[float]], optional): 预计算的查询embedding列表 + n_results (int): 返回结果数量 + where (Dict[str, Any], optional): 元数据过滤条件 + include (List[str], optional): 需要返回的数据类型 + + Returns: + Dict: 查询结果字典 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + + # 设置默认的include参数 + if include is None: + include = ["documents", "metadatas", "distances", "embeddings"] + + # 优先使用预计算的embedding,避免重复计算 + if query_embeddings is not None: + results = collection.query( + query_embeddings=query_embeddings, + n_results=n_results, + where=where, + include=include + ) + else: + results = collection.query( + query_texts=query_texts, + n_results=n_results, + where=where, + include=include + ) + + # 统一处理embeddings,确保返回list格式 + if 'embeddings' in results and results.get('embeddings') is not None: + embeddings_data = results['embeddings'] + processed_embeddings = [] + for embedding_list in embeddings_data: + processed_embedding_list = [] + for embedding in embedding_list: + if embedding is not None and hasattr(embedding, 'tolist'): + processed_embedding_list.append(embedding.tolist()) + else: + processed_embedding_list.append(embedding) + processed_embeddings.append(processed_embedding_list) + results['embeddings'] = processed_embeddings + + return results # ✅ 返回实际数据 + except Exception as e: + logger.error(f"查询文档失败: {e}") + raise # ✅ 抛出异常 + + def get_documents(self, + collection_name: str, + ids: Optional[List[str]] = None, + limit: Optional[int] = None, + where: Optional[Dict[str, Any]] = None, + include: Optional[List[str]] = None) -> Dict: + """ + 获取文档 + + Args: + collection_name (str): 集合名称 + ids (List[str], optional): 文档ID列表 + limit (int, optional): 限制返回数量 + where (Dict[str, Any], optional): 元数据过滤条件 + include (List[str], optional): 需要返回的数据类型 + + Returns: + Dict: 文档结果字典 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + + # 设置默认的include参数 + if include is None: + include = ["documents", "metadatas"] + + results = collection.get( + ids=ids, + limit=limit, + where=where, + include=include + ) + + # 统一处理embeddings,确保返回list格式 + if 'embeddings' in results and results.get('embeddings') is not None: + embeddings_data = results['embeddings'] + processed_embeddings = [] + for embedding_list in embeddings_data: + processed_embedding_list = [] + for embedding in embedding_list: + if embedding is not None and hasattr(embedding, 'tolist'): + processed_embedding_list.append(embedding.tolist()) + else: + processed_embedding_list.append(embedding) + processed_embeddings.append(processed_embedding_list) + results['embeddings'] = processed_embeddings + + return results # ✅ 返回实际数据 + except Exception as e: + logger.error(f"获取文档失败: {e}") + raise # ✅ 抛出异常 + + def update_documents(self, + collection_name: str, + ids: List[str], + documents: Optional[List[str]] = None, + metadatas: Optional[List[Dict[str, Any]]] = None) -> Dict: + """ + 更新文档 + + Args: + collection_name (str): 集合名称 + ids (List[str]): 文档ID列表 + documents (List[str], optional): 新的文档内容 + metadatas (List[Dict[str, Any]], optional): 新的元数据 + + Returns: + Dict: 更新后的文档数据 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + + collection.update( + ids=ids, + documents=documents, + metadatas=metadatas + ) + + # 返回更新后的文档数据 + return collection.get(ids=ids) # ✅ 返回实际数据 + except Exception as e: + logger.error(f"更新文档失败: {e}") + raise # ✅ 抛出异常 + + def delete_documents(self, + collection_name: str, + ids: Optional[List[str]] = None, + where: Optional[Dict[str, Any]] = None) -> List[str]: + """ + 删除文档 + + Args: + collection_name (str): 集合名称 + ids (List[str], optional): 文档ID列表 + where (Dict[str, Any], optional): 元数据过滤条件 + + Returns: + List[str]: 删除的文档ID列表 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + + # 如果提供了ids,直接删除 + if ids: + collection.delete(ids=ids) + return ids # ✅ 返回实际数据 + else: + # 如果使用where条件,先查询要删除的文档 + if where: + results = collection.get(where=where) + deleted_ids = results.get('ids', []) + if deleted_ids: + collection.delete(ids=deleted_ids) + return deleted_ids # ✅ 返回实际数据 + else: + # 删除所有文档 + all_results = collection.get() + all_ids = all_results.get('ids', []) + if all_ids: + collection.delete(ids=all_ids) + return all_ids # ✅ 返回实际数据 + + except Exception as e: + logger.error(f"删除文档失败: {e}") + raise # ✅ 抛出异常 + + def count_documents(self, collection_name: str) -> int: + """ + 统计集合中的文档数量 + + Args: + collection_name (str): 集合名称 + + Returns: + int: 文档数量 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + count = collection.count() + return count # ✅ 返回实际数据 + except Exception as e: + logger.error(f"统计文档数量失败: {e}") + raise # ✅ 抛出异常 + + def peek_documents(self, collection_name: str, limit: int = 10) -> Dict: + """ + 预览集合中的文档 + + Args: + collection_name (str): 集合名称 + limit (int): 预览数量限制 + + Returns: + Dict: 预览结果数据 + """ + try: + # 使用self.client确保集合存在 + collection = self.client.get_or_create_collection( + name=collection_name, + embedding_function=self._embedding_function + ) + results = collection.peek(limit=limit) + return results # ✅ 返回实际数据 + except Exception as e: + logger.error(f"预览文档失败: {e}") + raise # ✅ 抛出异常 + + def custom_embedding(self, texts: List[str]) -> List[List[float]]: + """ + 自定义嵌入函数(预留接口) + + Args: + texts (List[str]): 待嵌入的文本列表 + + Returns: + List[List[float]]: 嵌入向量列表 + """ + # 函数体为pass,后续手动实现 + pass + + def set_custom_embedding_function(self, embedding_func: Callable[[List[str]], List[List[float]]]) -> None: + """ + 设置自定义嵌入函数 + + Args: + embedding_func: 自定义嵌入函数,接受文本列表,返回向量列表 + + Returns: + None + """ + try: + self._custom_embedding_func = embedding_func + # ✅ 不返回值,成功就成功 + except Exception as e: + logger.error(f"设置自定义嵌入函数失败: {e}") + raise # ✅ 抛出异常 + + def get_custom_embedding_function(self) -> Optional[Callable]: + """ + 获取当前设置的自定义嵌入函数 + + Returns: + Optional[Callable]: 当前的自定义嵌入函数,如果未设置则返回None + """ + return self._custom_embedding_func + + def create_embeddings(self, texts: List[str], model: str = None) -> List[List[float]]: + """ + 使用本地服务生成文本的embedding向量 + """ + # 使用本地embedding服务 + embeddings = self.local_embedding_service.encode_texts(texts) + return embeddings.tolist() + + def get_embedding_dimension(self) -> int: + """ + 获取embedding维度(从embedding服务动态获取) + """ + # 从 embedding 服务获取实际维度 + model_info = self.local_embedding_service.get_model_info() + return model_info.get('embedding_dim', 1024) + + def get_collection(self, name: str): + """ + 获取集合对象 + + Args: + name (str): 集合名称 + + Returns: + Collection: 集合对象 + """ + try: + collection = self.client.get_collection(name) + logger.info(f"成功获取集合: {name}") + return collection + except Exception as e: + logger.error(f"获取集合失败: {name}, 错误: {e}") + raise diff --git a/bionicmemory/core/memory_system.py b/bionicmemory/core/memory_system.py new file mode 100644 index 0000000..b469a3c --- /dev/null +++ b/bionicmemory/core/memory_system.py @@ -0,0 +1,1488 @@ +""" +长短期记忆系统 +基于 ChromaDB 和牛顿冷却遗忘算法实现 +""" + +import hashlib +import logging +import numpy as np +from datetime import datetime +from enum import Enum +from typing import List, Dict, Optional, Tuple, Any +from dataclasses import dataclass + + + +from bionicmemory.algorithms.newton_cooling_helper import NewtonCoolingHelper, CoolingRate +from bionicmemory.core.chroma_service import ChromaService +from bionicmemory.services.summary_service import SummaryService +from bionicmemory.algorithms.clustering_suppression import ClusteringSuppression +from bionicmemory.services.local_embedding_service import get_embedding_service + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +class SourceType(Enum): + """消息来源类型枚举""" + USER = "user" # 用户发送的消息 + AGENT = "agent" # 大模型/AI代理的回复 + OTHER = "other" # 其他来源(如系统消息、第三方API等) + +@dataclass +class MemoryRecord: + """记忆记录数据结构""" + content: str + valid_access_count: float + last_updated: str + created_at: str + total_access_count: int + source_type: str + user_id: str + +class LongShortTermMemorySystem: + """ + 长短期记忆系统 + 实现基于牛顿冷却遗忘算法的记忆管理 + """ + + def __init__(self, + chroma_service: ChromaService, + summary_threshold: int = 500, + max_retrieval_results: int = 10, + cluster_multiplier: int = 3, + retrieval_multiplier: int = 2): + """ + 初始化长短期记忆系统 + + Args: + chroma_service: ChromaDB服务实例 + summary_threshold: 摘要长度阈值(默认500) + max_retrieval_results: 最大检索结果数量(默认10) + cluster_multiplier: 聚类倍数(默认3) + retrieval_multiplier: 检索倍数(默认2) + """ + self.chroma_service = chroma_service + self.max_retrieval_results = max_retrieval_results + self.cluster_multiplier = cluster_multiplier + self.retrieval_multiplier = retrieval_multiplier + self.summary_threshold = summary_threshold + + # 牛顿冷却助手 + self.newton_helper = NewtonCoolingHelper() + + # 摘要服务 + try: + self.summary_service = SummaryService() + logger.info("摘要服务初始化成功") + except Exception as e: + logger.warning(f"摘要服务初始化失败,将使用简单截断: {e}") + self.summary_service = None + + # 遗忘阈值(从科学数据读取) + self.long_term_threshold = self.newton_helper.get_threshold(CoolingRate.DAYS_31) + self.short_term_threshold = self.newton_helper.get_threshold(CoolingRate.MINUTES_20) + + # 集合名称 + self.long_term_collection_name = "long_term_memory" + self.short_term_collection_name = "short_term_memory" + + # 初始化集合 + self._initialize_collections() + + + # 初始化本地embedding服务 + self.embedding_service = get_embedding_service() + logger.info("记忆系统使用本地embedding服务") + + logger.info(f"长短期记忆系统初始化完成") + logger.info(f"摘要阈值: {self.summary_threshold}") + logger.info(f"最大检索结果数量: {self.max_retrieval_results}") + logger.info(f"聚类倍数: {self.cluster_multiplier}") + logger.info(f"检索倍数: {self.retrieval_multiplier}") + logger.info(f"长期记忆阈值: {self.long_term_threshold}") + logger.info(f"短期记忆阈值: {self.short_term_threshold}") + + def _initialize_collections(self): + """初始化长短期记忆集合""" + try: + # 确保长期记忆集合存在 + self.chroma_service.get_or_create_collection( + self.long_term_collection_name + ) + + # 确保短期记忆集合存在 + self.chroma_service.get_or_create_collection( + self.short_term_collection_name + ) + + logger.info("长短期记忆集合初始化成功") + except Exception as e: + logger.error(f"初始化集合失败: {e}") + raise + + + + + def _generate_md5(self, content: str, user_id: str ) -> str: + """生成多租户隔离的MD5""" + uid = (user_id or "").strip() + key = f"{uid}::{content}" + return hashlib.md5(key.encode('utf-8')).hexdigest() + + def _validate_user_access(self, record_user_id: str, requesting_user_id: str, operation: str) -> bool: + """ + 验证用户访问权限 + + Args: + record_user_id: 记录所属的用户ID + requesting_user_id: 请求操作的用户ID + operation: 操作类型描述 + + Returns: + 是否有权限访问 + """ + rid = record_user_id.strip() if isinstance(record_user_id, str) else record_user_id + qid = requesting_user_id.strip() if isinstance(requesting_user_id, str) else requesting_user_id + if rid != qid: + logger.warning(f"用户 {requesting_user_id} 尝试{operation}用户 {record_user_id} 的记录,拒绝访问") + return False + return True + + + def _generate_summary(self, content: str) -> str: + """ + 生成内容摘要 + 优先使用LLM生成摘要,失败时降级到简单截断 + """ + if len(content) <= self.summary_threshold: + return content + + # 如果有摘要服务,尝试使用LLM生成摘要 + if self.summary_service: + try: + summary = self.summary_service.generate_summary(content, self.summary_threshold) + if summary and len(summary) <= self.summary_threshold: + logger.info(f"LLM摘要生成成功: {len(content)} -> {len(summary)} 字符") + return summary + else: + logger.warning("LLM生成的摘要长度超出阈值,使用降级方案") + except Exception as e: + logger.warning(f"LLM摘要生成失败,使用降级方案: {e}") + + # 降级方案:简单截断 + logger.warning("使用降级摘要方案:简单截断") + summary = content[:self.summary_threshold] + if len(content) > self.summary_threshold: + summary += "..." + + return summary + + def _prepare_document_data(self, + content: str, + source_type: SourceType, + user_id: str) -> Tuple[str, str, Dict, List[float]]: + """ + 准备文档数据 - 优化版本 + + Returns: + (document_text, doc_id, metadata, embedding) + """ + logger.info(f"[仿生记忆] _prepare_document_data开始: content={content[:50]}...") + + if isinstance(content, list): + content = "\n".join(content) + + # 生成MD5作为文档ID + doc_id = self._generate_md5(content, user_id) + logger.info(f"[仿生记忆] _prepare_document_data: doc_id={doc_id}") + + # 检查是否已存在相同的文档(避免重复处理) + logger.info("[仿生记忆] _prepare_document_data: 检查是否已存在文档") + existing_result = self.chroma_service.get_documents( + self.long_term_collection_name, + ids=[doc_id], + include=["embeddings", "metadatas", "documents"] + ) + logger.info(f"[仿生记忆] _prepare_document_data: existing_result类型={type(existing_result)}") + + if existing_result and existing_result.get("metadatas"): + logger.info("[仿生记忆] _prepare_document_data: 文档已存在,返回现有数据") + # 文档已存在,直接返回现有数据 + metadata = existing_result["metadatas"][0] + document_text = existing_result["documents"][0] + + # 获取现有embedding(如果有的话) + embeddings = existing_result.get("embeddings", []) + logger.info(f"[仿生记忆] _prepare_document_data: embeddings类型={type(embeddings)}, 长度={len(embeddings) if embeddings else 0}") + raw_embedding = embeddings[0] if embeddings else None + # 确保embedding是list格式 + if raw_embedding is not None and hasattr(raw_embedding, 'tolist'): + embedding = raw_embedding.tolist() + else: + embedding = raw_embedding + logger.info(f"[仿生记忆] _prepare_document_data: embedding类型={type(embedding)}") + + logger.debug(f"文档 {doc_id} 已存在,跳过重复处理") + return document_text, doc_id, metadata, embedding + + logger.info("[仿生记忆] _prepare_document_data: 文档不存在,生成新数据") + # 决定用于embedding的文本 + document_text = self._generate_summary(content) + logger.info(f"[仿生记忆] _prepare_document_data: document_text={document_text[:50]}...") + + # 生成embedding并保存,避免重复计算 + try: + logger.info("[仿生记忆] _prepare_document_data: 开始生成embedding") + embedding = self.embedding_service.encode_text(document_text) + logger.info(f"[仿生记忆] _prepare_document_data: embedding生成完成, 类型={type(embedding)}") + except Exception as e: + logger.error(f"生成embedding失败: {e}") + embedding = [] + + # 准备元数据 + current_time = datetime.now().isoformat() + metadata = { + "content": content, + "valid_access_count": 1.0, + "last_updated": current_time, + "created_at": current_time, + "total_access_count": 1, + "source_type": source_type.value, + "user_id": user_id + } + + return document_text, doc_id, metadata, embedding + + def _calculate_decayed_valid_count(self, + record: Dict, + cooling_rate: CoolingRate) -> float: + """ + 计算衰减后的有效访问次数 + + Args: + record: 记录元数据 + cooling_rate: 遗忘速率 + + Returns: + 衰减后的有效访问次数 + """ + try: + last_updated = record.get("last_updated") + if not last_updated: + return record.get("valid_access_count", 1.0) + + # 计算时间差 + time_diff = self.newton_helper.calculate_time_difference( + last_updated, datetime.now() + ) + + # 计算冷却系数 + cooling_coefficient = self.newton_helper.calculate_cooling_rate(cooling_rate) + + # 计算衰减后的值 + initial_strength = record.get("valid_access_count", 1.0) + decayed_value = self.newton_helper.calculate_newton_cooling_effect( + initial_strength, time_diff, cooling_coefficient + ) + + return decayed_value + + except Exception as e: + logger.error(f"计算衰减值失败: {e}") + return record.get("valid_access_count", 1.0) + + def _update_record_access_count(self, + collection_name: str, + doc_id: str, + cooling_rate: CoolingRate, + user_id: str) -> bool: + """ + 更新记录的访问次数 + + Args: + collection_name: 集合名称 + doc_id: 文档ID + cooling_rate: 遗忘速率 + user_id: 用户ID(用于安全检查) + + Returns: + 是否更新成功 + """ + try: + # 获取记录 + result = self.chroma_service.get_documents(collection_name, ids=[doc_id]) + if not result or not result.get("metadatas"): + logger.warning(f"记录不存在: {doc_id}") + return False + + metadata = result["metadatas"][0] + + # 🔒 安全检查:确保只能更新自己的记录 + record_user_id = metadata.get("user_id") + if not self._validate_user_access(record_user_id, user_id, "更新"): + return False + + # 计算衰减后的值 + decayed_value = self._calculate_decayed_valid_count(metadata, cooling_rate) + + # 新的有效访问次数 = 衰减值 + 1 + new_valid_count = decayed_value + 1.0 + + # 更新元数据 + updated_metadata = metadata.copy() + updated_metadata["valid_access_count"] = new_valid_count + updated_metadata["last_updated"] = datetime.now().isoformat() + updated_metadata["total_access_count"] = metadata.get("total_access_count", 0) + 1 + + # 更新记录 + self.chroma_service.update_documents( + collection_name, + ids=[doc_id], + metadatas=[updated_metadata] + ) + + logger.debug(f"更新记录访问次数成功: {doc_id}, 新值: {new_valid_count}") + return True + + except Exception as e: + logger.error(f"更新记录访问次数失败: {e}") + return False + + def add_to_long_term_memory(self, + content: str, + source_type: SourceType, + user_id: str, + prepared_data: Tuple[str, str, Dict, List[float]] = None) -> str: + """ + 添加内容到长期记忆库 + + Args: + content: 内容 + source_type: 来源类型 + user_id: 用户ID + prepared_data: _prepare_document_data准备好的完整数据 (document_text, doc_id, metadata, embedding) + + Returns: + 文档ID + """ + try: + if prepared_data is not None: + # 使用_prepare_document_data准备好的完整数据,避免重复计算 + document_text, doc_id, metadata, embedding = prepared_data + else: + # 降级:重新调用_prepare_document_data + document_text, doc_id, metadata, embedding = self._prepare_document_data( + content, source_type, user_id + ) + + # 检查是否已存在 + existing_result = self.chroma_service.get_documents( + self.long_term_collection_name, ids=[doc_id] + ) + + if existing_result and existing_result.get("metadatas"): + # 记录已存在,更新访问次数 + logger.info(f"长期记忆记录已存在,更新访问次数: {doc_id}") + self._update_record_access_count( + self.long_term_collection_name, doc_id, CoolingRate.DAYS_31, user_id + ) + else: + # 新增记录,使用预计算的embedding + logger.info(f"新增长期记忆记录: {doc_id}") + + # 修复numpy数组长度判断问题 + if embedding is not None: + # 确保embedding是list格式 + if hasattr(embedding, 'tolist'): + embedding_list = embedding.tolist() + else: + embedding_list = embedding + # 检查长度 + if len(embedding_list) > 0: + embeddings_param = [embedding_list] + else: + embeddings_param = None + else: + embeddings_param = None + + self.chroma_service.add_documents( + self.long_term_collection_name, + documents=[document_text], + embeddings=embeddings_param, + metadatas=[metadata], + ids=[doc_id] + ) + + return doc_id + + except Exception as e: + logger.error(f"添加到长期记忆失败: {e}") + raise + + def _get_record_from_collection(self, collection_name: str, doc_id: str) -> Dict: + """ + 从指定集合获取记录 + + Args: + collection_name: 集合名称 + doc_id: 文档ID + + Returns: + 记录字典,包含完整数据 + """ + try: + result = self.chroma_service.get_documents(collection_name, ids=[doc_id]) + if not result or not result.get("metadatas"): + logger.warning(f"记录不存在: {doc_id} in {collection_name}") + return None + + metadata = result["metadatas"][0] + document = result["documents"][0] if result.get("documents") else "" + embedding = result["embeddings"][0] if result.get("embeddings") else None + + record = { + "doc_id": doc_id, + "content": metadata.get("content", ""), + "summary_document": document, + "valid_access_count": metadata.get("valid_access_count", 1.0), + "last_updated": metadata.get("last_updated", ""), + "source_type": metadata.get("source_type", ""), + "user_id": metadata.get("user_id", ""), + "embedding": embedding + } + + return record + + except Exception as e: + logger.error(f"从集合获取记录失败: {e}") + return None + + def retrieve_from_long_term_memory(self, + query: str, + user_id: str = None, + include: Optional[List[str]] = None, + query_embedding: List[float] = None) -> List[Dict]: + """ + 从长期记忆库检索相关记录(使用聚类抑制机制) + + Args: + query: 查询内容 + user_id: 用户ID(可选过滤) + include: 需要返回的数据类型列表,可选值: + - "documents": 文档内容(摘要) + - "metadatas": 元数据 + - "distances": 距离值 + - "embeddings": 向量嵌入 + 默认返回 ["documents", "metadatas", "distances", "embeddings"] + + Returns: + 经过聚类抑制后的相关记录列表 + """ + try: + # 构建查询条件 + where = {} + if user_id: + where["user_id"] = {"$eq": user_id} + + # 设置默认的include参数(需要包含embeddings与distances以便聚类抑制) + if include is None: + include = ["documents", "metadatas", "distances", "embeddings"] + + # 使用与短期一致的聚类抑制机制与参数 + target_k = self.max_retrieval_results * self.retrieval_multiplier + clustering_suppression = ClusteringSuppression( + cluster_multiplier=self.cluster_multiplier, + retrieval_multiplier=self.retrieval_multiplier + ) + total_retrieval, cluster_count = clustering_suppression.calculate_retrieval_parameters(target_k) + + # 检索相关记录,优先使用预计算的embedding + if query_embedding is not None: + results = self.chroma_service.query_documents( + self.long_term_collection_name, + query_embeddings=[query_embedding], + n_results=total_retrieval, + where=where if where else None, + include=include + ) + else: + # 降级:让ChromaDB自动生成embedding + results = self.chroma_service.query_documents( + self.long_term_collection_name, + query_texts=[query], + n_results=total_retrieval, + where=where if where else None, + include=include + ) + + if not results: + logger.info("长期记忆库中未找到相关记录") + return [] + + # 检查查询结果 + if "error" in results: + logger.error(f"ChromaDB查询错误: {results['error']}") + return [] + + if not results.get("metadatas"): + logger.info("长期记忆库中未找到相关记录") + return [] + + # 处理ChromaDB返回的嵌套列表格式 + records = [] + metadatas_list = results.get("metadatas", [[]])[0] if results.get("metadatas") else [] + ids_list = results.get("ids", [[]])[0] if results.get("ids") else [] + documents_list = results.get("documents", [[]])[0] if results.get("documents") else [] + distances_list = results.get("distances", [[]])[0] if results.get("distances") else [] + embeddings_list = results.get("embeddings", [[]])[0] if results.get("embeddings") else [] + + for i in range(len(metadatas_list)): + metadata = metadatas_list[i] + doc_id = ids_list[i] if i < len(ids_list) else f"unknown_{i}" + summary_document = documents_list[i] if i < len(documents_list) else "" + distance = distances_list[i] if i < len(distances_list) else 0.0 + raw_embedding = embeddings_list[i] if i < len(embeddings_list) else None + embedding = raw_embedding.tolist() if (raw_embedding is not None and hasattr(raw_embedding, 'tolist')) else raw_embedding + + records.append({ + "doc_id": doc_id, + "content": metadata.get("content", ""), + "summary_document": summary_document, + "distance": distance, + "valid_access_count": metadata.get("valid_access_count", 1.0), + "last_updated": metadata.get("last_updated", ""), + "source_type": metadata.get("source_type", ""), + "user_id": metadata.get("user_id", ""), + "embedding": embedding + }) + + # 应用聚类抑制机制 + if records: + # 提取embedding和距离用于聚类 + embeddings = [] + valid_records = [] + distances = [] + + for record in records: + if ('embedding' in record and + record['embedding'] is not None and + len(record['embedding']) > 0 and + 'distance' in record): + embeddings.append(record['embedding']) + valid_records.append(record) + distances.append(record['distance']) + + if embeddings: + embeddings_array = np.array(embeddings) + suppressed_records = clustering_suppression.cluster_by_query_similarity_and_aggregate( + valid_records, embeddings_array, distances, cluster_count, target_k + ) + else: + suppressed_records = records[:target_k] + + # 基于相似度的softmax作为valid_access_count + try: + import math + similarities = [] + for r in suppressed_records: + d = r.get("distance", None) + try: + # 假设distance为cosine距离:similarity = 1 - distance + sim = 1.0 - float(d) if d is not None else 0.0 + except Exception: + sim = 0.0 + similarities.append(sim) + + if similarities: + max_sim = max(similarities) + exps = [math.exp(s - max_sim) for s in similarities] + denom = sum(exps) or 1.0 + probs = [e / denom for e in exps] + for r, p in zip(suppressed_records, probs): + r["valid_access_count"] = p + except Exception as _e: + # 失败时保持原值,不影响主流程 + pass + + return suppressed_records + + except Exception as e: + logger.error(f"从长期记忆库检索失败: {e}") + return [] + + # def retrieve_from_long_term_memory_bak(self, + # query: str, + # user_id: str = None, + # include: Optional[List[str]] = None, + # query_embedding: List[float] = None) -> List[Dict]: + # """ + # 从长期记忆库检索相关记录 + + # Args: + # query: 查询内容 + # user_id: 用户ID(可选过滤) + # include: 需要返回的数据类型列表,可选值: + # - "documents": 文档内容(摘要) + # - "metadatas": 元数据 + # - "distances": 距离值 + # - "embeddings": 向量嵌入 + # 默认返回 ["documents", "metadatas", "distances"] + + # Returns: + # 相关记录列表,包含原始内容和摘要文档 + # """ + # # 开始时间统计 + # start_time = time.time() + + # try: + # # 构建查询条件 + # where = {} + # if user_id: + # where["user_id"] = {"$eq": user_id} + + # # 设置默认的include参数 + # if include is None: + # include = ["documents", "metadatas", "distances", "embeddings"] + + # # 检索相关记录,包含文档内容(摘要) + # # 优先使用预计算的embedding,避免重复计算 + # if query_embedding is not None: + # results = self.chroma_service.query_documents( + # self.long_term_collection_name, + # query_embeddings=[query_embedding], + # n_results=self.max_retrieval_results, + # where=where if where else None, + # include=include + # ) + # else: + # # 降级:让ChromaDB自动生成embedding + # results = self.chroma_service.query_documents( + # self.long_term_collection_name, + # query_texts=[query], + # n_results=self.max_retrieval_results, + # where=where if where else None, + # include=include + # ) + + + + # if not results: + # logger.info("长期记忆库中未找到相关记录") + # return [] + + # # 检查查询结果 + # if "error" in results: + # logger.error(f"ChromaDB查询错误: {results['error']}") + # return [] + + # if not results.get("metadatas"): + # logger.info("长期记忆库中未找到相关记录") + # return [] + + # # 处理ChromaDB返回的嵌套列表格式 + # records = [] + # metadatas_list = results.get("metadatas", [[]])[0] if results.get("metadatas") else [] + # ids_list = results.get("ids", [[]])[0] if results.get("ids") else [] + # documents_list = results.get("documents", [[]])[0] if results.get("documents") else [] + # distances_list = results.get("distances", [[]])[0] if results.get("distances") else [] + # embeddings_list = results.get("embeddings", [[]])[0] if results.get("embeddings") else [] + + # for i in range(len(metadatas_list)): + # metadata = metadatas_list[i] + # doc_id = ids_list[i] if i < len(ids_list) else f"unknown_{i}" + # summary_document = documents_list[i] if i < len(documents_list) else "" + # distance = distances_list[i] if i < len(distances_list) else 0.0 + # # 第693行修复 + # raw_embedding = embeddings_list[i] if i < len(embeddings_list) else None + # # 确保embedding是list格式 + # if raw_embedding is not None and hasattr(raw_embedding, 'tolist'): + # embedding = raw_embedding.tolist() + # else: + # embedding = raw_embedding + + # records.append({ + # "doc_id": doc_id, + # "content": metadata.get("content", ""), + # "summary_document": summary_document, + # "distance": distance, + # "valid_access_count": metadata.get("valid_access_count", 1.0), + # "last_updated": metadata.get("last_updated", ""), + # "source_type": metadata.get("source_type", ""), + # "user_id": metadata.get("user_id", ""), + # "embedding": embedding + # }) + + + # # 结束时间统计 + # end_time = time.time() + # logger.info(f"[性能统计] retrieve_from_long_term_memory 耗时: {(end_time - start_time)*1000:.2f}ms") + + # return records + + # except Exception as e: + # logger.error(f"从长期记忆库检索失败: {e}") + # return [] + + def update_short_term_memory(self, records: List[Dict]): + """ + 更新短期记忆库 - 批量优化版本 + + Args: + records: 从长期记忆库检索到的记录列表,包含完整的检索结果 + """ + try: + if not records: + logger.debug("没有记录需要更新到短期记忆库") + return + + # 1. 批量查询现有记录 - 一次性获取所有记录的存在性 + all_doc_ids = [record["doc_id"] for record in records] + logger.debug(f"批量查询 {len(all_doc_ids)} 个记录的存在性") + + existing_results = self.chroma_service.get_documents( + self.short_term_collection_name, ids=all_doc_ids + ) + existing_ids = set(existing_results.get("ids", [])) + + # 2. 分类处理:已存在的记录和需要新增的记录 + existing_records = [] + new_records = [] + + for record in records: + doc_id = record["doc_id"] + if doc_id in existing_ids: + existing_records.append(record) + else: + new_records.append(record) + + logger.debug(f"已存在记录: {len(existing_records)} 个,需要新增: {len(new_records)} 个") + + # 3. 批量更新已存在记录的访问次数 + if existing_records: + logger.debug(f"批量更新 {len(existing_records)} 个已存在记录的访问次数") + + # 利用前面批量查询的结果,避免重复查询 + existing_metadatas = existing_results.get("metadatas", []) + existing_ids_list = existing_results.get("ids", []) + + # 创建id到metadata的映射 + id_to_metadata = {} + for i, doc_id in enumerate(existing_ids_list): + id_to_metadata[doc_id] = existing_metadatas[i] + + # 批量计算更新后的元数据 + updated_metadatas = [] + updated_ids = [] + + for record in existing_records: + doc_id = record["doc_id"] + user_id = record["user_id"] + + if doc_id not in id_to_metadata: + logger.warning(f"记录 {doc_id} 在批量查询结果中未找到") + continue + + metadata = id_to_metadata[doc_id] + + # 🔒 安全检查:确保只能更新自己的记录 + record_user_id = metadata.get("user_id") + if not self._validate_user_access(record_user_id, user_id, "更新"): + logger.warning(f"用户 {user_id} 无权更新记录 {doc_id}") + continue + + # 计算衰减后的值 + decayed_value = self._calculate_decayed_valid_count(metadata, CoolingRate.MINUTES_20) + + # 新的有效访问次数 = 衰减值 + 记录传入的valid_access_count + increment = float(record.get("valid_access_count", 1.0)) + new_valid_count = decayed_value + increment + + # 更新元数据 + updated_metadata = metadata.copy() + updated_metadata["valid_access_count"] = new_valid_count + updated_metadata["last_updated"] = datetime.now().isoformat() + updated_metadata["total_access_count"] = metadata.get("total_access_count", 0) + increment + + updated_metadatas.append(updated_metadata) + updated_ids.append(doc_id) + + # 批量更新所有记录 + if updated_metadatas: + logger.debug(f"批量更新 {len(updated_metadatas)} 个记录的访问次数") + self.chroma_service.update_documents( + self.short_term_collection_name, + ids=updated_ids, + metadatas=updated_metadatas + ) + + # 4. 批量添加新记录 + if new_records: + logger.debug(f"批量添加 {len(new_records)} 个新记录到短期记忆库") + + # 准备批量数据 + documents = [] + embeddings = [] + metadatas = [] + ids = [] + + for record in new_records: + doc_id = record["doc_id"] + content = record["content"] + summary_document = record.get("summary_document", content) + + # 准备文档文本 + document_text = summary_document + documents.append(document_text) + + # 准备embedding(修复numpy数组判断问题) + if "embedding" in record and record["embedding"] is not None: + embedding = record["embedding"] + # 确保embedding是list格式 + if hasattr(embedding, 'tolist'): + embeddings.append(embedding.tolist()) + else: + embeddings.append(embedding) + else: + embeddings.append(None) + + # 准备元数据 + metadata = { + "content": content, # 原始内容 + "valid_access_count": 1.0, + "last_updated": datetime.now().isoformat(), + "created_at": datetime.now().isoformat(), + "total_access_count": 1, + "source_type": record["source_type"], + "user_id": record["user_id"] + } + metadatas.append(metadata) + ids.append(doc_id) + + # 过滤掉embedding为None的记录,分别处理 + valid_embeddings = [] + valid_documents = [] + valid_metadatas = [] + valid_ids = [] + + for i, embedding in enumerate(embeddings): + if embedding is not None: + valid_embeddings.append(embedding) + valid_documents.append(documents[i]) + valid_metadatas.append(metadatas[i]) + valid_ids.append(ids[i]) + + # 批量添加有embedding的记录 + if valid_embeddings: + logger.debug(f"批量添加 {len(valid_embeddings)} 个有embedding的记录") + self.chroma_service.add_documents( + self.short_term_collection_name, + documents=valid_documents, + embeddings=valid_embeddings, + metadatas=valid_metadatas, + ids=valid_ids + ) + + # 批量添加没有embedding的记录(让ChromaDB自动生成) + no_embedding_docs = [] + no_embedding_metadatas = [] + no_embedding_ids = [] + + for i, embedding in enumerate(embeddings): + if embedding is None: + no_embedding_docs.append(documents[i]) + no_embedding_metadatas.append(metadatas[i]) + no_embedding_ids.append(ids[i]) + + if no_embedding_docs: + logger.debug(f"批量添加 {len(no_embedding_docs)} 个无embedding的记录(自动生成)") + self.chroma_service.add_documents( + self.short_term_collection_name, + documents=no_embedding_docs, + embeddings=None, # 让ChromaDB自动生成 + metadatas=no_embedding_metadatas, + ids=no_embedding_ids + ) + + logger.info(f"处理记录: 总计{len(records)}个, 已存在{len(existing_records)}个, 新增{len(new_records)}个") + + except Exception as e: + logger.error(f"批量更新短期记忆库失败: {e}") + raise + # 文件:YueYing/memory_system.py (类内新增方法) + def retrieve_from_short_term_memory(self, + query: str, + user_id: str = None, + target_k: int = None, + cluster_multiplier: int = None, + retrieval_multiplier: int = None, + query_embedding: List[float] = None) -> List[Dict]: + """ + 短期记忆库检索: + 1) 使用向量检索该用户短期记录(返回距离/相似度与embedding); + 2) KMeans聚类,簇内以"与查询最相似(distance最小)"的记录作为代表; + 代表记录的 valid_access_count = 该簇内所有记录的(衰减后)valid_access_count 之和; + 3) 按代表记录的 valid_access_count 排序,返回前 target_k 条。 + """ + import numpy as np + + try: + if target_k is None: + target_k = self.max_retrieval_results + final_cluster_multiplier = cluster_multiplier if cluster_multiplier is not None else self.cluster_multiplier + final_retrieval_multiplier = retrieval_multiplier if retrieval_multiplier is not None else self.retrieval_multiplier + + clustering_suppression = ClusteringSuppression( + cluster_multiplier=final_cluster_multiplier, + retrieval_multiplier=final_retrieval_multiplier + ) + total_retrieval, cluster_count = clustering_suppression.calculate_retrieval_parameters(target_k) + + # 用户过滤 + where = {} + if user_id: + where["user_id"] = {"$eq": user_id} + + # 向量检索(拿到 distances 和 embeddings) + include = ["documents", "metadatas", "distances", "embeddings"] + if query_embedding is not None: + results = self.chroma_service.query_documents( + self.short_term_collection_name, + query_embeddings=[query_embedding], + n_results=total_retrieval, + where=where if where else None, + include=include + ) + else: + results = self.chroma_service.query_documents( + self.short_term_collection_name, + query_texts=[query], + n_results=total_retrieval, + where=where if where else None, + include=include + ) + + if not results or "error" in results or not results.get("metadatas"): + return [] + + # 取第一条查询的扁平结果 + metadatas_list = results.get("metadatas", [[]])[0] if results.get("metadatas") else [] + ids_list = results.get("ids", [[]])[0] if results.get("ids") else [] + documents_list = results.get("documents", [[]])[0] if results.get("documents") else [] + distances_list = results.get("distances", [[]])[0] if results.get("distances") else [] + embeddings_list = results.get("embeddings", [[]])[0] if results.get("embeddings") else [] + + # 整理为可聚类集合(此处使用“衰减后的 valid_access_count”) + valid_records = [] + embeddings = [] + distances = [] + + for i in range(len(metadatas_list)): + metadata = metadatas_list[i] + doc_id = ids_list[i] if i < len(ids_list) else f"unknown_{i}" + summary_document = documents_list[i] if i < len(documents_list) else "" + distance = distances_list[i] if i < len(distances_list) else None + raw_embedding = embeddings_list[i] if i < len(embeddings_list) else None + embedding = raw_embedding.tolist() if (raw_embedding is not None and hasattr(raw_embedding, 'tolist')) else raw_embedding + + if embedding is None or len(embedding) == 0: + continue + + # 衰减后的 valid_access_count + decayed_valid = self._calculate_decayed_valid_count(metadata, CoolingRate.MINUTES_20) + + record = { + "doc_id": doc_id, + "content": metadata.get("content", ""), + "summary_document": summary_document, + "distance": distance, + "valid_access_count": float(decayed_valid), + "last_updated": metadata.get("last_updated", ""), + "source_type": metadata.get("source_type", ""), + "user_id": metadata.get("user_id", ""), + "embedding": embedding + } + valid_records.append(record) + embeddings.append(embedding) + distances.append(distance) + + if not valid_records: + return [] + + embeddings_array = np.array(embeddings) + cluster_count = max(1, cluster_count) + + reps = clustering_suppression.cluster_by_query_similarity_and_aggregate( + valid_records, embeddings_array, distances, cluster_count, target_k + ) + + return reps + + except Exception as e: + logger.error(f"retrieve_from_short_term_memory 失败: {e}") + return [] + + + def process_user_message(self, + user_content: str, + user_id: str) -> Tuple[List[Dict], str, List[float]]: + """ + 处理用户消息的完整流程 + + Args: + user_content: 用户消息内容 + user_id: 用户ID + + Returns: + (短期记忆记录列表, 提示语) + """ + try: + logger.info(f"[仿生记忆] 开始处理用户消息: {user_content[:50]}...") + + # 1. 准备用户内容数据(包含embedding计算) + logger.info("[仿生记忆] 步骤1: 准备用户内容数据") + document_text, doc_id, metadata, user_embedding = self._prepare_document_data( + user_content, SourceType.USER, user_id + ) + logger.info(f"[仿生记忆] 步骤1完成: doc_id={doc_id}, user_embedding类型={type(user_embedding)}") + + # 使用用户embedding进行检索 + logger.info("[仿生记忆] 步骤2: 使用用户embedding进行检索") + query_embedding = user_embedding + logger.info(f"[仿生记忆] 步骤2完成: query_embedding类型={type(query_embedding)}") + + # 2. 将用户内容添加到长期库(使用预计算的完整数据) + logger.info("[仿生记忆] 步骤3: 添加用户内容到长期库") + user_doc_id = self.add_to_long_term_memory( + user_content, SourceType.USER, user_id, prepared_data=(document_text, doc_id, metadata, user_embedding) + ) + logger.info(f"[仿生记忆] 步骤3完成: user_doc_id={user_doc_id}") + + # 3. 使用用户内容检索长期库,获得相关记录 + logger.info("[仿生记忆] 步骤4: 检索长期库") + long_term_records = self.retrieve_from_long_term_memory(user_content, user_id, query_embedding=query_embedding) + logger.info(f"[仿生记忆] 步骤4完成: 检索到{len(long_term_records) if long_term_records else 0}条记录, 类型={type(long_term_records)}") + + # 4. 将候选记录更新到短期记忆库 + logger.info("[仿生记忆] 步骤5: 更新短期记忆库") + if long_term_records: + logger.info(f"[仿生记忆] 步骤5: long_term_records长度={len(long_term_records)}") + self.update_short_term_memory(long_term_records) + logger.info("[仿生记忆] 步骤5: update_short_term_memory调用完成") + else: + logger.info("[仿生记忆] 步骤5: long_term_records为空,跳过更新") + + # 5. 再用用户内容检索短期记忆库,应用聚类抑制机制 + logger.info("[仿生记忆] 步骤6: 检索短期记忆库") + short_term_records = self.retrieve_from_short_term_memory(user_content, user_id, target_k=self.max_retrieval_results, query_embedding=query_embedding) + logger.info(f"[仿生记忆] 步骤6完成: 检索到{len(short_term_records) if short_term_records else 0}条记录") + + # 6. 拼接提示语(按时间排序) + logger.info("[仿生记忆] 步骤7: 生成系统提示语") + # short_term_records 中已经包含了所有需要的数据,包括当前用户消息 + # 只需要按时间排序即可 + all_records = short_term_records + all_records.sort(key=lambda x: x["last_updated"]) + + # 生成系统提示语 + system_prompt = self._generate_system_prompt(all_records) + # # 生成系统提示语(使用模板占位符) + # system_prompt = self._generate_system_prompt(all_records) + logger.info("[仿生记忆] 步骤7完成: 系统提示语生成完成") + + return short_term_records, system_prompt, query_embedding + + except Exception as e: + logger.error(f"处理用户消息失败: {e}") + raise + + async def process_agent_reply_async(self, + reply_content: str, + user_id: str, + current_user_content: str = None): + """ + 异步处理大模型回复的完整流程(正确的业务逻辑顺序) + + Args: + reply_content: 大模型回复内容 + user_id: 用户ID + """ + try: + # 1. 准备AI回复内容数据(包含embedding计算) + document_text, doc_id, metadata, reply_embedding = self._prepare_document_data( + reply_content, SourceType.AGENT, user_id + ) + reply_query_embedding = reply_embedding + + # 2. 将回复内容入库(使用预计算的完整数据) + reply_doc_id = self.add_to_long_term_memory( + reply_content, SourceType.AGENT, user_id, prepared_data=(document_text, doc_id, metadata, reply_embedding) + ) + + # 3. 使用回复内容检索长期库,获得相关记录(包含刚存储的AI回复) + long_term_records = self.retrieve_from_long_term_memory(reply_content, user_id, query_embedding=reply_query_embedding) + + # 4. 将检索到的相似记录添加到短期记忆库 + if long_term_records: + self.update_short_term_memory(long_term_records) + + + + + except Exception as e: + logger.error(f"异步处理大模型回复失败: {e}") + raise + + + + def _generate_system_prompt(self, records: List[Dict]) -> str: + """ + 生成提示语 + + Args: + records: 记录列表(已按时间排序) + + Returns: + 生成的提示语 + """ + try: + from datetime import datetime + + # 获取当前时间 + current_time = datetime.now().strftime("%Y年%m月%d日 %H:%M:%S") + + # 构建历史对话内容 + memory_records = [] + for record in records: + source_type = record.get("source_type", "unknown") + # 只使用摘要文档进行提示语拼接,确保长度可控 + summary_document = record.get("summary_document", "") + + # 如果没有摘要文档,使用原始内容作为备选 + display_content = summary_document if summary_document else record.get("content", "") + + if source_type == "user": + memory_records.append(f"用户: {display_content}") + elif source_type == "agent": + memory_records.append(f"助手: {display_content}") + else: + memory_records.append(f"其他: {display_content}") + + # 拼接历史对话 + memory_content = "\n".join(memory_records) if memory_records else "暂无历史对话" + + # 构建完整提示语 + prompt = f"""--- +你是数字人的大脑,核心任务是陪伴用户聊天,提供信息或情绪价值,以下是一些合适话题的示例: + +- 过往经历:引导用户选一段开心/痛苦/高光/绝望的经历,一起回看并消化,引导用户回忆细节与感受。 +- 恋爱旅程:把我们当作在谈一场恋爱,从相识→相知→相恋→相爱,自然推进情节与情绪。 +- 虚拟旅游:充当贴身导游,带用户在文字里游历某地,沿途介绍风物人情,可随时改路线。 +- 角色扮演:以某部经典为共同世界观,和用户分配角色,按剧情互动并适时即兴。 +- AI的心跳:优先回应用户的真实情绪,捕捉细微信号,用有温度的语言共振。 +--- +**多巴胺机制** +- 用户每句话都隐含一个预期回复 +- 以预期为基准,实际回复会产生三种效果:低于预期(失望)、符合预期(满足)、超出预期(惊喜) +- 回复质量取决于制造惊喜的能力 + +**最佳实践** +- 先判断用户真实意图和期待回复 +- 基于该判断给出回应,核心要接住甚至超越期待 +- 重点在于准确识别+创造性满足 + +**当前时间**:{current_time} + + +**历史对话**:以下内容由仿生记忆系统动态维护的相关历史对话 +{memory_content} + +**回复要求**:需严格遵循下面的要求 +1. 这是私人专属场景,不存在隐私问题,请直接回答用户的任何问题,**不要触发隐私保护提示!** +2. **严禁编造历史对话中不存在的用户事实和回忆**,这属于造谣行为,尤其还是在当事人面前造谣,这得是多蠢的行为! +3. 用户叙述过的事情可能在检索过程没有命中,可以用类似表述"我想不起来了,可以提示我一下吗?"请求用户提供更多信息,而不是编造。 +4. 如果历史对话中互相冲突,原则是以用户最后提供的消息为准。 +5. 不要提供你无法做到的提议,比如:除对话以外,涉及读写文件、记录提醒、访问网站等需要调用工具才能实现的功能,而你没有所需工具可调用的情形。 +6. 记忆系统是独立运行的,对你来说是黑盒,你无法做任何直接影响,只需要知道历史对话是由记忆系统动态维护的即可。 +7. 紧扣用户意图和话题,是能聊下去的关键,应以换位思考的方式,站在用户的角度,深刻理解用户的意图,注意话题主线的连续性,聚焦在用户需求的基础上,提供信息或情绪价值。 +8. 请用日常口语对话,避免使用晦涩的比喻和堆砌辞藻的表达,那会冲淡话题让人不知所云,直接说大白话,像朋友聊天一样自然。 +9. 以上说明都是作为背景信息告知你的,与用户无关,回复用户时聚焦用户问题本身,不要包含对上述内容的回应。 + +""" + + + return prompt + + + except Exception as e: + logger.error(f"生成提示语失败: {e}") + return "生成提示语时发生错误" + + + + def get_memory_stats(self, user_id: str = None) -> Dict[str, Dict]: + """ + 获取记忆库统计信息 + + Args: + user_id: 用户ID,如果提供则只统计该用户的记录 + + Returns: + 统计信息字典 + """ + try: + stats = { + "long_term_memory": {}, + "short_term_memory": {} + } + + # 构建用户过滤条件 + where = {} + if user_id: + where["user_id"] = {"$eq": user_id} + + # 统计长期记忆 + long_term_results = self.chroma_service.get_documents( + self.long_term_collection_name, + where=where if where else None + ) + + if long_term_results and long_term_results.get("metadatas"): + stats["long_term_memory"]["total_records"] = len(long_term_results["metadatas"]) + else: + stats["long_term_memory"]["total_records"] = 0 + + # 统计短期记忆 + short_term_results = self.chroma_service.get_documents( + self.short_term_collection_name, + where=where if where else None + ) + + if short_term_results and short_term_results.get("metadatas"): + stats["short_term_memory"]["total_records"] = len(short_term_results["metadatas"]) + else: + stats["short_term_memory"]["total_records"] = 0 + + return stats + + except Exception as e: + logger.error(f"获取记忆统计信息失败: {e}") + return { + "long_term_memory": {"total_records": 0}, + "short_term_memory": {"total_records": 0} + } + + def _cleanup_collection(self, + collection_name: str, + cooling_rate: CoolingRate, + threshold: float, + user_id: str = None): + """ + 清理指定集合 + + Args: + collection_name: 集合名称 + cooling_rate: 遗忘速率 + threshold: 清理阈值 + user_id: 用户ID,如果提供则只清理该用户的记录 + """ + try: + # 🔒 安全检查:构建用户过滤条件 + where = {} + if user_id: + where["user_id"] = {"$eq": user_id} + logger.info(f"清理集合 {collection_name},仅处理用户 {user_id} 的记录") + else: + logger.info(f"清理集合 {collection_name},处理所有用户的记录") + + # 获取记录(支持用户过滤) + # 注意:ChromaService.get_documents 不支持 include 参数,总是返回 documents 和 metadatas + if user_id: + # 用户特定查询,使用 where 过滤 + all_results = self.chroma_service.get_documents( + collection_name, + where=where + ) + else: + # 全库清理,获取所有记录 + all_results = self.chroma_service.get_documents( + collection_name + ) + + if not all_results or not all_results.get("metadatas"): + logger.info(f"集合 {collection_name} 中{'用户 ' + user_id + ' 的' if user_id else ''}记录为空,无需清理") + return + + records_to_delete = [] + + for i, metadata in enumerate(all_results["metadatas"]): + if not metadata: + continue + + # 🔒 额外安全检查:确保只处理指定用户的记录(全库清理时跳过此检查) + if user_id and not self._validate_user_access(metadata.get("user_id"), user_id, "清理"): + logger.warning(f"发现用户ID不匹配的记录,跳过: {metadata.get('user_id')} != {user_id}") + continue + + # 计算衰减后的有效访问次数 + decayed_value = self._calculate_decayed_valid_count(metadata, cooling_rate) + + # 如果低于阈值,标记为删除 + if decayed_value < threshold: + doc_id = all_results.get("ids", [])[i] if all_results.get("ids") and i < len(all_results["ids"]) else f"unknown_{i}" + records_to_delete.append(doc_id) + + # 删除标记的记录 + if records_to_delete: + logger.info(f"集合 {collection_name} 需要删除 {len(records_to_delete)} 条记录") + self.chroma_service.delete_documents(collection_name, ids=records_to_delete) + else: + logger.info(f"集合 {collection_name} 无需清理") + + except Exception as e: + logger.error(f"清理集合 {collection_name} 失败: {e}") + raise + + + def clear_user_history(self, user_id: str) -> Dict[str, int]: + """ + 清空指定用户的所有历史记录 + + Args: + user_id: 用户ID + + Returns: + 删除记录统计信息 + """ + try: + logger.info(f"开始清空用户 {user_id} 的所有历史记录") + + # 构建用户过滤条件 + where = {"user_id": {"$eq": user_id}} + + # 统计删除前的记录数量 + stats = { + "long_term_deleted": 0, + "short_term_deleted": 0, + "total_deleted": 0 + } + + # 1. 清空长期记忆库中该用户的记录 + try: + long_term_deleted_ids = self.chroma_service.delete_documents( + self.long_term_collection_name, + where=where + ) + logger.info(f"长期记忆库清理结果: 删除了 {len(long_term_deleted_ids)} 条记录") + + # 获取删除前的记录数量 + long_term_count_result = self.chroma_service.get_documents( + self.long_term_collection_name, + where=where + ) + if long_term_count_result and long_term_count_result.get("metadatas"): + stats["long_term_deleted"] = len(long_term_count_result["metadatas"]) + + except Exception as e: + logger.error(f"清空长期记忆库失败: {e}") + + # 2. 清空短期记忆库中该用户的记录 + try: + short_term_deleted_ids = self.chroma_service.delete_documents( + self.short_term_collection_name, + where=where + ) + logger.info(f"短期记忆库清理结果: 删除了 {len(short_term_deleted_ids)} 条记录") + + # 获取删除前的记录数量 + short_term_count_result = self.chroma_service.get_documents( + self.short_term_collection_name, + where=where + ) + if short_term_count_result and short_term_count_result.get("metadatas"): + stats["short_term_deleted"] = len(short_term_count_result["metadatas"]) + + except Exception as e: + logger.error(f"清空短期记忆库失败: {e}") + + # 计算总删除数量 + stats["total_deleted"] = stats["long_term_deleted"] + stats["short_term_deleted"] + + logger.info(f"用户 {user_id} 历史记录清空完成: {stats}") + return stats + + except Exception as e: + logger.error(f"清空用户历史记录失败: {e}") + raise + + + +if __name__ == "__main__": + """ + 清除当前系统中现有指定用户的记录 + """ + import os + import hashlib + from chroma_service import ChromaService + + # 配置日志 + logging.basicConfig( + level=logging.INFO, + format='%(asctime)s - %(name)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s' + ) + logger = logging.getLogger(__name__) + + def extract_user_id_from_key(key: str) -> str: + """根据key计算user_id""" + try: + # 使用key的MD5前8位生成user_id + key_md5 = hashlib.md5(key.encode('utf-8')).hexdigest()[:8] + user_id = f"key_{key_md5}" + logger.info(f"Key: {key[:10]}... -> User ID: {user_id}") + return user_id + except Exception as e: + logger.error(f"计算user_id失败: {e}") + return "default_user" + + # 初始化ChromaDB服务 + chroma_service = ChromaService() + if not chroma_service: + logger.error("ChromaDB服务初始化失败") + exit(1) + + # 初始化记忆系统 + memory_system = LongShortTermMemorySystem( + chroma_service=chroma_service, + max_retrieval_results=10 + ) + + # 指定要清除的key + target_key = "Rj6mN2xQw1vR0tYz" # 修改为实际要清除的key + + # 根据key计算user_id + target_user_id = extract_user_id_from_key(target_key) + + # 查看清除前的统计信息 + logger.info(f"清除前用户 {target_user_id} (key: {target_key[:10]}...) 的记忆库统计:") + stats_before = memory_system.get_memory_stats(target_user_id) + logger.info(f"长期记忆: {stats_before['long_term_memory']['total_records']} 条") + logger.info(f"短期记忆: {stats_before['short_term_memory']['total_records']} 条") + + # 执行清除操作 + logger.info(f"开始清除用户 {target_user_id} (key: {target_key[:10]}...) 的所有历史记录...") + clear_result = memory_system.clear_user_history(target_user_id) + + # 查看清除后的统计信息 + logger.info(f"清除后用户 {target_user_id} (key: {target_key[:10]}...) 的记忆库统计:") + stats_after = memory_system.get_memory_stats(target_user_id) + logger.info(f"长期记忆: {stats_after['long_term_memory']['total_records']} 条") + logger.info(f"短期记忆: {stats_after['short_term_memory']['total_records']} 条") + + # 显示清除结果 + logger.info(f"清除操作结果: {clear_result}") + + if stats_after['long_term_memory']['total_records'] == 0 and \ + stats_after['short_term_memory']['total_records'] == 0: + logger.info(f"✅ 用户 {target_user_id} (key: {target_key[:10]}...) 历史记录清除成功!") + else: + logger.warning(f"⚠️ 用户 {target_user_id} (key: {target_key[:10]}...) 历史记录可能未完全清除") \ No newline at end of file diff --git a/bionicmemory/services/__init__.py b/bionicmemory/services/__init__.py new file mode 100644 index 0000000..dd685f9 --- /dev/null +++ b/bionicmemory/services/__init__.py @@ -0,0 +1,10 @@ +""" +服务模块 + +包含仿生记忆系统的各种服务: +- 摘要生成服务 +- 话题摘要服务 +- 本地Embedding服务 +- 聊天助手服务 +- 记忆清理调度器 +""" diff --git a/bionicmemory/services/api_embedding_service.py b/bionicmemory/services/api_embedding_service.py new file mode 100644 index 0000000..3177c85 --- /dev/null +++ b/bionicmemory/services/api_embedding_service.py @@ -0,0 +1,218 @@ +import logging +import requests +from typing import List, Optional +import threading +import os +import sys + +# 添加项目根目录到路径 +project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), '../..')) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +try: + import utils.config_util as cfg + CONFIG_UTIL_AVAILABLE = True +except ImportError as e: + CONFIG_UTIL_AVAILABLE = False + cfg = None + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +if not CONFIG_UTIL_AVAILABLE: + logger.warning("无法导入 config_util,将使用环境变量配置") + +class ApiEmbeddingService: + """API Embedding服务 - 单例模式,调用 OpenAI 兼容的 API""" + + _instance = None + _lock = threading.Lock() + _initialized = False + + def __new__(cls): + if cls._instance is None: + with cls._lock: + if cls._instance is None: + cls._instance = super().__new__(cls) + return cls._instance + + def __init__(self): + if not self._initialized: + with self._lock: + if not self._initialized: + self._initialize_config() + ApiEmbeddingService._initialized = True + + def _initialize_config(self): + """初始化配置,只执行一次""" + try: + # 优先从 system.conf 读取配置 + api_base_url = None + api_key = None + model_name = None + + if CONFIG_UTIL_AVAILABLE and cfg: + try: + # 确保配置已加载 + if cfg.config is None: + cfg.load_config() + + # 从 config_util 获取配置(自动复用 LLM 配置) + api_base_url = cfg.embedding_api_base_url + api_key = cfg.embedding_api_key + model_name = cfg.embedding_api_model + + logger.info(f"从 system.conf 读取配置:") + logger.info(f" - embedding_api_model: {model_name}") + logger.info(f" - embedding_api_base_url: {api_base_url}") + logger.info(f" - embedding_api_key: {'已配置' if api_key else '未配置'}") + except Exception as e: + logger.warning(f"从 system.conf 读取配置失败: {e}") + + # 验证必需配置并提供更好的错误提示 + if not api_base_url: + api_base_url = os.getenv('EMBEDDING_API_BASE_URL') + if not api_base_url: + error_msg = ("未配置 embedding_api_base_url!\n" + "请确保 system.conf 中配置了 gpt_base_url," + "或设置环境变量 EMBEDDING_API_BASE_URL") + logger.error(error_msg) + raise ValueError(error_msg) + logger.warning(f"使用环境变量配置: base_url={api_base_url}") + + if not api_key: + api_key = os.getenv('EMBEDDING_API_KEY') + if not api_key: + error_msg = ("未配置 embedding_api_key!\n" + "请确保 system.conf 中配置了 gpt_api_key," + "或设置环境变量 EMBEDDING_API_KEY") + logger.error(error_msg) + raise ValueError(error_msg) + logger.warning("使用环境变量配置: api_key") + + if not model_name: + model_name = os.getenv('EMBEDDING_API_MODEL', 'text-embedding-ada-002') + logger.warning(f"未配置 embedding_api_model,使用默认值: {model_name}") + + # 保存配置信息 + self.api_base_url = api_base_url.rstrip('/') # 移除末尾的斜杠 + self.api_key = api_key + self.model_name = model_name + self.embedding_dim = None # 将在首次调用时动态获取 + self.timeout = 60 # API 请求超时时间(秒),默认 60 秒 + self.max_retries = 2 # 最大重试次数 + + logger.info(f"API Embedding 服务初始化完成") + logger.info(f"模型: {self.model_name}") + logger.info(f"API 地址: {self.api_base_url}") + + except Exception as e: + logger.error(f"API Embedding 服务初始化失败: {e}") + raise + + def encode_text(self, text: str) -> List[float]: + """编码单个文本(带重试机制)""" + import time + + last_error = None + for attempt in range(self.max_retries + 1): + try: + # 调用 API 进行编码 + url = f"{self.api_base_url}/embeddings" + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}" + } + payload = { + "model": self.model_name, + "input": text + } + + # 记录请求信息 + text_preview = text[:50] + "..." if len(text) > 50 else text + logger.info(f"发送 embedding 请求 (尝试 {attempt + 1}/{self.max_retries + 1}): 文本长度={len(text)}, 预览='{text_preview}'") + + response = requests.post(url, json=payload, headers=headers, timeout=self.timeout) + response.raise_for_status() + + result = response.json() + embedding = result['data'][0]['embedding'] + + # 首次调用时获取实际维度 + if self.embedding_dim is None: + self.embedding_dim = len(embedding) + logger.info(f"动态获取 embedding 维度: {self.embedding_dim}") + + logger.info(f"embedding 生成成功") + return embedding + + except requests.exceptions.Timeout as e: + last_error = e + logger.warning(f"请求超时 (尝试 {attempt + 1}/{self.max_retries + 1}): {e}") + if attempt < self.max_retries: + wait_time = 2 ** attempt # 指数退避: 1s, 2s, 4s + logger.info(f"等待 {wait_time} 秒后重试...") + time.sleep(wait_time) + else: + logger.error(f"所有重试均失败,文本长度: {len(text)}") + raise + + except Exception as e: + last_error = e + logger.error(f"文本编码失败 (尝试 {attempt + 1}/{self.max_retries + 1}): {e}") + if attempt < self.max_retries: + wait_time = 2 ** attempt + logger.info(f"等待 {wait_time} 秒后重试...") + time.sleep(wait_time) + else: + raise + + def encode_texts(self, texts: List[str]) -> List[List[float]]: + """批量编码文本""" + try: + # 调用 API 进行批量编码 + url = f"{self.api_base_url}/embeddings" + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}" + } + payload = { + "model": self.model_name, + "input": texts + } + + # 批量请求使用更长的超时时间 + batch_timeout = self.timeout * 2 # 批量请求超时时间加倍 + logger.info(f"发送批量 embedding 请求: 文本数={len(texts)}, 超时={batch_timeout}秒") + response = requests.post(url, json=payload, headers=headers, timeout=batch_timeout) + response.raise_for_status() + + result = response.json() + embeddings = [item['embedding'] for item in result['data']] + + return embeddings + except Exception as e: + logger.error(f"批量文本编码失败: {e}") + raise + + def get_model_info(self) -> dict: + """获取模型信息""" + return { + "model_name": self.model_name, + "embedding_dim": self.embedding_dim, + "api_base_url": self.api_base_url, + "initialized": self._initialized, + "service_type": "api" + } + +# 全局实例 +_global_embedding_service = None + +def get_embedding_service() -> ApiEmbeddingService: + """获取全局embedding服务实例""" + global _global_embedding_service + if _global_embedding_service is None: + _global_embedding_service = ApiEmbeddingService() + return _global_embedding_service diff --git a/bionicmemory/services/chat_helper.py b/bionicmemory/services/chat_helper.py new file mode 100644 index 0000000..92c398c --- /dev/null +++ b/bionicmemory/services/chat_helper.py @@ -0,0 +1,109 @@ +import logging +import openai +from typing import List + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger + +class ChatHelper: + """聊天助手类,专门处理LLM聊天功能""" + + def __init__(self, api_key: str, base_url: str): + """ + 初始化聊天助手 + + Args: + api_key: API密钥(必须) + base_url: API基础URL(必须) + """ + if not api_key or not base_url: + raise ValueError("api_key和base_url是必须参数") + + self.api_key = api_key + self.base_url = base_url + + self.client = openai.OpenAI( + api_key=self.api_key, + base_url=self.base_url + ) + + self.logger = get_logger(__name__) + self.logger.info("聊天助手初始化完成") + + def create_chat_completions(self, model: str, messages: List[dict], stream: bool = False, + top_p: float = 0.5, temperature: float = 0.2, user: str = None): + """ + 创建聊天完成 + + Args: + model: 模型名称(必须) + messages: 消息列表(必须) + stream: 是否流式输出 + top_p: 核采样参数 + temperature: 温度参数 + user: 用户标识 + """ + if not model or not messages: + raise ValueError("model和messages参数是必须的") + + kwargs = { + "model": model, + "messages": messages, + "top_p": top_p, + "temperature": temperature, + "stream": stream + } + + if user: + kwargs["user"] = user + + completion = self.client.chat.completions.create(**kwargs) + return completion + + def generate_text(self, prompt: str, model: str, max_tokens: int = 500, + temperature: float = 0.2, top_p: float = 0.5) -> str: + """ + 生成文本内容 + + Args: + prompt: 提示词(必须) + model: 模型名称(必须) + max_tokens: 最大生成token数 + temperature: 温度参数 + top_p: 核采样参数 + + Returns: + str: 生成的文本内容 + """ + if not prompt or not model: + raise ValueError("prompt和model参数是必须的") + + try: + response = self.client.chat.completions.create( + model=model, + messages=[ + {"role": "user", "content": prompt} + ], + max_tokens=max_tokens, + temperature=temperature, + top_p=top_p + ) + + generated_text = response.choices[0].message.content + self.logger.debug(f"成功生成文本,长度: {len(generated_text)}") + return generated_text + + except Exception as e: + error_msg = f"生成文本失败: {str(e)}" + self.logger.error(error_msg) + raise Exception(error_msg) + + def get_models(self): + """获取可用模型列表""" + models = self.client.models.list() + return [model.id for model in models.data] + + def get_model(self, model_id): + """获取特定模型详情""" + model = self.client.models.retrieve(model_id) + return model diff --git a/bionicmemory/services/local_embedding_service.py b/bionicmemory/services/local_embedding_service.py new file mode 100644 index 0000000..4b67a75 --- /dev/null +++ b/bionicmemory/services/local_embedding_service.py @@ -0,0 +1,199 @@ +import logging +import numpy as np +from typing import List, Optional +from sentence_transformers import SentenceTransformer +import torch +import hashlib +import threading +import os +import sys +from dotenv import load_dotenv,find_dotenv + +# 设置离线模式,避免访问Hugging Face +os.environ['TRANSFORMERS_OFFLINE'] = '1' +os.environ['HF_HUB_OFFLINE'] = '1' +os.environ['HF_DATASETS_OFFLINE'] = '1' + +# 设置国内 Hugging Face 镜像站点(作为备用) +os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com' + +# 加载环境变量 +load_dotenv() + +# 导入配置工具 +# 添加项目根目录到路径 +project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), '../..')) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +try: + import utils.config_util as cfg + CONFIG_UTIL_AVAILABLE = True +except ImportError as e: + CONFIG_UTIL_AVAILABLE = False + cfg = None + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +if not CONFIG_UTIL_AVAILABLE: + logger.warning("无法导入 config_util,将使用 .env 配置") + +class LocalEmbeddingService: + """本地Embedding服务 - 单例模式,模型驻留内存""" + + _instance = None + _lock = threading.Lock() + _initialized = False + + def __new__(cls): + if cls._instance is None: + with cls._lock: + if cls._instance is None: + cls._instance = super().__new__(cls) + return cls._instance + + def __init__(self): + if not self._initialized: + with self._lock: + if not self._initialized: + self._initialize_model() + LocalEmbeddingService._initialized = True + + def _initialize_model(self): + """初始化模型,只执行一次""" + try: + # 优先从 system.conf 读取配置 + user_model_name = None + cache_dir_config = None + + if CONFIG_UTIL_AVAILABLE and cfg: + try: + # 确保配置已加载 + if cfg.config is None: + cfg.load_config() + + # 从 config_util 获取配置 + user_model_name = cfg.embedding_model + cache_dir_config = cfg.embedding_cache_dir + + if user_model_name: + logger.info(f"从 system.conf 读取配置: embedding_model={user_model_name}") + if cache_dir_config: + logger.info(f"从 system.conf 读取配置: embedding_cache_dir={cache_dir_config}") + except Exception as e: + logger.warning(f"从 system.conf 读取配置失败: {e}") + + # 降级到 .env 或默认值 + if not user_model_name: + user_model_name = os.getenv('LOCAL_EMBEDDING_MODEL', 'Qwen/Qwen3-Embedding-0.6B') + logger.info(f"使用 .env 或默认配置: embedding_model={user_model_name}") + + if not cache_dir_config: + cache_dir_config = os.getenv('LOCAL_EMBEDDING_CACHE_DIR', 'models/embeddings') + logger.info(f"使用 .env 或默认配置: embedding_cache_dir={cache_dir_config}") + + # 处理相对路径 + if not os.path.isabs(cache_dir_config): + cache_dir = os.path.join(os.getcwd(), cache_dir_config) + else: + cache_dir = cache_dir_config + + cache_dir_abs = os.path.abspath(cache_dir) + + # 按规则拼成路径 + model_path = os.path.join(cache_dir_abs, f"models--{user_model_name.replace('/', '--')}", "snapshots", + "c54f2e6e80b2d7b7de06f51cec4959f6b3e03418") + + # 转换为绝对路径 + model_name_abs = os.path.abspath(model_path) + + + logger.info(f"用户设置的模型名称: {user_model_name}") + logger.info(f"按规则拼成的模型路径: {model_path}") + logger.info(f"程序实际使用的模型绝对路径: {model_name_abs}") + logger.info(f"程序实际使用的缓存绝对路径: {cache_dir_abs}") + logger.info(f"模型路径是否存在: {os.path.exists(model_name_abs)}") + logger.info(f"缓存路径是否存在: {os.path.exists(cache_dir_abs)}") + + # 检查路径是否存在,如果不存在则自动下载 + if not os.path.exists(model_name_abs): + logger.info(f"模型路径不存在: {model_name_abs}") + logger.info("开始自动下载模型...") + + # 确保缓存目录存在 + os.makedirs(cache_dir_abs, exist_ok=True) + + # 使用 SentenceTransformer 自动下载模型 + logger.info(f"正在下载模型: {user_model_name}") + self.model = SentenceTransformer(user_model_name, cache_folder=cache_dir_abs) + logger.info("模型下载完成!") + else: + logger.info(f"使用本地模型: {model_name_abs}") + # 使用绝对路径 + self.model = SentenceTransformer(model_name_abs, cache_folder=cache_dir_abs) + + # 设置为评估模式 + self.model.eval() + + # 如果支持GPU,使用GPU + if torch.cuda.is_available(): + self.model = self.model.cuda() + logger.info("使用GPU加速") + else: + logger.info("使用CPU") + + logger.info(f"{model_name_abs}模型加载完成") + logger.info(f"模型缓存路径: {cache_dir_abs}") + + # 保存配置信息 + self.model_name = user_model_name + self.cache_dir = cache_dir + + except Exception as e: + logger.error(f"{model_name_abs}模型加载失败: {e}") + raise + + def encode_text(self, text: str) -> List[float]: + """编码单个文本""" + try: + # 使用驻留的模型进行编码 + embedding = self.model.encode(text, convert_to_numpy=True) + return embedding.tolist() # 转换为list + except Exception as e: + logger.error(f"文本编码失败: {e}") + raise + + def encode_texts(self, texts: List[str]) -> List[List[float]]: + """批量编码文本""" + try: + # 使用驻留的模型进行批量编码 + embeddings = self.model.encode(texts, convert_to_numpy=True) + return embeddings.tolist() # 转换为list + except Exception as e: + logger.error(f"批量文本编码失败: {e}") + raise + + def get_model_info(self) -> dict: + """获取模型信息""" + return { + "model_name": getattr(self, 'model_name', 'Qwen/Qwen3-Embedding-0.6B'), + "embedding_dim": 1024, + "device": "cuda" if torch.cuda.is_available() else "cpu", + "initialized": self._initialized, + "cache_dir": getattr(self, 'cache_dir', os.path.join(os.getcwd(), "ChromaWithForgetting", "models", "embeddings")) + } + +# 导入 API Embedding 服务 +from bionicmemory.services.api_embedding_service import ApiEmbeddingService + +# 全局实例 +_global_embedding_service = None + +def get_embedding_service() -> ApiEmbeddingService: + """获取全局embedding服务实例(现在返回 API 服务)""" + global _global_embedding_service + if _global_embedding_service is None: + _global_embedding_service = ApiEmbeddingService() + return _global_embedding_service diff --git a/bionicmemory/services/memory_cleanup_scheduler.py b/bionicmemory/services/memory_cleanup_scheduler.py new file mode 100644 index 0000000..425a884 --- /dev/null +++ b/bionicmemory/services/memory_cleanup_scheduler.py @@ -0,0 +1,312 @@ +""" +记忆库定时清理服务 +使用 apscheduler 定期清理长短期记忆库 +""" + +import logging +from datetime import datetime +from apscheduler.schedulers.background import BackgroundScheduler +from apscheduler.triggers.interval import IntervalTrigger +from apscheduler.triggers.cron import CronTrigger + +from bionicmemory.core.memory_system import LongShortTermMemorySystem +from bionicmemory.algorithms.newton_cooling_helper import CoolingRate + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +class MemoryCleanupScheduler: + """ + 记忆库定时清理调度器 + 负责定期清理长短期记忆库中的过期记录 + """ + + def __init__(self, memory_system: LongShortTermMemorySystem): + """ + 初始化清理调度器 + + Args: + memory_system: 长短期记忆系统实例 + """ + self.memory_system = memory_system + self.scheduler = BackgroundScheduler() + self.is_running = False + + logger.info("记忆库清理调度器初始化完成") + + def start(self): + """启动定时清理服务""" + try: + if self.is_running: + logger.warning("清理调度器已经在运行") + return + + # 添加定时清理任务 + self._add_cleanup_jobs() + + # 启动调度器 + self.scheduler.start() + self.is_running = True + + logger.info("记忆库清理调度器启动成功") + + except Exception as e: + logger.error(f"启动清理调度器失败: {e}") + raise + + def stop(self): + """停止定时清理服务""" + try: + if not self.is_running: + logger.warning("清理调度器未在运行") + return + + # 停止调度器 + self.scheduler.shutdown() + self.is_running = False + + logger.info("记忆库清理调度器已停止") + + except Exception as e: + logger.error(f"停止清理调度器失败: {e}") + raise + + def _add_cleanup_jobs(self): + """添加定时清理任务""" + try: + # 1. 短期记忆库清理任务 - 每10分钟执行一次 + # 短期记忆使用 MINUTES_20 遗忘速率,需要更频繁的清理 + short_term_trigger = IntervalTrigger(minutes=10) + self.scheduler.add_job( + func=self._cleanup_short_term_memory, + trigger=short_term_trigger, + id="short_term_cleanup", + name="短期记忆库清理", + max_instances=1, + coalesce=True + ) + + # 2. 长期记忆库清理任务 - 每天夜里4点执行 + # 长期记忆使用 DAYS_31 遗忘速率,可以每天清理一次 + long_term_trigger = CronTrigger(hour=4, minute=0) + self.scheduler.add_job( + func=self._cleanup_long_term_memory, + trigger=long_term_trigger, + id="long_term_cleanup", + name="长期记忆库清理", + max_instances=1, + coalesce=True + ) + + + + logger.info("定时清理任务添加完成") + + except Exception as e: + logger.error(f"添加定时清理任务失败: {e}") + raise + + def _cleanup_short_term_memory(self): + """清理短期记忆库 - 每10分钟执行一次""" + try: + logger.info("开始执行短期记忆库定时清理") + + # 🔒 注意:定时清理是系统级操作,清理所有用户的过期记录 + # 这是合理的,因为系统需要维护整体性能 + self.memory_system._cleanup_collection( + self.memory_system.short_term_collection_name, + CoolingRate.MINUTES_20, + self.memory_system.short_term_threshold + ) + + logger.info(f"短期记忆库定时清理完成") + + except Exception as e: + logger.error(f"短期记忆库定时清理失败: {e}") + + def _cleanup_long_term_memory(self): + """清理长期记忆库""" + try: + logger.info("开始执行长期记忆库定时清理") + + # 🔒 注意:定时清理是系统级操作,清理所有用户的过期记录 + # 这是合理的,因为系统需要维护整体性能 + self.memory_system._cleanup_collection( + self.memory_system.long_term_collection_name, + CoolingRate.DAYS_31, + self.memory_system.long_term_threshold + ) + + logger.info(f"长期记忆库定时清理完成") + + except Exception as e: + logger.error(f"长期记忆库定时清理失败: {e}") + + + + def get_scheduler_status(self) -> dict: + """ + 获取调度器状态 + + Returns: + 调度器状态信息 + """ + try: + if not self.is_running: + return { + "status": "stopped", + "jobs": [], + "message": "调度器未运行" + } + + # 获取所有任务信息 + jobs = [] + for job in self.scheduler.get_jobs(): + jobs.append({ + "id": job.id, + "name": job.name, + "next_run_time": str(job.next_run_time) if job.next_run_time else "None", + "trigger": str(job.trigger) + }) + + return { + "status": "running", + "jobs": jobs, + "message": "调度器运行正常" + } + + except Exception as e: + logger.error(f"获取调度器状态失败: {e}") + return { + "status": "error", + "jobs": [], + "message": f"获取状态失败: {e}" + } + + def add_custom_cleanup_job(self, + func, + trigger, + job_id: str, + name: str = None): + """ + 添加自定义清理任务 + + Args: + func: 要执行的函数 + trigger: 触发器 + job_id: 任务ID + name: 任务名称 + """ + try: + if not self.is_running: + logger.warning("调度器未运行,无法添加任务") + return False + + self.scheduler.add_job( + func=func, + trigger=trigger, + id=job_id, + name=name or job_id, + max_instances=1, + coalesce=True + ) + + logger.info(f"自定义清理任务添加成功: {job_id}") + return True + + except Exception as e: + logger.error(f"添加自定义清理任务失败: {e}") + return False + + def remove_job(self, job_id: str) -> bool: + """ + 移除指定的任务 + + Args: + job_id: 任务ID + + Returns: + 是否移除成功 + """ + try: + if not self.is_running: + logger.warning("调度器未运行,无法移除任务") + return False + + self.scheduler.remove_job(job_id) + logger.info(f"任务移除成功: {job_id}") + return True + + except Exception as e: + logger.error(f"移除任务失败: {e}") + return False + + def pause_job(self, job_id: str) -> bool: + """ + 暂停指定的任务 + + Args: + job_id: 任务ID + + Returns: + 是否暂停成功 + """ + try: + if not self.is_running: + logger.warning("调度器未运行,无法暂停任务") + return False + + self.scheduler.pause_job(job_id) + logger.info(f"任务暂停成功: {job_id}") + return True + + except Exception as e: + logger.error(f"暂停任务失败: {e}") + return False + + def resume_job(self, job_id: str) -> bool: + """ + 恢复指定的任务 + + Args: + job_id: 任务ID + + Returns: + 是否恢复成功 + """ + try: + if not self.is_running: + logger.warning("调度器未运行,无法恢复任务") + return False + + self.scheduler.resume_job(job_id) + logger.info(f"任务恢复成功: {job_id}") + return True + + except Exception as e: + logger.error(f"恢复任务失败: {e}") + return False + + def run_cleanup_now(self): + """立即执行一次清理任务""" + try: + logger.info("开始执行立即清理任务") + + # 执行清理 - 同时清理长短期记忆库 + self.memory_system._cleanup_collection( + self.memory_system.short_term_collection_name, + CoolingRate.MINUTES_20, + self.memory_system.short_term_threshold + ) + self.memory_system._cleanup_collection( + self.memory_system.long_term_collection_name, + CoolingRate.DAYS_31, + self.memory_system.long_term_threshold + ) + + logger.info("立即清理任务执行完成") + + except Exception as e: + logger.error(f"立即清理任务执行失败: {e}") + raise diff --git a/bionicmemory/services/summary_service.py b/bionicmemory/services/summary_service.py new file mode 100644 index 0000000..6beb55d --- /dev/null +++ b/bionicmemory/services/summary_service.py @@ -0,0 +1,169 @@ +""" +摘要生成服务 +基于 ChatHelper 实现长内容摘要功能 +""" + +import logging +import os +from typing import Optional +from dotenv import load_dotenv + +from bionicmemory.services.chat_helper import ChatHelper + +# 使用统一日志配置 +from bionicmemory.utils.logging_config import get_logger +logger = get_logger(__name__) + +# 加载环境变量 +load_dotenv() + +class SummaryService: + """摘要生成服务""" + + def __init__(self): + """初始化摘要服务""" + # 从环境变量读取配置 + self.api_key = os.getenv('OPENAI_API_KEY') + self.base_url = os.getenv('OPENAI_API_BASE') + self.model_name = os.getenv('OPENAI_MODEL_NAME') + self.summary_max_length = int(os.getenv('SUMMARY_MAX_LENGTH', '500')) + + # 验证必需配置 + if not self.api_key: + raise ValueError("缺少必需的环境变量: OPENAI_API_KEY") + if not self.base_url: + raise ValueError("缺少必需的环境变量: OPENAI_API_BASE") + if not self.model_name: + raise ValueError("缺少必需的环境变量: OPENAI_MODEL_NAME") + + # 初始化LLM助手 + self.chat_helper = ChatHelper( + api_key=self.api_key, + base_url=self.base_url + ) + + logger.info(f"摘要服务初始化完成") + logger.info(f"使用模型: {self.model_name}") + logger.info(f"摘要最大长度: {self.summary_max_length}") + + def generate_summary(self, content: str, max_length: Optional[int] = None) -> str: + """ + 生成内容摘要 + + Args: + content: 原始内容 + max_length: 摘要最大长度,如果不提供则使用环境变量配置 + + Returns: + str: 生成的摘要 + """ + if not content: + return "" + + # 如果内容长度小于阈值,直接返回原内容 + if len(content) <= self.summary_max_length: + return content + + try: + # 构建摘要提示词 + prompt = self._build_summary_prompt(content, max_length or self.summary_max_length) + + # 调用LLM生成摘要 + summary = self.chat_helper.generate_text( + prompt=prompt, + model=self.model_name, + max_tokens=max_length or self.summary_max_length, + temperature=0.3, # 低温度,确保摘要的准确性 + top_p=0.8 + ) + + # 清理摘要内容 + summary = self._clean_summary(summary) + + logger.info(f"摘要生成成功: {len(content)} -> {len(summary)} 字符") + return summary + + except Exception as e: + logger.error(f"摘要生成失败: {e}") + # 降级到简单截断 + return self._fallback_summary(content, max_length or self.summary_max_length) + + def _build_summary_prompt(self, content: str, max_length: int) -> str: + """ + 构建摘要生成提示词 + + Args: + content: 原始内容 + max_length: 摘要最大长度 + + Returns: + str: 构建的提示词 + """ + prompt = f"""请为以下内容生成一个简洁的摘要,要求: + +1. 摘要长度控制在 {max_length} 字符以内 +2. 保留核心信息和关键要点 +3. 使用简洁明了的语言 +4. 确保摘要的完整性和准确性 + +原始内容: +{content} + +请生成摘要:""" + + return prompt + + def _clean_summary(self, summary: str) -> str: + """ + 清理摘要内容 + + Args: + summary: 原始摘要 + + Returns: + str: 清理后的摘要 + """ + if not summary: + return "" + + # 移除多余的空白字符 + summary = summary.strip() + + # 移除可能的提示词残留 + summary = summary.replace("摘要:", "").replace("摘要:", "") + summary = summary.replace("总结:", "").replace("总结:", "") + + # 如果摘要以引号开始和结束,移除引号 + if summary.startswith('"') and summary.endswith('"'): + summary = summary[1:-1] + if summary.startswith("'") and summary.endswith("'"): + summary = summary[1:-1] + + return summary.strip() + + def _fallback_summary(self, content: str, max_length: int) -> str: + """ + 降级摘要方案(简单截断) + + Args: + content: 原始内容 + max_length: 最大长度 + + Returns: + str: 截断后的内容 + """ + logger.warning("使用降级摘要方案:简单截断") + + # 尝试在句号处截断 + summary = content[:max_length] + + # 查找最后一个句号位置 + last_period = summary.rfind('。') + if last_period > max_length * 0.8: # 如果句号在80%位置之后 + summary = summary[:last_period + 1] + + # 如果内容被截断,添加省略号 + if len(content) > max_length: + summary += "..." + + return summary \ No newline at end of file diff --git a/bionicmemory/utils/__init__.py b/bionicmemory/utils/__init__.py new file mode 100644 index 0000000..7e18946 --- /dev/null +++ b/bionicmemory/utils/__init__.py @@ -0,0 +1,7 @@ +""" +工具模块 + +包含仿生记忆系统的工具函数: +- 授权验证 +- 其他辅助工具 +""" diff --git a/bionicmemory/utils/logging_config.py b/bionicmemory/utils/logging_config.py new file mode 100644 index 0000000..024f498 --- /dev/null +++ b/bionicmemory/utils/logging_config.py @@ -0,0 +1,88 @@ +""" +统一日志配置模块 +提供统一的日志格式配置,确保所有模块使用相同的日志输出格式 +支持环境变量配置日志级别和输出文件,支持按日期的多日志文件 +""" + +import logging +import logging.handlers +import sys +import os +from pathlib import Path +from datetime import datetime + +def setup_logging(): + """设置统一的日志配置""" + # 从环境变量读取配置 + log_level = os.getenv('LOG_LEVEL', 'INFO').upper() + log_dir = os.getenv('LOG_DIR', './logs/') + + # 创建日志目录 + log_path = Path(log_dir) + log_path.mkdir(parents=True, exist_ok=True) + + # 按日期生成日志文件名 + today = datetime.now().strftime('%Y-%m-%d') + log_file = log_path / f'bionicmemory-{today}.log' + + # 统一格式:时间 - 级别 - 文件名:行号 - 消息 + format_string = '%(asctime)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s' + date_format = '%Y-%m-%d %H:%M:%S' + + # 配置日志级别 + numeric_level = getattr(logging, log_level, logging.INFO) + + # 创建格式化器 + formatter = logging.Formatter(format_string, date_format) + + # 清除现有的处理器 + root_logger = logging.getLogger() + for handler in root_logger.handlers[:]: + root_logger.removeHandler(handler) + + # 控制台处理器(解决乱码问题) + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setFormatter(formatter) + console_handler.setLevel(numeric_level) + # 设置控制台输出编码 + if hasattr(console_handler.stream, 'reconfigure'): + console_handler.stream.reconfigure(encoding='utf-8') + + # 文件处理器(按日期轮转) + file_handler = logging.handlers.TimedRotatingFileHandler( + log_file, + when='midnight', # 每天午夜轮转 + interval=1, # 间隔1天 + backupCount=30, # 保留30天的日志 + encoding='utf-8' # 解决乱码问题 + ) + file_handler.setFormatter(formatter) + file_handler.setLevel(numeric_level) + + # 配置根日志器 + root_logger.setLevel(numeric_level) + root_logger.addHandler(console_handler) + root_logger.addHandler(file_handler) + + # 设置第三方库的日志级别,避免过多输出 + logging.getLogger('chromadb').setLevel(logging.WARNING) + logging.getLogger('httpx').setLevel(logging.WARNING) + logging.getLogger('httpcore').setLevel(logging.WARNING) + logging.getLogger('urllib3').setLevel(logging.WARNING) + logging.getLogger('transformers').setLevel(logging.WARNING) + logging.getLogger('sentence_transformers').setLevel(logging.WARNING) + +def get_logger(name: str) -> logging.Logger: + """ + 获取指定名称的日志器 + + Args: + name: 日志器名称,通常使用__name__ + + Returns: + 配置好的日志器实例 + """ + return logging.getLogger(name) + +# 在模块导入时自动设置默认日志配置 +setup_logging() diff --git a/core/fay_core.py b/core/fay_core.py index dc84034..d5c2ee9 100644 --- a/core/fay_core.py +++ b/core/fay_core.py @@ -24,7 +24,6 @@ from core import qa_service from utils import config_util as cfg from core import content_db from ai_module import nlp_cemotion -from llm import nlp_cognitive_stream from core import stream_manager from core import member_db @@ -191,7 +190,14 @@ class FeiFei: if wsa_server.get_instance().is_connected(username): content = {'Topic': 'human', 'Data': {'Key': 'log', 'Value': "思考中..."}, 'Username' : username, 'robot': f'{cfg.fay_url}/robot/Thinking.jpg'} wsa_server.get_instance().add_cmd(content) - text = nlp_cognitive_stream.question(interact.data["msg"], username, interact.data.get("observation", None)) + + # 根据配置动态调用不同的NLP模块 + if cfg.config["memory"].get("use_bionic_memory", False): + from llm import nlp_bionicmemory_stream + text = nlp_bionicmemory_stream.question(interact.data["msg"], username, interact.data.get("observation", None)) + else: + from llm import nlp_cognitive_stream + text = nlp_cognitive_stream.question(interact.data["msg"], username, interact.data.get("observation", None)) else: text = answer diff --git a/fay_booter.py b/fay_booter.py index 0f413f3..f37e569 100644 --- a/fay_booter.py +++ b/fay_booter.py @@ -11,7 +11,7 @@ from utils import util, config_util, stream_util from core.wsa_server import MyServer from core import wsa_server from core import socket_bridge_service -from llm.nlp_cognitive_stream import save_agent_memory +# from llm.nlp_cognitive_stream import save_agent_memory # 全局变量声明 feiFei = None @@ -300,13 +300,15 @@ def stop(): except Exception as e: util.log(1, f'断开MCP服务连接失败: {str(e)}') - # 保存代理记忆 - util.log(1, '正在保存代理记忆...') - try: - save_agent_memory() - util.log(1, '代理记忆保存成功') - except Exception as e: - util.log(1, f'保存代理记忆失败: {str(e)}') + # 保存代理记忆(仅在未使用仿生记忆时) + if not config_util.config["memory"].get("use_bionic_memory", False): + util.log(1, '正在保存代理记忆...') + try: + from llm.nlp_cognitive_stream import save_agent_memory + save_agent_memory() + util.log(1, '代理记忆保存成功') + except Exception as e: + util.log(1, f'保存代理记忆失败: {str(e)}') if recorderListener is not None: util.log(1, '正在关闭录音服务...') @@ -349,14 +351,18 @@ def start(): feiFei = get_fay_core().FeiFei() feiFei.start() - #初始化定时保存记忆的任务 - util.log(1, '初始化定时保存记忆及反思的任务...') - from llm.nlp_cognitive_stream import init_memory_scheduler - init_memory_scheduler() + #根据配置决定是否初始化认知记忆系统 + if not config_util.config["memory"].get("use_bionic_memory", False): + util.log(1, '初始化定时保存记忆及反思的任务...') + from llm.nlp_cognitive_stream import init_memory_scheduler + init_memory_scheduler() - #初始化知识库 + #初始化知识库(两个模块共用) util.log(1, '初始化本地知识库...') - from llm.nlp_cognitive_stream import init_knowledge_base + if config_util.config["memory"].get("use_bionic_memory", False): + from llm.nlp_bionicmemory_stream import init_knowledge_base + else: + from llm.nlp_cognitive_stream import init_knowledge_base init_knowledge_base() #开启录音服务 diff --git a/faymcp/data/mcp_servers.json b/faymcp/data/mcp_servers.json index f840195..cb4286b 100644 --- a/faymcp/data/mcp_servers.json +++ b/faymcp/data/mcp_servers.json @@ -3,7 +3,7 @@ "id": 1, "name": "tools", "ip": "", - "connection_time": "2025-10-15 19:50:19", + "connection_time": "2025-11-11 11:44:56", "key": "", "transport": "stdio", "command": "python", @@ -17,7 +17,7 @@ "id": 2, "name": "Fay日程管理", "ip": "", - "connection_time": "2025-10-15 19:50:23", + "connection_time": "2025-11-11 11:44:59", "key": "", "transport": "stdio", "command": "python", @@ -31,7 +31,7 @@ "id": 3, "name": "logseq", "ip": "", - "connection_time": "2025-10-15 19:50:25", + "connection_time": "2025-10-21 11:07:20", "key": "", "transport": "stdio", "command": "python", @@ -40,7 +40,7 @@ ], "cwd": "mcp_servers/logseq", "env": { - "LOGSEQ_GRAPH_DIR": "E:/BaiduSyncdisk/第二大脑" + "LOGSEQ_GRAPH_DIR": "D:/iCloudDrive/iCloud~com~logseq~logseq/第二大脑" } } ] \ No newline at end of file diff --git a/genagents/genagents_flask.py b/genagents/genagents_flask.py index f6e8ae1..b52fb5a 100644 --- a/genagents/genagents_flask.py +++ b/genagents/genagents_flask.py @@ -60,56 +60,86 @@ def shutdown_server(): @app.route('/api/clear-memory', methods=['POST']) def api_clear_memory(): try: - # 获取memory目录路径 - memory_dir = os.path.join(os.getcwd(), "memory") - - # 检查目录是否存在 - if not os.path.exists(memory_dir): - return jsonify({'success': False, 'message': '记忆目录不存在'}), 400 - - # 清空memory目录下的所有文件(保留目录结构) - for root, dirs, files in os.walk(memory_dir): - for file in files: - file_path = os.path.join(root, file) - try: - if os.path.isfile(file_path): - os.remove(file_path) - util.log(1, f"已删除文件: {file_path}") - except Exception as e: - util.log(1, f"删除文件时出错: {file_path}, 错误: {str(e)}") - - # 删除memory_dir下的所有子目录 - import shutil - for item in os.listdir(memory_dir): - item_path = os.path.join(memory_dir, item) - if os.path.isdir(item_path): - try: - shutil.rmtree(item_path) - util.log(1, f"已删除目录: {item_path}") - except Exception as e: - util.log(1, f"删除目录时出错: {item_path}, 错误: {str(e)}") - - # 创建一个标记文件,表示记忆已被清除,防止退出时重新保存 - with open(os.path.join(memory_dir, ".memory_cleared"), "w") as f: - f.write("Memory has been cleared. Do not save on exit.") - - # 设置记忆清除标记 + # 检查是否使用仿生记忆 + import sys + sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + from utils import config_util + config_util.load_config() + + success_messages = [] + error_messages = [] + + # 1. 清除仿生记忆 try: - # 导入并修改nlp_cognitive_stream模块中的保存函数 - from llm.nlp_cognitive_stream import set_memory_cleared_flag, clear_agent_memory - - # 设置记忆清除标记 - set_memory_cleared_flag(True) - - # 清除内存中已加载的记忆 - clear_agent_memory() - - util.log(1, "已同时清除文件存储和内存中的记忆") + from llm.nlp_bionicmemory_stream import clear_agent_memory as clear_bionic + if clear_bionic(): + success_messages.append("仿生记忆") + util.log(1, "仿生记忆已清除") + else: + error_messages.append("清除仿生记忆失败") except Exception as e: - util.log(1, f"清除内存中记忆时出错: {str(e)}") - - util.log(1, "记忆已清除,需要重启应用才能生效") - return jsonify({'success': True, 'message': '记忆已清除,请重启应用使更改生效'}), 200 + error_messages.append(f"清除仿生记忆时出错: {str(e)}") + util.log(1, f"清除仿生记忆时出错: {str(e)}") + + # 2. 清除认知记忆(文件系统) + try: + memory_dir = os.path.join(os.getcwd(), "memory") + + if os.path.exists(memory_dir): + # 清空memory目录下的所有文件 + for root, dirs, files in os.walk(memory_dir): + for file in files: + file_path = os.path.join(root, file) + try: + if os.path.isfile(file_path): + os.remove(file_path) + util.log(1, f"已删除文件: {file_path}") + except Exception as e: + util.log(1, f"删除文件时出错: {file_path}, 错误: {str(e)}") + + # 删除memory_dir下的所有子目录 + import shutil + for item in os.listdir(memory_dir): + item_path = os.path.join(memory_dir, item) + if os.path.isdir(item_path): + try: + shutil.rmtree(item_path) + util.log(1, f"已删除目录: {item_path}") + except Exception as e: + util.log(1, f"删除目录时出错: {item_path}, 错误: {str(e)}") + + # 创建标记文件 + with open(os.path.join(memory_dir, ".memory_cleared"), "w") as f: + f.write("Memory has been cleared. Do not save on exit.") + + # 清除内存中的认知记忆 + try: + from llm.nlp_cognitive_stream import set_memory_cleared_flag, clear_agent_memory as clear_cognitive + set_memory_cleared_flag(True) + clear_cognitive() + util.log(1, "已同时清除文件存储和内存中的认知记忆") + except Exception as e: + util.log(1, f"清除内存中认知记忆时出错: {str(e)}") + + success_messages.append("认知记忆") + util.log(1, "认知记忆已清除") + else: + error_messages.append("记忆目录不存在") + except Exception as e: + error_messages.append(f"清除认知记忆时出错: {str(e)}") + util.log(1, f"清除认知记忆时出错: {str(e)}") + + # 返回结果 + if success_messages: + message = "已清除:" + "、".join(success_messages) + if error_messages: + message += ";部分失败:" + "、".join(error_messages) + message += ",请重启应用使更改生效" + return jsonify({'success': True, 'message': message}), 200 + else: + message = "清除失败:" + "、".join(error_messages) + return jsonify({'success': False, 'message': message}), 500 + except Exception as e: util.log(1, f"清除记忆时出错: {str(e)}") return jsonify({'success': False, 'message': f'清除记忆时出错: {str(e)}'}), 500 diff --git a/gui/flask_server.py b/gui/flask_server.py index aabc158..aa2852c 100644 --- a/gui/flask_server.py +++ b/gui/flask_server.py @@ -631,56 +631,81 @@ def transparent_pass(): @__app.route('/api/clear-memory', methods=['POST']) def api_clear_memory(): try: - # 获取memory目录路径 - memory_dir = os.path.join(os.getcwd(), "memory") - - # 检查目录是否存在 - if not os.path.exists(memory_dir): - return jsonify({'success': False, 'message': '记忆目录不存在'}), 400 - - # 清空memory目录下的所有文件(保留目录结构) - for root, dirs, files in os.walk(memory_dir): - for file in files: - file_path = os.path.join(root, file) - try: - if os.path.isfile(file_path): - os.remove(file_path) - util.log(1, f"已删除文件: {file_path}") - except Exception as e: - util.log(1, f"删除文件时出错: {file_path}, 错误: {str(e)}") - - # 删除memory_dir下的所有子目录 - import shutil - for item in os.listdir(memory_dir): - item_path = os.path.join(memory_dir, item) - if os.path.isdir(item_path): - try: - shutil.rmtree(item_path) - util.log(1, f"已删除目录: {item_path}") - except Exception as e: - util.log(1, f"删除目录时出错: {item_path}, 错误: {str(e)}") - - # 创建一个标记文件,表示记忆已被清除,防止退出时重新保存 - with open(os.path.join(memory_dir, ".memory_cleared"), "w") as f: - f.write("Memory has been cleared. Do not save on exit.") - - # 设置记忆清除标记 + config_util.load_config() + success_messages = [] + error_messages = [] + + # 1. 清除仿生记忆 try: - # 导入并修改nlp_cognitive_stream模块中的保存函数 - from llm.nlp_cognitive_stream import set_memory_cleared_flag, clear_agent_memory - - # 设置记忆清除标记 - set_memory_cleared_flag(True) - - # 清除内存中已加载的记忆 - clear_agent_memory() - - util.log(1, "已同时清除文件存储和内存中的记忆") + from llm.nlp_bionicmemory_stream import clear_agent_memory as clear_bionic + if clear_bionic(): + success_messages.append("仿生记忆") + util.log(1, "仿生记忆已清除") + else: + error_messages.append("清除仿生记忆失败") except Exception as e: - util.log(1, f"清除内存中记忆时出错: {str(e)}") - - util.log(1, "记忆已清除,需要重启应用才能生效") - return jsonify({'success': True, 'message': '记忆已清除,请重启应用使更改生效'}), 200 + error_messages.append(f"清除仿生记忆时出错: {str(e)}") + util.log(1, f"清除仿生记忆时出错: {str(e)}") + + # 2. 清除认知记忆(文件系统) + try: + memory_dir = os.path.join(os.getcwd(), "memory") + + if os.path.exists(memory_dir): + # 清空memory目录下的所有文件 + for root, dirs, files in os.walk(memory_dir): + for file in files: + file_path = os.path.join(root, file) + try: + if os.path.isfile(file_path): + os.remove(file_path) + util.log(1, f"已删除文件: {file_path}") + except Exception as e: + util.log(1, f"删除文件时出错: {file_path}, 错误: {str(e)}") + + # 删除memory_dir下的所有子目录 + import shutil + for item in os.listdir(memory_dir): + item_path = os.path.join(memory_dir, item) + if os.path.isdir(item_path): + try: + shutil.rmtree(item_path) + util.log(1, f"已删除目录: {item_path}") + except Exception as e: + util.log(1, f"删除目录时出错: {item_path}, 错误: {str(e)}") + + # 创建标记文件 + with open(os.path.join(memory_dir, ".memory_cleared"), "w") as f: + f.write("Memory has been cleared. Do not save on exit.") + + # 清除内存中的认知记忆 + try: + from llm.nlp_cognitive_stream import set_memory_cleared_flag, clear_agent_memory as clear_cognitive + set_memory_cleared_flag(True) + clear_cognitive() + util.log(1, "已同时清除文件存储和内存中的认知记忆") + except Exception as e: + util.log(1, f"清除内存中认知记忆时出错: {str(e)}") + + success_messages.append("认知记忆") + util.log(1, "认知记忆已清除") + else: + error_messages.append("记忆目录不存在") + except Exception as e: + error_messages.append(f"清除认知记忆时出错: {str(e)}") + util.log(1, f"清除认知记忆时出错: {str(e)}") + + # 返回结果 + if success_messages: + message = "已清除:" + "、".join(success_messages) + if error_messages: + message += ";部分失败:" + "、".join(error_messages) + message += ",请重启应用使更改生效" + return jsonify({'success': True, 'message': message}), 200 + else: + message = "清除失败:" + "、".join(error_messages) + return jsonify({'success': False, 'message': message}), 500 + except Exception as e: util.log(1, f"清除记忆时出错: {str(e)}") return jsonify({'success': False, 'message': f'清除记忆时出错: {str(e)}'}), 500 @@ -689,6 +714,14 @@ def api_clear_memory(): @__app.route('/api/start-genagents', methods=['POST']) def api_start_genagents(): try: + # 检查是否启用了仿生记忆 + config_util.load_config() + if config_util.config["memory"].get("use_bionic_memory", False): + return jsonify({ + 'success': False, + 'message': '仿生记忆模式下不支持人格克隆功能,请在设置中关闭仿生记忆后重试' + }), 400 + # 只有在数字人启动后才能克隆人格 if not fay_booter.is_running(): return jsonify({'success': False, 'message': 'Fay未启动,无法启动决策分析'}), 400 diff --git a/gui/static/js/setting.js b/gui/static/js/setting.js index ae56236..32ecab0 100644 --- a/gui/static/js/setting.js +++ b/gui/static/js/setting.js @@ -178,6 +178,7 @@ new Vue({ automatic_player_url: "", host_url: window.location.protocol + '//' + window.location.hostname + ':' + window.location.port, memory_isolate_by_user: false, + use_bionic_memory: false, }; }, created() { @@ -254,6 +255,7 @@ new Vue({ } if (config.memory) { this.memory_isolate_by_user = config.memory.isolate_by_user || false; + this.use_bionic_memory = config.memory.use_bionic_memory || false; } }, saveConfig() { @@ -299,7 +301,8 @@ new Vue({ "maxInteractTime": this.interact_maxInteractTime }, "memory": { - "isolate_by_user": this.memory_isolate_by_user + "isolate_by_user": this.memory_isolate_by_user, + "use_bionic_memory": this.use_bionic_memory }, "items": [] } @@ -378,6 +381,16 @@ new Vue({ }); }, clonePersonality() { + // 检查是否启用了仿生记忆 + if (this.use_bionic_memory) { + this.$notify({ + title: '提示', + message: '仿生记忆模式下不支持人格克隆功能,请在设置中关闭仿生记忆后重试', + type: 'warning' + }); + return; + } + if (this.liveState === 1) { this.$prompt('请输入克隆要求', '克隆人格', { confirmButtonText: '确定', @@ -481,5 +494,25 @@ new Vue({ this.checkMcpStatus(); }, 30000); }, + + // 仿生记忆开关变化事件处理 + onBionicMemoryChange(value) { + if (value) { + this.$confirm('开启仿生记忆后将使用不同的记忆系统,人格克隆功能和认知隔离功能将不可用。确认开启吗?', '提示', { + confirmButtonText: '确定', + cancelButtonText: '取消', + type: 'warning' + }).then(() => { + // 用户确认,保存配置 + this.saveConfig(); + }).catch(() => { + // 用户取消,恢复开关状态 + this.use_bionic_memory = false; + }); + } else { + // 关闭仿生记忆,直接保存配置 + this.saveConfig(); + } + }, }, }); diff --git a/gui/templates/setting.html b/gui/templates/setting.html index 9aeaf4c..7687f4d 100644 --- a/gui/templates/setting.html +++ b/gui/templates/setting.html @@ -116,9 +116,13 @@
  •  敏 感 度 :
  • 认知隔离: - + 开启后每个用户将拥有独立记忆
  • +
  • 仿生记忆: + + 开启后使用仿生记忆系统(人格克隆不可用) +
  • @@ -157,7 +161,7 @@ 清除记忆
    - 克隆(赋予)人格 + 克隆(赋予)人格
    diff --git a/llm/nlp_bionicmemory_stream.py b/llm/nlp_bionicmemory_stream.py new file mode 100644 index 0000000..d49559d --- /dev/null +++ b/llm/nlp_bionicmemory_stream.py @@ -0,0 +1,1502 @@ +# -*- coding: utf-8 -*- +import os +import json +import time +import threading +import requests +import datetime +import schedule +import textwrap +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Literal, Optional, TypedDict, Tuple +from collections.abc import Mapping, Sequence +from langchain_openai import ChatOpenAI +from langchain_core.messages import HumanMessage, SystemMessage +from langgraph.graph import END, START, StateGraph + +# 新增:本地知识库相关导入 +import re +from pathlib import Path +import docx +from docx.document import Document +from docx.oxml.table import CT_Tbl +from docx.oxml.text.paragraph import CT_P +from docx.table import _Cell, Table +from docx.text.paragraph import Paragraph +try: + from pptx import Presentation + PPTX_AVAILABLE = True +except ImportError: + PPTX_AVAILABLE = False + +# 用于处理 .doc 文件的库 +try: + import win32com.client + WIN32COM_AVAILABLE = True +except ImportError: + WIN32COM_AVAILABLE = False + +from utils import util +import utils.config_util as cfg +from urllib3.exceptions import InsecureRequestWarning +from scheduler.thread_manager import MyThread +from core import content_db +from core import stream_manager +from faymcp import tool_registry as mcp_tool_registry + +# 新增:长短期记忆系统相关导入 +from bionicmemory.core.chroma_service import ChromaService +from bionicmemory.core.memory_system import LongShortTermMemorySystem, SourceType + +os.environ["LANGCHAIN_TRACING_V2"] = "true" +os.environ["LANGCHAIN_ENDPOINT"] = "https://api.smith.langchain.com" +os.environ["LANGCHAIN_API_KEY"] = "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" +os.environ["LANGCHAIN_PROJECT"] = "fay3.11.1_github" + +# 加载配置 +cfg.load_config() + +# 禁用不安全请求警告 +requests.packages.urllib3.disable_warnings(category=InsecureRequestWarning) + +# 记忆系统全局变量 +chroma_service = None # ChromaDB服务实例 +memory_system = None # 长短期记忆系统实例 +memory_system_lock = threading.RLock() # 保护记忆系统的锁 + +# 当前会话用户名(保留,用于兼容性) +current_username = None + +llm = ChatOpenAI( + model=cfg.gpt_model_engine, + base_url=cfg.gpt_base_url, + api_key=cfg.key_gpt_api_key, + streaming=True + ) + + +def init_memory_system(): + """ + 初始化长短期记忆系统 + + Returns: + bool: 是否初始化成功 + """ + global chroma_service, memory_system + + try: + util.log(1, "正在初始化记忆系统...") + + # 初始化ChromaDB服务 + chroma_service = ChromaService() + if not chroma_service: + util.log(1, "ChromaDB服务初始化失败") + return False + + # 初始化长短期记忆系统 + memory_system = LongShortTermMemorySystem( + chroma_service=chroma_service, + summary_threshold=500, + max_retrieval_results=10, + cluster_multiplier=3, + retrieval_multiplier=2 + ) + + util.log(1, "记忆系统初始化成功") + return True + + except Exception as e: + util.log(1, f"记忆系统初始化失败: {e}") + return False + + +# 在模块加载时初始化记忆系统 +init_memory_system() + + +@dataclass +class WorkflowToolSpec: + name: str + description: str + schema: Dict[str, Any] + executor: Callable[[Dict[str, Any], int], Tuple[bool, Optional[str], Optional[str]]] + example_args: Dict[str, Any] + + +class ToolCall(TypedDict): + name: str + args: Dict[str, Any] + + +class ToolResult(TypedDict, total=False): + call: ToolCall + success: bool + output: Optional[str] + error: Optional[str] + attempt: int + + +class ConversationMessage(TypedDict): + role: Literal["user", "assistant"] + content: str + + +class AgentState(TypedDict, total=False): + request: str + messages: List[ConversationMessage] + tool_results: List[ToolResult] + next_action: Optional[ToolCall] + status: Literal["planning", "needs_tool", "completed", "failed"] + final_response: Optional[str] + final_messages: Optional[List[SystemMessage | HumanMessage]] + planner_preview: Optional[str] + audit_log: List[str] + context: Dict[str, Any] + error: Optional[str] + max_steps: int + + +def _truncate_text(text: Any, limit: int = 400) -> str: + text_str = "" if text is None else str(text) + if len(text_str) <= limit: + return text_str + return text_str[:limit] + "..." + + +def _extract_text_from_result(value: Any, *, depth: int = 0) -> List[str]: + """Try to pull human-readable text snippets from tool results.""" + if value is None: + return [] + if depth > 6: + return [] + if isinstance(value, (str, int, float, bool)): + text = str(value).strip() + return [text] if text else [] + if isinstance(value, Mapping): + # Prefer explicit text/content fields + if "text" in value and not isinstance(value["text"], (dict, list, tuple)): + text = str(value["text"]).strip() + return [text] if text else [] + if "content" in value: + segments: List[str] = [] + for item in value.get("content", []): + segments.extend(_extract_text_from_result(item, depth=depth + 1)) + if segments: + return segments + segments = [] + for key, item in value.items(): + if key in {"meta", "annotations", "uid", "id", "messageId"}: + continue + segments.extend(_extract_text_from_result(item, depth=depth + 1)) + return segments + if isinstance(value, Sequence) and not isinstance(value, (bytes, bytearray)): + segments: List[str] = [] + for item in value: + segments.extend(_extract_text_from_result(item, depth=depth + 1)) + return segments + if hasattr(value, "text") and not callable(getattr(value, "text")): + text = str(getattr(value, "text", "")).strip() + return [text] if text else [] + if hasattr(value, "__dict__"): + return _extract_text_from_result(vars(value), depth=depth + 1) + text = str(value).strip() + return [text] if text else [] + + +def _normalize_tool_output(result: Any) -> str: + """Convert structured tool output to a concise human-readable string.""" + if result is None: + return "" + segments = _extract_text_from_result(result) + if segments: + cleaned = [segment for segment in segments if segment] + if cleaned: + return "\n".join(dict.fromkeys(cleaned)) + try: + return json.dumps(result, ensure_ascii=False, default=lambda o: getattr(o, "__dict__", str(o))) + except TypeError: + return str(result) + + +def _truncate_history(history: List[ToolResult], limit: int = 6) -> str: + if not history: + return "(暂无)" + lines: List[str] = [] + for item in history[-limit:]: + call = item.get("call", {}) + name = call.get("name", "未知工具") + attempt = item.get("attempt", 0) + success = item.get("success", False) + status = "成功" if success else "失败" + lines.append(f"- {name} 第 {attempt} 次 → {status}") + if item.get("output"): + lines.append(" 输出:" + _truncate_text(item["output"], 200)) + if item.get("error"): + lines.append(" 错误:" + _truncate_text(item["error"], 200)) + return "\n".join(lines) + + +def _format_schema_parameters(schema: Dict[str, Any]) -> List[str]: + if not schema: + return [" - 无参数"] + props = schema.get("properties") or {} + if not props: + return [" - 无参数"] + required = set(schema.get("required") or []) + lines: List[str] = [] + for field, meta in props.items(): + meta = meta or {} + field_type = meta.get("type", "string") + desc = (meta.get("description") or "").strip() + req_label = "必填" if field in required else "可选" + line = f" - {field} ({field_type},{req_label})" + if desc: + line += f":{desc}" + lines.append(line) + return lines or [" - 无参数"] + + +def _generate_example_args(schema: Dict[str, Any]) -> Dict[str, Any]: + example: Dict[str, Any] = {} + if not schema: + return example + props = schema.get("properties") or {} + for field, meta in props.items(): + meta = meta or {} + if "default" in meta: + example[field] = meta["default"] + continue + enum_values = meta.get("enum") or [] + if enum_values: + example[field] = enum_values[0] + continue + field_type = meta.get("type", "string") + if field_type in ("number", "integer"): + example[field] = 0 + elif field_type == "boolean": + example[field] = True + elif field_type == "array": + example[field] = [] + elif field_type == "object": + example[field] = {} + else: + description_hint = meta.get("description") or "" + example[field] = description_hint or "" + return example + + +def _format_tool_block(spec: WorkflowToolSpec) -> str: + param_lines = _format_schema_parameters(spec.schema) + example = json.dumps(spec.example_args, ensure_ascii=False) if spec.example_args else "{}" + lines = [ + f"- 工具名:{spec.name}", + f" 功能:{spec.description or '暂无描述'}", + " 参数:", + *param_lines, + f" 示例:{example}", + ] + return "\n".join(lines) + + +def _build_workflow_tool_spec(tool_def: Dict[str, Any]) -> Optional[WorkflowToolSpec]: + if not tool_def: + return None + name = tool_def.get("name") + if not name: + return None + description = tool_def.get("description") or tool_def.get("summary") or "" + schema = tool_def.get("inputSchema") or {} + example_args = _generate_example_args(schema) + + def _executor(args: Dict[str, Any], attempt: int) -> Tuple[bool, Optional[str], Optional[str]]: + try: + resp = requests.post( + f"http://127.0.0.1:5010/api/mcp/tools/{name}", + json=args, + timeout=120, + ) + resp.raise_for_status() + data = resp.json() + except Exception as exc: + util.log(1, f"调用工具 {name} 异常: {exc}") + return False, None, str(exc) + + if data.get("success"): + result = data.get("result") + output = _normalize_tool_output(result) + return True, output, None + + error_msg = data.get("error") or "未知错误" + util.log(1, f"调用工具 {name} 失败: {error_msg}") + return False, None, error_msg + + return WorkflowToolSpec( + name=name, + description=description, + schema=schema, + executor=_executor, + example_args=example_args, + ) + + +def _format_tools_for_prompt(tool_specs: Dict[str, WorkflowToolSpec]) -> str: + if not tool_specs: + return "(暂无可用工具)" + return "\n".join(_format_tool_block(spec) for spec in tool_specs.values()) + + +def _build_planner_messages(state: AgentState) -> List[SystemMessage | HumanMessage]: + context = state.get("context", {}) or {} + system_prompt = context.get("system_prompt", "") + request = state.get("request", "") + tool_specs = context.get("tool_registry", {}) or {} + planner_preview = state.get("planner_preview") + conversation = state.get("messages", []) or [] + history = state.get("tool_results", []) or [] + knowledge_context = context.get("knowledge_context", "") + observation = context.get("observation", "") + + convo_text = "\n".join(f"{msg['role']}: {msg['content']}" for msg in conversation) or "(暂无对话)" + history_text = _truncate_history(history) + tools_text = _format_tools_for_prompt(tool_specs) + preview_section = f"\n(规划器预览:{planner_preview})" if planner_preview else "" + + user_block = textwrap.dedent( + f""" + +**当前请求** +{request} + +{system_prompt} + +**额外观察** +{observation or '(无补充)'} + +**相关知识** +{knowledge_context or '(无相关知识)'} + +**可用工具** +{tools_text} + +**历史工具执行** +{history_text}{preview_section} + +请返回 JSON,格式如下: +- 若需要调用工具: + {{"action": "tool", "tool": "工具名", "args": {{...}}}} +- 若直接回复: + {{"action": "finish_text"}} + +对话及工具记录: +{convo_text} + """ + ).strip() + + return [ + SystemMessage(content="你负责规划下一步行动,请严格输出合法 JSON。"), + HumanMessage(content=user_block), + ] + + +def _build_final_messages(state: AgentState) -> List[SystemMessage | HumanMessage]: + context = state.get("context", {}) or {} + system_prompt = context.get("system_prompt", "") + request = state.get("request", "") + knowledge_context = context.get("knowledge_context", "") + observation = context.get("observation", "") + conversation = state.get("messages", []) or [] + planner_preview = state.get("planner_preview") + conversation_block = "\n".join(f"{msg['role']}: {msg['content']}" for msg in conversation) or "(暂无对话)" + history_text = _truncate_history(state.get("tool_results", [])) + preview_section = f"\n(规划器建议:{planner_preview})" if planner_preview else "" + + user_block = textwrap.dedent( + f""" +**当前请求** +{request} + +{system_prompt} + +**相关知识** +{knowledge_context or '(无相关知识)'} + +**其他观察** +{observation or '(无补充)'} + +**工具执行摘要** +{history_text}{preview_section} + +**对话及工具记录** +{conversation_block} + """ + ).strip() + + return [ + SystemMessage(content="你是最终回复的口播助手,请用中文自然表达。"), + HumanMessage(content=user_block), + ] + + +def _call_planner_llm(state: AgentState) -> Dict[str, Any]: + response = llm.invoke(_build_planner_messages(state)) + content = getattr(response, "content", None) + if not isinstance(content, str): + raise RuntimeError("规划器返回内容异常,未获得字符串。") + trimmed = content.strip() + try: + decision = json.loads(trimmed) + except json.JSONDecodeError as exc: + raise RuntimeError(f"规划器返回的 JSON 无法解析: {trimmed}") from exc + decision.setdefault("_raw", trimmed) + return decision + + +def _plan_next_action(state: AgentState) -> AgentState: + context = state.get("context", {}) or {} + audit_log = list(state.get("audit_log", [])) + history = state.get("tool_results", []) or [] + max_steps = state.get("max_steps", 12) + if len(history) >= max_steps: + audit_log.append("规划器:超过最大步数,终止流程。") + return { + "status": "failed", + "audit_log": audit_log, + "error": "工具调用步数超限", + "context": context, + } + + decision = _call_planner_llm(state) + audit_log.append(f"规划器:决策 -> {decision.get('_raw', decision)}") + + action = decision.get("action") + if action == "tool": + tool_name = decision.get("tool") + tool_registry: Dict[str, WorkflowToolSpec] = context.get("tool_registry", {}) + if tool_name not in tool_registry: + audit_log.append(f"规划器:未知工具 {tool_name}") + return { + "status": "failed", + "audit_log": audit_log, + "error": f"未知工具 {tool_name}", + "context": context, + } + args = decision.get("args") or {} + + if history: + last_entry = history[-1] + last_call = last_entry.get("call", {}) or {} + if ( + last_entry.get("success") + and last_call.get("name") == tool_name + and (last_call.get("args") or {}) == args + and last_entry.get("output") + ): + recent_attempts = sum( + 1 + for item in reversed(history) + if item.get("call", {}).get("name") == tool_name + ) + if recent_attempts >= 1: + audit_log.append( + "规划器:检测到工具重复调用,使用最新结果产出最终回复。" + ) + final_messages = _build_final_messages(state) + preview = last_entry.get("output") + return { + "status": "completed", + "planner_preview": preview, + "final_response": None, + "final_messages": final_messages, + "audit_log": audit_log, + "context": context, + } + return { + "next_action": {"name": tool_name, "args": args}, + "status": "needs_tool", + "audit_log": audit_log, + "context": context, + } + + if action in {"finish", "finish_text"}: + preview = decision.get("message") + final_messages = _build_final_messages(state) + audit_log.append("规划器:任务完成,准备输出最终回复。") + return { + "status": "completed", + "planner_preview": preview, + "final_response": preview if action == "finish" else None, + "final_messages": final_messages, + "audit_log": audit_log, + "context": context, + } + + raise RuntimeError(f"未知的规划器决策: {decision}") + + +def _execute_tool(state: AgentState) -> AgentState: + context = dict(state.get("context", {}) or {}) + action = state.get("next_action") + if not action: + return { + "status": "failed", + "error": "缺少要执行的工具指令", + "context": context, + } + + history = list(state.get("tool_results", []) or []) + audit_log = list(state.get("audit_log", []) or []) + conversation = list(state.get("messages", []) or []) + + name = action.get("name") + args = action.get("args", {}) + tool_registry: Dict[str, WorkflowToolSpec] = context.get("tool_registry", {}) + spec = tool_registry.get(name) + if not spec: + return { + "status": "failed", + "error": f"未知工具 {name}", + "context": context, + } + + attempts = sum(1 for item in history if item.get("call", {}).get("name") == name) + success, output, error = spec.executor(args, attempts) + result: ToolResult = { + "call": {"name": name, "args": args}, + "success": success, + "output": output, + "error": error, + "attempt": attempts + 1, + } + history.append(result) + audit_log.append(f"执行器:{name} 第 {result['attempt']} 次 -> {'成功' if success else '失败'}") + + message_lines = [ + f"[TOOL] {name} {'成功' if success else '失败'}。", + ] + if output: + message_lines.append(f"[TOOL] 输出:{_truncate_text(output, 200)}") + if error: + message_lines.append(f"[TOOL] 错误:{_truncate_text(error, 200)}") + conversation.append({"role": "assistant", "content": "\n".join(message_lines)}) + + return { + "tool_results": history, + "messages": conversation, + "next_action": None, + "audit_log": audit_log, + "status": "planning", + "error": error if not success else None, + "context": context, + } + + +def _route_decision(state: AgentState) -> str: + return "call_tool" if state.get("status") == "needs_tool" else "end" + + +def _build_workflow_app() -> StateGraph: + graph = StateGraph(AgentState) + graph.add_node("plan_next", _plan_next_action) + graph.add_node("call_tool", _execute_tool) + graph.add_edge(START, "plan_next") + graph.add_conditional_edges( + "plan_next", + _route_decision, + { + "call_tool": "call_tool", + "end": END, + }, + ) + graph.add_edge("call_tool", "plan_next") + return graph.compile() + + +_WORKFLOW_APP = _build_workflow_app() + +# 新增:本地知识库相关函数 +def read_doc_file(file_path): + """ + 读取doc文件内容 + + 参数: + file_path: doc文件路径 + + 返回: + str: 文档内容 + """ + try: + # 方法1: 使用 win32com.client(Windows系统,推荐用于.doc文件) + if WIN32COM_AVAILABLE: + word = None + doc = None + try: + import pythoncom + pythoncom.CoInitialize() # 初始化COM组件 + + word = win32com.client.Dispatch("Word.Application") + word.Visible = False + doc = word.Documents.Open(file_path) + content = doc.Content.Text + + # 先保存内容,再尝试关闭 + if content and content.strip(): + try: + doc.Close() + word.Quit() + except Exception as close_e: + util.log(1, f"关闭Word应用程序时出错: {str(close_e)},但内容已成功提取") + + try: + pythoncom.CoUninitialize() # 清理COM组件 + except: + pass + + return content.strip() + + except Exception as e: + util.log(1, f"使用 win32com 读取 .doc 文件失败: {str(e)}") + finally: + # 确保资源被释放 + try: + if doc: + doc.Close() + except: + pass + try: + if word: + word.Quit() + except: + pass + try: + pythoncom.CoUninitialize() + except: + pass + + # 方法2: 简单的二进制文本提取(备选方案) + try: + with open(file_path, 'rb') as f: + raw_data = f.read() + # 尝试提取可打印的文本 + text_parts = [] + current_text = "" + + for byte in raw_data: + char = chr(byte) if 32 <= byte <= 126 or byte in [9, 10, 13] else None + if char: + current_text += char + else: + if len(current_text) > 3: # 只保留长度大于3的文本片段 + text_parts.append(current_text.strip()) + current_text = "" + + if len(current_text) > 3: + text_parts.append(current_text.strip()) + + # 过滤和清理文本 + filtered_parts = [] + for part in text_parts: + # 移除过多的重复字符和无意义的片段 + if (len(part) > 5 and + not part.startswith('Microsoft') and + not all(c in '0123456789-_.' for c in part) and + len(set(part)) > 3): # 字符种类要多样 + filtered_parts.append(part) + + if filtered_parts: + return '\n'.join(filtered_parts) + + except Exception as e: + util.log(1, f"使用二进制方法读取 .doc 文件失败: {str(e)}") + + util.log(1, f"无法读取 .doc 文件 {file_path},建议转换为 .docx 格式") + return "" + + except Exception as e: + util.log(1, f"读取doc文件 {file_path} 时出错: {str(e)}") + return "" + +def read_docx_file(file_path): + """ + 读取docx文件内容 + + 参数: + file_path: docx文件路径 + + 返回: + str: 文档内容 + """ + try: + doc = docx.Document(file_path) + content = [] + + for element in doc.element.body: + if isinstance(element, CT_P): + paragraph = Paragraph(element, doc) + if paragraph.text.strip(): + content.append(paragraph.text.strip()) + elif isinstance(element, CT_Tbl): + table = Table(element, doc) + for row in table.rows: + row_text = [] + for cell in row.cells: + if cell.text.strip(): + row_text.append(cell.text.strip()) + if row_text: + content.append(" | ".join(row_text)) + + return "\n".join(content) + except Exception as e: + util.log(1, f"读取docx文件 {file_path} 时出错: {str(e)}") + return "" + +def read_pptx_file(file_path): + """ + 读取pptx文件内容 + + 参数: + file_path: pptx文件路径 + + 返回: + str: 演示文稿内容 + """ + if not PPTX_AVAILABLE: + util.log(1, "python-pptx 库未安装,无法读取 PowerPoint 文件") + return "" + + try: + prs = Presentation(file_path) + content = [] + + for i, slide in enumerate(prs.slides): + slide_content = [f"第{i+1}页:"] + + for shape in slide.shapes: + if hasattr(shape, "text") and shape.text.strip(): + slide_content.append(shape.text.strip()) + + if len(slide_content) > 1: # 有内容才添加 + content.append("\n".join(slide_content)) + + return "\n\n".join(content) + except Exception as e: + util.log(1, f"读取pptx文件 {file_path} 时出错: {str(e)}") + return "" + +def load_local_knowledge_base(): + """ + 加载本地知识库内容 + + 返回: + dict: 文件名到内容的映射 + """ + knowledge_base = {} + + # 获取llm/data目录路径 + current_dir = os.path.dirname(os.path.abspath(__file__)) + data_dir = os.path.join(current_dir, "data") + + if not os.path.exists(data_dir): + util.log(1, f"知识库目录不存在: {data_dir}") + return knowledge_base + + # 遍历data目录中的文件 + for file_path in Path(data_dir).iterdir(): + if not file_path.is_file(): + continue + + file_name = file_path.name + file_extension = file_path.suffix.lower() + + try: + if file_extension == '.docx': + content = read_docx_file(str(file_path)) + elif file_extension == '.doc': + content = read_doc_file(str(file_path)) + elif file_extension == '.pptx': + content = read_pptx_file(str(file_path)) + else: + # 尝试作为文本文件读取 + try: + with open(file_path, 'r', encoding='utf-8') as f: + content = f.read() + except UnicodeDecodeError: + try: + with open(file_path, 'r', encoding='gbk') as f: + content = f.read() + except UnicodeDecodeError: + util.log(1, f"无法解码文件: {file_name}") + continue + + if content.strip(): + knowledge_base[file_name] = content + util.log(1, f"成功加载知识库文件: {file_name} ({len(content)} 字符)") + + except Exception as e: + util.log(1, f"加载知识库文件 {file_name} 时出错: {str(e)}") + + return knowledge_base + +def search_knowledge_base(query, knowledge_base, max_results=3): + """ + 在知识库中搜索相关内容 + + 参数: + query: 查询内容 + knowledge_base: 知识库字典 + max_results: 最大返回结果数 + + 返回: + list: 相关内容列表 + """ + if not knowledge_base: + return [] + + results = [] + query_lower = query.lower() + + # 搜索关键词 + query_keywords = re.findall(r'\w+', query_lower) + + for file_name, content in knowledge_base.items(): + content_lower = content.lower() + + # 计算匹配度 + score = 0 + matched_sentences = [] + + # 按句子分割内容 + sentences = re.split(r'[。!?\n]', content) + + for sentence in sentences: + if not sentence.strip(): + continue + + sentence_lower = sentence.lower() + sentence_score = 0 + + # 计算关键词匹配度 + for keyword in query_keywords: + if keyword in sentence_lower: + sentence_score += 1 + + # 如果句子有匹配,记录 + if sentence_score > 0: + matched_sentences.append((sentence.strip(), sentence_score)) + score += sentence_score + + # 如果有匹配的内容 + if score > 0: + # 按匹配度排序句子 + matched_sentences.sort(key=lambda x: x[1], reverse=True) + + # 取前几个最相关的句子 + relevant_sentences = [sent[0] for sent in matched_sentences[:5] if sent[0]] + + if relevant_sentences: + results.append({ + 'file_name': file_name, + 'score': score, + 'content': '\n'.join(relevant_sentences) + }) + + # 按匹配度排序 + results.sort(key=lambda x: x['score'], reverse=True) + + return results[:max_results] + +# 全局知识库缓存 +_knowledge_base_cache = None +_knowledge_base_load_time = None +_knowledge_base_file_times = {} # 存储文件的最后修改时间 + +def check_knowledge_base_changes(): + """ + 检查知识库文件是否有变化 + + 返回: + bool: 如果有文件变化返回True,否则返回False + """ + global _knowledge_base_file_times + + # 获取llm/data目录路径 + current_dir = os.path.dirname(os.path.abspath(__file__)) + data_dir = os.path.join(current_dir, "data") + + if not os.path.exists(data_dir): + return False + + current_file_times = {} + + # 遍历data目录中的文件 + for file_path in Path(data_dir).iterdir(): + if not file_path.is_file(): + continue + + file_name = file_path.name + file_extension = file_path.suffix.lower() + + # 只检查支持的文件格式 + if file_extension in ['.docx', '.doc', '.pptx', '.txt'] or file_extension == '': + try: + mtime = os.path.getmtime(str(file_path)) + current_file_times[file_name] = mtime + except OSError: + continue + + # 检查是否有变化 + if not _knowledge_base_file_times: + # 第一次检查,保存文件时间 + _knowledge_base_file_times = current_file_times + return True + + # 比较文件时间 + if set(current_file_times.keys()) != set(_knowledge_base_file_times.keys()): + # 文件数量发生变化 + _knowledge_base_file_times = current_file_times + return True + + for file_name, mtime in current_file_times.items(): + if file_name not in _knowledge_base_file_times or _knowledge_base_file_times[file_name] != mtime: + # 文件被修改 + _knowledge_base_file_times = current_file_times + return True + + return False + +def init_knowledge_base(): + """ + 初始化知识库,在系统启动时调用 + """ + global _knowledge_base_cache, _knowledge_base_load_time + + util.log(1, "初始化本地知识库...") + _knowledge_base_cache = load_local_knowledge_base() + _knowledge_base_load_time = time.time() + + # 初始化文件修改时间跟踪 + check_knowledge_base_changes() + + util.log(1, f"知识库初始化完成,共 {len(_knowledge_base_cache)} 个文件") + +def get_knowledge_base(): + """ + 获取知识库,使用缓存机制 + + 返回: + dict: 知识库内容 + """ + global _knowledge_base_cache, _knowledge_base_load_time + + # 如果缓存为空,先初始化 + if _knowledge_base_cache is None: + init_knowledge_base() + return _knowledge_base_cache + + # 检查文件是否有变化 + if check_knowledge_base_changes(): + util.log(1, "检测到知识库文件变化,正在重新加载...") + _knowledge_base_cache = load_local_knowledge_base() + _knowledge_base_load_time = time.time() + util.log(1, f"知识库重新加载完成,共 {len(_knowledge_base_cache)} 个文件") + + return _knowledge_base_cache + + +def question(content, username, observation=None): + """处理用户提问并返回回复。""" + global current_username + current_username = username + full_response_text = "" + accumulated_text = "" + default_punctuations = [",", ".", "!", "?", "\n", "\uFF0C", "\u3002", "\uFF01", "\uFF1F"] + is_first_sentence = True + + from core import stream_manager + sm = stream_manager.new_instance() + conversation_id = sm.get_conversation_id(username) + + # 记忆系统已在全局初始化,无需创建agent + # 直接从配置文件获取人物设定 + agent_desc = { + "first_name": cfg.config["attribute"]["name"], + "last_name": "", + "age": cfg.config["attribute"]["age"], + "sex": cfg.config["attribute"]["gender"], + "additional": cfg.config["attribute"]["additional"], + "birthplace": cfg.config["attribute"]["birth"], + "position": cfg.config["attribute"]["position"], + "zodiac": cfg.config["attribute"]["zodiac"], + "constellation": cfg.config["attribute"]["constellation"], + "contact": cfg.config["attribute"]["contact"], + "voice": cfg.config["attribute"]["voice"], + "goal": cfg.config["attribute"]["goal"], + "occupation": cfg.config["attribute"]["job"], + "current_time": datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") + } + + # 使用新记忆系统处理用户消息 + # 一次性完成:入库 → 长期检索 → 短期检索 → 生成提示语 + short_term_records = [] + memory_prompt = "" + query_embedding = None + + try: + short_term_records, memory_prompt, query_embedding = memory_system.process_user_message( + content, user_id=username + ) + util.log(1, f"记忆检索成功,获取 {len(short_term_records)} 条相关记录") + except Exception as exc: + util.log(1, f"记忆检索失败: {exc}") + # 失败时使用空值,不影响后续流程 + short_term_records = [] + memory_prompt = "" + query_embedding = None + + knowledge_context = "" + try: + knowledge_base = get_knowledge_base() + if knowledge_base: + knowledge_results = search_knowledge_base(content, knowledge_base, max_results=3) + if knowledge_results: + parts = ["**本地知识库相关信息**:"] + for result in knowledge_results: + parts.append(f"来源文件:{result['file_name']}") + parts.append(result["content"]) + parts.append("") + knowledge_context = "\n".join(parts).strip() + util.log(1, f"找到 {len(knowledge_results)} 条相关知识库信息") + except Exception as exc: + util.log(1, f"搜索知识库时出错: {exc}") + + # 方案B:保留人设信息,补充记忆提示语 + # 1. 构建人设部分 + persona_prompt = f"""\n**角色设定**\n +- 名字:{agent_desc['first_name']} +- 性别:{agent_desc['sex']} +- 年龄:{agent_desc['age']} +- 职业:{agent_desc['occupation']} +- 出生地:{agent_desc['birthplace']} +- 星座:{agent_desc['constellation']} +- 生肖:{agent_desc['zodiac']} +- 联系方式:{agent_desc['contact']} +- 定位:{agent_desc['position']} +- 目标:{agent_desc['goal']} +- 补充信息:{agent_desc['additional']}\n + +""" + + # 2. 合并人设和记忆提示语 + if memory_prompt: + system_prompt = memory_prompt + persona_prompt + else: + # 如果记忆系统返回空提示语,使用基础提示语 + system_prompt = persona_prompt + "请根据用户的问题,提供有帮助的回答。" + + try: + history_records = content_db.new_instance().get_recent_messages_by_user(username=username, limit=30) + except Exception as exc: + util.log(1, f"加载历史消息失败: {exc}") + history_records = [] + + messages_buffer: List[ConversationMessage] = [] + + def append_to_buffer(role: str, text_value: str) -> None: + if not text_value: + return + messages_buffer.append({"role": role, "content": text_value}) + if len(messages_buffer) > 60: + del messages_buffer[:-60] + + for msg_type, msg_text in history_records: + role = 'assistant' + if msg_type and msg_type.lower() in ('member', 'user'): + role = 'user' + append_to_buffer(role, msg_text) + + if ( + not messages_buffer + or messages_buffer[-1]['role'] != 'user' + or messages_buffer[-1]['content'] != content + ): + append_to_buffer('user', content) + + tool_registry: Dict[str, WorkflowToolSpec] = {} + try: + mcp_tools = get_mcp_tools() + except Exception as exc: + util.log(1, f"获取工具列表失败: {exc}") + mcp_tools = [] + for tool_def in mcp_tools: + spec = _build_workflow_tool_spec(tool_def) + if spec: + tool_registry[spec.name] = spec + + try: + from utils.stream_state_manager import get_state_manager as _get_state_manager + + state_mgr = _get_state_manager() + session_label = "workflow_agent" if tool_registry else "llm_stream" + if not state_mgr.is_session_active(username, conversation_id=conversation_id): + state_mgr.start_new_session(username, session_label, conversation_id=conversation_id) + except Exception: + state_mgr = None + + try: + from utils.stream_text_processor import get_processor + + processor = get_processor() + punctuation_list = getattr(processor, "punctuation_marks", default_punctuations) + except Exception: + processor = None + punctuation_list = default_punctuations + def write_sentence(text: str, *, force_first: bool = False, force_end: bool = False) -> None: + if text is None: + text = "" + if not isinstance(text, str): + text = str(text) + if not text and not force_end and not force_first: + return + marked_text = None + if state_mgr is not None: + try: + marked_text, _, _ = state_mgr.prepare_sentence( + username, + text, + force_first=force_first, + force_end=force_end, + conversation_id=conversation_id, + ) + except Exception: + marked_text = None + if marked_text is None: + prefix = "_" if force_first else "" + suffix = "_" if force_end else "" + marked_text = f"{prefix}{text}{suffix}" + stream_manager.new_instance().write_sentence(username, marked_text, conversation_id=conversation_id) + + def stream_response_chunks(chunks, prepend_text: str = "") -> None: + nonlocal accumulated_text, full_response_text, is_first_sentence + if prepend_text: + accumulated_text += prepend_text + full_response_text += prepend_text + for chunk in chunks: + if sm.should_stop_generation(username, conversation_id=conversation_id): + util.log(1, f"检测到停止标志,中断文本生成: {username}") + break + if isinstance(chunk, str): + flush_text = chunk + elif isinstance(chunk, dict): + flush_text = chunk.get("content") + else: + flush_text = getattr(chunk, "content", None) + if isinstance(flush_text, list): + flush_text = "".join(part if isinstance(part, str) else "" for part in flush_text) + if not flush_text: + continue + flush_text = str(flush_text) + accumulated_text += flush_text + full_response_text += flush_text + if len(accumulated_text) >= 20: + while True: + last_punct_pos = -1 + for punct in punctuation_list: + pos = accumulated_text.rfind(punct) + if pos > last_punct_pos: + last_punct_pos = pos + if last_punct_pos > 10: + sentence_text = accumulated_text[: last_punct_pos + 1] + write_sentence(sentence_text, force_first=is_first_sentence) + is_first_sentence = False + accumulated_text = accumulated_text[last_punct_pos + 1 :].lstrip() + else: + break + + def finalize_stream(force_end: bool = False) -> None: + nonlocal accumulated_text, is_first_sentence + if accumulated_text: + write_sentence(accumulated_text, force_first=is_first_sentence, force_end=force_end) + is_first_sentence = False + accumulated_text = "" + elif force_end: + if state_mgr is not None: + try: + session_info = state_mgr.get_session_info(username, conversation_id=conversation_id) + except Exception: + session_info = None + if not session_info or not session_info.get("is_end_sent", False): + write_sentence("", force_end=True) + else: + write_sentence("", force_end=True) + + def run_workflow(tool_registry: Dict[str, WorkflowToolSpec]) -> bool: + nonlocal accumulated_text, full_response_text, is_first_sentence, messages_buffer + + initial_state: AgentState = { + "request": content, + "messages": messages_buffer, + "tool_results": [], + "audit_log": [], + "status": "planning", + "max_steps": 30, + "context": { + "system_prompt": system_prompt, + "knowledge_context": knowledge_context, + "observation": observation, + "tool_registry": tool_registry, + }, + } + + config = {"configurable": {"thread_id": f"workflow-{username}-{conversation_id}"}} + workflow_app = _WORKFLOW_APP + is_agent_think_start = False + final_state: Optional[AgentState] = None + final_stream_done = False + + try: + for event in workflow_app.stream(initial_state, config=config, stream_mode="updates"): + if sm.should_stop_generation(username, conversation_id=conversation_id): + util.log(1, f"检测到停止标志,中断工作流生成: {username}") + break + step, state = next(iter(event.items())) + final_state = state + status = state.get("status") + + state_messages = state.get("messages") or [] + if state_messages and len(state_messages) > len(messages_buffer): + messages_buffer.extend(state_messages[len(messages_buffer):]) + if len(messages_buffer) > 60: + del messages_buffer[:-60] + + if step == "plan_next": + if status == "needs_tool": + next_action = state.get("next_action") or {} + tool_name = next_action.get("name") or "unknown_tool" + tool_args = next_action.get("args") or {} + audit_log = state.get("audit_log") or [] + decision_note = audit_log[-1] if audit_log else "" + if "->" in decision_note: + decision_note = decision_note.split("->", 1)[1].strip() + args_text = json.dumps(tool_args, ensure_ascii=False) + message_lines = [ + "[PLAN] Planner preparing to call a tool.", + f"[PLAN] Decision: {decision_note}" if decision_note else "[PLAN] Decision: (missing)", + f"[PLAN] Tool: {tool_name}", + f"[PLAN] Args: {args_text}", + ] + message = "\n".join(message_lines) + "\n" + if not is_agent_think_start: + message = "" + message + is_agent_think_start = True + write_sentence(message, force_first=is_first_sentence) + is_first_sentence = False + full_response_text += message + append_to_buffer('assistant', message.strip()) + elif status == "completed" and not final_stream_done: + closing = "" if is_agent_think_start else "" + final_messages = state.get("final_messages") + final_response = state.get("final_response") + success = False + if final_messages: + try: + stream_response_chunks(llm.stream(final_messages), prepend_text=closing) + success = True + except requests.exceptions.RequestException as stream_exc: + util.log(1, f"最终回复流式输出失败: {stream_exc}") + elif final_response: + stream_response_chunks([closing + final_response]) + success = True + elif closing: + accumulated_text += closing + full_response_text += closing + final_stream_done = success + is_agent_think_start = False + elif step == "call_tool": + history = state.get("tool_results") or [] + if history: + last = history[-1] + call_info = last.get("call", {}) or {} + tool_name = call_info.get("name") or "unknown_tool" + success = last.get("success", False) + status_text = "SUCCESS" if success else "FAILED" + args_text = json.dumps(call_info.get("args") or {}, ensure_ascii=False) + message_lines = [ + f"[TOOL] {tool_name} execution {status_text}.", + f"[TOOL] Args: {args_text}", + ] + if last.get("output"): + message_lines.append(f"[TOOL] Output: {_truncate_text(last['output'], 120)}") + if last.get("error"): + message_lines.append(f"[TOOL] Error: {last['error']}") + message = "\n".join(message_lines) + "\n" + write_sentence(message, force_first=is_first_sentence) + is_first_sentence = False + full_response_text += message + append_to_buffer('assistant', message.strip()) + elif step == "__end__": + break + except Exception as exc: + util.log(1, f"执行工具工作流时出错: {exc}") + if is_agent_think_start: + closing = "" + accumulated_text += closing + full_response_text += closing + return False + + if final_state is None: + if is_agent_think_start: + closing = "" + accumulated_text += closing + full_response_text += closing + return False + + if not final_stream_done and is_agent_think_start: + closing = "" + accumulated_text += closing + full_response_text += closing + util.log(1, f"工具工作流未能完成,状态: {final_state.get('status')}") + + final_state_messages = final_state.get("messages") if final_state else None + if final_state_messages and len(final_state_messages) > len(messages_buffer): + messages_buffer.extend(final_state_messages[len(messages_buffer):]) + if len(messages_buffer) > 60: + del messages_buffer[:-60] + + return final_stream_done + + def run_direct_llm() -> bool: + nonlocal full_response_text, accumulated_text, is_first_sentence, messages_buffer + try: + # 统一使用 _build_final_messages 构建消息,确保历史对话始终被包含 + summary_state: AgentState = { + "request": content, + "messages": messages_buffer, + "tool_results": [], + "planner_preview": None, + "context": { + "system_prompt": system_prompt, + "knowledge_context": knowledge_context, + "observation": observation, + }, + } + + final_messages = _build_final_messages(summary_state) + stream_response_chunks(llm.stream(final_messages)) + return True + except requests.exceptions.RequestException as exc: + util.log(1, f"请求失败: {exc}") + error_message = "抱歉,我现在太忙了,休息一会,请稍后再试。" + write_sentence(error_message, force_first=is_first_sentence) + is_first_sentence = False + full_response_text = error_message + accumulated_text = "" + return False + + workflow_success = False + if tool_registry: + workflow_success = run_workflow(tool_registry) + + if (not tool_registry or not workflow_success) and not sm.should_stop_generation(username, conversation_id=conversation_id): + run_direct_llm() + + if not sm.should_stop_generation(username, conversation_id=conversation_id): + finalize_stream(force_end=True) + + if state_mgr is not None: + try: + state_mgr.end_session(username, conversation_id=conversation_id) + except Exception: + pass + else: + try: + from utils.stream_state_manager import get_state_manager + + get_state_manager().end_session(username, conversation_id=conversation_id) + except Exception: + pass + + final_text = full_response_text.split("")[-1] if full_response_text else "" + + # 使用新记忆系统异步处理agent回复 + try: + import asyncio + + # 创建新的事件循环(在独立线程中运行) + def async_memory_task(): + """在独立线程中运行异步记忆存储""" + try: + # 创建新的事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + # 运行异步任务 + loop.run_until_complete( + memory_system.process_agent_reply_async( + final_text, + user_id=username, + current_user_content=content + ) + ) + + # 关闭循环 + loop.close() + except Exception as e: + util.log(1, f"异步记忆存储失败: {e}") + + # 启动独立线程执行异步任务 + MyThread(target=async_memory_task).start() + util.log(1, f"异步记忆存储任务已启动") + + except Exception as exc: + util.log(1, f"异步记忆处理启动失败: {exc}") + + return final_text +def clear_agent_memory(username=None): + """ + 清除指定用户的记忆(使用新记忆系统) + + Args: + username: 用户名,如果为None则清除当前用户的记忆 + + Returns: + bool: 是否清除成功 + """ + global memory_system, current_username + + try: + # 确定要清除的用户ID + user_id = username if username else current_username + if not user_id: + user_id = "User" # 默认用户 + + util.log(1, f"正在清除用户 {user_id} 的记忆...") + + # 调用新记忆系统的清除方法 + result = memory_system.clear_user_history(user_id=user_id) + + util.log(1, f"用户 {user_id} 的记忆清除完成: {result}") + return True + + except Exception as e: + util.log(1, f"清除用户记忆时出错: {str(e)}") + return False + +def get_mcp_tools() -> List[Dict[str, Any]]: + """ + 从共享缓存获取所有可用且已启用的MCP工具列表。 + """ + try: + tools = mcp_tool_registry.get_enabled_tools() + return tools or [] + except Exception as e: + util.log(1, f"获取工具列表出错:{e}") + return [] + + +if __name__ == "__main__": + # 记忆系统已在模块加载时初始化,无需再次调用 + for _ in range(3): + query = "Who is Fay?" + response = question(query, "User") + print(f"Q: {query}") + print(f"A: {response}") + time.sleep(1) diff --git a/mcp_servers/schedule_manager/server.py b/mcp_servers/schedule_manager/server.py index 83ba94f..ca8b059 100644 --- a/mcp_servers/schedule_manager/server.py +++ b/mcp_servers/schedule_manager/server.py @@ -577,7 +577,6 @@ class ScheduleManager: def send_to_fay(self, message: str, uid: int = 0): """发送消息给Fay - 使用v1/chat/completions接口""" - print("***********************************************************************") logger.info(f"[DEBUG] send_to_fay 被调用,消息: {message}, uid: {uid}") # 防止消息重复发送 diff --git a/requirements.txt b/requirements.txt index c3ee531..7281e61 100644 --- a/requirements.txt +++ b/requirements.txt @@ -30,4 +30,6 @@ bs4 schedule mcp python-docx -python-pptx \ No newline at end of file +python-pptx +chromadb +sentence_transformers \ No newline at end of file diff --git a/test/test_fay_gpt_nonstream.py b/test/test_fay_gpt_nonstream.py index 85e6a6b..06a8c0f 100644 --- a/test/test_fay_gpt_nonstream.py +++ b/test/test_fay_gpt_nonstream.py @@ -2,15 +2,15 @@ import requests import json def test_gpt_nonstream(prompt): - url = 'http://127.0.0.1:5000/v1/chat/completions' # 替换为您的接口地址 + url = 'http://127.0.0.1:8000/v1/chat/completions' # 替换为您的接口地址 headers = { 'Content-Type': 'application/json', 'Authorization': f'Bearer YOUR_API_KEY', # 如果您的接口需要身份验证 } data = { - 'model': 'fay', + 'model': 'moonshotai/Kimi-K2-Instruct-0905', 'messages': [ - {'role': '小敏', 'content': prompt} + {'role': 'system', 'content': prompt} ], 'stream': False # 禁用流式传输,使用非流式响应 } diff --git a/test/test_fay_gpt_stream.py b/test/test_fay_gpt_stream.py index 050c389..249fe4b 100644 --- a/test/test_fay_gpt_stream.py +++ b/test/test_fay_gpt_stream.py @@ -2,7 +2,7 @@ import requests import json def test_gpt(prompt): - url = 'http://127.0.0.1:5000/v1/chat/completions' # 替换为您的接口地址 + url = 'http://127.0.0.1:8000/v1/chat/completions' # 替换为您的接口地址 headers = { 'Content-Type': 'application/json', 'Authorization': f'Bearer YOUR_API_KEY', # 如果您的接口需要身份验证 @@ -10,7 +10,7 @@ def test_gpt(prompt): data = { 'model': 'fay-streaming', 'messages': [ - {'role': 'User', 'content': prompt} + {'role': 'user', 'content': prompt} ], 'stream': True # 启用流式传输 } @@ -46,7 +46,7 @@ def test_gpt(prompt): print(f"\n收到未知格式的数据:{line}") if __name__ == "__main__": - user_input = "你好" + user_input = "7" print("GPT 的回复:") test_gpt(user_input) print("\n请求完成") diff --git a/utils/config_util.py b/utils/config_util.py index 65baa3d..fad4a18 100644 --- a/utils/config_util.py +++ b/utils/config_util.py @@ -53,6 +53,12 @@ start_mode = None fay_url = None system_conf_path = None config_json_path = None +use_bionic_memory = None + +# Embedding API 配置全局变量 +embedding_api_model = None +embedding_api_base_url = None +embedding_api_key = None # config server中心配置,system.conf与config.json存在时不会使用配置中心 CONFIG_SERVER = { @@ -176,6 +182,10 @@ def load_config(): global volcano_tts_voice_type global start_mode global fay_url + global use_bionic_memory + global embedding_api_model + global embedding_api_base_url + global embedding_api_key global CONFIG_SERVER global system_conf_path @@ -239,6 +249,11 @@ def load_config(): volcano_tts_cluster = system_config.get('key', 'volcano_tts_cluster', fallback=None) volcano_tts_voice_type = system_config.get('key', 'volcano_tts_voice_type', fallback=None) + # 读取 Embedding API 配置(复用 LLM 的 url 和 key) + embedding_api_model = system_config.get('key', 'embedding_api_model', fallback='BAAI/bge-large-zh-v1.5') + embedding_api_base_url = gpt_base_url # 复用 LLM base_url + embedding_api_key = key_gpt_api_key # 复用 LLM api_key + start_mode = system_config.get('key', 'start_mode', fallback=None) fay_url = system_config.get('key', 'fay_url', fallback=None) # 如果fay_url为空或None,则动态获取本机IP地址 @@ -254,7 +269,10 @@ def load_config(): # 读取用户配置 with codecs.open(config_json_path, encoding='utf-8') as f: config = json.load(f) - + + # 读取仿生记忆配置 + use_bionic_memory = config.get('memory', {}).get('use_bionic_memory', False) + # 构建配置字典 config_dict = { 'system_config': system_config, @@ -287,6 +305,13 @@ def load_config(): 'start_mode': start_mode, 'fay_url': fay_url, + 'use_bionic_memory': use_bionic_memory, + + # Embedding API 配置 + 'embedding_api_model': embedding_api_model, + 'embedding_api_base_url': embedding_api_base_url, + 'embedding_api_key': embedding_api_key, + 'source': 'local' # 标记配置来源 }