feat: add neo4j search and chinese readme

This commit is contained in:
duwanze 2026-06-15 09:04:55 +08:00
parent e1241c8f5c
commit db5b958e50
36 changed files with 5389 additions and 152 deletions

View File

@ -5,12 +5,20 @@ LLM_API_KEY=your_api_key_here
LLM_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
LLM_MODEL_NAME=qwen-plus
# ===== ZEP记忆图谱配置 =====
# 每月免费额度即可支撑简单使用https://app.getzep.com/
# ===== 图数据库配置 =====
# 默认使用本地 Neo4j 替代 Zep Cloud
GRAPH_BACKEND=neo4j
NEO4J_URI=bolt://localhost:7687
NEO4J_USERNAME=neo4j
NEO4J_PASSWORD=password
NEO4J_DATABASE=neo4j
# ===== ZEP记忆图谱配置可选旧后端 =====
# 若设置 GRAPH_BACKEND=zep则需要配置此项
ZEP_API_KEY=your_zep_api_key_here
# ===== 加速 LLM 配置(可选)=====
# 注意如果不使用加速配置env文件中就不要出现下面的配置项
LLM_BOOST_API_KEY=your_api_key_here
LLM_BOOST_BASE_URL=your_base_url_here
LLM_BOOST_MODEL_NAME=your_model_name_here
LLM_BOOST_MODEL_NAME=your_model_name_here

8
.env.neo4j.example Normal file
View File

@ -0,0 +1,8 @@
# Neo4j 连接配置
NEO4J_URI=bolt://localhost:7687
NEO4J_USERNAME=neo4j
NEO4J_PASSWORD=password
NEO4J_DATABASE=neo4j
# 图数据库后端选择: 'zep' 或 'neo4j'
GRAPH_BACKEND=neo4j

13
.gitignore vendored
View File

@ -29,6 +29,8 @@ node_modules/
npm-debug.log*
yarn-debug.log*
yarn-error.log*
npm_cache/
frontend/npm_cache/
# IDE
.vscode/
@ -52,9 +54,18 @@ mytest/
# 日志文件
backend/logs/
*.log
log/
# 上传文件
backend/uploads/
# 本地依赖缓存
pip_cache/
backend/pip_cache/
# 本地环境配置
.env.dev
.env.neo4j
# Docker 数据
data/
data/

View File

