mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
feat: add a web UI server (#1345)
* update rdagent cmd * fix log error message * use multiProcessing.Process instead of subprocess.Popen * add traces to gitignore * add user interactor in RDLoop (finance scenarios) * add interactor (feedback, hypothesis) for quant scens * fix the test_end in qlib conf * add features init config, general instruction to qlib scenarios * set base features for based exp * fix bug when combine factors * move traces folder to git_ignore_folder * fix bug in features init * fix quant interact bug * fix logger warning error * bug fixes * modify rdagent logger, now it can set file output * adjust cli functions and fix logger bug * fix server port transport problem * update server_ui in cli * add web code * fix CI problem * black fix * update web ui README * update README * update readme
This commit is contained in:
+9
-2
@@ -19,9 +19,16 @@ class LogSettings(ExtendedBaseSettings):
|
||||
|
||||
storages: dict[str, list[int | str]] = {}
|
||||
|
||||
def set_ui_server_port(self, port: int | None) -> None:
|
||||
self.ui_server_port = port
|
||||
if port is None:
|
||||
self.storages.pop("rdagent.log.ui.storage.WebStorage", None)
|
||||
return
|
||||
|
||||
self.storages["rdagent.log.ui.storage.WebStorage"] = [port, self.trace_path]
|
||||
|
||||
def model_post_init(self, _context: Any, /) -> None:
|
||||
if self.ui_server_port is not None:
|
||||
self.storages["rdagent.log.ui.storage.WebStorage"] = [self.ui_server_port, self.trace_path]
|
||||
self.set_ui_server_port(self.ui_server_port)
|
||||
|
||||
|
||||
LOG_SETTINGS = LogSettings()
|
||||
|
||||
+34
-17
@@ -7,18 +7,12 @@ from pathlib import Path
|
||||
from typing import Generator
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .conf import LOG_SETTINGS
|
||||
|
||||
if LOG_SETTINGS.format_console is not None:
|
||||
logger.remove()
|
||||
logger.add(sys.stdout, format=LOG_SETTINGS.format_console)
|
||||
|
||||
from psutil import Process
|
||||
|
||||
from rdagent.core.utils import SingletonBaseClass, import_class
|
||||
|
||||
from .base import Storage
|
||||
from .conf import LOG_SETTINGS
|
||||
from .storage import FileStorage
|
||||
from .utils import get_caller_info
|
||||
|
||||
@@ -48,6 +42,18 @@ class RDAgentLog(SingletonBaseClass):
|
||||
|
||||
# Thread-/coroutine-local tag; In Linux forked subprocess, it will be copied to the subprocess.
|
||||
_tag_ctx: ContextVar[str] = ContextVar("_tag_ctx", default="")
|
||||
_raw_log_key = "_rdagent_raw"
|
||||
|
||||
@classmethod
|
||||
def _configure_console_sinks(cls) -> None:
|
||||
raw_filter = lambda record: bool(record["extra"].get(cls._raw_log_key, False))
|
||||
normal_filter = lambda record: not raw_filter(record)
|
||||
|
||||
if LOG_SETTINGS.format_console is not None:
|
||||
logger.add(sys.stdout, format=LOG_SETTINGS.format_console, filter=normal_filter)
|
||||
else:
|
||||
logger.add(sys.stdout, filter=normal_filter)
|
||||
logger.add(sys.stdout, format="{message}", filter=raw_filter)
|
||||
|
||||
@property
|
||||
def _tag(self) -> str: # Get current tag
|
||||
@@ -58,13 +64,29 @@ class RDAgentLog(SingletonBaseClass):
|
||||
self._tag_ctx.set(value)
|
||||
|
||||
def __init__(self) -> None:
|
||||
logger.remove()
|
||||
self._configure_console_sinks()
|
||||
|
||||
self.storage = FileStorage(LOG_SETTINGS.trace_path)
|
||||
self.other_storages: list[Storage] = []
|
||||
self.refresh_storages_from_settings()
|
||||
|
||||
self.main_pid = os.getpid()
|
||||
|
||||
def refresh_storages_from_settings(self) -> None:
|
||||
self.other_storages = []
|
||||
for storage, args in LOG_SETTINGS.storages.items():
|
||||
storage_cls = import_class(storage)
|
||||
self.other_storages.append(storage_cls(*args))
|
||||
|
||||
self.main_pid = os.getpid()
|
||||
def rebind_console_to_current_streams(self) -> None:
|
||||
"""Rebind loguru sinks to the current stdio objects.
|
||||
|
||||
This is needed in forked/spawned subprocesses after stdout/stderr have been
|
||||
redirected, because loguru keeps references to the original stream objects.
|
||||
"""
|
||||
logger.remove()
|
||||
self._configure_console_sinks()
|
||||
|
||||
@contextmanager
|
||||
def tag(self, tag: str) -> Generator[None, None, None]:
|
||||
@@ -82,6 +104,8 @@ class RDAgentLog(SingletonBaseClass):
|
||||
self._tag_ctx.reset(token)
|
||||
|
||||
def set_storages_path(self, path: str | Path) -> None:
|
||||
if isinstance(path, str):
|
||||
path = Path(path)
|
||||
for storage in [self.storage] + self.other_storages:
|
||||
if hasattr(storage, "path"):
|
||||
storage.path = path
|
||||
@@ -115,17 +139,10 @@ class RDAgentLog(SingletonBaseClass):
|
||||
caller_info = get_caller_info(level=3)
|
||||
tag = f"{self._tag}.{tag}.{self.get_pids()}".strip(".")
|
||||
|
||||
if raw:
|
||||
logger.remove()
|
||||
logger.add(sys.stderr, format=lambda r: "{message}")
|
||||
|
||||
log_func = getattr(logger.patch(lambda r: r.update(caller_info)), level)
|
||||
patched_logger = logger.patch(lambda r: r.update(caller_info)).bind(**{self._raw_log_key: raw}).opt(raw=raw)
|
||||
log_func = getattr(patched_logger, level)
|
||||
log_func(msg)
|
||||
|
||||
if raw:
|
||||
logger.remove()
|
||||
logger.add(sys.stderr)
|
||||
|
||||
def info(self, msg: str, *, tag: str = "", raw: bool = False) -> None:
|
||||
self._log("info", msg, tag=tag, raw=raw)
|
||||
|
||||
|
||||
+378
-85
@@ -1,69 +1,287 @@
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import signal
|
||||
import subprocess
|
||||
import traceback
|
||||
from collections import defaultdict
|
||||
from contextlib import redirect_stderr, redirect_stdout
|
||||
from datetime import datetime, timezone
|
||||
from multiprocessing import Process, Queue
|
||||
from pathlib import Path
|
||||
from queue import Empty
|
||||
|
||||
import randomname
|
||||
import typer
|
||||
from flask import Flask, jsonify, request, send_from_directory
|
||||
from flask import Flask, jsonify, request, send_file, send_from_directory
|
||||
from flask_cors import CORS
|
||||
from werkzeug.utils import secure_filename
|
||||
|
||||
from rdagent.log.storage import FileStorage
|
||||
from rdagent.log.ui.conf import UI_SETTING
|
||||
from rdagent.log.ui.storage import WebStorage
|
||||
from rdagent.log.utils import is_valid_session
|
||||
|
||||
app = Flask(__name__, static_folder=UI_SETTING.static_path)
|
||||
app = Flask(__name__, static_folder=str(Path(UI_SETTING.static_path).resolve()))
|
||||
CORS(app)
|
||||
app.config["UI_SERVER_PORT"] = 19899
|
||||
|
||||
rdagent_processes = defaultdict()
|
||||
server_port = 19899
|
||||
_YELLOW = "\033[33m"
|
||||
_RESET = "\033[0m"
|
||||
|
||||
|
||||
class _YellowWarningFormatter(logging.Formatter):
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
if record.levelno == logging.WARNING:
|
||||
record.levelname = f"{_YELLOW}{record.levelname}{_RESET}"
|
||||
return super().format(record)
|
||||
|
||||
|
||||
def _configure_app_logger() -> None:
|
||||
formatter = _YellowWarningFormatter(
|
||||
fmt="[%(asctime)s] %(levelname)s in %(module)s: %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
for handler in app.logger.handlers:
|
||||
handler.setFormatter(formatter)
|
||||
|
||||
|
||||
_configure_app_logger()
|
||||
|
||||
|
||||
_TARGETS_WITHOUT_USER_INTERACTION = {"general_model", "fin_factor_report"}
|
||||
|
||||
|
||||
class RDAgentTask:
|
||||
def __init__(
|
||||
self,
|
||||
target_name: str,
|
||||
kwargs: dict,
|
||||
stdout_path: str,
|
||||
log_trace_path: str,
|
||||
scenario: str,
|
||||
trace_name: str,
|
||||
ui_server_port: int | None = None,
|
||||
create_process: bool = True,
|
||||
) -> None:
|
||||
self.target_name = target_name
|
||||
self.kwargs = kwargs
|
||||
self.stdout_path = stdout_path
|
||||
self.log_trace_path = log_trace_path
|
||||
self.scenario = scenario
|
||||
self.trace_name = trace_name
|
||||
self.ui_server_port = ui_server_port
|
||||
self.process: Process | None = None
|
||||
|
||||
# Two IPC queues for user interaction.
|
||||
# - `user_request_q`: rdagent subprocess -> server (dicts to render on frontend)
|
||||
# - `user_response_q`: server -> rdagent subprocess (user input dicts)
|
||||
# NOTE: Use multiprocessing.Queue because rdagent is started as a separate process.
|
||||
self.user_request_q: Queue = Queue(maxsize=1024)
|
||||
self.user_response_q: Queue = Queue(maxsize=1024)
|
||||
|
||||
if create_process:
|
||||
self.process = Process(
|
||||
target=self._run,
|
||||
name=f"rdagent:{self.scenario}:{self.trace_name}",
|
||||
)
|
||||
self.messages: list[dict] = []
|
||||
self.pointers: defaultdict[str, int] = defaultdict(int)
|
||||
|
||||
def start(self) -> None:
|
||||
if self.process is not None:
|
||||
self.process.start()
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
return self.process is not None and self.process.is_alive()
|
||||
|
||||
def get_end_code(self) -> int:
|
||||
if self.process is None or self.process.exitcode is None:
|
||||
return 0
|
||||
return self.process.exitcode
|
||||
|
||||
def stop(self) -> None:
|
||||
if self.process is not None and self.process.is_alive():
|
||||
self.process.terminate()
|
||||
self.process.join()
|
||||
|
||||
# Best-effort cleanup for IPC queues.
|
||||
for q in (self.user_request_q, self.user_response_q):
|
||||
try:
|
||||
q.cancel_join_thread()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
q.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _run(self) -> None:
|
||||
from rdagent.log.conf import LOG_SETTINGS
|
||||
|
||||
LOG_SETTINGS.set_ui_server_port(self.ui_server_port)
|
||||
|
||||
from rdagent.log import rdagent_logger
|
||||
|
||||
rdagent_logger.refresh_storages_from_settings()
|
||||
rdagent_logger.set_storages_path(self.log_trace_path)
|
||||
Path(self.stdout_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(self.stdout_path, "w") as log_file:
|
||||
with redirect_stdout(log_file), redirect_stderr(log_file):
|
||||
rdagent_logger.rebind_console_to_current_streams()
|
||||
try:
|
||||
# Only interactive targets should receive IPC queues.
|
||||
if self.target_name not in _TARGETS_WITHOUT_USER_INTERACTION:
|
||||
self.kwargs.setdefault(
|
||||
"user_interaction_queues",
|
||||
(self.user_request_q, self.user_response_q),
|
||||
)
|
||||
|
||||
if self.target_name == "data_science":
|
||||
from rdagent.app.data_science.loop import main as data_science
|
||||
|
||||
data_science(**self.kwargs)
|
||||
elif self.target_name == "general_model":
|
||||
from rdagent.app.general_model.general_model import (
|
||||
extract_models_and_implement as general_model,
|
||||
)
|
||||
|
||||
general_model(**self.kwargs)
|
||||
elif self.target_name == "fin_factor":
|
||||
from rdagent.app.qlib_rd_loop.factor import main as fin_factor
|
||||
|
||||
fin_factor(**self.kwargs)
|
||||
elif self.target_name == "fin_factor_report":
|
||||
from rdagent.app.qlib_rd_loop.factor_from_report import (
|
||||
main as fin_factor_report,
|
||||
)
|
||||
|
||||
fin_factor_report(**self.kwargs)
|
||||
elif self.target_name == "fin_model":
|
||||
from rdagent.app.qlib_rd_loop.model import main as fin_model
|
||||
|
||||
fin_model(**self.kwargs)
|
||||
elif self.target_name == "fin_quant":
|
||||
from rdagent.app.qlib_rd_loop.quant import main as fin_quant
|
||||
|
||||
fin_quant(**self.kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown target: {self.target_name}")
|
||||
except Exception:
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
rdagent_processes: dict[str, RDAgentTask] = {}
|
||||
log_folder_path = Path(UI_SETTING.trace_folder).absolute()
|
||||
|
||||
|
||||
def _drain_user_requests_into_messages(task: RDAgentTask) -> None:
|
||||
"""Move a single pending user-interaction request into `task.messages`.
|
||||
|
||||
Assumption: each rdagent process only has one active request at a time.
|
||||
"""
|
||||
|
||||
try:
|
||||
req = task.user_request_q.get_nowait()
|
||||
except Empty:
|
||||
return
|
||||
except Exception:
|
||||
return
|
||||
|
||||
# Standardize the message shape for the frontend.
|
||||
# The agent can send either a full message dict, or a raw content dict.
|
||||
if isinstance(req, dict) and {"tag", "timestamp", "content"}.issubset(req.keys()):
|
||||
msg = req
|
||||
else:
|
||||
msg = {
|
||||
"tag": "user_interaction.request",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"content": req,
|
||||
}
|
||||
task.messages.append(msg)
|
||||
|
||||
|
||||
@app.route("/favicon.ico")
|
||||
def favicon():
|
||||
return send_from_directory(app.static_folder, "favicon.ico", mimetype="image/vnd.microsoft.icon")
|
||||
|
||||
|
||||
msgs_for_frontend = defaultdict(list)
|
||||
pointers = defaultdict(lambda: defaultdict(int)) # pointers[trace_id][user_ip]
|
||||
def _normalize_static_request_path(fn: str) -> str:
|
||||
static_prefix = UI_SETTING.static_path.strip("./")
|
||||
if static_prefix and fn.startswith(f"{static_prefix}/"):
|
||||
return fn[len(static_prefix) + 1 :]
|
||||
return fn
|
||||
|
||||
|
||||
def _get_or_create_task(trace_id: str) -> RDAgentTask:
|
||||
task = rdagent_processes.get(trace_id)
|
||||
if task is None:
|
||||
task = RDAgentTask(
|
||||
target_name="",
|
||||
kwargs={},
|
||||
stdout_path="",
|
||||
log_trace_path=trace_id,
|
||||
scenario="",
|
||||
trace_name="",
|
||||
ui_server_port=None,
|
||||
create_process=False,
|
||||
)
|
||||
rdagent_processes[trace_id] = task
|
||||
return task
|
||||
|
||||
|
||||
def _resolve_stdout_path(trace_id: str) -> Path | None:
|
||||
normalized_trace_id = str(trace_id or "").strip()
|
||||
if not normalized_trace_id:
|
||||
return None
|
||||
|
||||
task = rdagent_processes.get(str(log_folder_path / normalized_trace_id))
|
||||
if task is None or not task.stdout_path:
|
||||
return None
|
||||
|
||||
stdout_path = Path(task.stdout_path).resolve()
|
||||
|
||||
try:
|
||||
if os.path.commonpath([str(stdout_path), str(log_folder_path)]) != str(log_folder_path):
|
||||
return None
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
return stdout_path
|
||||
|
||||
|
||||
def read_trace(log_path: Path, id: str = "") -> None:
|
||||
fs = FileStorage(log_path)
|
||||
ws = WebStorage(port=1, path=log_path)
|
||||
msgs_for_frontend[id] = []
|
||||
task = _get_or_create_task(id)
|
||||
task.messages = []
|
||||
last_timestamp = None
|
||||
for msg in fs.iter_msg():
|
||||
data = ws._obj_to_json(obj=msg.content, tag=msg.tag, id=id, timestamp=msg.timestamp.isoformat())
|
||||
if data:
|
||||
if isinstance(data, list):
|
||||
for d in data:
|
||||
msgs_for_frontend[id].append(d["msg"])
|
||||
task.messages.append(d["msg"])
|
||||
last_timestamp = msg.timestamp
|
||||
else:
|
||||
msgs_for_frontend[id].append(data["msg"])
|
||||
task.messages.append(data["msg"])
|
||||
last_timestamp = msg.timestamp
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
if last_timestamp and (now - last_timestamp).total_seconds() > 1800:
|
||||
msgs_for_frontend[id].append({"tag": "END", "timestamp": now.isoformat(), "content": {}})
|
||||
task.messages.append(
|
||||
{
|
||||
"tag": "END",
|
||||
"timestamp": now.isoformat(),
|
||||
"content": {"error_msg": "Trace session has ended.", "end_code": 0},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# load all traces from the log folder
|
||||
for p in log_folder_path.glob("*/*/"):
|
||||
if is_valid_session(p):
|
||||
read_trace(p, id=str(p))
|
||||
# for p in log_folder_path.glob("*/*/"):
|
||||
# read_trace(p, id=str(p))
|
||||
|
||||
|
||||
@app.route("/trace", methods=["POST"])
|
||||
def update_trace():
|
||||
global pointers, msgs_for_frontend
|
||||
data = request.get_json()
|
||||
trace_id = data.get("id")
|
||||
return_all = data.get("all")
|
||||
@@ -75,28 +293,64 @@ def update_trace():
|
||||
return jsonify({"error": "Trace ID is required"}), 400
|
||||
trace_id = str(log_folder_path / trace_id)
|
||||
|
||||
task = _get_or_create_task(trace_id)
|
||||
|
||||
# Make sure any pending user-interaction requests are visible to the frontend.
|
||||
_drain_user_requests_into_messages(task)
|
||||
|
||||
if task.process is not None and not task.is_alive():
|
||||
if not task.messages or task.messages[-1].get("tag") != "END":
|
||||
task.messages.append(
|
||||
{
|
||||
"tag": "END",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"content": {
|
||||
"error_msg": "RD-Agent process has completed.",
|
||||
"end_code": task.get_end_code(),
|
||||
},
|
||||
}
|
||||
)
|
||||
app.logger.warning(f"Process for {trace_id} has ended.")
|
||||
|
||||
user_ip = request.remote_addr
|
||||
|
||||
if reset:
|
||||
pointers[trace_id][user_ip] = 0
|
||||
task.pointers[user_ip] = 0
|
||||
|
||||
start_pointer = pointers[trace_id][user_ip]
|
||||
start_pointer = task.pointers[user_ip]
|
||||
end_pointer = start_pointer + msg_num
|
||||
if end_pointer > len(msgs_for_frontend[trace_id]) or return_all:
|
||||
end_pointer = len(msgs_for_frontend[trace_id])
|
||||
if end_pointer > len(task.messages) or return_all:
|
||||
end_pointer = len(task.messages)
|
||||
|
||||
returned_msgs = msgs_for_frontend[trace_id][start_pointer:end_pointer]
|
||||
|
||||
pointers[trace_id][user_ip] = end_pointer
|
||||
returned_msgs = task.messages[start_pointer:end_pointer]
|
||||
task.pointers[user_ip] = end_pointer
|
||||
if returned_msgs:
|
||||
app.logger.info([msg["tag"] for msg in returned_msgs])
|
||||
return jsonify(returned_msgs), 200
|
||||
|
||||
|
||||
@app.route("/stdout", methods=["GET"])
|
||||
def download_stdout_file():
|
||||
trace_id = request.args.get("id", "")
|
||||
stdout_path = _resolve_stdout_path(trace_id)
|
||||
|
||||
if stdout_path is None:
|
||||
return jsonify({"error": "Trace ID is required or invalid"}), 400
|
||||
if not stdout_path.exists() or not stdout_path.is_file():
|
||||
return jsonify({"error": "Stdout file not found"}), 404
|
||||
|
||||
return send_file(
|
||||
stdout_path,
|
||||
as_attachment=True,
|
||||
download_name=stdout_path.name,
|
||||
mimetype="text/plain",
|
||||
)
|
||||
|
||||
|
||||
@app.route("/upload", methods=["POST"])
|
||||
def upload_file():
|
||||
# 获取请求体中的字段
|
||||
global rdagent_processes, server_port
|
||||
global rdagent_processes
|
||||
scenario = request.form.get("scenario")
|
||||
files = request.files.getlist("files")
|
||||
competition = request.form.get("competition")
|
||||
@@ -109,21 +363,19 @@ def upload_file():
|
||||
trace_name = f"{competition}-{randomname.get_name()}"
|
||||
else:
|
||||
trace_name = randomname.get_name()
|
||||
trace_files_path = log_folder_path / scenario / "uploads" / trace_name
|
||||
trace_files_path = log_folder_path / "uploads" / scenario / trace_name
|
||||
|
||||
log_trace_path = (log_folder_path / scenario / trace_name).absolute()
|
||||
stdout_path = log_folder_path / scenario / f"{trace_name}.stdout"
|
||||
stdout_path = log_folder_path / scenario / f"{trace_name}.log"
|
||||
if not stdout_path.exists():
|
||||
stdout_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# save files
|
||||
for file in files:
|
||||
if file:
|
||||
p = (log_folder_path / scenario / "uploads" / trace_name).resolve()
|
||||
p = (log_folder_path / "uploads" / scenario / trace_name).resolve()
|
||||
sanitized_filename = secure_filename(file.filename) # Sanitize filename
|
||||
target_path = (p / sanitized_filename).resolve() # Normalize target path
|
||||
if not sanitized_filename.lower().endswith(".pdf"):
|
||||
return jsonify({"error": "Invalid file type"}), 400
|
||||
# Ensure target_path is within the allowed base directory
|
||||
if os.path.commonpath([str(target_path), str(p)]) == str(p) and target_path.is_file() == False:
|
||||
if not p.exists():
|
||||
@@ -132,42 +384,62 @@ def upload_file():
|
||||
else:
|
||||
return jsonify({"error": "Invalid file path"}), 400
|
||||
|
||||
target_name = None
|
||||
kwargs = {}
|
||||
loop_n_val = int(loop_n) if loop_n else None
|
||||
all_duration_val = f"{all_duration}h" if all_duration else None
|
||||
|
||||
if scenario == "Finance Data Building":
|
||||
cmds = ["rdagent", "fin_factor"]
|
||||
if scenario == "Finance Data Building (Reports)":
|
||||
cmds = ["rdagent", "fin_factor_report", "--report_folder", str(trace_files_path)]
|
||||
target_name = "fin_factor"
|
||||
kwargs = {
|
||||
"loop_n": loop_n_val,
|
||||
"all_duration": all_duration_val,
|
||||
"base_features_path": str(trace_files_path),
|
||||
}
|
||||
if scenario == "Finance Model Implementation":
|
||||
cmds = ["rdagent", "fin_model"]
|
||||
target_name = "fin_model"
|
||||
kwargs = {
|
||||
"loop_n": loop_n_val,
|
||||
"all_duration": all_duration_val,
|
||||
"base_features_path": str(trace_files_path),
|
||||
}
|
||||
if scenario == "Finance Whole Pipeline":
|
||||
target_name = "fin_quant"
|
||||
kwargs = {
|
||||
"loop_n": loop_n_val,
|
||||
"all_duration": all_duration_val,
|
||||
"base_features_path": str(trace_files_path),
|
||||
}
|
||||
if scenario == "Finance Data Building (Reports)":
|
||||
target_name = "fin_factor_report"
|
||||
kwargs = {"report_folder": str(trace_files_path), "all_duration": all_duration_val}
|
||||
if scenario == "General Model Implementation":
|
||||
if len(files) == 0: # files is one link
|
||||
rfp = request.form.get("files")[0]
|
||||
else: # one file is uploaded
|
||||
rfp = str(trace_files_path / files[0].filename)
|
||||
cmds = ["rdagent", "general_model", "--report_file_path", rfp]
|
||||
if scenario == "Finance Whole Pipeline":
|
||||
cmds = ["rdagent", "fin_quant"]
|
||||
target_name = "general_model"
|
||||
kwargs = {"report_file_path": rfp}
|
||||
if scenario == "Data Science":
|
||||
cmds = ["rdagent", "data_science", "--competition", competition]
|
||||
target_name = "data_science"
|
||||
kwargs = {"competition": competition, "loop_n": loop_n_val, "timeout": all_duration_val}
|
||||
|
||||
# time control parameters
|
||||
if scenario != "Finance Data Building (Reports)":
|
||||
if loop_n:
|
||||
cmds += ["--loop_n", loop_n]
|
||||
if all_duration:
|
||||
cmds += ["--timeout", f"{all_duration}h"]
|
||||
if target_name is None:
|
||||
return jsonify({"error": "Unknown scenario"}), 400
|
||||
|
||||
app.logger.info(f"Started process for {log_trace_path} with parameters: {cmds}")
|
||||
with stdout_path.open("w") as log_file:
|
||||
rdagent_processes[str(log_trace_path)] = subprocess.Popen(
|
||||
cmds,
|
||||
stdout=log_file,
|
||||
stderr=log_file,
|
||||
env={
|
||||
**os.environ,
|
||||
"LOG_TRACE_PATH": str(log_trace_path),
|
||||
"LOG_UI_SERVER_PORT": str(server_port),
|
||||
},
|
||||
)
|
||||
app.logger.info(f"Started process for {log_trace_path} with target: {target_name}, kwargs: {kwargs}")
|
||||
task = RDAgentTask(
|
||||
target_name=target_name,
|
||||
kwargs=kwargs,
|
||||
stdout_path=str(stdout_path),
|
||||
log_trace_path=str(log_trace_path),
|
||||
scenario=scenario,
|
||||
trace_name=trace_name,
|
||||
ui_server_port=app.config["UI_SERVER_PORT"],
|
||||
)
|
||||
task.start()
|
||||
app.logger.warning(f"Task {log_trace_path} started.")
|
||||
rdagent_processes[str(log_trace_path)] = task
|
||||
return (
|
||||
jsonify(
|
||||
{
|
||||
@@ -182,7 +454,6 @@ def upload_file():
|
||||
def receive_msgs():
|
||||
try:
|
||||
data = request.get_json()
|
||||
# app.logger.info(data["msg"]["tag"])
|
||||
if not data:
|
||||
return jsonify({"error": "No JSON data received"}), 400
|
||||
except Exception as e:
|
||||
@@ -190,16 +461,41 @@ def receive_msgs():
|
||||
|
||||
if isinstance(data, list):
|
||||
for d in data:
|
||||
msgs_for_frontend[d["id"]].append(d["msg"])
|
||||
task = _get_or_create_task(d["id"])
|
||||
task.messages.append(d["msg"])
|
||||
else:
|
||||
msgs_for_frontend[data["id"]].append(data["msg"])
|
||||
task = _get_or_create_task(data["id"])
|
||||
task.messages.append(data["msg"])
|
||||
|
||||
return jsonify({"status": "success"}), 200
|
||||
|
||||
|
||||
@app.route("/user_interaction/submit", methods=["POST"])
|
||||
def submit_user_interaction_response():
|
||||
"""Frontend submits a user response; server forwards it to the rdagent subprocess via IPC queue."""
|
||||
data = request.get_json(silent=True) or {}
|
||||
trace_id = data.get("id")
|
||||
payload = data.get("payload")
|
||||
|
||||
if not trace_id:
|
||||
return jsonify({"error": "Trace ID is required"}), 400
|
||||
if payload is None:
|
||||
return jsonify({"error": "Missing 'payload'"}), 400
|
||||
|
||||
trace_id = str(log_folder_path / trace_id)
|
||||
task = _get_or_create_task(trace_id)
|
||||
|
||||
try:
|
||||
task.user_response_q.put(payload, block=False)
|
||||
except Exception as e:
|
||||
return jsonify({"error": f"Failed to enqueue user response: {e}"}), 500
|
||||
|
||||
return jsonify({"status": "success"}), 200
|
||||
|
||||
|
||||
@app.route("/control", methods=["POST"])
|
||||
def control_process():
|
||||
global rdagent_processes, msgs_for_frontend
|
||||
global rdagent_processes
|
||||
data = request.get_json()
|
||||
app.logger.info(data)
|
||||
if not data or "id" not in data or "action" not in data:
|
||||
@@ -208,32 +504,31 @@ def control_process():
|
||||
id = str(log_folder_path / data["id"])
|
||||
action = data["action"]
|
||||
|
||||
if action != "stop":
|
||||
return jsonify({"error": "Only 'stop' action is supported"}), 400
|
||||
|
||||
if id not in rdagent_processes or rdagent_processes[id] is None:
|
||||
return jsonify({"error": "No running process for given id"}), 400
|
||||
|
||||
process = rdagent_processes[id]
|
||||
task = rdagent_processes[id]
|
||||
|
||||
if process.poll() is not None:
|
||||
msgs_for_frontend[id].append({"tag": "END", "timestamp": datetime.now(timezone.utc).isoformat(), "content": {}})
|
||||
return jsonify({"error": "Process has already terminated"}), 400
|
||||
if task.process is None:
|
||||
return jsonify({"error": "No running process for given id"}), 400
|
||||
|
||||
try:
|
||||
if action == "pause":
|
||||
os.kill(process.pid, signal.SIGSTOP)
|
||||
return jsonify({"status": "paused"}), 200
|
||||
elif action == "resume":
|
||||
os.kill(process.pid, signal.SIGCONT)
|
||||
return jsonify({"status": "resumed"}), 200
|
||||
elif action == "stop":
|
||||
process.terminate()
|
||||
process.wait()
|
||||
del rdagent_processes[id]
|
||||
msgs_for_frontend[id].append(
|
||||
{"tag": "END", "timestamp": datetime.now(timezone.utc).isoformat(), "content": {}}
|
||||
if task.is_alive():
|
||||
task.stop()
|
||||
|
||||
if not task.messages or task.messages[-1].get("tag") != "END":
|
||||
task.messages.append(
|
||||
{
|
||||
"tag": "END",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"content": {"error_msg": "RD-Agent process was stopped by user.", "end_code": -1},
|
||||
}
|
||||
)
|
||||
return jsonify({"status": "stopped"}), 200
|
||||
else:
|
||||
return jsonify({"error": "Unknown action"}), 400
|
||||
app.logger.warning(f"Process for {id} has been stopped.")
|
||||
return jsonify({"status": "stopped"}), 200
|
||||
except Exception as e:
|
||||
return jsonify({"error": f"Failed to {action} process, {e}"}), 500
|
||||
|
||||
@@ -241,9 +536,8 @@ def control_process():
|
||||
@app.route("/test", methods=["GET"])
|
||||
def test():
|
||||
# return 'Hello, World!'
|
||||
global msgs_for_frontend, pointers
|
||||
msgs = {k: [i["tag"] for i in v] for k, v in msgs_for_frontend.items()}
|
||||
pointers = pointers
|
||||
msgs = {k: [i["tag"] for i in task.messages] for k, task in rdagent_processes.items()}
|
||||
pointers = {k: dict(task.pointers) for k, task in rdagent_processes.items()}
|
||||
return jsonify({"msgs": msgs, "pointers": pointers}), 200
|
||||
|
||||
|
||||
@@ -256,12 +550,11 @@ def index():
|
||||
|
||||
@app.route("/<path:fn>", methods=["GET"])
|
||||
def server_static_files(fn):
|
||||
return send_from_directory(app.static_folder, fn)
|
||||
return send_from_directory(app.static_folder, _normalize_static_request_path(fn))
|
||||
|
||||
|
||||
def main(port: int = 19899):
|
||||
global server_port
|
||||
server_port = port
|
||||
app.config["UI_SERVER_PORT"] = port
|
||||
app.run(debug=False, host="0.0.0.0", port=port)
|
||||
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ class UIBasePropSetting(ExtendedBaseSettings):
|
||||
|
||||
static_path: str = "./git_ignore_folder/static"
|
||||
|
||||
trace_folder: str = "./traces"
|
||||
trace_folder: str = "./git_ignore_folder/traces"
|
||||
|
||||
enable_cache: bool = True
|
||||
|
||||
|
||||
@@ -33,10 +33,11 @@ class WebStorage(Storage):
|
||||
def log(self, obj: object, tag: str, timestamp: datetime | None = None, **kwargs: Any) -> str | Path:
|
||||
timestamp = gen_datetime(timestamp)
|
||||
if "pdf_image" in tag or "load_pdf_screenshot" in tag:
|
||||
obj.save(f"{UI_SETTING.static_path}/{timestamp.isoformat()}.jpg")
|
||||
Path(f"{UI_SETTING.static_path}/pdf_images").mkdir(parents=True, exist_ok=True)
|
||||
obj.save(f"{UI_SETTING.static_path}/pdf_images/{timestamp.isoformat()}.jpg")
|
||||
|
||||
try:
|
||||
data = self._obj_to_json(obj=obj, tag=tag, id=self.path, timestamp=timestamp.isoformat())
|
||||
data = self._obj_to_json(obj=obj, tag=tag, id=str(self.path), timestamp=timestamp.isoformat())
|
||||
if not data:
|
||||
return "Normal log, skipped"
|
||||
if isinstance(data, list):
|
||||
@@ -48,7 +49,7 @@ class WebStorage(Storage):
|
||||
resp = requests.post(f"{self.url}/receive", json=data, headers=headers, timeout=1)
|
||||
return f"{resp.status_code} {resp.text}"
|
||||
except (requests.ConnectionError, requests.Timeout) as e:
|
||||
pass
|
||||
print(f"Failed to connect to the web storage server at {self.url}: {e}")
|
||||
|
||||
def truncate(self, time: datetime) -> None:
|
||||
self.msgs = [m for m in self.msgs if datetime.fromisoformat(m["msg"]["timestamp"]) <= time]
|
||||
@@ -100,7 +101,7 @@ class WebStorage(Storage):
|
||||
"tag": "research.pdf_image",
|
||||
"timestamp": timestamp,
|
||||
"loop_id": li,
|
||||
"content": {"image": f"{timestamp}.jpg"},
|
||||
"content": {"image": f"pdf_images/{timestamp}.jpg"},
|
||||
},
|
||||
}
|
||||
elif "experiment generation" in tag or "load_experiment" in tag:
|
||||
|
||||
@@ -92,7 +92,7 @@ def extract_loopid_func_name(tag: str) -> tuple[str, str] | tuple[None, None]:
|
||||
|
||||
def extract_evoid(tag: str) -> str | None:
|
||||
"""extract evo id from the tag in Message"""
|
||||
match = re.search(r"\.evo_loop_(\d+)\.", tag)
|
||||
match = re.search(r"evo_loop_(\d+)\.", tag)
|
||||
return cast(str, match.group(1)) if match else None
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user