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 @@