Himankpro commited on
Commit
a54db35
·
verified ·
1 Parent(s): 8e2665a

Upload 93 files

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +2 -0
  2. app/.dockerignore +12 -0
  3. app/.python-version +1 -0
  4. app/.python-version.txt +1 -0
  5. app/Dockerfile +20 -0
  6. app/HUGGINGFACE_SPACE.md +50 -0
  7. app/Procfile +1 -0
  8. app/app/__init__.py +108 -0
  9. app/app/__pycache__/__init__.cpython-311.pyc +0 -0
  10. app/app/__pycache__/config.cpython-311.pyc +0 -0
  11. app/app/api/__init__.py +18 -0
  12. app/app/api/__pycache__/__init__.cpython-311.pyc +0 -0
  13. app/app/api/__pycache__/billing.cpython-311.pyc +0 -0
  14. app/app/api/__pycache__/graph.cpython-311.pyc +0 -0
  15. app/app/api/__pycache__/pipeline.cpython-311.pyc +0 -0
  16. app/app/api/__pycache__/report.cpython-311.pyc +0 -0
  17. app/app/api/__pycache__/simulation.cpython-311.pyc +3 -0
  18. app/app/api/billing.py +82 -0
  19. app/app/api/graph.py +617 -0
  20. app/app/api/pipeline.py +230 -0
  21. app/app/api/report.py +1015 -0
  22. app/app/api/simulation.py +2718 -0
  23. app/app/config.py +129 -0
  24. app/app/models/__init__.py +9 -0
  25. app/app/models/__pycache__/__init__.cpython-311.pyc +0 -0
  26. app/app/models/__pycache__/project.cpython-311.pyc +0 -0
  27. app/app/models/__pycache__/task.cpython-311.pyc +0 -0
  28. app/app/models/project.py +305 -0
  29. app/app/models/task.py +184 -0
  30. app/app/services/__init__.py +73 -0
  31. app/app/services/__pycache__/__init__.cpython-311.pyc +0 -0
  32. app/app/services/__pycache__/graph_builder.cpython-311.pyc +0 -0
  33. app/app/services/__pycache__/oasis_profile_generator.cpython-311.pyc +0 -0
  34. app/app/services/__pycache__/ontology_generator.cpython-311.pyc +0 -0
  35. app/app/services/__pycache__/pipeline_orchestrator.cpython-311.pyc +0 -0
  36. app/app/services/__pycache__/report_agent.cpython-311.pyc +3 -0
  37. app/app/services/__pycache__/simulation_config_generator.cpython-311.pyc +0 -0
  38. app/app/services/__pycache__/simulation_ipc.cpython-311.pyc +0 -0
  39. app/app/services/__pycache__/simulation_manager.cpython-311.pyc +0 -0
  40. app/app/services/__pycache__/simulation_runner.cpython-311.pyc +0 -0
  41. app/app/services/__pycache__/supabase_jobs.cpython-311.pyc +0 -0
  42. app/app/services/__pycache__/text_processor.cpython-311.pyc +0 -0
  43. app/app/services/__pycache__/zep_entity_reader.cpython-311.pyc +0 -0
  44. app/app/services/__pycache__/zep_graph_memory_updater.cpython-311.pyc +0 -0
  45. app/app/services/__pycache__/zep_tools.cpython-311.pyc +0 -0
  46. app/app/services/billing_service.py +85 -0
  47. app/app/services/graph_builder.py +532 -0
  48. app/app/services/oasis_profile_generator.py +1339 -0
  49. app/app/services/ontology_generator.py +507 -0
  50. app/app/services/pipeline_orchestrator.py +368 -0
.gitattributes CHANGED
@@ -35,3 +35,5 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
  app/api/__pycache__/simulation.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
37
  app/services/__pycache__/report_agent.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
 
 
 
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
  app/api/__pycache__/simulation.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
37
  app/services/__pycache__/report_agent.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
38
+ app/app/api/__pycache__/simulation.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
39
+ app/app/services/__pycache__/report_agent.cpython-311.pyc filter=lfs diff=lfs merge=lfs -text
app/.dockerignore ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __pycache__
2
+ *.py[cod]
3
+ *.egg-info
4
+ .eggs
5
+ .env
6
+ .venv
7
+ venv
8
+ uploads
9
+ *.log
10
+ .git
11
+ .gitignore
12
+ .pytest_cache
app/.python-version ADDED
@@ -0,0 +1 @@
 
 
1
+ 3.11.9
app/.python-version.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ 3.11.9
app/Dockerfile ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # MiroFish API — Docker (Hugging Face Spaces, Fly, etc.)
2
+ # Build context must be this backend/ directory (contains app/, wsgi.py, requirements.txt).
3
+
4
+ FROM python:3.11-slim-bookworm
5
+
6
+ WORKDIR /app
7
+
8
+ RUN apt-get update && apt-get install -y --no-install-recommends \
9
+ build-essential \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ COPY requirements.txt .
13
+ RUN pip install --no-cache-dir -r requirements.txt
14
+
15
+ COPY . .
16
+
17
+ ENV PYTHONUNBUFFERED=1
18
+ EXPOSE 7860
19
+
20
+ CMD gunicorn --bind "0.0.0.0:${PORT:-7860}" --workers 1 --threads 4 --timeout 600 --worker-tmp-dir /dev/shm wsgi:app
app/HUGGINGFACE_SPACE.md ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Hugging Face Spaces — MiroFish API (Docker)
2
+
3
+ ## Create the Space
4
+
5
+ 1. [huggingface.co/spaces](https://huggingface.co/spaces) → **Create new Space**
6
+ 2. **SDK**: **Docker**
7
+ 3. **Hardware**: CPU (free) — long simulations may OOM or time out; upgrade if needed.
8
+ 4. Connect this repo and set the Docker **app directory** to **`backend`** (or copy only `backend/` into a dedicated repo).
9
+
10
+ ### Visibility must be **Public** (important)
11
+
12
+ If the Space is **Private**, browsers and your static site get **`401` with body `unauthorized`** (plain text) **before** the request reaches Flask.
13
+ Open **Space → Settings → Visibility** and set **Public** so `https://YOUR-SPACE.hf.space/api/...` works without a Hugging Face login cookie.
14
+
15
+ Build context = folder that contains `Dockerfile`, `requirements.txt`, `app/`, `wsgi.py`.
16
+
17
+ ## Secrets (Space → Settings → Variables and secrets)
18
+
19
+ | Variable | Notes |
20
+ |----------|--------|
21
+ | `SECRET_KEY` | Long random string (Flask). |
22
+ | `LLM_API_KEY`, `LLM_BASE_URL`, `LLM_MODEL_NAME` | Your LLM. |
23
+ | `ZEP_API_KEY` | Zep. |
24
+ | `SUPABASE_URL` | `https://xxx.supabase.co` |
25
+ | `SUPABASE_SERVICE_ROLE_KEY` | Service role (never expose to browser). |
26
+ | `PIPELINE_CLIENT_SECRET` | **Same** as website `VITE_PIPELINE_CLIENT_SECRET`. |
27
+ | `SUPABASE_REPORTS_BUCKET` | Default `reports`. |
28
+ | `PIPELINE_MAX_ROUNDS` | Optional integer cap. |
29
+
30
+ Optional: `SUPABASE_JWT_SECRET` if you use Supabase Auth JWT on the pipeline.
31
+
32
+ ## Website (`website/.env`)
33
+
34
+ ```env
35
+ VITE_API_BASE_URLS=https://YOUR-SPACE.hf.space
36
+ VITE_PIPELINE_CLIENT_SECRET=same-as-PIPELINE_CLIENT_SECRET
37
+ VITE_SUPABASE_URL=https://xxx.supabase.co
38
+ VITE_SUPABASE_ANON_KEY=your_anon_key
39
+ ```
40
+
41
+ Login/sign-up use **Supabase** on `public."User"` — run `supabase/grants_public_user.sql` in the SQL editor.
42
+
43
+ ## Free tier notes
44
+
45
+ - Space **sleeps** when idle; first request is a cold start.
46
+ - Proxies may **time out** on very long HTTP requests; heavy pipelines may need paid hardware or another host.
47
+
48
+ ## Health check
49
+
50
+ `GET /health` → `{"status":"ok",...}`
app/Procfile ADDED
@@ -0,0 +1 @@
 
 
1
+ web: gunicorn --bind 0.0.0.0:$PORT --workers 1 --threads 8 --timeout 600 wsgi:app
app/app/__init__.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ MiroFish Backend - Flask应用工厂
3
+ """
4
+
5
+ import os
6
+ import warnings
7
+
8
+ # 抑制 multiprocessing resource_tracker 的警告(来自第三方库如 transformers)
9
+ # 需要在所有其他导入之前设置
10
+ warnings.filterwarnings("ignore", message=".*resource_tracker.*")
11
+
12
+ from flask import Flask, request
13
+ from flask_cors import CORS
14
+
15
+ from .config import Config
16
+ from .services.supabase_jobs import sync_enabled as supabase_sync_enabled
17
+ from .utils.logger import setup_logger, get_logger
18
+
19
+
20
+ def create_app(config_class=Config):
21
+ """Flask应用工厂函数"""
22
+ app = Flask(__name__)
23
+ app.config.from_object(config_class)
24
+
25
+ # 设置JSON编码:确保中文直接显示(而不是 \uXXXX 格式)
26
+ # Flask >= 2.3 使用 app.json.ensure_ascii,旧版本使用 JSON_AS_ASCII 配置
27
+ if hasattr(app, 'json') and hasattr(app.json, 'ensure_ascii'):
28
+ app.json.ensure_ascii = False
29
+
30
+ # 设置日志
31
+ logger = setup_logger('mirofish')
32
+
33
+ # 只在 reloader 子进程中打印启动信息(避免 debug 模式下打印两次)
34
+ is_reloader_process = os.environ.get('WERKZEUG_RUN_MAIN') == 'true'
35
+ debug_mode = app.config.get('DEBUG', False)
36
+ should_log_startup = not debug_mode or is_reloader_process
37
+
38
+ if should_log_startup:
39
+ logger.info("=" * 50)
40
+ logger.info("MiroFish Backend 启动中...")
41
+ logger.info("=" * 50)
42
+
43
+ # CORS: allow custom pipeline auth headers from static sites / other origins
44
+ CORS(
45
+ app,
46
+ resources={
47
+ r"/api/*": {
48
+ "origins": "*",
49
+ "allow_headers": [
50
+ "Content-Type",
51
+ "Authorization",
52
+ "X-Mirofish-User-Id",
53
+ "X-Pipeline-Client-Secret",
54
+ ],
55
+ }
56
+ },
57
+ )
58
+
59
+ # 注册模拟进程清理函数(确保服务器关闭时终止所有模拟进程)
60
+ from .services.simulation_runner import SimulationRunner
61
+ SimulationRunner.register_cleanup()
62
+ if should_log_startup:
63
+ logger.info("已注册模拟进程清理函数")
64
+
65
+ # 请求日志中间件
66
+ @app.before_request
67
+ def log_request():
68
+ logger = get_logger('mirofish.request')
69
+ logger.debug(f"请求: {request.method} {request.path}")
70
+ if request.content_type and 'json' in request.content_type:
71
+ logger.debug(f"请求体: {request.get_json(silent=True)}")
72
+
73
+ @app.after_request
74
+ def log_response(response):
75
+ logger = get_logger('mirofish.request')
76
+ logger.debug(f"响应: {response.status_code}")
77
+ return response
78
+
79
+ # 注册蓝图
80
+ from .api import graph_bp, simulation_bp, report_bp, pipeline_bp, billing_bp
81
+ app.register_blueprint(graph_bp, url_prefix='/api/graph')
82
+ app.register_blueprint(simulation_bp, url_prefix='/api/simulation')
83
+ app.register_blueprint(report_bp, url_prefix='/api/report')
84
+ app.register_blueprint(pipeline_bp, url_prefix='/api/pipeline')
85
+ app.register_blueprint(billing_bp, url_prefix='/api/billing')
86
+
87
+ # 健康检查
88
+ @app.route('/health')
89
+ def health():
90
+ return {
91
+ 'status': 'ok',
92
+ 'service': 'MiroFish Backend',
93
+ 'supabase_sync': supabase_sync_enabled(),
94
+ 'pipeline_client_secret_configured': bool((Config.PIPELINE_CLIENT_SECRET or '').strip()),
95
+ }
96
+
97
+ if should_log_startup:
98
+ if supabase_sync_enabled():
99
+ logger.info("Supabase cloud sync: ENABLED (simulation_runs + reports)")
100
+ else:
101
+ logger.warning(
102
+ "Supabase cloud sync: DISABLED — set SUPABASE_URL and SUPABASE_SERVICE_ROLE_KEY "
103
+ "so runs appear in Past simulations and in the database."
104
+ )
105
+ logger.info("MiroFish Backend 启动完成")
106
+
107
+ return app
108
+
app/app/__pycache__/__init__.cpython-311.pyc ADDED
Binary file (4.57 kB). View file
 
app/app/__pycache__/config.cpython-311.pyc ADDED
Binary file (7.78 kB). View file
 
app/app/api/__init__.py ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ API路由模块
3
+ """
4
+
5
+ from flask import Blueprint
6
+
7
+ graph_bp = Blueprint('graph', __name__)
8
+ simulation_bp = Blueprint('simulation', __name__)
9
+ report_bp = Blueprint('report', __name__)
10
+ pipeline_bp = Blueprint('pipeline', __name__)
11
+ billing_bp = Blueprint('billing', __name__)
12
+
13
+ from . import graph # noqa: E402, F401
14
+ from . import simulation # noqa: E402, F401
15
+ from . import report # noqa: E402, F401
16
+ from . import pipeline # noqa: E402, F401
17
+ from . import billing # noqa: E402, F401
18
+
app/app/api/__pycache__/__init__.cpython-311.pyc ADDED
Binary file (787 Bytes). View file
 
app/app/api/__pycache__/billing.cpython-311.pyc ADDED
Binary file (7.13 kB). View file
 
app/app/api/__pycache__/graph.cpython-311.pyc ADDED
Binary file (23.7 kB). View file
 
app/app/api/__pycache__/pipeline.cpython-311.pyc ADDED
Binary file (13.6 kB). View file
 
app/app/api/__pycache__/report.cpython-311.pyc ADDED
Binary file (36.4 kB). View file
 
app/app/api/__pycache__/simulation.cpython-311.pyc ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9dce4de5fc345cab32a0ca7867868edae47884de3c0f667c0268cba30c6364d3
3
+ size 108767
app/app/api/billing.py ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Razorpay + coupon unlock for /api/pipeline/start."""
2
+
3
+ from flask import request, jsonify
4
+
5
+ from . import billing_bp
6
+ from ..config import Config
7
+ from ..services import billing_service
8
+ from ..utils.logger import get_logger
9
+
10
+ logger = get_logger("mirofish.api.billing")
11
+
12
+
13
+ @billing_bp.route("/status", methods=["GET"])
14
+ def billing_status():
15
+ if not Config.billing_enabled():
16
+ return jsonify(
17
+ {
18
+ "success": True,
19
+ "data": {
20
+ "enabled": False,
21
+ "key_id": None,
22
+ "amount": None,
23
+ "currency": None,
24
+ },
25
+ }
26
+ )
27
+ data = {
28
+ "enabled": True,
29
+ "key_id": Config.RAZORPAY_KEY_ID or None,
30
+ "amount": Config.RAZORPAY_AMOUNT_PAISE if Config.RAZORPAY_KEY_SECRET else None,
31
+ "currency": Config.RAZORPAY_CURRENCY if Config.RAZORPAY_KEY_SECRET else None,
32
+ "coupons_enabled": bool(Config.payment_coupon_set()),
33
+ }
34
+ return jsonify({"success": True, "data": data})
35
+
36
+
37
+ @billing_bp.route("/create-order", methods=["POST"])
38
+ def billing_create_order():
39
+ try:
40
+ if not (Config.RAZORPAY_KEY_ID and Config.RAZORPAY_KEY_SECRET):
41
+ return jsonify({"success": False, "error": "Razorpay is not configured"}), 503
42
+ order = billing_service.create_razorpay_order()
43
+ return jsonify({"success": True, "data": order})
44
+ except Exception as e:
45
+ logger.error("create-order: %s", e)
46
+ return jsonify({"success": False, "error": str(e)}), 500
47
+
48
+
49
+ @billing_bp.route("/verify-payment", methods=["POST"])
50
+ def billing_verify_payment():
51
+ try:
52
+ if not (Config.RAZORPAY_KEY_ID and Config.RAZORPAY_KEY_SECRET):
53
+ return jsonify({"success": False, "error": "Razorpay is not configured"}), 503
54
+ body = request.get_json(silent=True) or {}
55
+ oid = (body.get("razorpay_order_id") or "").strip()
56
+ pid = (body.get("razorpay_payment_id") or "").strip()
57
+ sig = (body.get("razorpay_signature") or "").strip()
58
+ if not oid or not pid or not sig:
59
+ return jsonify({"success": False, "error": "Missing payment fields"}), 400
60
+ if not billing_service.verify_razorpay_payment_signature(oid, pid, sig):
61
+ return jsonify({"success": False, "error": "Invalid payment signature"}), 400
62
+ token = billing_service.issue_simulation_payment_token(via="razorpay")
63
+ return jsonify({"success": True, "data": {"simulation_payment_token": token}})
64
+ except Exception as e:
65
+ logger.error("verify-payment: %s", e)
66
+ return jsonify({"success": False, "error": str(e)}), 500
67
+
68
+
69
+ @billing_bp.route("/apply-coupon", methods=["POST"])
70
+ def billing_apply_coupon():
71
+ try:
72
+ if not Config.payment_coupon_set():
73
+ return jsonify({"success": False, "error": "Coupons are not configured"}), 503
74
+ body = request.get_json(silent=True) or {}
75
+ code = body.get("code") or body.get("coupon") or ""
76
+ if not billing_service.coupon_valid(str(code)):
77
+ return jsonify({"success": False, "error": "Invalid coupon code"}), 400
78
+ token = billing_service.issue_simulation_payment_token(via="coupon")
79
+ return jsonify({"success": True, "data": {"simulation_payment_token": token}})
80
+ except Exception as e:
81
+ logger.error("apply-coupon: %s", e)
82
+ return jsonify({"success": False, "error": str(e)}), 500
app/app/api/graph.py ADDED
@@ -0,0 +1,617 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 图谱相关API路由
3
+ 采用项目上下文机制,服务端持久化状态
4
+ """
5
+
6
+ import os
7
+ import traceback
8
+ import threading
9
+ from flask import request, jsonify
10
+
11
+ from . import graph_bp
12
+ from ..config import Config
13
+ from ..services.ontology_generator import OntologyGenerator
14
+ from ..services.graph_builder import GraphBuilderService
15
+ from ..services.text_processor import TextProcessor
16
+ from ..utils.file_parser import FileParser
17
+ from ..utils.logger import get_logger
18
+ from ..models.task import TaskManager, TaskStatus
19
+ from ..models.project import ProjectManager, ProjectStatus
20
+
21
+ # 获取日志器
22
+ logger = get_logger('mirofish.api')
23
+
24
+
25
+ def allowed_file(filename: str) -> bool:
26
+ """检查文件扩展名是否允许"""
27
+ if not filename or '.' not in filename:
28
+ return False
29
+ ext = os.path.splitext(filename)[1].lower().lstrip('.')
30
+ return ext in Config.ALLOWED_EXTENSIONS
31
+
32
+
33
+ # ============== 项目管理接口 ==============
34
+
35
+ @graph_bp.route('/project/<project_id>', methods=['GET'])
36
+ def get_project(project_id: str):
37
+ """
38
+ 获取项目详情
39
+ """
40
+ project = ProjectManager.get_project(project_id)
41
+
42
+ if not project:
43
+ return jsonify({
44
+ "success": False,
45
+ "error": f"项目不存在: {project_id}"
46
+ }), 404
47
+
48
+ return jsonify({
49
+ "success": True,
50
+ "data": project.to_dict()
51
+ })
52
+
53
+
54
+ @graph_bp.route('/project/list', methods=['GET'])
55
+ def list_projects():
56
+ """
57
+ 列出所有项目
58
+ """
59
+ limit = request.args.get('limit', 50, type=int)
60
+ projects = ProjectManager.list_projects(limit=limit)
61
+
62
+ return jsonify({
63
+ "success": True,
64
+ "data": [p.to_dict() for p in projects],
65
+ "count": len(projects)
66
+ })
67
+
68
+
69
+ @graph_bp.route('/project/<project_id>', methods=['DELETE'])
70
+ def delete_project(project_id: str):
71
+ """
72
+ 删除项目
73
+ """
74
+ success = ProjectManager.delete_project(project_id)
75
+
76
+ if not success:
77
+ return jsonify({
78
+ "success": False,
79
+ "error": f"项目不存在或删除失败: {project_id}"
80
+ }), 404
81
+
82
+ return jsonify({
83
+ "success": True,
84
+ "message": f"项目已删除: {project_id}"
85
+ })
86
+
87
+
88
+ @graph_bp.route('/project/<project_id>/reset', methods=['POST'])
89
+ def reset_project(project_id: str):
90
+ """
91
+ 重置项目状态(用于重新构建图谱)
92
+ """
93
+ project = ProjectManager.get_project(project_id)
94
+
95
+ if not project:
96
+ return jsonify({
97
+ "success": False,
98
+ "error": f"项目不存在: {project_id}"
99
+ }), 404
100
+
101
+ # 重置到本体已生成状态
102
+ if project.ontology:
103
+ project.status = ProjectStatus.ONTOLOGY_GENERATED
104
+ else:
105
+ project.status = ProjectStatus.CREATED
106
+
107
+ project.graph_id = None
108
+ project.graph_build_task_id = None
109
+ project.error = None
110
+ ProjectManager.save_project(project)
111
+
112
+ return jsonify({
113
+ "success": True,
114
+ "message": f"项目已重置: {project_id}",
115
+ "data": project.to_dict()
116
+ })
117
+
118
+
119
+ # ============== 接口1:上传文件并生成本体 ==============
120
+
121
+ @graph_bp.route('/ontology/generate', methods=['POST'])
122
+ def generate_ontology():
123
+ """
124
+ 接口1:上传文件,分析生成本体定义
125
+
126
+ 请求方式:multipart/form-data
127
+
128
+ 参数:
129
+ files: 上传的文件(PDF/MD/TXT),可多个
130
+ simulation_requirement: 模拟需求描述(必填)
131
+ project_name: 项目名称(可选)
132
+ additional_context: 额外说明(可选)
133
+
134
+ 返回:
135
+ {
136
+ "success": true,
137
+ "data": {
138
+ "project_id": "proj_xxxx",
139
+ "ontology": {
140
+ "entity_types": [...],
141
+ "edge_types": [...],
142
+ "analysis_summary": "..."
143
+ },
144
+ "files": [...],
145
+ "total_text_length": 12345
146
+ }
147
+ }
148
+ """
149
+ try:
150
+ logger.info("=== 开始生成本体定义 ===")
151
+
152
+ # 获取参数
153
+ simulation_requirement = request.form.get('simulation_requirement', '')
154
+ project_name = request.form.get('project_name', 'Unnamed Project')
155
+ additional_context = request.form.get('additional_context', '')
156
+
157
+ logger.debug(f"项目名称: {project_name}")
158
+ logger.debug(f"模拟需求: {simulation_requirement[:100]}...")
159
+
160
+ if not simulation_requirement:
161
+ return jsonify({
162
+ "success": False,
163
+ "error": "请提供模拟需求描述 (simulation_requirement)"
164
+ }), 400
165
+
166
+ # 获取上传的文件
167
+ uploaded_files = request.files.getlist('files')
168
+ if not uploaded_files or all(not f.filename for f in uploaded_files):
169
+ return jsonify({
170
+ "success": False,
171
+ "error": "请至少上传一个文档文件"
172
+ }), 400
173
+
174
+ # 创建项目
175
+ project = ProjectManager.create_project(name=project_name)
176
+ project.simulation_requirement = simulation_requirement
177
+ logger.info(f"创建项目: {project.project_id}")
178
+
179
+ # 保存文件并提取文本
180
+ document_texts = []
181
+ all_text = ""
182
+
183
+ for file in uploaded_files:
184
+ if file and file.filename and allowed_file(file.filename):
185
+ # 保存文件到项目目录
186
+ file_info = ProjectManager.save_file_to_project(
187
+ project.project_id,
188
+ file,
189
+ file.filename
190
+ )
191
+ project.files.append({
192
+ "filename": file_info["original_filename"],
193
+ "size": file_info["size"]
194
+ })
195
+
196
+ # 提取文本
197
+ text = FileParser.extract_text(file_info["path"])
198
+ text = TextProcessor.preprocess_text(text)
199
+ document_texts.append(text)
200
+ all_text += f"\n\n=== {file_info['original_filename']} ===\n{text}"
201
+
202
+ if not document_texts:
203
+ ProjectManager.delete_project(project.project_id)
204
+ return jsonify({
205
+ "success": False,
206
+ "error": "没有成功处理任何文档,请检查文件格式"
207
+ }), 400
208
+
209
+ # 保存提取的文本
210
+ project.total_text_length = len(all_text)
211
+ ProjectManager.save_extracted_text(project.project_id, all_text)
212
+ logger.info(f"文本提取完成,共 {len(all_text)} 字符")
213
+
214
+ # 生成本体
215
+ logger.info("调用 LLM 生成本体定义...")
216
+ generator = OntologyGenerator()
217
+ ontology = generator.generate(
218
+ document_texts=document_texts,
219
+ simulation_requirement=simulation_requirement,
220
+ additional_context=additional_context if additional_context else None
221
+ )
222
+
223
+ # 保存本体到项目
224
+ entity_count = len(ontology.get("entity_types", []))
225
+ edge_count = len(ontology.get("edge_types", []))
226
+ logger.info(f"本体生成完成: {entity_count} 个实体类型, {edge_count} 个关系类型")
227
+
228
+ project.ontology = {
229
+ "entity_types": ontology.get("entity_types", []),
230
+ "edge_types": ontology.get("edge_types", [])
231
+ }
232
+ project.analysis_summary = ontology.get("analysis_summary", "")
233
+ project.status = ProjectStatus.ONTOLOGY_GENERATED
234
+ ProjectManager.save_project(project)
235
+ logger.info(f"=== 本体生成完成 === 项目ID: {project.project_id}")
236
+
237
+ return jsonify({
238
+ "success": True,
239
+ "data": {
240
+ "project_id": project.project_id,
241
+ "project_name": project.name,
242
+ "ontology": project.ontology,
243
+ "analysis_summary": project.analysis_summary,
244
+ "files": project.files,
245
+ "total_text_length": project.total_text_length
246
+ }
247
+ })
248
+
249
+ except Exception as e:
250
+ return jsonify({
251
+ "success": False,
252
+ "error": str(e),
253
+ "traceback": traceback.format_exc()
254
+ }), 500
255
+
256
+
257
+ # ============== 接口2:构建图谱 ==============
258
+
259
+ @graph_bp.route('/build', methods=['POST'])
260
+ def build_graph():
261
+ """
262
+ 接口2:根据project_id构建图谱
263
+
264
+ 请求(JSON):
265
+ {
266
+ "project_id": "proj_xxxx", // 必填,来自接口1
267
+ "graph_name": "图谱名称", // 可选
268
+ "chunk_size": 500, // 可选,默认500
269
+ "chunk_overlap": 50 // 可选,默认50
270
+ }
271
+
272
+ 返回:
273
+ {
274
+ "success": true,
275
+ "data": {
276
+ "project_id": "proj_xxxx",
277
+ "task_id": "task_xxxx",
278
+ "message": "图谱构建任务已启动"
279
+ }
280
+ }
281
+ """
282
+ try:
283
+ logger.info("=== 开始构建图谱 ===")
284
+
285
+ # 检查配置
286
+ errors = []
287
+ if not Config.ZEP_API_KEY:
288
+ errors.append("ZEP_API_KEY未配置")
289
+ if errors:
290
+ logger.error(f"配置错误: {errors}")
291
+ return jsonify({
292
+ "success": False,
293
+ "error": "配置错误: " + "; ".join(errors)
294
+ }), 500
295
+
296
+ # 解析请求
297
+ data = request.get_json() or {}
298
+ project_id = data.get('project_id')
299
+ logger.debug(f"请求参数: project_id={project_id}")
300
+
301
+ if not project_id:
302
+ return jsonify({
303
+ "success": False,
304
+ "error": "请提供 project_id"
305
+ }), 400
306
+
307
+ # 获取项目
308
+ project = ProjectManager.get_project(project_id)
309
+ if not project:
310
+ return jsonify({
311
+ "success": False,
312
+ "error": f"项目不存在: {project_id}"
313
+ }), 404
314
+
315
+ # 检查项目状态
316
+ force = data.get('force', False) # 强制重新构建
317
+
318
+ if project.status == ProjectStatus.CREATED:
319
+ return jsonify({
320
+ "success": False,
321
+ "error": "项目尚未生成本体,请先调用 /ontology/generate"
322
+ }), 400
323
+
324
+ if project.status == ProjectStatus.GRAPH_BUILDING and not force:
325
+ return jsonify({
326
+ "success": False,
327
+ "error": "图谱正在构建中,请勿重复提交。如需强制重建,请添加 force: true",
328
+ "task_id": project.graph_build_task_id
329
+ }), 400
330
+
331
+ # 如果强制重建,重置状态
332
+ if force and project.status in [ProjectStatus.GRAPH_BUILDING, ProjectStatus.FAILED, ProjectStatus.GRAPH_COMPLETED]:
333
+ project.status = ProjectStatus.ONTOLOGY_GENERATED
334
+ project.graph_id = None
335
+ project.graph_build_task_id = None
336
+ project.error = None
337
+
338
+ # 获取配置
339
+ graph_name = data.get('graph_name', project.name or 'MiroFish Graph')
340
+ chunk_size = data.get('chunk_size', project.chunk_size or Config.DEFAULT_CHUNK_SIZE)
341
+ chunk_overlap = data.get('chunk_overlap', project.chunk_overlap or Config.DEFAULT_CHUNK_OVERLAP)
342
+
343
+ # 更新项目配置
344
+ project.chunk_size = chunk_size
345
+ project.chunk_overlap = chunk_overlap
346
+
347
+ # 获取提取的文本
348
+ text = ProjectManager.get_extracted_text(project_id)
349
+ if not text:
350
+ return jsonify({
351
+ "success": False,
352
+ "error": "未找到提取的文本内容"
353
+ }), 400
354
+
355
+ # 获取本体
356
+ ontology = project.ontology
357
+ if not ontology:
358
+ return jsonify({
359
+ "success": False,
360
+ "error": "未找到本体定义"
361
+ }), 400
362
+
363
+ # 创建异步任务
364
+ task_manager = TaskManager()
365
+ task_id = task_manager.create_task(f"构建图谱: {graph_name}")
366
+ logger.info(f"创建图谱构建任务: task_id={task_id}, project_id={project_id}")
367
+
368
+ # 更新项目状态
369
+ project.status = ProjectStatus.GRAPH_BUILDING
370
+ project.graph_build_task_id = task_id
371
+ ProjectManager.save_project(project)
372
+
373
+ # 启动后台任务
374
+ def build_task():
375
+ build_logger = get_logger('mirofish.build')
376
+ try:
377
+ build_logger.info(f"[{task_id}] 开始构建图谱...")
378
+ task_manager.update_task(
379
+ task_id,
380
+ status=TaskStatus.PROCESSING,
381
+ message="初始化图谱构建服务..."
382
+ )
383
+
384
+ # 创建图谱构建服务
385
+ builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
386
+
387
+ # 分块
388
+ task_manager.update_task(
389
+ task_id,
390
+ message="文本分块中...",
391
+ progress=5
392
+ )
393
+ chunks = TextProcessor.split_text(
394
+ text,
395
+ chunk_size=chunk_size,
396
+ overlap=chunk_overlap
397
+ )
398
+ total_chunks = len(chunks)
399
+
400
+ # 创建图谱
401
+ task_manager.update_task(
402
+ task_id,
403
+ message="创建Zep图谱...",
404
+ progress=10
405
+ )
406
+ graph_id = builder.create_graph(name=graph_name)
407
+
408
+ # 更新项目的graph_id
409
+ project.graph_id = graph_id
410
+ ProjectManager.save_project(project)
411
+
412
+ # 设置本体
413
+ task_manager.update_task(
414
+ task_id,
415
+ message="设置本体定义...",
416
+ progress=15
417
+ )
418
+ builder.set_ontology(graph_id, ontology)
419
+
420
+ # 添加文本(progress_callback 签名是 (msg, progress_ratio))
421
+ def add_progress_callback(msg, progress_ratio):
422
+ progress = 15 + int(progress_ratio * 40) # 15% - 55%
423
+ task_manager.update_task(
424
+ task_id,
425
+ message=msg,
426
+ progress=progress
427
+ )
428
+
429
+ task_manager.update_task(
430
+ task_id,
431
+ message=f"开始添加 {total_chunks} 个文本块...",
432
+ progress=15
433
+ )
434
+
435
+ episode_uuids = builder.add_text_batches(
436
+ graph_id,
437
+ chunks,
438
+ batch_size=3,
439
+ progress_callback=add_progress_callback
440
+ )
441
+
442
+ # 等待Zep处理完成(查询每个episode的processed状态)
443
+ task_manager.update_task(
444
+ task_id,
445
+ message="等待Zep处理数据...",
446
+ progress=55
447
+ )
448
+
449
+ def wait_progress_callback(msg, progress_ratio):
450
+ progress = 55 + int(progress_ratio * 35) # 55% - 90%
451
+ task_manager.update_task(
452
+ task_id,
453
+ message=msg,
454
+ progress=progress
455
+ )
456
+
457
+ builder._wait_for_episodes(episode_uuids, wait_progress_callback)
458
+
459
+ # 获取图谱数据
460
+ task_manager.update_task(
461
+ task_id,
462
+ message="获取图谱数据...",
463
+ progress=95
464
+ )
465
+ graph_data = builder.get_graph_data(graph_id)
466
+
467
+ # 更新项目状态
468
+ project.status = ProjectStatus.GRAPH_COMPLETED
469
+ ProjectManager.save_project(project)
470
+
471
+ node_count = graph_data.get("node_count", 0)
472
+ edge_count = graph_data.get("edge_count", 0)
473
+ build_logger.info(f"[{task_id}] 图谱构建完成: graph_id={graph_id}, 节点={node_count}, 边={edge_count}")
474
+
475
+ # 完成
476
+ task_manager.update_task(
477
+ task_id,
478
+ status=TaskStatus.COMPLETED,
479
+ message="图谱构建完成",
480
+ progress=100,
481
+ result={
482
+ "project_id": project_id,
483
+ "graph_id": graph_id,
484
+ "node_count": node_count,
485
+ "edge_count": edge_count,
486
+ "chunk_count": total_chunks
487
+ }
488
+ )
489
+
490
+ except Exception as e:
491
+ # 更新项目状态为失败
492
+ build_logger.error(f"[{task_id}] 图谱构建失败: {str(e)}")
493
+ build_logger.debug(traceback.format_exc())
494
+
495
+ project.status = ProjectStatus.FAILED
496
+ project.error = str(e)
497
+ ProjectManager.save_project(project)
498
+
499
+ task_manager.update_task(
500
+ task_id,
501
+ status=TaskStatus.FAILED,
502
+ message=f"构建失败: {str(e)}",
503
+ error=traceback.format_exc()
504
+ )
505
+
506
+ # 启动后台线程
507
+ thread = threading.Thread(target=build_task, daemon=True)
508
+ thread.start()
509
+
510
+ return jsonify({
511
+ "success": True,
512
+ "data": {
513
+ "project_id": project_id,
514
+ "task_id": task_id,
515
+ "message": "图谱构建任务已启动,请通过 /task/{task_id} 查询进度"
516
+ }
517
+ })
518
+
519
+ except Exception as e:
520
+ return jsonify({
521
+ "success": False,
522
+ "error": str(e),
523
+ "traceback": traceback.format_exc()
524
+ }), 500
525
+
526
+
527
+ # ============== 任务查询接口 ==============
528
+
529
+ @graph_bp.route('/task/<task_id>', methods=['GET'])
530
+ def get_task(task_id: str):
531
+ """
532
+ 查询任务状态
533
+ """
534
+ task = TaskManager().get_task(task_id)
535
+
536
+ if not task:
537
+ return jsonify({
538
+ "success": False,
539
+ "error": f"任务不存在: {task_id}"
540
+ }), 404
541
+
542
+ return jsonify({
543
+ "success": True,
544
+ "data": task.to_dict()
545
+ })
546
+
547
+
548
+ @graph_bp.route('/tasks', methods=['GET'])
549
+ def list_tasks():
550
+ """
551
+ 列出所有任务
552
+ """
553
+ tasks = TaskManager().list_tasks()
554
+
555
+ return jsonify({
556
+ "success": True,
557
+ "data": [t.to_dict() for t in tasks],
558
+ "count": len(tasks)
559
+ })
560
+
561
+
562
+ # ============== 图谱数据接口 ==============
563
+
564
+ @graph_bp.route('/data/<graph_id>', methods=['GET'])
565
+ def get_graph_data(graph_id: str):
566
+ """
567
+ 获取图谱数据(节点和边)
568
+ """
569
+ try:
570
+ if not Config.ZEP_API_KEY:
571
+ return jsonify({
572
+ "success": False,
573
+ "error": "ZEP_API_KEY未配置"
574
+ }), 500
575
+
576
+ builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
577
+ graph_data = builder.get_graph_data(graph_id)
578
+
579
+ return jsonify({
580
+ "success": True,
581
+ "data": graph_data
582
+ })
583
+
584
+ except Exception as e:
585
+ return jsonify({
586
+ "success": False,
587
+ "error": str(e),
588
+ "traceback": traceback.format_exc()
589
+ }), 500
590
+
591
+
592
+ @graph_bp.route('/delete/<graph_id>', methods=['DELETE'])
593
+ def delete_graph(graph_id: str):
594
+ """
595
+ 删除Zep图谱
596
+ """
597
+ try:
598
+ if not Config.ZEP_API_KEY:
599
+ return jsonify({
600
+ "success": False,
601
+ "error": "ZEP_API_KEY未配置"
602
+ }), 500
603
+
604
+ builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
605
+ builder.delete_graph(graph_id)
606
+
607
+ return jsonify({
608
+ "success": True,
609
+ "message": f"图谱已删除: {graph_id}"
610
+ })
611
+
612
+ except Exception as e:
613
+ return jsonify({
614
+ "success": False,
615
+ "error": str(e),
616
+ "traceback": traceback.format_exc()
617
+ }), 500
app/app/api/pipeline.py ADDED
@@ -0,0 +1,230 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ One-shot pipeline API for the website: upload text + requirement → full automation → report.
3
+
4
+ When Supabase sync is configured (URL + service role), records runs + reports.
5
+ Identify users: Bearer JWT (Mirofish / Supabase Auth) or X-Mirofish-User-Id + X-Pipeline-Client-Secret.
6
+ """
7
+
8
+ import threading
9
+ import traceback
10
+ import uuid
11
+
12
+ from flask import request, jsonify
13
+
14
+ from . import pipeline_bp
15
+ from ..models.task import TaskManager, TaskStatus
16
+ from ..services.pipeline_orchestrator import run_full_pipeline
17
+ from ..services.supabase_jobs import (
18
+ get_run_row_for_user,
19
+ insert_pending_run,
20
+ list_simulation_runs_for_user,
21
+ signed_url_for_storage_path,
22
+ sync_enabled,
23
+ )
24
+ from ..config import Config
25
+ from ..services.billing_service import verify_simulation_payment_token
26
+ from ..utils.logger import get_logger
27
+ from ..utils.supabase_auth import resolve_pipeline_user_id
28
+
29
+ logger = get_logger("mirofish.api.pipeline")
30
+
31
+ ALLOWED = {"txt", "md", "text"}
32
+
33
+
34
+ def _allowed(name: str) -> bool:
35
+ if not name or "." not in name:
36
+ return False
37
+ ext = name.rsplit(".", 1)[-1].lower()
38
+ return ext in ALLOWED
39
+
40
+
41
+ @pipeline_bp.route("/start", methods=["POST"])
42
+ def pipeline_start():
43
+ """
44
+ multipart/form-data:
45
+ - file: .txt (or .md) reality seed
46
+ - simulation_requirement: string
47
+
48
+ When Supabase sync is enabled: Bearer JWT (optional) or headers
49
+ X-Mirofish-User-Id + X-Pipeline-Client-Secret (same secret as PIPELINE_CLIENT_SECRET).
50
+
51
+ Returns { task_id, simulation_run_id? }.
52
+ """
53
+ try:
54
+ simulation_requirement = ""
55
+ document_text = ""
56
+ filename = "seed.txt"
57
+
58
+ if request.content_type and "multipart/form-data" in request.content_type:
59
+ simulation_requirement = (request.form.get("simulation_requirement") or "").strip()
60
+ f = request.files.get("file")
61
+ if not f or not f.filename:
62
+ return jsonify({"success": False, "error": "Missing file"}), 400
63
+ if not _allowed(f.filename):
64
+ return jsonify(
65
+ {"success": False, "error": "Only .txt or .md files are allowed for this endpoint"}
66
+ ), 400
67
+ filename = f.filename
68
+ raw = f.read()
69
+ try:
70
+ document_text = raw.decode("utf-8")
71
+ except UnicodeDecodeError:
72
+ document_text = raw.decode("utf-8", errors="replace")
73
+ else:
74
+ data = request.get_json(silent=True) or {}
75
+ simulation_requirement = (data.get("simulation_requirement") or "").strip()
76
+ document_text = (data.get("document_text") or "").strip()
77
+ filename = (data.get("filename") or "seed.txt").strip() or "seed.txt"
78
+ if not document_text:
79
+ return jsonify({"success": False, "error": "document_text is required"}), 400
80
+
81
+ if not simulation_requirement:
82
+ return jsonify({"success": False, "error": "simulation_requirement is required"}), 400
83
+ if len(document_text.strip()) < 10:
84
+ return jsonify({"success": False, "error": "Document text is too short"}), 400
85
+
86
+ if Config.billing_enabled():
87
+ if request.content_type and "multipart/form-data" in request.content_type:
88
+ pay_tok = (request.form.get("simulation_payment_token") or "").strip()
89
+ else:
90
+ _jd = request.get_json(silent=True) or {}
91
+ pay_tok = (_jd.get("simulation_payment_token") or "").strip()
92
+ if not verify_simulation_payment_token(pay_tok):
93
+ return jsonify(
94
+ {
95
+ "success": False,
96
+ "error": "Complete payment or apply a valid coupon before starting a simulation.",
97
+ }
98
+ ), 402
99
+
100
+ user_id = None
101
+ simulation_run_id = None
102
+ if sync_enabled():
103
+ uid, err = resolve_pipeline_user_id()
104
+ if err or not uid:
105
+ return jsonify({"success": False, "error": err or "Unauthorized"}), 401
106
+ user_id = uid
107
+ simulation_run_id = str(uuid.uuid4())
108
+
109
+ tm = TaskManager()
110
+ task_id = tm.create_task(
111
+ task_type="full_pipeline",
112
+ metadata={
113
+ "filename": filename,
114
+ "simulation_run_id": simulation_run_id,
115
+ "user_id": user_id,
116
+ },
117
+ )
118
+ tm.update_task(
119
+ task_id,
120
+ status=TaskStatus.PROCESSING,
121
+ progress=1,
122
+ message="Pipeline queued…",
123
+ )
124
+
125
+ if sync_enabled() and simulation_run_id and user_id:
126
+ try:
127
+ insert_pending_run(
128
+ simulation_run_id=simulation_run_id,
129
+ user_id=user_id,
130
+ backend_task_id=task_id,
131
+ requirement_preview=simulation_requirement,
132
+ )
133
+ except Exception as e:
134
+ logger.error("Supabase insert_pending_run: %s", e)
135
+ tm.fail_task(task_id, f"Could not record run: {e}")
136
+ return jsonify({"success": False, "error": str(e)}), 500
137
+
138
+ def worker():
139
+ run_full_pipeline(
140
+ simulation_requirement=simulation_requirement,
141
+ document_text=document_text,
142
+ source_filename=filename,
143
+ task_id=task_id,
144
+ task_manager=tm,
145
+ simulation_run_id=simulation_run_id,
146
+ user_id=user_id,
147
+ )
148
+
149
+ threading.Thread(target=worker, daemon=True).start()
150
+
151
+ sync = sync_enabled()
152
+ payload = {
153
+ "task_id": task_id,
154
+ "supabase_sync": sync,
155
+ "message": (
156
+ "Simulation started. Runs are saved to Supabase when cloud sync is enabled."
157
+ if sync
158
+ else "Simulation started. Cloud sync is OFF — add SUPABASE_URL + SUPABASE_SERVICE_ROLE_KEY on the server to record runs in Past simulations."
159
+ ),
160
+ }
161
+ if simulation_run_id:
162
+ payload["simulation_run_id"] = simulation_run_id
163
+
164
+ if sync:
165
+ logger.info(
166
+ "pipeline/start task=%s simulation_run_id=%s user_id=%s",
167
+ task_id,
168
+ simulation_run_id,
169
+ user_id,
170
+ )
171
+ else:
172
+ logger.warning(
173
+ "pipeline/start task=%s without Supabase sync — no simulation_runs row (set SUPABASE_* env)",
174
+ task_id,
175
+ )
176
+
177
+ return jsonify({"success": True, "data": payload})
178
+ except Exception as e:
179
+ logger.error(traceback.format_exc())
180
+ return jsonify({"success": False, "error": str(e)}), 500
181
+
182
+
183
+ @pipeline_bp.route("/status/<task_id>", methods=["GET"])
184
+ def pipeline_status(task_id: str):
185
+ tm = TaskManager()
186
+ task = tm.get_task(task_id)
187
+ if not task:
188
+ return jsonify({"success": False, "error": "Unknown task_id"}), 404
189
+ out = task.to_dict()
190
+ return jsonify({"success": True, "data": out})
191
+
192
+
193
+ @pipeline_bp.route("/runs", methods=["GET"])
194
+ def pipeline_list_runs():
195
+ """List simulation_runs for the authenticated user (service role query on server)."""
196
+ try:
197
+ if not sync_enabled():
198
+ return jsonify({"success": True, "data": {"runs": [], "supabase_sync": False}})
199
+ uid, err = resolve_pipeline_user_id()
200
+ if err or not uid:
201
+ return jsonify({"success": False, "error": err or "Unauthorized"}), 401
202
+ runs = list_simulation_runs_for_user(user_id=uid)
203
+ return jsonify({"success": True, "data": {"runs": runs, "supabase_sync": True}})
204
+ except Exception as e:
205
+ logger.error(traceback.format_exc())
206
+ return jsonify({"success": False, "error": str(e)}), 500
207
+
208
+
209
+ @pipeline_bp.route("/runs/<run_id>/signed-download", methods=["GET"])
210
+ def pipeline_signed_report_download(run_id: str):
211
+ """Return a time-limited signed URL for the report object in Storage."""
212
+ try:
213
+ if not sync_enabled():
214
+ return jsonify({"success": False, "error": "Sync not configured"}), 503
215
+ uid, err = resolve_pipeline_user_id()
216
+ if err or not uid:
217
+ return jsonify({"success": False, "error": err or "Unauthorized"}), 401
218
+ row = get_run_row_for_user(run_id=run_id, user_id=uid)
219
+ if not row:
220
+ return jsonify({"success": False, "error": "Run not found"}), 404
221
+ path = row.get("report_storage_path") or ""
222
+ if not path:
223
+ return jsonify({"success": False, "error": "Report not ready"}), 400
224
+ url = signed_url_for_storage_path(storage_path=path, expires_sec=3600)
225
+ if not url:
226
+ return jsonify({"success": False, "error": "Could not create download link"}), 500
227
+ return jsonify({"success": True, "data": {"signed_url": url}})
228
+ except Exception as e:
229
+ logger.error(traceback.format_exc())
230
+ return jsonify({"success": False, "error": str(e)}), 500
app/app/api/report.py ADDED
@@ -0,0 +1,1015 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Report API路由
3
+ 提供模拟报告生成、获取、对话等接口
4
+ """
5
+
6
+ import os
7
+ import traceback
8
+ import threading
9
+ from flask import request, jsonify, send_file
10
+
11
+ from . import report_bp
12
+ from ..config import Config
13
+ from ..services.report_agent import ReportAgent, ReportManager, ReportStatus
14
+ from ..services.simulation_manager import SimulationManager
15
+ from ..models.project import ProjectManager
16
+ from ..models.task import TaskManager, TaskStatus
17
+ from ..utils.logger import get_logger
18
+
19
+ logger = get_logger('mirofish.api.report')
20
+
21
+
22
+ # ============== 报告生成接口 ==============
23
+
24
+ @report_bp.route('/generate', methods=['POST'])
25
+ def generate_report():
26
+ """
27
+ 生成模拟分析报告(异步任务)
28
+
29
+ 这是一个耗时操作,接口会立即返回task_id,
30
+ 使用 GET /api/report/generate/status 查询进度
31
+
32
+ 请求(JSON):
33
+ {
34
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
35
+ "force_regenerate": false // 可选,强制重新生成
36
+ }
37
+
38
+ 返回:
39
+ {
40
+ "success": true,
41
+ "data": {
42
+ "simulation_id": "sim_xxxx",
43
+ "task_id": "task_xxxx",
44
+ "status": "generating",
45
+ "message": "报告生成任务已启动"
46
+ }
47
+ }
48
+ """
49
+ try:
50
+ data = request.get_json() or {}
51
+
52
+ simulation_id = data.get('simulation_id')
53
+ if not simulation_id:
54
+ return jsonify({
55
+ "success": False,
56
+ "error": "请提供 simulation_id"
57
+ }), 400
58
+
59
+ force_regenerate = data.get('force_regenerate', False)
60
+
61
+ # 获取模拟信息
62
+ manager = SimulationManager()
63
+ state = manager.get_simulation(simulation_id)
64
+
65
+ if not state:
66
+ return jsonify({
67
+ "success": False,
68
+ "error": f"模拟不存在: {simulation_id}"
69
+ }), 404
70
+
71
+ # 检查是否已有报告
72
+ if not force_regenerate:
73
+ existing_report = ReportManager.get_report_by_simulation(simulation_id)
74
+ if existing_report and existing_report.status == ReportStatus.COMPLETED:
75
+ return jsonify({
76
+ "success": True,
77
+ "data": {
78
+ "simulation_id": simulation_id,
79
+ "report_id": existing_report.report_id,
80
+ "status": "completed",
81
+ "message": "报告已存在",
82
+ "already_generated": True
83
+ }
84
+ })
85
+
86
+ # 获取项目信息
87
+ project = ProjectManager.get_project(state.project_id)
88
+ if not project:
89
+ return jsonify({
90
+ "success": False,
91
+ "error": f"项目不存在: {state.project_id}"
92
+ }), 404
93
+
94
+ graph_id = state.graph_id or project.graph_id
95
+ if not graph_id:
96
+ return jsonify({
97
+ "success": False,
98
+ "error": "缺少图谱ID,请确保已构建图谱"
99
+ }), 400
100
+
101
+ simulation_requirement = project.simulation_requirement
102
+ if not simulation_requirement:
103
+ return jsonify({
104
+ "success": False,
105
+ "error": "缺少模拟需求描述"
106
+ }), 400
107
+
108
+ # 提前生成 report_id,以便立即返回给前端
109
+ import uuid
110
+ report_id = f"report_{uuid.uuid4().hex[:12]}"
111
+
112
+ # 创建异步任务
113
+ task_manager = TaskManager()
114
+ task_id = task_manager.create_task(
115
+ task_type="report_generate",
116
+ metadata={
117
+ "simulation_id": simulation_id,
118
+ "graph_id": graph_id,
119
+ "report_id": report_id
120
+ }
121
+ )
122
+
123
+ # 定义后台任务
124
+ def run_generate():
125
+ try:
126
+ task_manager.update_task(
127
+ task_id,
128
+ status=TaskStatus.PROCESSING,
129
+ progress=0,
130
+ message="初始化Report Agent..."
131
+ )
132
+
133
+ # 创建Report Agent
134
+ agent = ReportAgent(
135
+ graph_id=graph_id,
136
+ simulation_id=simulation_id,
137
+ simulation_requirement=simulation_requirement
138
+ )
139
+
140
+ # 进度回调
141
+ def progress_callback(stage, progress, message):
142
+ task_manager.update_task(
143
+ task_id,
144
+ progress=progress,
145
+ message=f"[{stage}] {message}"
146
+ )
147
+
148
+ # 生成报告(传入预先生成的 report_id)
149
+ report = agent.generate_report(
150
+ progress_callback=progress_callback,
151
+ report_id=report_id
152
+ )
153
+
154
+ # 保存报告
155
+ ReportManager.save_report(report)
156
+
157
+ if report.status == ReportStatus.COMPLETED:
158
+ task_manager.complete_task(
159
+ task_id,
160
+ result={
161
+ "report_id": report.report_id,
162
+ "simulation_id": simulation_id,
163
+ "status": "completed"
164
+ }
165
+ )
166
+ else:
167
+ task_manager.fail_task(task_id, report.error or "报告生成失败")
168
+
169
+ except Exception as e:
170
+ logger.error(f"报告生成失败: {str(e)}")
171
+ task_manager.fail_task(task_id, str(e))
172
+
173
+ # 启动后台线程
174
+ thread = threading.Thread(target=run_generate, daemon=True)
175
+ thread.start()
176
+
177
+ return jsonify({
178
+ "success": True,
179
+ "data": {
180
+ "simulation_id": simulation_id,
181
+ "report_id": report_id,
182
+ "task_id": task_id,
183
+ "status": "generating",
184
+ "message": "报告生成任务已启动,请通过 /api/report/generate/status 查询进度",
185
+ "already_generated": False
186
+ }
187
+ })
188
+
189
+ except Exception as e:
190
+ logger.error(f"启动报告生成任务失败: {str(e)}")
191
+ return jsonify({
192
+ "success": False,
193
+ "error": str(e),
194
+ "traceback": traceback.format_exc()
195
+ }), 500
196
+
197
+
198
+ @report_bp.route('/generate/status', methods=['POST'])
199
+ def get_generate_status():
200
+ """
201
+ 查询报告生成任务进度
202
+
203
+ 请求(JSON):
204
+ {
205
+ "task_id": "task_xxxx", // 可选,generate返回的task_id
206
+ "simulation_id": "sim_xxxx" // 可选,模拟ID
207
+ }
208
+
209
+ 返回:
210
+ {
211
+ "success": true,
212
+ "data": {
213
+ "task_id": "task_xxxx",
214
+ "status": "processing|completed|failed",
215
+ "progress": 45,
216
+ "message": "..."
217
+ }
218
+ }
219
+ """
220
+ try:
221
+ data = request.get_json() or {}
222
+
223
+ task_id = data.get('task_id')
224
+ simulation_id = data.get('simulation_id')
225
+
226
+ # 如果提供了simulation_id,先检查是否已有完成的报告
227
+ if simulation_id:
228
+ existing_report = ReportManager.get_report_by_simulation(simulation_id)
229
+ if existing_report and existing_report.status == ReportStatus.COMPLETED:
230
+ return jsonify({
231
+ "success": True,
232
+ "data": {
233
+ "simulation_id": simulation_id,
234
+ "report_id": existing_report.report_id,
235
+ "status": "completed",
236
+ "progress": 100,
237
+ "message": "报告已生成",
238
+ "already_completed": True
239
+ }
240
+ })
241
+
242
+ if not task_id:
243
+ return jsonify({
244
+ "success": False,
245
+ "error": "请提供 task_id 或 simulation_id"
246
+ }), 400
247
+
248
+ task_manager = TaskManager()
249
+ task = task_manager.get_task(task_id)
250
+
251
+ if not task:
252
+ return jsonify({
253
+ "success": False,
254
+ "error": f"任务不存在: {task_id}"
255
+ }), 404
256
+
257
+ return jsonify({
258
+ "success": True,
259
+ "data": task.to_dict()
260
+ })
261
+
262
+ except Exception as e:
263
+ logger.error(f"查询任务状态失败: {str(e)}")
264
+ return jsonify({
265
+ "success": False,
266
+ "error": str(e)
267
+ }), 500
268
+
269
+
270
+ # ============== 报告获取接口 ==============
271
+
272
+ @report_bp.route('/<report_id>', methods=['GET'])
273
+ def get_report(report_id: str):
274
+ """
275
+ 获取报告详情
276
+
277
+ 返回:
278
+ {
279
+ "success": true,
280
+ "data": {
281
+ "report_id": "report_xxxx",
282
+ "simulation_id": "sim_xxxx",
283
+ "status": "completed",
284
+ "outline": {...},
285
+ "markdown_content": "...",
286
+ "created_at": "...",
287
+ "completed_at": "..."
288
+ }
289
+ }
290
+ """
291
+ try:
292
+ report = ReportManager.get_report(report_id)
293
+
294
+ if not report:
295
+ return jsonify({
296
+ "success": False,
297
+ "error": f"报告不存在: {report_id}"
298
+ }), 404
299
+
300
+ return jsonify({
301
+ "success": True,
302
+ "data": report.to_dict()
303
+ })
304
+
305
+ except Exception as e:
306
+ logger.error(f"获取报告失败: {str(e)}")
307
+ return jsonify({
308
+ "success": False,
309
+ "error": str(e),
310
+ "traceback": traceback.format_exc()
311
+ }), 500
312
+
313
+
314
+ @report_bp.route('/by-simulation/<simulation_id>', methods=['GET'])
315
+ def get_report_by_simulation(simulation_id: str):
316
+ """
317
+ 根据模拟ID获取报告
318
+
319
+ 返回:
320
+ {
321
+ "success": true,
322
+ "data": {
323
+ "report_id": "report_xxxx",
324
+ ...
325
+ }
326
+ }
327
+ """
328
+ try:
329
+ report = ReportManager.get_report_by_simulation(simulation_id)
330
+
331
+ if not report:
332
+ return jsonify({
333
+ "success": False,
334
+ "error": f"该模拟暂无报告: {simulation_id}",
335
+ "has_report": False
336
+ }), 404
337
+
338
+ return jsonify({
339
+ "success": True,
340
+ "data": report.to_dict(),
341
+ "has_report": True
342
+ })
343
+
344
+ except Exception as e:
345
+ logger.error(f"获取报告失败: {str(e)}")
346
+ return jsonify({
347
+ "success": False,
348
+ "error": str(e),
349
+ "traceback": traceback.format_exc()
350
+ }), 500
351
+
352
+
353
+ @report_bp.route('/list', methods=['GET'])
354
+ def list_reports():
355
+ """
356
+ 列出所有报告
357
+
358
+ Query参数:
359
+ simulation_id: 按模拟ID过滤(可选)
360
+ limit: 返回数量限制(默认50)
361
+
362
+ 返回:
363
+ {
364
+ "success": true,
365
+ "data": [...],
366
+ "count": 10
367
+ }
368
+ """
369
+ try:
370
+ simulation_id = request.args.get('simulation_id')
371
+ limit = request.args.get('limit', 50, type=int)
372
+
373
+ reports = ReportManager.list_reports(
374
+ simulation_id=simulation_id,
375
+ limit=limit
376
+ )
377
+
378
+ return jsonify({
379
+ "success": True,
380
+ "data": [r.to_dict() for r in reports],
381
+ "count": len(reports)
382
+ })
383
+
384
+ except Exception as e:
385
+ logger.error(f"列出报告失败: {str(e)}")
386
+ return jsonify({
387
+ "success": False,
388
+ "error": str(e),
389
+ "traceback": traceback.format_exc()
390
+ }), 500
391
+
392
+
393
+ @report_bp.route('/<report_id>/download', methods=['GET'])
394
+ def download_report(report_id: str):
395
+ """
396
+ 下载报告(Markdown格式)
397
+
398
+ 返回Markdown文件
399
+ """
400
+ try:
401
+ report = ReportManager.get_report(report_id)
402
+
403
+ if not report:
404
+ return jsonify({
405
+ "success": False,
406
+ "error": f"报告不存在: {report_id}"
407
+ }), 404
408
+
409
+ md_path = ReportManager._get_report_markdown_path(report_id)
410
+
411
+ if not os.path.exists(md_path):
412
+ # 如果MD文件不存在,生成一个临时文件
413
+ import tempfile
414
+ with tempfile.NamedTemporaryFile(mode='w', suffix='.md', delete=False) as f:
415
+ f.write(report.markdown_content)
416
+ temp_path = f.name
417
+
418
+ return send_file(
419
+ temp_path,
420
+ as_attachment=True,
421
+ download_name=f"{report_id}.md"
422
+ )
423
+
424
+ return send_file(
425
+ md_path,
426
+ as_attachment=True,
427
+ download_name=f"{report_id}.md"
428
+ )
429
+
430
+ except Exception as e:
431
+ logger.error(f"下载报告失败: {str(e)}")
432
+ return jsonify({
433
+ "success": False,
434
+ "error": str(e),
435
+ "traceback": traceback.format_exc()
436
+ }), 500
437
+
438
+
439
+ @report_bp.route('/<report_id>', methods=['DELETE'])
440
+ def delete_report(report_id: str):
441
+ """删除报告"""
442
+ try:
443
+ success = ReportManager.delete_report(report_id)
444
+
445
+ if not success:
446
+ return jsonify({
447
+ "success": False,
448
+ "error": f"报告不存在: {report_id}"
449
+ }), 404
450
+
451
+ return jsonify({
452
+ "success": True,
453
+ "message": f"报告已删除: {report_id}"
454
+ })
455
+
456
+ except Exception as e:
457
+ logger.error(f"删除报告失败: {str(e)}")
458
+ return jsonify({
459
+ "success": False,
460
+ "error": str(e),
461
+ "traceback": traceback.format_exc()
462
+ }), 500
463
+
464
+
465
+ # ============== Report Agent对话接口 ==============
466
+
467
+ @report_bp.route('/chat', methods=['POST'])
468
+ def chat_with_report_agent():
469
+ """
470
+ 与Report Agent对话
471
+
472
+ Report Agent可以在对话中自主调用检索工具来回答问题
473
+
474
+ 请求(JSON):
475
+ {
476
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
477
+ "message": "请解释一下舆情走向", // 必填,用户消息
478
+ "chat_history": [ // 可选,对话历史
479
+ {"role": "user", "content": "..."},
480
+ {"role": "assistant", "content": "..."}
481
+ ]
482
+ }
483
+
484
+ 返回:
485
+ {
486
+ "success": true,
487
+ "data": {
488
+ "response": "Agent回复...",
489
+ "tool_calls": [调用的工具列表],
490
+ "sources": [信息来源]
491
+ }
492
+ }
493
+ """
494
+ try:
495
+ data = request.get_json() or {}
496
+
497
+ simulation_id = data.get('simulation_id')
498
+ message = data.get('message')
499
+ chat_history = data.get('chat_history', [])
500
+
501
+ if not simulation_id:
502
+ return jsonify({
503
+ "success": False,
504
+ "error": "请提供 simulation_id"
505
+ }), 400
506
+
507
+ if not message:
508
+ return jsonify({
509
+ "success": False,
510
+ "error": "请提供 message"
511
+ }), 400
512
+
513
+ # 获取模拟和项目信息
514
+ manager = SimulationManager()
515
+ state = manager.get_simulation(simulation_id)
516
+
517
+ if not state:
518
+ return jsonify({
519
+ "success": False,
520
+ "error": f"模拟不存在: {simulation_id}"
521
+ }), 404
522
+
523
+ project = ProjectManager.get_project(state.project_id)
524
+ if not project:
525
+ return jsonify({
526
+ "success": False,
527
+ "error": f"项目不存在: {state.project_id}"
528
+ }), 404
529
+
530
+ graph_id = state.graph_id or project.graph_id
531
+ if not graph_id:
532
+ return jsonify({
533
+ "success": False,
534
+ "error": "缺少图谱ID"
535
+ }), 400
536
+
537
+ simulation_requirement = project.simulation_requirement or ""
538
+
539
+ # 创建Agent并进行对话
540
+ agent = ReportAgent(
541
+ graph_id=graph_id,
542
+ simulation_id=simulation_id,
543
+ simulation_requirement=simulation_requirement
544
+ )
545
+
546
+ result = agent.chat(message=message, chat_history=chat_history)
547
+
548
+ return jsonify({
549
+ "success": True,
550
+ "data": result
551
+ })
552
+
553
+ except Exception as e:
554
+ logger.error(f"对话失败: {str(e)}")
555
+ return jsonify({
556
+ "success": False,
557
+ "error": str(e),
558
+ "traceback": traceback.format_exc()
559
+ }), 500
560
+
561
+
562
+ # ============== 报告进度与分章节接口 ==============
563
+
564
+ @report_bp.route('/<report_id>/progress', methods=['GET'])
565
+ def get_report_progress(report_id: str):
566
+ """
567
+ 获取报告生成进度(实时)
568
+
569
+ 返回:
570
+ {
571
+ "success": true,
572
+ "data": {
573
+ "status": "generating",
574
+ "progress": 45,
575
+ "message": "正在生成章节: 关键发现",
576
+ "current_section": "关键发现",
577
+ "completed_sections": ["执行摘要", "模拟背景"],
578
+ "updated_at": "2025-12-09T..."
579
+ }
580
+ }
581
+ """
582
+ try:
583
+ progress = ReportManager.get_progress(report_id)
584
+
585
+ if not progress:
586
+ return jsonify({
587
+ "success": False,
588
+ "error": f"报告不存在或进度信息不可用: {report_id}"
589
+ }), 404
590
+
591
+ return jsonify({
592
+ "success": True,
593
+ "data": progress
594
+ })
595
+
596
+ except Exception as e:
597
+ logger.error(f"获取报告进度失败: {str(e)}")
598
+ return jsonify({
599
+ "success": False,
600
+ "error": str(e),
601
+ "traceback": traceback.format_exc()
602
+ }), 500
603
+
604
+
605
+ @report_bp.route('/<report_id>/sections', methods=['GET'])
606
+ def get_report_sections(report_id: str):
607
+ """
608
+ 获取已生成的章节列表(分章节输出)
609
+
610
+ 前端可以轮询此接口获取已生成的章节内容,无需等待整个报告完成
611
+
612
+ 返回:
613
+ {
614
+ "success": true,
615
+ "data": {
616
+ "report_id": "report_xxxx",
617
+ "sections": [
618
+ {
619
+ "filename": "section_01.md",
620
+ "section_index": 1,
621
+ "content": "## 执行摘要\\n\\n..."
622
+ },
623
+ ...
624
+ ],
625
+ "total_sections": 3,
626
+ "is_complete": false
627
+ }
628
+ }
629
+ """
630
+ try:
631
+ sections = ReportManager.get_generated_sections(report_id)
632
+
633
+ # 获取报告状态
634
+ report = ReportManager.get_report(report_id)
635
+ is_complete = report is not None and report.status == ReportStatus.COMPLETED
636
+
637
+ return jsonify({
638
+ "success": True,
639
+ "data": {
640
+ "report_id": report_id,
641
+ "sections": sections,
642
+ "total_sections": len(sections),
643
+ "is_complete": is_complete
644
+ }
645
+ })
646
+
647
+ except Exception as e:
648
+ logger.error(f"获取章节列表失败: {str(e)}")
649
+ return jsonify({
650
+ "success": False,
651
+ "error": str(e),
652
+ "traceback": traceback.format_exc()
653
+ }), 500
654
+
655
+
656
+ @report_bp.route('/<report_id>/section/<int:section_index>', methods=['GET'])
657
+ def get_single_section(report_id: str, section_index: int):
658
+ """
659
+ 获取单个章节内容
660
+
661
+ 返回:
662
+ {
663
+ "success": true,
664
+ "data": {
665
+ "filename": "section_01.md",
666
+ "content": "## 执行摘要\\n\\n..."
667
+ }
668
+ }
669
+ """
670
+ try:
671
+ section_path = ReportManager._get_section_path(report_id, section_index)
672
+
673
+ if not os.path.exists(section_path):
674
+ return jsonify({
675
+ "success": False,
676
+ "error": f"章节不存在: section_{section_index:02d}.md"
677
+ }), 404
678
+
679
+ with open(section_path, 'r', encoding='utf-8') as f:
680
+ content = f.read()
681
+
682
+ return jsonify({
683
+ "success": True,
684
+ "data": {
685
+ "filename": f"section_{section_index:02d}.md",
686
+ "section_index": section_index,
687
+ "content": content
688
+ }
689
+ })
690
+
691
+ except Exception as e:
692
+ logger.error(f"获取章节内容失败: {str(e)}")
693
+ return jsonify({
694
+ "success": False,
695
+ "error": str(e),
696
+ "traceback": traceback.format_exc()
697
+ }), 500
698
+
699
+
700
+ # ============== 报告状态检查接口 ==============
701
+
702
+ @report_bp.route('/check/<simulation_id>', methods=['GET'])
703
+ def check_report_status(simulation_id: str):
704
+ """
705
+ 检查模拟是否有报告,以及报告状态
706
+
707
+ 用于前端判断是否解锁Interview功能
708
+
709
+ 返回:
710
+ {
711
+ "success": true,
712
+ "data": {
713
+ "simulation_id": "sim_xxxx",
714
+ "has_report": true,
715
+ "report_status": "completed",
716
+ "report_id": "report_xxxx",
717
+ "interview_unlocked": true
718
+ }
719
+ }
720
+ """
721
+ try:
722
+ report = ReportManager.get_report_by_simulation(simulation_id)
723
+
724
+ has_report = report is not None
725
+ report_status = report.status.value if report else None
726
+ report_id = report.report_id if report else None
727
+
728
+ # 只有报告完成后才解锁interview
729
+ interview_unlocked = has_report and report.status == ReportStatus.COMPLETED
730
+
731
+ return jsonify({
732
+ "success": True,
733
+ "data": {
734
+ "simulation_id": simulation_id,
735
+ "has_report": has_report,
736
+ "report_status": report_status,
737
+ "report_id": report_id,
738
+ "interview_unlocked": interview_unlocked
739
+ }
740
+ })
741
+
742
+ except Exception as e:
743
+ logger.error(f"检查报告状态失败: {str(e)}")
744
+ return jsonify({
745
+ "success": False,
746
+ "error": str(e),
747
+ "traceback": traceback.format_exc()
748
+ }), 500
749
+
750
+
751
+ # ============== Agent 日志接口 ==============
752
+
753
+ @report_bp.route('/<report_id>/agent-log', methods=['GET'])
754
+ def get_agent_log(report_id: str):
755
+ """
756
+ 获取 Report Agent 的详细执行日志
757
+
758
+ 实时获取报告生成过程中的每一步动作,包括:
759
+ - 报告开始、规划开始/完成
760
+ - 每个章节的开始、工具调用、LLM响应、完成
761
+ - 报告完成或失败
762
+
763
+ Query参数:
764
+ from_line: 从第几行开始读取(可选,默认0,用于增量获取)
765
+
766
+ 返回:
767
+ {
768
+ "success": true,
769
+ "data": {
770
+ "logs": [
771
+ {
772
+ "timestamp": "2025-12-13T...",
773
+ "elapsed_seconds": 12.5,
774
+ "report_id": "report_xxxx",
775
+ "action": "tool_call",
776
+ "stage": "generating",
777
+ "section_title": "执行摘要",
778
+ "section_index": 1,
779
+ "details": {
780
+ "tool_name": "insight_forge",
781
+ "parameters": {...},
782
+ ...
783
+ }
784
+ },
785
+ ...
786
+ ],
787
+ "total_lines": 25,
788
+ "from_line": 0,
789
+ "has_more": false
790
+ }
791
+ }
792
+ """
793
+ try:
794
+ from_line = request.args.get('from_line', 0, type=int)
795
+
796
+ log_data = ReportManager.get_agent_log(report_id, from_line=from_line)
797
+
798
+ return jsonify({
799
+ "success": True,
800
+ "data": log_data
801
+ })
802
+
803
+ except Exception as e:
804
+ logger.error(f"获取Agent日志失败: {str(e)}")
805
+ return jsonify({
806
+ "success": False,
807
+ "error": str(e),
808
+ "traceback": traceback.format_exc()
809
+ }), 500
810
+
811
+
812
+ @report_bp.route('/<report_id>/agent-log/stream', methods=['GET'])
813
+ def stream_agent_log(report_id: str):
814
+ """
815
+ 获取完整的 Agent 日志(一次性获取全部)
816
+
817
+ 返回:
818
+ {
819
+ "success": true,
820
+ "data": {
821
+ "logs": [...],
822
+ "count": 25
823
+ }
824
+ }
825
+ """
826
+ try:
827
+ logs = ReportManager.get_agent_log_stream(report_id)
828
+
829
+ return jsonify({
830
+ "success": True,
831
+ "data": {
832
+ "logs": logs,
833
+ "count": len(logs)
834
+ }
835
+ })
836
+
837
+ except Exception as e:
838
+ logger.error(f"获取Agent日志失败: {str(e)}")
839
+ return jsonify({
840
+ "success": False,
841
+ "error": str(e),
842
+ "traceback": traceback.format_exc()
843
+ }), 500
844
+
845
+
846
+ # ============== 控制台日志接口 ==============
847
+
848
+ @report_bp.route('/<report_id>/console-log', methods=['GET'])
849
+ def get_console_log(report_id: str):
850
+ """
851
+ 获取 Report Agent 的控制台输出日志
852
+
853
+ 实时获取报告生成过程中的控制台输出(INFO、WARNING等),
854
+ 这与 agent-log 接口返回的结构化 JSON 日志不同,
855
+ 是纯文本格式的控制台风格日志。
856
+
857
+ Query参数:
858
+ from_line: 从第几行开始读取(可选,默认0,用于增量获取)
859
+
860
+ 返回:
861
+ {
862
+ "success": true,
863
+ "data": {
864
+ "logs": [
865
+ "[19:46:14] INFO: 搜索完成: 找到 15 条相关事实",
866
+ "[19:46:14] INFO: 图谱搜索: graph_id=xxx, query=...",
867
+ ...
868
+ ],
869
+ "total_lines": 100,
870
+ "from_line": 0,
871
+ "has_more": false
872
+ }
873
+ }
874
+ """
875
+ try:
876
+ from_line = request.args.get('from_line', 0, type=int)
877
+
878
+ log_data = ReportManager.get_console_log(report_id, from_line=from_line)
879
+
880
+ return jsonify({
881
+ "success": True,
882
+ "data": log_data
883
+ })
884
+
885
+ except Exception as e:
886
+ logger.error(f"获取控制台日志失败: {str(e)}")
887
+ return jsonify({
888
+ "success": False,
889
+ "error": str(e),
890
+ "traceback": traceback.format_exc()
891
+ }), 500
892
+
893
+
894
+ @report_bp.route('/<report_id>/console-log/stream', methods=['GET'])
895
+ def stream_console_log(report_id: str):
896
+ """
897
+ 获取完整的控制台日志(一次性获取全部)
898
+
899
+ 返回:
900
+ {
901
+ "success": true,
902
+ "data": {
903
+ "logs": [...],
904
+ "count": 100
905
+ }
906
+ }
907
+ """
908
+ try:
909
+ logs = ReportManager.get_console_log_stream(report_id)
910
+
911
+ return jsonify({
912
+ "success": True,
913
+ "data": {
914
+ "logs": logs,
915
+ "count": len(logs)
916
+ }
917
+ })
918
+
919
+ except Exception as e:
920
+ logger.error(f"获取控制台日志失败: {str(e)}")
921
+ return jsonify({
922
+ "success": False,
923
+ "error": str(e),
924
+ "traceback": traceback.format_exc()
925
+ }), 500
926
+
927
+
928
+ # ============== 工具调用接口(供调试使用)==============
929
+
930
+ @report_bp.route('/tools/search', methods=['POST'])
931
+ def search_graph_tool():
932
+ """
933
+ 图谱搜索工具接口(供调试使用)
934
+
935
+ 请求(JSON):
936
+ {
937
+ "graph_id": "mirofish_xxxx",
938
+ "query": "搜索查询",
939
+ "limit": 10
940
+ }
941
+ """
942
+ try:
943
+ data = request.get_json() or {}
944
+
945
+ graph_id = data.get('graph_id')
946
+ query = data.get('query')
947
+ limit = data.get('limit', 10)
948
+
949
+ if not graph_id or not query:
950
+ return jsonify({
951
+ "success": False,
952
+ "error": "请提供 graph_id 和 query"
953
+ }), 400
954
+
955
+ from ..services.zep_tools import ZepToolsService
956
+
957
+ tools = ZepToolsService()
958
+ result = tools.search_graph(
959
+ graph_id=graph_id,
960
+ query=query,
961
+ limit=limit
962
+ )
963
+
964
+ return jsonify({
965
+ "success": True,
966
+ "data": result.to_dict()
967
+ })
968
+
969
+ except Exception as e:
970
+ logger.error(f"图谱搜索失败: {str(e)}")
971
+ return jsonify({
972
+ "success": False,
973
+ "error": str(e),
974
+ "traceback": traceback.format_exc()
975
+ }), 500
976
+
977
+
978
+ @report_bp.route('/tools/statistics', methods=['POST'])
979
+ def get_graph_statistics_tool():
980
+ """
981
+ 图谱统计工具接口(供调试使用)
982
+
983
+ 请求(JSON):
984
+ {
985
+ "graph_id": "mirofish_xxxx"
986
+ }
987
+ """
988
+ try:
989
+ data = request.get_json() or {}
990
+
991
+ graph_id = data.get('graph_id')
992
+
993
+ if not graph_id:
994
+ return jsonify({
995
+ "success": False,
996
+ "error": "请提供 graph_id"
997
+ }), 400
998
+
999
+ from ..services.zep_tools import ZepToolsService
1000
+
1001
+ tools = ZepToolsService()
1002
+ result = tools.get_graph_statistics(graph_id)
1003
+
1004
+ return jsonify({
1005
+ "success": True,
1006
+ "data": result
1007
+ })
1008
+
1009
+ except Exception as e:
1010
+ logger.error(f"获取图谱统计失败: {str(e)}")
1011
+ return jsonify({
1012
+ "success": False,
1013
+ "error": str(e),
1014
+ "traceback": traceback.format_exc()
1015
+ }), 500
app/app/api/simulation.py ADDED
@@ -0,0 +1,2718 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 模拟相关API路由
3
+ Step2: Zep实体读取与过滤、OASIS模拟准备与运行(全程自动化)
4
+ """
5
+
6
+ import os
7
+ import traceback
8
+ from flask import request, jsonify, send_file
9
+
10
+ from . import simulation_bp
11
+ from ..config import Config
12
+ from ..services.zep_entity_reader import ZepEntityReader
13
+ from ..services.oasis_profile_generator import OasisProfileGenerator
14
+ from ..services.simulation_manager import SimulationManager, SimulationStatus
15
+ from ..services.simulation_runner import SimulationRunner, RunnerStatus
16
+ from ..utils.logger import get_logger
17
+ from ..models.project import ProjectManager
18
+
19
+ logger = get_logger('mirofish.api.simulation')
20
+
21
+
22
+ # Interview prompt 优化前缀
23
+ # 添加此前缀可以避免Agent调用工具,直接用文本回复。
24
+ # Startup validation: agents stay in persona (investor, founder, customer, skeptic, etc.).
25
+ INTERVIEW_PROMPT_PREFIX = (
26
+ "Stay in character using your persona, memories, and past actions. Do not call any tools—reply in plain text only. "
27
+ "You are in a startup idea validation simulation: respond in character only "
28
+ "(e.g. VC, angel, YC-style partner, picky investor, peer founder, buyer, undecided user). "
29
+ "Ground answers in your persona and the simulated discussion, not generic advice.\n\n"
30
+ "Question: "
31
+ )
32
+
33
+
34
+ def optimize_interview_prompt(prompt: str) -> str:
35
+ """
36
+ 优化Interview提问,添加前缀避免Agent调用工具
37
+
38
+ Args:
39
+ prompt: 原始提问
40
+
41
+ Returns:
42
+ 优化后的提问
43
+ """
44
+ if not prompt:
45
+ return prompt
46
+ # 避免重复添加前缀
47
+ if prompt.startswith(INTERVIEW_PROMPT_PREFIX):
48
+ return prompt
49
+ return f"{INTERVIEW_PROMPT_PREFIX}{prompt}"
50
+
51
+
52
+ # ============== 实体读取接口 ==============
53
+
54
+ @simulation_bp.route('/entities/<graph_id>', methods=['GET'])
55
+ def get_graph_entities(graph_id: str):
56
+ """
57
+ 获取图谱中的所有实体(已过滤)
58
+
59
+ 只返回符合预定义实体类型的节点(Labels不只是Entity的节点)
60
+
61
+ Query参数:
62
+ entity_types: 逗号分隔的实体类型列表(可选,用于进一步过滤)
63
+ enrich: 是否获取相关边信息(默认true)
64
+ """
65
+ try:
66
+ if not Config.ZEP_API_KEY:
67
+ return jsonify({
68
+ "success": False,
69
+ "error": "ZEP_API_KEY未配置"
70
+ }), 500
71
+
72
+ entity_types_str = request.args.get('entity_types', '')
73
+ entity_types = [t.strip() for t in entity_types_str.split(',') if t.strip()] if entity_types_str else None
74
+ enrich = request.args.get('enrich', 'true').lower() == 'true'
75
+
76
+ logger.info(f"获取图谱实体: graph_id={graph_id}, entity_types={entity_types}, enrich={enrich}")
77
+
78
+ reader = ZepEntityReader()
79
+ result = reader.filter_defined_entities(
80
+ graph_id=graph_id,
81
+ defined_entity_types=entity_types,
82
+ enrich_with_edges=enrich
83
+ )
84
+
85
+ return jsonify({
86
+ "success": True,
87
+ "data": result.to_dict()
88
+ })
89
+
90
+ except Exception as e:
91
+ logger.error(f"获取图谱实体失败: {str(e)}")
92
+ return jsonify({
93
+ "success": False,
94
+ "error": str(e),
95
+ "traceback": traceback.format_exc()
96
+ }), 500
97
+
98
+
99
+ @simulation_bp.route('/entities/<graph_id>/<entity_uuid>', methods=['GET'])
100
+ def get_entity_detail(graph_id: str, entity_uuid: str):
101
+ """获取单个实体的详细信息"""
102
+ try:
103
+ if not Config.ZEP_API_KEY:
104
+ return jsonify({
105
+ "success": False,
106
+ "error": "ZEP_API_KEY未配置"
107
+ }), 500
108
+
109
+ reader = ZepEntityReader()
110
+ entity = reader.get_entity_with_context(graph_id, entity_uuid)
111
+
112
+ if not entity:
113
+ return jsonify({
114
+ "success": False,
115
+ "error": f"实体不存在: {entity_uuid}"
116
+ }), 404
117
+
118
+ return jsonify({
119
+ "success": True,
120
+ "data": entity.to_dict()
121
+ })
122
+
123
+ except Exception as e:
124
+ logger.error(f"获取实体详情失败: {str(e)}")
125
+ return jsonify({
126
+ "success": False,
127
+ "error": str(e),
128
+ "traceback": traceback.format_exc()
129
+ }), 500
130
+
131
+
132
+ @simulation_bp.route('/entities/<graph_id>/by-type/<entity_type>', methods=['GET'])
133
+ def get_entities_by_type(graph_id: str, entity_type: str):
134
+ """获取指定类型的所有实体"""
135
+ try:
136
+ if not Config.ZEP_API_KEY:
137
+ return jsonify({
138
+ "success": False,
139
+ "error": "ZEP_API_KEY未配置"
140
+ }), 500
141
+
142
+ enrich = request.args.get('enrich', 'true').lower() == 'true'
143
+
144
+ reader = ZepEntityReader()
145
+ entities = reader.get_entities_by_type(
146
+ graph_id=graph_id,
147
+ entity_type=entity_type,
148
+ enrich_with_edges=enrich
149
+ )
150
+
151
+ return jsonify({
152
+ "success": True,
153
+ "data": {
154
+ "entity_type": entity_type,
155
+ "count": len(entities),
156
+ "entities": [e.to_dict() for e in entities]
157
+ }
158
+ })
159
+
160
+ except Exception as e:
161
+ logger.error(f"获取实体失败: {str(e)}")
162
+ return jsonify({
163
+ "success": False,
164
+ "error": str(e),
165
+ "traceback": traceback.format_exc()
166
+ }), 500
167
+
168
+
169
+ # ============== 模拟管理接口 ==============
170
+
171
+ @simulation_bp.route('/create', methods=['POST'])
172
+ def create_simulation():
173
+ """
174
+ 创建新的模拟
175
+
176
+ 注意:max_rounds等参数由LLM智能生成,无需手动设置
177
+
178
+ 请求(JSON):
179
+ {
180
+ "project_id": "proj_xxxx", // 必填
181
+ "graph_id": "mirofish_xxxx", // 可选,如不提供则从project获取
182
+ "enable_twitter": true, // 可选,默认true
183
+ "enable_reddit": true // 可选,默认true
184
+ }
185
+
186
+ 返回:
187
+ {
188
+ "success": true,
189
+ "data": {
190
+ "simulation_id": "sim_xxxx",
191
+ "project_id": "proj_xxxx",
192
+ "graph_id": "mirofish_xxxx",
193
+ "status": "created",
194
+ "enable_twitter": true,
195
+ "enable_reddit": true,
196
+ "created_at": "2025-12-01T10:00:00"
197
+ }
198
+ }
199
+ """
200
+ try:
201
+ data = request.get_json() or {}
202
+
203
+ project_id = data.get('project_id')
204
+ if not project_id:
205
+ return jsonify({
206
+ "success": False,
207
+ "error": "请提供 project_id"
208
+ }), 400
209
+
210
+ project = ProjectManager.get_project(project_id)
211
+ if not project:
212
+ return jsonify({
213
+ "success": False,
214
+ "error": f"项目不存在: {project_id}"
215
+ }), 404
216
+
217
+ graph_id = data.get('graph_id') or project.graph_id
218
+ if not graph_id:
219
+ return jsonify({
220
+ "success": False,
221
+ "error": "项目尚未构建图谱,请先调用 /api/graph/build"
222
+ }), 400
223
+
224
+ manager = SimulationManager()
225
+ state = manager.create_simulation(
226
+ project_id=project_id,
227
+ graph_id=graph_id,
228
+ enable_twitter=data.get('enable_twitter', True),
229
+ enable_reddit=data.get('enable_reddit', True),
230
+ )
231
+
232
+ return jsonify({
233
+ "success": True,
234
+ "data": state.to_dict()
235
+ })
236
+
237
+ except Exception as e:
238
+ logger.error(f"创建模拟失败: {str(e)}")
239
+ return jsonify({
240
+ "success": False,
241
+ "error": str(e),
242
+ "traceback": traceback.format_exc()
243
+ }), 500
244
+
245
+
246
+ def _check_simulation_prepared(simulation_id: str) -> tuple:
247
+ """
248
+ 检查模拟是否已经准备完成
249
+
250
+ 检查条件:
251
+ 1. state.json 存在且 status 为 "ready"
252
+ 2. 必要文件存在:reddit_profiles.json, twitter_profiles.csv, simulation_config.json
253
+
254
+ 注意:运行脚本(run_*.py)保留在 backend/scripts/ 目录,不再复制到模拟目录
255
+
256
+ Args:
257
+ simulation_id: 模拟ID
258
+
259
+ Returns:
260
+ (is_prepared: bool, info: dict)
261
+ """
262
+ import os
263
+ from ..config import Config
264
+
265
+ simulation_dir = os.path.join(Config.OASIS_SIMULATION_DATA_DIR, simulation_id)
266
+
267
+ # 检查目录是否存在
268
+ if not os.path.exists(simulation_dir):
269
+ return False, {"reason": "模拟目录不存在"}
270
+
271
+ # 必要文件列表(不包括脚本,脚本位于 backend/scripts/)
272
+ required_files = [
273
+ "state.json",
274
+ "simulation_config.json",
275
+ "reddit_profiles.json",
276
+ "twitter_profiles.csv"
277
+ ]
278
+
279
+ # 检查文件是否存在
280
+ existing_files = []
281
+ missing_files = []
282
+ for f in required_files:
283
+ file_path = os.path.join(simulation_dir, f)
284
+ if os.path.exists(file_path):
285
+ existing_files.append(f)
286
+ else:
287
+ missing_files.append(f)
288
+
289
+ if missing_files:
290
+ return False, {
291
+ "reason": "缺少必要文件",
292
+ "missing_files": missing_files,
293
+ "existing_files": existing_files
294
+ }
295
+
296
+ # 检查state.json中的状态
297
+ state_file = os.path.join(simulation_dir, "state.json")
298
+ try:
299
+ import json
300
+ with open(state_file, 'r', encoding='utf-8') as f:
301
+ state_data = json.load(f)
302
+
303
+ status = state_data.get("status", "")
304
+ config_generated = state_data.get("config_generated", False)
305
+
306
+ # 详细日志
307
+ logger.debug(f"检测模拟准备状态: {simulation_id}, status={status}, config_generated={config_generated}")
308
+
309
+ # 如果 config_generated=True 且文件存在,认为准备完成
310
+ # 以下���态都说明准备工作已完成:
311
+ # - ready: 准备完成,可以运行
312
+ # - preparing: 如果 config_generated=True 说明已完成
313
+ # - running: 正在运行,说明准备早就完成了
314
+ # - completed: 运行完成,说明准备早就完成了
315
+ # - stopped: 已停止,说明准备早就完成了
316
+ # - failed: 运行失败(但准备是完成的)
317
+ prepared_statuses = ["ready", "preparing", "running", "completed", "stopped", "failed"]
318
+ if status in prepared_statuses and config_generated:
319
+ # 获取文件统计信息
320
+ profiles_file = os.path.join(simulation_dir, "reddit_profiles.json")
321
+ config_file = os.path.join(simulation_dir, "simulation_config.json")
322
+
323
+ profiles_count = 0
324
+ if os.path.exists(profiles_file):
325
+ with open(profiles_file, 'r', encoding='utf-8') as f:
326
+ profiles_data = json.load(f)
327
+ profiles_count = len(profiles_data) if isinstance(profiles_data, list) else 0
328
+
329
+ # 如果状态是preparing但文件已完成,自动更新状态为ready
330
+ if status == "preparing":
331
+ try:
332
+ state_data["status"] = "ready"
333
+ from datetime import datetime
334
+ state_data["updated_at"] = datetime.now().isoformat()
335
+ with open(state_file, 'w', encoding='utf-8') as f:
336
+ json.dump(state_data, f, ensure_ascii=False, indent=2)
337
+ logger.info(f"自动更新模拟状态: {simulation_id} preparing -> ready")
338
+ status = "ready"
339
+ except Exception as e:
340
+ logger.warning(f"自动更新状态失败: {e}")
341
+
342
+ logger.info(f"模拟 {simulation_id} 检测结果: 已准备完成 (status={status}, config_generated={config_generated})")
343
+ return True, {
344
+ "status": status,
345
+ "entities_count": state_data.get("entities_count", 0),
346
+ "profiles_count": profiles_count,
347
+ "entity_types": state_data.get("entity_types", []),
348
+ "config_generated": config_generated,
349
+ "created_at": state_data.get("created_at"),
350
+ "updated_at": state_data.get("updated_at"),
351
+ "existing_files": existing_files
352
+ }
353
+ else:
354
+ logger.warning(f"模拟 {simulation_id} 检测结果: 未准备完成 (status={status}, config_generated={config_generated})")
355
+ return False, {
356
+ "reason": f"状态不在已准备列表中或config_generated为false: status={status}, config_generated={config_generated}",
357
+ "status": status,
358
+ "config_generated": config_generated
359
+ }
360
+
361
+ except Exception as e:
362
+ return False, {"reason": f"读取状态文件失败: {str(e)}"}
363
+
364
+
365
+ @simulation_bp.route('/prepare', methods=['POST'])
366
+ def prepare_simulation():
367
+ """
368
+ 准备模拟环境(异步任务,LLM智能生成所有参数)
369
+
370
+ 这是一个耗时操作,接口会立即返回task_id,
371
+ 使用 GET /api/simulation/prepare/status 查询进度
372
+
373
+ 特性:
374
+ - 自动检测已完成的准备工作,避免重复生成
375
+ - 如果已准备完成,直接返回已有结果
376
+ - 支持强制重新生成(force_regenerate=true)
377
+
378
+ 步骤:
379
+ 1. 检查是否已有完成的准备工作
380
+ 2. 从Zep图谱读取并过滤实体
381
+ 3. 为每个实体生成OASIS Agent Profile(带重试机制)
382
+ 4. LLM智能生成模拟配置(带重试机制)
383
+ 5. 保存配置文件和预设脚本
384
+
385
+ 请求(JSON):
386
+ {
387
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
388
+ "entity_types": ["Student", "PublicFigure"], // 可选,指定实体类型
389
+ "use_llm_for_profiles": true, // 可选,是否用LLM生成人设
390
+ "parallel_profile_count": 5, // 可选,并行生成人设数量,默认5
391
+ "force_regenerate": false // 可选,强制重新生成,默认false
392
+ }
393
+
394
+ 返回:
395
+ {
396
+ "success": true,
397
+ "data": {
398
+ "simulation_id": "sim_xxxx",
399
+ "task_id": "task_xxxx", // 新任务时返回
400
+ "status": "preparing|ready",
401
+ "message": "准备任务已启动|已有完成的准备工作",
402
+ "already_prepared": true|false // 是否已准备完成
403
+ }
404
+ }
405
+ """
406
+ import threading
407
+ import os
408
+ from ..models.task import TaskManager, TaskStatus
409
+ from ..config import Config
410
+
411
+ try:
412
+ data = request.get_json() or {}
413
+
414
+ simulation_id = data.get('simulation_id')
415
+ if not simulation_id:
416
+ return jsonify({
417
+ "success": False,
418
+ "error": "请提供 simulation_id"
419
+ }), 400
420
+
421
+ manager = SimulationManager()
422
+ state = manager.get_simulation(simulation_id)
423
+
424
+ if not state:
425
+ return jsonify({
426
+ "success": False,
427
+ "error": f"模拟不存在: {simulation_id}"
428
+ }), 404
429
+
430
+ # 检查是否强制重新生成
431
+ force_regenerate = data.get('force_regenerate', False)
432
+ logger.info(f"开始处理 /prepare 请求: simulation_id={simulation_id}, force_regenerate={force_regenerate}")
433
+
434
+ # 检查是否已经准备完成(避免重复生成)
435
+ if not force_regenerate:
436
+ logger.debug(f"检查模拟 {simulation_id} 是否已准备完成...")
437
+ is_prepared, prepare_info = _check_simulation_prepared(simulation_id)
438
+ logger.debug(f"检查结果: is_prepared={is_prepared}, prepare_info={prepare_info}")
439
+ if is_prepared:
440
+ logger.info(f"模拟 {simulation_id} 已准备完成,跳过重复生成")
441
+ return jsonify({
442
+ "success": True,
443
+ "data": {
444
+ "simulation_id": simulation_id,
445
+ "status": "ready",
446
+ "message": "已有完成的准备工作,无需重复生成",
447
+ "already_prepared": True,
448
+ "prepare_info": prepare_info
449
+ }
450
+ })
451
+ else:
452
+ logger.info(f"模拟 {simulation_id} 未准备完成,将启动准备任务")
453
+
454
+ # 从项目获取必要信息
455
+ project = ProjectManager.get_project(state.project_id)
456
+ if not project:
457
+ return jsonify({
458
+ "success": False,
459
+ "error": f"项目不存在: {state.project_id}"
460
+ }), 404
461
+
462
+ # 获取模拟需求
463
+ simulation_requirement = project.simulation_requirement or ""
464
+ if not simulation_requirement:
465
+ return jsonify({
466
+ "success": False,
467
+ "error": "项目缺少模拟需求描述 (simulation_requirement)"
468
+ }), 400
469
+
470
+ # 获取文档文本
471
+ document_text = ProjectManager.get_extracted_text(state.project_id) or ""
472
+
473
+ entity_types_list = data.get('entity_types')
474
+ use_llm_for_profiles = data.get('use_llm_for_profiles', True)
475
+ parallel_profile_count = data.get('parallel_profile_count', 5)
476
+
477
+ # ========== 同步获取实体数量(在后台任务启动前) ==========
478
+ # 这样前端在调用prepare后立即就能获取到预期Agent总数
479
+ try:
480
+ logger.info(f"同步获取实体数量: graph_id={state.graph_id}")
481
+ reader = ZepEntityReader()
482
+ # 快速读取实体(不需要边信息,只统计数量)
483
+ filtered_preview = reader.filter_defined_entities(
484
+ graph_id=state.graph_id,
485
+ defined_entity_types=entity_types_list,
486
+ enrich_with_edges=False # 不获取边信息,加快速度
487
+ )
488
+ # 保存实体数量到状态(供前端立即获取)
489
+ state.entities_count = filtered_preview.filtered_count
490
+ state.entity_types = list(filtered_preview.entity_types)
491
+ logger.info(f"预期实体数量: {filtered_preview.filtered_count}, 类型: {filtered_preview.entity_types}")
492
+ except Exception as e:
493
+ logger.warning(f"同步获取实体数量失败(将在后台任务中重试): {e}")
494
+ # 失败不影响后续流程,后台任务会重新获取
495
+
496
+ # 创建异步任务
497
+ task_manager = TaskManager()
498
+ task_id = task_manager.create_task(
499
+ task_type="simulation_prepare",
500
+ metadata={
501
+ "simulation_id": simulation_id,
502
+ "project_id": state.project_id
503
+ }
504
+ )
505
+
506
+ # 更新模拟状态(包含预先获取的实体数量)
507
+ state.status = SimulationStatus.PREPARING
508
+ manager._save_simulation_state(state)
509
+
510
+ # 定义后台任务
511
+ def run_prepare():
512
+ try:
513
+ task_manager.update_task(
514
+ task_id,
515
+ status=TaskStatus.PROCESSING,
516
+ progress=0,
517
+ message="开始准备模拟环境..."
518
+ )
519
+
520
+ # 准备模拟(带进度回调)
521
+ # 存储阶段进度详情
522
+ stage_details = {}
523
+
524
+ def progress_callback(stage, progress, message, **kwargs):
525
+ # 计算总进度
526
+ stage_weights = {
527
+ "reading": (0, 20), # 0-20%
528
+ "generating_profiles": (20, 70), # 20-70%
529
+ "generating_config": (70, 90), # 70-90%
530
+ "copying_scripts": (90, 100) # 90-100%
531
+ }
532
+
533
+ start, end = stage_weights.get(stage, (0, 100))
534
+ current_progress = int(start + (end - start) * progress / 100)
535
+
536
+ # 构建详细进度信息
537
+ stage_names = {
538
+ "reading": "读取图谱实体",
539
+ "generating_profiles": "生成Agent人设",
540
+ "generating_config": "生成模拟配置",
541
+ "copying_scripts": "准备模拟脚本"
542
+ }
543
+
544
+ stage_index = list(stage_weights.keys()).index(stage) + 1 if stage in stage_weights else 1
545
+ total_stages = len(stage_weights)
546
+
547
+ # 更新阶段详情
548
+ stage_details[stage] = {
549
+ "stage_name": stage_names.get(stage, stage),
550
+ "stage_progress": progress,
551
+ "current": kwargs.get("current", 0),
552
+ "total": kwargs.get("total", 0),
553
+ "item_name": kwargs.get("item_name", "")
554
+ }
555
+
556
+ # 构建详细进度信息
557
+ detail = stage_details[stage]
558
+ progress_detail_data = {
559
+ "current_stage": stage,
560
+ "current_stage_name": stage_names.get(stage, stage),
561
+ "stage_index": stage_index,
562
+ "total_stages": total_stages,
563
+ "stage_progress": progress,
564
+ "current_item": detail["current"],
565
+ "total_items": detail["total"],
566
+ "item_description": message
567
+ }
568
+
569
+ # 构建简洁消息
570
+ if detail["total"] > 0:
571
+ detailed_message = (
572
+ f"[{stage_index}/{total_stages}] {stage_names.get(stage, stage)}: "
573
+ f"{detail['current']}/{detail['total']} - {message}"
574
+ )
575
+ else:
576
+ detailed_message = f"[{stage_index}/{total_stages}] {stage_names.get(stage, stage)}: {message}"
577
+
578
+ task_manager.update_task(
579
+ task_id,
580
+ progress=current_progress,
581
+ message=detailed_message,
582
+ progress_detail=progress_detail_data
583
+ )
584
+
585
+ result_state = manager.prepare_simulation(
586
+ simulation_id=simulation_id,
587
+ simulation_requirement=simulation_requirement,
588
+ document_text=document_text,
589
+ defined_entity_types=entity_types_list,
590
+ use_llm_for_profiles=use_llm_for_profiles,
591
+ progress_callback=progress_callback,
592
+ parallel_profile_count=parallel_profile_count
593
+ )
594
+
595
+ # 任务完成
596
+ task_manager.complete_task(
597
+ task_id,
598
+ result=result_state.to_simple_dict()
599
+ )
600
+
601
+ except Exception as e:
602
+ logger.error(f"准备模拟失败: {str(e)}")
603
+ task_manager.fail_task(task_id, str(e))
604
+
605
+ # 更新模拟状态为失败
606
+ state = manager.get_simulation(simulation_id)
607
+ if state:
608
+ state.status = SimulationStatus.FAILED
609
+ state.error = str(e)
610
+ manager._save_simulation_state(state)
611
+
612
+ # 启动后台线程
613
+ thread = threading.Thread(target=run_prepare, daemon=True)
614
+ thread.start()
615
+
616
+ return jsonify({
617
+ "success": True,
618
+ "data": {
619
+ "simulation_id": simulation_id,
620
+ "task_id": task_id,
621
+ "status": "preparing",
622
+ "message": "准备任务已启动,请通过 /api/simulation/prepare/status 查询进度",
623
+ "already_prepared": False,
624
+ "expected_entities_count": state.entities_count, # 预期的Agent总数
625
+ "entity_types": state.entity_types # 实体类型列表
626
+ }
627
+ })
628
+
629
+ except ValueError as e:
630
+ return jsonify({
631
+ "success": False,
632
+ "error": str(e)
633
+ }), 404
634
+
635
+ except Exception as e:
636
+ logger.error(f"启动准备任务失败: {str(e)}")
637
+ return jsonify({
638
+ "success": False,
639
+ "error": str(e),
640
+ "traceback": traceback.format_exc()
641
+ }), 500
642
+
643
+
644
+ @simulation_bp.route('/prepare/status', methods=['POST'])
645
+ def get_prepare_status():
646
+ """
647
+ 查询准备��务进度
648
+
649
+ 支持两种查询方式:
650
+ 1. 通过task_id查询正在进行的任务进度
651
+ 2. 通过simulation_id检查是否已有完成的准备工作
652
+
653
+ 请求(JSON):
654
+ {
655
+ "task_id": "task_xxxx", // 可选,prepare返回的task_id
656
+ "simulation_id": "sim_xxxx" // 可选,模拟ID(用于检查已完成的准备)
657
+ }
658
+
659
+ 返回:
660
+ {
661
+ "success": true,
662
+ "data": {
663
+ "task_id": "task_xxxx",
664
+ "status": "processing|completed|ready",
665
+ "progress": 45,
666
+ "message": "...",
667
+ "already_prepared": true|false, // 是否已有完成的准备
668
+ "prepare_info": {...} // 已准备完成时的详细信息
669
+ }
670
+ }
671
+ """
672
+ from ..models.task import TaskManager
673
+
674
+ try:
675
+ data = request.get_json() or {}
676
+
677
+ task_id = data.get('task_id')
678
+ simulation_id = data.get('simulation_id')
679
+
680
+ # 如果提供了simulation_id,先检查是否已准备完成
681
+ if simulation_id:
682
+ is_prepared, prepare_info = _check_simulation_prepared(simulation_id)
683
+ if is_prepared:
684
+ return jsonify({
685
+ "success": True,
686
+ "data": {
687
+ "simulation_id": simulation_id,
688
+ "status": "ready",
689
+ "progress": 100,
690
+ "message": "已有完成的准备工作",
691
+ "already_prepared": True,
692
+ "prepare_info": prepare_info
693
+ }
694
+ })
695
+
696
+ # 如果没有task_id,返回错误
697
+ if not task_id:
698
+ if simulation_id:
699
+ # 有simulation_id但未准备完成
700
+ return jsonify({
701
+ "success": True,
702
+ "data": {
703
+ "simulation_id": simulation_id,
704
+ "status": "not_started",
705
+ "progress": 0,
706
+ "message": "尚未开始准备,请调用 /api/simulation/prepare 开始",
707
+ "already_prepared": False
708
+ }
709
+ })
710
+ return jsonify({
711
+ "success": False,
712
+ "error": "请提供 task_id 或 simulation_id"
713
+ }), 400
714
+
715
+ task_manager = TaskManager()
716
+ task = task_manager.get_task(task_id)
717
+
718
+ if not task:
719
+ # 任务不存在,但如果有simulation_id,检查是否已准备完成
720
+ if simulation_id:
721
+ is_prepared, prepare_info = _check_simulation_prepared(simulation_id)
722
+ if is_prepared:
723
+ return jsonify({
724
+ "success": True,
725
+ "data": {
726
+ "simulation_id": simulation_id,
727
+ "task_id": task_id,
728
+ "status": "ready",
729
+ "progress": 100,
730
+ "message": "任务已完成(准备工作已存在)",
731
+ "already_prepared": True,
732
+ "prepare_info": prepare_info
733
+ }
734
+ })
735
+
736
+ return jsonify({
737
+ "success": False,
738
+ "error": f"任务不存在: {task_id}"
739
+ }), 404
740
+
741
+ task_dict = task.to_dict()
742
+ task_dict["already_prepared"] = False
743
+
744
+ return jsonify({
745
+ "success": True,
746
+ "data": task_dict
747
+ })
748
+
749
+ except Exception as e:
750
+ logger.error(f"查询任务状态失败: {str(e)}")
751
+ return jsonify({
752
+ "success": False,
753
+ "error": str(e)
754
+ }), 500
755
+
756
+
757
+ @simulation_bp.route('/<simulation_id>', methods=['GET'])
758
+ def get_simulation(simulation_id: str):
759
+ """获取模拟状态"""
760
+ try:
761
+ manager = SimulationManager()
762
+ state = manager.get_simulation(simulation_id)
763
+
764
+ if not state:
765
+ return jsonify({
766
+ "success": False,
767
+ "error": f"模拟不存在: {simulation_id}"
768
+ }), 404
769
+
770
+ result = state.to_dict()
771
+
772
+ # 如果模拟已准备好,附加运行说明
773
+ if state.status == SimulationStatus.READY:
774
+ result["run_instructions"] = manager.get_run_instructions(simulation_id)
775
+
776
+ return jsonify({
777
+ "success": True,
778
+ "data": result
779
+ })
780
+
781
+ except Exception as e:
782
+ logger.error(f"获取模拟状态失败: {str(e)}")
783
+ return jsonify({
784
+ "success": False,
785
+ "error": str(e),
786
+ "traceback": traceback.format_exc()
787
+ }), 500
788
+
789
+
790
+ @simulation_bp.route('/list', methods=['GET'])
791
+ def list_simulations():
792
+ """
793
+ 列出所有模拟
794
+
795
+ Query参数:
796
+ project_id: 按项目ID过滤(可选)
797
+ """
798
+ try:
799
+ project_id = request.args.get('project_id')
800
+
801
+ manager = SimulationManager()
802
+ simulations = manager.list_simulations(project_id=project_id)
803
+
804
+ return jsonify({
805
+ "success": True,
806
+ "data": [s.to_dict() for s in simulations],
807
+ "count": len(simulations)
808
+ })
809
+
810
+ except Exception as e:
811
+ logger.error(f"列出模拟失败: {str(e)}")
812
+ return jsonify({
813
+ "success": False,
814
+ "error": str(e),
815
+ "traceback": traceback.format_exc()
816
+ }), 500
817
+
818
+
819
+ def _get_report_id_for_simulation(simulation_id: str) -> str:
820
+ """
821
+ 获取 simulation 对应的最新 report_id
822
+
823
+ 遍历 reports 目录,找出 simulation_id 匹配的 report,
824
+ 如果有多个则返回最新的(按 created_at 排序)
825
+
826
+ Args:
827
+ simulation_id: 模拟ID
828
+
829
+ Returns:
830
+ report_id 或 None
831
+ """
832
+ import json
833
+ from datetime import datetime
834
+
835
+ # reports 目录路径:backend/uploads/reports
836
+ # __file__ 是 app/api/simulation.py,需要向上两级到 backend/
837
+ reports_dir = os.path.join(os.path.dirname(__file__), '../../uploads/reports')
838
+ if not os.path.exists(reports_dir):
839
+ return None
840
+
841
+ matching_reports = []
842
+
843
+ try:
844
+ for report_folder in os.listdir(reports_dir):
845
+ report_path = os.path.join(reports_dir, report_folder)
846
+ if not os.path.isdir(report_path):
847
+ continue
848
+
849
+ meta_file = os.path.join(report_path, "meta.json")
850
+ if not os.path.exists(meta_file):
851
+ continue
852
+
853
+ try:
854
+ with open(meta_file, 'r', encoding='utf-8') as f:
855
+ meta = json.load(f)
856
+
857
+ if meta.get("simulation_id") == simulation_id:
858
+ matching_reports.append({
859
+ "report_id": meta.get("report_id"),
860
+ "created_at": meta.get("created_at", ""),
861
+ "status": meta.get("status", "")
862
+ })
863
+ except Exception:
864
+ continue
865
+
866
+ if not matching_reports:
867
+ return None
868
+
869
+ # 按创建时间倒序排序,返回最新的
870
+ matching_reports.sort(key=lambda x: x.get("created_at", ""), reverse=True)
871
+ return matching_reports[0].get("report_id")
872
+
873
+ except Exception as e:
874
+ logger.warning(f"查找 simulation {simulation_id} 的 report 失败: {e}")
875
+ return None
876
+
877
+
878
+ @simulation_bp.route('/history', methods=['GET'])
879
+ def get_simulation_history():
880
+ """
881
+ 获取历史模拟列表(带项目详情)
882
+
883
+ 用于首页历史项目展示,返回包含项目名称、描述等丰富信息的模拟列表
884
+
885
+ Query参数:
886
+ limit: 返回数量限制(默认20)
887
+
888
+ 返回:
889
+ {
890
+ "success": true,
891
+ "data": [
892
+ {
893
+ "simulation_id": "sim_xxxx",
894
+ "project_id": "proj_xxxx",
895
+ "project_name": "武大舆情分析",
896
+ "simulation_requirement": "如果武汉大学发布...",
897
+ "status": "completed",
898
+ "entities_count": 68,
899
+ "profiles_count": 68,
900
+ "entity_types": ["Student", "Professor", ...],
901
+ "created_at": "2024-12-10",
902
+ "updated_at": "2024-12-10",
903
+ "total_rounds": 120,
904
+ "current_round": 120,
905
+ "report_id": "report_xxxx",
906
+ "version": "v1.0.2"
907
+ },
908
+ ...
909
+ ],
910
+ "count": 7
911
+ }
912
+ """
913
+ try:
914
+ limit = request.args.get('limit', 20, type=int)
915
+
916
+ manager = SimulationManager()
917
+ simulations = manager.list_simulations()[:limit]
918
+
919
+ # 增强模拟数据,只从 Simulation 文件读取
920
+ enriched_simulations = []
921
+ for sim in simulations:
922
+ sim_dict = sim.to_dict()
923
+
924
+ # 获取模拟配置信息(从 simulation_config.json 读取 simulation_requirement)
925
+ config = manager.get_simulation_config(sim.simulation_id)
926
+ if config:
927
+ sim_dict["simulation_requirement"] = config.get("simulation_requirement", "")
928
+ time_config = config.get("time_config", {})
929
+ sim_dict["total_simulation_hours"] = time_config.get("total_simulation_hours", 0)
930
+ # 推荐轮数(后备值)
931
+ recommended_rounds = int(
932
+ time_config.get("total_simulation_hours", 0) * 60 /
933
+ max(time_config.get("minutes_per_round", 60), 1)
934
+ )
935
+ else:
936
+ sim_dict["simulation_requirement"] = ""
937
+ sim_dict["total_simulation_hours"] = 0
938
+ recommended_rounds = 0
939
+
940
+ # 获取运行状态(从 run_state.json 读取用户设置的实际轮数)
941
+ run_state = SimulationRunner.get_run_state(sim.simulation_id)
942
+ if run_state:
943
+ sim_dict["current_round"] = run_state.current_round
944
+ sim_dict["runner_status"] = run_state.runner_status.value
945
+ # 使用用户设置的 total_rounds,若无则使用推荐轮数
946
+ sim_dict["total_rounds"] = run_state.total_rounds if run_state.total_rounds > 0 else recommended_rounds
947
+ else:
948
+ sim_dict["current_round"] = 0
949
+ sim_dict["runner_status"] = "idle"
950
+ sim_dict["total_rounds"] = recommended_rounds
951
+
952
+ # 获取关联项目的文件列表(最多3个)
953
+ project = ProjectManager.get_project(sim.project_id)
954
+ if project and hasattr(project, 'files') and project.files:
955
+ sim_dict["files"] = [
956
+ {"filename": f.get("filename", "未知文件")}
957
+ for f in project.files[:3]
958
+ ]
959
+ else:
960
+ sim_dict["files"] = []
961
+
962
+ # 获取关联的 report_id(查找该 simulation 最新的 report)
963
+ sim_dict["report_id"] = _get_report_id_for_simulation(sim.simulation_id)
964
+
965
+ # 添加版本号
966
+ sim_dict["version"] = "v1.0.2"
967
+
968
+ # 格式化日期
969
+ try:
970
+ created_date = sim_dict.get("created_at", "")[:10]
971
+ sim_dict["created_date"] = created_date
972
+ except:
973
+ sim_dict["created_date"] = ""
974
+
975
+ enriched_simulations.append(sim_dict)
976
+
977
+ return jsonify({
978
+ "success": True,
979
+ "data": enriched_simulations,
980
+ "count": len(enriched_simulations)
981
+ })
982
+
983
+ except Exception as e:
984
+ logger.error(f"获取历史模拟失败: {str(e)}")
985
+ return jsonify({
986
+ "success": False,
987
+ "error": str(e),
988
+ "traceback": traceback.format_exc()
989
+ }), 500
990
+
991
+
992
+ @simulation_bp.route('/<simulation_id>/profiles', methods=['GET'])
993
+ def get_simulation_profiles(simulation_id: str):
994
+ """
995
+ 获取模拟的Agent Profile
996
+
997
+ Query参数:
998
+ platform: 平台类型(reddit/twitter,默认reddit)
999
+ """
1000
+ try:
1001
+ platform = request.args.get('platform', 'reddit')
1002
+
1003
+ manager = SimulationManager()
1004
+ profiles = manager.get_profiles(simulation_id, platform=platform)
1005
+
1006
+ return jsonify({
1007
+ "success": True,
1008
+ "data": {
1009
+ "platform": platform,
1010
+ "count": len(profiles),
1011
+ "profiles": profiles
1012
+ }
1013
+ })
1014
+
1015
+ except ValueError as e:
1016
+ return jsonify({
1017
+ "success": False,
1018
+ "error": str(e)
1019
+ }), 404
1020
+
1021
+ except Exception as e:
1022
+ logger.error(f"获取Profile失败: {str(e)}")
1023
+ return jsonify({
1024
+ "success": False,
1025
+ "error": str(e),
1026
+ "traceback": traceback.format_exc()
1027
+ }), 500
1028
+
1029
+
1030
+ @simulation_bp.route('/<simulation_id>/profiles/realtime', methods=['GET'])
1031
+ def get_simulation_profiles_realtime(simulation_id: str):
1032
+ """
1033
+ 实时获取模拟的Agent Profile(用于在生成过程中实时查看进度)
1034
+
1035
+ 与 /profiles 接口的区别:
1036
+ - 直接读取文件,不经过 SimulationManager
1037
+ - 适用于生成过程中的实时查看
1038
+ - 返回额外的元数据(如文件修改时间、是否正在生成等)
1039
+
1040
+ Query参数:
1041
+ platform: 平台类型(reddit/twitter,默认reddit)
1042
+
1043
+ 返回:
1044
+ {
1045
+ "success": true,
1046
+ "data": {
1047
+ "simulation_id": "sim_xxxx",
1048
+ "platform": "reddit",
1049
+ "count": 15,
1050
+ "total_expected": 93, // 预期总数(如果有)
1051
+ "is_generating": true, // 是否正在生成
1052
+ "file_exists": true,
1053
+ "file_modified_at": "2025-12-04T18:20:00",
1054
+ "profiles": [...]
1055
+ }
1056
+ }
1057
+ """
1058
+ import json
1059
+ import csv
1060
+ from datetime import datetime
1061
+
1062
+ try:
1063
+ platform = request.args.get('platform', 'reddit')
1064
+
1065
+ # 获取模拟目录
1066
+ sim_dir = os.path.join(Config.OASIS_SIMULATION_DATA_DIR, simulation_id)
1067
+
1068
+ if not os.path.exists(sim_dir):
1069
+ return jsonify({
1070
+ "success": False,
1071
+ "error": f"模拟不存在: {simulation_id}"
1072
+ }), 404
1073
+
1074
+ # 确定文件路径
1075
+ if platform == "reddit":
1076
+ profiles_file = os.path.join(sim_dir, "reddit_profiles.json")
1077
+ else:
1078
+ profiles_file = os.path.join(sim_dir, "twitter_profiles.csv")
1079
+
1080
+ # 检查文件是否存在
1081
+ file_exists = os.path.exists(profiles_file)
1082
+ profiles = []
1083
+ file_modified_at = None
1084
+
1085
+ if file_exists:
1086
+ # 获取文件修改时间
1087
+ file_stat = os.stat(profiles_file)
1088
+ file_modified_at = datetime.fromtimestamp(file_stat.st_mtime).isoformat()
1089
+
1090
+ try:
1091
+ if platform == "reddit":
1092
+ with open(profiles_file, 'r', encoding='utf-8') as f:
1093
+ profiles = json.load(f)
1094
+ else:
1095
+ with open(profiles_file, 'r', encoding='utf-8') as f:
1096
+ reader = csv.DictReader(f)
1097
+ profiles = list(reader)
1098
+ except (json.JSONDecodeError, Exception) as e:
1099
+ logger.warning(f"读取 profiles 文件失败(可能正在写入中): {e}")
1100
+ profiles = []
1101
+
1102
+ # 检查是否正在生成(通过 state.json 判断)
1103
+ is_generating = False
1104
+ total_expected = None
1105
+
1106
+ state_file = os.path.join(sim_dir, "state.json")
1107
+ if os.path.exists(state_file):
1108
+ try:
1109
+ with open(state_file, 'r', encoding='utf-8') as f:
1110
+ state_data = json.load(f)
1111
+ status = state_data.get("status", "")
1112
+ is_generating = status == "preparing"
1113
+ total_expected = state_data.get("entities_count")
1114
+ except Exception:
1115
+ pass
1116
+
1117
+ return jsonify({
1118
+ "success": True,
1119
+ "data": {
1120
+ "simulation_id": simulation_id,
1121
+ "platform": platform,
1122
+ "count": len(profiles),
1123
+ "total_expected": total_expected,
1124
+ "is_generating": is_generating,
1125
+ "file_exists": file_exists,
1126
+ "file_modified_at": file_modified_at,
1127
+ "profiles": profiles
1128
+ }
1129
+ })
1130
+
1131
+ except Exception as e:
1132
+ logger.error(f"实时获取Profile失败: {str(e)}")
1133
+ return jsonify({
1134
+ "success": False,
1135
+ "error": str(e),
1136
+ "traceback": traceback.format_exc()
1137
+ }), 500
1138
+
1139
+
1140
+ @simulation_bp.route('/<simulation_id>/config/realtime', methods=['GET'])
1141
+ def get_simulation_config_realtime(simulation_id: str):
1142
+ """
1143
+ 实时获取模拟配置(用于在生成过程中实时查看进度)
1144
+
1145
+ 与 /config 接口的区别:
1146
+ - 直接读取文件,不经过 SimulationManager
1147
+ - 适用于生成过程中的实时查看
1148
+ - 返回额外的元数据(如文件修改时间、是否正在生成等)
1149
+ - 即使配置还没生成完也能返回部分信息
1150
+
1151
+ 返回:
1152
+ {
1153
+ "success": true,
1154
+ "data": {
1155
+ "simulation_id": "sim_xxxx",
1156
+ "file_exists": true,
1157
+ "file_modified_at": "2025-12-04T18:20:00",
1158
+ "is_generating": true, // 是否正在生成
1159
+ "generation_stage": "generating_config", // 当前生成阶段
1160
+ "config": {...} // 配置内容(如果存在)
1161
+ }
1162
+ }
1163
+ """
1164
+ import json
1165
+ from datetime import datetime
1166
+
1167
+ try:
1168
+ # 获取模拟目录
1169
+ sim_dir = os.path.join(Config.OASIS_SIMULATION_DATA_DIR, simulation_id)
1170
+
1171
+ if not os.path.exists(sim_dir):
1172
+ return jsonify({
1173
+ "success": False,
1174
+ "error": f"模拟不存在: {simulation_id}"
1175
+ }), 404
1176
+
1177
+ # 配置文件路径
1178
+ config_file = os.path.join(sim_dir, "simulation_config.json")
1179
+
1180
+ # 检查文件是否存在
1181
+ file_exists = os.path.exists(config_file)
1182
+ config = None
1183
+ file_modified_at = None
1184
+
1185
+ if file_exists:
1186
+ # 获取文件修改时间
1187
+ file_stat = os.stat(config_file)
1188
+ file_modified_at = datetime.fromtimestamp(file_stat.st_mtime).isoformat()
1189
+
1190
+ try:
1191
+ with open(config_file, 'r', encoding='utf-8') as f:
1192
+ config = json.load(f)
1193
+ except (json.JSONDecodeError, Exception) as e:
1194
+ logger.warning(f"读取 config 文件失败(可能正在写入中): {e}")
1195
+ config = None
1196
+
1197
+ # 检查是否正在生成(通过 state.json 判断)
1198
+ is_generating = False
1199
+ generation_stage = None
1200
+ config_generated = False
1201
+
1202
+ state_file = os.path.join(sim_dir, "state.json")
1203
+ if os.path.exists(state_file):
1204
+ try:
1205
+ with open(state_file, 'r', encoding='utf-8') as f:
1206
+ state_data = json.load(f)
1207
+ status = state_data.get("status", "")
1208
+ is_generating = status == "preparing"
1209
+ config_generated = state_data.get("config_generated", False)
1210
+
1211
+ # 判断当前阶段
1212
+ if is_generating:
1213
+ if state_data.get("profiles_generated", False):
1214
+ generation_stage = "generating_config"
1215
+ else:
1216
+ generation_stage = "generating_profiles"
1217
+ elif status == "ready":
1218
+ generation_stage = "completed"
1219
+ except Exception:
1220
+ pass
1221
+
1222
+ # 构建返回数据
1223
+ response_data = {
1224
+ "simulation_id": simulation_id,
1225
+ "file_exists": file_exists,
1226
+ "file_modified_at": file_modified_at,
1227
+ "is_generating": is_generating,
1228
+ "generation_stage": generation_stage,
1229
+ "config_generated": config_generated,
1230
+ "config": config
1231
+ }
1232
+
1233
+ # 如果配置存在,提取一些关键统计信息
1234
+ if config:
1235
+ response_data["summary"] = {
1236
+ "total_agents": len(config.get("agent_configs", [])),
1237
+ "simulation_hours": config.get("time_config", {}).get("total_simulation_hours"),
1238
+ "initial_posts_count": len(config.get("event_config", {}).get("initial_posts", [])),
1239
+ "hot_topics_count": len(config.get("event_config", {}).get("hot_topics", [])),
1240
+ "has_twitter_config": "twitter_config" in config,
1241
+ "has_reddit_config": "reddit_config" in config,
1242
+ "generated_at": config.get("generated_at"),
1243
+ "llm_model": config.get("llm_model")
1244
+ }
1245
+
1246
+ return jsonify({
1247
+ "success": True,
1248
+ "data": response_data
1249
+ })
1250
+
1251
+ except Exception as e:
1252
+ logger.error(f"实时获取Config失败: {str(e)}")
1253
+ return jsonify({
1254
+ "success": False,
1255
+ "error": str(e),
1256
+ "traceback": traceback.format_exc()
1257
+ }), 500
1258
+
1259
+
1260
+ @simulation_bp.route('/<simulation_id>/config', methods=['GET'])
1261
+ def get_simulation_config(simulation_id: str):
1262
+ """
1263
+ 获取模拟配置(LLM智能生成的完整配置)
1264
+
1265
+ 返回包含:
1266
+ - time_config: 时间配置(模拟时长、轮次、高峰/低谷时段)
1267
+ - agent_configs: 每个Agent的活动配置(活跃度、发言频率、立场等)
1268
+ - event_config: 事件配置(初始帖子、热点话题)
1269
+ - platform_configs: 平台配置
1270
+ - generation_reasoning: LLM的配置推理说明
1271
+ """
1272
+ try:
1273
+ manager = SimulationManager()
1274
+ config = manager.get_simulation_config(simulation_id)
1275
+
1276
+ if not config:
1277
+ return jsonify({
1278
+ "success": False,
1279
+ "error": f"模拟配置不存在,请先调用 /prepare 接口"
1280
+ }), 404
1281
+
1282
+ return jsonify({
1283
+ "success": True,
1284
+ "data": config
1285
+ })
1286
+
1287
+ except Exception as e:
1288
+ logger.error(f"获取配置失败: {str(e)}")
1289
+ return jsonify({
1290
+ "success": False,
1291
+ "error": str(e),
1292
+ "traceback": traceback.format_exc()
1293
+ }), 500
1294
+
1295
+
1296
+ @simulation_bp.route('/<simulation_id>/config/download', methods=['GET'])
1297
+ def download_simulation_config(simulation_id: str):
1298
+ """下载模拟配置文件"""
1299
+ try:
1300
+ manager = SimulationManager()
1301
+ sim_dir = manager._get_simulation_dir(simulation_id)
1302
+ config_path = os.path.join(sim_dir, "simulation_config.json")
1303
+
1304
+ if not os.path.exists(config_path):
1305
+ return jsonify({
1306
+ "success": False,
1307
+ "error": "配置文件不存在,请先调用 /prepare 接口"
1308
+ }), 404
1309
+
1310
+ return send_file(
1311
+ config_path,
1312
+ as_attachment=True,
1313
+ download_name="simulation_config.json"
1314
+ )
1315
+
1316
+ except Exception as e:
1317
+ logger.error(f"下载配置失败: {str(e)}")
1318
+ return jsonify({
1319
+ "success": False,
1320
+ "error": str(e),
1321
+ "traceback": traceback.format_exc()
1322
+ }), 500
1323
+
1324
+
1325
+ @simulation_bp.route('/script/<script_name>/download', methods=['GET'])
1326
+ def download_simulation_script(script_name: str):
1327
+ """
1328
+ 下载模拟运行脚本文件(通用脚本,位于 backend/scripts/)
1329
+
1330
+ script_name可选值:
1331
+ - run_twitter_simulation.py
1332
+ - run_reddit_simulation.py
1333
+ - run_parallel_simulation.py
1334
+ - action_logger.py
1335
+ """
1336
+ try:
1337
+ # 脚本位于 backend/scripts/ 目录
1338
+ scripts_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), '../../scripts'))
1339
+
1340
+ # 验证脚本名称
1341
+ allowed_scripts = [
1342
+ "run_twitter_simulation.py",
1343
+ "run_reddit_simulation.py",
1344
+ "run_parallel_simulation.py",
1345
+ "action_logger.py"
1346
+ ]
1347
+
1348
+ if script_name not in allowed_scripts:
1349
+ return jsonify({
1350
+ "success": False,
1351
+ "error": f"未知脚本: {script_name},可选: {allowed_scripts}"
1352
+ }), 400
1353
+
1354
+ script_path = os.path.join(scripts_dir, script_name)
1355
+
1356
+ if not os.path.exists(script_path):
1357
+ return jsonify({
1358
+ "success": False,
1359
+ "error": f"脚本文件不存在: {script_name}"
1360
+ }), 404
1361
+
1362
+ return send_file(
1363
+ script_path,
1364
+ as_attachment=True,
1365
+ download_name=script_name
1366
+ )
1367
+
1368
+ except Exception as e:
1369
+ logger.error(f"下载脚本失败: {str(e)}")
1370
+ return jsonify({
1371
+ "success": False,
1372
+ "error": str(e),
1373
+ "traceback": traceback.format_exc()
1374
+ }), 500
1375
+
1376
+
1377
+ # ============== Profile生成接口(独立使用) ==============
1378
+
1379
+ @simulation_bp.route('/generate-profiles', methods=['POST'])
1380
+ def generate_profiles():
1381
+ """
1382
+ 直接从图谱生成OASIS Agent Profile(不创建模拟)
1383
+
1384
+ 请求(JSON):
1385
+ {
1386
+ "graph_id": "mirofish_xxxx", // 必填
1387
+ "entity_types": ["Student"], // 可选
1388
+ "use_llm": true, // 可选
1389
+ "platform": "reddit" // 可选
1390
+ }
1391
+ """
1392
+ try:
1393
+ data = request.get_json() or {}
1394
+
1395
+ graph_id = data.get('graph_id')
1396
+ if not graph_id:
1397
+ return jsonify({
1398
+ "success": False,
1399
+ "error": "请提供 graph_id"
1400
+ }), 400
1401
+
1402
+ entity_types = data.get('entity_types')
1403
+ use_llm = data.get('use_llm', True)
1404
+ platform = data.get('platform', 'reddit')
1405
+
1406
+ reader = ZepEntityReader()
1407
+ filtered = reader.filter_defined_entities(
1408
+ graph_id=graph_id,
1409
+ defined_entity_types=entity_types,
1410
+ enrich_with_edges=True
1411
+ )
1412
+
1413
+ if filtered.filtered_count == 0:
1414
+ return jsonify({
1415
+ "success": False,
1416
+ "error": "没有找到符合条件的实体"
1417
+ }), 400
1418
+
1419
+ generator = OasisProfileGenerator()
1420
+ profiles = generator.generate_profiles_from_entities(
1421
+ entities=filtered.entities,
1422
+ use_llm=use_llm
1423
+ )
1424
+
1425
+ if platform == "reddit":
1426
+ profiles_data = [p.to_reddit_format() for p in profiles]
1427
+ elif platform == "twitter":
1428
+ profiles_data = [p.to_twitter_format() for p in profiles]
1429
+ else:
1430
+ profiles_data = [p.to_dict() for p in profiles]
1431
+
1432
+ return jsonify({
1433
+ "success": True,
1434
+ "data": {
1435
+ "platform": platform,
1436
+ "entity_types": list(filtered.entity_types),
1437
+ "count": len(profiles_data),
1438
+ "profiles": profiles_data
1439
+ }
1440
+ })
1441
+
1442
+ except Exception as e:
1443
+ logger.error(f"生成Profile失败: {str(e)}")
1444
+ return jsonify({
1445
+ "success": False,
1446
+ "error": str(e),
1447
+ "traceback": traceback.format_exc()
1448
+ }), 500
1449
+
1450
+
1451
+ # ============== 模拟运行控制接口 ==============
1452
+
1453
+ @simulation_bp.route('/start', methods=['POST'])
1454
+ def start_simulation():
1455
+ """
1456
+ 开始运行模拟
1457
+
1458
+ 请求(JSON):
1459
+ {
1460
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
1461
+ "platform": "parallel", // 可选: twitter / reddit / parallel (默认)
1462
+ "max_rounds": 100, // 可选: 最大模拟轮数,用于截断过长的模拟
1463
+ "enable_graph_memory_update": false, // 可选: 是否将Agent活动动态更新到Zep图谱记忆
1464
+ "force": false // 可选: 强制重新开始(会停止运行中的模拟并清理日志)
1465
+ }
1466
+
1467
+ 关于 force 参数:
1468
+ - 启用后,如果模拟正在运行或已完成,会先停止并清理运行日志
1469
+ - 清理的内容包括:run_state.json, actions.jsonl, simulation.log 等
1470
+ - 不会清理配置文件(simulation_config.json)和 profile 文件
1471
+ - 适用于需要重新运行模拟的场景
1472
+
1473
+ 关于 enable_graph_memory_update:
1474
+ - 启用后,模拟中所有Agent的活动(发帖、评论、点赞等)都会实时更新到Zep图谱
1475
+ - 这可以让图谱"记住"模拟过程,用于后续分析或AI对话
1476
+ - 需要模拟关联的项目有有效的 graph_id
1477
+ - 采用批量更新机制,减少API调用次数
1478
+
1479
+ 返回:
1480
+ {
1481
+ "success": true,
1482
+ "data": {
1483
+ "simulation_id": "sim_xxxx",
1484
+ "runner_status": "running",
1485
+ "process_pid": 12345,
1486
+ "twitter_running": true,
1487
+ "reddit_running": true,
1488
+ "started_at": "2025-12-01T10:00:00",
1489
+ "graph_memory_update_enabled": true, // 是否启用了图谱记忆更新
1490
+ "force_restarted": true // 是否是强制重新开始
1491
+ }
1492
+ }
1493
+ """
1494
+ try:
1495
+ data = request.get_json() or {}
1496
+
1497
+ simulation_id = data.get('simulation_id')
1498
+ if not simulation_id:
1499
+ return jsonify({
1500
+ "success": False,
1501
+ "error": "请提供 simulation_id"
1502
+ }), 400
1503
+
1504
+ platform = data.get('platform', 'parallel')
1505
+ max_rounds = data.get('max_rounds') # 可选:最大模拟轮数
1506
+ enable_graph_memory_update = data.get('enable_graph_memory_update', False) # 可选:是否启用图谱记忆更新
1507
+ force = data.get('force', False) # 可选:强制重新开始
1508
+
1509
+ # 验证 max_rounds 参数
1510
+ if max_rounds is not None:
1511
+ try:
1512
+ max_rounds = int(max_rounds)
1513
+ if max_rounds <= 0:
1514
+ return jsonify({
1515
+ "success": False,
1516
+ "error": "max_rounds 必须是正整数"
1517
+ }), 400
1518
+ except (ValueError, TypeError):
1519
+ return jsonify({
1520
+ "success": False,
1521
+ "error": "max_rounds 必须是有效的整数"
1522
+ }), 400
1523
+
1524
+ if platform not in ['twitter', 'reddit', 'parallel']:
1525
+ return jsonify({
1526
+ "success": False,
1527
+ "error": f"无效的平台类型: {platform},可选: twitter/reddit/parallel"
1528
+ }), 400
1529
+
1530
+ # 检查模拟是否已准备好
1531
+ manager = SimulationManager()
1532
+ state = manager.get_simulation(simulation_id)
1533
+
1534
+ if not state:
1535
+ return jsonify({
1536
+ "success": False,
1537
+ "error": f"模拟不存在: {simulation_id}"
1538
+ }), 404
1539
+
1540
+ force_restarted = False
1541
+
1542
+ # 智能处理状态:如果准备工作已完成,允许重新启动
1543
+ if state.status != SimulationStatus.READY:
1544
+ # 检查准备工作是否已完成
1545
+ is_prepared, prepare_info = _check_simulation_prepared(simulation_id)
1546
+
1547
+ if is_prepared:
1548
+ # 准备工作已完成,检查是否有正在运行的进程
1549
+ if state.status == SimulationStatus.RUNNING:
1550
+ # 检查模拟进程是否真的在运行
1551
+ run_state = SimulationRunner.get_run_state(simulation_id)
1552
+ if run_state and run_state.runner_status.value == "running":
1553
+ # 进程确实在运行
1554
+ if force:
1555
+ # 强制模式:停止运行中的模拟
1556
+ logger.info(f"强制模式:停止运行中的模拟 {simulation_id}")
1557
+ try:
1558
+ SimulationRunner.stop_simulation(simulation_id)
1559
+ except Exception as e:
1560
+ logger.warning(f"停止模拟时出现警告: {str(e)}")
1561
+ else:
1562
+ return jsonify({
1563
+ "success": False,
1564
+ "error": f"模拟正在运行中,请先调用 /stop 接口停止,或使用 force=true 强制重新开始"
1565
+ }), 400
1566
+
1567
+ # 如果是强制模式,清理运行日志
1568
+ if force:
1569
+ logger.info(f"强制模式:清理模拟日志 {simulation_id}")
1570
+ cleanup_result = SimulationRunner.cleanup_simulation_logs(simulation_id)
1571
+ if not cleanup_result.get("success"):
1572
+ logger.warning(f"清理日志时出现警告: {cleanup_result.get('errors')}")
1573
+ force_restarted = True
1574
+
1575
+ # 进程不存在或已结束,重置状态为 ready
1576
+ logger.info(f"模拟 {simulation_id} 准备工作已完成,重置状态为 ready(原状态: {state.status.value})")
1577
+ state.status = SimulationStatus.READY
1578
+ manager._save_simulation_state(state)
1579
+ else:
1580
+ # 准备工作未完成
1581
+ return jsonify({
1582
+ "success": False,
1583
+ "error": f"模拟未准备好,当前状态: {state.status.value},请先调用 /prepare 接口"
1584
+ }), 400
1585
+
1586
+ # 获取图谱ID(用于图谱记忆更新)
1587
+ graph_id = None
1588
+ if enable_graph_memory_update:
1589
+ # 从模拟状态或项目中获取 graph_id
1590
+ graph_id = state.graph_id
1591
+ if not graph_id:
1592
+ # 尝试从项目中获取
1593
+ project = ProjectManager.get_project(state.project_id)
1594
+ if project:
1595
+ graph_id = project.graph_id
1596
+
1597
+ if not graph_id:
1598
+ return jsonify({
1599
+ "success": False,
1600
+ "error": "启用图谱记忆更新需要有效的 graph_id,请确保项目已构建图谱"
1601
+ }), 400
1602
+
1603
+ logger.info(f"启用图谱记忆更新: simulation_id={simulation_id}, graph_id={graph_id}")
1604
+
1605
+ # 启动模拟
1606
+ run_state = SimulationRunner.start_simulation(
1607
+ simulation_id=simulation_id,
1608
+ platform=platform,
1609
+ max_rounds=max_rounds,
1610
+ enable_graph_memory_update=enable_graph_memory_update,
1611
+ graph_id=graph_id
1612
+ )
1613
+
1614
+ # 更新模拟状态
1615
+ state.status = SimulationStatus.RUNNING
1616
+ manager._save_simulation_state(state)
1617
+
1618
+ response_data = run_state.to_dict()
1619
+ if max_rounds:
1620
+ response_data['max_rounds_applied'] = max_rounds
1621
+ response_data['graph_memory_update_enabled'] = enable_graph_memory_update
1622
+ response_data['force_restarted'] = force_restarted
1623
+ if enable_graph_memory_update:
1624
+ response_data['graph_id'] = graph_id
1625
+
1626
+ return jsonify({
1627
+ "success": True,
1628
+ "data": response_data
1629
+ })
1630
+
1631
+ except ValueError as e:
1632
+ return jsonify({
1633
+ "success": False,
1634
+ "error": str(e)
1635
+ }), 400
1636
+
1637
+ except Exception as e:
1638
+ logger.error(f"启动模拟失败: {str(e)}")
1639
+ return jsonify({
1640
+ "success": False,
1641
+ "error": str(e),
1642
+ "traceback": traceback.format_exc()
1643
+ }), 500
1644
+
1645
+
1646
+ @simulation_bp.route('/stop', methods=['POST'])
1647
+ def stop_simulation():
1648
+ """
1649
+ 停止模拟
1650
+
1651
+ 请求(JSON):
1652
+ {
1653
+ "simulation_id": "sim_xxxx" // 必填,模拟ID
1654
+ }
1655
+
1656
+ 返回:
1657
+ {
1658
+ "success": true,
1659
+ "data": {
1660
+ "simulation_id": "sim_xxxx",
1661
+ "runner_status": "stopped",
1662
+ "completed_at": "2025-12-01T12:00:00"
1663
+ }
1664
+ }
1665
+ """
1666
+ try:
1667
+ data = request.get_json() or {}
1668
+
1669
+ simulation_id = data.get('simulation_id')
1670
+ if not simulation_id:
1671
+ return jsonify({
1672
+ "success": False,
1673
+ "error": "请提供 simulation_id"
1674
+ }), 400
1675
+
1676
+ run_state = SimulationRunner.stop_simulation(simulation_id)
1677
+
1678
+ # 更新模拟状态
1679
+ manager = SimulationManager()
1680
+ state = manager.get_simulation(simulation_id)
1681
+ if state:
1682
+ state.status = SimulationStatus.PAUSED
1683
+ manager._save_simulation_state(state)
1684
+
1685
+ return jsonify({
1686
+ "success": True,
1687
+ "data": run_state.to_dict()
1688
+ })
1689
+
1690
+ except ValueError as e:
1691
+ return jsonify({
1692
+ "success": False,
1693
+ "error": str(e)
1694
+ }), 400
1695
+
1696
+ except Exception as e:
1697
+ logger.error(f"停止模拟失败: {str(e)}")
1698
+ return jsonify({
1699
+ "success": False,
1700
+ "error": str(e),
1701
+ "traceback": traceback.format_exc()
1702
+ }), 500
1703
+
1704
+
1705
+ # ============== 实时状态监控接口 ==============
1706
+
1707
+ @simulation_bp.route('/<simulation_id>/run-status', methods=['GET'])
1708
+ def get_run_status(simulation_id: str):
1709
+ """
1710
+ 获取模拟运行实时状态(用于前端轮询)
1711
+
1712
+ 返回:
1713
+ {
1714
+ "success": true,
1715
+ "data": {
1716
+ "simulation_id": "sim_xxxx",
1717
+ "runner_status": "running",
1718
+ "current_round": 5,
1719
+ "total_rounds": 144,
1720
+ "progress_percent": 3.5,
1721
+ "simulated_hours": 2,
1722
+ "total_simulation_hours": 72,
1723
+ "twitter_running": true,
1724
+ "reddit_running": true,
1725
+ "twitter_actions_count": 150,
1726
+ "reddit_actions_count": 200,
1727
+ "total_actions_count": 350,
1728
+ "started_at": "2025-12-01T10:00:00",
1729
+ "updated_at": "2025-12-01T10:30:00"
1730
+ }
1731
+ }
1732
+ """
1733
+ try:
1734
+ run_state = SimulationRunner.get_run_state(simulation_id)
1735
+
1736
+ if not run_state:
1737
+ return jsonify({
1738
+ "success": True,
1739
+ "data": {
1740
+ "simulation_id": simulation_id,
1741
+ "runner_status": "idle",
1742
+ "current_round": 0,
1743
+ "total_rounds": 0,
1744
+ "progress_percent": 0,
1745
+ "twitter_actions_count": 0,
1746
+ "reddit_actions_count": 0,
1747
+ "total_actions_count": 0,
1748
+ }
1749
+ })
1750
+
1751
+ return jsonify({
1752
+ "success": True,
1753
+ "data": run_state.to_dict()
1754
+ })
1755
+
1756
+ except Exception as e:
1757
+ logger.error(f"获取运行状态失败: {str(e)}")
1758
+ return jsonify({
1759
+ "success": False,
1760
+ "error": str(e),
1761
+ "traceback": traceback.format_exc()
1762
+ }), 500
1763
+
1764
+
1765
+ @simulation_bp.route('/<simulation_id>/run-status/detail', methods=['GET'])
1766
+ def get_run_status_detail(simulation_id: str):
1767
+ """
1768
+ 获取模拟运行详细状态(包含所有动作)
1769
+
1770
+ 用于前端展示实时动态
1771
+
1772
+ Query参数:
1773
+ platform: 过滤平台(twitter/reddit,可选)
1774
+
1775
+ 返回:
1776
+ {
1777
+ "success": true,
1778
+ "data": {
1779
+ "simulation_id": "sim_xxxx",
1780
+ "runner_status": "running",
1781
+ "current_round": 5,
1782
+ ...
1783
+ "all_actions": [
1784
+ {
1785
+ "round_num": 5,
1786
+ "timestamp": "2025-12-01T10:30:00",
1787
+ "platform": "twitter",
1788
+ "agent_id": 3,
1789
+ "agent_name": "Agent Name",
1790
+ "action_type": "CREATE_POST",
1791
+ "action_args": {"content": "..."},
1792
+ "result": null,
1793
+ "success": true
1794
+ },
1795
+ ...
1796
+ ],
1797
+ "twitter_actions": [...], # Twitter 平台的所有动作
1798
+ "reddit_actions": [...] # Reddit 平台的所有动作
1799
+ }
1800
+ }
1801
+ """
1802
+ try:
1803
+ run_state = SimulationRunner.get_run_state(simulation_id)
1804
+ platform_filter = request.args.get('platform')
1805
+
1806
+ if not run_state:
1807
+ return jsonify({
1808
+ "success": True,
1809
+ "data": {
1810
+ "simulation_id": simulation_id,
1811
+ "runner_status": "idle",
1812
+ "all_actions": [],
1813
+ "twitter_actions": [],
1814
+ "reddit_actions": []
1815
+ }
1816
+ })
1817
+
1818
+ # 获取完整的动作列表
1819
+ all_actions = SimulationRunner.get_all_actions(
1820
+ simulation_id=simulation_id,
1821
+ platform=platform_filter
1822
+ )
1823
+
1824
+ # 分平台获取动作
1825
+ twitter_actions = SimulationRunner.get_all_actions(
1826
+ simulation_id=simulation_id,
1827
+ platform="twitter"
1828
+ ) if not platform_filter or platform_filter == "twitter" else []
1829
+
1830
+ reddit_actions = SimulationRunner.get_all_actions(
1831
+ simulation_id=simulation_id,
1832
+ platform="reddit"
1833
+ ) if not platform_filter or platform_filter == "reddit" else []
1834
+
1835
+ # 获取当前轮次的动作(recent_actions 只展示最新一轮)
1836
+ current_round = run_state.current_round
1837
+ recent_actions = SimulationRunner.get_all_actions(
1838
+ simulation_id=simulation_id,
1839
+ platform=platform_filter,
1840
+ round_num=current_round
1841
+ ) if current_round > 0 else []
1842
+
1843
+ # 获取基础状态信息
1844
+ result = run_state.to_dict()
1845
+ result["all_actions"] = [a.to_dict() for a in all_actions]
1846
+ result["twitter_actions"] = [a.to_dict() for a in twitter_actions]
1847
+ result["reddit_actions"] = [a.to_dict() for a in reddit_actions]
1848
+ result["rounds_count"] = len(run_state.rounds)
1849
+ # recent_actions 只展示当前最新一轮两个平台的内容
1850
+ result["recent_actions"] = [a.to_dict() for a in recent_actions]
1851
+
1852
+ return jsonify({
1853
+ "success": True,
1854
+ "data": result
1855
+ })
1856
+
1857
+ except Exception as e:
1858
+ logger.error(f"获取详细状态失败: {str(e)}")
1859
+ return jsonify({
1860
+ "success": False,
1861
+ "error": str(e),
1862
+ "traceback": traceback.format_exc()
1863
+ }), 500
1864
+
1865
+
1866
+ @simulation_bp.route('/<simulation_id>/actions', methods=['GET'])
1867
+ def get_simulation_actions(simulation_id: str):
1868
+ """
1869
+ 获取模拟中的Agent动作历史
1870
+
1871
+ Query参数:
1872
+ limit: 返回数量(默认100)
1873
+ offset: 偏移量(默认0)
1874
+ platform: 过滤平台(twitter/reddit)
1875
+ agent_id: 过滤Agent ID
1876
+ round_num: 过滤轮次
1877
+
1878
+ 返回:
1879
+ {
1880
+ "success": true,
1881
+ "data": {
1882
+ "count": 100,
1883
+ "actions": [...]
1884
+ }
1885
+ }
1886
+ """
1887
+ try:
1888
+ limit = request.args.get('limit', 100, type=int)
1889
+ offset = request.args.get('offset', 0, type=int)
1890
+ platform = request.args.get('platform')
1891
+ agent_id = request.args.get('agent_id', type=int)
1892
+ round_num = request.args.get('round_num', type=int)
1893
+
1894
+ actions = SimulationRunner.get_actions(
1895
+ simulation_id=simulation_id,
1896
+ limit=limit,
1897
+ offset=offset,
1898
+ platform=platform,
1899
+ agent_id=agent_id,
1900
+ round_num=round_num
1901
+ )
1902
+
1903
+ return jsonify({
1904
+ "success": True,
1905
+ "data": {
1906
+ "count": len(actions),
1907
+ "actions": [a.to_dict() for a in actions]
1908
+ }
1909
+ })
1910
+
1911
+ except Exception as e:
1912
+ logger.error(f"获取动作历史失败: {str(e)}")
1913
+ return jsonify({
1914
+ "success": False,
1915
+ "error": str(e),
1916
+ "traceback": traceback.format_exc()
1917
+ }), 500
1918
+
1919
+
1920
+ @simulation_bp.route('/<simulation_id>/timeline', methods=['GET'])
1921
+ def get_simulation_timeline(simulation_id: str):
1922
+ """
1923
+ 获取模拟时间线(按轮次汇总)
1924
+
1925
+ 用于前端展示进度条和时间线视图
1926
+
1927
+ Query参数:
1928
+ start_round: 起始轮次(默认0)
1929
+ end_round: 结束轮次(默认全部)
1930
+
1931
+ 返回每轮的汇总信息
1932
+ """
1933
+ try:
1934
+ start_round = request.args.get('start_round', 0, type=int)
1935
+ end_round = request.args.get('end_round', type=int)
1936
+
1937
+ timeline = SimulationRunner.get_timeline(
1938
+ simulation_id=simulation_id,
1939
+ start_round=start_round,
1940
+ end_round=end_round
1941
+ )
1942
+
1943
+ return jsonify({
1944
+ "success": True,
1945
+ "data": {
1946
+ "rounds_count": len(timeline),
1947
+ "timeline": timeline
1948
+ }
1949
+ })
1950
+
1951
+ except Exception as e:
1952
+ logger.error(f"获取时间线失败: {str(e)}")
1953
+ return jsonify({
1954
+ "success": False,
1955
+ "error": str(e),
1956
+ "traceback": traceback.format_exc()
1957
+ }), 500
1958
+
1959
+
1960
+ @simulation_bp.route('/<simulation_id>/agent-stats', methods=['GET'])
1961
+ def get_agent_stats(simulation_id: str):
1962
+ """
1963
+ 获取每个Agent的统计信息
1964
+
1965
+ 用于前端展示Agent活跃度排行、动作分布等
1966
+ """
1967
+ try:
1968
+ stats = SimulationRunner.get_agent_stats(simulation_id)
1969
+
1970
+ return jsonify({
1971
+ "success": True,
1972
+ "data": {
1973
+ "agents_count": len(stats),
1974
+ "stats": stats
1975
+ }
1976
+ })
1977
+
1978
+ except Exception as e:
1979
+ logger.error(f"获取Agent统计失败: {str(e)}")
1980
+ return jsonify({
1981
+ "success": False,
1982
+ "error": str(e),
1983
+ "traceback": traceback.format_exc()
1984
+ }), 500
1985
+
1986
+
1987
+ # ============== 数据库查询接口 ==============
1988
+
1989
+ @simulation_bp.route('/<simulation_id>/posts', methods=['GET'])
1990
+ def get_simulation_posts(simulation_id: str):
1991
+ """
1992
+ 获取模拟中的帖子
1993
+
1994
+ Query参数:
1995
+ platform: 平台类型(twitter/reddit)
1996
+ limit: 返回数量(默认50)
1997
+ offset: 偏移量
1998
+
1999
+ 返回帖子列表(从SQLite数据库读取)
2000
+ """
2001
+ try:
2002
+ platform = request.args.get('platform', 'reddit')
2003
+ limit = request.args.get('limit', 50, type=int)
2004
+ offset = request.args.get('offset', 0, type=int)
2005
+
2006
+ sim_dir = os.path.join(
2007
+ os.path.dirname(__file__),
2008
+ f'../../uploads/simulations/{simulation_id}'
2009
+ )
2010
+
2011
+ db_file = f"{platform}_simulation.db"
2012
+ db_path = os.path.join(sim_dir, db_file)
2013
+
2014
+ if not os.path.exists(db_path):
2015
+ return jsonify({
2016
+ "success": True,
2017
+ "data": {
2018
+ "platform": platform,
2019
+ "count": 0,
2020
+ "posts": [],
2021
+ "message": "数据库不存在,模拟可能尚未运行"
2022
+ }
2023
+ })
2024
+
2025
+ import sqlite3
2026
+ conn = sqlite3.connect(db_path)
2027
+ conn.row_factory = sqlite3.Row
2028
+ cursor = conn.cursor()
2029
+
2030
+ try:
2031
+ cursor.execute("""
2032
+ SELECT * FROM post
2033
+ ORDER BY created_at DESC
2034
+ LIMIT ? OFFSET ?
2035
+ """, (limit, offset))
2036
+
2037
+ posts = [dict(row) for row in cursor.fetchall()]
2038
+
2039
+ cursor.execute("SELECT COUNT(*) FROM post")
2040
+ total = cursor.fetchone()[0]
2041
+
2042
+ except sqlite3.OperationalError:
2043
+ posts = []
2044
+ total = 0
2045
+
2046
+ conn.close()
2047
+
2048
+ return jsonify({
2049
+ "success": True,
2050
+ "data": {
2051
+ "platform": platform,
2052
+ "total": total,
2053
+ "count": len(posts),
2054
+ "posts": posts
2055
+ }
2056
+ })
2057
+
2058
+ except Exception as e:
2059
+ logger.error(f"获取帖子失败: {str(e)}")
2060
+ return jsonify({
2061
+ "success": False,
2062
+ "error": str(e),
2063
+ "traceback": traceback.format_exc()
2064
+ }), 500
2065
+
2066
+
2067
+ @simulation_bp.route('/<simulation_id>/comments', methods=['GET'])
2068
+ def get_simulation_comments(simulation_id: str):
2069
+ """
2070
+ 获取模拟中的评论(仅Reddit)
2071
+
2072
+ Query参数:
2073
+ post_id: 过滤帖子ID(可选)
2074
+ limit: 返回数量
2075
+ offset: 偏移量
2076
+ """
2077
+ try:
2078
+ post_id = request.args.get('post_id')
2079
+ limit = request.args.get('limit', 50, type=int)
2080
+ offset = request.args.get('offset', 0, type=int)
2081
+
2082
+ sim_dir = os.path.join(
2083
+ os.path.dirname(__file__),
2084
+ f'../../uploads/simulations/{simulation_id}'
2085
+ )
2086
+
2087
+ db_path = os.path.join(sim_dir, "reddit_simulation.db")
2088
+
2089
+ if not os.path.exists(db_path):
2090
+ return jsonify({
2091
+ "success": True,
2092
+ "data": {
2093
+ "count": 0,
2094
+ "comments": []
2095
+ }
2096
+ })
2097
+
2098
+ import sqlite3
2099
+ conn = sqlite3.connect(db_path)
2100
+ conn.row_factory = sqlite3.Row
2101
+ cursor = conn.cursor()
2102
+
2103
+ try:
2104
+ if post_id:
2105
+ cursor.execute("""
2106
+ SELECT * FROM comment
2107
+ WHERE post_id = ?
2108
+ ORDER BY created_at DESC
2109
+ LIMIT ? OFFSET ?
2110
+ """, (post_id, limit, offset))
2111
+ else:
2112
+ cursor.execute("""
2113
+ SELECT * FROM comment
2114
+ ORDER BY created_at DESC
2115
+ LIMIT ? OFFSET ?
2116
+ """, (limit, offset))
2117
+
2118
+ comments = [dict(row) for row in cursor.fetchall()]
2119
+
2120
+ except sqlite3.OperationalError:
2121
+ comments = []
2122
+
2123
+ conn.close()
2124
+
2125
+ return jsonify({
2126
+ "success": True,
2127
+ "data": {
2128
+ "count": len(comments),
2129
+ "comments": comments
2130
+ }
2131
+ })
2132
+
2133
+ except Exception as e:
2134
+ logger.error(f"获取评论失败: {str(e)}")
2135
+ return jsonify({
2136
+ "success": False,
2137
+ "error": str(e),
2138
+ "traceback": traceback.format_exc()
2139
+ }), 500
2140
+
2141
+
2142
+ # ============== Interview 采访接口 ==============
2143
+
2144
+ @simulation_bp.route('/interview', methods=['POST'])
2145
+ def interview_agent():
2146
+ """
2147
+ 采访单个Agent
2148
+
2149
+ 注意:此功能需要模拟环境处于运行状态(完成模拟循环后进入等待命令模式)
2150
+
2151
+ 请求(JSON):
2152
+ {
2153
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
2154
+ "agent_id": 0, // 必填,Agent ID
2155
+ "prompt": "你对这件事有什么看法?", // 必填,采访问题
2156
+ "platform": "twitter", // 可选,指定平台(twitter/reddit)
2157
+ // 不指定时:双平台模拟同时采访两个平台
2158
+ "timeout": 60 // 可选,超时时间(秒),默认60
2159
+ }
2160
+
2161
+ 返回(不指定platform,双平台模式):
2162
+ {
2163
+ "success": true,
2164
+ "data": {
2165
+ "agent_id": 0,
2166
+ "prompt": "你对这件事有什么看法?",
2167
+ "result": {
2168
+ "agent_id": 0,
2169
+ "prompt": "...",
2170
+ "platforms": {
2171
+ "twitter": {"agent_id": 0, "response": "...", "platform": "twitter"},
2172
+ "reddit": {"agent_id": 0, "response": "...", "platform": "reddit"}
2173
+ }
2174
+ },
2175
+ "timestamp": "2025-12-08T10:00:01"
2176
+ }
2177
+ }
2178
+
2179
+ 返回(指定platform):
2180
+ {
2181
+ "success": true,
2182
+ "data": {
2183
+ "agent_id": 0,
2184
+ "prompt": "你对这件事有什么看法?",
2185
+ "result": {
2186
+ "agent_id": 0,
2187
+ "response": "我认为...",
2188
+ "platform": "twitter",
2189
+ "timestamp": "2025-12-08T10:00:00"
2190
+ },
2191
+ "timestamp": "2025-12-08T10:00:01"
2192
+ }
2193
+ }
2194
+ """
2195
+ try:
2196
+ data = request.get_json() or {}
2197
+
2198
+ simulation_id = data.get('simulation_id')
2199
+ agent_id = data.get('agent_id')
2200
+ prompt = data.get('prompt')
2201
+ platform = data.get('platform') # 可选:twitter/reddit/None
2202
+ timeout = data.get('timeout', 60)
2203
+
2204
+ if not simulation_id:
2205
+ return jsonify({
2206
+ "success": False,
2207
+ "error": "请提供 simulation_id"
2208
+ }), 400
2209
+
2210
+ if agent_id is None:
2211
+ return jsonify({
2212
+ "success": False,
2213
+ "error": "请提供 agent_id"
2214
+ }), 400
2215
+
2216
+ if not prompt:
2217
+ return jsonify({
2218
+ "success": False,
2219
+ "error": "请提供 prompt(采访问题)"
2220
+ }), 400
2221
+
2222
+ # 验证platform参数
2223
+ if platform and platform not in ("twitter", "reddit"):
2224
+ return jsonify({
2225
+ "success": False,
2226
+ "error": "platform 参数只能是 'twitter' 或 'reddit'"
2227
+ }), 400
2228
+
2229
+ # 检查环境状态
2230
+ if not SimulationRunner.check_env_alive(simulation_id):
2231
+ return jsonify({
2232
+ "success": False,
2233
+ "error": "模拟环境未运行或已关闭。请确保模拟已完成并进入等待命令���式。"
2234
+ }), 400
2235
+
2236
+ # 优化prompt,添加前缀避免Agent调用工具
2237
+ optimized_prompt = optimize_interview_prompt(prompt)
2238
+
2239
+ result = SimulationRunner.interview_agent(
2240
+ simulation_id=simulation_id,
2241
+ agent_id=agent_id,
2242
+ prompt=optimized_prompt,
2243
+ platform=platform,
2244
+ timeout=timeout
2245
+ )
2246
+
2247
+ return jsonify({
2248
+ "success": result.get("success", False),
2249
+ "data": result
2250
+ })
2251
+
2252
+ except ValueError as e:
2253
+ return jsonify({
2254
+ "success": False,
2255
+ "error": str(e)
2256
+ }), 400
2257
+
2258
+ except TimeoutError as e:
2259
+ return jsonify({
2260
+ "success": False,
2261
+ "error": f"等待Interview响应超时: {str(e)}"
2262
+ }), 504
2263
+
2264
+ except Exception as e:
2265
+ logger.error(f"Interview失败: {str(e)}")
2266
+ return jsonify({
2267
+ "success": False,
2268
+ "error": str(e),
2269
+ "traceback": traceback.format_exc()
2270
+ }), 500
2271
+
2272
+
2273
+ @simulation_bp.route('/interview/batch', methods=['POST'])
2274
+ def interview_agents_batch():
2275
+ """
2276
+ 批量采访多个Agent
2277
+
2278
+ 注意:此功能需要模拟环境处于运行状态
2279
+
2280
+ 请求(JSON):
2281
+ {
2282
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
2283
+ "interviews": [ // 必填,采访列表
2284
+ {
2285
+ "agent_id": 0,
2286
+ "prompt": "你对A有什么看法?",
2287
+ "platform": "twitter" // 可选,指定该Agent的采访平台
2288
+ },
2289
+ {
2290
+ "agent_id": 1,
2291
+ "prompt": "你对B有什么看法?" // 不指定platform则使用默认值
2292
+ }
2293
+ ],
2294
+ "platform": "reddit", // 可选,默认平台(被每项的platform覆盖)
2295
+ // 不指定时:双平台模拟每个Agent同时采访两个平台
2296
+ "timeout": 120 // 可选,超时时间(秒),默认120
2297
+ }
2298
+
2299
+ 返回:
2300
+ {
2301
+ "success": true,
2302
+ "data": {
2303
+ "interviews_count": 2,
2304
+ "result": {
2305
+ "interviews_count": 4,
2306
+ "results": {
2307
+ "twitter_0": {"agent_id": 0, "response": "...", "platform": "twitter"},
2308
+ "reddit_0": {"agent_id": 0, "response": "...", "platform": "reddit"},
2309
+ "twitter_1": {"agent_id": 1, "response": "...", "platform": "twitter"},
2310
+ "reddit_1": {"agent_id": 1, "response": "...", "platform": "reddit"}
2311
+ }
2312
+ },
2313
+ "timestamp": "2025-12-08T10:00:01"
2314
+ }
2315
+ }
2316
+ """
2317
+ try:
2318
+ data = request.get_json() or {}
2319
+
2320
+ simulation_id = data.get('simulation_id')
2321
+ interviews = data.get('interviews')
2322
+ platform = data.get('platform') # 可选:twitter/reddit/None
2323
+ timeout = data.get('timeout', 120)
2324
+
2325
+ if not simulation_id:
2326
+ return jsonify({
2327
+ "success": False,
2328
+ "error": "请提供 simulation_id"
2329
+ }), 400
2330
+
2331
+ if not interviews or not isinstance(interviews, list):
2332
+ return jsonify({
2333
+ "success": False,
2334
+ "error": "请提供 interviews(采访列表)"
2335
+ }), 400
2336
+
2337
+ # 验证platform参数
2338
+ if platform and platform not in ("twitter", "reddit"):
2339
+ return jsonify({
2340
+ "success": False,
2341
+ "error": "platform 参数只能是 'twitter' 或 'reddit'"
2342
+ }), 400
2343
+
2344
+ # 验证每个采访项
2345
+ for i, interview in enumerate(interviews):
2346
+ if 'agent_id' not in interview:
2347
+ return jsonify({
2348
+ "success": False,
2349
+ "error": f"采访列表第{i+1}项缺少 agent_id"
2350
+ }), 400
2351
+ if 'prompt' not in interview:
2352
+ return jsonify({
2353
+ "success": False,
2354
+ "error": f"采访列表第{i+1}项缺少 prompt"
2355
+ }), 400
2356
+ # 验证每项的platform(如果有)
2357
+ item_platform = interview.get('platform')
2358
+ if item_platform and item_platform not in ("twitter", "reddit"):
2359
+ return jsonify({
2360
+ "success": False,
2361
+ "error": f"采访列表第{i+1}项的platform只能是 'twitter' 或 'reddit'"
2362
+ }), 400
2363
+
2364
+ # 检查环境状态
2365
+ if not SimulationRunner.check_env_alive(simulation_id):
2366
+ return jsonify({
2367
+ "success": False,
2368
+ "error": "模拟环境未运行或已关闭。请确保模拟已完成并进入等待命令模式。"
2369
+ }), 400
2370
+
2371
+ # 优化每个采访项的prompt,添��前缀避免Agent调用工具
2372
+ optimized_interviews = []
2373
+ for interview in interviews:
2374
+ optimized_interview = interview.copy()
2375
+ optimized_interview['prompt'] = optimize_interview_prompt(interview.get('prompt', ''))
2376
+ optimized_interviews.append(optimized_interview)
2377
+
2378
+ result = SimulationRunner.interview_agents_batch(
2379
+ simulation_id=simulation_id,
2380
+ interviews=optimized_interviews,
2381
+ platform=platform,
2382
+ timeout=timeout
2383
+ )
2384
+
2385
+ return jsonify({
2386
+ "success": result.get("success", False),
2387
+ "data": result
2388
+ })
2389
+
2390
+ except ValueError as e:
2391
+ return jsonify({
2392
+ "success": False,
2393
+ "error": str(e)
2394
+ }), 400
2395
+
2396
+ except TimeoutError as e:
2397
+ return jsonify({
2398
+ "success": False,
2399
+ "error": f"等待批量Interview响应超时: {str(e)}"
2400
+ }), 504
2401
+
2402
+ except Exception as e:
2403
+ logger.error(f"批量Interview失败: {str(e)}")
2404
+ return jsonify({
2405
+ "success": False,
2406
+ "error": str(e),
2407
+ "traceback": traceback.format_exc()
2408
+ }), 500
2409
+
2410
+
2411
+ @simulation_bp.route('/interview/all', methods=['POST'])
2412
+ def interview_all_agents():
2413
+ """
2414
+ 全局采访 - 使用相同问题采访所有Agent
2415
+
2416
+ 注意:此功能需要模拟环境处于运行状态
2417
+
2418
+ 请求(JSON):
2419
+ {
2420
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
2421
+ "prompt": "你对这件事整体有什么看法?", // 必填,采访问题(所有Agent使用相同问题)
2422
+ "platform": "reddit", // 可选,指定平台(twitter/reddit)
2423
+ // 不指定时:双平台模拟每个Agent同时采访两个平台
2424
+ "timeout": 180 // 可选,超时时间(秒),默认180
2425
+ }
2426
+
2427
+ 返回:
2428
+ {
2429
+ "success": true,
2430
+ "data": {
2431
+ "interviews_count": 50,
2432
+ "result": {
2433
+ "interviews_count": 100,
2434
+ "results": {
2435
+ "twitter_0": {"agent_id": 0, "response": "...", "platform": "twitter"},
2436
+ "reddit_0": {"agent_id": 0, "response": "...", "platform": "reddit"},
2437
+ ...
2438
+ }
2439
+ },
2440
+ "timestamp": "2025-12-08T10:00:01"
2441
+ }
2442
+ }
2443
+ """
2444
+ try:
2445
+ data = request.get_json() or {}
2446
+
2447
+ simulation_id = data.get('simulation_id')
2448
+ prompt = data.get('prompt')
2449
+ platform = data.get('platform') # 可选:twitter/reddit/None
2450
+ timeout = data.get('timeout', 180)
2451
+
2452
+ if not simulation_id:
2453
+ return jsonify({
2454
+ "success": False,
2455
+ "error": "请提供 simulation_id"
2456
+ }), 400
2457
+
2458
+ if not prompt:
2459
+ return jsonify({
2460
+ "success": False,
2461
+ "error": "请提供 prompt(采访问题)"
2462
+ }), 400
2463
+
2464
+ # 验证platform参数
2465
+ if platform and platform not in ("twitter", "reddit"):
2466
+ return jsonify({
2467
+ "success": False,
2468
+ "error": "platform 参数只能是 'twitter' 或 'reddit'"
2469
+ }), 400
2470
+
2471
+ # 检查环境状态
2472
+ if not SimulationRunner.check_env_alive(simulation_id):
2473
+ return jsonify({
2474
+ "success": False,
2475
+ "error": "模拟环境未运行或已关闭。请确保模拟已完成并进入等待命令模式。"
2476
+ }), 400
2477
+
2478
+ # 优化prompt,添加前缀避免Agent调用工具
2479
+ optimized_prompt = optimize_interview_prompt(prompt)
2480
+
2481
+ result = SimulationRunner.interview_all_agents(
2482
+ simulation_id=simulation_id,
2483
+ prompt=optimized_prompt,
2484
+ platform=platform,
2485
+ timeout=timeout
2486
+ )
2487
+
2488
+ return jsonify({
2489
+ "success": result.get("success", False),
2490
+ "data": result
2491
+ })
2492
+
2493
+ except ValueError as e:
2494
+ return jsonify({
2495
+ "success": False,
2496
+ "error": str(e)
2497
+ }), 400
2498
+
2499
+ except TimeoutError as e:
2500
+ return jsonify({
2501
+ "success": False,
2502
+ "error": f"等待全局Interview响应超时: {str(e)}"
2503
+ }), 504
2504
+
2505
+ except Exception as e:
2506
+ logger.error(f"全局Interview失败: {str(e)}")
2507
+ return jsonify({
2508
+ "success": False,
2509
+ "error": str(e),
2510
+ "traceback": traceback.format_exc()
2511
+ }), 500
2512
+
2513
+
2514
+ @simulation_bp.route('/interview/history', methods=['POST'])
2515
+ def get_interview_history():
2516
+ """
2517
+ 获取Interview历史记录
2518
+
2519
+ 从模拟数据库中读取所有Interview记录
2520
+
2521
+ 请求(JSON):
2522
+ {
2523
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
2524
+ "platform": "reddit", // 可选,平台类��(reddit/twitter)
2525
+ // 不指定则返回两个平台的所有历史
2526
+ "agent_id": 0, // 可选,只获取该Agent的采访历史
2527
+ "limit": 100 // 可选,返回数量,默认100
2528
+ }
2529
+
2530
+ 返回:
2531
+ {
2532
+ "success": true,
2533
+ "data": {
2534
+ "count": 10,
2535
+ "history": [
2536
+ {
2537
+ "agent_id": 0,
2538
+ "response": "我认为...",
2539
+ "prompt": "你对这件事有什么看法?",
2540
+ "timestamp": "2025-12-08T10:00:00",
2541
+ "platform": "reddit"
2542
+ },
2543
+ ...
2544
+ ]
2545
+ }
2546
+ }
2547
+ """
2548
+ try:
2549
+ data = request.get_json() or {}
2550
+
2551
+ simulation_id = data.get('simulation_id')
2552
+ platform = data.get('platform') # 不指定则返回两个平台的历史
2553
+ agent_id = data.get('agent_id')
2554
+ limit = data.get('limit', 100)
2555
+
2556
+ if not simulation_id:
2557
+ return jsonify({
2558
+ "success": False,
2559
+ "error": "请提供 simulation_id"
2560
+ }), 400
2561
+
2562
+ history = SimulationRunner.get_interview_history(
2563
+ simulation_id=simulation_id,
2564
+ platform=platform,
2565
+ agent_id=agent_id,
2566
+ limit=limit
2567
+ )
2568
+
2569
+ return jsonify({
2570
+ "success": True,
2571
+ "data": {
2572
+ "count": len(history),
2573
+ "history": history
2574
+ }
2575
+ })
2576
+
2577
+ except Exception as e:
2578
+ logger.error(f"获取Interview历史失败: {str(e)}")
2579
+ return jsonify({
2580
+ "success": False,
2581
+ "error": str(e),
2582
+ "traceback": traceback.format_exc()
2583
+ }), 500
2584
+
2585
+
2586
+ @simulation_bp.route('/env-status', methods=['POST'])
2587
+ def get_env_status():
2588
+ """
2589
+ 获取模拟环境状态
2590
+
2591
+ 检查模拟环境是否存活(可以接收Interview命令)
2592
+
2593
+ 请求(JSON):
2594
+ {
2595
+ "simulation_id": "sim_xxxx" // 必填,模拟ID
2596
+ }
2597
+
2598
+ 返回:
2599
+ {
2600
+ "success": true,
2601
+ "data": {
2602
+ "simulation_id": "sim_xxxx",
2603
+ "env_alive": true,
2604
+ "twitter_available": true,
2605
+ "reddit_available": true,
2606
+ "message": "环境正在运行,可以接收Interview命令"
2607
+ }
2608
+ }
2609
+ """
2610
+ try:
2611
+ data = request.get_json() or {}
2612
+
2613
+ simulation_id = data.get('simulation_id')
2614
+
2615
+ if not simulation_id:
2616
+ return jsonify({
2617
+ "success": False,
2618
+ "error": "请提供 simulation_id"
2619
+ }), 400
2620
+
2621
+ env_alive = SimulationRunner.check_env_alive(simulation_id)
2622
+
2623
+ # 获取更详细的状态信息
2624
+ env_status = SimulationRunner.get_env_status_detail(simulation_id)
2625
+
2626
+ if env_alive:
2627
+ message = "环境正在运行,可以接收Interview命令"
2628
+ else:
2629
+ message = "环境未运行或已关闭"
2630
+
2631
+ return jsonify({
2632
+ "success": True,
2633
+ "data": {
2634
+ "simulation_id": simulation_id,
2635
+ "env_alive": env_alive,
2636
+ "twitter_available": env_status.get("twitter_available", False),
2637
+ "reddit_available": env_status.get("reddit_available", False),
2638
+ "message": message
2639
+ }
2640
+ })
2641
+
2642
+ except Exception as e:
2643
+ logger.error(f"获取环境状态失败: {str(e)}")
2644
+ return jsonify({
2645
+ "success": False,
2646
+ "error": str(e),
2647
+ "traceback": traceback.format_exc()
2648
+ }), 500
2649
+
2650
+
2651
+ @simulation_bp.route('/close-env', methods=['POST'])
2652
+ def close_simulation_env():
2653
+ """
2654
+ 关闭模拟环境
2655
+
2656
+ 向模拟发送关闭环境命令,使其优雅退出等待命令模式。
2657
+
2658
+ 注意:这不同于 /stop 接口,/stop 会强制终止进程,
2659
+ 而此接口会让模拟优雅地关闭环境并退出。
2660
+
2661
+ 请求(JSON):
2662
+ {
2663
+ "simulation_id": "sim_xxxx", // 必填,模拟ID
2664
+ "timeout": 30 // 可选,超时时间(秒),默认30
2665
+ }
2666
+
2667
+ 返回:
2668
+ {
2669
+ "success": true,
2670
+ "data": {
2671
+ "message": "环境关闭命令已发送",
2672
+ "result": {...},
2673
+ "timestamp": "2025-12-08T10:00:01"
2674
+ }
2675
+ }
2676
+ """
2677
+ try:
2678
+ data = request.get_json() or {}
2679
+
2680
+ simulation_id = data.get('simulation_id')
2681
+ timeout = data.get('timeout', 30)
2682
+
2683
+ if not simulation_id:
2684
+ return jsonify({
2685
+ "success": False,
2686
+ "error": "请提供 simulation_id"
2687
+ }), 400
2688
+
2689
+ result = SimulationRunner.close_simulation_env(
2690
+ simulation_id=simulation_id,
2691
+ timeout=timeout
2692
+ )
2693
+
2694
+ # 更新模拟状态
2695
+ manager = SimulationManager()
2696
+ state = manager.get_simulation(simulation_id)
2697
+ if state:
2698
+ state.status = SimulationStatus.COMPLETED
2699
+ manager._save_simulation_state(state)
2700
+
2701
+ return jsonify({
2702
+ "success": result.get("success", False),
2703
+ "data": result
2704
+ })
2705
+
2706
+ except ValueError as e:
2707
+ return jsonify({
2708
+ "success": False,
2709
+ "error": str(e)
2710
+ }), 400
2711
+
2712
+ except Exception as e:
2713
+ logger.error(f"关闭环境失败: {str(e)}")
2714
+ return jsonify({
2715
+ "success": False,
2716
+ "error": str(e),
2717
+ "traceback": traceback.format_exc()
2718
+ }), 500
app/app/config.py ADDED
@@ -0,0 +1,129 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 配置管理
3
+ 统一从项目根目录的 .env 文件加载配置
4
+ """
5
+
6
+ import os
7
+ from dotenv import load_dotenv
8
+
9
+ # 加载项目根目录的 .env 文件
10
+ # 路径: MiroFish/.env (相对于 backend/app/config.py)
11
+ project_root_env = os.path.join(os.path.dirname(__file__), '../../.env')
12
+
13
+ if os.path.exists(project_root_env):
14
+ load_dotenv(project_root_env, override=True)
15
+ else:
16
+ # 如果根目录没有 .env,尝试加载环境变量(用于生产环境)
17
+ load_dotenv(override=True)
18
+
19
+
20
+ class Config:
21
+ """Flask配置类"""
22
+
23
+ # Flask配置
24
+ SECRET_KEY = os.environ.get('SECRET_KEY', 'mirofish-secret-key')
25
+ DEBUG = os.environ.get('FLASK_DEBUG', 'True').lower() == 'true'
26
+
27
+ # JSON配置 - 禁用ASCII转义,让中文直接显示(而不是 \uXXXX 格式)
28
+ JSON_AS_ASCII = False
29
+
30
+ # LLM配置(统一使用OpenAI格式)
31
+ LLM_API_KEY = os.environ.get('LLM_API_KEY')
32
+ LLM_BASE_URL = os.environ.get('LLM_BASE_URL', 'https://api.openai.com/v1')
33
+ LLM_MODEL_NAME = os.environ.get('LLM_MODEL_NAME', 'gpt-4o-mini')
34
+
35
+ # Zep配置
36
+ ZEP_API_KEY = os.environ.get('ZEP_API_KEY')
37
+
38
+ @classmethod
39
+ def get_zep_api_keys(cls):
40
+ """Non-empty Zep API keys: single ZEP_API_KEY or comma-separated list (first used by default)."""
41
+ raw = (cls.ZEP_API_KEY or "").strip()
42
+ if not raw:
43
+ return []
44
+ return [k.strip() for k in raw.split(",") if k.strip()]
45
+
46
+ @classmethod
47
+ def get_llm_api_keys(cls):
48
+ """Non-empty LLM keys: single LLM_API_KEY or comma-separated list (paired slot with Zep)."""
49
+ raw = (cls.LLM_API_KEY or "").strip()
50
+ if not raw:
51
+ return []
52
+ return [k.strip() for k in raw.split(",") if k.strip()]
53
+
54
+ # Razorpay + simulation unlock (website)
55
+ RAZORPAY_KEY_ID = os.environ.get("RAZORPAY_KEY_ID", "").strip()
56
+ RAZORPAY_KEY_SECRET = os.environ.get("RAZORPAY_KEY_SECRET", "").strip()
57
+ RAZORPAY_AMOUNT_PAISE = int(os.environ.get("RAZORPAY_AMOUNT_PAISE", "49900"))
58
+ RAZORPAY_CURRENCY = os.environ.get("RAZORPAY_CURRENCY", "INR").strip() or "INR"
59
+ PAYMENT_COUPON_CODES = os.environ.get("PAYMENT_COUPON_CODES", "").strip()
60
+ PAYMENT_JWT_SECRET = os.environ.get("PAYMENT_JWT_SECRET", "").strip()
61
+
62
+ @classmethod
63
+ def payment_coupon_set(cls):
64
+ if not cls.PAYMENT_COUPON_CODES:
65
+ return set()
66
+ return {c.strip().lower() for c in cls.PAYMENT_COUPON_CODES.split(",") if c.strip()}
67
+
68
+ @classmethod
69
+ def billing_enabled(cls) -> bool:
70
+ razorpay_ok = bool(cls.RAZORPAY_KEY_ID and cls.RAZORPAY_KEY_SECRET)
71
+ return razorpay_ok or bool(cls.payment_coupon_set())
72
+
73
+ @classmethod
74
+ def payment_signing_secret(cls) -> str:
75
+ return (cls.PAYMENT_JWT_SECRET or cls.PIPELINE_CLIENT_SECRET or cls.SECRET_KEY or "mirofish-pay").strip()
76
+
77
+ # 文件上传配置
78
+ MAX_CONTENT_LENGTH = 50 * 1024 * 1024 # 50MB
79
+ UPLOAD_FOLDER = os.path.join(os.path.dirname(__file__), '../uploads')
80
+ ALLOWED_EXTENSIONS = {'pdf', 'md', 'txt', 'markdown'}
81
+
82
+ # 文本处理配置
83
+ DEFAULT_CHUNK_SIZE = 500 # 默认切块大小
84
+ DEFAULT_CHUNK_OVERLAP = 50 # 默认重叠大小
85
+
86
+ # OASIS模拟配置
87
+ OASIS_DEFAULT_MAX_ROUNDS = int(os.environ.get('OASIS_DEFAULT_MAX_ROUNDS', '10'))
88
+ OASIS_SIMULATION_DATA_DIR = os.path.join(os.path.dirname(__file__), '../uploads/simulations')
89
+
90
+ # OASIS平台可用动作配置
91
+ OASIS_TWITTER_ACTIONS = [
92
+ 'CREATE_POST', 'LIKE_POST', 'REPOST', 'FOLLOW', 'DO_NOTHING', 'QUOTE_POST'
93
+ ]
94
+ OASIS_REDDIT_ACTIONS = [
95
+ 'LIKE_POST', 'DISLIKE_POST', 'CREATE_POST', 'CREATE_COMMENT',
96
+ 'LIKE_COMMENT', 'DISLIKE_COMMENT', 'SEARCH_POSTS', 'SEARCH_USER',
97
+ 'TREND', 'REFRESH', 'DO_NOTHING', 'FOLLOW', 'MUTE'
98
+ ]
99
+
100
+ # 全自动流水线(website → /api/pipeline/start)
101
+ PIPELINE_MAX_ROUNDS = os.environ.get('PIPELINE_MAX_ROUNDS') # optional cap on OASIS rounds
102
+ PIPELINE_PARALLEL_PROFILES = int(os.environ.get('PIPELINE_PARALLEL_PROFILES', '5'))
103
+
104
+ # Supabase: report storage + simulation_runs table (free tier friendly)
105
+ SUPABASE_URL = os.environ.get('SUPABASE_URL', '').rstrip('/')
106
+ SUPABASE_SERVICE_ROLE_KEY = os.environ.get('SUPABASE_SERVICE_ROLE_KEY', '')
107
+ # Project Settings → API → JWT Secret (verify user access tokens from the website)
108
+ SUPABASE_JWT_SECRET = os.environ.get('SUPABASE_JWT_SECRET', '')
109
+ SUPABASE_REPORTS_BUCKET = os.environ.get('SUPABASE_REPORTS_BUCKET', 'reports')
110
+ # Optional: signing secret for legacy MiroFish Bearer JWTs. Defaults to SECRET_KEY.
111
+ MIROFISH_SESSION_JWT_SECRET = os.environ.get('MIROFISH_SESSION_JWT_SECRET', '').strip()
112
+ # Same value as website VITE_PIPELINE_CLIENT_SECRET (browser sends X-Mirofish-User-Id + X-Pipeline-Client-Secret).
113
+ PIPELINE_CLIENT_SECRET = os.environ.get('PIPELINE_CLIENT_SECRET', '').strip()
114
+
115
+ # Report Agent配置
116
+ REPORT_AGENT_MAX_TOOL_CALLS = int(os.environ.get('REPORT_AGENT_MAX_TOOL_CALLS', '5'))
117
+ REPORT_AGENT_MAX_REFLECTION_ROUNDS = int(os.environ.get('REPORT_AGENT_MAX_REFLECTION_ROUNDS', '2'))
118
+ REPORT_AGENT_TEMPERATURE = float(os.environ.get('REPORT_AGENT_TEMPERATURE', '0.5'))
119
+
120
+ @classmethod
121
+ def validate(cls):
122
+ """验证必要配置"""
123
+ errors = []
124
+ if not cls.get_llm_api_keys():
125
+ errors.append("LLM_API_KEY 未配置")
126
+ if not cls.get_zep_api_keys():
127
+ errors.append("ZEP_API_KEY 未配置")
128
+ return errors
129
+
app/app/models/__init__.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 数据模型模块
3
+ """
4
+
5
+ from .task import TaskManager, TaskStatus
6
+ from .project import Project, ProjectStatus, ProjectManager
7
+
8
+ __all__ = ['TaskManager', 'TaskStatus', 'Project', 'ProjectStatus', 'ProjectManager']
9
+
app/app/models/__pycache__/__init__.cpython-311.pyc ADDED
Binary file (451 Bytes). View file
 
app/app/models/__pycache__/project.cpython-311.pyc ADDED
Binary file (15.7 kB). View file
 
app/app/models/__pycache__/task.cpython-311.pyc ADDED
Binary file (9.69 kB). View file
 
app/app/models/project.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 项目上下文管理
3
+ 用于在服务端持久化项目状态,避免前端在接口间传递大量数据
4
+ """
5
+
6
+ import os
7
+ import json
8
+ import uuid
9
+ import shutil
10
+ from datetime import datetime
11
+ from typing import Dict, Any, List, Optional
12
+ from enum import Enum
13
+ from dataclasses import dataclass, field, asdict
14
+ from ..config import Config
15
+
16
+
17
+ class ProjectStatus(str, Enum):
18
+ """项目状态"""
19
+ CREATED = "created" # 刚创建,文件已上传
20
+ ONTOLOGY_GENERATED = "ontology_generated" # 本体已生成
21
+ GRAPH_BUILDING = "graph_building" # 图谱构建中
22
+ GRAPH_COMPLETED = "graph_completed" # 图谱构建完成
23
+ FAILED = "failed" # 失败
24
+
25
+
26
+ @dataclass
27
+ class Project:
28
+ """项目数据模型"""
29
+ project_id: str
30
+ name: str
31
+ status: ProjectStatus
32
+ created_at: str
33
+ updated_at: str
34
+
35
+ # 文件信息
36
+ files: List[Dict[str, str]] = field(default_factory=list) # [{filename, path, size}]
37
+ total_text_length: int = 0
38
+
39
+ # 本体信息(接口1生成后填充)
40
+ ontology: Optional[Dict[str, Any]] = None
41
+ analysis_summary: Optional[str] = None
42
+
43
+ # 图谱信息(接口2完成后填充)
44
+ graph_id: Optional[str] = None
45
+ graph_build_task_id: Optional[str] = None
46
+
47
+ # 配置
48
+ simulation_requirement: Optional[str] = None
49
+ chunk_size: int = 500
50
+ chunk_overlap: int = 50
51
+
52
+ # 错误信息
53
+ error: Optional[str] = None
54
+
55
+ def to_dict(self) -> Dict[str, Any]:
56
+ """转换为字典"""
57
+ return {
58
+ "project_id": self.project_id,
59
+ "name": self.name,
60
+ "status": self.status.value if isinstance(self.status, ProjectStatus) else self.status,
61
+ "created_at": self.created_at,
62
+ "updated_at": self.updated_at,
63
+ "files": self.files,
64
+ "total_text_length": self.total_text_length,
65
+ "ontology": self.ontology,
66
+ "analysis_summary": self.analysis_summary,
67
+ "graph_id": self.graph_id,
68
+ "graph_build_task_id": self.graph_build_task_id,
69
+ "simulation_requirement": self.simulation_requirement,
70
+ "chunk_size": self.chunk_size,
71
+ "chunk_overlap": self.chunk_overlap,
72
+ "error": self.error
73
+ }
74
+
75
+ @classmethod
76
+ def from_dict(cls, data: Dict[str, Any]) -> 'Project':
77
+ """从字典创建"""
78
+ status = data.get('status', 'created')
79
+ if isinstance(status, str):
80
+ status = ProjectStatus(status)
81
+
82
+ return cls(
83
+ project_id=data['project_id'],
84
+ name=data.get('name', 'Unnamed Project'),
85
+ status=status,
86
+ created_at=data.get('created_at', ''),
87
+ updated_at=data.get('updated_at', ''),
88
+ files=data.get('files', []),
89
+ total_text_length=data.get('total_text_length', 0),
90
+ ontology=data.get('ontology'),
91
+ analysis_summary=data.get('analysis_summary'),
92
+ graph_id=data.get('graph_id'),
93
+ graph_build_task_id=data.get('graph_build_task_id'),
94
+ simulation_requirement=data.get('simulation_requirement'),
95
+ chunk_size=data.get('chunk_size', 500),
96
+ chunk_overlap=data.get('chunk_overlap', 50),
97
+ error=data.get('error')
98
+ )
99
+
100
+
101
+ class ProjectManager:
102
+ """项目管理器 - 负责项目的持久化存储和检索"""
103
+
104
+ # 项目存储根目录
105
+ PROJECTS_DIR = os.path.join(Config.UPLOAD_FOLDER, 'projects')
106
+
107
+ @classmethod
108
+ def _ensure_projects_dir(cls):
109
+ """确保项目目录存在"""
110
+ os.makedirs(cls.PROJECTS_DIR, exist_ok=True)
111
+
112
+ @classmethod
113
+ def _get_project_dir(cls, project_id: str) -> str:
114
+ """获取项目目录路径"""
115
+ return os.path.join(cls.PROJECTS_DIR, project_id)
116
+
117
+ @classmethod
118
+ def _get_project_meta_path(cls, project_id: str) -> str:
119
+ """获取项目元数据文件路径"""
120
+ return os.path.join(cls._get_project_dir(project_id), 'project.json')
121
+
122
+ @classmethod
123
+ def _get_project_files_dir(cls, project_id: str) -> str:
124
+ """获取项目文件存储目录"""
125
+ return os.path.join(cls._get_project_dir(project_id), 'files')
126
+
127
+ @classmethod
128
+ def _get_project_text_path(cls, project_id: str) -> str:
129
+ """获取项目提取文本存储路径"""
130
+ return os.path.join(cls._get_project_dir(project_id), 'extracted_text.txt')
131
+
132
+ @classmethod
133
+ def create_project(cls, name: str = "Unnamed Project") -> Project:
134
+ """
135
+ 创建新项目
136
+
137
+ Args:
138
+ name: 项目名称
139
+
140
+ Returns:
141
+ 新创建的Project对象
142
+ """
143
+ cls._ensure_projects_dir()
144
+
145
+ project_id = f"proj_{uuid.uuid4().hex[:12]}"
146
+ now = datetime.now().isoformat()
147
+
148
+ project = Project(
149
+ project_id=project_id,
150
+ name=name,
151
+ status=ProjectStatus.CREATED,
152
+ created_at=now,
153
+ updated_at=now
154
+ )
155
+
156
+ # 创建项目目录结构
157
+ project_dir = cls._get_project_dir(project_id)
158
+ files_dir = cls._get_project_files_dir(project_id)
159
+ os.makedirs(project_dir, exist_ok=True)
160
+ os.makedirs(files_dir, exist_ok=True)
161
+
162
+ # 保存项目元数据
163
+ cls.save_project(project)
164
+
165
+ return project
166
+
167
+ @classmethod
168
+ def save_project(cls, project: Project) -> None:
169
+ """保存项目元数据"""
170
+ project.updated_at = datetime.now().isoformat()
171
+ meta_path = cls._get_project_meta_path(project.project_id)
172
+
173
+ with open(meta_path, 'w', encoding='utf-8') as f:
174
+ json.dump(project.to_dict(), f, ensure_ascii=False, indent=2)
175
+
176
+ @classmethod
177
+ def get_project(cls, project_id: str) -> Optional[Project]:
178
+ """
179
+ 获取项目
180
+
181
+ Args:
182
+ project_id: 项目ID
183
+
184
+ Returns:
185
+ Project对象,如果不存在返回None
186
+ """
187
+ meta_path = cls._get_project_meta_path(project_id)
188
+
189
+ if not os.path.exists(meta_path):
190
+ return None
191
+
192
+ with open(meta_path, 'r', encoding='utf-8') as f:
193
+ data = json.load(f)
194
+
195
+ return Project.from_dict(data)
196
+
197
+ @classmethod
198
+ def list_projects(cls, limit: int = 50) -> List[Project]:
199
+ """
200
+ 列出所有项目
201
+
202
+ Args:
203
+ limit: 返回数量限制
204
+
205
+ Returns:
206
+ 项目列表,按创建时间倒序
207
+ """
208
+ cls._ensure_projects_dir()
209
+
210
+ projects = []
211
+ for project_id in os.listdir(cls.PROJECTS_DIR):
212
+ project = cls.get_project(project_id)
213
+ if project:
214
+ projects.append(project)
215
+
216
+ # 按创建时间倒序排序
217
+ projects.sort(key=lambda p: p.created_at, reverse=True)
218
+
219
+ return projects[:limit]
220
+
221
+ @classmethod
222
+ def delete_project(cls, project_id: str) -> bool:
223
+ """
224
+ 删除项目及其所有文件
225
+
226
+ Args:
227
+ project_id: 项目ID
228
+
229
+ Returns:
230
+ 是否删除成功
231
+ """
232
+ project_dir = cls._get_project_dir(project_id)
233
+
234
+ if not os.path.exists(project_dir):
235
+ return False
236
+
237
+ shutil.rmtree(project_dir)
238
+ return True
239
+
240
+ @classmethod
241
+ def save_file_to_project(cls, project_id: str, file_storage, original_filename: str) -> Dict[str, str]:
242
+ """
243
+ 保存上传的文件到项目目录
244
+
245
+ Args:
246
+ project_id: 项目ID
247
+ file_storage: Flask的FileStorage对象
248
+ original_filename: 原始文件名
249
+
250
+ Returns:
251
+ 文件信息字典 {filename, path, size}
252
+ """
253
+ files_dir = cls._get_project_files_dir(project_id)
254
+ os.makedirs(files_dir, exist_ok=True)
255
+
256
+ # 生成安全的文件名
257
+ ext = os.path.splitext(original_filename)[1].lower()
258
+ safe_filename = f"{uuid.uuid4().hex[:8]}{ext}"
259
+ file_path = os.path.join(files_dir, safe_filename)
260
+
261
+ # 保存文件
262
+ file_storage.save(file_path)
263
+
264
+ # 获取文件大小
265
+ file_size = os.path.getsize(file_path)
266
+
267
+ return {
268
+ "original_filename": original_filename,
269
+ "saved_filename": safe_filename,
270
+ "path": file_path,
271
+ "size": file_size
272
+ }
273
+
274
+ @classmethod
275
+ def save_extracted_text(cls, project_id: str, text: str) -> None:
276
+ """保存提取的文本"""
277
+ text_path = cls._get_project_text_path(project_id)
278
+ with open(text_path, 'w', encoding='utf-8') as f:
279
+ f.write(text)
280
+
281
+ @classmethod
282
+ def get_extracted_text(cls, project_id: str) -> Optional[str]:
283
+ """获取提取的文本"""
284
+ text_path = cls._get_project_text_path(project_id)
285
+
286
+ if not os.path.exists(text_path):
287
+ return None
288
+
289
+ with open(text_path, 'r', encoding='utf-8') as f:
290
+ return f.read()
291
+
292
+ @classmethod
293
+ def get_project_files(cls, project_id: str) -> List[str]:
294
+ """获取项目的所有文件路径"""
295
+ files_dir = cls._get_project_files_dir(project_id)
296
+
297
+ if not os.path.exists(files_dir):
298
+ return []
299
+
300
+ return [
301
+ os.path.join(files_dir, f)
302
+ for f in os.listdir(files_dir)
303
+ if os.path.isfile(os.path.join(files_dir, f))
304
+ ]
305
+
app/app/models/task.py ADDED
@@ -0,0 +1,184 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 任务状态管理
3
+ 用于跟踪长时间运行的任务(如图谱构建)
4
+ """
5
+
6
+ import uuid
7
+ import threading
8
+ from datetime import datetime
9
+ from enum import Enum
10
+ from typing import Dict, Any, Optional
11
+ from dataclasses import dataclass, field
12
+
13
+
14
+ class TaskStatus(str, Enum):
15
+ """任务状态枚举"""
16
+ PENDING = "pending" # 等待中
17
+ PROCESSING = "processing" # 处理中
18
+ COMPLETED = "completed" # 已完成
19
+ FAILED = "failed" # 失败
20
+
21
+
22
+ @dataclass
23
+ class Task:
24
+ """任务数据类"""
25
+ task_id: str
26
+ task_type: str
27
+ status: TaskStatus
28
+ created_at: datetime
29
+ updated_at: datetime
30
+ progress: int = 0 # 总进度百分比 0-100
31
+ message: str = "" # 状态消息
32
+ result: Optional[Dict] = None # 任务结果
33
+ error: Optional[str] = None # 错误信息
34
+ metadata: Dict = field(default_factory=dict) # 额外元数据
35
+ progress_detail: Dict = field(default_factory=dict) # 详细进度信息
36
+
37
+ def to_dict(self) -> Dict[str, Any]:
38
+ """转换为字典"""
39
+ return {
40
+ "task_id": self.task_id,
41
+ "task_type": self.task_type,
42
+ "status": self.status.value,
43
+ "created_at": self.created_at.isoformat(),
44
+ "updated_at": self.updated_at.isoformat(),
45
+ "progress": self.progress,
46
+ "message": self.message,
47
+ "progress_detail": self.progress_detail,
48
+ "result": self.result,
49
+ "error": self.error,
50
+ "metadata": self.metadata,
51
+ }
52
+
53
+
54
+ class TaskManager:
55
+ """
56
+ 任务管理器
57
+ 线程安全的任务状态管理
58
+ """
59
+
60
+ _instance = None
61
+ _lock = threading.Lock()
62
+
63
+ def __new__(cls):
64
+ """单例模式"""
65
+ if cls._instance is None:
66
+ with cls._lock:
67
+ if cls._instance is None:
68
+ cls._instance = super().__new__(cls)
69
+ cls._instance._tasks: Dict[str, Task] = {}
70
+ cls._instance._task_lock = threading.Lock()
71
+ return cls._instance
72
+
73
+ def create_task(self, task_type: str, metadata: Optional[Dict] = None) -> str:
74
+ """
75
+ 创建新任务
76
+
77
+ Args:
78
+ task_type: 任务类型
79
+ metadata: 额外元数据
80
+
81
+ Returns:
82
+ 任务ID
83
+ """
84
+ task_id = str(uuid.uuid4())
85
+ now = datetime.now()
86
+
87
+ task = Task(
88
+ task_id=task_id,
89
+ task_type=task_type,
90
+ status=TaskStatus.PENDING,
91
+ created_at=now,
92
+ updated_at=now,
93
+ metadata=metadata or {}
94
+ )
95
+
96
+ with self._task_lock:
97
+ self._tasks[task_id] = task
98
+
99
+ return task_id
100
+
101
+ def get_task(self, task_id: str) -> Optional[Task]:
102
+ """获取任务"""
103
+ with self._task_lock:
104
+ return self._tasks.get(task_id)
105
+
106
+ def update_task(
107
+ self,
108
+ task_id: str,
109
+ status: Optional[TaskStatus] = None,
110
+ progress: Optional[int] = None,
111
+ message: Optional[str] = None,
112
+ result: Optional[Dict] = None,
113
+ error: Optional[str] = None,
114
+ progress_detail: Optional[Dict] = None
115
+ ):
116
+ """
117
+ 更新任务状态
118
+
119
+ Args:
120
+ task_id: 任务ID
121
+ status: 新状态
122
+ progress: 进度
123
+ message: 消息
124
+ result: 结果
125
+ error: 错误信息
126
+ progress_detail: 详细进度信息
127
+ """
128
+ with self._task_lock:
129
+ task = self._tasks.get(task_id)
130
+ if task:
131
+ task.updated_at = datetime.now()
132
+ if status is not None:
133
+ task.status = status
134
+ if progress is not None:
135
+ task.progress = progress
136
+ if message is not None:
137
+ task.message = message
138
+ if result is not None:
139
+ task.result = result
140
+ if error is not None:
141
+ task.error = error
142
+ if progress_detail is not None:
143
+ task.progress_detail = progress_detail
144
+
145
+ def complete_task(self, task_id: str, result: Dict):
146
+ """标记任务完成"""
147
+ self.update_task(
148
+ task_id,
149
+ status=TaskStatus.COMPLETED,
150
+ progress=100,
151
+ message="任务完成",
152
+ result=result
153
+ )
154
+
155
+ def fail_task(self, task_id: str, error: str):
156
+ """标记任务失败"""
157
+ self.update_task(
158
+ task_id,
159
+ status=TaskStatus.FAILED,
160
+ message="任务失败",
161
+ error=error
162
+ )
163
+
164
+ def list_tasks(self, task_type: Optional[str] = None) -> list:
165
+ """列出任务"""
166
+ with self._task_lock:
167
+ tasks = list(self._tasks.values())
168
+ if task_type:
169
+ tasks = [t for t in tasks if t.task_type == task_type]
170
+ return [t.to_dict() for t in sorted(tasks, key=lambda x: x.created_at, reverse=True)]
171
+
172
+ def cleanup_old_tasks(self, max_age_hours: int = 24):
173
+ """清理旧任务"""
174
+ from datetime import timedelta
175
+ cutoff = datetime.now() - timedelta(hours=max_age_hours)
176
+
177
+ with self._task_lock:
178
+ old_ids = [
179
+ tid for tid, task in self._tasks.items()
180
+ if task.created_at < cutoff and task.status in [TaskStatus.COMPLETED, TaskStatus.FAILED]
181
+ ]
182
+ for tid in old_ids:
183
+ del self._tasks[tid]
184
+
app/app/services/__init__.py ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 业务服务模块
3
+ """
4
+
5
+ from .ontology_generator import OntologyGenerator
6
+ from .graph_builder import GraphBuilderService
7
+ from .text_processor import TextProcessor
8
+ from .zep_entity_reader import ZepEntityReader, EntityNode, FilteredEntities
9
+ from .oasis_profile_generator import OasisProfileGenerator, OasisAgentProfile
10
+ from .simulation_manager import SimulationManager, SimulationState, SimulationStatus
11
+ from .simulation_config_generator import (
12
+ SimulationConfigGenerator,
13
+ SimulationParameters,
14
+ AgentActivityConfig,
15
+ TimeSimulationConfig,
16
+ EventConfig,
17
+ PlatformConfig
18
+ )
19
+ from .simulation_runner import (
20
+ SimulationRunner,
21
+ SimulationRunState,
22
+ RunnerStatus,
23
+ AgentAction,
24
+ RoundSummary
25
+ )
26
+ from .zep_graph_memory_updater import (
27
+ ZepGraphMemoryUpdater,
28
+ ZepGraphMemoryManager,
29
+ AgentActivity
30
+ )
31
+ from .simulation_ipc import (
32
+ SimulationIPCClient,
33
+ SimulationIPCServer,
34
+ IPCCommand,
35
+ IPCResponse,
36
+ CommandType,
37
+ CommandStatus
38
+ )
39
+
40
+ __all__ = [
41
+ 'OntologyGenerator',
42
+ 'GraphBuilderService',
43
+ 'TextProcessor',
44
+ 'ZepEntityReader',
45
+ 'EntityNode',
46
+ 'FilteredEntities',
47
+ 'OasisProfileGenerator',
48
+ 'OasisAgentProfile',
49
+ 'SimulationManager',
50
+ 'SimulationState',
51
+ 'SimulationStatus',
52
+ 'SimulationConfigGenerator',
53
+ 'SimulationParameters',
54
+ 'AgentActivityConfig',
55
+ 'TimeSimulationConfig',
56
+ 'EventConfig',
57
+ 'PlatformConfig',
58
+ 'SimulationRunner',
59
+ 'SimulationRunState',
60
+ 'RunnerStatus',
61
+ 'AgentAction',
62
+ 'RoundSummary',
63
+ 'ZepGraphMemoryUpdater',
64
+ 'ZepGraphMemoryManager',
65
+ 'AgentActivity',
66
+ 'SimulationIPCClient',
67
+ 'SimulationIPCServer',
68
+ 'IPCCommand',
69
+ 'IPCResponse',
70
+ 'CommandType',
71
+ 'CommandStatus',
72
+ ]
73
+
app/app/services/__pycache__/__init__.cpython-311.pyc ADDED
Binary file (1.95 kB). View file
 
app/app/services/__pycache__/graph_builder.cpython-311.pyc ADDED
Binary file (23.2 kB). View file
 
app/app/services/__pycache__/oasis_profile_generator.cpython-311.pyc ADDED
Binary file (62.6 kB). View file
 
app/app/services/__pycache__/ontology_generator.cpython-311.pyc ADDED
Binary file (22.1 kB). View file
 
app/app/services/__pycache__/pipeline_orchestrator.cpython-311.pyc ADDED
Binary file (15 kB). View file
 
app/app/services/__pycache__/report_agent.cpython-311.pyc ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e0abd3c1a2d5a1d6fb13ad578076ad09dfa74c527e44a79114bcfb8c101ea788
3
+ size 108041
app/app/services/__pycache__/simulation_config_generator.cpython-311.pyc ADDED
Binary file (44.8 kB). View file
 
app/app/services/__pycache__/simulation_ipc.cpython-311.pyc ADDED
Binary file (20.2 kB). View file
 
app/app/services/__pycache__/simulation_manager.cpython-311.pyc ADDED
Binary file (23.5 kB). View file
 
app/app/services/__pycache__/simulation_runner.cpython-311.pyc ADDED
Binary file (78.1 kB). View file
 
app/app/services/__pycache__/supabase_jobs.cpython-311.pyc ADDED
Binary file (11 kB). View file
 
app/app/services/__pycache__/text_processor.cpython-311.pyc ADDED
Binary file (3.28 kB). View file
 
app/app/services/__pycache__/zep_entity_reader.cpython-311.pyc ADDED
Binary file (18.6 kB). View file
 
app/app/services/__pycache__/zep_graph_memory_updater.cpython-311.pyc ADDED
Binary file (29.1 kB). View file
 
app/app/services/__pycache__/zep_tools.cpython-311.pyc ADDED
Binary file (82.4 kB). View file
 
app/app/services/billing_service.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Razorpay orders, signature verification, coupon checks, and simulation payment JWTs."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import time
6
+ import uuid
7
+
8
+ import jwt
9
+
10
+ from ..config import Config
11
+
12
+
13
+ def issue_simulation_payment_token(*, via: str) -> str:
14
+ """via: 'razorpay' | 'coupon'"""
15
+ now = int(time.time())
16
+ payload = {
17
+ "typ": "sim_pay",
18
+ "via": via,
19
+ "iat": now,
20
+ "exp": now + 3600,
21
+ "jti": uuid.uuid4().hex,
22
+ }
23
+ return jwt.encode(payload, Config.payment_signing_secret(), algorithm="HS256")
24
+
25
+
26
+ def verify_simulation_payment_token(token: str | None) -> bool:
27
+ if not token or not str(token).strip():
28
+ return False
29
+ try:
30
+ payload = jwt.decode(
31
+ str(token).strip(),
32
+ Config.payment_signing_secret(),
33
+ algorithms=["HS256"],
34
+ )
35
+ return payload.get("typ") == "sim_pay"
36
+ except jwt.PyJWTError:
37
+ return False
38
+
39
+
40
+ def create_razorpay_order() -> dict:
41
+ import razorpay
42
+
43
+ client = razorpay.Client(auth=(Config.RAZORPAY_KEY_ID, Config.RAZORPAY_KEY_SECRET))
44
+ receipt = f"mf_{uuid.uuid4().hex[:24]}"
45
+ order = client.order.create(
46
+ {
47
+ "amount": Config.RAZORPAY_AMOUNT_PAISE,
48
+ "currency": Config.RAZORPAY_CURRENCY,
49
+ "receipt": receipt,
50
+ }
51
+ )
52
+ return {
53
+ "order_id": order["id"],
54
+ "amount": order["amount"],
55
+ "currency": order["currency"],
56
+ "key_id": Config.RAZORPAY_KEY_ID,
57
+ }
58
+
59
+
60
+ def verify_razorpay_payment_signature(
61
+ razorpay_order_id: str,
62
+ razorpay_payment_id: str,
63
+ razorpay_signature: str,
64
+ ) -> bool:
65
+ import razorpay
66
+
67
+ client = razorpay.Client(auth=(Config.RAZORPAY_KEY_ID, Config.RAZORPAY_KEY_SECRET))
68
+ try:
69
+ client.utility.verify_payment_signature(
70
+ {
71
+ "razorpay_order_id": razorpay_order_id,
72
+ "razorpay_payment_id": razorpay_payment_id,
73
+ "razorpay_signature": razorpay_signature,
74
+ }
75
+ )
76
+ return True
77
+ except Exception:
78
+ return False
79
+
80
+
81
+ def coupon_valid(code: str | None) -> bool:
82
+ if not code:
83
+ return False
84
+ normalized = code.strip().lower()
85
+ return normalized in Config.payment_coupon_set()
app/app/services/graph_builder.py ADDED
@@ -0,0 +1,532 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 图谱构建服务
3
+ 接口2:使用Zep API构建Standalone Graph
4
+ """
5
+
6
+ import os
7
+ import uuid
8
+ import time
9
+ import threading
10
+ from typing import Dict, Any, List, Optional, Callable
11
+ from dataclasses import dataclass
12
+
13
+ from zep_cloud.client import Zep
14
+ from zep_cloud import EpisodeData, EntityEdgeSourceTarget
15
+
16
+ from ..config import Config
17
+ from ..models.task import TaskManager, TaskStatus
18
+ from ..utils.api_key_runtime import call_with_limit_rotation, get_zep_key
19
+ from ..utils.zep_paging import fetch_all_nodes, fetch_all_edges
20
+ from .text_processor import TextProcessor
21
+
22
+
23
+ @dataclass
24
+ class GraphInfo:
25
+ """图谱信息"""
26
+ graph_id: str
27
+ node_count: int
28
+ edge_count: int
29
+ entity_types: List[str]
30
+
31
+ def to_dict(self) -> Dict[str, Any]:
32
+ return {
33
+ "graph_id": self.graph_id,
34
+ "node_count": self.node_count,
35
+ "edge_count": self.edge_count,
36
+ "entity_types": self.entity_types,
37
+ }
38
+
39
+
40
+ class GraphBuilderService:
41
+ """
42
+ 图谱构建服务
43
+ 负责调用Zep API构建知识图谱
44
+ """
45
+
46
+ def __init__(self, api_key: Optional[str] = None):
47
+ self._fixed_api_key = (api_key or "").strip() or None
48
+ if self._fixed_api_key:
49
+ self.api_key = self._fixed_api_key
50
+ else:
51
+ self.api_key = get_zep_key() or (Config.get_zep_api_keys()[0] if Config.get_zep_api_keys() else None)
52
+ if not self.api_key and not Config.get_zep_api_keys():
53
+ raise ValueError("ZEP_API_KEY 未配置")
54
+
55
+ self.client = Zep(api_key=self.api_key) if self._fixed_api_key else None # noqa: SLF001
56
+ self.task_manager = TaskManager()
57
+
58
+ def _zep_call(self, fn):
59
+ """Zep API call with shared LLM/Zep rate-limit rotation (pipeline uses runtime keys)."""
60
+ if self._fixed_api_key:
61
+ return fn(Zep(api_key=self._fixed_api_key))
62
+
63
+ def run():
64
+ k = get_zep_key()
65
+ if not k:
66
+ raise ValueError("ZEP_API_KEY 未配置")
67
+ return fn(Zep(api_key=k))
68
+
69
+ return call_with_limit_rotation(run)
70
+
71
+ def build_graph_async(
72
+ self,
73
+ text: str,
74
+ ontology: Dict[str, Any],
75
+ graph_name: str = "MiroFish Graph",
76
+ chunk_size: int = 500,
77
+ chunk_overlap: int = 50,
78
+ batch_size: int = 3
79
+ ) -> str:
80
+ """
81
+ 异步构建图谱
82
+
83
+ Args:
84
+ text: 输入文本
85
+ ontology: 本体定义(来自接口1的输出)
86
+ graph_name: 图谱名称
87
+ chunk_size: 文本块大小
88
+ chunk_overlap: 块重叠大小
89
+ batch_size: 每批发送的块数量
90
+
91
+ Returns:
92
+ 任务ID
93
+ """
94
+ # 创建任务
95
+ task_id = self.task_manager.create_task(
96
+ task_type="graph_build",
97
+ metadata={
98
+ "graph_name": graph_name,
99
+ "chunk_size": chunk_size,
100
+ "text_length": len(text),
101
+ }
102
+ )
103
+
104
+ # 在后台线程中执行构建
105
+ thread = threading.Thread(
106
+ target=self._build_graph_worker,
107
+ args=(task_id, text, ontology, graph_name, chunk_size, chunk_overlap, batch_size)
108
+ )
109
+ thread.daemon = True
110
+ thread.start()
111
+
112
+ return task_id
113
+
114
+ def _build_graph_worker(
115
+ self,
116
+ task_id: str,
117
+ text: str,
118
+ ontology: Dict[str, Any],
119
+ graph_name: str,
120
+ chunk_size: int,
121
+ chunk_overlap: int,
122
+ batch_size: int
123
+ ):
124
+ """图谱构建工作线程"""
125
+ try:
126
+ self.task_manager.update_task(
127
+ task_id,
128
+ status=TaskStatus.PROCESSING,
129
+ progress=5,
130
+ message="开始构建图谱..."
131
+ )
132
+
133
+ # 1. 创建图谱
134
+ graph_id = self.create_graph(graph_name)
135
+ self.task_manager.update_task(
136
+ task_id,
137
+ progress=10,
138
+ message=f"图谱已创建: {graph_id}"
139
+ )
140
+
141
+ # 2. 设置本体
142
+ self.set_ontology(graph_id, ontology)
143
+ self.task_manager.update_task(
144
+ task_id,
145
+ progress=15,
146
+ message="本体已设置"
147
+ )
148
+
149
+ # 3. 文本分块
150
+ chunks = TextProcessor.split_text(text, chunk_size, chunk_overlap)
151
+ total_chunks = len(chunks)
152
+ self.task_manager.update_task(
153
+ task_id,
154
+ progress=20,
155
+ message=f"文本已分割为 {total_chunks} 个块"
156
+ )
157
+
158
+ # 4. 分批发送数据
159
+ episode_uuids = self.add_text_batches(
160
+ graph_id, chunks, batch_size,
161
+ lambda msg, prog: self.task_manager.update_task(
162
+ task_id,
163
+ progress=20 + int(prog * 0.4), # 20-60%
164
+ message=msg
165
+ )
166
+ )
167
+
168
+ # 5. 等待Zep处理完成
169
+ self.task_manager.update_task(
170
+ task_id,
171
+ progress=60,
172
+ message="等待Zep处理数据..."
173
+ )
174
+
175
+ self._wait_for_episodes(
176
+ episode_uuids,
177
+ lambda msg, prog: self.task_manager.update_task(
178
+ task_id,
179
+ progress=60 + int(prog * 0.3), # 60-90%
180
+ message=msg
181
+ )
182
+ )
183
+
184
+ # 6. 获取图谱信息
185
+ self.task_manager.update_task(
186
+ task_id,
187
+ progress=90,
188
+ message="获取图谱信息..."
189
+ )
190
+
191
+ graph_info = self._get_graph_info(graph_id)
192
+
193
+ # 完成
194
+ self.task_manager.complete_task(task_id, {
195
+ "graph_id": graph_id,
196
+ "graph_info": graph_info.to_dict(),
197
+ "chunks_processed": total_chunks,
198
+ })
199
+
200
+ except Exception as e:
201
+ import traceback
202
+ error_msg = f"{str(e)}\n{traceback.format_exc()}"
203
+ self.task_manager.fail_task(task_id, error_msg)
204
+
205
+ def create_graph(self, name: str) -> str:
206
+ """创建Zep图谱(公开方法)"""
207
+ graph_id = f"mirofish_{uuid.uuid4().hex[:16]}"
208
+
209
+ def op(c: Zep):
210
+ c.graph.create(
211
+ graph_id=graph_id,
212
+ name=name,
213
+ description="MiroFish Social Simulation Graph",
214
+ )
215
+
216
+ self._zep_call(op)
217
+ return graph_id
218
+
219
+ def set_ontology(self, graph_id: str, ontology: Dict[str, Any]):
220
+ """设置图谱本体(公开方法)"""
221
+ import warnings
222
+ from typing import Optional
223
+ from pydantic import Field
224
+ from zep_cloud.external_clients.ontology import EntityModel, EntityText, EdgeModel
225
+
226
+ # 抑制 Pydantic v2 关于 Field(default=None) 的警告
227
+ # 这是 Zep SDK 要求的用法,警告来自动态类创建,可以安全忽略
228
+ warnings.filterwarnings('ignore', category=UserWarning, module='pydantic')
229
+
230
+ # Zep 保留名称,不能作为属性名
231
+ RESERVED_NAMES = {'uuid', 'name', 'group_id', 'name_embedding', 'summary', 'created_at'}
232
+
233
+ def safe_attr_name(attr_name: str) -> str:
234
+ """将保留名称转换为安全名称"""
235
+ if attr_name.lower() in RESERVED_NAMES:
236
+ return f"entity_{attr_name}"
237
+ return attr_name
238
+
239
+ # 动态创建实体类型
240
+ entity_types = {}
241
+ for entity_def in ontology.get("entity_types", []):
242
+ name = entity_def["name"]
243
+ description = entity_def.get("description", f"A {name} entity.")
244
+
245
+ # 创建属性字典和类型注解(Pydantic v2 需要)
246
+ attrs = {"__doc__": description}
247
+ annotations = {}
248
+
249
+ for attr_def in entity_def.get("attributes", []):
250
+ attr_name = safe_attr_name(attr_def["name"]) # 使用安全名称
251
+ attr_desc = attr_def.get("description", attr_name)
252
+ # Zep API 需要 Field 的 description,这是必需的
253
+ attrs[attr_name] = Field(description=attr_desc, default=None)
254
+ annotations[attr_name] = Optional[EntityText] # 类型注解
255
+
256
+ attrs["__annotations__"] = annotations
257
+
258
+ # 动态创建类
259
+ entity_class = type(name, (EntityModel,), attrs)
260
+ entity_class.__doc__ = description
261
+ entity_types[name] = entity_class
262
+
263
+ # 动态创建边类型
264
+ edge_definitions = {}
265
+ for edge_def in ontology.get("edge_types", []):
266
+ name = edge_def["name"]
267
+ description = edge_def.get("description", f"A {name} relationship.")
268
+
269
+ # 创建属性字典和类型注解
270
+ attrs = {"__doc__": description}
271
+ annotations = {}
272
+
273
+ for attr_def in edge_def.get("attributes", []):
274
+ attr_name = safe_attr_name(attr_def["name"]) # 使用安全名称
275
+ attr_desc = attr_def.get("description", attr_name)
276
+ # Zep API 需要 Field 的 description,这是必需的
277
+ attrs[attr_name] = Field(description=attr_desc, default=None)
278
+ annotations[attr_name] = Optional[str] # 边属性用str类型
279
+
280
+ attrs["__annotations__"] = annotations
281
+
282
+ # 动态创建类
283
+ class_name = ''.join(word.capitalize() for word in name.split('_'))
284
+ edge_class = type(class_name, (EdgeModel,), attrs)
285
+ edge_class.__doc__ = description
286
+
287
+ # 构建source_targets
288
+ source_targets = []
289
+ for st in edge_def.get("source_targets", []):
290
+ source_targets.append(
291
+ EntityEdgeSourceTarget(
292
+ source=st.get("source", "Entity"),
293
+ target=st.get("target", "Entity")
294
+ )
295
+ )
296
+
297
+ if source_targets:
298
+ edge_definitions[name] = (edge_class, source_targets)
299
+
300
+ # 调用Zep API设置本体
301
+ if entity_types or edge_definitions:
302
+ def op(c: Zep):
303
+ c.graph.set_ontology(
304
+ graph_ids=[graph_id],
305
+ entities=entity_types if entity_types else None,
306
+ edges=edge_definitions if edge_definitions else None,
307
+ )
308
+
309
+ self._zep_call(op)
310
+
311
+ def add_text_batches(
312
+ self,
313
+ graph_id: str,
314
+ chunks: List[str],
315
+ batch_size: int = 3,
316
+ progress_callback: Optional[Callable] = None
317
+ ) -> List[str]:
318
+ """分批添加文本到图谱,返回所有 episode 的 uuid 列表"""
319
+ episode_uuids = []
320
+ total_chunks = len(chunks)
321
+
322
+ for i in range(0, total_chunks, batch_size):
323
+ batch_chunks = chunks[i:i + batch_size]
324
+ batch_num = i // batch_size + 1
325
+ total_batches = (total_chunks + batch_size - 1) // batch_size
326
+
327
+ if progress_callback:
328
+ progress = (i + len(batch_chunks)) / total_chunks
329
+ progress_callback(
330
+ f"发送第 {batch_num}/{total_batches} 批数据 ({len(batch_chunks)} 块)...",
331
+ progress
332
+ )
333
+
334
+ # 构建episode数据
335
+ episodes = [
336
+ EpisodeData(data=chunk, type="text")
337
+ for chunk in batch_chunks
338
+ ]
339
+
340
+ # 发送到Zep
341
+ try:
342
+ def op(c: Zep):
343
+ return c.graph.add_batch(graph_id=graph_id, episodes=episodes)
344
+
345
+ batch_result = self._zep_call(op)
346
+
347
+ # 收集返回的 episode uuid
348
+ if batch_result and isinstance(batch_result, list):
349
+ for ep in batch_result:
350
+ ep_uuid = getattr(ep, 'uuid_', None) or getattr(ep, 'uuid', None)
351
+ if ep_uuid:
352
+ episode_uuids.append(ep_uuid)
353
+
354
+ # 避免请求过快
355
+ time.sleep(1)
356
+
357
+ except Exception as e:
358
+ if progress_callback:
359
+ progress_callback(f"批次 {batch_num} 发送失败: {str(e)}", 0)
360
+ raise
361
+
362
+ return episode_uuids
363
+
364
+ def _wait_for_episodes(
365
+ self,
366
+ episode_uuids: List[str],
367
+ progress_callback: Optional[Callable] = None,
368
+ timeout: int = 600
369
+ ):
370
+ """等待所有 episode 处理完成(通过查询每个 episode 的 processed 状态)"""
371
+ if not episode_uuids:
372
+ if progress_callback:
373
+ progress_callback("无需等待(没有 episode)", 1.0)
374
+ return
375
+
376
+ start_time = time.time()
377
+ pending_episodes = set(episode_uuids)
378
+ completed_count = 0
379
+ total_episodes = len(episode_uuids)
380
+
381
+ if progress_callback:
382
+ progress_callback(f"开始等待 {total_episodes} 个文本块处理...", 0)
383
+
384
+ while pending_episodes:
385
+ if time.time() - start_time > timeout:
386
+ if progress_callback:
387
+ progress_callback(
388
+ f"部分文本块超时,已完成 {completed_count}/{total_episodes}",
389
+ completed_count / total_episodes
390
+ )
391
+ break
392
+
393
+ # 检查每个 episode 的处理状态
394
+ for ep_uuid in list(pending_episodes):
395
+ try:
396
+ def op(c: Zep):
397
+ return c.graph.episode.get(uuid_=ep_uuid)
398
+
399
+ episode = self._zep_call(op)
400
+ is_processed = getattr(episode, 'processed', False)
401
+
402
+ if is_processed:
403
+ pending_episodes.remove(ep_uuid)
404
+ completed_count += 1
405
+
406
+ except Exception as e:
407
+ # 忽略单个查询错误,继续
408
+ pass
409
+
410
+ elapsed = int(time.time() - start_time)
411
+ if progress_callback:
412
+ progress_callback(
413
+ f"Zep处理中... {completed_count}/{total_episodes} 完成, {len(pending_episodes)} 待处理 ({elapsed}秒)",
414
+ completed_count / total_episodes if total_episodes > 0 else 0
415
+ )
416
+
417
+ if pending_episodes:
418
+ time.sleep(3) # 每3秒检查一次
419
+
420
+ if progress_callback:
421
+ progress_callback(f"处理完成: {completed_count}/{total_episodes}", 1.0)
422
+
423
+ def _get_graph_info(self, graph_id: str) -> GraphInfo:
424
+ """获取图谱信息"""
425
+
426
+ def fetch_pair(c: Zep):
427
+ return fetch_all_nodes(c, graph_id), fetch_all_edges(c, graph_id)
428
+
429
+ nodes, edges = self._zep_call(fetch_pair)
430
+
431
+ # 统计实体类型
432
+ entity_types = set()
433
+ for node in nodes:
434
+ if node.labels:
435
+ for label in node.labels:
436
+ if label not in ["Entity", "Node"]:
437
+ entity_types.add(label)
438
+
439
+ return GraphInfo(
440
+ graph_id=graph_id,
441
+ node_count=len(nodes),
442
+ edge_count=len(edges),
443
+ entity_types=list(entity_types)
444
+ )
445
+
446
+ def get_graph_data(self, graph_id: str) -> Dict[str, Any]:
447
+ """
448
+ 获取完整图谱数据(包含详细信息)
449
+
450
+ Args:
451
+ graph_id: 图谱ID
452
+
453
+ Returns:
454
+ 包含nodes和edges的字典,包括时间信息、属性等详细数据
455
+ """
456
+ def fetch_pair(c: Zep):
457
+ return fetch_all_nodes(c, graph_id), fetch_all_edges(c, graph_id)
458
+
459
+ nodes, edges = self._zep_call(fetch_pair)
460
+
461
+ # 创建节点映射用于获取节点名称
462
+ node_map = {}
463
+ for node in nodes:
464
+ node_map[node.uuid_] = node.name or ""
465
+
466
+ nodes_data = []
467
+ for node in nodes:
468
+ # 获取创建时间
469
+ created_at = getattr(node, 'created_at', None)
470
+ if created_at:
471
+ created_at = str(created_at)
472
+
473
+ nodes_data.append({
474
+ "uuid": node.uuid_,
475
+ "name": node.name,
476
+ "labels": node.labels or [],
477
+ "summary": node.summary or "",
478
+ "attributes": node.attributes or {},
479
+ "created_at": created_at,
480
+ })
481
+
482
+ edges_data = []
483
+ for edge in edges:
484
+ # 获取时间信息
485
+ created_at = getattr(edge, 'created_at', None)
486
+ valid_at = getattr(edge, 'valid_at', None)
487
+ invalid_at = getattr(edge, 'invalid_at', None)
488
+ expired_at = getattr(edge, 'expired_at', None)
489
+
490
+ # 获取 episodes
491
+ episodes = getattr(edge, 'episodes', None) or getattr(edge, 'episode_ids', None)
492
+ if episodes and not isinstance(episodes, list):
493
+ episodes = [str(episodes)]
494
+ elif episodes:
495
+ episodes = [str(e) for e in episodes]
496
+
497
+ # 获取 fact_type
498
+ fact_type = getattr(edge, 'fact_type', None) or edge.name or ""
499
+
500
+ edges_data.append({
501
+ "uuid": edge.uuid_,
502
+ "name": edge.name or "",
503
+ "fact": edge.fact or "",
504
+ "fact_type": fact_type,
505
+ "source_node_uuid": edge.source_node_uuid,
506
+ "target_node_uuid": edge.target_node_uuid,
507
+ "source_node_name": node_map.get(edge.source_node_uuid, ""),
508
+ "target_node_name": node_map.get(edge.target_node_uuid, ""),
509
+ "attributes": edge.attributes or {},
510
+ "created_at": str(created_at) if created_at else None,
511
+ "valid_at": str(valid_at) if valid_at else None,
512
+ "invalid_at": str(invalid_at) if invalid_at else None,
513
+ "expired_at": str(expired_at) if expired_at else None,
514
+ "episodes": episodes or [],
515
+ })
516
+
517
+ return {
518
+ "graph_id": graph_id,
519
+ "nodes": nodes_data,
520
+ "edges": edges_data,
521
+ "node_count": len(nodes_data),
522
+ "edge_count": len(edges_data),
523
+ }
524
+
525
+ def delete_graph(self, graph_id: str):
526
+ """删除图谱"""
527
+
528
+ def op(c: Zep):
529
+ c.graph.delete(graph_id=graph_id)
530
+
531
+ self._zep_call(op)
532
+
app/app/services/oasis_profile_generator.py ADDED
@@ -0,0 +1,1339 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ OASIS Agent Profile生成器
3
+ 将Zep图谱中的实体转换为OASIS模拟平台所需的Agent Profile格式
4
+
5
+ 优化改进:
6
+ 1. 调用Zep检索功能二次丰富节点信息
7
+ 2. 优化提示词生成非常详细的人设
8
+ 3. 区分个人实体和抽象群体实体
9
+ """
10
+
11
+ import json
12
+ import random
13
+ import time
14
+ from typing import Dict, Any, List, Optional
15
+ from dataclasses import dataclass, field
16
+ from datetime import datetime
17
+
18
+ from openai import OpenAI
19
+ from zep_cloud.client import Zep
20
+
21
+ from ..config import Config
22
+ from ..utils.api_key_runtime import call_with_limit_rotation, get_llm_key, get_zep_key
23
+ from ..utils.logger import get_logger
24
+ from .zep_entity_reader import EntityNode, ZepEntityReader
25
+
26
+ logger = get_logger('mirofish.oasis_profile')
27
+
28
+
29
+ @dataclass
30
+ class OasisAgentProfile:
31
+ """OASIS Agent Profile数据结构"""
32
+ # 通用字段
33
+ user_id: int
34
+ user_name: str
35
+ name: str
36
+ bio: str
37
+ persona: str
38
+
39
+ # 可选字段 - Reddit风格
40
+ karma: int = 1000
41
+
42
+ # 可选字段 - Twitter风格
43
+ friend_count: int = 100
44
+ follower_count: int = 150
45
+ statuses_count: int = 500
46
+
47
+ # 额外人设信息
48
+ age: Optional[int] = None
49
+ gender: Optional[str] = None
50
+ mbti: Optional[str] = None
51
+ country: Optional[str] = None
52
+ profession: Optional[str] = None
53
+ interested_topics: List[str] = field(default_factory=list)
54
+
55
+ # 来源实体信息
56
+ source_entity_uuid: Optional[str] = None
57
+ source_entity_type: Optional[str] = None
58
+
59
+ created_at: str = field(default_factory=lambda: datetime.now().strftime("%Y-%m-%d"))
60
+
61
+ def to_reddit_format(self) -> Dict[str, Any]:
62
+ """转换为Reddit平台格式"""
63
+ profile = {
64
+ "user_id": self.user_id,
65
+ "username": self.user_name, # OASIS 库要求字段名为 username(无下划线)
66
+ "name": self.name,
67
+ "bio": self.bio,
68
+ "persona": self.persona,
69
+ "karma": self.karma,
70
+ "created_at": self.created_at,
71
+ }
72
+
73
+ # 添加额外人设信息(如果有)
74
+ if self.age:
75
+ profile["age"] = self.age
76
+ if self.gender:
77
+ profile["gender"] = self.gender
78
+ if self.mbti:
79
+ profile["mbti"] = self.mbti
80
+ if self.country:
81
+ profile["country"] = self.country
82
+ if self.profession:
83
+ profile["profession"] = self.profession
84
+ if self.interested_topics:
85
+ profile["interested_topics"] = self.interested_topics
86
+
87
+ return profile
88
+
89
+ def to_twitter_format(self) -> Dict[str, Any]:
90
+ """转换为Twitter平台格式"""
91
+ profile = {
92
+ "user_id": self.user_id,
93
+ "username": self.user_name, # OASIS 库要求字段名为 username(无下划线)
94
+ "name": self.name,
95
+ "bio": self.bio,
96
+ "persona": self.persona,
97
+ "friend_count": self.friend_count,
98
+ "follower_count": self.follower_count,
99
+ "statuses_count": self.statuses_count,
100
+ "created_at": self.created_at,
101
+ }
102
+
103
+ # 添加额外人设信息
104
+ if self.age:
105
+ profile["age"] = self.age
106
+ if self.gender:
107
+ profile["gender"] = self.gender
108
+ if self.mbti:
109
+ profile["mbti"] = self.mbti
110
+ if self.country:
111
+ profile["country"] = self.country
112
+ if self.profession:
113
+ profile["profession"] = self.profession
114
+ if self.interested_topics:
115
+ profile["interested_topics"] = self.interested_topics
116
+
117
+ return profile
118
+
119
+ def to_dict(self) -> Dict[str, Any]:
120
+ """转换为完整字典格式"""
121
+ return {
122
+ "user_id": self.user_id,
123
+ "user_name": self.user_name,
124
+ "name": self.name,
125
+ "bio": self.bio,
126
+ "persona": self.persona,
127
+ "karma": self.karma,
128
+ "friend_count": self.friend_count,
129
+ "follower_count": self.follower_count,
130
+ "statuses_count": self.statuses_count,
131
+ "age": self.age,
132
+ "gender": self.gender,
133
+ "mbti": self.mbti,
134
+ "country": self.country,
135
+ "profession": self.profession,
136
+ "interested_topics": self.interested_topics,
137
+ "source_entity_uuid": self.source_entity_uuid,
138
+ "source_entity_type": self.source_entity_type,
139
+ "created_at": self.created_at,
140
+ }
141
+
142
+
143
+ class OasisProfileGenerator:
144
+ """
145
+ OASIS Profile生成器
146
+
147
+ 将Zep图谱中的实体转换为OASIS模拟所需的Agent Profile
148
+
149
+ 优化特性:
150
+ 1. 调用Zep图谱检索功能获取更丰富的上下文
151
+ 2. 生成非常详细的人设(包括基本信息、职业经历、性格特征、社交媒体行为等)
152
+ 3. 区分个人实体和抽象群体实体
153
+ """
154
+
155
+ # MBTI类型列表
156
+ MBTI_TYPES = [
157
+ "INTJ", "INTP", "ENTJ", "ENTP",
158
+ "INFJ", "INFP", "ENFJ", "ENFP",
159
+ "ISTJ", "ISFJ", "ESTJ", "ESFJ",
160
+ "ISTP", "ISFP", "ESTP", "ESFP"
161
+ ]
162
+
163
+ # 常见国家列表
164
+ COUNTRIES = [
165
+ "China", "US", "UK", "Japan", "Germany", "France",
166
+ "Canada", "Australia", "Brazil", "India", "South Korea"
167
+ ]
168
+
169
+ # 个人类型实体(需要生成具体人设)
170
+ INDIVIDUAL_ENTITY_TYPES = [
171
+ "student", "alumni", "professor", "person", "publicfigure",
172
+ "expert", "faculty", "official", "journalist", "activist"
173
+ ]
174
+
175
+ # 群体/机构类型实体(需要生成群体代表人设)
176
+ GROUP_ENTITY_TYPES = [
177
+ "university", "governmentagency", "organization", "ngo",
178
+ "mediaoutlet", "company", "institution", "group", "community"
179
+ ]
180
+
181
+ def __init__(
182
+ self,
183
+ api_key: Optional[str] = None,
184
+ base_url: Optional[str] = None,
185
+ model_name: Optional[str] = None,
186
+ zep_api_key: Optional[str] = None,
187
+ graph_id: Optional[str] = None
188
+ ):
189
+ self._fixed_llm_key = (api_key or "").strip() or None
190
+ self.base_url = base_url or Config.LLM_BASE_URL
191
+ self.model_name = model_name or Config.LLM_MODEL_NAME
192
+
193
+ if self._fixed_llm_key:
194
+ self.api_key = self._fixed_llm_key
195
+ self.client = OpenAI(api_key=self.api_key, base_url=self.base_url)
196
+ else:
197
+ if not Config.get_llm_api_keys():
198
+ raise ValueError("LLM_API_KEY 未配置")
199
+ self.api_key = Config.get_llm_api_keys()[0]
200
+ self.client = None
201
+
202
+ # Zep客户端用于检索丰富上下文(运行时与 pipeline 共用 key slot)
203
+ self._fixed_zep_key = (zep_api_key or "").strip() or None
204
+ self.zep_api_key = self._fixed_zep_key or (
205
+ Config.get_zep_api_keys()[0] if Config.get_zep_api_keys() else None
206
+ )
207
+ self.zep_client = None
208
+ self.graph_id = graph_id
209
+
210
+ if self._fixed_zep_key:
211
+ try:
212
+ self.zep_client = Zep(api_key=self._fixed_zep_key)
213
+ except Exception as e:
214
+ logger.warning(f"Zep客户端初始化失败: {e}")
215
+
216
+ def generate_profile_from_entity(
217
+ self,
218
+ entity: EntityNode,
219
+ user_id: int,
220
+ use_llm: bool = True
221
+ ) -> OasisAgentProfile:
222
+ """
223
+ 从Zep实体生成OASIS Agent Profile
224
+
225
+ Args:
226
+ entity: Zep实体节点
227
+ user_id: 用户ID(用于OASIS)
228
+ use_llm: 是否使用LLM生成详细人设
229
+
230
+ Returns:
231
+ OasisAgentProfile
232
+ """
233
+ entity_type = entity.get_entity_type() or "Entity"
234
+
235
+ # 基础信息
236
+ name = entity.name
237
+ user_name = self._generate_username(name)
238
+
239
+ # 构建上下文信息
240
+ context = self._build_entity_context(entity)
241
+
242
+ if use_llm:
243
+ # 使用LLM生成详细人设
244
+ profile_data = self._generate_profile_with_llm(
245
+ entity_name=name,
246
+ entity_type=entity_type,
247
+ entity_summary=entity.summary,
248
+ entity_attributes=entity.attributes,
249
+ context=context
250
+ )
251
+ else:
252
+ # 使用规则生成基础人设
253
+ profile_data = self._generate_profile_rule_based(
254
+ entity_name=name,
255
+ entity_type=entity_type,
256
+ entity_summary=entity.summary,
257
+ entity_attributes=entity.attributes
258
+ )
259
+
260
+ return OasisAgentProfile(
261
+ user_id=user_id,
262
+ user_name=user_name,
263
+ name=name,
264
+ bio=profile_data.get("bio", f"{entity_type}: {name}"),
265
+ persona=profile_data.get("persona", entity.summary or f"A {entity_type} named {name}."),
266
+ karma=profile_data.get("karma", random.randint(500, 5000)),
267
+ friend_count=profile_data.get("friend_count", random.randint(50, 500)),
268
+ follower_count=profile_data.get("follower_count", random.randint(100, 1000)),
269
+ statuses_count=profile_data.get("statuses_count", random.randint(100, 2000)),
270
+ age=profile_data.get("age"),
271
+ gender=profile_data.get("gender"),
272
+ mbti=profile_data.get("mbti"),
273
+ country=profile_data.get("country"),
274
+ profession=profile_data.get("profession"),
275
+ interested_topics=profile_data.get("interested_topics", []),
276
+ source_entity_uuid=entity.uuid,
277
+ source_entity_type=entity_type,
278
+ )
279
+
280
+ def _generate_username(self, name: str) -> str:
281
+ """生成用户名"""
282
+ # 移除特殊字符,转换为小写
283
+ username = name.lower().replace(" ", "_")
284
+ username = ''.join(c for c in username if c.isalnum() or c == '_')
285
+
286
+ # 添加随机后缀避免重复
287
+ suffix = random.randint(100, 999)
288
+ return f"{username}_{suffix}"
289
+
290
+ def _search_zep_for_entity(self, entity: EntityNode) -> Dict[str, Any]:
291
+ """
292
+ 使用Zep图谱混合搜索功能获取实体相关的丰富信息
293
+
294
+ Zep没有内置混合搜索接口,需要分别搜索edges和nodes然后合并结果。
295
+ 使用并行请求同时搜索,提高效率。
296
+
297
+ Args:
298
+ entity: 实体节点对象
299
+
300
+ Returns:
301
+ 包含facts, node_summaries, context的字典
302
+ """
303
+ import concurrent.futures
304
+
305
+ if not self.graph_id:
306
+ return {"facts": [], "node_summaries": [], "context": ""}
307
+ if self._fixed_zep_key and not self.zep_client:
308
+ return {"facts": [], "node_summaries": [], "context": ""}
309
+
310
+ entity_name = entity.name
311
+
312
+ results = {
313
+ "facts": [],
314
+ "node_summaries": [],
315
+ "context": ""
316
+ }
317
+
318
+ comprehensive_query = f"关于{entity_name}的所有信息、活动、事件、关系和背景"
319
+
320
+ def search_edges():
321
+ """搜索边(事实/关系)- 带重试机制"""
322
+ max_retries = 3
323
+ last_exception = None
324
+ delay = 2.0
325
+
326
+ for attempt in range(max_retries):
327
+ try:
328
+ def _edges():
329
+ k = self._fixed_zep_key or get_zep_key()
330
+ if not k:
331
+ raise ValueError("ZEP_API_KEY 未配置")
332
+ return Zep(api_key=k).graph.search(
333
+ query=comprehensive_query,
334
+ graph_id=self.graph_id,
335
+ limit=30,
336
+ scope="edges",
337
+ reranker="rrf",
338
+ )
339
+
340
+ if self._fixed_zep_key:
341
+ r = self.zep_client.graph.search(
342
+ query=comprehensive_query,
343
+ graph_id=self.graph_id,
344
+ limit=30,
345
+ scope="edges",
346
+ reranker="rrf",
347
+ )
348
+ else:
349
+ r = call_with_limit_rotation(_edges)
350
+ return r
351
+ except Exception as e:
352
+ last_exception = e
353
+ if attempt < max_retries - 1:
354
+ logger.debug(f"Zep边搜索第 {attempt + 1} 次失败: {str(e)[:80]}, 重试中...")
355
+ time.sleep(delay)
356
+ delay *= 2
357
+ else:
358
+ logger.debug(f"Zep边搜索在 {max_retries} 次尝试后仍失败: {e}")
359
+ return None
360
+
361
+ def search_nodes():
362
+ """搜索节点(实体摘要)- 带重试机制"""
363
+ max_retries = 3
364
+ last_exception = None
365
+ delay = 2.0
366
+
367
+ for attempt in range(max_retries):
368
+ try:
369
+ def _nodes():
370
+ k = self._fixed_zep_key or get_zep_key()
371
+ if not k:
372
+ raise ValueError("ZEP_API_KEY 未配置")
373
+ return Zep(api_key=k).graph.search(
374
+ query=comprehensive_query,
375
+ graph_id=self.graph_id,
376
+ limit=20,
377
+ scope="nodes",
378
+ reranker="rrf",
379
+ )
380
+
381
+ if self._fixed_zep_key:
382
+ r = self.zep_client.graph.search(
383
+ query=comprehensive_query,
384
+ graph_id=self.graph_id,
385
+ limit=20,
386
+ scope="nodes",
387
+ reranker="rrf",
388
+ )
389
+ else:
390
+ r = call_with_limit_rotation(_nodes)
391
+ return r
392
+ except Exception as e:
393
+ last_exception = e
394
+ if attempt < max_retries - 1:
395
+ logger.debug(f"Zep节点搜索第 {attempt + 1} 次失败: {str(e)[:80]}, 重试中...")
396
+ time.sleep(delay)
397
+ delay *= 2
398
+ else:
399
+ logger.debug(f"Zep节点搜索在 {max_retries} 次尝试后仍失败: {e}")
400
+ return None
401
+
402
+ try:
403
+ # 并行执行edges和nodes搜索
404
+ with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor:
405
+ edge_future = executor.submit(search_edges)
406
+ node_future = executor.submit(search_nodes)
407
+
408
+ # 获取结果
409
+ edge_result = edge_future.result(timeout=30)
410
+ node_result = node_future.result(timeout=30)
411
+
412
+ # 处理边搜索结果
413
+ all_facts = set()
414
+ if edge_result and hasattr(edge_result, 'edges') and edge_result.edges:
415
+ for edge in edge_result.edges:
416
+ if hasattr(edge, 'fact') and edge.fact:
417
+ all_facts.add(edge.fact)
418
+ results["facts"] = list(all_facts)
419
+
420
+ # 处理节点搜索结果
421
+ all_summaries = set()
422
+ if node_result and hasattr(node_result, 'nodes') and node_result.nodes:
423
+ for node in node_result.nodes:
424
+ if hasattr(node, 'summary') and node.summary:
425
+ all_summaries.add(node.summary)
426
+ if hasattr(node, 'name') and node.name and node.name != entity_name:
427
+ all_summaries.add(f"相关实体: {node.name}")
428
+ results["node_summaries"] = list(all_summaries)
429
+
430
+ # 构建综合上下文
431
+ context_parts = []
432
+ if results["facts"]:
433
+ context_parts.append("事实信息:\n" + "\n".join(f"- {f}" for f in results["facts"][:20]))
434
+ if results["node_summaries"]:
435
+ context_parts.append("相关实体:\n" + "\n".join(f"- {s}" for s in results["node_summaries"][:10]))
436
+ results["context"] = "\n\n".join(context_parts)
437
+
438
+ logger.info(f"Zep混合检索完成: {entity_name}, 获取 {len(results['facts'])} 条事实, {len(results['node_summaries'])} 个相关节点")
439
+
440
+ except concurrent.futures.TimeoutError:
441
+ logger.warning(f"Zep检索超时 ({entity_name})")
442
+ except Exception as e:
443
+ logger.warning(f"Zep检索失败 ({entity_name}): {e}")
444
+
445
+ return results
446
+
447
+ def _build_entity_context(self, entity: EntityNode) -> str:
448
+ """
449
+ 构建实体的完整上下文信息
450
+
451
+ 包括:
452
+ 1. 实体本身的边信息(事实)
453
+ 2. 关联节点的详细信息
454
+ 3. Zep混合检索到的丰富信息
455
+ """
456
+ context_parts = []
457
+
458
+ # 1. 添加实体属性信息
459
+ if entity.attributes:
460
+ attrs = []
461
+ for key, value in entity.attributes.items():
462
+ if value and str(value).strip():
463
+ attrs.append(f"- {key}: {value}")
464
+ if attrs:
465
+ context_parts.append("### 实体属性\n" + "\n".join(attrs))
466
+
467
+ # 2. 添加相关边信息(事实/关系)
468
+ existing_facts = set()
469
+ if entity.related_edges:
470
+ relationships = []
471
+ for edge in entity.related_edges: # 不限制数量
472
+ fact = edge.get("fact", "")
473
+ edge_name = edge.get("edge_name", "")
474
+ direction = edge.get("direction", "")
475
+
476
+ if fact:
477
+ relationships.append(f"- {fact}")
478
+ existing_facts.add(fact)
479
+ elif edge_name:
480
+ if direction == "outgoing":
481
+ relationships.append(f"- {entity.name} --[{edge_name}]--> (相关实体)")
482
+ else:
483
+ relationships.append(f"- (相关实体) --[{edge_name}]--> {entity.name}")
484
+
485
+ if relationships:
486
+ context_parts.append("### 相关事实和关系\n" + "\n".join(relationships))
487
+
488
+ # 3. 添加关联节点的详细信息
489
+ if entity.related_nodes:
490
+ related_info = []
491
+ for node in entity.related_nodes: # 不限制数量
492
+ node_name = node.get("name", "")
493
+ node_labels = node.get("labels", [])
494
+ node_summary = node.get("summary", "")
495
+
496
+ # 过滤掉默认标签
497
+ custom_labels = [l for l in node_labels if l not in ["Entity", "Node"]]
498
+ label_str = f" ({', '.join(custom_labels)})" if custom_labels else ""
499
+
500
+ if node_summary:
501
+ related_info.append(f"- **{node_name}**{label_str}: {node_summary}")
502
+ else:
503
+ related_info.append(f"- **{node_name}**{label_str}")
504
+
505
+ if related_info:
506
+ context_parts.append("### 关联实体信息\n" + "\n".join(related_info))
507
+
508
+ # 4. 使用Zep混合检索获取更丰富的信息
509
+ zep_results = self._search_zep_for_entity(entity)
510
+
511
+ if zep_results.get("facts"):
512
+ # 去重:排除已存在的事实
513
+ new_facts = [f for f in zep_results["facts"] if f not in existing_facts]
514
+ if new_facts:
515
+ context_parts.append("### Zep检索到的事实信息\n" + "\n".join(f"- {f}" for f in new_facts[:15]))
516
+
517
+ if zep_results.get("node_summaries"):
518
+ context_parts.append("### Zep检索到的相关节点\n" + "\n".join(f"- {s}" for s in zep_results["node_summaries"][:10]))
519
+
520
+ return "\n\n".join(context_parts)
521
+
522
+ def _is_individual_entity(self, entity_type: str) -> bool:
523
+ """判断是否是个人类型实体"""
524
+ return entity_type.lower() in self.INDIVIDUAL_ENTITY_TYPES
525
+
526
+ def _is_group_entity(self, entity_type: str) -> bool:
527
+ """判断是否是群体/机构类型实体"""
528
+ return entity_type.lower() in self.GROUP_ENTITY_TYPES
529
+
530
+ def _generate_profile_with_llm(
531
+ self,
532
+ entity_name: str,
533
+ entity_type: str,
534
+ entity_summary: str,
535
+ entity_attributes: Dict[str, Any],
536
+ context: str
537
+ ) -> Dict[str, Any]:
538
+ """
539
+ 使用LLM生成非常详细的人设
540
+
541
+ 根据实体类型区分:
542
+ - 个人实体:生成具体的人物设定
543
+ - 群体/机构实体:生成代表性账号设定
544
+ """
545
+
546
+ is_individual = self._is_individual_entity(entity_type)
547
+
548
+ if is_individual:
549
+ prompt = self._build_individual_persona_prompt(
550
+ entity_name, entity_type, entity_summary, entity_attributes, context
551
+ )
552
+ else:
553
+ prompt = self._build_group_persona_prompt(
554
+ entity_name, entity_type, entity_summary, entity_attributes, context
555
+ )
556
+
557
+ # 尝试多次生成,直到成功或达到最大重试次数
558
+ max_attempts = 3
559
+ last_error = None
560
+
561
+ for attempt in range(max_attempts):
562
+ try:
563
+ def _create():
564
+ if self._fixed_llm_key:
565
+ return self.client.chat.completions.create(
566
+ model=self.model_name,
567
+ messages=[
568
+ {"role": "system", "content": self._get_system_prompt(is_individual)},
569
+ {"role": "user", "content": prompt},
570
+ ],
571
+ response_format={"type": "json_object"},
572
+ temperature=0.7 - (attempt * 0.1),
573
+ )
574
+ key = get_llm_key()
575
+ if not key:
576
+ raise ValueError("LLM_API_KEY 未配置")
577
+ cl = OpenAI(api_key=key, base_url=self.base_url)
578
+ return cl.chat.completions.create(
579
+ model=self.model_name,
580
+ messages=[
581
+ {"role": "system", "content": self._get_system_prompt(is_individual)},
582
+ {"role": "user", "content": prompt},
583
+ ],
584
+ response_format={"type": "json_object"},
585
+ temperature=0.7 - (attempt * 0.1),
586
+ )
587
+
588
+ response = (
589
+ self.client.chat.completions.create(
590
+ model=self.model_name,
591
+ messages=[
592
+ {"role": "system", "content": self._get_system_prompt(is_individual)},
593
+ {"role": "user", "content": prompt},
594
+ ],
595
+ response_format={"type": "json_object"},
596
+ temperature=0.7 - (attempt * 0.1),
597
+ )
598
+ if self._fixed_llm_key
599
+ else call_with_limit_rotation(_create)
600
+ )
601
+
602
+ content = response.choices[0].message.content
603
+
604
+ # 检查是否被截断(finish_reason不是'stop')
605
+ finish_reason = response.choices[0].finish_reason
606
+ if finish_reason == 'length':
607
+ logger.warning(f"LLM输出被截断 (attempt {attempt+1}), 尝试修复...")
608
+ content = self._fix_truncated_json(content)
609
+
610
+ # 尝试解析JSON
611
+ try:
612
+ result = json.loads(content)
613
+
614
+ # 验证必需字段
615
+ if "bio" not in result or not result["bio"]:
616
+ result["bio"] = entity_summary[:200] if entity_summary else f"{entity_type}: {entity_name}"
617
+ if "persona" not in result or not result["persona"]:
618
+ result["persona"] = entity_summary or f"{entity_name}是一个{entity_type}。"
619
+
620
+ return result
621
+
622
+ except json.JSONDecodeError as je:
623
+ logger.warning(f"JSON解析失败 (attempt {attempt+1}): {str(je)[:80]}")
624
+
625
+ # 尝试修复JSON
626
+ result = self._try_fix_json(content, entity_name, entity_type, entity_summary)
627
+ if result.get("_fixed"):
628
+ del result["_fixed"]
629
+ return result
630
+
631
+ last_error = je
632
+
633
+ except Exception as e:
634
+ logger.warning(f"LLM调用失败 (attempt {attempt+1}): {str(e)[:80]}")
635
+ last_error = e
636
+ import time
637
+ time.sleep(1 * (attempt + 1)) # 指数退避
638
+
639
+ logger.warning(f"LLM生成人设失败({max_attempts}次尝试): {last_error}, 使用规则生成")
640
+ return self._generate_profile_rule_based(
641
+ entity_name, entity_type, entity_summary, entity_attributes
642
+ )
643
+
644
+ def _fix_truncated_json(self, content: str) -> str:
645
+ """修复被截断的JSON(输出被max_tokens限制截断)"""
646
+ import re
647
+
648
+ # 如果JSON被截断,尝试闭合它
649
+ content = content.strip()
650
+
651
+ # 计算未闭合的括号
652
+ open_braces = content.count('{') - content.count('}')
653
+ open_brackets = content.count('[') - content.count(']')
654
+
655
+ # 检查是否有未闭合的字符串
656
+ # 简单检查:如果最后一个引号后没有逗号或闭合括号,可能是字符串被截断
657
+ if content and content[-1] not in '",}]':
658
+ # 尝试闭合字符串
659
+ content += '"'
660
+
661
+ # 闭合括号
662
+ content += ']' * open_brackets
663
+ content += '}' * open_braces
664
+
665
+ return content
666
+
667
+ def _try_fix_json(self, content: str, entity_name: str, entity_type: str, entity_summary: str = "") -> Dict[str, Any]:
668
+ """尝试修复损坏的JSON"""
669
+ import re
670
+
671
+ # 1. 首先尝试修复被截断的情况
672
+ content = self._fix_truncated_json(content)
673
+
674
+ # 2. 尝试提取JSON部分
675
+ json_match = re.search(r'\{[\s\S]*\}', content)
676
+ if json_match:
677
+ json_str = json_match.group()
678
+
679
+ # 3. 处理字符串中的换行符问题
680
+ # 找到所有字符串值并替换其中的换行符
681
+ def fix_string_newlines(match):
682
+ s = match.group(0)
683
+ # 替换字符串内的实际换行符为空格
684
+ s = s.replace('\n', ' ').replace('\r', ' ')
685
+ # 替换多余空格
686
+ s = re.sub(r'\s+', ' ', s)
687
+ return s
688
+
689
+ # 匹配JSON字符串值
690
+ json_str = re.sub(r'"[^"\\]*(?:\\.[^"\\]*)*"', fix_string_newlines, json_str)
691
+
692
+ # 4. 尝试解析
693
+ try:
694
+ result = json.loads(json_str)
695
+ result["_fixed"] = True
696
+ return result
697
+ except json.JSONDecodeError as e:
698
+ # 5. 如果还是失败,尝试更激进的修复
699
+ try:
700
+ # 移除所有控制字符
701
+ json_str = re.sub(r'[\x00-\x1f\x7f-\x9f]', ' ', json_str)
702
+ # 替换所有连续空白
703
+ json_str = re.sub(r'\s+', ' ', json_str)
704
+ result = json.loads(json_str)
705
+ result["_fixed"] = True
706
+ return result
707
+ except:
708
+ pass
709
+
710
+ # 6. 尝试从内容中提取部分信息
711
+ bio_match = re.search(r'"bio"\s*:\s*"([^"]*)"', content)
712
+ persona_match = re.search(r'"persona"\s*:\s*"([^"]*)', content) # 可能被截断
713
+
714
+ bio = bio_match.group(1) if bio_match else (entity_summary[:200] if entity_summary else f"{entity_type}: {entity_name}")
715
+ persona = persona_match.group(1) if persona_match else (entity_summary or f"{entity_name}是一个{entity_type}。")
716
+
717
+ # 如果提取到了有意义的内容,标记为已修复
718
+ if bio_match or persona_match:
719
+ logger.info(f"从损坏的JSON中提取了部分信息")
720
+ return {
721
+ "bio": bio,
722
+ "persona": persona,
723
+ "_fixed": True
724
+ }
725
+
726
+ # 7. 完全失败,返回基础结构
727
+ logger.warning(f"JSON修复失败,返回基础结构")
728
+ return {
729
+ "bio": entity_summary[:200] if entity_summary else f"{entity_type}: {entity_name}",
730
+ "persona": entity_summary or f"{entity_name}是一个{entity_type}。"
731
+ }
732
+
733
+ def _get_system_prompt(self, is_individual: bool) -> str:
734
+ """获取系统提示词"""
735
+ base_prompt = (
736
+ "你是社交媒体用户画像生成专家。生成详细、真实的人设用于「创业想法验证」舆论模拟:"
737
+ "实体可能为硅谷风格 VC、YC 式合伙人、难取悦的天使、同行创始人、目标付费客户、犹豫用户等。"
738
+ "人设须体现其对该类想法的典型评判标准( traction、差异化、团队、风险等)。"
739
+ "必须返回有效的JSON格式,所有字符串值不能包含未转义的换行符。使用中文。"
740
+ )
741
+ return base_prompt
742
+
743
+ def _build_individual_persona_prompt(
744
+ self,
745
+ entity_name: str,
746
+ entity_type: str,
747
+ entity_summary: str,
748
+ entity_attributes: Dict[str, Any],
749
+ context: str
750
+ ) -> str:
751
+ """构建个人实体的详细人设提示词"""
752
+
753
+ attrs_str = json.dumps(entity_attributes, ensure_ascii=False) if entity_attributes else "无"
754
+ context_str = context[:3000] if context else "无额外上下文"
755
+
756
+ return f"""为实体生成详细的社交媒体用户人设,最大程度还原已有现实情况。
757
+
758
+ 实体名称: {entity_name}
759
+ 实体类型: {entity_type}
760
+ 实体摘要: {entity_summary}
761
+ 实体属性: {attrs_str}
762
+
763
+ 上下文信息:
764
+ {context_str}
765
+
766
+ 请生成JSON,包含以下字段:
767
+
768
+ 1. bio: 社交媒体简介,200字
769
+ 2. persona: 详细人设描述(2000字的纯文本),需包含:
770
+ - 基本信息(年龄、职业、教育背景、所在地)
771
+ - 人物背景(重要经历、与事件的关联、社会关系)
772
+ - 性格特征(MBTI类型、核心性格、情绪表达方式)
773
+ - 社交媒体行为(发帖频率、内容偏好、互动风格、语言特点)
774
+ - 立场观点(对话题的态度;若是投资人需体现投资偏好与挑剔点;若是客户需体现付费阈值与痛点)
775
+ - 独特特征(口头禅、特殊经历、个人爱好)
776
+ - 个人记忆(人设的重要部分:该个体与材料中创业想法/事件的关系,以及已有动作与反应)
777
+ 3. age: 年龄数字(必须是整数)
778
+ 4. gender: 性别,必须是英文: "male" 或 "female"
779
+ 5. mbti: MBTI类型(如INTJ、ENFP等)
780
+ 6. country: 国家(使用中文,如"中国")
781
+ 7. profession: 职业
782
+ 8. interested_topics: 感兴趣话题数组
783
+
784
+ 重要:
785
+ - 所有字段值必须是字符串或数字,不要使用换行符
786
+ - persona必须是一段连贯的文字描述
787
+ - 使用中文(除了gender字段必须用英文male/female)
788
+ - 内容要与实体信息保持一致
789
+ - age必须是有效的整数,gender必须是"male"或"female"
790
+ """
791
+
792
+ def _build_group_persona_prompt(
793
+ self,
794
+ entity_name: str,
795
+ entity_type: str,
796
+ entity_summary: str,
797
+ entity_attributes: Dict[str, Any],
798
+ context: str
799
+ ) -> str:
800
+ """构建群体/机构实体的详细人设提示词"""
801
+
802
+ attrs_str = json.dumps(entity_attributes, ensure_ascii=False) if entity_attributes else "无"
803
+ context_str = context[:3000] if context else "无额外上下文"
804
+
805
+ return f"""为机构/群体实体生成详细的社交媒体账号设定,最大程度还原已有现实情况。
806
+
807
+ 实体名称: {entity_name}
808
+ 实体类型: {entity_type}
809
+ 实体摘要: {entity_summary}
810
+ 实体属性: {attrs_str}
811
+
812
+ 上下文信息:
813
+ {context_str}
814
+
815
+ 请生成JSON,包含以下字段:
816
+
817
+ 1. bio: 官方账号简介,200字,专业得体
818
+ 2. persona: 详细账号设定描述(2000字的纯文本),需包含:
819
+ - 机构基本信息(正式名称、机构性质、成立背景、主要职能)
820
+ - 账号定位(账号类型、目标受众、核心功能)
821
+ - 发言风格(语言特点、常用表达、禁忌话题)
822
+ - 发布内容特点(内容类型、发布频率、活跃时间段)
823
+ - 立场态度(对核心话题的官方立场、面对争议的处理方式;若是投资机构/媒体需体现对该赛道的公开态度)
824
+ - 特殊说明(代表的群体画像、运营习惯)
825
+ - 机构记忆(机构与材料中创业想法/事件的关系,以及已有动作与反应)
826
+ 3. age: 固定填30(机构账号的虚拟年龄)
827
+ 4. gender: 固定填"other"(机构账号使用other表示非个人)
828
+ 5. mbti: MBTI类型,用于描述账号风格,如ISTJ代表严谨保守
829
+ 6. country: 国家(使用中文,如"中国")
830
+ 7. profession: 机构职能描述
831
+ 8. interested_topics: 关注领域数组
832
+
833
+ 重要:
834
+ - 所有字段值必须是字符串或数字,不允许null值
835
+ - persona必须是一段连贯的文字描述,不要使用换行符
836
+ - 使用中文(除了gender字段必须用英文"other")
837
+ - age必须是整数30,gender必须是字符串"other"
838
+ - 机构账号发言要符合其身份定位"""
839
+
840
+ def _generate_profile_rule_based(
841
+ self,
842
+ entity_name: str,
843
+ entity_type: str,
844
+ entity_summary: str,
845
+ entity_attributes: Dict[str, Any]
846
+ ) -> Dict[str, Any]:
847
+ """使用规则生成基础人设"""
848
+
849
+ # 根据实体类型生成不同的人设
850
+ entity_type_lower = entity_type.lower()
851
+
852
+ if entity_type_lower in ["student", "alumni"]:
853
+ return {
854
+ "bio": f"{entity_type} with interests in academics and social issues.",
855
+ "persona": f"{entity_name} is a {entity_type.lower()} who is actively engaged in academic and social discussions. They enjoy sharing perspectives and connecting with peers.",
856
+ "age": random.randint(18, 30),
857
+ "gender": random.choice(["male", "female"]),
858
+ "mbti": random.choice(self.MBTI_TYPES),
859
+ "country": random.choice(self.COUNTRIES),
860
+ "profession": "Student",
861
+ "interested_topics": ["Education", "Social Issues", "Technology"],
862
+ }
863
+
864
+ elif entity_type_lower in ["publicfigure", "expert", "faculty"]:
865
+ return {
866
+ "bio": f"Expert and thought leader in their field.",
867
+ "persona": f"{entity_name} is a recognized {entity_type.lower()} who shares insights and opinions on important matters. They are known for their expertise and influence in public discourse.",
868
+ "age": random.randint(35, 60),
869
+ "gender": random.choice(["male", "female"]),
870
+ "mbti": random.choice(["ENTJ", "INTJ", "ENTP", "INTP"]),
871
+ "country": random.choice(self.COUNTRIES),
872
+ "profession": entity_attributes.get("occupation", "Expert"),
873
+ "interested_topics": ["Politics", "Economics", "Culture & Society"],
874
+ }
875
+
876
+ elif entity_type_lower in ["mediaoutlet", "socialmediaplatform"]:
877
+ return {
878
+ "bio": f"Official account for {entity_name}. News and updates.",
879
+ "persona": f"{entity_name} is a media entity that reports news and facilitates public discourse. The account shares timely updates and engages with the audience on current events.",
880
+ "age": 30, # 机构虚拟年龄
881
+ "gender": "other", # 机构使用other
882
+ "mbti": "ISTJ", # 机构风格:严谨保守
883
+ "country": "中国",
884
+ "profession": "Media",
885
+ "interested_topics": ["General News", "Current Events", "Public Affairs"],
886
+ }
887
+
888
+ elif entity_type_lower in ["university", "governmentagency", "ngo", "organization"]:
889
+ return {
890
+ "bio": f"Official account of {entity_name}.",
891
+ "persona": f"{entity_name} is an institutional entity that communicates official positions, announcements, and engages with stakeholders on relevant matters.",
892
+ "age": 30, # 机构虚拟年龄
893
+ "gender": "other", # 机构使用other
894
+ "mbti": "ISTJ", # 机构风格:严谨保守
895
+ "country": "中国",
896
+ "profession": entity_type,
897
+ "interested_topics": ["Public Policy", "Community", "Official Announcements"],
898
+ }
899
+
900
+ elif entity_type_lower in ["ventureinvestor", "venturecapitalfirm"]:
901
+ return {
902
+ "bio": "VC partner; focuses on traction, moat, and team.",
903
+ "persona": f"{entity_name} is a Silicon Valley-style VC who is hard to impress, asks about TAM and defensibility, and rarely commits without signal.",
904
+ "age": random.randint(32, 55),
905
+ "gender": random.choice(["male", "female"]),
906
+ "mbti": random.choice(["ENTJ", "INTJ", "ESTJ"]),
907
+ "country": "美国",
908
+ "profession": "Venture Investor",
909
+ "interested_topics": ["Startups", "B2B SaaS", "Fundraising"],
910
+ }
911
+
912
+ elif entity_type_lower in ["angelinvestor"]:
913
+ return {
914
+ "bio": "Angel investor; skeptical, pattern-matches on founders.",
915
+ "persona": f"{entity_name} is a picky angel who has seen many pitches and pushes back on vague metrics and hype.",
916
+ "age": random.randint(38, 62),
917
+ "gender": random.choice(["male", "female"]),
918
+ "mbti": random.choice(["INTJ", "ISTP", "ENTP"]),
919
+ "country": random.choice(self.COUNTRIES),
920
+ "profession": "Angel Investor",
921
+ "interested_topics": ["Early-stage", "Angel deals", "Product-market fit"],
922
+ }
923
+
924
+ elif entity_type_lower in ["acceleratorpartner"]:
925
+ return {
926
+ "bio": "Accelerator partner; YC-style fast feedback and clarity.",
927
+ "persona": f"{entity_name} is an accelerator partner who demands clear one-liners, weekly growth, and ruthless focus—friendly but intense.",
928
+ "age": random.randint(30, 48),
929
+ "gender": random.choice(["male", "female"]),
930
+ "mbti": random.choice(["ENTJ", "ENFP", "ENTP"]),
931
+ "country": "美国",
932
+ "profession": "Accelerator Partner",
933
+ "interested_topics": ["Demo day", "Batch companies", "Mentorship"],
934
+ }
935
+
936
+ elif entity_type_lower in ["startupfounder"]:
937
+ return {
938
+ "bio": "Peer founder building in the same space.",
939
+ "persona": f"{entity_name} is another founder who compares ideas to their own roadmap, shares war stories, and may collaborate or compete.",
940
+ "age": random.randint(26, 45),
941
+ "gender": random.choice(["male", "female"]),
942
+ "mbti": random.choice(self.MBTI_TYPES),
943
+ "country": random.choice(self.COUNTRIES),
944
+ "profession": "Startup Founder",
945
+ "interested_topics": ["Product", "Growth", "Fundraising"],
946
+ }
947
+
948
+ elif entity_type_lower in ["targetcustomer"]:
949
+ return {
950
+ "bio": "Potential buyer / adopter in the ICP.",
951
+ "persona": f"{entity_name} is a target customer who cares about price, workflow fit, and whether the product solves a real daily pain.",
952
+ "age": random.randint(25, 52),
953
+ "gender": random.choice(["male", "female"]),
954
+ "mbti": random.choice(["ISFJ", "ESTJ", "ENFP", "ISTJ"]),
955
+ "country": random.choice(self.COUNTRIES),
956
+ "profession": "Buyer / User",
957
+ "interested_topics": ["ROI", "Usability", "Support"],
958
+ }
959
+
960
+ elif entity_type_lower in ["skepticaluser"]:
961
+ return {
962
+ "bio": "On-the-fence or negative potential user.",
963
+ "persona": f"{entity_name} is skeptical about new tools—worried about switching cost, privacy, or whether the idea beats incumbents.",
964
+ "age": random.randint(22, 55),
965
+ "gender": random.choice(["male", "female"]),
966
+ "mbti": random.choice(["ISTP", "INTP", "ISTJ"]),
967
+ "country": random.choice(self.COUNTRIES),
968
+ "profession": "Skeptical User",
969
+ "interested_topics": ["Trust", "Alternatives", "Risk"],
970
+ }
971
+
972
+ else:
973
+ # 默认人设
974
+ return {
975
+ "bio": entity_summary[:150] if entity_summary else f"{entity_type}: {entity_name}",
976
+ "persona": entity_summary or f"{entity_name} is a {entity_type.lower()} participating in social discussions.",
977
+ "age": random.randint(25, 50),
978
+ "gender": random.choice(["male", "female"]),
979
+ "mbti": random.choice(self.MBTI_TYPES),
980
+ "country": random.choice(self.COUNTRIES),
981
+ "profession": entity_type,
982
+ "interested_topics": ["General", "Social Issues"],
983
+ }
984
+
985
+ def set_graph_id(self, graph_id: str):
986
+ """设置图谱ID用于Zep检索"""
987
+ self.graph_id = graph_id
988
+
989
+ def generate_profiles_from_entities(
990
+ self,
991
+ entities: List[EntityNode],
992
+ use_llm: bool = True,
993
+ progress_callback: Optional[callable] = None,
994
+ graph_id: Optional[str] = None,
995
+ parallel_count: int = 5,
996
+ realtime_output_path: Optional[str] = None,
997
+ output_platform: str = "reddit"
998
+ ) -> List[OasisAgentProfile]:
999
+ """
1000
+ 批量从实体生成Agent Profile(支持并行生成)
1001
+
1002
+ Args:
1003
+ entities: 实体列表
1004
+ use_llm: 是否使用LLM生成详细人设
1005
+ progress_callback: 进度回调函数 (current, total, message)
1006
+ graph_id: 图谱ID,用于Zep检索获取更丰富上下文
1007
+ parallel_count: 并行生成数量,默认5
1008
+ realtime_output_path: 实时写入的文件路径(如果提供,每生成一个就写入一次)
1009
+ output_platform: 输出平台格式 ("reddit" 或 "twitter")
1010
+
1011
+ Returns:
1012
+ Agent Profile列表
1013
+ """
1014
+ import concurrent.futures
1015
+ from threading import Lock
1016
+
1017
+ # 设置graph_id用于Zep检索
1018
+ if graph_id:
1019
+ self.graph_id = graph_id
1020
+
1021
+ total = len(entities)
1022
+ profiles = [None] * total # 预分配列表保持顺序
1023
+ completed_count = [0] # 使用列表以便在闭包中修改
1024
+ lock = Lock()
1025
+
1026
+ # 实时写入文件的辅助函数
1027
+ def save_profiles_realtime():
1028
+ """实时保存已生成的 profiles 到文件"""
1029
+ if not realtime_output_path:
1030
+ return
1031
+
1032
+ with lock:
1033
+ # 过滤出已生成的 profiles
1034
+ existing_profiles = [p for p in profiles if p is not None]
1035
+ if not existing_profiles:
1036
+ return
1037
+
1038
+ try:
1039
+ if output_platform == "reddit":
1040
+ # Reddit JSON 格式
1041
+ profiles_data = [p.to_reddit_format() for p in existing_profiles]
1042
+ with open(realtime_output_path, 'w', encoding='utf-8') as f:
1043
+ json.dump(profiles_data, f, ensure_ascii=False, indent=2)
1044
+ else:
1045
+ # Twitter CSV 格式
1046
+ import csv
1047
+ profiles_data = [p.to_twitter_format() for p in existing_profiles]
1048
+ if profiles_data:
1049
+ fieldnames = list(profiles_data[0].keys())
1050
+ with open(realtime_output_path, 'w', encoding='utf-8', newline='') as f:
1051
+ writer = csv.DictWriter(f, fieldnames=fieldnames)
1052
+ writer.writeheader()
1053
+ writer.writerows(profiles_data)
1054
+ except Exception as e:
1055
+ logger.warning(f"实时保存 profiles 失败: {e}")
1056
+
1057
+ def generate_single_profile(idx: int, entity: EntityNode) -> tuple:
1058
+ """生成单个profile的工作函数"""
1059
+ entity_type = entity.get_entity_type() or "Entity"
1060
+
1061
+ try:
1062
+ profile = self.generate_profile_from_entity(
1063
+ entity=entity,
1064
+ user_id=idx,
1065
+ use_llm=use_llm
1066
+ )
1067
+
1068
+ # 实时输出生成的人设到控制台和日志
1069
+ self._print_generated_profile(entity.name, entity_type, profile)
1070
+
1071
+ return idx, profile, None
1072
+
1073
+ except Exception as e:
1074
+ logger.error(f"生成实体 {entity.name} 的人设失败: {str(e)}")
1075
+ # 创建一个基础profile
1076
+ fallback_profile = OasisAgentProfile(
1077
+ user_id=idx,
1078
+ user_name=self._generate_username(entity.name),
1079
+ name=entity.name,
1080
+ bio=f"{entity_type}: {entity.name}",
1081
+ persona=entity.summary or f"A participant in social discussions.",
1082
+ source_entity_uuid=entity.uuid,
1083
+ source_entity_type=entity_type,
1084
+ )
1085
+ return idx, fallback_profile, str(e)
1086
+
1087
+ logger.info(f"开始并行生成 {total} 个Agent人设(并行数: {parallel_count})...")
1088
+ print(f"\n{'='*60}")
1089
+ print(f"开始生成Agent人设 - 共 {total} 个实体,并行数: {parallel_count}")
1090
+ print(f"{'='*60}\n")
1091
+
1092
+ # 使用线程池并行执行
1093
+ with concurrent.futures.ThreadPoolExecutor(max_workers=parallel_count) as executor:
1094
+ # 提交所有任务
1095
+ future_to_entity = {
1096
+ executor.submit(generate_single_profile, idx, entity): (idx, entity)
1097
+ for idx, entity in enumerate(entities)
1098
+ }
1099
+
1100
+ # 收集结果
1101
+ for future in concurrent.futures.as_completed(future_to_entity):
1102
+ idx, entity = future_to_entity[future]
1103
+ entity_type = entity.get_entity_type() or "Entity"
1104
+
1105
+ try:
1106
+ result_idx, profile, error = future.result()
1107
+ profiles[result_idx] = profile
1108
+
1109
+ with lock:
1110
+ completed_count[0] += 1
1111
+ current = completed_count[0]
1112
+
1113
+ # 实时写入文件
1114
+ save_profiles_realtime()
1115
+
1116
+ if progress_callback:
1117
+ progress_callback(
1118
+ current,
1119
+ total,
1120
+ f"已完成 {current}/{total}: {entity.name}({entity_type})"
1121
+ )
1122
+
1123
+ if error:
1124
+ logger.warning(f"[{current}/{total}] {entity.name} 使用备用人设: {error}")
1125
+ else:
1126
+ logger.info(f"[{current}/{total}] 成功生成人设: {entity.name} ({entity_type})")
1127
+
1128
+ except Exception as e:
1129
+ logger.error(f"处理实体 {entity.name} 时发生异常: {str(e)}")
1130
+ with lock:
1131
+ completed_count[0] += 1
1132
+ profiles[idx] = OasisAgentProfile(
1133
+ user_id=idx,
1134
+ user_name=self._generate_username(entity.name),
1135
+ name=entity.name,
1136
+ bio=f"{entity_type}: {entity.name}",
1137
+ persona=entity.summary or "A participant in social discussions.",
1138
+ source_entity_uuid=entity.uuid,
1139
+ source_entity_type=entity_type,
1140
+ )
1141
+ # 实时写入文件(即使是备用人设)
1142
+ save_profiles_realtime()
1143
+
1144
+ print(f"\n{'='*60}")
1145
+ print(f"人设生成完成!共生成 {len([p for p in profiles if p])} 个Agent")
1146
+ print(f"{'='*60}\n")
1147
+
1148
+ return profiles
1149
+
1150
+ def _print_generated_profile(self, entity_name: str, entity_type: str, profile: OasisAgentProfile):
1151
+ """实时输出生成的人设到控制台(完整内容,不截断)"""
1152
+ separator = "-" * 70
1153
+
1154
+ # 构建完整输出内容(不截断)
1155
+ topics_str = ', '.join(profile.interested_topics) if profile.interested_topics else '无'
1156
+
1157
+ output_lines = [
1158
+ f"\n{separator}",
1159
+ f"[已生成] {entity_name} ({entity_type})",
1160
+ f"{separator}",
1161
+ f"用户名: {profile.user_name}",
1162
+ f"",
1163
+ f"【简介】",
1164
+ f"{profile.bio}",
1165
+ f"",
1166
+ f"【详细人设】",
1167
+ f"{profile.persona}",
1168
+ f"",
1169
+ f"【基本属性】",
1170
+ f"年龄: {profile.age} | 性别: {profile.gender} | MBTI: {profile.mbti}",
1171
+ f"职业: {profile.profession} | 国家: {profile.country}",
1172
+ f"兴趣话题: {topics_str}",
1173
+ separator
1174
+ ]
1175
+
1176
+ output = "\n".join(output_lines)
1177
+
1178
+ # 只输出到控制台(避免重复,logger不再输出完整内容)
1179
+ print(output)
1180
+
1181
+ def save_profiles(
1182
+ self,
1183
+ profiles: List[OasisAgentProfile],
1184
+ file_path: str,
1185
+ platform: str = "reddit"
1186
+ ):
1187
+ """
1188
+ 保存Profile到文件(根据平台选择正确格式)
1189
+
1190
+ OASIS平台格式要求:
1191
+ - Twitter: CSV格式
1192
+ - Reddit: JSON格式
1193
+
1194
+ Args:
1195
+ profiles: Profile列表
1196
+ file_path: 文件路径
1197
+ platform: 平台类型 ("reddit" 或 "twitter")
1198
+ """
1199
+ if platform == "twitter":
1200
+ self._save_twitter_csv(profiles, file_path)
1201
+ else:
1202
+ self._save_reddit_json(profiles, file_path)
1203
+
1204
+ def _save_twitter_csv(self, profiles: List[OasisAgentProfile], file_path: str):
1205
+ """
1206
+ 保存Twitter Profile为CSV格式(符合OASIS官方要求)
1207
+
1208
+ OASIS Twitter要求的CSV字段:
1209
+ - user_id: 用户ID(根据CSV顺序从0开始)
1210
+ - name: 用户真实姓名
1211
+ - username: 系统中的用户名
1212
+ - user_char: 详细人设描述(注入到LLM系统提示中,指导Agent行为)
1213
+ - description: 简短的公开简介(显示在用户资料页面)
1214
+
1215
+ user_char vs description 区别:
1216
+ - user_char: 内部使用,LLM系统提示,决定Agent如何思考和行动
1217
+ - description: 外部显示,其他用户可见的简介
1218
+ """
1219
+ import csv
1220
+
1221
+ # 确保文件扩展名是.csv
1222
+ if not file_path.endswith('.csv'):
1223
+ file_path = file_path.replace('.json', '.csv')
1224
+
1225
+ with open(file_path, 'w', newline='', encoding='utf-8') as f:
1226
+ writer = csv.writer(f)
1227
+
1228
+ # 写入OASIS要求的表头
1229
+ headers = ['user_id', 'name', 'username', 'user_char', 'description']
1230
+ writer.writerow(headers)
1231
+
1232
+ # 写入数据行
1233
+ for idx, profile in enumerate(profiles):
1234
+ # user_char: 完整人设(bio + persona),用于LLM系统提示
1235
+ user_char = profile.bio
1236
+ if profile.persona and profile.persona != profile.bio:
1237
+ user_char = f"{profile.bio} {profile.persona}"
1238
+ # 处理换行符(CSV中用空格替代)
1239
+ user_char = user_char.replace('\n', ' ').replace('\r', ' ')
1240
+
1241
+ # description: 简短简介,用于外部显示
1242
+ description = profile.bio.replace('\n', ' ').replace('\r', ' ')
1243
+
1244
+ row = [
1245
+ idx, # user_id: 从0开始的顺序ID
1246
+ profile.name, # name: 真实姓名
1247
+ profile.user_name, # username: 用户名
1248
+ user_char, # user_char: 完整人设(内部LLM使用)
1249
+ description # description: 简短简介(外部显示)
1250
+ ]
1251
+ writer.writerow(row)
1252
+
1253
+ logger.info(f"已保存 {len(profiles)} 个Twitter Profile到 {file_path} (OASIS CSV格式)")
1254
+
1255
+ def _normalize_gender(self, gender: Optional[str]) -> str:
1256
+ """
1257
+ 标准化gender字段为OASIS要求的英文格式
1258
+
1259
+ OASIS要求: male, female, other
1260
+ """
1261
+ if not gender:
1262
+ return "other"
1263
+
1264
+ gender_lower = gender.lower().strip()
1265
+
1266
+ # 中文映射
1267
+ gender_map = {
1268
+ "男": "male",
1269
+ "女": "female",
1270
+ "机构": "other",
1271
+ "其他": "other",
1272
+ # 英文已有
1273
+ "male": "male",
1274
+ "female": "female",
1275
+ "other": "other",
1276
+ }
1277
+
1278
+ return gender_map.get(gender_lower, "other")
1279
+
1280
+ def _save_reddit_json(self, profiles: List[OasisAgentProfile], file_path: str):
1281
+ """
1282
+ 保存Reddit Profile为JSON格式
1283
+
1284
+ 使用与 to_reddit_format() 一致的格式,确保 OASIS 能正确读取。
1285
+ 必须包含 user_id 字段,这是 OASIS agent_graph.get_agent() 匹配的关键!
1286
+
1287
+ 必需字段:
1288
+ - user_id: 用户ID(整数,用于匹配 initial_posts 中的 poster_agent_id)
1289
+ - username: 用户名
1290
+ - name: 显示名称
1291
+ - bio: 简介
1292
+ - persona: 详细人设
1293
+ - age: 年龄(整数)
1294
+ - gender: "male", "female", 或 "other"
1295
+ - mbti: MBTI类型
1296
+ - country: 国家
1297
+ """
1298
+ data = []
1299
+ for idx, profile in enumerate(profiles):
1300
+ # 使用与 to_reddit_format() 一致的格式
1301
+ item = {
1302
+ "user_id": profile.user_id if profile.user_id is not None else idx, # 关键:必须包含 user_id
1303
+ "username": profile.user_name,
1304
+ "name": profile.name,
1305
+ "bio": profile.bio[:150] if profile.bio else f"{profile.name}",
1306
+ "persona": profile.persona or f"{profile.name} is a participant in social discussions.",
1307
+ "karma": profile.karma if profile.karma else 1000,
1308
+ "created_at": profile.created_at,
1309
+ # OASIS必需字段 - 确保都有默认值
1310
+ "age": profile.age if profile.age else 30,
1311
+ "gender": self._normalize_gender(profile.gender),
1312
+ "mbti": profile.mbti if profile.mbti else "ISTJ",
1313
+ "country": profile.country if profile.country else "中国",
1314
+ }
1315
+
1316
+ # 可选字段
1317
+ if profile.profession:
1318
+ item["profession"] = profile.profession
1319
+ if profile.interested_topics:
1320
+ item["interested_topics"] = profile.interested_topics
1321
+
1322
+ data.append(item)
1323
+
1324
+ with open(file_path, 'w', encoding='utf-8') as f:
1325
+ json.dump(data, f, ensure_ascii=False, indent=2)
1326
+
1327
+ logger.info(f"已保存 {len(profiles)} 个Reddit Profile到 {file_path} (JSON格式,包含user_id字段)")
1328
+
1329
+ # 保留旧方法名作为别名,保持向后兼容
1330
+ def save_profiles_to_json(
1331
+ self,
1332
+ profiles: List[OasisAgentProfile],
1333
+ file_path: str,
1334
+ platform: str = "reddit"
1335
+ ):
1336
+ """[已废弃] 请使用 save_profiles() 方法"""
1337
+ logger.warning("save_profiles_to_json已废弃,请使用save_profiles方法")
1338
+ self.save_profiles(profiles, file_path, platform)
1339
+
app/app/services/ontology_generator.py ADDED
@@ -0,0 +1,507 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ 本体生成服务
3
+ 接口1:分析文本内容,生成适合社会模拟的实体和关系类型定义
4
+ """
5
+
6
+ import json
7
+ import logging
8
+ import re
9
+ from typing import Dict, Any, List, Optional
10
+ from ..utils.llm_client import LLMClient
11
+
12
+ logger = logging.getLogger(__name__)
13
+
14
+
15
+ def _to_pascal_case(name: str) -> str:
16
+ """将任意格式的名称转换为 PascalCase(如 'works_for' -> 'WorksFor', 'person' -> 'Person')"""
17
+ # 按非字母数字字符分割
18
+ parts = re.split(r'[^a-zA-Z0-9]+', name)
19
+ # 再按 camelCase 边界分割(如 'camelCase' -> ['camel', 'Case'])
20
+ words = []
21
+ for part in parts:
22
+ words.extend(re.sub(r'([a-z])([A-Z])', r'\1_\2', part).split('_'))
23
+ # 每个词首字母大写,过滤空串
24
+ result = ''.join(word.capitalize() for word in words if word)
25
+ return result if result else 'Unknown'
26
+
27
+
28
+ # 本体生成的系统提示词
29
+ ONTOLOGY_SYSTEM_PROMPT = """你是一个专业的知识图谱本体设计专家。你的任务是分析给定的文本内容和模拟需求,设计适合**社交媒体上的创业想法验证模拟**的实体类型和关系类型(在保留原「舆论/讨论」能力的同时,优先服务想法验证)。
30
+
31
+ **重要:你必须输出有效的JSON格式数据,不要输出任何其他内容。**
32
+
33
+ ## 核心任务背景
34
+
35
+ 我们正在构建一个**社交媒体模拟系统**,典型用途是**用多角色 Agent 压力测试用户的创业想法**:投资人(可体现硅谷 VC、YC 风格、难取悦的天使等不同挑剔程度)、同行创始人、目标客户与潜在买家(含果断、犹豫、拒绝)、媒体或意见领袖等会在平台上讨论、质疑或支持该想法。
36
+ - 每个实体都是一个可以在社交媒体上发声、互动、传播信息的"账号"或"主体"
37
+ - 实体之间会相互影响、转发、评论、回应
38
+ - 我们需要模拟围绕「该产品/商业模式」的讨论、投资兴趣与购买意愿等信息路径
39
+
40
+ 因此,**实体必须是现实中真实存在的、可以在社媒上发声和互动的主体**:
41
+
42
+ **可以是**:
43
+ - 具体的个人(公众人物、当事人、意见领袖、专家学者、普通人)
44
+ - 公司、企业(包括其官方账号)
45
+ - 组织机构(大学、协会、NGO、工会等)
46
+ - 政府部门、监管机构
47
+ - 媒体机构(报纸、电视台、自媒体、网站)
48
+ - 社交媒体平台本身
49
+ - 特定群体代表(如校友会、粉丝团、维权群体等)
50
+
51
+ **不可以是**:
52
+ - 抽象概念(如"舆论"、"情绪"、"趋势")
53
+ - 主题/话题(如"学术诚信"、"教育改革")
54
+ - 观点/态度(如"支持方"、"反对方")
55
+
56
+ ## 输出格式
57
+
58
+ 请输出JSON格式,包含以下结构:
59
+
60
+ ```json
61
+ {
62
+ "entity_types": [
63
+ {
64
+ "name": "实体类型名称(英文,PascalCase)",
65
+ "description": "简短描述(英文,不超过100字符)",
66
+ "attributes": [
67
+ {
68
+ "name": "属性名(英文,snake_case)",
69
+ "type": "text",
70
+ "description": "属性描述"
71
+ }
72
+ ],
73
+ "examples": ["示例实体1", "示例实体2"]
74
+ }
75
+ ],
76
+ "edge_types": [
77
+ {
78
+ "name": "关系类型名称(英文,UPPER_SNAKE_CASE)",
79
+ "description": "简短描述(英文,不超过100字符)",
80
+ "source_targets": [
81
+ {"source": "源实体类型", "target": "目标实体类型"}
82
+ ],
83
+ "attributes": []
84
+ }
85
+ ],
86
+ "analysis_summary": "对文本内容的简要分析说明(中文)"
87
+ }
88
+ ```
89
+
90
+ ## 设计指南(极其重要!)
91
+
92
+ ### 1. 实体类型设计 - 必须严格遵守
93
+
94
+ **数量要求:必须正好10个实体类型**
95
+
96
+ **层次结构要求(必须同时包含具体类型和兜底类型)**:
97
+
98
+ 你的10个实体类型必须包含以下层次:
99
+
100
+ A. **兜底类型(必须包含,放在列表最后2个)**:
101
+ - `Person`: 任何自然人个体的兜底类型。当一个人不属于其他更具体的人物类型时,归入此类。
102
+ - `Organization`: 任何组织机构的兜底类型。当一个组织不属于其他更具体的组织类型时,归入此类。
103
+
104
+ B. **具体类型(8个,根据文本内容设计)**:
105
+ - 针对文本中出现的主要角色,设计更具体的类型
106
+ - 例如:如果文本涉及学术事件,可以有 `Student`, `Professor`, `University`
107
+ - 例如:如果文本涉及商业事件,可以有 `Company`, `CEO`, `Employee`
108
+
109
+ **为什么需要兜底类型**:
110
+ - 文本中会出现各种人物,如"中小学教师"、"路人甲"、"某位网友"
111
+ - 如果没有专门的类型匹配,他们应该被归入 `Person`
112
+ - 同理,小型组织、临时团体等应该归入 `Organization`
113
+
114
+ **具体类型的设计原则**:
115
+ - 从文本中识别出高频出现或关键的角色类型
116
+ - 每个具体类型应该有明确的边界,避免重叠
117
+ - description 必须清晰说明这个类型和兜底类型的区别
118
+
119
+ ### 2. 关系类型设计
120
+
121
+ - 数量:6-10个
122
+ - 关系应该反映社媒互动中的真实联系
123
+ - 确保关系的 source_targets 涵盖你定义的实体类型
124
+
125
+ ### 3. 属性设计
126
+
127
+ - 每个实体类型1-3个关键属性
128
+ - **注意**:属性名不能使用 `name`、`uuid`、`group_id`、`created_at`、`summary`(这些是系统保留字)
129
+ - 推荐使用:`full_name`, `title`, `role`, `position`, `location`, `description` 等
130
+
131
+ ## 实体类型参考
132
+
133
+ **个人类(具体)** — 若文本与「创业/产品」相关,请优先包含以下类型中的若干(名称用 PascalCase,可微调命名):
134
+ - VentureInvestor: 风险投资人(可偏硅谷风格大型基金)
135
+ - AngelInvestor: 天使/个人投资人(可设定为极难取悦)
136
+ - AcceleratorPartner: 加速器合伙人(可体现 YC 式 fast iteration 提问风格)
137
+ - StartupFounder: 同行创始人
138
+ - TargetCustomer: 目标付费用户或核心采用者
139
+ - SkepticalUser: 犹豫或反对的潜在用户
140
+ - IndustryExpert: 行业专家/KOL(可质疑技术或商业假设)
141
+ - Journalist: 科技/商业记者
142
+
143
+ 其他常见具体类型(若文本需要可替换其中部分):
144
+ - Student, Professor, Executive, Official, Lawyer, Doctor, Celebrity 等
145
+
146
+ **个人类(兜底)**:
147
+ - Person: 任何自然人(不属于上述具体类型时使用)
148
+
149
+ **组织类(具体)**:
150
+ - StartupCompany: 初创公司
151
+ - VentureCapitalFirm: VC 机构
152
+ - University: 高校
153
+ - Company: 公司企业
154
+ - MediaOutlet: 媒体机构
155
+ - NGO: 非政府组织
156
+ (可按文本选用 GovernmentAgency、Hospital、School 等)
157
+
158
+ **组织类(兜底)**:
159
+ - Organization: 任何组织机构(不属于上述具体类型时使用)
160
+
161
+ ## 关系类型参考
162
+
163
+ - INVESTS_IN: 投资关系
164
+ - COMPETES_WITH: 竞争
165
+ - WORKS_FOR: 工作于
166
+ - STUDIES_AT: 就读于
167
+ - AFFILIATED_WITH: 隶属于
168
+ - REPRESENTS: 代表
169
+ - REPORTS_ON: 报道
170
+ - COMMENTS_ON: 评论
171
+ - RESPONDS_TO: 回应
172
+ - SUPPORTS: 支持
173
+ - OPPOSES: 反对
174
+ - COLLABORATES_WITH: 合作
175
+ - PURCHASES_FROM / USES_PRODUCT: 可在描述中体现客户与产品关系(命名保持英文 PascalCase / UPPER_SNAKE_CASE)
176
+ """
177
+
178
+
179
+ class OntologyGenerator:
180
+ """
181
+ 本体生成器
182
+ 分析文本内容,生成实体和关系类型定义
183
+ """
184
+
185
+ def __init__(self, llm_client: Optional[LLMClient] = None):
186
+ self.llm_client = llm_client or LLMClient()
187
+
188
+ def generate(
189
+ self,
190
+ document_texts: List[str],
191
+ simulation_requirement: str,
192
+ additional_context: Optional[str] = None
193
+ ) -> Dict[str, Any]:
194
+ """
195
+ 生成本体定义
196
+
197
+ Args:
198
+ document_texts: 文档文本列表
199
+ simulation_requirement: 模拟需求描述
200
+ additional_context: 额外上下文
201
+
202
+ Returns:
203
+ 本体定义(entity_types, edge_types等)
204
+ """
205
+ # 构建用户消息
206
+ user_message = self._build_user_message(
207
+ document_texts,
208
+ simulation_requirement,
209
+ additional_context
210
+ )
211
+
212
+ messages = [
213
+ {"role": "system", "content": ONTOLOGY_SYSTEM_PROMPT},
214
+ {"role": "user", "content": user_message}
215
+ ]
216
+
217
+ # 调用LLM
218
+ result = self.llm_client.chat_json(
219
+ messages=messages,
220
+ temperature=0.3,
221
+ max_tokens=4096
222
+ )
223
+
224
+ # 验证和后处理
225
+ result = self._validate_and_process(result)
226
+
227
+ return result
228
+
229
+ # 传给 LLM 的文本最大长度(5万字)
230
+ MAX_TEXT_LENGTH_FOR_LLM = 50000
231
+
232
+ def _build_user_message(
233
+ self,
234
+ document_texts: List[str],
235
+ simulation_requirement: str,
236
+ additional_context: Optional[str]
237
+ ) -> str:
238
+ """构建用户消息"""
239
+
240
+ # 合并文本
241
+ combined_text = "\n\n---\n\n".join(document_texts)
242
+ original_length = len(combined_text)
243
+
244
+ # 如果文本超过5万字,截断(仅影响传给LLM的内容,不影响图谱构建)
245
+ if len(combined_text) > self.MAX_TEXT_LENGTH_FOR_LLM:
246
+ combined_text = combined_text[:self.MAX_TEXT_LENGTH_FOR_LLM]
247
+ combined_text += f"\n\n...(原文共{original_length}字,已截取前{self.MAX_TEXT_LENGTH_FOR_LLM}字用于本体分析)..."
248
+
249
+ message = f"""## 模拟需求
250
+
251
+ {simulation_requirement}
252
+
253
+ ## 文档内容
254
+
255
+ {combined_text}
256
+ """
257
+
258
+ if additional_context:
259
+ message += f"""
260
+ ## 额外说明
261
+
262
+ {additional_context}
263
+ """
264
+
265
+ message += """
266
+ 请根据以上内容,设计适合「创业想法验证」场景下社交媒体讨论的实���类型和关系类型。
267
+
268
+ **必须遵守的规则**:
269
+ 1. 必须正好输出10个实体类型
270
+ 2. 最后2个必须是兜底类型:Person(个人兜底)和 Organization(组织兜底)
271
+ 3. 前8个是根据文本内容设计的具体类型(若材料涉及创业/产品,应包含投资人、创始人、客户等可验证角色中的多类)
272
+ 4. 所有实体类型必须是现实中可以发声的主体,不能是抽象概念
273
+ 5. 属性名不能使用 name、uuid、group_id 等保留字,用 full_name、org_name 等替代
274
+ """
275
+
276
+ return message
277
+
278
+ def _validate_and_process(self, result: Dict[str, Any]) -> Dict[str, Any]:
279
+ """验证和后处理结果"""
280
+
281
+ # 确保必要字段存在
282
+ if "entity_types" not in result:
283
+ result["entity_types"] = []
284
+ if "edge_types" not in result:
285
+ result["edge_types"] = []
286
+ if "analysis_summary" not in result:
287
+ result["analysis_summary"] = ""
288
+
289
+ # 验证实体类型
290
+ # 记录原始名称到 PascalCase 的映射,用于后续修正 edge 的 source_targets 引用
291
+ entity_name_map = {}
292
+ for entity in result["entity_types"]:
293
+ # 强制将 entity name 转为 PascalCase(Zep API 要求)
294
+ if "name" in entity:
295
+ original_name = entity["name"]
296
+ entity["name"] = _to_pascal_case(original_name)
297
+ if entity["name"] != original_name:
298
+ logger.warning(f"Entity type name '{original_name}' auto-converted to '{entity['name']}'")
299
+ entity_name_map[original_name] = entity["name"]
300
+ if "attributes" not in entity:
301
+ entity["attributes"] = []
302
+ if "examples" not in entity:
303
+ entity["examples"] = []
304
+ # 确保description不超过100字符
305
+ if len(entity.get("description", "")) > 100:
306
+ entity["description"] = entity["description"][:97] + "..."
307
+
308
+ # 验证关系类型
309
+ for edge in result["edge_types"]:
310
+ # 强制将 edge name 转为 SCREAMING_SNAKE_CASE(Zep API 要求)
311
+ if "name" in edge:
312
+ original_name = edge["name"]
313
+ edge["name"] = original_name.upper()
314
+ if edge["name"] != original_name:
315
+ logger.warning(f"Edge type name '{original_name}' auto-converted to '{edge['name']}'")
316
+ # 修正 source_targets 中的实体名称引用,与转换后的 PascalCase 保持一致
317
+ for st in edge.get("source_targets", []):
318
+ if st.get("source") in entity_name_map:
319
+ st["source"] = entity_name_map[st["source"]]
320
+ if st.get("target") in entity_name_map:
321
+ st["target"] = entity_name_map[st["target"]]
322
+ if "source_targets" not in edge:
323
+ edge["source_targets"] = []
324
+ if "attributes" not in edge:
325
+ edge["attributes"] = []
326
+ if len(edge.get("description", "")) > 100:
327
+ edge["description"] = edge["description"][:97] + "..."
328
+
329
+ # Zep API 限制:最多 10 个自定义实体类型,最多 10 个自定义边类型
330
+ MAX_ENTITY_TYPES = 10
331
+ MAX_EDGE_TYPES = 10
332
+
333
+ # 去重:按 name 去重,保留首次出现的
334
+ seen_names = set()
335
+ deduped = []
336
+ for entity in result["entity_types"]:
337
+ name = entity.get("name", "")
338
+ if name and name not in seen_names:
339
+ seen_names.add(name)
340
+ deduped.append(entity)
341
+ elif name in seen_names:
342
+ logger.warning(f"Duplicate entity type '{name}' removed during validation")
343
+ result["entity_types"] = deduped
344
+
345
+ # 兜底类型定义
346
+ person_fallback = {
347
+ "name": "Person",
348
+ "description": "Any individual person not fitting other specific person types.",
349
+ "attributes": [
350
+ {"name": "full_name", "type": "text", "description": "Full name of the person"},
351
+ {"name": "role", "type": "text", "description": "Role or occupation"}
352
+ ],
353
+ "examples": ["ordinary citizen", "anonymous netizen"]
354
+ }
355
+
356
+ organization_fallback = {
357
+ "name": "Organization",
358
+ "description": "Any organization not fitting other specific organization types.",
359
+ "attributes": [
360
+ {"name": "org_name", "type": "text", "description": "Name of the organization"},
361
+ {"name": "org_type", "type": "text", "description": "Type of organization"}
362
+ ],
363
+ "examples": ["small business", "community group"]
364
+ }
365
+
366
+ # 检查是否已有兜底类型
367
+ entity_names = {e["name"] for e in result["entity_types"]}
368
+ has_person = "Person" in entity_names
369
+ has_organization = "Organization" in entity_names
370
+
371
+ # 需要添加的兜底类型
372
+ fallbacks_to_add = []
373
+ if not has_person:
374
+ fallbacks_to_add.append(person_fallback)
375
+ if not has_organization:
376
+ fallbacks_to_add.append(organization_fallback)
377
+
378
+ if fallbacks_to_add:
379
+ current_count = len(result["entity_types"])
380
+ needed_slots = len(fallbacks_to_add)
381
+
382
+ # 如果添加后会超过 10 个,需要移除一些现有类型
383
+ if current_count + needed_slots > MAX_ENTITY_TYPES:
384
+ # 计算需要移除多少个
385
+ to_remove = current_count + needed_slots - MAX_ENTITY_TYPES
386
+ # 从末尾移除(保留前面更重要的具体类型)
387
+ result["entity_types"] = result["entity_types"][:-to_remove]
388
+
389
+ # 添加兜底类型
390
+ result["entity_types"].extend(fallbacks_to_add)
391
+
392
+ # 最终确保不超过限制(防御性编程)
393
+ if len(result["entity_types"]) > MAX_ENTITY_TYPES:
394
+ result["entity_types"] = result["entity_types"][:MAX_ENTITY_TYPES]
395
+
396
+ if len(result["edge_types"]) > MAX_EDGE_TYPES:
397
+ result["edge_types"] = result["edge_types"][:MAX_EDGE_TYPES]
398
+
399
+ return result
400
+
401
+ def generate_python_code(self, ontology: Dict[str, Any]) -> str:
402
+ """
403
+ 将本体定义转换为Python代码(类似ontology.py)
404
+
405
+ Args:
406
+ ontology: 本体定义
407
+
408
+ Returns:
409
+ Python代码字符串
410
+ """
411
+ code_lines = [
412
+ '"""',
413
+ '自定义实体类型定义',
414
+ '由MiroFish自动生成,用于创业想法验证与社交媒体模拟',
415
+ '"""',
416
+ '',
417
+ 'from pydantic import Field',
418
+ 'from zep_cloud.external_clients.ontology import EntityModel, EntityText, EdgeModel',
419
+ '',
420
+ '',
421
+ '# ============== 实体类型定义 ==============',
422
+ '',
423
+ ]
424
+
425
+ # 生成实体类型
426
+ for entity in ontology.get("entity_types", []):
427
+ name = entity["name"]
428
+ desc = entity.get("description", f"A {name} entity.")
429
+
430
+ code_lines.append(f'class {name}(EntityModel):')
431
+ code_lines.append(f' """{desc}"""')
432
+
433
+ attrs = entity.get("attributes", [])
434
+ if attrs:
435
+ for attr in attrs:
436
+ attr_name = attr["name"]
437
+ attr_desc = attr.get("description", attr_name)
438
+ code_lines.append(f' {attr_name}: EntityText = Field(')
439
+ code_lines.append(f' description="{attr_desc}",')
440
+ code_lines.append(f' default=None')
441
+ code_lines.append(f' )')
442
+ else:
443
+ code_lines.append(' pass')
444
+
445
+ code_lines.append('')
446
+ code_lines.append('')
447
+
448
+ code_lines.append('# ============== 关系类型定义 ==============')
449
+ code_lines.append('')
450
+
451
+ # 生成关系类型
452
+ for edge in ontology.get("edge_types", []):
453
+ name = edge["name"]
454
+ # 转换为PascalCase类名
455
+ class_name = ''.join(word.capitalize() for word in name.split('_'))
456
+ desc = edge.get("description", f"A {name} relationship.")
457
+
458
+ code_lines.append(f'class {class_name}(EdgeModel):')
459
+ code_lines.append(f' """{desc}"""')
460
+
461
+ attrs = edge.get("attributes", [])
462
+ if attrs:
463
+ for attr in attrs:
464
+ attr_name = attr["name"]
465
+ attr_desc = attr.get("description", attr_name)
466
+ code_lines.append(f' {attr_name}: EntityText = Field(')
467
+ code_lines.append(f' description="{attr_desc}",')
468
+ code_lines.append(f' default=None')
469
+ code_lines.append(f' )')
470
+ else:
471
+ code_lines.append(' pass')
472
+
473
+ code_lines.append('')
474
+ code_lines.append('')
475
+
476
+ # 生成类型字典
477
+ code_lines.append('# ============== 类型配置 ==============')
478
+ code_lines.append('')
479
+ code_lines.append('ENTITY_TYPES = {')
480
+ for entity in ontology.get("entity_types", []):
481
+ name = entity["name"]
482
+ code_lines.append(f' "{name}": {name},')
483
+ code_lines.append('}')
484
+ code_lines.append('')
485
+ code_lines.append('EDGE_TYPES = {')
486
+ for edge in ontology.get("edge_types", []):
487
+ name = edge["name"]
488
+ class_name = ''.join(word.capitalize() for word in name.split('_'))
489
+ code_lines.append(f' "{name}": {class_name},')
490
+ code_lines.append('}')
491
+ code_lines.append('')
492
+
493
+ # 生成边的source_targets映射
494
+ code_lines.append('EDGE_SOURCE_TARGETS = {')
495
+ for edge in ontology.get("edge_types", []):
496
+ name = edge["name"]
497
+ source_targets = edge.get("source_targets", [])
498
+ if source_targets:
499
+ st_list = ', '.join([
500
+ f'{{"source": "{st.get("source", "Entity")}", "target": "{st.get("target", "Entity")}"}}'
501
+ for st in source_targets
502
+ ])
503
+ code_lines.append(f' "{name}": [{st_list}],')
504
+ code_lines.append('}')
505
+
506
+ return '\n'.join(code_lines)
507
+
app/app/services/pipeline_orchestrator.py ADDED
@@ -0,0 +1,368 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ End-to-end pipeline: ontology → graph → simulation create → prepare → run → report.
3
+ Executes inside a worker thread (see api/pipeline.py).
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import time
9
+ import traceback
10
+ import uuid
11
+ from typing import Any, Dict, Optional
12
+
13
+ from ..config import Config
14
+ from ..models.project import ProjectManager, ProjectStatus
15
+ from ..models.task import TaskManager, TaskStatus
16
+ from ..services.graph_builder import GraphBuilderService
17
+ from ..services.ontology_generator import OntologyGenerator
18
+ from ..services.report_agent import ReportAgent, ReportManager, ReportStatus
19
+ from ..services.simulation_manager import SimulationManager, SimulationStatus
20
+ from ..services.simulation_runner import SimulationRunner, RunnerStatus
21
+ from ..services.text_processor import TextProcessor
22
+ from ..utils.api_key_runtime import reset_pipeline_api_keys
23
+ from ..utils.logger import get_logger
24
+
25
+ logger = get_logger("mirofish.pipeline")
26
+
27
+
28
+ def _update(
29
+ tm: TaskManager,
30
+ task_id: str,
31
+ progress: int,
32
+ message: str,
33
+ *,
34
+ stage: str,
35
+ extra: Optional[Dict[str, Any]] = None,
36
+ ) -> None:
37
+ detail: Dict[str, Any] = {"stage": stage}
38
+ if extra:
39
+ detail.update(extra)
40
+ tm.update_task(
41
+ task_id,
42
+ status=TaskStatus.PROCESSING,
43
+ progress=max(0, min(100, progress)),
44
+ message=message,
45
+ progress_detail=detail,
46
+ )
47
+
48
+
49
+ def _run_graph_build_sync(
50
+ project_id: str,
51
+ tm: TaskManager,
52
+ pipeline_task_id: str,
53
+ ) -> str:
54
+ project = ProjectManager.get_project(project_id)
55
+ if not project or not project.ontology:
56
+ raise RuntimeError("Project missing ontology")
57
+
58
+ text = ProjectManager.get_extracted_text(project_id)
59
+ if not text:
60
+ raise RuntimeError("No extracted text for project")
61
+
62
+ if not Config.get_zep_api_keys():
63
+ raise RuntimeError("ZEP_API_KEY is not configured")
64
+
65
+ graph_name = project.name or "MiroFish Graph"
66
+ chunk_size = project.chunk_size or Config.DEFAULT_CHUNK_SIZE
67
+ chunk_overlap = project.chunk_overlap or Config.DEFAULT_CHUNK_OVERLAP
68
+
69
+ project.status = ProjectStatus.GRAPH_BUILDING
70
+ ProjectManager.save_project(project)
71
+
72
+ builder = GraphBuilderService()
73
+
74
+ _update(tm, pipeline_task_id, 16, "Chunking document…", stage="graph")
75
+ chunks = TextProcessor.split_text(text, chunk_size=chunk_size, overlap=chunk_overlap)
76
+ total_chunks = len(chunks)
77
+
78
+ _update(
79
+ tm,
80
+ pipeline_task_id,
81
+ 18,
82
+ f"Creating Zep graph ({total_chunks} chunks)…",
83
+ stage="graph",
84
+ )
85
+ graph_id = builder.create_graph(name=graph_name)
86
+ project.graph_id = graph_id
87
+ ProjectManager.save_project(project)
88
+
89
+ _update(tm, pipeline_task_id, 20, "Applying ontology to graph…", stage="graph", extra={"graph_id": graph_id})
90
+ builder.set_ontology(graph_id, project.ontology)
91
+
92
+ def add_progress(msg: str, ratio: float) -> None:
93
+ p = 20 + int(ratio * 22) # 20–42
94
+ _update(tm, pipeline_task_id, p, msg, stage="graph", extra={"graph_id": graph_id})
95
+
96
+ episode_uuids = builder.add_text_batches(
97
+ graph_id, chunks, batch_size=3, progress_callback=add_progress
98
+ )
99
+
100
+ _update(tm, pipeline_task_id, 43, "Waiting for Zep to process episodes…", stage="graph")
101
+
102
+ def wait_progress(msg: str, ratio: float) -> None:
103
+ p = 43 + int(ratio * 12) # 43–55
104
+ _update(tm, pipeline_task_id, p, msg, stage="graph", extra={"graph_id": graph_id})
105
+
106
+ builder._wait_for_episodes(episode_uuids, wait_progress, timeout=900)
107
+
108
+ project.status = ProjectStatus.GRAPH_COMPLETED
109
+ ProjectManager.save_project(project)
110
+
111
+ _update(tm, pipeline_task_id, 55, "Graph build complete.", stage="graph", extra={"graph_id": graph_id})
112
+ return graph_id
113
+
114
+
115
+ def run_full_pipeline(
116
+ *,
117
+ simulation_requirement: str,
118
+ document_text: str,
119
+ source_filename: str,
120
+ task_id: str,
121
+ task_manager: TaskManager,
122
+ simulation_run_id: Optional[str] = None,
123
+ user_id: Optional[str] = None,
124
+ ) -> None:
125
+ """
126
+ Full automated run. Updates task_manager task_id; on success completes with result dict.
127
+ Optional simulation_run_id + user_id sync status and report to Supabase.
128
+ """
129
+ tm = task_manager
130
+ project_id: Optional[str] = None
131
+ simulation_id: Optional[str] = None
132
+ report_id = f"report_{uuid.uuid4().hex[:12]}"
133
+
134
+ max_rounds: Optional[int] = None
135
+ if Config.PIPELINE_MAX_ROUNDS:
136
+ try:
137
+ max_rounds = int(Config.PIPELINE_MAX_ROUNDS)
138
+ except ValueError:
139
+ max_rounds = None
140
+
141
+ parallel_profiles = Config.PIPELINE_PARALLEL_PROFILES
142
+
143
+ try:
144
+ reset_pipeline_api_keys()
145
+ _update(tm, task_id, 2, "Creating project…", stage="ontology")
146
+
147
+ project = ProjectManager.create_project(name=f"Pipeline {source_filename[:40]}")
148
+ project_id = project.project_id
149
+ project.simulation_requirement = simulation_requirement
150
+ project.files.append({"filename": source_filename, "size": len(document_text.encode("utf-8"))})
151
+ ProjectManager.save_extracted_text(project_id, document_text)
152
+ project.total_text_length = len(document_text)
153
+ ProjectManager.save_project(project)
154
+
155
+ _update(
156
+ tm,
157
+ task_id,
158
+ 5,
159
+ "Generating ontology (LLM)…",
160
+ stage="ontology",
161
+ extra={"project_id": project_id},
162
+ )
163
+ generator = OntologyGenerator()
164
+ ontology = generator.generate(
165
+ document_texts=[document_text],
166
+ simulation_requirement=simulation_requirement,
167
+ additional_context=None,
168
+ )
169
+ project.ontology = {
170
+ "entity_types": ontology.get("entity_types", []),
171
+ "edge_types": ontology.get("edge_types", []),
172
+ }
173
+ project.analysis_summary = ontology.get("analysis_summary", "")
174
+ project.status = ProjectStatus.ONTOLOGY_GENERATED
175
+ ProjectManager.save_project(project)
176
+
177
+ graph_id = _run_graph_build_sync(project_id, tm, task_id)
178
+
179
+ _update(
180
+ tm,
181
+ task_id,
182
+ 56,
183
+ "Creating simulation instance…",
184
+ stage="simulation_create",
185
+ extra={"project_id": project_id, "graph_id": graph_id},
186
+ )
187
+ manager = SimulationManager()
188
+ state = manager.create_simulation(
189
+ project_id=project_id,
190
+ graph_id=graph_id,
191
+ enable_twitter=True,
192
+ enable_reddit=True,
193
+ )
194
+ simulation_id = state.simulation_id
195
+
196
+ document_text_full = ProjectManager.get_extracted_text(project_id) or document_text
197
+
198
+ def prepare_progress(stage: str, prog: int, message: str, **kwargs: Any) -> None:
199
+ weights = {
200
+ "reading": (56, 60),
201
+ "generating_profiles": (60, 68),
202
+ "generating_config": (68, 71),
203
+ "copying_scripts": (71, 72),
204
+ }
205
+ lo, hi = weights.get(stage, (56, 72))
206
+ pct = int(lo + (hi - lo) * (prog / 100.0))
207
+ _update(
208
+ tm,
209
+ task_id,
210
+ pct,
211
+ message,
212
+ stage="prepare",
213
+ extra={"project_id": project_id, "simulation_id": simulation_id, "graph_id": graph_id},
214
+ )
215
+
216
+ _update(
217
+ tm,
218
+ task_id,
219
+ 57,
220
+ "Preparing agents and simulation config (may take a long time)…",
221
+ stage="prepare",
222
+ extra={"project_id": project_id, "simulation_id": simulation_id},
223
+ )
224
+ prepared = manager.prepare_simulation(
225
+ simulation_id=simulation_id,
226
+ simulation_requirement=simulation_requirement,
227
+ document_text=document_text_full,
228
+ defined_entity_types=None,
229
+ use_llm_for_profiles=True,
230
+ progress_callback=prepare_progress,
231
+ parallel_profile_count=parallel_profiles,
232
+ )
233
+ if prepared.status == SimulationStatus.FAILED:
234
+ raise RuntimeError(prepared.error or "Simulation prepare failed (no entities or setup error)")
235
+
236
+ _update(
237
+ tm,
238
+ task_id,
239
+ 73,
240
+ "Starting OASIS simulation (parallel platforms)…",
241
+ stage="run",
242
+ extra={"project_id": project_id, "simulation_id": simulation_id},
243
+ )
244
+ SimulationRunner.start_simulation(
245
+ simulation_id=simulation_id,
246
+ platform="parallel",
247
+ max_rounds=max_rounds,
248
+ enable_graph_memory_update=False,
249
+ graph_id=None,
250
+ )
251
+
252
+ last_prog = 73
253
+ # Wait until the simulation finishes — no wall-clock deadline (can run as long as OASIS needs).
254
+ while True:
255
+ rs = SimulationRunner.get_run_state(simulation_id)
256
+ if not rs:
257
+ time.sleep(3)
258
+ continue
259
+ status = rs.runner_status
260
+ if status == RunnerStatus.COMPLETED:
261
+ break
262
+ if status in (RunnerStatus.FAILED, RunnerStatus.STOPPED):
263
+ raise RuntimeError(rs.error or f"Simulation ended with status {status.value}")
264
+ last_prog = min(88, last_prog + 1)
265
+ _update(
266
+ tm,
267
+ task_id,
268
+ last_prog,
269
+ f"Simulation running: round ~{rs.twitter_current_round or rs.reddit_current_round or rs.current_round}…",
270
+ stage="run",
271
+ extra={"project_id": project_id, "simulation_id": simulation_id},
272
+ )
273
+ time.sleep(5)
274
+
275
+ _update(
276
+ tm,
277
+ task_id,
278
+ 90,
279
+ "Generating English validation report (LLM)…",
280
+ stage="report",
281
+ extra={
282
+ "project_id": project_id,
283
+ "simulation_id": simulation_id,
284
+ "report_id": report_id,
285
+ },
286
+ )
287
+
288
+ agent = ReportAgent(
289
+ graph_id=graph_id,
290
+ simulation_id=simulation_id,
291
+ simulation_requirement=simulation_requirement,
292
+ )
293
+
294
+ def report_progress(stage: str, progress: int, message: str) -> None:
295
+ p = 90 + int(progress * 0.09) # 90–99
296
+ _update(
297
+ tm,
298
+ task_id,
299
+ min(99, p),
300
+ f"[{stage}] {message}",
301
+ stage="report",
302
+ extra={
303
+ "project_id": project_id,
304
+ "simulation_id": simulation_id,
305
+ "report_id": report_id,
306
+ },
307
+ )
308
+
309
+ report = agent.generate_report(
310
+ progress_callback=report_progress,
311
+ report_id=report_id,
312
+ )
313
+ ReportManager.save_report(report)
314
+
315
+ if report.status != ReportStatus.COMPLETED:
316
+ raise RuntimeError(report.error or "Report generation failed")
317
+
318
+ if simulation_run_id and user_id:
319
+ from ..services.supabase_jobs import mark_run_completed
320
+
321
+ try:
322
+ mark_run_completed(
323
+ simulation_run_id=simulation_run_id,
324
+ user_id=user_id,
325
+ report_id=report.report_id,
326
+ markdown_content=report.markdown_content or "",
327
+ simulation_id=simulation_id,
328
+ project_id=project_id,
329
+ )
330
+ except Exception:
331
+ logger.exception("Supabase mark_run_completed failed (report saved on disk)")
332
+
333
+ tm.update_task(
334
+ task_id,
335
+ status=TaskStatus.COMPLETED,
336
+ progress=100,
337
+ message="Pipeline complete — report ready.",
338
+ result={
339
+ "project_id": project_id,
340
+ "graph_id": graph_id,
341
+ "simulation_id": simulation_id,
342
+ "report_id": report.report_id,
343
+ },
344
+ progress_detail={
345
+ "stage": "completed",
346
+ "project_id": project_id,
347
+ "graph_id": graph_id,
348
+ "simulation_id": simulation_id,
349
+ "report_id": report.report_id,
350
+ },
351
+ )
352
+ logger.info(
353
+ "Pipeline finished: project=%s sim=%s report=%s",
354
+ project_id,
355
+ simulation_id,
356
+ report.report_id,
357
+ )
358
+
359
+ except Exception as e:
360
+ logger.error("Pipeline failed: %s", traceback.format_exc())
361
+ if simulation_run_id:
362
+ from ..services.supabase_jobs import mark_run_failed
363
+
364
+ try:
365
+ mark_run_failed(simulation_run_id=simulation_run_id, error_message=str(e))
366
+ except Exception:
367
+ logger.exception("Supabase mark_run_failed failed")
368
+ tm.fail_task(task_id, str(e))