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:
XianBW
2026-03-18 14:04:52 +08:00
committed by GitHub
parent 7cd64a26fd
commit 14395488b9
141 changed files with 12852 additions and 279 deletions
+9 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+1 -1
View File
@@ -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
+5 -4
View File
@@ -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:
+1 -1
View File
@@ -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