@ -2,6 +2,8 @@
[English](./README.md) | 中文
<a href="https://trendshift.io/repositories/16144" target="_blank"><img src="https://trendshift.io/api/badge/repositories/16144" alt="MiroFish | Trendshift" style="width: 250px; height: 55px;" width="250" height="55"/></a>
本项目基于原始开源仓库 [666ghj/MiroFish](https://github.com/666ghj/MiroFish) 进行二次开发。
本版本的核心改动是:将原项目中的 Zep Cloud 图谱记忆与检索依赖替换为本地 Neo4j 后端,使项目可以使用本地图数据库完成实体关系存储、图谱检索、报告工具调用和模拟后的记忆更新,同时保留原 MiroFish 的多智能体模拟流程。
@ -233,4 +235,4 @@ docker compose -f docker-compose.neo4j.yml ps
## 致谢
本项目基于 [666ghj/MiroFish](https://github.com/666ghj/MiroFish) 进行二次开发。感谢原作者及所有贡献者的开源工作。
本项目基于 [666ghj/MiroFish](https://github.com/666ghj/MiroFish) 进行二次开发。感谢原作者及所有贡献者的开源工作。

View File

@ -11,13 +11,13 @@ from flask import request, jsonify
from . import graph_bp
from ..config import Config
from ..services.ontology_generator import OntologyGenerator
from ..services.graph_builder import GraphBuilderService
from ..services.text_processor import TextProcessor
from ..services.graph_service_factory import get_graph_factory
from ..utils.file_parser import FileParser
from ..utils.logger import get_logger
from ..utils.locale import t, get_locale, set_locale
from ..models.task import TaskManager, TaskStatus
from ..models.project import ProjectManager, ProjectStatus
from ..models.project import ProjectManager, ProjectStatus, validate_simulation_requirement
# 获取日志器
logger = get_logger('mirofish.api')
@ -163,6 +163,13 @@ def generate_ontology():
"success": False,
"error": t('api.requireSimulationRequirement')
}), 400
requirement_error = validate_simulation_requirement(simulation_requirement)
if requirement_error:
return jsonify({
"success": False,
"error": requirement_error
}), 400
# 获取上传的文件
uploaded_files = request.files.getlist('files')
@ -283,10 +290,15 @@ def build_graph():
try:
logger.info("=== 开始构建图谱 ===")
# 检查配置
# 检查配置 - 根据后端检查相应配置
errors = []
if not Config.ZEP_API_KEY:
errors.append(t('api.zepApiKeyMissing'))
factory = get_graph_factory()
health = factory.check_backend_health()
if not health.get('healthy'):
if factory.backend == 'zep':
errors.append(t('api.zepApiKeyMissing'))
else:
errors.append(f"Neo4j 连接失败,请检查配置")
if errors:
logger.error(f"配置错误: {errors}")
return jsonify({
@ -387,7 +399,7 @@ def build_graph():
)
# 创建图谱构建服务
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
builder = get_graph_factory().get_graph_builder()
# 分块
task_manager.update_task(
@ -572,20 +584,15 @@ def get_graph_data(graph_id: str):
获取图谱数据节点和边
"""
try:
if not Config.ZEP_API_KEY:
return jsonify({
"success": False,
"error": t('api.zepApiKeyMissing')
}), 500
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
factory = get_graph_factory()
builder = factory.get_graph_builder()
graph_data = builder.get_graph_data(graph_id)
return jsonify({
"success": True,
"data": graph_data
})
except Exception as e:
return jsonify({
"success": False,
@ -600,20 +607,15 @@ def delete_graph(graph_id: str):
删除Zep图谱
"""
try:
if not Config.ZEP_API_KEY:
return jsonify({
"success": False,
"error": t('api.zepApiKeyMissing')
}), 500
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
factory = get_graph_factory()
builder = factory.get_graph_builder()
builder.delete_graph(graph_id)
return jsonify({
"success": True,
"message": t('api.graphDeleted', id=graph_id)
})
except Exception as e:
return jsonify({
"success": False,

View File

@ -957,9 +957,9 @@ def search_graph_tool():
"error": t('api.requireGraphIdAndQuery')
}), 400
from ..services.zep_tools import ZepToolsService
tools = ZepToolsService()
from ..services.graph_service_factory import get_graph_factory
tools = get_graph_factory().get_search_service()
result = tools.search_graph(
graph_id=graph_id,
query=query,
@ -1001,9 +1001,9 @@ def get_graph_statistics_tool():
"error": t('api.requireGraphId')
}), 400
from ..services.zep_tools import ZepToolsService
tools = ZepToolsService()
from ..services.graph_service_factory import get_graph_factory
tools = get_graph_factory().get_search_service()
result = tools.get_graph_statistics(graph_id)
return jsonify({

View File

@ -9,13 +9,13 @@ from flask import request, jsonify, send_file
from . import simulation_bp
from ..config import Config
from ..services.zep_entity_reader import ZepEntityReader
from ..services.oasis_profile_generator import OasisProfileGenerator
from ..services.simulation_manager import SimulationManager, SimulationStatus
from ..services.simulation_runner import SimulationRunner, RunnerStatus
from ..services.graph_service_factory import get_graph_factory
from ..utils.logger import get_logger
from ..utils.locale import t, get_locale, set_locale
from ..models.project import ProjectManager
from ..models.project import ProjectManager, validate_simulation_requirement
logger = get_logger('mirofish.api.simulation')
@ -25,6 +25,19 @@ logger = get_logger('mirofish.api.simulation')
INTERVIEW_PROMPT_PREFIX = "结合你的人设、所有的过往记忆与行动,不调用任何工具直接用文本回复我:"
def _graph_backend_error_response():
factory = get_graph_factory()
health = factory.check_backend_health()
if health.get("healthy"):
return None
if factory.backend == "zep":
error = t('api.zepApiKeyMissing')
else:
error = "Neo4j 连接失败,请检查 NEO4J_URI/NEO4J_USERNAME/NEO4J_PASSWORD 配置"
return jsonify({"success": False, "error": error}), 500
def optimize_interview_prompt(prompt: str) -> str:
"""
优化Interview提问添加前缀避免Agent调用工具
@ -57,11 +70,9 @@ def get_graph_entities(graph_id: str):
enrich: 是否获取相关边信息默认true
"""
try:
if not Config.ZEP_API_KEY:
return jsonify({
"success": False,
"error": t('api.zepApiKeyMissing')
}), 500
backend_error = _graph_backend_error_response()
if backend_error:
return backend_error
entity_types_str = request.args.get('entity_types', '')
entity_types = [t.strip() for t in entity_types_str.split(',') if t.strip()] if entity_types_str else None
@ -69,7 +80,7 @@ def get_graph_entities(graph_id: str):
logger.info(f"获取图谱实体: graph_id={graph_id}, entity_types={entity_types}, enrich={enrich}")
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
result = reader.filter_defined_entities(
graph_id=graph_id,
defined_entity_types=entity_types,
@ -94,13 +105,11 @@ def get_graph_entities(graph_id: str):
def get_entity_detail(graph_id: str, entity_uuid: str):
"""获取单个实体的详细信息"""
try:
if not Config.ZEP_API_KEY:
return jsonify({
"success": False,
"error": t('api.zepApiKeyMissing')
}), 500
backend_error = _graph_backend_error_response()
if backend_error:
return backend_error
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
entity = reader.get_entity_with_context(graph_id, entity_uuid)
if not entity:
@ -127,15 +136,13 @@ def get_entity_detail(graph_id: str, entity_uuid: str):
def get_entities_by_type(graph_id: str, entity_type: str):
"""获取指定类型的所有实体"""
try:
if not Config.ZEP_API_KEY:
return jsonify({
"success": False,
"error": t('api.zepApiKeyMissing')
}), 500
backend_error = _graph_backend_error_response()
if backend_error:
return backend_error
enrich = request.args.get('enrich', 'true').lower() == 'true'
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
entities = reader.get_entities_by_type(
graph_id=graph_id,
entity_type=entity_type,
@ -207,6 +214,13 @@ def create_simulation():
"success": False,
"error": t('api.projectNotFound', id=project_id)
}), 404
requirement_error = validate_simulation_requirement(project.simulation_requirement or "")
if requirement_error:
return jsonify({
"success": False,
"error": requirement_error
}), 400
graph_id = data.get('graph_id') or project.graph_id
if not graph_id:
@ -472,7 +486,7 @@ def prepare_simulation():
# 这样前端在调用prepare后立即就能获取到预期Agent总数
try:
logger.info(f"同步获取实体数量: graph_id={state.graph_id}")
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
# 快速读取实体(不需要边信息,只统计数量)
filtered_preview = reader.filter_defined_entities(
graph_id=state.graph_id,
@ -1401,7 +1415,7 @@ def generate_profiles():
use_llm = data.get('use_llm', True)
platform = data.get('platform', 'reddit')
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
filtered = reader.filter_defined_entities(
graph_id=graph_id,
defined_entity_types=entity_types,

View File

@ -31,9 +31,23 @@ class Config:
LLM_API_KEY = os.environ.get('LLM_API_KEY')
LLM_BASE_URL = os.environ.get('LLM_BASE_URL', 'https://api.openai.com/v1')
LLM_MODEL_NAME = os.environ.get('LLM_MODEL_NAME', 'gpt-4o-mini')
LLM_RATE_LIMIT_MAX_ATTEMPTS = int(os.environ.get('LLM_RATE_LIMIT_MAX_ATTEMPTS', '20'))
LLM_RATE_LIMIT_INITIAL_DELAY = float(os.environ.get('LLM_RATE_LIMIT_INITIAL_DELAY', '30'))
LLM_RATE_LIMIT_MAX_DELAY = float(os.environ.get('LLM_RATE_LIMIT_MAX_DELAY', '180'))
LLM_RATE_LIMIT_BACKOFF_FACTOR = float(os.environ.get('LLM_RATE_LIMIT_BACKOFF_FACTOR', '1.5'))
# Zep配置
ZEP_API_KEY = os.environ.get('ZEP_API_KEY')
# Neo4j配置替代Zep的本地图数据库
NEO4J_URI = os.environ.get('NEO4J_URI', 'bolt://localhost:7687')
NEO4J_USERNAME = os.environ.get('NEO4J_USERNAME', 'neo4j')
NEO4J_PASSWORD = os.environ.get('NEO4J_PASSWORD', '')
NEO4J_DATABASE = os.environ.get('NEO4J_DATABASE', 'neo4j')
NEO4J_MAX_POOL_SIZE = int(os.environ.get('NEO4J_MAX_POOL_SIZE', '50'))
# 选择图数据库后端:'zep' 或 'neo4j'
GRAPH_BACKEND = os.environ.get('GRAPH_BACKEND', 'neo4j').lower()
# 文件上传配置
MAX_CONTENT_LENGTH = 50 * 1024 * 1024 # 50MB
@ -69,7 +83,9 @@ class Config:
errors: list[str] = []
if not cls.LLM_API_KEY:
errors.append("LLM_API_KEY 未配置")
if not cls.ZEP_API_KEY:
if cls.GRAPH_BACKEND == 'zep' and not cls.ZEP_API_KEY:
errors.append("ZEP_API_KEY 未配置")
if cls.GRAPH_BACKEND == 'neo4j' and not cls.NEO4J_PASSWORD:
errors.append("NEO4J_PASSWORD 未配置")
return errors

View File

@ -14,6 +14,41 @@ from dataclasses import dataclass, field, asdict
from ..config import Config
CONTAMINATED_REQUIREMENT_KEYWORDS = (
"No module named 'camel'",
"缺少依赖",
"进程退出码",
"runner_status",
"insight_forge",
"panorama_search",
"quick_search",
"interview_agents",
"Session.run() got multiple values",
"to_text",
"Neo4jSearchService",
"模拟系统故障",
"模拟启动失败",
"系统启动失败",
)
def validate_simulation_requirement(requirement: str) -> Optional[str]:
"""拒绝把内部错误日志或故障报告误当作模拟需求。"""
text = (requirement or "").strip()
if not text:
return "模拟需求不能为空"
for keyword in CONTAMINATED_REQUIREMENT_KEYWORDS:
if keyword.lower() in text.lower():
return (
"模拟需求疑似包含内部错误日志或故障报告,请重新填写真实预测需求。"
"例如:预测 Qatar vs Switzerland 足球比赛的舆论反应、双方支持者观点、"
"媒体关注点、潜在争议和胜负倾向。"
)
return None
class ProjectStatus(str, Enum):
"""项目状态"""
CREATED = "created" # 刚创建,文件已上传

View File

@ -1,15 +1,16 @@
"""
业务服务模块
注意Zep 相关模块已移至 adapters/通过 graph_service_factory 按需加载
不需要安装 zep_cloud Neo4j 模式下直接使用本地 Neo4j
"""
from .ontology_generator import OntologyGenerator
from .graph_builder import GraphBuilderService
from .text_processor import TextProcessor
from .zep_entity_reader import ZepEntityReader, EntityNode, FilteredEntities
from .oasis_profile_generator import OasisProfileGenerator, OasisAgentProfile
from .simulation_manager import SimulationManager, SimulationState, SimulationStatus
from .simulation_config_generator import (
SimulationConfigGenerator,
SimulationConfigGenerator,
SimulationParameters,
AgentActivityConfig,
TimeSimulationConfig,
@ -23,11 +24,6 @@ from .simulation_runner import (
AgentAction,
RoundSummary
)
from .zep_graph_memory_updater import (
ZepGraphMemoryUpdater,
ZepGraphMemoryManager,
AgentActivity
)
from .simulation_ipc import (
SimulationIPCClient,
SimulationIPCServer,
@ -38,12 +34,8 @@ from .simulation_ipc import (
)
__all__ = [
'OntologyGenerator',
'GraphBuilderService',
'OntologyGenerator',
'TextProcessor',
'ZepEntityReader',
'EntityNode',
'FilteredEntities',
'OasisProfileGenerator',
'OasisAgentProfile',
'SimulationManager',
@ -60,14 +52,10 @@ __all__ = [
'RunnerStatus',
'AgentAction',
'RoundSummary',
'ZepGraphMemoryUpdater',
'ZepGraphMemoryManager',
'AgentActivity',
'SimulationIPCClient',
'SimulationIPCServer',
'IPCCommand',
'IPCResponse',
'CommandType',
'CommandStatus',
]
]

View File

@ -0,0 +1,49 @@
"""
Adapters Package
图谱适配器实现支持切换底层图数据库Zep / Neo4j
"""
from .graph_adapter import GraphAdapter, GraphNode, GraphEdge, GraphInfo, SearchResult
from .llm_extractor import LLMExtractionPipeline, LLMEntityEnricher, ExtractedEntity, ExtractedEdge
from .neo4j_graph_builder import Neo4jGraphBuilder, Neo4jAsyncGraphBuilder
from .neo4j_entity_reader import Neo4jEntityReader, EntityNode, FilteredEntities
from .neo4j_search_service import Neo4jSearchService, NodeInfo, EdgeInfo
from .neo4j_graph_memory_updater import (
Neo4jGraphMemoryUpdater,
Neo4jGraphMemoryManager,
AgentActivity
)
__all__ = [
# 核心接口
'GraphAdapter',
'GraphNode',
'GraphEdge',
'GraphInfo',
'SearchResult',
# LLM 提取
'LLMExtractionPipeline',
'LLMEntityEnricher',
'ExtractedEntity',
'ExtractedEdge',
# Neo4j 图谱构建
'Neo4jGraphBuilder',
'Neo4jAsyncGraphBuilder',
# Neo4j 实体读取
'Neo4jEntityReader',
'EntityNode',
'FilteredEntities',
# Neo4j 搜索服务
'Neo4jSearchService',
'NodeInfo',
'EdgeInfo',
# Neo4j 记忆更新
'Neo4jGraphMemoryUpdater',
'Neo4jGraphMemoryManager',
'AgentActivity',
]

View File

@ -0,0 +1,275 @@
"""
图谱适配器接口定义
抽象 Zep Cloud Neo4j 的共同操作支持底层图数据库切换
"""
from abc import ABC, abstractmethod
from typing import Dict, Any, List, Optional, Callable
from dataclasses import dataclass
@dataclass
class GraphNode:
"""图谱节点"""
uuid: str
name: str
labels: List[str]
summary: str
attributes: Dict[str, Any]
created_at: Optional[str] = None
@dataclass
class GraphEdge:
"""图谱边"""
uuid: str
name: str # 关系类型名称
fact: str # 事实描述
source_node_uuid: str
target_node_uuid: str
attributes: Dict[str, Any]
# 时间信息
created_at: Optional[str] = None
valid_at: Optional[str] = None
invalid_at: Optional[str] = None
expired_at: Optional[str] = None
@dataclass
class GraphInfo:
"""图谱信息"""
graph_id: str
node_count: int
edge_count: int
entity_types: List[str]
def to_dict(self) -> Dict[str, Any]:
return {
"graph_id": self.graph_id,
"node_count": self.node_count,
"edge_count": self.edge_count,
"entity_types": self.entity_types,
}
@dataclass
class SearchResult:
"""搜索结果"""
facts: List[str]
edges: List[Dict[str, Any]]
nodes: List[Dict[str, Any]]
query: str
total_count: int
def to_dict(self) -> Dict[str, Any]:
return {
"facts": self.facts,
"edges": self.edges,
"nodes": self.nodes,
"query": self.query,
"total_count": self.total_count,
}
def to_text(self) -> str:
text_parts = [
f"搜索查询: {self.query}",
f"找到 {self.total_count} 条相关信息",
]
if self.facts:
text_parts.append("\n### 相关事实:")
for i, fact in enumerate(self.facts, 1):
text_parts.append(f"{i}. {fact}")
if self.nodes:
text_parts.append("\n### 相关节点:")
for i, node in enumerate(self.nodes, 1):
name = node.get("name", "未知实体")
labels = ", ".join(node.get("labels") or [])
summary = node.get("summary") or ""
text_parts.append(f"- **{name}** ({labels})")
if summary:
text_parts.append(f" 摘要: {summary}")
if self.edges:
text_parts.append("\n### 相关边:")
for edge in self.edges:
name = edge.get("name", "关系")
fact = edge.get("fact") or ""
source = edge.get("source_name") or edge.get("source_node_name") or edge.get("source_node_uuid", "")[:8]
target = edge.get("target_name") or edge.get("target_node_name") or edge.get("target_node_uuid", "")[:8]
text_parts.append(f"- {source} --[{name}]--> {target}")
if fact:
text_parts.append(f" 事实: {fact}")
return "\n".join(text_parts)
class GraphAdapter(ABC):
"""
图谱适配器接口
定义所有图谱操作的抽象接口Zep Neo4j 都必须实现这些方法
这样上层服务可以透明地切换底层图数据库
"""
@abstractmethod
def create_graph(self, name: str) -> str:
"""
创建新图谱
Args:
name: 图谱名称
Returns:
graph_id: 新创建的图谱ID
"""
pass
@abstractmethod
def set_ontology(
self,
graph_id: str,
entity_types: List[Dict[str, Any]],
edge_types: List[Dict[str, Any]]
) -> None:
"""
设置图谱本体Schema
Args:
graph_id: 图谱ID
entity_types: 实体类型定义列表
edge_types: 关系类型定义列表
"""
pass
@abstractmethod
def add_text_batches(
self,
graph_id: str,
chunks: List[str],
batch_size: int = 3,
progress_callback: Optional[Callable] = None
) -> List[str]:
"""
批量添加文本到图谱
Args:
graph_id: 图谱ID
chunks: 文本块列表
batch_size: 每批发送的块数量
progress_callback: 进度回调函数
Returns:
episode_uuids: 所有文本块的UUID列表
"""
pass
@abstractmethod
def get_all_nodes(self, graph_id: str) -> List[GraphNode]:
"""
获取图谱的所有节点
Args:
graph_id: 图谱ID
Returns:
节点列表
"""
pass
@abstractmethod
def get_all_edges(self, graph_id: str) -> List[GraphEdge]:
"""
获取图谱的所有边
Args:
graph_id: 图谱ID
Returns:
边列表
"""
pass
@abstractmethod
def get_node(self, node_uuid: str) -> Optional[GraphNode]:
"""
获取单个节点
Args:
node_uuid: 节点UUID
Returns:
节点对象或None
"""
pass
@abstractmethod
def get_node_edges(self, node_uuid: str) -> List[GraphEdge]:
"""
获取指定节点的所有相关边
Args:
node_uuid: 节点UUID
Returns:
边列表
"""
pass
@abstractmethod
def search(
self,
graph_id: str,
query: str,
limit: int = 10,
scope: str = "edges"
) -> SearchResult:
"""
图谱语义搜索
Args:
graph_id: 图谱ID
query: 搜索查询
limit: 返回结果数量
scope: 搜索范围 ("edges" / "nodes" / "both")
Returns:
SearchResult: 搜索结果
"""
pass
@abstractmethod
def add_activities(
self,
graph_id: str,
activities: List[str]
) -> None:
"""
添加活动记录到图谱
Args:
graph_id: 图谱ID
activities: 活动描述文本列表
"""
pass
@abstractmethod
def delete_graph(self, graph_id: str) -> None:
"""
删除图谱
Args:
graph_id: 图谱ID
"""
pass
@abstractmethod
def health_check(self) -> bool:
"""
健康检查
Returns:
True 如果连接正常
"""
pass

View File

@ -0,0 +1,416 @@
"""
LLM 文本提取 Pipeline
使用 LLM 将文本转换为结构化的实体和关系
替代 Zep Cloud add_batch() 自动提取功能
"""
import json
import time
import uuid
from typing import Dict, Any, List, Optional, Callable, Tuple
from dataclasses import dataclass, field
from concurrent.futures import ThreadPoolExecutor, as_completed
from ...config import Config
from ...utils.llm_client import LLMClient
from ...utils.logger import get_logger
logger = get_logger('mirofish.llm_extractor')
@dataclass
class ExtractedEntity:
"""提取的实体"""
name: str
entity_type: str # e.g., "Student", "Person"
summary: str = ""
attributes: Dict[str, Any] = field(default_factory=dict)
source_chunk: str = "" # 来源文本块
@dataclass
class ExtractedEdge:
"""提取的关系"""
name: str # e.g., "STUDIES_AT", "COMMENTS_ON"
source_name: str # 源实体名称
target_name: str # 目标实体名称
fact: str = "" # 事实描述
source_type: str = "" # 源实体类型
target_type: str = "" # 目标实体类型
attributes: Dict[str, Any] = field(default_factory=dict)
@dataclass
class ExtractionResult:
"""提取结果"""
entities: List[ExtractedEntity]
edges: List[ExtractedEdge]
chunk_index: int = 0
chunk_text: str = ""
error: Optional[str] = None
def is_valid(self) -> bool:
return len(self.entities) > 0 or len(self.edges) > 0
class LLMExtractionPipeline:
"""
LLM 文本提取 Pipeline
将非结构化文本通过 LLM 提取为结构化的实体和关系
适配 Neo4j 的节点和边格式
"""
# 每次 LLM 调用提取的实体/关系数量限制
MAX_ENTITIES_PER_CALL = 15
MAX_EDGES_PER_CALL = 10
# 并行处理配置
DEFAULT_PARALLEL_WORKERS = 3
def __init__(self, llm_client: Optional[LLMClient] = None):
"""
初始化提取 Pipeline
Args:
llm_client: LLM 客户端可选默认创建新实例
"""
self._llm_client = llm_client
@property
def llm(self) -> LLMClient:
"""延迟初始化 LLM 客户端"""
if self._llm_client is None:
self._llm_client = LLMClient()
return self._llm_client
def extract_from_chunks(
self,
chunks: List[str],
ontology: Dict[str, Any],
graph_id: str,
progress_callback: Optional[Callable] = None,
parallel_workers: int = None
) -> Tuple[List[ExtractedEntity], List[ExtractedEdge]]:
"""
从多个文本块中提取实体和关系
Args:
chunks: 文本块列表
ontology: 本体定义包含 entity_types edge_types
graph_id: 图谱ID用于日志
progress_callback: 进度回调
parallel_workers: 并行工作线程数
Returns:
(所有实体列表, 所有边列表)
"""
parallel_workers = parallel_workers or self.DEFAULT_PARALLEL_WORKERS
entity_types = ontology.get("entity_types", [])
edge_types = ontology.get("edge_types", [])
# 构建类型名称列表
entity_type_names = [e["name"] for e in entity_types]
edge_type_names = [e["name"] for e in edge_types]
all_entities = []
all_edges = []
total_chunks = len(chunks)
logger.info(f"开始 LLM 提取: {total_chunks} 个文本块, 并行数: {parallel_workers}")
# 使用线程池并行处理
with ThreadPoolExecutor(max_workers=parallel_workers) as executor:
futures = {}
for idx, chunk in enumerate(chunks):
future = executor.submit(
self._extract_single_chunk,
chunk=chunk,
chunk_index=idx,
entity_type_names=entity_type_names,
edge_type_names=edge_type_names,
entity_types=entity_types,
edge_types=edge_types
)
futures[future] = idx
completed = 0
for future in as_completed(futures):
idx = futures[future]
try:
result = future.result()
if result.is_valid():
all_entities.extend(result.entities)
all_edges.extend(result.edges)
completed += 1
if progress_callback:
progress = completed / total_chunks
progress_callback(
f"已处理 {completed}/{total_chunks} 个文本块",
progress
)
except Exception as e:
logger.error(f"处理文本块 {idx} 失败: {e}")
completed += 1
# 去重处理
all_entities = self._deduplicate_entities(all_entities)
all_edges = self._deduplicate_edges(all_edges)
logger.info(f"LLM 提取完成: {len(all_entities)} 个实体, {len(all_edges)} 条边")
return all_entities, all_edges
def _extract_single_chunk(
self,
chunk: str,
chunk_index: int,
entity_type_names: List[str],
edge_type_names: List[str],
entity_types: List[Dict],
edge_types: List[Dict]
) -> ExtractionResult:
"""
从单个文本块提取实体和关系
"""
try:
result = self._call_llm_extraction(
chunk=chunk,
entity_type_names=entity_type_names,
edge_type_names=edge_type_names
)
# 补充来源信息
for entity in result.get("entities", []):
entity["source_chunk"] = chunk[:200] # 保留前200字符作为来源
return ExtractionResult(
entities=[
ExtractedEntity(
name=e.get("name", ""),
entity_type=e.get("entity_type", ""),
summary=e.get("summary", ""),
attributes=e.get("attributes", {}),
source_chunk=chunk[:200]
)
for e in result.get("entities", [])
if e.get("name") and e.get("entity_type")
],
edges=[
ExtractedEdge(
name=e.get("name", ""),
source_name=e.get("source_name", ""),
target_name=e.get("target_name", ""),
fact=e.get("fact", ""),
source_type=e.get("source_type", ""),
target_type=e.get("target_type", ""),
attributes=e.get("attributes", {})
)
for e in result.get("edges", [])
if e.get("name") and e.get("source_name") and e.get("target_name")
],
chunk_index=chunk_index,
chunk_text=chunk
)
except Exception as e:
logger.error(f"LLM 提取失败 (chunk {chunk_index}): {e}")
return ExtractionResult(
entities=[],
edges=[],
chunk_index=chunk_index,
chunk_text=chunk,
error=str(e)
)
def _call_llm_extraction(
self,
chunk: str,
entity_type_names: List[str],
edge_type_names: List[str]
) -> Dict[str, Any]:
"""
调用 LLM 提取实体和关系
"""
system_prompt = f"""你是一个专业的知识图谱提取专家。你的任务是从给定的文本中提取实体和关系。
实体类型必须严格使用这些类型:
{json.dumps(entity_type_names, ensure_ascii=False)}
关系类型必须严格使用这些类型:
{json.dumps(edge_type_names, ensure_ascii=False)}
提取规则
1. 实体只提取符合预定义类型之一的实体
2. 关系只提取符合预定义关系类型之一的关系
3. 每个实体需要有 name名称 entity_type类型
4. 每个关系需要有 name关系类型source_name源实体名target_name目标实体名
5. 实体描述summary应简洁明了50字以内
6. 关系事实fact应描述性的句子说明 source target 之间的关系
返回JSON格式
{{
"entities": [
{{"name": "实体名", "entity_type": "类型", "summary": "描述"}}
],
"edges": [
{{"name": "关系类型", "source_name": "源实体", "target_name": "目标实体", "fact": "事实描述"}}
]
}}
只返回有效的实体和关系不要编造内容"""
user_prompt = f"""请从以下文本中提取实体和关系:
{chunk}
只返回JSON格式的结果不要包含其他内容"""
try:
response = self.llm.chat_json(
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
temperature=0.1, # 低温度保证一致性
max_tokens=2000
)
# 确保返回格式正确
return {
"entities": response.get("entities", []),
"edges": response.get("edges", [])
}
except Exception as e:
logger.warning(f"LLM 调用失败: {e}")
return {"entities": [], "edges": []}
def _deduplicate_entities(
self,
entities: List[ExtractedEntity]
) -> List[ExtractedEntity]:
"""实体去重(按 name + entity_type 组合)"""
seen = set()
result = []
for entity in entities:
key = (entity.name, entity.entity_type)
if key not in seen:
seen.add(key)
result.append(entity)
return result
def _deduplicate_edges(
self,
edges: List[ExtractedEdge]
) -> List[ExtractedEdge]:
"""边去重(按 name + source_name + target_name 组合)"""
seen = set()
result = []
for edge in edges:
key = (edge.name, edge.source_name, edge.target_name)
if key not in seen:
seen.add(key)
result.append(edge)
return result
class LLMEntityEnricher:
"""
LLM 实体增强器
用于在生成 Agent Profile 对单个实体进行深度检索和上下文丰富
这个功能替代了 Zep entity enrichment 能力
"""
def __init__(self, llm_client: Optional[LLMClient] = None):
self._llm_client = llm_client
@property
def llm(self) -> LLMClient:
if self._llm_client is None:
self._llm_client = LLMClient()
return self._llm_client
def enrich_entity(
self,
entity_name: str,
entity_type: str,
related_facts: List[str],
related_nodes: List[Dict[str, Any]]
) -> Dict[str, Any]:
"""
增强单个实体的上下文信息
Args:
entity_name: 实体名称
entity_type: 实体类型
related_facts: 相关事实列表
related_nodes: 相关节点列表
Returns:
增强后的实体信息包含更丰富的上下文描述
"""
facts_text = "\n".join([f"- {f}" for f in related_facts[:10]]) if related_facts else "无相关事实"
nodes_text = "\n".join([
f"- {n.get('name', '未知')} ({n.get('type', '未知')})"
for n in related_nodes[:5]
]) if related_nodes else "无关联实体"
system_prompt = """你是一个角色设定专家。根据给定的实体信息和关联内容,为该实体生成一个丰富的人设描述。
要求
1. 结合实体的基本信息和关联内容生成符合社交媒体场景的人设描述
2. 人设描述应该包含基本信息性格特点可能的立场观点社交媒体行为特征
3. 如果实体是人物类型应包含 MBTI 性格年龄范围国家/地区等
4. 如果实体是组织类型应包含组织性质规模立场等
5. 描述应简洁但有信息量总长度200-500
6. 只基于提供的信息生成不要编造额外细节"""
user_prompt = f"""实体信息:
- 名称{entity_name}
- 类型{entity_type}
相关事实
{facts_text}
关联实体
{nodes_text}
请生成该实体的人设描述包含性格特点可能观点行为特征等"""
try:
description = self.llm.chat(
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
temperature=0.5,
max_tokens=500
)
return {
"name": entity_name,
"entity_type": entity_type,
"enriched_description": description,
"related_facts_count": len(related_facts),
"related_nodes_count": len(related_nodes)
}
except Exception as e:
logger.warning(f"实体增强失败 {entity_name}: {e}")
return {
"name": entity_name,
"entity_type": entity_type,
"enriched_description": "",
"related_facts_count": len(related_facts),
"related_nodes_count": len(related_nodes),
"error": str(e)
}

View File

@ -0,0 +1,558 @@
"""
Neo4j 实体读取器
Neo4j 图谱中读取和过滤实体
替代 ZepEntityReader
"""
import time
from typing import Dict, Any, List, Optional, Set, Callable, TypeVar, TYPE_CHECKING
if TYPE_CHECKING:
from neo4j import Driver
from .graph_adapter import GraphAdapter, GraphNode, GraphEdge
from ...utils.logger import get_logger
logger = get_logger('mirofish.neo4j_entity_reader')
T = TypeVar('T')
def _domain_label(label: str) -> str:
"""Map Neo4j storage labels like Entity_Student back to Student."""
if label.startswith("Entity_"):
return label[len("Entity_"):]
return label
class EntityNode:
"""
实体节点数据结构
与原始的 ZepEntityReader.EntityNode 保持兼容
"""
uuid: str
name: str
labels: List[str]
summary: str
attributes: Dict[str, Any]
related_edges: List[Dict[str, Any]]
related_nodes: List[Dict[str, Any]]
def __init__(
self,
uuid: str,
name: str,
labels: List[str],
summary: str,
attributes: Dict[str, Any],
related_edges: List[Dict[str, Any]] = None,
related_nodes: List[Dict[str, Any]] = None
):
self.uuid = uuid
self.name = name
self.labels = labels
self.summary = summary
self.attributes = attributes
self.related_edges = related_edges or []
self.related_nodes = related_nodes or []
def to_dict(self) -> Dict[str, Any]:
return {
"uuid": self.uuid,
"name": self.name,
"labels": self.labels,
"summary": self.summary,
"attributes": self.attributes,
"related_edges": self.related_edges,
"related_nodes": self.related_nodes,
}
def get_entity_type(self) -> Optional[str]:
"""获取实体类型(排除默认的 Entity 标签)"""
for label in self.labels:
if label not in ("Entity", "Node"):
return _domain_label(label)
return None
class FilteredEntities:
"""过滤后的实体集合"""
entities: List[EntityNode]
entity_types: Set[str]
total_count: int
filtered_count: int
def __init__(
self,
entities: List[EntityNode],
entity_types: Set[str],
total_count: int,
filtered_count: int
):
self.entities = entities
self.entity_types = entity_types
self.total_count = total_count
self.filtered_count = filtered_count
def to_dict(self) -> Dict[str, Any]:
return {
"entities": [e.to_dict() for e in self.entities],
"entity_types": list(self.entity_types),
"total_count": self.total_count,
"filtered_count": self.filtered_count,
}
class Neo4jEntityReader:
"""
Neo4j 实体读取与过滤服务
主要功能
1. Neo4j 图谱读取所有节点
2. 筛选出符合预定义实体类型的节点
3. 获取每个实体的相关边和关联节点信息
"""
# 重试配置
MAX_RETRIES = 3
RETRY_DELAY = 2.0
def __init__(self, driver: 'Driver' = None):
"""
初始化读取器
Args:
driver: Neo4j 驱动可选默认使用全局驱动
"""
from ...utils.neo4j.driver import get_neo4j_driver
self.driver = driver or get_neo4j_driver()
def _call_with_retry(
self,
func: Callable[[], T],
operation_name: str,
max_retries: int = None,
initial_delay: float = None
) -> T:
"""
带重试机制的 Neo4j 查询
Args:
func: 要执行的函数
operation_name: 操作名称用于日志
max_retries: 最大重试次数
initial_delay: 初始延迟秒数
Returns:
查询结果
"""
max_retries = max_retries or self.MAX_RETRIES
initial_delay = initial_delay or self.RETRY_DELAY
last_exception = None
delay = initial_delay
for attempt in range(max_retries):
try:
return func()
except Exception as e:
last_exception = e
if attempt < max_retries - 1:
logger.warning(
f"Neo4j {operation_name}{attempt + 1} 次尝试失败: {str(e)[:100]}, "
f"{delay:.1f}秒后重试..."
)
time.sleep(delay)
delay *= 2
else:
logger.error(
f"Neo4j {operation_name}{max_retries} 次尝试后仍失败: {str(e)}"
)
raise last_exception
def get_all_nodes(self, graph_id: str) -> List[Dict[str, Any]]:
"""
获取图谱的所有节点分页获取
Args:
graph_id: 图谱ID
Returns:
节点列表
"""
logger.info(f"获取图谱 {graph_id} 的所有节点...")
def _query():
nodes = []
with self.driver.session() as session:
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
n.created_at AS created_at, properties(n) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
nodes.append({
"uuid": record["uuid"],
"name": record["name"],
"labels": record["labels"],
"summary": record["summary"] or "",
"entity_type": record.get("entity_type"),
"attributes": record["attributes"] or {},
})
return nodes
nodes = self._call_with_retry(_query, f"获取所有节点({graph_id})")
logger.info(f"共获取 {len(nodes)} 个节点")
return nodes
def get_all_edges(self, graph_id: str) -> List[Dict[str, Any]]:
"""
获取图谱的所有边
Args:
graph_id: 图谱ID
Returns:
边列表
"""
logger.info(f"获取图谱 {graph_id} 的所有边...")
def _query():
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
r.source_node_uuid AS source_node_uuid,
r.target_node_uuid AS target_node_uuid,
r.created_at AS created_at, properties(r) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
edges.append({
"uuid": record["uuid"],
"name": record["name"],
"fact": record["fact"] or "",
"source_node_uuid": record["source_node_uuid"],
"target_node_uuid": record["target_node_uuid"],
"attributes": record["attributes"] or {},
})
return edges
edges = self._call_with_retry(_query, f"获取所有边({graph_id})")
logger.info(f"共获取 {len(edges)} 条边")
return edges
def get_node_edges(self, node_uuid: str) -> List[Dict[str, Any]]:
"""
获取指定节点的所有相关边
Args:
node_uuid: 节点UUID
Returns:
边列表
"""
def _query():
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.source_node_uuid = $uuid OR r.target_node_uuid = $uuid
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
r.source_node_uuid AS source_node_uuid,
r.target_node_uuid AS target_node_uuid,
source.name AS source_name, target.name AS target_name,
properties(r) AS attributes
"""
result = session.run(cypher, uuid=node_uuid)
for record in result:
edges.append({
"uuid": record["uuid"],
"name": record["name"],
"fact": record["fact"] or "",
"source_node_uuid": record["source_node_uuid"],
"target_node_uuid": record["target_node_uuid"],
"attributes": record["attributes"] or {},
})
return edges
return self._call_with_retry(
_query, f"获取节点边({node_uuid[:8]}...)"
)
def filter_defined_entities(
self,
graph_id: str,
defined_entity_types: Optional[List[str]] = None,
enrich_with_edges: bool = True
) -> FilteredEntities:
"""
筛选出符合预定义实体类型的节点
筛选逻辑
- 如果节点的 Labels 只有一个 "Entity"说明这个实体不符合预定义类型跳过
- 如果节点的 Labels 包含除 "Entity" "Node" 之外的标签说明符合预定义类型保留
Args:
graph_id: 图谱ID
defined_entity_types: 预定义的实体类型列表可选
enrich_with_edges: 是否获取每个实体的相关边信息
Returns:
FilteredEntities: 过滤后的实体集合
"""
logger.info(f"开始筛选图谱 {graph_id} 的实体...")
# 获取所有节点
all_nodes = self.get_all_nodes(graph_id)
total_count = len(all_nodes)
# 获取所有边(用于后续关联查找)
all_edges = self.get_all_edges(graph_id) if enrich_with_edges else []
# 构建节点UUID到节点数据的映射
node_map = {n["uuid"]: n for n in all_nodes}
# 筛选符合条件的实体
filtered_entities = []
entity_types_found = set()
for node in all_nodes:
labels = node.get("labels", [])
# 筛选逻辑Labels 必须包含除 "Entity" 和 "Node" 之外的标签
custom_labels = [_domain_label(l) for l in labels if l not in ("Entity", "Node")]
if not custom_labels:
continue
# 如果指定了预定义类型,检查是否匹配
if defined_entity_types:
matching_labels = [l for l in custom_labels if l in defined_entity_types]
if not matching_labels:
continue
entity_type = matching_labels[0]
else:
entity_type = custom_labels[0]
entity_types_found.add(entity_type)
# 创建实体节点对象
display_labels = ["Entity", *custom_labels]
entity = EntityNode(
uuid=node["uuid"],
name=node["name"],
labels=display_labels,
summary=node["summary"],
attributes=node["attributes"],
)
# 获取相关边和节点
if enrich_with_edges:
related_edges = []
related_node_uuids = set()
for edge in all_edges:
if edge["source_node_uuid"] == node["uuid"]:
related_edges.append({
"direction": "outgoing",
"edge_name": edge["name"],
"fact": edge["fact"],
"target_node_uuid": edge["target_node_uuid"],
})
related_node_uuids.add(edge["target_node_uuid"])
elif edge["target_node_uuid"] == node["uuid"]:
related_edges.append({
"direction": "incoming",
"edge_name": edge["name"],
"fact": edge["fact"],
"source_node_uuid": edge["source_node_uuid"],
})
related_node_uuids.add(edge["source_node_uuid"])
entity.related_edges = related_edges
# 获取关联节点的基本信息
related_nodes = []
for related_uuid in related_node_uuids:
if related_uuid in node_map:
related_node = node_map[related_uuid]
related_labels = [
_domain_label(l)
for l in related_node["labels"]
if l not in ("Entity", "Node")
]
related_nodes.append({
"uuid": related_node["uuid"],
"name": related_node["name"],
"labels": ["Entity", *related_labels],
"summary": related_node.get("summary", ""),
})
entity.related_nodes = related_nodes
filtered_entities.append(entity)
logger.info(
f"筛选完成: 总节点 {total_count}, 符合条件 {len(filtered_entities)}, "
f"实体类型: {entity_types_found}"
)
return FilteredEntities(
entities=filtered_entities,
entity_types=entity_types_found,
total_count=total_count,
filtered_count=len(filtered_entities),
)
def get_entity_with_context(
self,
graph_id: str,
entity_uuid: str
) -> Optional[EntityNode]:
"""
获取单个实体及其完整上下文边和关联节点
Args:
graph_id: 图谱ID
entity_uuid: 实体UUID
Returns:
EntityNode None
"""
def _query():
with self.driver.session() as session:
# 获取节点
cypher = """
MATCH (n:Entity)
WHERE n.uuid = $uuid AND n.graph_id = $graph_id
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
properties(n) AS attributes
"""
result = session.run(cypher, uuid=entity_uuid, graph_id=graph_id)
record = result.single()
if not record:
return None
# 获取节点的边
edges_cypher = """
MATCH (source)-[r]->(target)
WHERE (r.source_node_uuid = $uuid OR r.target_node_uuid = $uuid)
AND r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
r.source_node_uuid AS source_node_uuid,
r.target_node_uuid AS target_node_uuid
"""
edges_result = session.run(edges_cypher, uuid=entity_uuid, graph_id=graph_id)
edges = []
for edge_record in edges_result:
edges.append({
"uuid": edge_record["uuid"],
"name": edge_record["name"],
"fact": edge_record["fact"] or "",
"source_node_uuid": edge_record["source_node_uuid"],
"target_node_uuid": edge_record["target_node_uuid"],
})
return record, edges
try:
result = self._call_with_retry(
_query,
f"获取实体详情({entity_uuid[:8]}...)"
)
if not result:
return None
record, edges = result
# 获取所有节点用于关联查找
all_nodes = self.get_all_nodes(graph_id)
node_map = {n["uuid"]: n for n in all_nodes}
# 处理相关边和节点
related_edges = []
related_node_uuids = set()
for edge in edges:
if edge["source_node_uuid"] == entity_uuid:
related_edges.append({
"direction": "outgoing",
"edge_name": edge["name"],
"fact": edge["fact"],
"target_node_uuid": edge["target_node_uuid"],
})
related_node_uuids.add(edge["target_node_uuid"])
else:
related_edges.append({
"direction": "incoming",
"edge_name": edge["name"],
"fact": edge["fact"],
"source_node_uuid": edge["source_node_uuid"],
})
related_node_uuids.add(edge["source_node_uuid"])
# 获取关联节点信息
related_nodes = []
for related_uuid in related_node_uuids:
if related_uuid in node_map:
related_node = node_map[related_uuid]
related_nodes.append({
"uuid": related_node["uuid"],
"name": related_node["name"],
"labels": related_node["labels"],
"summary": related_node.get("summary", ""),
})
labels = record["labels"] or []
custom_labels = [_domain_label(l) for l in labels if l not in ("Entity", "Node")]
return EntityNode(
uuid=record["uuid"],
name=record["name"],
labels=["Entity", *custom_labels],
summary=record["summary"] or "",
attributes=record["attributes"] or {},
related_edges=related_edges,
related_nodes=related_nodes,
)
except Exception as e:
logger.error(f"获取实体 {entity_uuid} 失败: {str(e)}")
return None
def get_entities_by_type(
self,
graph_id: str,
entity_type: str,
enrich_with_edges: bool = True
) -> List[EntityNode]:
"""
获取指定类型的所有实体
Args:
graph_id: 图谱ID
entity_type: 实体类型 "Student", "PublicFigure"
enrich_with_edges: 是否获取相关边信息
Returns:
实体列表
"""
result = self.filter_defined_entities(
graph_id=graph_id,
defined_entity_types=[entity_type],
enrich_with_edges=enrich_with_edges
)
return result.entities

View File

@ -0,0 +1,864 @@
"""
Neo4j 图谱构建器
实现 GraphAdapter 接口替代 Zep Cloud GraphBuilderService
"""
import os
import uuid
import time
import threading
import json
from typing import Dict, Any, List, Optional, Callable
from neo4j import Driver
from .graph_adapter import GraphAdapter, GraphInfo, GraphNode, GraphEdge, SearchResult
from .llm_extractor import LLMExtractionPipeline
from ...utils.neo4j.driver import get_neo4j_driver
from ...utils.neo4j.schema import Neo4jSchemaManager
from ...utils.logger import get_logger
from ...models.task import TaskManager, TaskStatus
from ...utils.locale import t, get_locale, set_locale
logger = get_logger('mirofish.neo4j_graph_builder')
def _safe_neo4j_identifier(value: str, fallback: str = "RELATED_TO") -> str:
"""Return a conservative Neo4j label/relationship identifier."""
cleaned = "".join(ch if ch.isalnum() or ch == "_" else "_" for ch in str(value or ""))
if not cleaned:
cleaned = fallback
if cleaned[0].isdigit():
cleaned = f"_{cleaned}"
return cleaned
class Neo4jGraphBuilder(GraphAdapter):
"""
Neo4j 图谱构建器
使用 LLM Extraction Pipeline + Neo4j 实现知识图谱构建
替代 Zep Cloud 的自动文本提取功能
"""
def __init__(
self,
driver: Optional[Driver] = None,
llm_extractor: Optional[LLMExtractionPipeline] = None
):
"""
初始化 Neo4j 图谱构建器
Args:
driver: Neo4j 驱动可选默认使用全局驱动
llm_extractor: LLM 提取器可选默认创建新实例
"""
self.driver = driver or get_neo4j_driver()
self.extractor = llm_extractor or LLMExtractionPipeline()
self.schema_manager = Neo4jSchemaManager(self.driver)
self.task_manager = TaskManager()
self._ontology_cache: Dict[str, Dict[str, Any]] = {}
def create_graph(self, name: str) -> str:
"""
创建新图谱
Args:
name: 图谱名称
Returns:
graph_id: 新创建的图谱ID
"""
graph_id = f"mirofish_{uuid.uuid4().hex[:16]}"
with self.driver.session() as session:
# 创建图谱元数据节点
cypher = """
MERGE (g:_GraphMetadata {graph_id: $graph_id})
ON CREATE SET
g.name = $name,
g.created_at = datetime(),
g.entity_count = 0,
g.edge_count = 0
RETURN g.graph_id AS graph_id
"""
result = session.run(cypher, graph_id=graph_id, name=name)
created_id = result.single()["graph_id"]
logger.info(f"创建 Neo4j 图谱: {graph_id}, 名称: {name}")
return created_id
def set_ontology(
self,
graph_id: str,
entity_types,
edge_types: Optional[List[Dict[str, Any]]] = None
) -> None:
"""
设置图谱本体Schema
Args:
graph_id: 图谱ID
entity_types: 实体类型定义列表
edge_types: 关系类型定义列表
"""
if isinstance(entity_types, dict):
ontology = entity_types
entity_types = ontology.get("entity_types", [])
edge_types = ontology.get("edge_types", [])
else:
ontology = {
"entity_types": entity_types or [],
"edge_types": edge_types or [],
}
self._ontology_cache[graph_id] = ontology
self.schema_manager.setup_graph_schema(graph_id, entity_types or [], edge_types or [])
with self.driver.session() as session:
session.run(
"""
MERGE (g:_GraphMetadata {graph_id: $graph_id})
SET g.ontology_json = $ontology_json,
g.updated_at = datetime()
""",
graph_id=graph_id,
ontology_json=json.dumps(ontology, ensure_ascii=False)
)
logger.info(f"本体已设置: {graph_id}")
def add_text_batches(
self,
graph_id: str,
chunks: List[str],
batch_size: int = 3,
progress_callback: Optional[Callable] = None
) -> List[str]:
"""
批量添加文本到图谱通过 LLM 提取
Args:
graph_id: 图谱ID
chunks: 文本块列表
batch_size: 批处理大小LLM 并行数
progress_callback: 进度回调
Returns:
chunk_ids: 所有文本块的 ID 列表
"""
# 构建 ontology 字典(用于 LLM 提取)
ontology = self._build_ontology_for_extraction(graph_id)
# 使用 LLM 提取实体和关系
entities, edges = self.extractor.extract_from_chunks(
chunks=chunks,
ontology=ontology,
graph_id=graph_id,
progress_callback=progress_callback,
parallel_workers=batch_size
)
# 将提取的结果写入 Neo4j
self._write_extracted_data(graph_id, entities, edges)
# 返回 chunk IDs这里用索引代替 UUID
chunk_ids = [f"chunk_{i}" for i in range(len(chunks))]
return chunk_ids
def _build_ontology_for_extraction(self, graph_id: str) -> Dict[str, Any]:
"""从图谱 Schema 构建 LLM 提取用的 ontology"""
if graph_id in self._ontology_cache:
return self._ontology_cache[graph_id]
with self.driver.session() as session:
record = session.run(
"""
MATCH (g:_GraphMetadata {graph_id: $graph_id})
RETURN g.ontology_json AS ontology_json
""",
graph_id=graph_id
).single()
if record and record["ontology_json"]:
try:
ontology = json.loads(record["ontology_json"])
self._ontology_cache[graph_id] = ontology
return ontology
except json.JSONDecodeError:
logger.warning(f"图谱本体解析失败,使用空本体: {graph_id}")
return {"entity_types": [], "edge_types": []}
def _write_extracted_data(
self,
graph_id: str,
entities: List,
edges: List
) -> None:
"""
将提取的实体和关系写入 Neo4j
Args:
graph_id: 图谱ID
entities: 实体列表
edges: 边列表
"""
with self.driver.session() as session:
# 写入实体
for entity in entities:
self._write_node(session, graph_id, entity)
# 写入关系
# 先建立节点名称到UUID的映射
name_to_uuid = self._get_name_to_uuid_mapping(session, graph_id)
for edge in edges:
self._write_edge(session, graph_id, edge, name_to_uuid)
# 更新图谱统计
self._update_graph_stats(session, graph_id, len(entities), len(edges))
def _write_node(self, session, graph_id: str, entity) -> None:
"""写入单个节点"""
entity_type = _safe_neo4j_identifier(entity.entity_type, "Entity")
# Neo4j Label 格式Entity_Student
labels = ["Entity", f"Entity_{entity_type}"] if entity_type else ["Entity"]
# 构建属性
properties = {
"uuid": str(uuid.uuid4()),
"name": entity.name,
"summary": entity.summary,
"graph_id": graph_id,
"created_at": time.strftime("%Y-%m-%dT%H:%M:%S"),
"entity_type": entity_type
}
# 添加自定义属性
for key, value in entity.attributes.items():
properties[key] = value
# 构建 MERGE 查询
label_str = ":".join(labels)
set_clause = ", ".join([f"n.{k} = ${k}" for k in properties.keys()])
cypher = f"""
MERGE (n:{label_str} {{name: $name, graph_id: $graph_id}})
ON CREATE SET {set_clause}
ON MATCH SET {set_clause}
RETURN n.uuid AS uuid
"""
try:
session.run(cypher, **properties)
except Exception as e:
logger.warning(f"写入节点失败: {entity.name}, {e}")
def _write_edge(
self,
session,
graph_id: str,
edge,
name_to_uuid: Dict[str, str]
) -> None:
"""写入单条边"""
source_uuid = name_to_uuid.get(edge.source_name)
target_uuid = name_to_uuid.get(edge.target_name)
if not source_uuid or not target_uuid:
logger.debug(f"跳过边(找不到节点): {edge.source_name} -> {edge.target_name}")
return
properties = {
"uuid": str(uuid.uuid4()),
"name": edge.name,
"fact": edge.fact,
"graph_id": graph_id,
"source_node_uuid": source_uuid,
"target_node_uuid": target_uuid,
"created_at": time.strftime("%Y-%m-%dT%H:%M:%S")
}
rel_type = _safe_neo4j_identifier(edge.name)
cypher = f"""
MATCH (source:Entity {{uuid: $source_node_uuid}})
MATCH (target:Entity {{uuid: $target_node_uuid}})
MERGE (source)-[rel:`{rel_type}` {{
graph_id: $graph_id,
source_node_uuid: $source_node_uuid,
target_node_uuid: $target_node_uuid
}}]->(target)
SET rel.uuid = coalesce(rel.uuid, $uuid),
rel.name = $name,
rel.fact = $fact,
rel.created_at = coalesce(rel.created_at, $created_at)
RETURN rel.uuid AS uuid
"""
try:
session.run(cypher, **properties)
except Exception as e:
logger.warning(f"写入边失败: {edge.name}, {e}")
def _get_name_to_uuid_mapping(
self,
session,
graph_id: str
) -> Dict[str, str]:
"""获取节点名称到UUID的映射"""
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN n.name AS name, n.uuid AS uuid
"""
result = session.run(cypher, graph_id=graph_id)
return {record["name"]: record["uuid"] for record in result}
def _update_graph_stats(
self,
session,
graph_id: str,
entities_count: int,
edges_count: int
) -> None:
"""更新图谱统计信息"""
cypher = """
MATCH (g:_GraphMetadata {graph_id: $graph_id})
SET g.entity_count = g.entity_count + $entities,
g.edge_count = g.edge_count + $edges,
g.updated_at = datetime()
"""
session.run(cypher, graph_id=graph_id, entities=entities_count, edges=edges_count)
def get_all_nodes(self, graph_id: str) -> List[GraphNode]:
"""
获取图谱的所有节点
Args:
graph_id: 图谱ID
Returns:
节点列表
"""
nodes = []
with self.driver.session() as session:
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
n.created_at AS created_at,
properties(n) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
# 过滤掉保留的 Label
labels = [l for l in record["labels"] if l not in ("Entity", "Node")]
nodes.append(GraphNode(
uuid=record["uuid"],
name=record["name"],
labels=labels,
summary=record["summary"] or "",
attributes=record["attributes"] or {},
created_at=record["created_at"]
))
logger.info(f"获取节点: {graph_id}, 共 {len(nodes)}")
return nodes
def get_all_edges(self, graph_id: str) -> List[GraphEdge]:
"""
获取图谱的所有边
Args:
graph_id: 图谱ID
Returns:
边列表
"""
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
r.created_at AS created_at, r.valid_at AS valid_at,
r.invalid_at AS invalid_at, r.expired_at AS expired_at,
properties(r) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
edges.append(GraphEdge(
uuid=record["uuid"],
name=record["name"],
fact=record["fact"] or "",
source_node_uuid=record["source_node_uuid"],
target_node_uuid=record["target_node_uuid"],
attributes=record["attributes"] or {},
created_at=record["created_at"],
valid_at=record["valid_at"],
invalid_at=record["invalid_at"],
expired_at=record["expired_at"]
))
logger.info(f"获取边: {graph_id}, 共 {len(edges)}")
return edges
def get_node(self, node_uuid: str) -> Optional[GraphNode]:
"""
获取单个节点
Args:
node_uuid: 节点UUID
Returns:
节点对象或None
"""
with self.driver.session() as session:
cypher = """
MATCH (n:Entity)
WHERE n.uuid = $uuid
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
n.created_at AS created_at, properties(n) AS attributes
"""
result = session.run(cypher, uuid=node_uuid)
record = result.single()
if not record:
return None
labels = [l for l in record["labels"] if l not in ("Entity", "Node")]
return GraphNode(
uuid=record["uuid"],
name=record["name"],
labels=labels,
summary=record["summary"] or "",
attributes=record["attributes"] or {},
created_at=record["created_at"]
)
def get_node_edges(self, node_uuid: str) -> List[GraphEdge]:
"""
获取指定节点的所有相关边
Args:
node_uuid: 节点UUID
Returns:
边列表
"""
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.source_node_uuid = $uuid OR r.target_node_uuid = $uuid
OR source.uuid = $uuid OR target.uuid = $uuid
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
r.created_at AS created_at, properties(r) AS attributes
"""
result = session.run(cypher, uuid=node_uuid)
for record in result:
edges.append(GraphEdge(
uuid=record["uuid"],
name=record["name"],
fact=record["fact"] or "",
source_node_uuid=record["source_node_uuid"],
target_node_uuid=record["target_node_uuid"],
attributes=record["attributes"] or {},
created_at=record["created_at"]
))
return edges
def search(
self,
graph_id: str,
query: str,
limit: int = 10,
scope: str = "edges"
) -> SearchResult:
"""
图谱搜索
使用 Neo4j 的全文索引或标签/属性搜索
Args:
graph_id: 图谱ID
query: 搜索查询
limit: 返回结果数量
scope: 搜索范围 ("edges" / "nodes" / "both")
Returns:
SearchResult: 搜索结果
"""
facts = []
edges_result = []
nodes_result = []
with self.driver.session() as session:
if scope in ("edges", "both"):
# 搜索边(通过 fact 属性)
edge_cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
AND (r.fact CONTAINS $query OR r.name CONTAINS $query)
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
source.name AS source_name, target.name AS target_name,
properties(r) AS attributes
LIMIT $limit
"""
result = session.run(edge_cypher, graph_id=graph_id, query=query, limit=limit)
for record in result:
if record["fact"]:
facts.append(record["fact"])
edges_result.append({
"uuid": record["uuid"],
"name": record["name"],
"fact": record["fact"],
"source_node_uuid": record["source_node_uuid"],
"target_node_uuid": record["target_node_uuid"],
})
if scope in ("nodes", "both"):
# 搜索节点
node_cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
AND (n.name CONTAINS $query OR n.summary CONTAINS $query)
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
properties(n) AS attributes
LIMIT $limit
"""
result = session.run(node_cypher, graph_id=graph_id, query=query, limit=limit)
for record in result:
labels = [l for l in record["labels"] if l not in ("Entity", "Node")]
if record["summary"]:
facts.append(f"[{record['name']}]: {record['summary']}")
nodes_result.append({
"uuid": record["uuid"],
"name": record["name"],
"labels": labels,
"summary": record["summary"],
})
return SearchResult(
facts=facts,
edges=edges_result,
nodes=nodes_result,
query=query,
total_count=len(facts)
)
def add_activities(
self,
graph_id: str,
activities: List[str]
) -> None:
"""
添加活动记录到图谱
将活动文本作为新的文本块处理提取实体和关系
Args:
graph_id: 图谱ID
activities: 活动描述文本列表
"""
# 将活动文本当作新的文本块,使用 LLM 提取
ontology = self._build_ontology_for_extraction(graph_id)
# 并行提取
entities, edges = self.extractor.extract_from_chunks(
chunks=activities,
ontology=ontology,
graph_id=graph_id,
parallel_workers=1 # 活动添加通常较小,单线程即可
)
# 写入 Neo4j
self._write_extracted_data(graph_id, entities, edges)
logger.info(f"添加活动: {graph_id}, {len(activities)}")
def _wait_for_episodes(
self,
episode_uuids: List[str],
progress_callback: Optional[Callable] = None,
timeout: int = 600
) -> None:
"""
Neo4j 写入是同步完成的保留该方法以兼容原 Zep 构建流程
"""
if progress_callback:
progress_callback("Neo4j 图谱写入已完成", 1.0)
def get_graph_data(self, graph_id: str) -> Dict[str, Any]:
"""
获取完整图谱数据返回格式与 GraphBuilderService 保持一致
"""
nodes = self.get_all_nodes(graph_id)
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
source.name AS source_node_name,
target.name AS target_node_name,
r.created_at AS created_at,
r.valid_at AS valid_at,
r.invalid_at AS invalid_at,
r.expired_at AS expired_at,
properties(r) AS attributes
"""
edge_records = list(session.run(cypher, graph_id=graph_id))
nodes_data = [
{
"uuid": node.uuid,
"name": node.name,
"labels": node.labels,
"summary": node.summary,
"attributes": node.attributes,
"created_at": str(node.created_at) if node.created_at else None,
}
for node in nodes
]
edges_data = []
for record in edge_records:
created_at = record["created_at"]
valid_at = record["valid_at"]
invalid_at = record["invalid_at"]
expired_at = record["expired_at"]
edge_name = record["name"] or ""
edges_data.append({
"uuid": record["uuid"],
"name": edge_name,
"fact": record["fact"] or "",
"fact_type": edge_name,
"source_node_uuid": record["source_node_uuid"],
"target_node_uuid": record["target_node_uuid"],
"source_node_name": record["source_node_name"] or "",
"target_node_name": record["target_node_name"] or "",
"attributes": record["attributes"] or {},
"created_at": str(created_at) if created_at else None,
"valid_at": str(valid_at) if valid_at else None,
"invalid_at": str(invalid_at) if invalid_at else None,
"expired_at": str(expired_at) if expired_at else None,
"episodes": [],
})
return {
"graph_id": graph_id,
"nodes": nodes_data,
"edges": edges_data,
"node_count": len(nodes_data),
"edge_count": len(edges_data),
}
def delete_graph(self, graph_id: str) -> None:
"""
删除图谱
Args:
graph_id: 图谱ID
"""
self.schema_manager.drop_graph_data(graph_id)
logger.info(f"删除图谱: {graph_id}")
def health_check(self) -> bool:
"""
健康检查
Returns:
True 如果连接正常
"""
try:
with self.driver.session() as session:
result = session.run("RETURN 1 AS test")
result.single()
return True
except Exception as e:
logger.error(f"Neo4j 健康检查失败: {e}")
return False
class Neo4jAsyncGraphBuilder(Neo4jGraphBuilder):
"""
Neo4j 异步图谱构建器
支持异步操作适用于大规模图谱构建
"""
async def build_graph_async(
self,
text: str,
ontology: Dict[str, Any],
graph_name: str = "MiroFish Graph",
chunk_size: int = 500,
chunk_overlap: int = 50,
batch_size: int = 3
) -> str:
"""
异步构建图谱
Args:
text: 输入文本
ontology: 本体定义
graph_name: 图谱名称
chunk_size: 文本块大小
chunk_overlap: 块重叠大小
batch_size: LLM 并行处理数
Returns:
task_id: 任务ID
"""
from ...utils.text_processor import TextProcessor
# 创建任务
task_id = self.task_manager.create_task(
task_type="neo4j_graph_build",
metadata={
"graph_name": graph_name,
"chunk_size": chunk_size,
"text_length": len(text),
}
)
# 在后台线程执行
thread = threading.Thread(
target=self._build_graph_worker,
args=(
task_id, text, ontology, graph_name,
chunk_size, chunk_overlap, batch_size,
get_locale()
)
)
thread.daemon = True
thread.start()
return task_id
def _build_graph_worker(
self,
task_id: str,
text: str,
ontology: Dict[str, Any],
graph_name: str,
chunk_size: int,
chunk_overlap: int,
batch_size: int,
locale: str
):
"""图谱构建工作线程"""
set_locale(locale)
try:
self.task_manager.update_task(
task_id,
status=TaskStatus.PROCESSING,
progress=5,
message=t('progress.startBuildingGraph')
)
# 1. 创建图谱
graph_id = self.create_graph(graph_name)
self.task_manager.update_task(
task_id,
progress=10,
message=t('progress.graphCreated', graphId=graph_id)
)
# 2. 设置本体
self.set_ontology(
graph_id,
ontology.get("entity_types", []),
ontology.get("edge_types", [])
)
self.task_manager.update_task(
task_id,
progress=15,
message=t('progress.ontologySet')
)
# 3. 文本分块
from ...utils.text_processor import TextProcessor
chunks = TextProcessor.split_text(text, chunk_size, chunk_overlap)
total_chunks = len(chunks)
self.task_manager.update_task(
task_id,
progress=20,
message=t('progress.textSplit', count=total_chunks)
)
# 4. 提取并写入
self.task_manager.update_task(
task_id,
progress=30,
message=t('progress.extractingEntities')
)
# 定义进度回调
def progress_callback(msg: str, prog: float):
self.task_manager.update_task(
task_id,
progress=30 + int(prog * 50), # 30-80%
message=msg
)
self.add_text_batches(
graph_id, chunks, batch_size, progress_callback
)
# 5. 完成
self.task_manager.update_task(
task_id,
progress=95,
message=t('progress.fetchingGraphInfo')
)
graph_info = self._get_graph_info(graph_id)
self.task_manager.complete_task(task_id, {
"graph_id": graph_id,
"graph_info": graph_info.to_dict(),
"chunks_processed": total_chunks,
})
except Exception as e:
import traceback
error_msg = f"{str(e)}\n{traceback.format_exc()}"
self.task_manager.fail_task(task_id, error_msg)
def _get_graph_info(self, graph_id: str) -> GraphInfo:
"""获取图谱信息"""
stats = self.schema_manager.get_graph_stats(graph_id)
entity_types = list(stats.get("entity_types", {}).keys())
return GraphInfo(
graph_id=graph_id,
node_count=stats.get("node_count", 0),
edge_count=stats.get("edge_count", 0),
entity_types=entity_types
)

View File

@ -0,0 +1,668 @@
"""
Neo4j 图谱记忆更新服务
将模拟中的 Agent 活动动态更新到 Neo4j 图谱中
替代 ZepGraphMemoryUpdater
"""
import time
import threading
from typing import Dict, Any, List, Optional, TYPE_CHECKING
from dataclasses import dataclass, field
from datetime import datetime
from queue import Queue, Empty
if TYPE_CHECKING:
from neo4j import Driver
from .llm_extractor import LLMExtractionPipeline
from ...utils.logger import get_logger
from ...utils.locale import get_locale, set_locale
logger = get_logger('mirofish.neo4j_memory_updater')
def _safe_neo4j_identifier(value: str, fallback: str = "RELATED_TO") -> str:
cleaned = "".join(ch if ch.isalnum() or ch == "_" else "_" for ch in str(value or ""))
if not cleaned:
cleaned = fallback
if cleaned[0].isdigit():
cleaned = f"_{cleaned}"
return cleaned
@dataclass
class AgentActivity:
"""Agent 活动记录"""
platform: str # twitter / reddit
agent_id: int
agent_name: str
action_type: str # CREATE_POST, LIKE_POST, etc.
action_args: Dict[str, Any]
round_num: int
timestamp: str
def to_episode_text(self) -> str:
"""
将活动转换为可以写入 Neo4j 的文本描述
采用自然语言描述格式 LLM 能够从中提取实体和关系
"""
action_descriptions = {
"CREATE_POST": self._describe_create_post,
"LIKE_POST": self._describe_like_post,
"DISLIKE_POST": self._describe_dislike_post,
"REPOST": self._describe_repost,
"QUOTE_POST": self._describe_quote_post,
"FOLLOW": self._describe_follow,
"CREATE_COMMENT": self._describe_create_comment,
"LIKE_COMMENT": self._describe_like_comment,
"DISLIKE_COMMENT": self._describe_dislike_comment,
"SEARCH_POSTS": self._describe_search,
"SEARCH_USER": self._describe_search_user,
"MUTE": self._describe_mute,
}
describe_func = action_descriptions.get(
self.action_type,
self._describe_generic
)
description = describe_func()
return f"{self.agent_name}: {description}"
def _describe_create_post(self) -> str:
content = self.action_args.get("content", "")
if content:
return f"发布了一条帖子:「{content}"
return "发布了一条帖子"
def _describe_like_post(self) -> str:
post_content = self.action_args.get("post_content", "")
post_author = self.action_args.get("post_author_name", "")
if post_content and post_author:
return f"点赞了{post_author}的帖子:「{post_content}"
elif post_content:
return f"点赞了一条帖子:「{post_content}"
elif post_author:
return f"点赞了{post_author}的一条帖子"
return "点赞了一条帖子"
def _describe_dislike_post(self) -> str:
post_content = self.action_args.get("post_content", "")
post_author = self.action_args.get("post_author_name", "")
if post_content and post_author:
return f"踩了{post_author}的帖子:「{post_content}"
elif post_content:
return f"踩了一条帖子:「{post_content}"
elif post_author:
return f"踩了{post_author}的一条帖子"
return "踩了一条帖子"
def _describe_repost(self) -> str:
original_content = self.action_args.get("original_content", "")
original_author = self.action_args.get("original_author_name", "")
if original_content and original_author:
return f"转发了{original_author}的帖子:「{original_content}"
elif original_content:
return f"转发了一条帖子:「{original_content}"
elif original_author:
return f"转发了{original_author}的一条帖子"
return "转发了一条帖子"
def _describe_quote_post(self) -> str:
original_content = self.action_args.get("original_content", "")
original_author = self.action_args.get("original_author_name", "")
quote_content = self.action_args.get("quote_content", "") or self.action_args.get("content", "")
base = ""
if original_content and original_author:
base = f"引用了{original_author}的帖子「{original_content}"
elif original_content:
base = f"引用了一条帖子「{original_content}"
elif original_author:
base = f"引用了{original_author}的一条帖子"
else:
base = "引用了一条帖子"
if quote_content:
base += f",并评论道:「{quote_content}"
return base
def _describe_follow(self) -> str:
target_user_name = self.action_args.get("target_user_name", "")
if target_user_name:
return f"关注了用户「{target_user_name}"
return "关注了一个用户"
def _describe_create_comment(self) -> str:
content = self.action_args.get("content", "")
post_content = self.action_args.get("post_content", "")
post_author = self.action_args.get("post_author_name", "")
if content:
if post_content and post_author:
return f"{post_author}的帖子「{post_content}」下评论道:「{content}"
elif post_content:
return f"在帖子「{post_content}」下评论道:「{content}"
elif post_author:
return f"{post_author}的帖子下评论道:「{content}"
return f"评论道:「{content}"
return "发表了评论"
def _describe_like_comment(self) -> str:
comment_content = self.action_args.get("comment_content", "")
comment_author = self.action_args.get("comment_author_name", "")
if comment_content and comment_author:
return f"点赞了{comment_author}的评论:「{comment_content}"
elif comment_content:
return f"点赞了一条评论:「{comment_content}"
elif comment_author:
return f"点赞了{comment_author}的一条评论"
return "点赞了一条评论"
def _describe_dislike_comment(self) -> str:
comment_content = self.action_args.get("comment_content", "")
comment_author = self.action_args.get("comment_author_name", "")
if comment_content and comment_author:
return f"踩了{comment_author}的评论:「{comment_content}"
elif comment_content:
return f"踩了一条评论:「{comment_content}"
elif comment_author:
return f"踩了{comment_author}的一条评论"
return "踩了一条评论"
def _describe_search(self) -> str:
query = self.action_args.get("query", "") or self.action_args.get("keyword", "")
return f"搜索了「{query}" if query else "进行了搜索"
def _describe_search_user(self) -> str:
query = self.action_args.get("query", "") or self.action_args.get("username", "")
return f"搜索了用户「{query}" if query else "搜索了用户"
def _describe_mute(self) -> str:
target_user_name = self.action_args.get("target_user_name", "")
if target_user_name:
return f"屏蔽了用户「{target_user_name}"
return "屏蔽了一个用户"
def _describe_generic(self) -> str:
return f"执行了{self.action_type}操作"
class Neo4jGraphMemoryUpdater:
"""
Neo4j 图谱记忆更新器
监控模拟的 actions 日志文件将新的 agent 活动实时更新到 Neo4j 图谱中
按平台分组每累积 BATCH_SIZE 条活动后批量处理
"""
# 批量发送大小
BATCH_SIZE = 5
# 平台显示名称
PLATFORM_DISPLAY_NAMES = {
'twitter': '世界1',
'reddit': '世界2',
}
# 发送间隔(秒)
SEND_INTERVAL = 0.5
# 重试配置
MAX_RETRIES = 3
RETRY_DELAY = 2
def __init__(
self,
graph_id: str,
driver: 'Driver' = None,
llm_extractor: LLMExtractionPipeline = None
):
"""
初始化更新器
Args:
graph_id: 图谱ID
driver: Neo4j 驱动
llm_extractor: LLM 提取器
"""
from ...utils.neo4j.driver import get_neo4j_driver
self.graph_id = graph_id
self.driver = driver or get_neo4j_driver()
self.extractor = llm_extractor or LLMExtractionPipeline()
# 活动队列
self._activity_queue: Queue = Queue()
# 按平台分组的活动缓冲区
self._platform_buffers: Dict[str, List[AgentActivity]] = {
'twitter': [],
'reddit': [],
}
self._buffer_lock = threading.Lock()
# 控制标志
self._running = False
self._worker_thread: Optional[threading.Thread] = None
# 统计
self._total_activities = 0
self._total_sent = 0
self._total_items_sent = 0
self._failed_count = 0
self._skipped_count = 0
logger.info(
f"Neo4jGraphMemoryUpdater 初始化完成: "
f"graph_id={graph_id}, batch_size={self.BATCH_SIZE}"
)
def _get_platform_display_name(self, platform: str) -> str:
return self.PLATFORM_DISPLAY_NAMES.get(platform.lower(), platform)
def start(self):
"""启动后台工作线程"""
if self._running:
return
current_locale = get_locale()
self._running = True
self._worker_thread = threading.Thread(
target=self._worker_loop,
args=(current_locale,),
daemon=True,
name=f"Neo4jMemoryUpdater-{self.graph_id[:8]}"
)
self._worker_thread.start()
logger.info(f"Neo4jGraphMemoryUpdater 已启动: graph_id={self.graph_id}")
def stop(self):
"""停止后台工作线程"""
self._running = False
# 发送剩余的活动
self._flush_remaining()
if self._worker_thread and self._worker_thread.is_alive():
self._worker_thread.join(timeout=10)
logger.info(
f"Neo4jGraphMemoryUpdater 已停止: graph_id={self.graph_id}, "
f"total_activities={self._total_activities}, "
f"batches_sent={self._total_sent}, "
f"items_sent={self._total_items_sent}, "
f"failed={self._failed_count}, "
f"skipped={self._skipped_count}"
)
def add_activity(self, activity: AgentActivity):
"""
添加一个 agent 活动到队列
Args:
activity: Agent 活动记录
"""
if activity.action_type == "DO_NOTHING":
self._skipped_count += 1
return
self._activity_queue.put(activity)
self._total_activities += 1
logger.debug(
f"添加活动到队列: {activity.agent_name} - {activity.action_type}"
)
def add_activity_from_dict(self, data: Dict[str, Any], platform: str):
"""
从字典数据添加活动
Args:
data: actions.jsonl 解析的字典数据
platform: 平台名称 (twitter/reddit)
"""
# 跳过事件类型的条目
if "event_type" in data:
return
activity = AgentActivity(
platform=platform,
agent_id=data.get("agent_id", 0),
agent_name=data.get("agent_name", ""),
action_type=data.get("action_type", ""),
action_args=data.get("action_args", {}),
round_num=data.get("round", 0),
timestamp=data.get("timestamp", datetime.now().isoformat()),
)
self.add_activity(activity)
def _worker_loop(self, locale: str = 'zh'):
"""后台工作循环 - 按平台批量发送活动到 Neo4j"""
set_locale(locale)
while self._running or not self._activity_queue.empty():
try:
# 尝试从队列获取活动超时1秒
try:
activity = self._activity_queue.get(timeout=1)
# 将活动添加到对应平台的缓冲区
platform = activity.platform.lower()
with self._buffer_lock:
if platform not in self._platform_buffers:
self._platform_buffers[platform] = []
self._platform_buffers[platform].append(activity)
# 检查该平台是否达到批量大小
if len(self._platform_buffers[platform]) >= self.BATCH_SIZE:
batch = self._platform_buffers[platform][:self.BATCH_SIZE]
self._platform_buffers[platform] = self._platform_buffers[platform][self.BATCH_SIZE:]
# 释放锁后再发送
self._send_batch_activities(batch, platform)
# 发送间隔,避免请求过快
time.sleep(self.SEND_INTERVAL)
except Empty:
pass
except Exception as e:
logger.error(f"工作循环异常: {e}")
time.sleep(1)
def _send_batch_activities(
self,
activities: List[AgentActivity],
platform: str
):
"""
批量发送活动到 Neo4j 图谱
Args:
activities: Agent 活动列表
platform: 平台名称
"""
if not activities:
return
# 将活动转换为文本
activity_texts = [activity.to_episode_text() for activity in activities]
# 带重试的发送
for attempt in range(self.MAX_RETRIES):
try:
# 使用 LLM 提取实体和关系
ontology = self._get_activity_ontology()
entities, edges = self.extractor.extract_from_chunks(
chunks=activity_texts,
ontology=ontology,
graph_id=self.graph_id,
parallel_workers=1
)
# 写入 Neo4j
with self.driver.session() as session:
# 写入实体
for entity in entities:
self._write_activity_node(session, entity)
# 建立映射并写入边
name_to_uuid = self._get_name_to_uuid_mapping(session)
for edge in edges:
self._write_activity_edge(session, edge, name_to_uuid)
self._total_sent += 1
self._total_items_sent += len(activities)
display_name = self._get_platform_display_name(platform)
logger.info(
f"成功批量发送 {len(activities)}{display_name}活动到图谱 {self.graph_id}"
)
return
except Exception as e:
if attempt < self.MAX_RETRIES - 1:
logger.warning(
f"批量发送到 Neo4j 失败 (尝试 {attempt + 1}/{self.MAX_RETRIES}): {e}"
)
time.sleep(self.RETRY_DELAY * (attempt + 1))
else:
logger.error(
f"批量发送到 Neo4j 失败,已重试{self.MAX_RETRIES}次: {e}"
)
self._failed_count += 1
def _get_activity_ontology(self) -> Dict[str, Any]:
"""获取活动提取用的简化本体"""
return {
"entity_types": [
{"name": "Agent", "description": "模拟中的 AI Agent"},
{"name": "Post", "description": "社交媒体帖子"},
{"name": "Comment", "description": "评论"},
{"name": "User", "description": "用户"},
],
"edge_types": [
{"name": "POSTED", "description": "发布"},
{"name": "LIKED", "description": "点赞"},
{"name": "COMMENTED_ON", "description": "评论"},
{"name": "FOLLOWED", "description": "关注"},
{"name": "REPOSTED", "description": "转发"},
]
}
def _write_activity_node(self, session, entity) -> None:
"""写入活动节点"""
entity_type = _safe_neo4j_identifier(entity.entity_type or "Agent", "Agent")
labels = ["Entity", f"Entity_{entity_type}", "Activity"]
properties = {
"uuid": str(uuid.uuid4()),
"name": entity.name,
"summary": entity.summary,
"graph_id": self.graph_id,
"created_at": time.strftime("%Y-%m-%dT%H:%M:%S"),
"entity_type": entity_type
}
label_str = ":".join(labels)
set_clause = ", ".join([f"n.{k} = ${k}" for k in properties.keys()])
cypher = f"""
MERGE (n:{label_str} {{name: $name, graph_id: $graph_id}})
ON CREATE SET {set_clause}
ON MATCH SET {set_clause}
"""
try:
session.run(cypher, **properties)
except Exception as e:
logger.warning(f"写入活动节点失败: {entity.name}, {e}")
def _write_activity_edge(
self,
session,
edge,
name_to_uuid: Dict[str, str]
) -> None:
"""写入活动边"""
source_uuid = name_to_uuid.get(edge.source_name)
target_uuid = name_to_uuid.get(edge.target_name)
if not source_uuid or not target_uuid:
return
properties = {
"uuid": str(uuid.uuid4()),
"name": edge.name,
"fact": edge.fact,
"graph_id": self.graph_id,
"source_node_uuid": source_uuid,
"target_node_uuid": target_uuid,
"created_at": time.strftime("%Y-%m-%dT%H:%M:%S")
}
rel_type = _safe_neo4j_identifier(edge.name)
cypher = f"""
MATCH (source:Entity {{uuid: $source_node_uuid}})
MATCH (target:Entity {{uuid: $target_node_uuid}})
MERGE (source)-[r:`{rel_type}` {{
graph_id: $graph_id,
source_node_uuid: $source_node_uuid,
target_node_uuid: $target_node_uuid
}}]->(target)
SET r.uuid = coalesce(r.uuid, $uuid),
r.name = $name,
r.fact = $fact,
r.created_at = coalesce(r.created_at, $created_at)
"""
try:
session.run(cypher, **properties)
except Exception as e:
logger.warning(f"写入活动边失败: {edge.name}, {e}")
def _get_name_to_uuid_mapping(self, session) -> Dict[str, str]:
"""获取节点名称到 UUID 的映射"""
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN n.name AS name, n.uuid AS uuid
"""
result = session.run(cypher, graph_id=self.graph_id)
return {record["name"]: record["uuid"] for record in result}
def _flush_remaining(self):
"""发送队列和缓冲区中剩余的活动"""
# 首先处理队列中剩余的活动,添加到缓冲区
while not self._activity_queue.empty():
try:
activity = self._activity_queue.get_nowait()
platform = activity.platform.lower()
with self._buffer_lock:
if platform not in self._platform_buffers:
self._platform_buffers[platform] = []
self._platform_buffers[platform].append(activity)
except Empty:
break
# 发送各平台缓冲区中剩余的活动
with self._buffer_lock:
for platform, buffer in self._platform_buffers.items():
if buffer:
display_name = self._get_platform_display_name(platform)
logger.info(f"发送{display_name}平台剩余的 {len(buffer)} 条活动")
self._send_batch_activities(buffer, platform)
# 清空所有缓冲区
for platform in self._platform_buffers:
self._platform_buffers[platform] = []
def get_stats(self) -> Dict[str, Any]:
"""获取统计信息"""
with self._buffer_lock:
buffer_sizes = {p: len(b) for p, b in self._platform_buffers.items()}
return {
"graph_id": self.graph_id,
"batch_size": self.BATCH_SIZE,
"total_activities": self._total_activities,
"batches_sent": self._total_sent,
"items_sent": self._total_items_sent,
"failed_count": self._failed_count,
"skipped_count": self._skipped_count,
"queue_size": self._activity_queue.qsize(),
"buffer_sizes": buffer_sizes,
"running": self._running,
}
class Neo4jGraphMemoryManager:
"""
管理多个模拟的 Neo4j 图谱记忆更新器
每个模拟可以有自己的更新器实例
"""
_updaters: Dict[str, Neo4jGraphMemoryUpdater] = {}
_lock = threading.Lock()
@classmethod
def create_updater(
cls,
simulation_id: str,
graph_id: str,
driver=None,
llm_extractor=None
) -> Neo4jGraphMemoryUpdater:
"""
为模拟创建图谱记忆更新器
Args:
simulation_id: 模拟ID
graph_id: 图谱ID
driver: Neo4j 驱动
llm_extractor: LLM 提取器
Returns:
Neo4jGraphMemoryUpdater 实例
"""
with cls._lock:
# 如果已存在,先停止旧的
if simulation_id in cls._updaters:
cls._updaters[simulation_id].stop()
updater = Neo4jGraphMemoryUpdater(
graph_id=graph_id,
driver=driver,
llm_extractor=llm_extractor
)
updater.start()
cls._updaters[simulation_id] = updater
logger.info(
f"创建图谱记忆更新器: simulation_id={simulation_id}, graph_id={graph_id}"
)
return updater
@classmethod
def get_updater(cls, simulation_id: str) -> Optional[Neo4jGraphMemoryUpdater]:
"""获取模拟的更新器"""
return cls._updaters.get(simulation_id)
@classmethod
def stop_updater(cls, simulation_id: str):
"""停止并移除模拟的更新器"""
with cls._lock:
if simulation_id in cls._updaters:
cls._updaters[simulation_id].stop()
del cls._updaters[simulation_id]
logger.info(f"已停止图谱记忆更新器: simulation_id={simulation_id}")
@classmethod
def stop_all(cls):
"""停止所有更新器"""
if cls._updaters:
for simulation_id, updater in list(cls._updaters.items()):
try:
updater.stop()
except Exception as e:
logger.error(f"停止更新器失败: simulation_id={simulation_id}, error={e}")
cls._updaters.clear()
logger.info("已停止所有图谱记忆更新器")
@classmethod
def get_all_stats(cls) -> Dict[str, Dict[str, Any]]:
"""获取所有更新器的统计信息"""
return {
sim_id: updater.get_stats()
for sim_id, updater in cls._updaters.items()
}
# 导入 uuid
import uuid

View File

@ -0,0 +1,965 @@
"""
Neo4j 检索工具服务
封装图谱搜索节点读取边查询等工具 Report Agent 使用
替代 ZepToolsService
"""
import json
import re
import time
from typing import Dict, Any, List, Optional, TYPE_CHECKING
if TYPE_CHECKING:
from neo4j import Driver
from .graph_adapter import SearchResult
from .neo4j_entity_reader import Neo4jEntityReader, EntityNode
from ...utils.llm_client import LLMClient
from ...utils.logger import get_logger
from ...utils.locale import get_locale
logger = get_logger('mirofish.neo4j_search')
class NodeInfo:
"""节点信息"""
uuid: str
name: str
labels: List[str]
summary: str
attributes: Dict[str, Any]
def __init__(
self,
uuid: str,
name: str,
labels: List[str],
summary: str,
attributes: Dict[str, Any]
):
self.uuid = uuid
self.name = name
self.labels = labels
self.summary = summary
self.attributes = attributes
def to_dict(self) -> Dict[str, Any]:
return {
"uuid": self.uuid,
"name": self.name,
"labels": self.labels,
"summary": self.summary,
"attributes": self.attributes
}
def to_text(self) -> str:
entity_type = next(
(l for l in self.labels if l not in ("Entity", "Node")),
"未知类型"
)
return f"实体: {self.name} (类型: {entity_type})\n摘要: {self.summary}"
class EdgeInfo:
"""边信息"""
uuid: str
name: str
fact: str
source_node_uuid: str
target_node_uuid: str
source_node_name: Optional[str] = None
target_node_name: Optional[str] = None
created_at: Optional[str] = None
valid_at: Optional[str] = None
invalid_at: Optional[str] = None
expired_at: Optional[str] = None
def __init__(self, **kwargs):
for key, value in kwargs.items():
setattr(self, key, value)
def to_dict(self) -> Dict[str, Any]:
return {
"uuid": self.uuid,
"name": self.name,
"fact": self.fact,
"source_node_uuid": self.source_node_uuid,
"target_node_uuid": self.target_node_uuid,
"source_node_name": self.source_node_name,
"target_node_name": self.target_node_name,
"created_at": self.created_at,
"valid_at": self.valid_at,
"invalid_at": self.invalid_at,
"expired_at": self.expired_at
}
def to_text(self, include_temporal: bool = False) -> str:
source = self.source_node_name or self.source_node_uuid[:8]
target = self.target_node_name or self.target_node_uuid[:8]
base_text = f"关系: {source} --[{self.name}]--> {target}\n事实: {self.fact}"
if include_temporal:
valid_at = self.valid_at or "未知"
invalid_at = self.invalid_at or "至今"
base_text += f"\n时效: {valid_at} - {invalid_at}"
if self.expired_at:
base_text += f" (已过期: {self.expired_at})"
return base_text
@property
def is_expired(self) -> bool:
return self.expired_at is not None
@property
def is_invalid(self) -> bool:
return self.invalid_at is not None
class Neo4jSearchService:
"""
Neo4j 检索工具服务
提供图谱搜索节点查询边查询等功能
适配 Report Agent 的工具调用需求
"""
MAX_RETRIES = 3
RETRY_DELAY = 2.0
STOPWORDS = {
"", "", "", "", "", "", "", "", "", "", "",
"分析", "结果", "预测", "模拟", "引擎", "机制", "方法", "能力",
"the", "and", "or", "of", "for", "to", "in", "on", "with",
}
def __init__(
self,
driver: 'Driver' = None,
llm_client: Optional[LLMClient] = None
):
"""
初始化检索服务
Args:
driver: Neo4j 驱动
llm_client: LLM 客户端用于 InsightForge 等需要 LLM 的功能
"""
from ...utils.neo4j.driver import get_neo4j_driver
self.driver = driver or get_neo4j_driver()
self.entity_reader = Neo4jEntityReader(driver=self.driver)
self._llm_client = llm_client
@property
def llm(self) -> LLMClient:
"""延迟初始化 LLM 客户端"""
if self._llm_client is None:
self._llm_client = LLMClient()
return self._llm_client
def _call_with_retry(self, func, operation_name: str, max_retries: int = None):
"""带重试机制的查询"""
max_retries = max_retries or self.MAX_RETRIES
last_exception = None
delay = self.RETRY_DELAY
for attempt in range(max_retries):
try:
return func()
except Exception as e:
last_exception = e
if attempt < max_retries - 1:
logger.warning(
f"Neo4j {operation_name}{attempt + 1} 次尝试失败: "
f"{str(e)[:100]}, {delay:.1f}秒后重试..."
)
time.sleep(delay)
delay *= 2
else:
logger.error(
f"Neo4j {operation_name}{max_retries} 次尝试后仍失败: {str(e)}"
)
raise last_exception
def _extract_search_terms(self, query: str) -> List[str]:
"""Split an LLM-style long query into useful searchable terms."""
raw_terms = re.findall(r"[A-Za-z0-9_.+-]+|[\u4e00-\u9fff]{2,}", query or "")
terms: List[str] = []
for term in raw_terms:
normalized = term.strip()
if not normalized:
continue
lower = normalized.lower()
if lower in self.STOPWORDS or len(normalized) < 2:
continue
terms.append(normalized)
# Add short Chinese fragments for long concatenated phrases.
for term in list(terms):
if re.fullmatch(r"[\u4e00-\u9fff]{3,}", term):
for size in (2, 3, 4):
for index in range(0, len(term) - size + 1):
fragment = term[index:index + size]
if fragment not in self.STOPWORDS:
terms.append(fragment)
deduped: List[str] = []
seen = set()
for term in terms:
key = term.lower()
if key not in seen:
seen.add(key)
deduped.append(term)
return deduped[:40]
def search_graph(
self,
graph_id: str,
query: str,
limit: int = 10,
scope: str = "edges"
) -> SearchResult:
"""
图谱语义搜索
Args:
graph_id: 图谱ID
query: 搜索查询
limit: 返回结果数量
scope: 搜索范围 ("edges" / "nodes" / "both")
Returns:
SearchResult: 搜索结果
"""
logger.info(f"搜索图谱 {graph_id}: query={query[:50]}")
facts = []
edges_result = []
nodes_result = []
terms = self._extract_search_terms(query)
with self.driver.session() as session:
if scope in ("edges", "both"):
# 搜索边
edge_cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
AND (
toLower(coalesce(r.fact, '')) CONTAINS toLower($search_query)
OR toLower(coalesce(r.name, '')) CONTAINS toLower($search_query)
OR any(term IN $terms WHERE
toLower(coalesce(r.fact, '')) CONTAINS toLower(term)
OR toLower(coalesce(r.name, '')) CONTAINS toLower(term)
OR toLower(coalesce(source.name, '')) CONTAINS toLower(term)
OR toLower(coalesce(target.name, '')) CONTAINS toLower(term)
)
)
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
source.name AS source_name, target.name AS target_name,
properties(r) AS attributes,
reduce(score = 0, term IN $terms |
score
+ CASE WHEN toLower(coalesce(r.fact, '')) CONTAINS toLower(term) THEN 10 ELSE 0 END
+ CASE WHEN toLower(coalesce(r.name, '')) CONTAINS toLower(term) THEN 5 ELSE 0 END
+ CASE WHEN toLower(coalesce(source.name, '')) CONTAINS toLower(term) THEN 3 ELSE 0 END
+ CASE WHEN toLower(coalesce(target.name, '')) CONTAINS toLower(term) THEN 3 ELSE 0 END
) AS score
ORDER BY score DESC
LIMIT $limit
"""
result = session.run(
edge_cypher,
graph_id=graph_id,
search_query=query,
terms=terms,
limit=limit
)
for record in result:
if record["fact"]:
facts.append(record["fact"])
edges_result.append({
"uuid": record["uuid"],
"name": record["name"],
"fact": record["fact"],
"source_node_uuid": record["source_node_uuid"],
"target_node_uuid": record["target_node_uuid"],
"source_name": record["source_name"],
"target_name": record["target_name"],
})
if scope in ("nodes", "both"):
# 搜索节点
node_cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
AND (
toLower(coalesce(n.name, '')) CONTAINS toLower($search_query)
OR toLower(coalesce(n.summary, '')) CONTAINS toLower($search_query)
OR any(term IN $terms WHERE
toLower(coalesce(n.name, '')) CONTAINS toLower(term)
OR toLower(coalesce(n.summary, '')) CONTAINS toLower(term)
OR toLower(coalesce(n.entity_type, '')) CONTAINS toLower(term)
)
)
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, n.entity_type AS entity_type,
properties(n) AS attributes,
reduce(score = 0, term IN $terms |
score
+ CASE WHEN toLower(coalesce(n.name, '')) CONTAINS toLower(term) THEN 10 ELSE 0 END
+ CASE WHEN toLower(coalesce(n.summary, '')) CONTAINS toLower(term) THEN 6 ELSE 0 END
+ CASE WHEN toLower(coalesce(n.entity_type, '')) CONTAINS toLower(term) THEN 3 ELSE 0 END
) AS score
ORDER BY score DESC
LIMIT $limit
"""
result = session.run(
node_cypher,
graph_id=graph_id,
search_query=query,
terms=terms,
limit=limit
)
for record in result:
labels = [l for l in record["labels"] if l not in ("Entity", "Node")]
if record["summary"]:
facts.append(f"[{record['name']}]: {record['summary']}")
nodes_result.append({
"uuid": record["uuid"],
"name": record["name"],
"labels": labels,
"summary": record["summary"],
})
if not facts and not edges_result and not nodes_result:
fallback_cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id AND coalesce(n.summary, '') <> ''
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, properties(n) AS attributes
LIMIT $limit
"""
result = session.run(fallback_cypher, graph_id=graph_id, limit=limit)
for record in result:
labels = [l for l in record["labels"] if l not in ("Entity", "Node")]
if record["summary"]:
facts.append(f"[{record['name']}]: {record['summary']}")
nodes_result.append({
"uuid": record["uuid"],
"name": record["name"],
"labels": labels,
"summary": record["summary"],
})
logger.info(f"搜索完成: {len(facts)} 条相关结果")
return SearchResult(
facts=facts,
edges=edges_result,
nodes=nodes_result,
query=query,
total_count=len(facts) or (len(edges_result) + len(nodes_result))
)
def get_all_nodes(self, graph_id: str) -> List[NodeInfo]:
"""
获取图谱的所有节点
Args:
graph_id: 图谱ID
Returns:
节点列表
"""
logger.info(f"获取图谱 {graph_id} 的所有节点")
def _query():
nodes = []
with self.driver.session() as session:
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, properties(n) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
nodes.append(NodeInfo(
uuid=record["uuid"],
name=record["name"],
labels=record["labels"],
summary=record["summary"] or "",
attributes=record["attributes"] or {}
))
return nodes
return self._call_with_retry(_query, f"获取所有节点({graph_id})")
def get_all_edges(self, graph_id: str, include_temporal: bool = True) -> List[EdgeInfo]:
"""
获取图谱的所有边
Args:
graph_id: 图谱ID
include_temporal: 是否包含时间信息
Returns:
边列表
"""
logger.info(f"获取图谱 {graph_id} 的所有边")
def _query():
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
source.name AS source_node_name,
target.name AS target_node_name,
properties(r) AS attributes
"""
result = session.run(cypher, graph_id=graph_id)
for record in result:
attributes = record["attributes"] or {}
edges.append(EdgeInfo(
uuid=record["uuid"],
name=record["name"],
fact=record["fact"] or "",
source_node_uuid=record["source_node_uuid"],
target_node_uuid=record["target_node_uuid"],
source_node_name=record.get("source_node_name"),
target_node_name=record.get("target_node_name"),
created_at=str(attributes.get("created_at")) if attributes.get("created_at") else None,
valid_at=str(attributes.get("valid_at")) if attributes.get("valid_at") else None,
invalid_at=str(attributes.get("invalid_at")) if attributes.get("invalid_at") else None,
expired_at=str(attributes.get("expired_at")) if attributes.get("expired_at") else None,
attributes=attributes
))
return edges
return self._call_with_retry(_query, f"获取所有边({graph_id})")
def get_node_detail(self, node_uuid: str) -> Optional[NodeInfo]:
"""
获取单个节点的详细信息
Args:
node_uuid: 节点UUID
Returns:
节点信息或 None
"""
def _query():
with self.driver.session() as session:
cypher = """
MATCH (n:Entity)
WHERE n.uuid = $uuid
RETURN n.uuid AS uuid, n.name AS name, labels(n) AS labels,
n.summary AS summary, properties(n) AS attributes
"""
result = session.run(cypher, uuid=node_uuid)
record = result.single()
if not record:
return None
return NodeInfo(
uuid=record["uuid"],
name=record["name"],
labels=record["labels"],
summary=record["summary"] or "",
attributes=record["attributes"] or {}
)
try:
return self._call_with_retry(_query, f"获取节点详情({node_uuid[:8]}...)")
except Exception as e:
logger.error(f"获取节点详情失败: {e}")
return None
def get_node_edges(self, graph_id: str, node_uuid: str) -> List[EdgeInfo]:
"""
获取节点相关的所有边
Args:
graph_id: 图谱ID
node_uuid: 节点UUID
Returns:
边列表
"""
def _query():
edges = []
with self.driver.session() as session:
cypher = """
MATCH (source)-[r]->(target)
WHERE (r.source_node_uuid = $uuid OR r.target_node_uuid = $uuid
OR source.uuid = $uuid OR target.uuid = $uuid)
AND r.graph_id = $graph_id
RETURN r.uuid AS uuid, type(r) AS name, r.fact AS fact,
coalesce(r.source_node_uuid, source.uuid) AS source_node_uuid,
coalesce(r.target_node_uuid, target.uuid) AS target_node_uuid,
source.name AS source_node_name,
target.name AS target_node_name,
r.created_at AS created_at
"""
result = session.run(cypher, uuid=node_uuid, graph_id=graph_id)
for record in result:
edges.append(EdgeInfo(
uuid=record["uuid"],
name=record["name"],
fact=record["fact"] or "",
source_node_uuid=record["source_node_uuid"],
target_node_uuid=record["target_node_uuid"],
source_node_name=record.get("source_node_name"),
target_node_name=record.get("target_node_name"),
created_at=str(record["created_at"]) if record.get("created_at") else None,
))
return edges
return self._call_with_retry(
_query, f"获取节点边({node_uuid[:8]}...)"
)
def get_entities_by_type(
self,
graph_id: str,
entity_type: str
) -> List[NodeInfo]:
"""
按类型获取实体
Args:
graph_id: 图谱ID
entity_type: 实体类型
Returns:
符合类型的实体列表
"""
all_nodes = self.get_all_nodes(graph_id)
filtered = [n for n in all_nodes if entity_type in n.labels]
logger.info(f"按类型 {entity_type} 获取实体: {len(filtered)}")
return filtered
def get_entity_summary(
self,
graph_id: str,
entity_name: str
) -> Dict[str, Any]:
"""
获取指定实体的关系摘要
Args:
graph_id: 图谱ID
entity_name: 实体名称
Returns:
实体摘要信息
"""
logger.info(f"获取实体摘要: {entity_name}")
# 搜索相关实体
search_result = self.search_graph(
graph_id=graph_id,
query=entity_name,
limit=20
)
# 在所有节点中查找该实体
all_nodes = self.get_all_nodes(graph_id)
entity_node = None
for node in all_nodes:
if node.name.lower() == entity_name.lower():
entity_node = node
break
# 获取关联边
related_edges = []
if entity_node:
related_edges = self.get_node_edges(graph_id, entity_node.uuid)
return {
"entity_name": entity_name,
"entity_info": entity_node.to_dict() if entity_node else None,
"related_facts": search_result.facts,
"related_edges": [e.to_dict() for e in related_edges],
"total_relations": len(related_edges)
}
def get_graph_statistics(self, graph_id: str) -> Dict[str, Any]:
"""
获取图谱的统计信息
Args:
graph_id: 图谱ID
Returns:
统计信息
"""
logger.info(f"获取图谱统计: {graph_id}")
nodes = self.get_all_nodes(graph_id)
edges = self.get_all_edges(graph_id)
# 统计实体类型分布
entity_types = {}
for node in nodes:
for label in node.labels:
if label not in ("Entity", "Node"):
entity_types[label] = entity_types.get(label, 0) + 1
# 统计关系类型分布
relation_types = {}
for edge in edges:
relation_types[edge.name] = relation_types.get(edge.name, 0) + 1
return {
"graph_id": graph_id,
"total_nodes": len(nodes),
"total_edges": len(edges),
"entity_types": entity_types,
"relation_types": relation_types
}
def get_simulation_context(
self,
graph_id: str,
simulation_requirement: str = "",
max_facts: int = 80
) -> Dict[str, Any]:
"""
构建报告生成所需的图谱上下文兼容 ZepToolsService
"""
stats = self.get_graph_statistics(graph_id)
edges = self.get_all_edges(graph_id)
nodes = self.get_all_nodes(graph_id)
facts = [edge.fact for edge in edges if edge.fact][:max_facts]
return {
"graph_id": graph_id,
"simulation_requirement": simulation_requirement,
"statistics": stats,
"summary": (
f"图谱包含 {stats.get('total_nodes', 0)} 个实体、"
f"{stats.get('total_edges', 0)} 条关系。"
),
"key_facts": facts,
"entity_types": stats.get("entity_types", {}),
"relation_types": stats.get("relation_types", {}),
"sample_entities": [node.to_dict() for node in nodes[:20]],
}
def interview_agents(
self,
graph_id: str = None,
question: str = None,
entity_names: Optional[List[str]] = None,
max_agents: int = 5,
simulation_id: str = None,
interview_requirement: str = None,
simulation_requirement: str = "",
custom_questions: Optional[List[str]] = None
) -> Dict[str, Any]:
"""
报告工具兼容方法真实采访由 simulation_runner 负责这里提供图谱上下文降级
"""
question = question or interview_requirement or ""
targets = entity_names or []
if not targets:
targets = [node.name for node in self.get_all_nodes(graph_id)[:max_agents]]
responses = []
for name in targets[:max_agents]:
summary = self.get_entity_summary(graph_id, name)
responses.append({
"agent": name,
"question": question,
"response": summary,
"source": "neo4j_graph_context"
})
return {
"question": question,
"responses": responses,
"count": len(responses)
}
def insight_forge(
self,
graph_id: str,
query: str,
simulation_requirement: str,
report_context: str = "",
max_sub_queries: int = 5
) -> Dict[str, Any]:
"""
深度洞察检索
使用 LLM 将问题分解为多个子问题然后对每个子问题进行搜索
Args:
graph_id: 图谱ID
query: 用户问题
simulation_requirement: 模拟需求描述
report_context: 报告上下文
max_sub_queries: 最大子问题数量
Returns:
深度洞察检索结果
"""
logger.info(f"InsightForge: query={query[:50]}")
# Step 1: 生成子问题
sub_queries = self._generate_sub_queries(
query=query,
simulation_requirement=simulation_requirement,
report_context=report_context,
max_queries=max_sub_queries
)
# Step 2: 对每个子问题进行搜索
all_facts = []
all_edges = []
seen_facts = set()
for sub_query in sub_queries:
result = self.search_graph(
graph_id=graph_id,
query=sub_query,
limit=15,
scope="both"
)
for fact in result.facts:
if fact not in seen_facts:
all_facts.append(fact)
seen_facts.add(fact)
all_edges.extend(result.edges)
# 对原始问题也进行搜索
main_result = self.search_graph(
graph_id=graph_id,
query=query,
limit=20,
scope="both"
)
for fact in main_result.facts:
if fact not in seen_facts:
all_facts.append(fact)
seen_facts.add(fact)
all_edges.extend(main_result.edges)
if not all_facts:
panorama = self.panorama_search(
graph_id=graph_id,
query=query,
include_expired=True,
limit=20
)
for fact in panorama.get("active_facts", []) + panorama.get("historical_facts", []):
if fact not in seen_facts:
all_facts.append(fact)
seen_facts.add(fact)
for node in panorama.get("all_nodes", [])[:20]:
summary = node.get("summary") or ""
name = node.get("name") or "未知实体"
if summary:
fact = f"{name}: {summary}"
if fact not in seen_facts:
all_facts.append(fact)
seen_facts.add(fact)
# Step 3: 获取相关实体
entity_uuids = set()
for edge_data in all_edges:
if isinstance(edge_data, dict):
entity_uuids.add(edge_data.get('source_node_uuid', ''))
entity_uuids.add(edge_data.get('target_node_uuid', ''))
entity_insights = []
for uuid in entity_uuids:
if not uuid:
continue
node = self.get_node_detail(uuid)
if node:
related_facts = [f for f in all_facts if node.name.lower() in f.lower()]
entity_insights.append({
"uuid": node.uuid,
"name": node.name,
"type": next((l for l in node.labels if l not in ("Entity", "Node")), "实体"),
"summary": node.summary,
"related_facts": related_facts
})
# Step 4: 构建关系链
relationship_chains = []
node_map = {e["uuid"]: e for e in entity_insights}
for edge_data in all_edges:
if isinstance(edge_data, dict):
source_uuid = edge_data.get('source_node_uuid', '')
target_uuid = edge_data.get('target_node_uuid', '')
relation_name = edge_data.get('name', '')
source_name = node_map.get(source_uuid, {}).get('name', '') or source_uuid[:8]
target_name = node_map.get(target_uuid, {}).get('name', '') or target_uuid[:8]
chain = f"{source_name} --[{relation_name}]--> {target_name}"
if chain not in relationship_chains:
relationship_chains.append(chain)
return {
"query": query,
"simulation_requirement": simulation_requirement,
"sub_queries": sub_queries,
"semantic_facts": all_facts,
"entity_insights": entity_insights,
"relationship_chains": relationship_chains,
"total_facts": len(all_facts),
"total_entities": len(entity_insights),
"total_relationships": len(relationship_chains)
}
def _generate_sub_queries(
self,
query: str,
simulation_requirement: str,
report_context: str = "",
max_queries: int = 5
) -> List[str]:
"""使用 LLM 生成子问题"""
system_prompt = """你是一个专业的问题分析专家。你的任务是将一个复杂问题分解为多个可以在模拟世界中独立观察的子问题。
要求
1. 每个子问题应该足够具体可以在模拟世界中找到相关的Agent行为或事件
2. 子问题应该覆盖原问题的不同维度什么为什么怎么样何时何地
3. 子问题应该与模拟场景相关
4. 返回JSON格式{"sub_queries": ["子问题1", "子问题2", ...]}"""
user_prompt = f"""模拟需求背景:
{simulation_requirement}
{f"报告上下文:{report_context[:500]}" if report_context else ""}
请将以下问题分解为{max_queries}个子问题
{query}
返回JSON格式的子问题列表"""
try:
response = self.llm.chat_json(
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
temperature=0.3
)
sub_queries = response.get("sub_queries", [])
return [str(sq) for sq in sub_queries[:max_queries]]
except Exception as e:
logger.warning(f"生成子问题失败: {e}")
# 降级:返回基于原问题的变体
return [
query,
f"{query} 的主要参与者",
f"{query} 的原因和影响",
f"{query} 的发展过程"
][:max_queries]
def panorama_search(
self,
graph_id: str,
query: str,
include_expired: bool = True,
limit: int = 50
) -> Dict[str, Any]:
"""
广度搜索
获取全貌视图包括所有相关内容和历史/过期信息
Args:
graph_id: 图谱ID
query: 搜索查询
include_expired: 是否包含过期内容
limit: 返回结果数量限制
Returns:
广度搜索结果
"""
logger.info(f"PanoramaSearch: query={query[:50]}")
# 获取所有节点
all_nodes = self.get_all_nodes(graph_id)
node_map = {n.uuid: n for n in all_nodes}
# 获取所有边
all_edges = self.get_all_edges(graph_id, include_temporal=True)
# 分类事实
active_facts = []
historical_facts = []
for edge in all_edges:
if not edge.fact:
continue
# 判断是否过期/失效
is_historical = edge.is_expired or edge.is_invalid
if is_historical:
valid_at = edge.valid_at or "未知"
invalid_at = edge.invalid_at or edge.expired_at or "未知"
fact_with_time = f"[{valid_at} - {invalid_at}] {edge.fact}"
historical_facts.append(fact_with_time)
else:
active_facts.append(edge.fact)
# 排序并限制数量
active_facts.sort(key=lambda x: query.lower() in x.lower(), reverse=True)
historical_facts.sort(key=lambda x: query.lower() in x.lower(), reverse=True)
return {
"query": query,
"all_nodes": [n.to_dict() for n in all_nodes],
"all_edges": [e.to_dict() for e in all_edges],
"active_facts": active_facts[:limit],
"historical_facts": historical_facts[:limit] if include_expired else [],
"total_nodes": len(all_nodes),
"total_edges": len(all_edges),
"active_count": len(active_facts),
"historical_count": len(historical_facts)
}
def quick_search(
self,
graph_id: str,
query: str,
limit: int = 10
) -> SearchResult:
"""
快速搜索
Args:
graph_id: 图谱ID
query: 搜索查询
limit: 返回结果数量
Returns:
搜索结果
"""
logger.info(f"QuickSearch: query={query[:50]}")
return self.search_graph(graph_id=graph_id, query=query, limit=limit, scope="both")

View File

@ -0,0 +1,371 @@
"""
Football score probability simulation utilities.
This module extracts football prediction inputs from a simulation folder and
runs a reproducible bivariate-Poisson Monte Carlo simulation.
"""
import hashlib
import json
import math
import os
import random
import re
from collections import Counter
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Tuple
from ..config import Config
@dataclass
class FootballPredictionInputs:
home_team: str
away_team: str
lambda_home: float
lambda_away: float
samples: int
correlation: float
source: str
warnings: List[str]
class FootballProbabilitySimulator:
"""Extracts inputs and runs score probability simulations for football reports."""
DEFAULT_SAMPLES = 100_000
SCORE_MATRIX_MAX = 6
FOOTBALL_KEYWORDS = (
"足球", "football", "泊松", "poisson", "比分", "score",
"lambda_home", "lambda_away", "expected_goals", "xg",
)
@classmethod
def should_run(cls, simulation_requirement: str) -> bool:
text = (simulation_requirement or "").lower()
return any(keyword.lower() in text for keyword in cls.FOOTBALL_KEYWORDS)
@classmethod
def simulate_from_simulation(
cls,
simulation_id: str,
simulation_requirement: str,
samples: int = DEFAULT_SAMPLES,
) -> Optional[Dict[str, Any]]:
if not cls.should_run(simulation_requirement):
return None
config = cls._load_simulation_config(simulation_id)
texts = cls._collect_texts(simulation_id, config, simulation_requirement)
inputs = cls._extract_inputs(config, texts, samples=samples)
if not inputs:
return None
return cls._run_bivariate_poisson(inputs)
@classmethod
def _load_simulation_config(cls, simulation_id: str) -> Dict[str, Any]:
path = os.path.join(Config.OASIS_SIMULATION_DATA_DIR, simulation_id, "simulation_config.json")
if not os.path.exists(path):
return {}
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
@classmethod
def _collect_texts(
cls,
simulation_id: str,
config: Dict[str, Any],
simulation_requirement: str,
) -> List[str]:
texts: List[str] = [simulation_requirement or ""]
for post in config.get("event_config", {}).get("initial_posts", []) or []:
content = post.get("content")
if content:
texts.append(str(content))
sim_dir = os.path.join(Config.OASIS_SIMULATION_DATA_DIR, simulation_id)
for platform in ("twitter", "reddit"):
actions_path = os.path.join(sim_dir, platform, "actions.jsonl")
if not os.path.exists(actions_path):
continue
try:
with open(actions_path, "r", encoding="utf-8") as f:
for index, line in enumerate(f):
if index >= 200:
break
try:
item = json.loads(line)
except json.JSONDecodeError:
continue
args = item.get("action_args") or {}
content = args.get("content") or item.get("content")
if content:
texts.append(str(content))
except OSError:
continue
return texts
@classmethod
def _extract_inputs(
cls,
config: Dict[str, Any],
texts: List[str],
samples: int,
) -> Optional[FootballPredictionInputs]:
warnings: List[str] = []
joined = "\n".join(texts)
home_team, away_team = cls._extract_teams(config, joined)
lambda_home, lambda_away, source = cls._extract_lambdas(joined, home_team, away_team)
if lambda_home is None or lambda_away is None:
warnings.append("未找到明确的 lambda_home/lambda_away 数值,无法生成比分概率分布。")
return None
correlation = cls._infer_correlation(joined, lambda_home, lambda_away)
return FootballPredictionInputs(
home_team=home_team,
away_team=away_team,
lambda_home=lambda_home,
lambda_away=lambda_away,
samples=samples,
correlation=correlation,
source=source,
warnings=warnings,
)
@classmethod
def _extract_teams(cls, config: Dict[str, Any], text: str) -> Tuple[str, str]:
configured = [
agent.get("entity_name")
for agent in config.get("agent_configs", []) or []
if agent.get("entity_type") == "Team" and agent.get("entity_name")
]
match = re.search(r"([\w\u4e00-\u9fff]+)\s*(?:vs|VS|对阵|迎战|挑战)\s*([\w\u4e00-\u9fff]+)", text)
if match:
first, second = match.group(1), match.group(2)
if "主场" in text[max(0, match.start() - 30):match.end() + 60]:
return first, second
if "Qatar" in text or "卡塔尔" in text:
if "Switzerland" in text or "瑞士" in text:
return "Qatar", "Switzerland"
if len(configured) >= 2:
return configured[1], configured[0]
return "Home", "Away"
@classmethod
def _extract_lambdas(
cls,
text: str,
home_team: str,
away_team: str,
) -> Tuple[Optional[float], Optional[float], str]:
patterns = [
(r"lambda_home\s*[:=]\s*([0-9]+(?:\.[0-9]+)?)", "home"),
(r"lambda_away\s*[:=]\s*([0-9]+(?:\.[0-9]+)?)", "away"),
(r"主队[^0-9]{0,20}(?:预期进球|expected_goals|lambda)[^0-9]{0,10}([0-9]+(?:\.[0-9]+)?)", "home"),
(r"客队[^0-9]{0,20}(?:预期进球|expected_goals|lambda)[^0-9]{0,10}([0-9]+(?:\.[0-9]+)?)", "away"),
(r"卡塔尔主场\s*([0-9]+(?:\.[0-9]+)?)", "home"),
(r"瑞士客场(?:预期进球)?\s*([0-9]+(?:\.[0-9]+)?)", "away"),
]
values: Dict[str, float] = {}
for pattern, key in patterns:
match = re.search(pattern, text, re.IGNORECASE)
if match:
values[key] = float(match.group(1))
if "home" in values and "away" in values:
return values["home"], values["away"], "explicit_lambda"
# Common prose: "瑞士客场预期进球1.71卡塔尔主场0.87".
away_match = re.search(r"瑞士[^。\n]{0,20}(?:预期进球|expected_goals|lambda值?)[^0-9]{0,8}([0-9]+(?:\.[0-9]+)?)", text, re.IGNORECASE)
home_match = re.search(r"卡塔尔[^。\n]{0,20}(?:预期进球|expected_goals|主场)[^0-9]{0,8}([0-9]+(?:\.[0-9]+)?)", text, re.IGNORECASE)
if home_match and away_match:
return float(home_match.group(1)), float(away_match.group(1)), "prose_lambda"
# Fallback for "expected_goals 0.98 vs 瑞士的1.63".
eg_match = re.search(
r"expected_goals\s*([0-9]+(?:\.[0-9]+)?)\s*(?:vs|VS|对)\s*(?:瑞士|Switzerland)[^0-9]{0,8}([0-9]+(?:\.[0-9]+)?)",
text,
re.IGNORECASE,
)
if eg_match:
return float(eg_match.group(1)), float(eg_match.group(2)), "expected_goals_pair"
return None, None, "missing"
@classmethod
def _infer_correlation(cls, text: str, lambda_home: float, lambda_away: float) -> float:
defense_values = [float(v) for v in re.findall(r"防守(?:指数|评分)?[^0-9]{0,6}([0-9]+(?:\.[0-9]+)?)", text)]
if defense_values:
strongest = max(defense_values)
return min(0.12, max(0.03, (strongest - 70.0) / 250.0))
gap = abs(lambda_home - lambda_away)
return min(0.08, max(0.03, gap / 20.0))
@classmethod
def _run_bivariate_poisson(cls, inputs: FootballPredictionInputs) -> Dict[str, Any]:
common_lambda = min(inputs.lambda_home, inputs.lambda_away) * inputs.correlation
home_base = max(0.001, inputs.lambda_home - common_lambda)
away_base = max(0.001, inputs.lambda_away - common_lambda)
seed_basis = f"{inputs.home_team}|{inputs.away_team}|{inputs.lambda_home}|{inputs.lambda_away}|{inputs.samples}|{inputs.correlation}"
seed = int(hashlib.sha256(seed_basis.encode("utf-8")).hexdigest()[:16], 16)
rng = random.Random(seed)
scores: Counter[Tuple[int, int]] = Counter()
home_wins = draws = away_wins = 0
home_goals_total = away_goals_total = 0
for _ in range(inputs.samples):
home_goals = cls._sample_poisson(rng, home_base) + cls._sample_poisson(rng, common_lambda)
away_goals = cls._sample_poisson(rng, away_base) + cls._sample_poisson(rng, common_lambda)
scores[(home_goals, away_goals)] += 1
home_goals_total += home_goals
away_goals_total += away_goals
if home_goals > away_goals:
home_wins += 1
elif home_goals == away_goals:
draws += 1
else:
away_wins += 1
top_scores = [
{"score": f"{home}-{away}", "prob": round(count / inputs.samples, 4)}
for (home, away), count in scores.most_common(8)
]
matrix = []
for home in range(cls.SCORE_MATRIX_MAX + 1):
row = []
for away in range(cls.SCORE_MATRIX_MAX + 1):
row.append(round(scores.get((home, away), 0) / inputs.samples, 4))
matrix.append(row)
overflow = sum(
count for (home, away), count in scores.items()
if home > cls.SCORE_MATRIX_MAX or away > cls.SCORE_MATRIX_MAX
)
return {
"kind": "football_score_prediction",
"method": "bivariate_poisson_monte_carlo",
"result": {
"win_prob": {
"home": round(home_wins / inputs.samples, 4),
"draw": round(draws / inputs.samples, 4),
"away": round(away_wins / inputs.samples, 4),
},
"top_scores": top_scores,
"expected_goals": {
"home": round(home_goals_total / inputs.samples, 3),
"away": round(away_goals_total / inputs.samples, 3),
},
"score_distribution_matrix": {
"home_goals": list(range(cls.SCORE_MATRIX_MAX + 1)),
"away_goals": list(range(cls.SCORE_MATRIX_MAX + 1)),
"probabilities": matrix,
"overflow_prob": round(overflow / inputs.samples, 4),
},
},
"inputs": {
"home_team": inputs.home_team,
"away_team": inputs.away_team,
"lambda_home": inputs.lambda_home,
"lambda_away": inputs.lambda_away,
"samples": inputs.samples,
"correlation": round(inputs.correlation, 4),
"source": inputs.source,
},
"warnings": inputs.warnings,
}
@staticmethod
def _sample_poisson(rng: random.Random, lam: float) -> int:
if lam <= 0:
return 0
threshold = math.exp(-lam)
k = 0
product = 1.0
while product > threshold:
k += 1
product *= rng.random()
return k - 1
def football_prediction_to_markdown(prediction: Dict[str, Any]) -> str:
"""Render a prediction result as a report section."""
result = prediction["result"]
inputs = prediction["inputs"]
win_prob = result["win_prob"]
expected = result["expected_goals"]
matrix = result["score_distribution_matrix"]
def pct(value: float) -> str:
return f"{value * 100:.1f}%"
top_scores = result.get("top_scores", [])
top_text = "".join(f"{item['score']}{pct(item['prob'])}" for item in top_scores[:5])
lines = [
"本章节直接给出本次足球概率模拟的核心输出。系统已从模拟种子和初始动作中抽取到可用的泊松参数并按双变量泊松模型完成100,000次蒙特卡洛采样因此本次报告不再停留在方法论描述。",
"",
"**核心输入**",
"",
f"- 主队: {inputs['home_team']}",
f"- 客队: {inputs['away_team']}",
f"- lambda_home: {inputs['lambda_home']}",
f"- lambda_away: {inputs['lambda_away']}",
f"- 攻防相关性修正: {inputs['correlation']}",
f"- 采样次数: {inputs['samples']:,}",
"",
"**胜平负概率**",
"",
f"- {inputs['home_team']} 主胜: {pct(win_prob['home'])}",
f"- 平局: {pct(win_prob['draw'])}",
f"- {inputs['away_team']} 客胜: {pct(win_prob['away'])}",
"",
"**最可能比分**",
"",
f"最高概率比分集中在 {top_text}。从分布看,{inputs['away_team']} 的胜率显著高于 {inputs['home_team']},但平局和主队低比分抢分仍保留可观概率。",
"",
"**期望进球**",
"",
f"- {inputs['home_team']}: {expected['home']}",
f"- {inputs['away_team']}: {expected['away']}",
"",
"**比分概率分布矩阵**",
"",
"下表按“主队进球-客队进球”展示0到6球范围内的概率更高比分被计入溢出概率。",
"",
]
header = "| 主队\\客队 | " + " | ".join(str(goal) for goal in matrix["away_goals"]) + " |"
separator = "|---" * (len(matrix["away_goals"]) + 1) + "|"
lines.append(header)
lines.append(separator)
for home_goal, row in zip(matrix["home_goals"], matrix["probabilities"]):
cells = " | ".join(pct(value) for value in row)
lines.append(f"| {home_goal} | {cells} |")
lines.extend([
"",
f"矩阵外溢出概率: {pct(matrix['overflow_prob'])}",
"",
"该结果是概率分布,不是单一确定比分。若需要给出一个最可能比分,当前模拟的首选比分为 "
f"{top_scores[0]['score']},对应概率 {pct(top_scores[0]['prob'])}",
])
if prediction.get("warnings"):
lines.append("")
lines.append("**数据提示**")
lines.extend(f"- {warning}" for warning in prediction["warnings"])
return "\n".join(lines)

View File

@ -0,0 +1,176 @@
"""
图谱服务工厂
根据配置动态选择 Zep Neo4j 后端实现
"""
from typing import Optional, TYPE_CHECKING
from ..config import Config
from ..utils.logger import get_logger
if TYPE_CHECKING:
from .adapters.graph_adapter import GraphAdapter
from .adapters.neo4j_graph_builder import Neo4jGraphBuilder
logger = get_logger('mirofish.graph_factory')
class GraphServiceFactory:
"""
图谱服务工厂
根据 Config.GRAPH_BACKEND 配置返回对应的服务实现
- 'zep': 使用 Zep Cloud 服务
- 'neo4j': 使用 Neo4j 本地服务
"""
_instance: Optional['GraphServiceFactory'] = None
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialized = False
return cls._instance
def __init__(self):
if not self._initialized:
self._initialized = True
self._backend = Config.GRAPH_BACKEND or 'zep'
logger.info(f"GraphServiceFactory 初始化: backend={self._backend}")
@property
def backend(self) -> str:
return self._backend
def get_graph_builder(self):
"""
获取图谱构建服务
Returns:
GraphBuilderService Neo4jGraphBuilder
"""
if self._backend == 'neo4j':
from .adapters.neo4j_graph_builder import Neo4jGraphBuilder
logger.info("使用 Neo4jGraphBuilder")
return Neo4jGraphBuilder()
else:
from .graph_builder import GraphBuilderService
logger.info("使用 GraphBuilderService (Zep)")
return GraphBuilderService(api_key=Config.ZEP_API_KEY)
def get_entity_reader(self):
"""
获取实体读取服务
Returns:
ZepEntityReader Neo4jEntityReader
"""
if self._backend == 'neo4j':
from .adapters.neo4j_entity_reader import Neo4jEntityReader
logger.info("使用 Neo4jEntityReader")
return Neo4jEntityReader()
else:
from .zep_entity_reader import ZepEntityReader
logger.info("使用 ZepEntityReader")
return ZepEntityReader(api_key=Config.ZEP_API_KEY)
def get_search_service(self, llm_client=None):
"""
获取检索工具服务
Args:
llm_client: LLM 客户端用于需要 LLM 的功能
Returns:
ZepToolsService Neo4jSearchService
"""
if self._backend == 'neo4j':
from .adapters.neo4j_search_service import Neo4jSearchService
logger.info("使用 Neo4jSearchService")
return Neo4jSearchService(llm_client=llm_client)
else:
from .zep_tools import ZepToolsService
logger.info("使用 ZepToolsService")
return ZepToolsService(api_key=Config.ZEP_API_KEY, llm_client=llm_client)
def get_memory_updater(self, simulation_id: str, graph_id: str):
"""
获取图谱记忆更新器
Args:
simulation_id: 模拟ID
graph_id: 图谱ID
Returns:
ZepGraphMemoryUpdater Neo4jGraphMemoryUpdater
"""
if self._backend == 'neo4j':
from .adapters.neo4j_graph_memory_updater import Neo4jGraphMemoryManager
logger.info("使用 Neo4jGraphMemoryManager")
return Neo4jGraphMemoryManager.create_updater(simulation_id, graph_id)
else:
from .zep_graph_memory_updater import ZepGraphMemoryManager
logger.info("使用 ZepGraphMemoryManager")
return ZepGraphMemoryManager.create_updater(simulation_id, graph_id)
def get_existing_memory_updater(self, simulation_id: str):
"""获取已存在的图谱记忆更新器。"""
if self._backend == 'neo4j':
from .adapters.neo4j_graph_memory_updater import Neo4jGraphMemoryManager
return Neo4jGraphMemoryManager.get_updater(simulation_id)
else:
from .zep_graph_memory_updater import ZepGraphMemoryManager
return ZepGraphMemoryManager.get_updater(simulation_id)
def stop_memory_updater(self, simulation_id: str) -> None:
"""停止指定模拟的图谱记忆更新器。"""
if self._backend == 'neo4j':
from .adapters.neo4j_graph_memory_updater import Neo4jGraphMemoryManager
Neo4jGraphMemoryManager.stop_updater(simulation_id)
else:
from .zep_graph_memory_updater import ZepGraphMemoryManager
ZepGraphMemoryManager.stop_updater(simulation_id)
def stop_all_memory_updaters(self) -> None:
"""停止所有图谱记忆更新器。"""
if self._backend == 'neo4j':
from .adapters.neo4j_graph_memory_updater import Neo4jGraphMemoryManager
Neo4jGraphMemoryManager.stop_all()
else:
from .zep_graph_memory_updater import ZepGraphMemoryManager
ZepGraphMemoryManager.stop_all()
def check_backend_health(self) -> dict:
"""
检查后端健康状态
Returns:
健康状态字典
"""
if self._backend == 'neo4j':
from ..utils.neo4j import neo4j_health_check
healthy = neo4j_health_check()
return {
"backend": "neo4j",
"healthy": healthy,
"uri": Config.NEO4J_URI
}
else:
# Zep 没有内置健康检查,简单返回 True
return {
"backend": "zep",
"healthy": True,
"api_key_configured": bool(Config.ZEP_API_KEY)
}
# 全局工厂实例
_graph_factory: Optional[GraphServiceFactory] = None
def get_graph_factory() -> GraphServiceFactory:
"""获取图谱服务工厂实例"""
global _graph_factory
if _graph_factory is None:
_graph_factory = GraphServiceFactory()
return _graph_factory

View File

@ -16,12 +16,36 @@ from dataclasses import dataclass, field
from datetime import datetime
from openai import OpenAI
from zep_cloud.client import Zep
# Zep 是可选依赖,如果未安装则使用 Neo4j 适配器
try:
from zep_cloud.client import Zep as ZepClient
except ImportError:
ZepClient = None
from ..config import Config
from ..utils.logger import get_logger
from ..utils.locale import get_language_instruction, get_locale, set_locale, t
from .zep_entity_reader import EntityNode, ZepEntityReader
from ..utils.llm_rate_limit import call_llm_with_rate_limit_retry
# 实体节点类型 - 延迟导入以支持 Neo4j 后端
EntityNode = None
ZepEntityReader = None
def _ensure_entity_types():
"""延迟导入实体类型,确保在需要时可用"""
global EntityNode, ZepEntityReader
if EntityNode is None:
# 根据后端选择合适的实体读取器
if Config.GRAPH_BACKEND == 'neo4j':
from ..services.adapters.neo4j_entity_reader import EntityNode as Neo4jEntityNode
EntityNode = Neo4jEntityNode
else:
from .zep_entity_reader import EntityNode as ZepEntityNode
from .zep_entity_reader import ZepEntityReader as ZepEntityReaderCls
EntityNode = ZepEntityNode
ZepEntityReader = ZepEntityReaderCls
logger = get_logger('mirofish.oasis_profile')
@ -198,20 +222,20 @@ class OasisProfileGenerator:
base_url=self.base_url
)
# Zep客户端用于检索丰富上下文
# Zep客户端用于检索丰富上下文(可选)
self.zep_api_key = zep_api_key or Config.ZEP_API_KEY
self.zep_client = None
self.graph_id = graph_id
if self.zep_api_key:
if ZepClient and self.zep_api_key:
try:
self.zep_client = Zep(api_key=self.zep_api_key)
self.zep_client = ZepClient(api_key=self.zep_api_key)
except Exception as e:
logger.warning(f"Zep客户端初始化失败: {e}")
def generate_profile_from_entity(
self,
entity: EntityNode,
self,
entity, # EntityNode - 类型在运行时确定
user_id: int,
use_llm: bool = True
) -> OasisAgentProfile:
@ -283,7 +307,7 @@ class OasisProfileGenerator:
suffix = random.randint(100, 999)
return f"{username}_{suffix}"
def _search_zep_for_entity(self, entity: EntityNode) -> Dict[str, Any]:
def _search_zep_for_entity(self, entity) -> Dict[str, Any]: # EntityNode
"""
使用Zep图谱混合搜索功能获取实体相关的丰富信息
@ -411,7 +435,7 @@ class OasisProfileGenerator:
return results
def _build_entity_context(self, entity: EntityNode) -> str:
def _build_entity_context(self, entity) -> str: # EntityNode
"""
构建实体的完整上下文信息
@ -541,7 +565,10 @@ class OasisProfileGenerator:
if "MiniMax" not in self.model_name and "minimax" not in self.model_name:
kwargs["response_format"] = {"type": "json_object"}
response = self.client.chat.completions.create(**kwargs)
response = call_llm_with_rate_limit_retry(
lambda: self.client.chat.completions.create(**kwargs),
operation_name="OASIS profile LLM"
)
content = response.choices[0].message.content
@ -855,7 +882,7 @@ class OasisProfileGenerator:
def generate_profiles_from_entities(
self,
entities: List[EntityNode],
entities, # List[EntityNode]
use_llm: bool = True,
progress_callback: Optional[callable] = None,
graph_id: Optional[str] = None,

View File

@ -22,14 +22,28 @@ from ..config import Config
from ..utils.llm_client import LLMClient
from ..utils.logger import get_logger
from ..utils.locale import get_language_instruction, t
from .zep_tools import (
ZepToolsService,
SearchResult,
InsightForgeResult,
PanoramaResult,
InterviewResult
from .graph_service_factory import get_graph_factory
from .football_probability import (
FootballProbabilitySimulator,
football_prediction_to_markdown,
)
# 根据后端选择合适的工具类型
if Config.GRAPH_BACKEND == 'neo4j':
from .adapters.graph_adapter import SearchResult
from .adapters.neo4j_search_service import Neo4jSearchService
# Neo4j 返回的是 dict需要包装
InsightForgeResult = dict
PanoramaResult = dict
InterviewResult = dict
else:
from .zep_tools import (
SearchResult,
InsightForgeResult,
PanoramaResult,
InterviewResult
)
logger = get_logger('mirofish.report_agent')
@ -592,6 +606,9 @@ PLAN_USER_PROMPT_TEMPLATE = """\
预测场景设定
我们向模拟世界注入的变量模拟需求{simulation_requirement}
后端已计算的结构化预测结果
{computed_prediction_context}
模拟世界规模
- 参与模拟的实体数量: {total_nodes}
- 实体间产生的关系数量: {total_edges}
@ -608,6 +625,8 @@ PLAN_USER_PROMPT_TEMPLATE = """\
根据预测结果设计最合适的报告章节结构
如果上方存在足球比分概率预测结果必须把比分预测与概率分布作为第一章并在摘要中直接体现胜平负概率最可能比分和期望进球禁止写缺乏实际输入数据
再次提醒报告章节数量最少2个最多5个内容要精炼聚焦于核心预测发现"""
# ── 章节生成 prompt ──
@ -619,6 +638,9 @@ SECTION_SYSTEM_PROMPT_TEMPLATE = """\
报告摘要: {report_summary}
预测场景模拟需求: {simulation_requirement}
后端已计算的结构化预测结果:
{computed_prediction_context}
当前要撰写的章节: {section_title}
@ -887,7 +909,7 @@ class ReportAgent:
simulation_id: str,
simulation_requirement: str,
llm_client: Optional[LLMClient] = None,
zep_tools: Optional[ZepToolsService] = None
zep_tools: Optional[object] = None
):
"""
初始化Report Agent
@ -904,7 +926,11 @@ class ReportAgent:
self.simulation_requirement = simulation_requirement
self.llm = llm_client or LLMClient()
self.zep_tools = zep_tools or ZepToolsService()
self.zep_tools = zep_tools or get_graph_factory().get_search_service(llm_client=llm_client)
self.computed_prediction = FootballProbabilitySimulator.simulate_from_simulation(
simulation_id=self.simulation_id,
simulation_requirement=self.simulation_requirement
)
# 工具定义
self.tools = self._define_tools()
@ -915,6 +941,51 @@ class ReportAgent:
self.console_logger: Optional[ReportConsoleLogger] = None
logger.info(t('report.agentInitDone', graphId=graph_id, simulationId=simulation_id))
def _get_computed_prediction_context(self) -> str:
"""Return compact JSON context for deterministic prediction results."""
if not self.computed_prediction:
return "(无结构化预测结果)"
compact = {
"kind": self.computed_prediction.get("kind"),
"method": self.computed_prediction.get("method"),
"inputs": self.computed_prediction.get("inputs"),
"result": self.computed_prediction.get("result"),
"warnings": self.computed_prediction.get("warnings", []),
}
return json.dumps(compact, ensure_ascii=False, indent=2)
def _is_computed_prediction_section(self, section: ReportSection, section_index: int) -> bool:
if not self.computed_prediction:
return False
title = section.title or ""
if section_index == 1:
return True
return any(keyword in title for keyword in ("比分", "概率分布", "胜平负", "Score", "score"))
def _ensure_prediction_outline(self, outline: ReportOutline) -> ReportOutline:
"""Ensure football reports lead with the computed score distribution."""
if not self.computed_prediction:
return outline
prediction_title = "比分预测与概率分布"
sections = [s for s in outline.sections if s.title != prediction_title]
sections.insert(0, ReportSection(title=prediction_title))
outline.sections = sections[:5]
result = self.computed_prediction["result"]
inputs = self.computed_prediction["inputs"]
top_score = result["top_scores"][0]["score"] if result.get("top_scores") else "未知"
outline.summary = (
f"基于{inputs['samples']:,}次双变量泊松蒙特卡洛采样,"
f"{inputs['home_team']}主胜{result['win_prob']['home'] * 100:.1f}%"
f"平局{result['win_prob']['draw'] * 100:.1f}%"
f"{inputs['away_team']}客胜{result['win_prob']['away'] * 100:.1f}%"
f"最可能比分为{top_score}"
)
if "比分" not in outline.title and "足球" in self.simulation_requirement:
outline.title = "MiroFish足球比分概率模拟预测报告"
return outline
def _define_tools(self) -> Dict[str, Dict[str, Any]]:
"""定义可用工具"""
@ -952,6 +1023,119 @@ class ReportAgent:
}
}
}
def _tool_result_to_text(self, result: Any) -> str:
"""将不同图数据库后端的工具结果统一转换为文本。"""
if hasattr(result, "to_text"):
return result.to_text()
if isinstance(result, dict):
if "semantic_facts" in result or "entity_insights" in result:
return self._format_insight_result(result)
if "active_facts" in result or "all_nodes" in result or "historical_facts" in result:
return self._format_panorama_result(result)
if "responses" in result and "question" in result:
return self._format_interview_fallback(result)
return json.dumps(result, ensure_ascii=False, indent=2)
return str(result)
def _format_insight_result(self, result: Dict[str, Any]) -> str:
"""Format Neo4j insight dict in the same text shape expected by the frontend."""
text_parts = [
"## 未来预测深度分析",
f"分析问题: {result.get('query', '')}",
f"预测场景: {result.get('simulation_requirement', '')}",
"",
"### 预测数据统计",
f"- 相关预测事实: {result.get('total_facts', len(result.get('semantic_facts', [])))}",
f"- 涉及实体: {result.get('total_entities', len(result.get('entity_insights', [])))}",
f"- 关系链: {result.get('total_relationships', len(result.get('relationship_chains', [])))}",
]
sub_queries = result.get("sub_queries") or []
if sub_queries:
text_parts.append("\n### 分析的子问题")
for i, query in enumerate(sub_queries, 1):
text_parts.append(f"{i}. {query}")
semantic_facts = result.get("semantic_facts") or []
if semantic_facts:
text_parts.append("\n### 【关键事实】(请在报告中引用这些原文)")
for i, fact in enumerate(semantic_facts, 1):
text_parts.append(f'{i}. "{fact}"')
entities = result.get("entity_insights") or []
if entities:
text_parts.append("\n### 【核心实体】")
for entity in entities:
text_parts.append(f"- **{entity.get('name', '未知')}** ({entity.get('type', '实体')})")
if entity.get("summary"):
text_parts.append(f' 摘要: "{entity.get("summary")}"')
if entity.get("related_facts") is not None:
text_parts.append(f" 相关事实: {len(entity.get('related_facts') or [])}")
chains = result.get("relationship_chains") or []
if chains:
text_parts.append("\n### 【关系链】")
for chain in chains:
text_parts.append(f"- {chain}")
return "\n".join(text_parts)
def _format_panorama_result(self, result: Dict[str, Any]) -> str:
"""Format Neo4j panorama dict in the same text shape expected by the frontend."""
text_parts = [
"## 广度搜索结果(未来全景视图)",
f"查询: {result.get('query', '')}",
"",
"### 统计信息",
f"- 总节点数: {result.get('total_nodes', len(result.get('all_nodes', [])))}",
f"- 总边数: {result.get('total_edges', len(result.get('all_edges', [])))}",
f"- 当前有效事实: {result.get('active_count', len(result.get('active_facts', [])))}",
f"- 历史/过期事实: {result.get('historical_count', len(result.get('historical_facts', [])))}",
]
active_facts = result.get("active_facts") or []
if active_facts:
text_parts.append("\n### 【当前有效事实】(模拟结果原文)")
for i, fact in enumerate(active_facts, 1):
text_parts.append(f'{i}. "{fact}"')
historical_facts = result.get("historical_facts") or []
if historical_facts:
text_parts.append("\n### 【历史/过期事实】(演变过程记录)")
for i, fact in enumerate(historical_facts, 1):
text_parts.append(f'{i}. "{fact}"')
nodes = result.get("all_nodes") or []
if nodes:
text_parts.append("\n### 【涉及实体】")
for node in nodes:
labels = node.get("labels") or []
entity_type = next((l for l in labels if l not in ("Entity", "Node")), node.get("attributes", {}).get("entity_type", "实体"))
text_parts.append(f"- **{node.get('name', '未知')}** ({entity_type})")
return "\n".join(text_parts)
def _format_interview_fallback(self, result: Dict[str, Any]) -> str:
"""Readable fallback for Neo4j graph-context interview responses."""
lines = [
f"**采访主题:** {result.get('question', '')}",
f"**采访人数:** {result.get('count', 0)} / {result.get('count', 0)} 位模拟Agent",
"",
"### 采访实录",
]
for index, response in enumerate(result.get("responses", []), 1):
agent = response.get("agent", f"Agent {index}")
summary = response.get("response", {})
lines.extend([
"",
f"#### 采访 #{index}: {agent}",
f"**{agent}** (neo4j_graph_context)",
"",
"**A:**",
json.dumps(summary, ensure_ascii=False, indent=2),
])
return "\n".join(lines)
def _execute_tool(self, tool_name: str, parameters: Dict[str, Any], report_context: str = "") -> str:
"""
@ -977,7 +1161,7 @@ class ReportAgent:
simulation_requirement=self.simulation_requirement,
report_context=ctx
)
return result.to_text()
return self._tool_result_to_text(result)
elif tool_name == "panorama_search":
# 广度搜索 - 获取全貌
@ -990,7 +1174,7 @@ class ReportAgent:
query=query,
include_expired=include_expired
)
return result.to_text()
return self._tool_result_to_text(result)
elif tool_name == "quick_search":
# 简单搜索 - 快速检索
@ -1003,7 +1187,7 @@ class ReportAgent:
query=query,
limit=limit
)
return result.to_text()
return self._tool_result_to_text(result)
elif tool_name == "interview_agents":
# 深度采访 - 调用真实的OASIS采访API获取模拟Agent的回答双平台
@ -1013,12 +1197,13 @@ class ReportAgent:
max_agents = int(max_agents)
max_agents = min(max_agents, 10)
result = self.zep_tools.interview_agents(
graph_id=self.graph_id,
simulation_id=self.simulation_id,
interview_requirement=interview_topic,
simulation_requirement=self.simulation_requirement,
max_agents=max_agents
)
return result.to_text()
return self._tool_result_to_text(result)
# ========== 向后兼容的旧工具(内部重定向到新工具) ==========
@ -1166,6 +1351,7 @@ class ReportAgent:
system_prompt = f"{PLAN_SYSTEM_PROMPT}\n\n{get_language_instruction()}"
user_prompt = PLAN_USER_PROMPT_TEMPLATE.format(
simulation_requirement=self.simulation_requirement,
computed_prediction_context=self._get_computed_prediction_context(),
total_nodes=context.get('graph_statistics', {}).get('total_nodes', 0),
total_edges=context.get('graph_statistics', {}).get('total_edges', 0),
entity_types=list(context.get('graph_statistics', {}).get('entity_types', {}).keys()),
@ -1198,6 +1384,7 @@ class ReportAgent:
summary=response.get("summary", ""),
sections=sections
)
outline = self._ensure_prediction_outline(outline)
if progress_callback:
progress_callback("planning", 100, t('progress.outlinePlanComplete'))
@ -1208,7 +1395,7 @@ class ReportAgent:
except Exception as e:
logger.error(t('report.outlinePlanFailed', error=str(e)))
# 返回默认大纲3个章节作为fallback
return ReportOutline(
return self._ensure_prediction_outline(ReportOutline(
title="未来预测报告",
summary="基于模拟预测的未来趋势与风险分析",
sections=[
@ -1216,7 +1403,7 @@ class ReportAgent:
ReportSection(title="人群行为预测分析"),
ReportSection(title="趋势展望与风险提示")
]
)
))
def _generate_section_react(
self,
@ -1256,6 +1443,7 @@ class ReportAgent:
report_title=outline.title,
report_summary=outline.summary,
simulation_requirement=self.simulation_requirement,
computed_prediction_context=self._get_computed_prediction_context(),
section_title=section.title,
tools_description=self._get_tools_description(),
)
@ -1651,20 +1839,31 @@ class ReportAgent:
base_progress,
t('progress.generatingSection', title=section.title, current=section_num, total=total_sections)
)
# 生成主章节内容
section_content = self._generate_section_react(
section=section,
outline=outline,
previous_sections=generated_sections,
progress_callback=lambda stage, prog, msg:
progress_callback(
stage,
base_progress + int(prog * 0.7 / total_sections),
msg
) if progress_callback else None,
section_index=section_num
)
# 足球比分概率章节使用后端已计算结果避免核心预测被LLM漏写。
if self._is_computed_prediction_section(section, section_num):
section_content = football_prediction_to_markdown(self.computed_prediction)
if self.report_logger:
self.report_logger.log_section_content(
section_title=section.title,
section_index=section_num,
content=section_content,
tool_calls_count=0
)
else:
# 生成主章节内容
section_content = self._generate_section_react(
section=section,
outline=outline,
previous_sections=generated_sections,
progress_callback=lambda stage, prog, msg:
progress_callback(
stage,
base_progress + int(prog * 0.7 / total_sections),
msg
) if progress_callback else None,
section_index=section_num
)
section.content = section_content
generated_sections.append(f"## {section.title}\n\n{section_content}")

View File

@ -21,7 +21,13 @@ from openai import OpenAI
from ..config import Config
from ..utils.logger import get_logger
from ..utils.locale import get_language_instruction, t
from .zep_entity_reader import EntityNode, ZepEntityReader
from ..utils.llm_rate_limit import call_llm_with_rate_limit_retry
# 根据后端选择实体读取器
if Config.GRAPH_BACKEND == 'neo4j':
from .adapters.neo4j_entity_reader import EntityNode, FilteredEntities
else:
from .zep_entity_reader import EntityNode, FilteredEntities
logger = get_logger('mirofish.simulation_config')
@ -440,15 +446,18 @@ class SimulationConfigGenerator:
for attempt in range(max_attempts):
try:
response = self.client.chat.completions.create(
model=self.model_name,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt}
],
response_format={"type": "json_object"},
temperature=0.7 - (attempt * 0.1) # 每次重试降低温度
# 不设置max_tokens让LLM自由发挥
response = call_llm_with_rate_limit_retry(
lambda: self.client.chat.completions.create(
model=self.model_name,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt}
],
response_format={"type": "json_object"},
temperature=0.7 - (attempt * 0.1) # 每次重试降低温度
# 不设置max_tokens让LLM自由发挥
),
operation_name="Simulation config LLM"
)
content = response.choices[0].message.content

View File

@ -14,7 +14,9 @@ from enum import Enum
from ..config import Config
from ..utils.logger import get_logger
from .zep_entity_reader import ZepEntityReader, FilteredEntities
# FilteredEntities 是 Zep 实体读取器的类型,现在通过 graph_service_factory 延迟加载
# from .zep_entity_reader import FilteredEntities
from .graph_service_factory import get_graph_factory
from .oasis_profile_generator import OasisProfileGenerator, OasisAgentProfile
from .simulation_config_generator import SimulationConfigGenerator, SimulationParameters
from ..utils.locale import t
@ -273,7 +275,7 @@ class SimulationManager:
if progress_callback:
progress_callback("reading", 0, t('progress.connectingZepGraph'))
reader = ZepEntityReader()
reader = get_graph_factory().get_entity_reader()
if progress_callback:
progress_callback("reading", 30, t('progress.readingNodeData'))

View File

@ -21,7 +21,7 @@ from queue import Queue
from ..config import Config
from ..utils.logger import get_logger
from ..utils.locale import get_locale, set_locale
from .zep_graph_memory_updater import ZepGraphMemoryManager
from .graph_service_factory import get_graph_factory
from .simulation_ipc import SimulationIPCClient, CommandType, IPCResponse
logger = get_logger('mirofish.simulation_runner')
@ -375,7 +375,7 @@ class SimulationRunner:
raise ValueError("启用图谱记忆更新时必须提供 graph_id")
try:
ZepGraphMemoryManager.create_updater(simulation_id, graph_id)
get_graph_factory().get_memory_updater(simulation_id, graph_id)
cls._graph_memory_enabled[simulation_id] = True
logger.info(f"已启用图谱记忆更新: simulation_id={simulation_id}, graph_id={graph_id}")
except Exception as e:
@ -414,7 +414,7 @@ class SimulationRunner:
# simulation.log - 主进程日志
cmd = [
sys.executable, # Python解释器
"D:\\conda\\python.exe", # 使用包含 camel-ai 的 Python
script_path,
"--config", config_path, # 使用完整配置文件路径
]
@ -432,6 +432,13 @@ class SimulationRunner:
env = os.environ.copy()
env['PYTHONUTF8'] = '1' # Python 3.7+ 支持,让所有 open() 默认使用 UTF-8
env['PYTHONIOENCODING'] = 'utf-8' # 确保 stdout/stderr 使用 UTF-8
# 确保使用 conda Python优先搜索其路径
conda_python_dir = r"D:\conda"
if 'PATH' in env:
env['PATH'] = conda_python_dir + os.pathsep + env['PATH']
else:
env['PATH'] = conda_python_dir
# 设置工作目录为模拟目录(数据库等文件会生成在此)
# 使用 start_new_session=True 创建新的进程组,确保可以通过 os.killpg 终止所有子进程
@ -556,7 +563,7 @@ class SimulationRunner:
# 停止图谱记忆更新器
if cls._graph_memory_enabled.get(simulation_id, False):
try:
ZepGraphMemoryManager.stop_updater(simulation_id)
get_graph_factory().stop_memory_updater(simulation_id)
logger.info(f"已停止图谱记忆更新: simulation_id={simulation_id}")
except Exception as e:
logger.error(f"停止图谱记忆更新器失败: {e}")
@ -604,7 +611,7 @@ class SimulationRunner:
graph_memory_enabled = cls._graph_memory_enabled.get(state.simulation_id, False)
graph_updater = None
if graph_memory_enabled:
graph_updater = ZepGraphMemoryManager.get_updater(state.simulation_id)
graph_updater = get_graph_factory().get_existing_memory_updater(state.simulation_id)
try:
with open(log_path, 'r', encoding='utf-8') as f:
@ -812,7 +819,7 @@ class SimulationRunner:
# 停止图谱记忆更新器
if cls._graph_memory_enabled.get(simulation_id, False):
try:
ZepGraphMemoryManager.stop_updater(simulation_id)
get_graph_factory().stop_memory_updater(simulation_id)
logger.info(f"已停止图谱记忆更新: simulation_id={simulation_id}")
except Exception as e:
logger.error(f"停止图谱记忆更新器失败: {e}")
@ -1206,7 +1213,7 @@ class SimulationRunner:
# 首先停止所有图谱记忆更新器stop_all 内部会打印日志)
try:
ZepGraphMemoryManager.stop_all()
get_graph_factory().stop_all_memory_updaters()
except Exception as e:
logger.error(f"停止图谱记忆更新器失败: {e}")
cls._graph_memory_enabled.clear()

View File

@ -9,6 +9,7 @@ from typing import Optional, Dict, Any, List
from openai import OpenAI
from ..config import Config
from .llm_rate_limit import call_llm_with_rate_limit_retry
class LLMClient:
@ -62,7 +63,10 @@ class LLMClient:
if response_format and "MiniMax" not in self.model and "minimax" not in self.model:
kwargs["response_format"] = response_format
response = self.client.chat.completions.create(**kwargs)
response = call_llm_with_rate_limit_retry(
lambda: self.client.chat.completions.create(**kwargs),
operation_name="LLM chat"
)
content = response.choices[0].message.content
# 部分模型如MiniMax M2.5会在content中包含<think>思考内容,需要移除
content = re.sub(r'<think>[\s\S]*?</think>', '', content).strip()

View File

@ -0,0 +1,83 @@
"""
Shared helpers for LLM rate-limit handling.
"""
import random
import re
import time
from typing import Any, Callable, Optional
from openai import RateLimitError
from ..config import Config
from ..utils.logger import get_logger
logger = get_logger('mirofish.llm_rate_limit')
def is_rate_limit_error(error: Exception) -> bool:
"""Return True when an exception represents an LLM 429/rate-limit response."""
if isinstance(error, RateLimitError):
return True
message = str(error).lower()
return "429" in message or "rate_limit" in message or "速率限制" in message
def retry_after_seconds(error: Exception) -> Optional[float]:
"""Extract retry-after seconds from OpenAI-compatible errors when available."""
response = getattr(error, "response", None)
headers = getattr(response, "headers", None)
if headers:
retry_after = headers.get("retry-after") or headers.get("Retry-After")
if retry_after:
try:
return max(1.0, float(retry_after))
except ValueError:
pass
message = str(error)
match = re.search(r"retry[- ]after[:= ]+([0-9]+(?:\.[0-9]+)?)", message, re.IGNORECASE)
if match:
return max(1.0, float(match.group(1)))
return None
def call_llm_with_rate_limit_retry(
call: Callable[[], Any],
operation_name: str,
max_attempts: Optional[int] = None,
initial_delay: Optional[float] = None,
max_delay: Optional[float] = None,
) -> Any:
"""
Run an LLM call. On HTTP 429, sleep and retry instead of dropping the task.
"""
attempts = max_attempts or Config.LLM_RATE_LIMIT_MAX_ATTEMPTS
delay = initial_delay or Config.LLM_RATE_LIMIT_INITIAL_DELAY
max_sleep = max_delay or Config.LLM_RATE_LIMIT_MAX_DELAY
last_error: Optional[Exception] = None
for attempt in range(1, attempts + 1):
try:
return call()
except Exception as error:
last_error = error
if not is_rate_limit_error(error):
raise
if attempt >= attempts:
logger.error(
f"{operation_name} 遇到 429 限流,已等待重试 {attempts} 次仍失败: {error}"
)
raise
retry_after = retry_after_seconds(error)
sleep_seconds = retry_after if retry_after is not None else min(delay, max_sleep)
sleep_seconds = sleep_seconds * (0.9 + random.random() * 0.2)
logger.warning(
f"{operation_name} 遇到 429 限流,暂停 {sleep_seconds:.1f} 秒后继续 "
f"({attempt}/{attempts})"
)
time.sleep(sleep_seconds)
delay = min(delay * Config.LLM_RATE_LIMIT_BACKOFF_FACTOR, max_sleep)
raise last_error or RuntimeError(f"{operation_name} 调用失败")

View File

@ -0,0 +1,21 @@
"""
Neo4j Utilities Package
"""
from .driver import (
Neo4jDriverManager,
Neo4jConfig,
get_neo4j_driver,
close_neo4j_driver,
neo4j_health_check
)
from .schema import Neo4jSchemaManager
__all__ = [
'Neo4jDriverManager',
'Neo4jConfig',
'get_neo4j_driver',
'close_neo4j_driver',
'neo4j_health_check',
'Neo4jSchemaManager',
]

View File

@ -0,0 +1,166 @@
"""
Neo4j 数据库连接管理
提供 Neo4j 驱动的初始化连接池管理和健康检查
"""
import os
from typing import Optional
from dataclasses import dataclass
from neo4j import GraphDatabase
from neo4j import Driver
from neo4j.exceptions import ServiceUnavailable, AuthError
from ...config import Config
from ..logger import get_logger
logger = get_logger('mirofish.neo4j')
@dataclass
class Neo4jConfig:
"""Neo4j 配置"""
uri: str
username: str
password: str
database: str = "neo4j"
max_connection_pool_size: int = 50
connection_acquisition_timeout: int = 60
class Neo4jDriverManager:
"""
Neo4j 驱动管理器
管理 Neo4j 连接池支持单例模式
"""
_instance: Optional['Neo4jDriverManager'] = None
_driver: Optional[Driver] = None
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
if not hasattr(self, '_initialized'):
self._initialized = True
self._config: Optional[Neo4jConfig] = None
self._driver: Optional[Driver] = None
def _load_config(self) -> Neo4jConfig:
"""从环境变量加载配置"""
uri = Config.NEO4J_URI
username = Config.NEO4J_USERNAME
password = Config.NEO4J_PASSWORD
database = Config.NEO4J_DATABASE
return Neo4jConfig(
uri=uri,
username=username,
password=password,
database=database,
max_connection_pool_size=Config.NEO4J_MAX_POOL_SIZE,
connection_acquisition_timeout=int(os.environ.get('NEO4J_CONNECT_TIMEOUT', '60'))
)
def get_driver(self) -> Driver:
"""
获取 Neo4j 驱动单例
Returns:
Neo4j Driver 实例
Raises:
ValueError: 如果配置不完整
"""
if self._driver is None:
self._config = self._load_config()
if not self._config.password:
raise ValueError(
"Neo4j 密码未配置,请设置 NEO4J_PASSWORD 环境变量"
)
logger.info(f"创建 Neo4j 驱动: {self._config.uri}")
self._driver = GraphDatabase.driver(
self._config.uri,
auth=(self._config.username, self._config.password),
max_connection_pool_size=self._config.max_connection_pool_size,
connection_acquisition_timeout=self._config.connection_acquisition_timeout
)
return self._driver
def close(self):
"""关闭驱动"""
if self._driver is not None:
self._driver.close()
self._driver = None
logger.info("Neo4j 驱动已关闭")
def health_check(self) -> bool:
"""
健康检查
Returns:
True 如果连接正常
"""
try:
driver = self.get_driver()
with driver.session(database=self._config.database) as session:
result = session.run("RETURN 1 AS test")
result.single()
return True
except AuthError as e:
logger.error(f"Neo4j 认证失败: {e}")
return False
except ServiceUnavailable as e:
logger.error(f"Neo4j 服务不可用: {e}")
return False
except Exception as e:
logger.error(f"Neo4j 健康检查失败: {e}")
return False
def verify_connectivity(self) -> bool:
"""
验证连接带重试
Returns:
True 如果连接成功
"""
import time
max_retries = 3
for attempt in range(max_retries):
if self.health_check():
logger.info("Neo4j 连接验证成功")
return True
if attempt < max_retries - 1:
wait_time = 2 ** attempt # 指数退避
logger.warning(f"Neo4j 连接验证失败,{wait_time}秒后重试...")
time.sleep(wait_time)
return False
# 全局实例
_neo4j_manager = Neo4jDriverManager()
def get_neo4j_driver() -> Driver:
"""获取 Neo4j 驱动的便捷函数"""
return _neo4j_manager.get_driver()
def close_neo4j_driver():
"""关闭 Neo4j 驱动的便捷函数"""
_neo4j_manager.close()
def neo4j_health_check() -> bool:
"""Neo4j 健康检查的便捷函数"""
return _neo4j_manager.health_check()

View File

@ -0,0 +1,258 @@
"""
Neo4j Schema 管理器
处理 LabelRelationship TypeIndex 的创建和管理
"""
from typing import Dict, Any, List, Optional
from neo4j import Driver, Session
from .driver import get_neo4j_driver
from ..logger import get_logger
logger = get_logger('mirofish.neo4j.schema')
class Neo4jSchemaManager:
"""
Neo4j Schema 管理器
负责创建和管理
1. 节点 Label
2. 关系类型 (Relationship Type)
3. 索引 (Index)
4. 约束 (Constraint)
"""
# 保留的 Label 名称(不能用作自定义实体类型)
RESERVED_LABELS = {"Entity", "Node", "_GraphMetadata"}
def __init__(self, driver: Optional[Driver] = None):
self.driver = driver or get_neo4j_driver()
def setup_graph_schema(
self,
graph_id: str,
entity_types: List[Dict[str, Any]],
edge_types: List[Dict[str, Any]]
) -> None:
"""
设置图谱的 Schema本体
Args:
graph_id: 图谱ID
entity_types: 实体类型定义列表
edge_types: 关系类型定义列表
"""
with self.driver.session() as session:
# 1. 创建图谱根节点(用于隔离不同图谱的数据)
self._create_graph_root_node(session, graph_id)
# 2. 创建实体 Label以 Entity 开头)
for entity_def in entity_types:
label_name = entity_def["name"]
if label_name in self.RESERVED_LABELS:
logger.warning(f"实体类型 {label_name} 是保留名称,跳过")
continue
full_label = f"Entity_{label_name}"
self._create_entity_label(session, full_label, entity_def)
# 3. 创建关系类型
for edge_def in edge_types:
rel_type = edge_def["name"]
self._create_relationship_type(session, rel_type, edge_def)
# 4. 创建索引
self._create_indexes(session, graph_id)
logger.info(
f"Schema 设置完成: {len(entity_types)} 个实体类型, "
f"{len(edge_types)} 个关系类型"
)
def _create_graph_root_node(self, session: Session, graph_id: str) -> None:
"""创建图谱根节点,用于数据隔离"""
cypher = """
MERGE (g:_GraphMetadata {graph_id: $graph_id})
ON CREATE SET
g.created_at = datetime(),
g.entity_count = 0,
g.edge_count = 0
RETURN g
"""
session.run(cypher, graph_id=graph_id)
def _create_entity_label(
self,
session: Session,
label: str,
entity_def: Dict[str, Any]
) -> None:
"""
创建实体 Label
Args:
session: Neo4j 会话
label: Label 名称 (e.g., "Entity_Student")
entity_def: 实体定义
"""
description = entity_def.get("description", f"A {label} entity")
# 构建属性定义
properties = entity_def.get("attributes", [])
# 基本属性始终存在
base_props = {
"uuid": "STRING",
"name": "STRING",
"summary": "STRING",
"graph_id": "STRING",
"created_at": "STRING"
}
# 构建 CREATE LABEL 语句Neo4j 不支持程序化创建 Label
# 所以我们在节点创建时直接使用动态 Label
# 这里主要是记录日志
logger.debug(f"实体 Label: {label}, 描述: {description}")
def _create_relationship_type(
self,
session: Session,
rel_type: str,
edge_def: Dict[str, Any]
) -> None:
"""
创建关系类型
Args:
session: Neo4j 会话
rel_type: 关系类型名称 (e.g., "STUDIES_AT")
edge_def: 关系定义
"""
description = edge_def.get("description", f"A {rel_type} relationship")
logger.debug(f"关系类型: {rel_type}, 描述: {description}")
def _create_indexes(self, session: Session, graph_id: str) -> None:
"""
创建索引
为常用的查询字段创建索引以提高性能
"""
indexes = [
# 节点索引
("entity_uuid_index", "INDEX FOR (n:Entity) ON (n.uuid)"),
("entity_name_index", "INDEX FOR (n:Entity) ON (n.name)"),
("entity_graph_id_index", "INDEX FOR (n:Entity) ON (n.graph_id)"),
# 关系索引
("rel_source_index", "INDEX FOR ()-[r]-() ON (r.source_node_uuid)"),
("rel_target_index", "INDEX FOR ()-[r]-() ON (r.target_node_uuid)"),
]
for index_name, index_query in indexes:
try:
# 使用 IF NOT EXISTS 避免重复创建
session.run(f"CREATE INDEX {index_name} IF NOT EXISTS FOR {index_query.split(' FOR ')[1]}")
logger.debug(f"索引已创建: {index_name}")
except Exception as e:
# 忽略已存在的索引错误
logger.debug(f"索引 {index_name} 创建跳过: {e}")
def create_fulltext_index(self, index_name: str, node_labels: List[str], properties: List[str]) -> None:
"""
创建全文索引
Args:
index_name: 索引名称
node_labels: 节点 Label 列表
properties: 属性列表
"""
with self.driver.session() as session:
labels_str = ":".join(node_labels)
props_str = ", ".join([f"n.{p}" for p in properties])
cypher = f"""
CREATE FULLTEXT INDEX {index_name}
FOR (n:{labels_str}) ON EACH [{props_str}]
"""
try:
session.run(cypher)
logger.info(f"全文索引已创建: {index_name}")
except Exception as e:
logger.warning(f"全文索引创建失败: {index_name}, {e}")
def drop_graph_data(self, graph_id: str) -> None:
"""
删除图谱的所有数据保留 Schema
Args:
graph_id: 图谱ID
"""
with self.driver.session() as session:
# 删除该图谱的所有节点(级联删除关系)
cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
DETACH DELETE n
"""
result = session.run(cypher, graph_id=graph_id)
summary = result.consume()
nodes_deleted = summary.counters.nodes_deleted
# 更新图谱元数据
cypher2 = """
MATCH (g:_GraphMetadata {graph_id: $graph_id})
SET g.entity_count = 0, g.edge_count = 0, g.deleted_at = datetime()
"""
session.run(cypher2, graph_id=graph_id)
logger.info(f"图谱数据已删除: {graph_id}, 删除了 {nodes_deleted} 个节点")
def get_graph_stats(self, graph_id: str) -> Dict[str, Any]:
"""
获取图谱统计信息
Args:
graph_id: 图谱ID
Returns:
统计信息字典
"""
with self.driver.session() as session:
# 统计节点
node_cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
RETURN count(n) AS node_count
"""
node_record = session.run(node_cypher, graph_id=graph_id).single()
node_count = node_record["node_count"] if node_record else 0
# 统计边
edge_cypher = """
MATCH ()-[r]->()
WHERE r.graph_id = $graph_id
RETURN count(r) AS edge_count
"""
edge_record = session.run(edge_cypher, graph_id=graph_id).single()
edge_count = edge_record["edge_count"] if edge_record else 0
# 统计实体类型分布
type_cypher = """
MATCH (n:Entity)
WHERE n.graph_id = $graph_id
WITH labels(n) AS lbs, count(*) AS cnt
UNWIND lbs AS label
WITH label, cnt WHERE NOT label IN ['Entity', 'Node']
RETURN label AS entity_type, sum(cnt) AS count
"""
type_result = session.run(type_cypher, graph_id=graph_id)
entity_types = {record["entity_type"]: record["count"] for record in type_result}
return {
"graph_id": graph_id,
"node_count": node_count,
"edge_count": edge_count,
"entity_types": entity_types
}

View File

@ -16,7 +16,10 @@ dependencies = [
# LLM 相关
"openai>=1.0.0",
# Zep Cloud
# 图数据库
"neo4j==5.23.0",
# Zep Cloud可选旧后端
"zep-cloud==3.13.0",
# OASIS 社交媒体模拟

View File

@ -13,7 +13,10 @@ flask-cors>=6.0.0
# OpenAI SDK统一使用 OpenAI 格式调用 LLM
openai>=1.0.0
# ============= Zep Cloud =============
# ============= 图数据库 =============
neo4j==5.23.0
# ============= Zep Cloud可选旧后端 =============
zep-cloud==3.13.0
# ============= OASIS 社交媒体模拟 =============

33
docker-compose.neo4j.yml Normal file
View File

@ -0,0 +1,33 @@
# Neo4j 本地开发数据库
# 使用方式: docker compose -f docker-compose.neo4j.yml up -d
services:
neo4j:
image: neo4j:5-community
container_name: mirofish-neo4j
ports:
- "7474:7474" # Neo4j Browser (HTTP)
- "7687:7687" # Bolt 协议 (用于驱动连接)
environment:
# Neo4j 认证 (neo4j/密码)
- NEO4J_AUTH=neo4j/password
# 内存配置
- NEO4J_server.memory.heap.initial_size=512m
- NEO4J_server.memory.heap.max_size=2G
# 允许远程连接
- NEO4J_server.default_listen_address=0.0.0.0
volumes:
# 持久化数据
- neo4j_data:/data
- neo4j_logs:/logs
restart: unless-stopped
healthcheck:
test: ["CMD-SHELL", "wget -q --spider localhost:7474 || exit 1"]
interval: 30s
timeout: 10s
retries: 5
start_period: 30s
volumes:
neo4j_data:
neo4j_logs:

View File

@ -1,8 +1,8 @@
services:
mirofish:
image: ghcr.io/666ghj/mirofish:latest
# image: ghcr.io/666ghj/mirofish:latest
# 加速镜像(如拉取缓慢可替换上方地址)
# image: ghcr.nju.edu.cn/666ghj/mirofish:latest
image: ghcr.nju.edu.cn/666ghj/mirofish:latest
container_name: mirofish
env_file:
- .env

View File

@ -1435,7 +1435,6 @@
"resolved": "https://registry.npmjs.org/d3-selection/-/d3-selection-3.0.0.tgz",
"integrity": "sha512-fmTRWbNMmsmWq6xJV8D19U/gw/bwrHfNXxrIN+HfZgnzqTHp9jOmKMhsTUjXOJnZOdZY9Q28y4yebKzqDKlxlQ==",
"license": "ISC",
"peer": true,
"engines": {
"node": ">=12"
}
@ -1913,7 +1912,6 @@
"integrity": "sha512-QP88BAKvMam/3NxH6vj2o21R6MjxZUAd6nlwAS/pnGvN9IVLocLHxGYIzFhg6fUQ+5th6P4dv4eW9jX3DSIj7A==",
"dev": true,
"license": "MIT",
"peer": true,
"engines": {
"node": ">=12"
},
@ -2053,7 +2051,6 @@
"integrity": "sha512-ITcnkFeR3+fI8P1wMgItjGrR10170d8auB4EpMLPqmx6uxElH3a/hHGQabSHKdqd4FXWO1nFIp9rRn7JQ34ACQ==",
"dev": true,
"license": "MIT",
"peer": true,
"dependencies": {
"esbuild": "^0.25.0",
"fdir": "^6.5.0",
@ -2128,7 +2125,6 @@
"resolved": "https://registry.npmjs.org/vue/-/vue-3.5.25.tgz",
"integrity": "sha512-YLVdgv2K13WJ6n+kD5owehKtEXwdwXuj2TTyJMsO7pSeKw2bfRNZGjhB7YzrpbMYj5b5QsUebHpOqR3R3ziy/g==",
"license": "MIT",
"peer": true,
"dependencies": {
"@vue/compiler-dom": "3.5.25",
"@vue/compiler-sfc": "3.5.25",

View File

@ -393,10 +393,10 @@
"progress": {
"initGraphService": "Initializing graph build service...",
"textChunking": "Chunking text...",
"creatingZepGraph": "Creating Zep graph...",
"creatingZepGraph": "Creating graph...",
"settingOntology": "Setting ontology definition...",
"addingChunks": "Adding {count} text chunks...",
"waitingZepProcess": "Waiting for Zep to process data...",
"waitingZepProcess": "Waiting for data processing...",
"fetchingGraphData": "Fetching graph data...",
"graphBuildComplete": "Graph build complete",
"buildFailed": "Build failed: {error}",
@ -410,12 +410,12 @@
"noEpisodesWait": "No episodes to wait for",
"waitingEpisodes": "Waiting for {count} text chunks to process...",
"episodesTimeout": "Some chunks timed out, {completed}/{total} completed",
"zepProcessing": "Zep processing... {completed}/{total} done, {pending} pending ({elapsed}s)",
"zepProcessing": "Processing... {completed}/{total} done, {pending} pending ({elapsed}s)",
"processingComplete": "Processing complete: {completed}/{total}",
"taskComplete": "Task complete",
"taskFailed": "Task failed",
"startPreparingEnv": "Preparing simulation environment...",
"connectingZepGraph": "Connecting to Zep graph...",
"connectingZepGraph": "Connecting to graph...",
"readingNodeData": "Reading node data...",
"readingComplete": "Done, {count} entities found",
"startGenerating": "Starting generation...",

View File

@ -393,10 +393,10 @@
"progress": {
"initGraphService": "初始化图谱构建服务...",
"textChunking": "文本分块中...",
"creatingZepGraph": "创建Zep图谱...",
"creatingZepGraph": "创建图谱...",
"settingOntology": "设置本体定义...",
"addingChunks": "开始添加 {count} 个文本块...",
"waitingZepProcess": "等待Zep处理数据...",
"waitingZepProcess": "等待数据处理...",
"fetchingGraphData": "获取图谱数据...",
"graphBuildComplete": "图谱构建完成",
"buildFailed": "构建失败: {error}",
@ -410,12 +410,12 @@
"noEpisodesWait": "无需等待(没有 episode",
"waitingEpisodes": "开始等待 {count} 个文本块处理...",
"episodesTimeout": "部分文本块超时,已完成 {completed}/{total}",
"zepProcessing": "Zep处理中... {completed}/{total} 完成, {pending} 待处理 ({elapsed}秒)",
"zepProcessing": "数据处理中... {completed}/{total} 完成, {pending} 待处理 ({elapsed}秒)",
"processingComplete": "处理完成: {completed}/{total}",
"taskComplete": "任务完成",
"taskFailed": "任务失败",
"startPreparingEnv": "开始准备模拟环境...",
"connectingZepGraph": "正在连接Zep图谱...",
"connectingZepGraph": "正在连接图谱...",
"readingNodeData": "正在读取节点数据...",
"readingComplete": "完成,共 {count} 个实体",
"startGenerating": "开始生成...",