From 01fd47a3638dd775f3d610a90fac64223790a88e Mon Sep 17 00:00:00 2001 From: Hao OUYANG Date: Tue, 7 Jul 2026 22:46:34 +0800 Subject: [PATCH] fix: harden runtime API and simulation handling --- backend/app/api/graph.py | 2 +- backend/app/api/report.py | 10 ++-- backend/app/api/simulation.py | 30 +++++------ .../app/services/oasis_profile_generator.py | 8 +-- .../services/simulation_config_generator.py | 3 +- backend/app/services/simulation_ipc.py | 2 +- backend/app/services/simulation_manager.py | 2 +- backend/app/services/simulation_runner.py | 54 +++++++++++++------ backend/app/services/zep_tools.py | 3 +- backend/app/utils/llm_client.py | 3 +- backend/pyproject.toml | 1 + backend/requirements.txt | 1 + backend/scripts/run_parallel_simulation.py | 46 ++++++++++++++-- backend/scripts/run_reddit_simulation.py | 6 ++- backend/scripts/run_twitter_simulation.py | 14 ++++- backend/uv.lock | 2 + 16 files changed, 135 insertions(+), 52 deletions(-) diff --git a/backend/app/api/graph.py b/backend/app/api/graph.py index 759ff48b..05aa158d 100644 --- a/backend/app/api/graph.py +++ b/backend/app/api/graph.py @@ -295,7 +295,7 @@ def build_graph(): }), 500 # 解析请求 - data = request.get_json() or {} + data = request.get_json(silent=True) or {} project_id = data.get('project_id') logger.debug(f"请求参数: project_id={project_id}") diff --git a/backend/app/api/report.py b/backend/app/api/report.py index d7f2a4d0..54bd8ee1 100644 --- a/backend/app/api/report.py +++ b/backend/app/api/report.py @@ -48,7 +48,7 @@ def generate_report(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') if not simulation_id: @@ -223,7 +223,7 @@ def get_generate_status(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} task_id = data.get('task_id') simulation_id = data.get('simulation_id') @@ -497,7 +497,7 @@ def chat_with_report_agent(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') message = data.get('message') @@ -945,7 +945,7 @@ def search_graph_tool(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} graph_id = data.get('graph_id') query = data.get('query') @@ -991,7 +991,7 @@ def get_graph_statistics_tool(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} graph_id = data.get('graph_id') diff --git a/backend/app/api/simulation.py b/backend/app/api/simulation.py index 3a8e1e3f..126732c5 100644 --- a/backend/app/api/simulation.py +++ b/backend/app/api/simulation.py @@ -192,7 +192,7 @@ def create_simulation(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} project_id = data.get('project_id') if not project_id: @@ -291,7 +291,7 @@ def _check_simulation_prepared(simulation_id: str) -> tuple: state_file = os.path.join(simulation_dir, "state.json") try: import json - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: state_data = json.load(f) status = state_data.get("status", "") @@ -403,7 +403,7 @@ def prepare_simulation(): from ..config import Config try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') if not simulation_id: @@ -670,7 +670,7 @@ def get_prepare_status(): from ..models.task import TaskManager try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} task_id = data.get('task_id') simulation_id = data.get('simulation_id') @@ -1104,7 +1104,7 @@ def get_simulation_profiles_realtime(simulation_id: str): state_file = os.path.join(sim_dir, "state.json") if os.path.exists(state_file): try: - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: state_data = json.load(f) status = state_data.get("status", "") is_generating = status == "preparing" @@ -1200,7 +1200,7 @@ def get_simulation_config_realtime(simulation_id: str): state_file = os.path.join(sim_dir, "state.json") if os.path.exists(state_file): try: - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: state_data = json.load(f) status = state_data.get("status", "") is_generating = status == "preparing" @@ -1388,7 +1388,7 @@ def generate_profiles(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} graph_id = data.get('graph_id') if not graph_id: @@ -1490,7 +1490,7 @@ def start_simulation(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') if not simulation_id: @@ -1662,7 +1662,7 @@ def stop_simulation(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') if not simulation_id: @@ -2191,7 +2191,7 @@ def interview_agent(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') agent_id = data.get('agent_id') @@ -2313,7 +2313,7 @@ def interview_agents_batch(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') interviews = data.get('interviews') @@ -2440,7 +2440,7 @@ def interview_all_agents(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') prompt = data.get('prompt') @@ -2544,7 +2544,7 @@ def get_interview_history(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') platform = data.get('platform') # 不指定则返回两个平台的历史 @@ -2606,7 +2606,7 @@ def get_env_status(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') @@ -2673,7 +2673,7 @@ def close_simulation_env(): } """ try: - data = request.get_json() or {} + data = request.get_json(silent=True) or {} simulation_id = data.get('simulation_id') timeout = data.get('timeout', 30) diff --git a/backend/app/services/oasis_profile_generator.py b/backend/app/services/oasis_profile_generator.py index 7704a627..d5b13f58 100644 --- a/backend/app/services/oasis_profile_generator.py +++ b/backend/app/services/oasis_profile_generator.py @@ -195,7 +195,8 @@ class OasisProfileGenerator: self.client = OpenAI( api_key=self.api_key, - base_url=self.base_url + base_url=self.base_url, + default_headers={"User-Agent": "python-requests/2.32.5"} ) # Zep客户端用于检索丰富上下文 @@ -315,6 +316,7 @@ class OasisProfileGenerator: return results comprehensive_query = t('progress.zepSearchQuery', name=entity_name) + zep_query = comprehensive_query[:400] def search_edges(): """搜索边(事实/关系)- 带重试机制""" @@ -325,7 +327,7 @@ class OasisProfileGenerator: for attempt in range(max_retries): try: return self.zep_client.graph.search( - query=comprehensive_query, + query=zep_query, graph_id=self.graph_id, limit=30, scope="edges", @@ -350,7 +352,7 @@ class OasisProfileGenerator: for attempt in range(max_retries): try: return self.zep_client.graph.search( - query=comprehensive_query, + query=zep_query, graph_id=self.graph_id, limit=20, scope="nodes", diff --git a/backend/app/services/simulation_config_generator.py b/backend/app/services/simulation_config_generator.py index cb77f6b6..22669e50 100644 --- a/backend/app/services/simulation_config_generator.py +++ b/backend/app/services/simulation_config_generator.py @@ -237,7 +237,8 @@ class SimulationConfigGenerator: self.client = OpenAI( api_key=self.api_key, - base_url=self.base_url + base_url=self.base_url, + default_headers={"User-Agent": "python-requests/2.32.5"} ) def generate_config( diff --git a/backend/app/services/simulation_ipc.py b/backend/app/services/simulation_ipc.py index 9d70d0be..67b8f91d 100644 --- a/backend/app/services/simulation_ipc.py +++ b/backend/app/services/simulation_ipc.py @@ -278,7 +278,7 @@ class SimulationIPCClient: return False try: - with open(status_file, 'r', encoding='utf-8') as f: + with open(status_file, 'r', encoding='utf-8-sig') as f: status = json.load(f) return status.get("status") == "alive" except (json.JSONDecodeError, OSError): diff --git a/backend/app/services/simulation_manager.py b/backend/app/services/simulation_manager.py index 0d161a90..ae69ee57 100644 --- a/backend/app/services/simulation_manager.py +++ b/backend/app/services/simulation_manager.py @@ -165,7 +165,7 @@ class SimulationManager: if not os.path.exists(state_file): return None - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: data = json.load(f) state = SimulationState( diff --git a/backend/app/services/simulation_runner.py b/backend/app/services/simulation_runner.py index e86021f8..ae864db5 100644 --- a/backend/app/services/simulation_runner.py +++ b/backend/app/services/simulation_runner.py @@ -12,6 +12,7 @@ import threading import subprocess import signal import atexit +import psutil from typing import Dict, Any, List, Optional, Union from dataclasses import dataclass, field from datetime import datetime @@ -247,7 +248,7 @@ class SimulationRunner: return None try: - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: data = json.load(f) state = SimulationRunState( @@ -1232,12 +1233,18 @@ class SimulationRunner: # 更新 run_state.json state = cls.get_run_state(simulation_id) + keep_completed = bool( + state and ( + state.runner_status == RunnerStatus.COMPLETED or + (state.progress_percent >= 100 and state.twitter_completed and state.reddit_completed) + ) + ) if state: - state.runner_status = RunnerStatus.STOPPED + state.runner_status = RunnerStatus.COMPLETED if keep_completed else RunnerStatus.STOPPED state.twitter_running = False state.reddit_running = False - state.completed_at = datetime.now().isoformat() - state.error = "服务器关闭,模拟被终止" + state.completed_at = state.completed_at or datetime.now().isoformat() + state.error = None if keep_completed else "服务器关闭,模拟被终止" cls._save_run_state(state) # 同时更新 state.json,将状态设为 stopped @@ -1246,13 +1253,13 @@ class SimulationRunner: state_file = os.path.join(sim_dir, "state.json") logger.info(f"尝试更新 state.json: {state_file}") if os.path.exists(state_file): - with open(state_file, 'r', encoding='utf-8') as f: + with open(state_file, 'r', encoding='utf-8-sig') as f: state_data = json.load(f) - state_data['status'] = 'stopped' + state_data['status'] = 'completed' if keep_completed else 'stopped' state_data['updated_at'] = datetime.now().isoformat() with open(state_file, 'w', encoding='utf-8') as f: json.dump(state_data, f, indent=2, ensure_ascii=False) - logger.info(f"已更新 state.json 状态为 stopped: {simulation_id}") + logger.info(f"已更新 state.json 状态为 {state_data['status']}: {simulation_id}") else: logger.warning(f"state.json 不存在: {state_file}") except Exception as state_err: @@ -1386,7 +1393,23 @@ class SimulationRunner: return False ipc_client = SimulationIPCClient(sim_dir) - return ipc_client.check_env_alive() + if not ipc_client.check_env_alive(): + return False + + process = cls._processes.get(simulation_id) + if process is not None: + return process.poll() is None + + state = cls.get_run_state(simulation_id) + pid = state.process_pid if state else None + if not pid: + return False + + try: + proc = psutil.Process(pid) + return proc.is_running() and proc.status() != psutil.STATUS_ZOMBIE + except psutil.Error: + return False @classmethod def get_env_status_detail(cls, simulation_id: str) -> Dict[str, Any]: @@ -1413,12 +1436,13 @@ class SimulationRunner: return default_status try: - with open(status_file, 'r', encoding='utf-8') as f: + with open(status_file, 'r', encoding='utf-8-sig') as f: status = json.load(f) + env_alive = cls.check_env_alive(simulation_id) return { - "status": status.get("status", "stopped"), - "twitter_available": status.get("twitter_available", False), - "reddit_available": status.get("reddit_available", False), + "status": status.get("status", "stopped") if env_alive else "stopped", + "twitter_available": status.get("twitter_available", False) if env_alive else False, + "reddit_available": status.get("reddit_available", False) if env_alive else False, "timestamp": status.get("timestamp") } except (json.JSONDecodeError, OSError): @@ -1459,7 +1483,7 @@ class SimulationRunner: ipc_client = SimulationIPCClient(sim_dir) - if not ipc_client.check_env_alive(): + if not cls.check_env_alive(simulation_id): raise ValueError(f"模拟环境未运行或已关闭,无法执行Interview: {simulation_id}") logger.info(f"发送Interview命令: simulation_id={simulation_id}, agent_id={agent_id}, platform={platform}") @@ -1521,7 +1545,7 @@ class SimulationRunner: ipc_client = SimulationIPCClient(sim_dir) - if not ipc_client.check_env_alive(): + if not cls.check_env_alive(simulation_id): raise ValueError(f"模拟环境未运行或已关闭,无法执行Interview: {simulation_id}") logger.info(f"发送批量Interview命令: simulation_id={simulation_id}, count={len(interviews)}, platform={platform}") @@ -1631,7 +1655,7 @@ class SimulationRunner: ipc_client = SimulationIPCClient(sim_dir) - if not ipc_client.check_env_alive(): + if not cls.check_env_alive(simulation_id): return { "success": True, "message": "环境已经关闭" diff --git a/backend/app/services/zep_tools.py b/backend/app/services/zep_tools.py index 3bc8a57a..ae883698 100644 --- a/backend/app/services/zep_tools.py +++ b/backend/app/services/zep_tools.py @@ -484,13 +484,14 @@ class ZepToolsService: SearchResult: 搜索结果 """ logger.info(t("console.graphSearch", graphId=graph_id, query=query[:50])) + zep_query = query[:400] # 尝试使用Zep Cloud Search API try: search_results = self._call_with_retry( func=lambda: self.client.graph.search( graph_id=graph_id, - query=query, + query=zep_query, limit=limit, scope=scope, reranker="cross_encoder" diff --git a/backend/app/utils/llm_client.py b/backend/app/utils/llm_client.py index 6c1a81f4..ddd69e0f 100644 --- a/backend/app/utils/llm_client.py +++ b/backend/app/utils/llm_client.py @@ -29,7 +29,8 @@ class LLMClient: self.client = OpenAI( api_key=self.api_key, - base_url=self.base_url + base_url=self.base_url, + default_headers={"User-Agent": "python-requests/2.32.5"} ) def chat( diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 8c65b729..d0f04cfc 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -32,6 +32,7 @@ dependencies = [ # 工具库 "python-dotenv>=1.0.0", "pydantic>=2.0.0", + "psutil>=5.9.0", ] [project.optional-dependencies] diff --git a/backend/requirements.txt b/backend/requirements.txt index 4f146296..c6a2fe99 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -33,3 +33,4 @@ python-dotenv>=1.0.0 # 数据验证 pydantic>=2.0.0 +psutil>=5.9.0 diff --git a/backend/scripts/run_parallel_simulation.py b/backend/scripts/run_parallel_simulation.py index 2a627ffd..9ec72f72 100644 --- a/backend/scripts/run_parallel_simulation.py +++ b/backend/scripts/run_parallel_simulation.py @@ -678,8 +678,9 @@ def fetch_new_actions_from_db( if not os.path.exists(db_path): return actions, new_last_rowid + conn = None try: - conn = sqlite3.connect(db_path) + conn = sqlite3.connect(db_path, timeout=5) cursor = conn.cursor() # 使用 rowid 来追踪已处理的记录(rowid 是 SQLite 的内置自增字段) @@ -739,13 +740,35 @@ def fetch_new_actions_from_db( 'action_args': simplified_args, }) - conn.close() except Exception as e: print(f"读取数据库动作失败: {e}") + finally: + if conn is not None: + conn.close() return actions, new_last_rowid +def get_latest_trace_rowid(db_path: str) -> int: + """Return the latest processed trace row id for a platform DB.""" + if not os.path.exists(db_path): + return 0 + + conn = None + try: + conn = sqlite3.connect(db_path, timeout=5) + cursor = conn.cursor() + cursor.execute("SELECT COALESCE(MAX(rowid), 0) FROM trace") + rowid = int(cursor.fetchone()[0] or 0) + return rowid + except Exception as e: + print(f"读取最新动作游标失败: {e}") + return 0 + finally: + if conn is not None: + conn.close() + + def _enrich_action_context( cursor, action_type: str, @@ -1034,6 +1057,9 @@ def create_model(config: Dict[str, Any], use_boost: bool = False): return ModelFactory.create( model_platform=ModelPlatformType.OPENAI, model_type=llm_model, + api_key=llm_api_key or None, + url=llm_base_url or None, + default_headers={"User-Agent": "python-requests/2.32.5"}, ) @@ -1184,10 +1210,16 @@ async def run_twitter_simulation( content = post.get("content", "") try: agent = result.env.agent_graph.get_agent(agent_id) - initial_actions[agent] = ManualAction( + manual_action = ManualAction( action_type=ActionType.CREATE_POST, action_args={"content": content} ) + if agent in initial_actions: + if not isinstance(initial_actions[agent], list): + initial_actions[agent] = [initial_actions[agent]] + initial_actions[agent].append(manual_action) + else: + initial_actions[agent] = manual_action if action_logger: action_logger.log_action( @@ -1204,7 +1236,9 @@ async def run_twitter_simulation( if initial_actions: await result.env.step(initial_actions) - log_info(f"已发布 {len(initial_actions)} 条初始帖子") + last_rowid = get_latest_trace_rowid(db_path) + posted_count = sum(len(action) if isinstance(action, list) else 1 for action in initial_actions.values()) + log_info(f"已发布 {posted_count} 条初始帖子") # 记录 round 0 结束 if action_logger: @@ -1403,7 +1437,9 @@ async def run_reddit_simulation( if initial_actions: await result.env.step(initial_actions) - log_info(f"已发布 {len(initial_actions)} 条初始帖子") + last_rowid = get_latest_trace_rowid(db_path) + posted_count = sum(len(action) if isinstance(action, list) else 1 for action in initial_actions.values()) + log_info(f"已发布 {posted_count} 条初始帖子") # 记录 round 0 结束 if action_logger: diff --git a/backend/scripts/run_reddit_simulation.py b/backend/scripts/run_reddit_simulation.py index 14907cbd..4df88b02 100644 --- a/backend/scripts/run_reddit_simulation.py +++ b/backend/scripts/run_reddit_simulation.py @@ -464,6 +464,9 @@ class RedditSimulationRunner: return ModelFactory.create( model_platform=ModelPlatformType.OPENAI, model_type=llm_model, + api_key=llm_api_key or None, + url=llm_base_url or None, + default_headers={"User-Agent": "python-requests/2.32.5"}, ) def _get_active_agents_for_round( @@ -617,7 +620,8 @@ class RedditSimulationRunner: if initial_actions: await self.env.step(initial_actions) - print(f" 已发布 {len(initial_actions)} 条初始帖子") + posted_count = sum(len(action) if isinstance(action, list) else 1 for action in initial_actions.values()) + print(f" 已发布 {posted_count} 条初始帖子") # 主模拟循环 print("\n开始模拟循环...") diff --git a/backend/scripts/run_twitter_simulation.py b/backend/scripts/run_twitter_simulation.py index caab9e9d..7cb55a87 100644 --- a/backend/scripts/run_twitter_simulation.py +++ b/backend/scripts/run_twitter_simulation.py @@ -457,6 +457,9 @@ class TwitterSimulationRunner: return ModelFactory.create( model_platform=ModelPlatformType.OPENAI, model_type=llm_model, + api_key=llm_api_key or None, + url=llm_base_url or None, + default_headers={"User-Agent": "python-requests/2.32.5"}, ) def _get_active_agents_for_round( @@ -615,16 +618,23 @@ class TwitterSimulationRunner: content = post.get("content", "") try: agent = self.env.agent_graph.get_agent(agent_id) - initial_actions[agent] = ManualAction( + manual_action = ManualAction( action_type=ActionType.CREATE_POST, action_args={"content": content} ) + if agent in initial_actions: + if not isinstance(initial_actions[agent], list): + initial_actions[agent] = [initial_actions[agent]] + initial_actions[agent].append(manual_action) + else: + initial_actions[agent] = manual_action except Exception as e: print(f" 警告: 无法为Agent {agent_id}创建初始帖子: {e}") if initial_actions: await self.env.step(initial_actions) - print(f" 已发布 {len(initial_actions)} 条初始帖子") + posted_count = sum(len(action) if isinstance(action, list) else 1 for action in initial_actions.values()) + print(f" 已发布 {posted_count} 条初始帖子") # 主模拟循环 print("\n开始模拟循环...") diff --git a/backend/uv.lock b/backend/uv.lock index 642dd9c3..20f0f8e6 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -995,6 +995,7 @@ dependencies = [ { name = "flask" }, { name = "flask-cors" }, { name = "openai" }, + { name = "psutil" }, { name = "pydantic" }, { name = "pymupdf" }, { name = "python-dotenv" }, @@ -1024,6 +1025,7 @@ requires-dist = [ { name = "flask-cors", specifier = ">=6.0.0" }, { name = "openai", specifier = ">=1.0.0" }, { name = "pipreqs", marker = "extra == 'dev'", specifier = ">=0.5.0" }, + { name = "psutil", specifier = ">=5.9.0" }, { name = "pydantic", specifier = ">=2.0.0" }, { name = "pymupdf", specifier = ">=1.24.0" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0.0" },