diff --git a/backend/app/api/experiments.py b/backend/app/api/experiments.py index e2d6bb5..2ded92b 100644 --- a/backend/app/api/experiments.py +++ b/backend/app/api/experiments.py @@ -8,8 +8,7 @@ from datetime import datetime from fastapi import APIRouter, BackgroundTasks, HTTPException from app.api.deps import DbSession, ExperimentRepoDep, JobRepoDep -from app.api.jobs import _bg_factories -from app.application.services.job_executor import execute_job, new_id +from app.application.services.job_executor import new_id, run_job_background from app.domain.entities.research import ( ExperimentRecord, JobRecord, @@ -77,5 +76,5 @@ def rerun_experiment( ) job_repo.create(job) session.commit() - background.add_task(execute_job, job.id, **_bg_factories()) + background.add_task(run_job_background, job.id) return {"job_id": job.id, "status": job.status, "origin_experiment": exp.id} diff --git a/backend/app/api/jobs.py b/backend/app/api/jobs.py index 6cc4ac7..12c86da 100644 --- a/backend/app/api/jobs.py +++ b/backend/app/api/jobs.py @@ -1,8 +1,11 @@ """异步研究 Job API(Phase 4):提交 / 查询 / SSE 进度。 -POST /api/jobs 创建 Job(BackgroundTasks 本地执行),立即返回 job_id +POST /api/jobs 创建 Job(BackgroundTasks 后台执行),立即返回 job_id GET /api/jobs/{id} 状态 + 结果(成功时内嵌 result) GET /api/jobs/{id}/events SSE 进度(queued→running→success|failed) + +执行模式见 config job.mode:subprocess 时研究任务在独立子进程跑(内存隔离), +API worker 不被重任务拖垮(内存优化专项)。 """ from __future__ import annotations @@ -15,7 +18,7 @@ from fastapi import APIRouter, BackgroundTasks, HTTPException from fastapi.responses import StreamingResponse from app.api.deps import DbSession, JobRepoDep -from app.application.services.job_executor import execute_job, new_id +from app.application.services.job_executor import new_id, run_job_background from app.domain.entities.research import ( BacktestResult, FactorTestReport, @@ -31,12 +34,6 @@ from app.infrastructure.persistence.sqlalchemy.session import SessionLocal router = APIRouter(prefix="/jobs", tags=["jobs"]) -def _bg_factories(): - from app.application.services.job_executor import default_factories - - return default_factories() - - def _decode_result(kind: str, result_json: str | None): if result_json is None: return None @@ -68,7 +65,7 @@ def create_job( ) job_repo.create(job) session.commit() - background.add_task(execute_job, job.id, **_bg_factories()) + background.add_task(run_job_background, job.id) return {"job_id": job.id, "status": job.status} diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py index 613be81..0fa2965 100644 --- a/backend/app/application/services/job_executor.py +++ b/backend/app/application/services/job_executor.py @@ -1,14 +1,23 @@ -"""Job 执行编排与 Experiment 归档(Phase 4)。 +"""Job 执行编排与 Experiment 归档(Phase 4)+ 内存隔离执行(内存优化专项)。 -execute_job 在后台运行:状态机 queued→running→(success|failed), +execute_job 运行状态机 queued→running→(success|failed), 成功时自动把 spec/result 存为 Experiment(含代码版本),实现「研究可复现」(AGENT.md §21)。 独立 Session 生命周期(不依赖请求 scope),可被 FastAPI BackgroundTasks 或测试直接调用。 + +执行模式(config job.mode): +- local —— 本进程内执行(开发 / 单测;无内存隔离) +- subprocess —— 独立子进程执行(python -m app.cli.run_job ), + 子进程内施加 RLIMIT_AS 上限(config job.max_memory_gb),研究任务 OOM 只会 + MemoryError 失败归档,不会拖垮 API worker;父进程在子进程异常退出时补记 failed。 """ from __future__ import annotations import json +import os import subprocess +import sys +import threading import uuid from collections.abc import Callable from datetime import datetime @@ -24,7 +33,13 @@ from app.domain.entities.research import ( ) from app.quant.service import ResearchService -PROJECT_ROOT = Path(__file__).resolve().parents[3] +# backend/app/application/services/job_executor.py → parents[4] = 项目根 +PROJECT_ROOT = Path(__file__).resolve().parents[4] +BACKEND_ROOT = PROJECT_ROOT / "backend" + +# 本进程内正在运行的子进程 Job(并发上限保护宿主内存) +_active_jobs: dict[str, subprocess.Popen | None] = {} +_active_lock = threading.Lock() def new_id(prefix: str) -> str: @@ -171,7 +186,10 @@ def default_factories() -> dict: def submit_and_run(spec: ResearchSpec, *, factories: dict | None = None) -> JobRecord: - """同步创建并执行一个 Job(复用 Job 状态机与 Experiment 归档),返回终态 Job。""" + """创建并执行一个 Job(复用 Job 状态机与 Experiment 归档),返回终态 Job。 + + 按 job.mode 调度:subprocess 模式在独立进程内执行(内存隔离),否则本进程。 + """ facts = factories or default_factories() session_factory = facts["session_factory"] job_repo = facts["job_repo_factory"] @@ -185,16 +203,178 @@ def submit_and_run(spec: ResearchSpec, *, factories: dict | None = None) -> JobR with session_factory() as session: job_repo(session).create(job) session.commit() - execute_job( - job.id, - session_factory=session_factory, - job_repo_factory=job_repo, - experiment_repo_factory=facts["experiment_repo_factory"], - stock_repo_factory=facts["stock_repo_factory"], - daily_repo_factory=facts["daily_repo_factory"], - engine=facts["engine"], - ) + if _job_mode() == "subprocess": + _run_in_subprocess(job.id, session_factory=session_factory, job_repo_factory=job_repo) + else: + execute_job( + job.id, + session_factory=session_factory, + job_repo_factory=job_repo, + experiment_repo_factory=facts["experiment_repo_factory"], + stock_repo_factory=facts["stock_repo_factory"], + daily_repo_factory=facts["daily_repo_factory"], + engine=facts["engine"], + ) with session_factory() as session: done = job_repo(session).get(job.id) assert done is not None return done + + +# --------------------------------------------------------------------------- +# 执行模式调度(config job.mode:local | subprocess) +# --------------------------------------------------------------------------- + + +def _job_mode() -> str: + from app.core.config import get_settings + + return get_settings().job_mode + + +def _job_memory_limit_gb() -> int: + from app.core.config import get_settings + + return get_settings().job_memory_limit_gb + + +def _job_max_concurrent() -> int: + from app.core.config import get_settings + + return get_settings().job_max_concurrent + + +def _job_worker_cmd(job_id: str) -> list[str]: + """研究子进程命令行(独立进程执行,cwd=backend 使 `-m app.cli.run_job` 可导入)。""" + return [sys.executable, "-m", "app.cli.run_job", job_id] + + +def run_job_background( + job_id: str, + *, + session_factory: Callable | None = None, + job_repo_factory: Callable | None = None, +) -> None: + """API 后台任务入口:按 job.mode 调度 Job 执行。""" + if _job_mode() == "subprocess": + _run_in_subprocess( + job_id, session_factory=session_factory, job_repo_factory=job_repo_factory + ) + else: + execute_job(job_id, **default_factories()) + + +def _acquire_slot(job_id: str) -> bool: + with _active_lock: + if len(_active_jobs) >= max(_job_max_concurrent(), 1): + return False + _active_jobs[job_id] = None + return True + + +def _release_slot(job_id: str) -> None: + with _active_lock: + _active_jobs.pop(job_id, None) + + +def _mark_failed( + job_id: str, + error: str, + *, + session_factory: Callable | None = None, + job_repo_factory: Callable | None = None, +) -> None: + """把非终态 Job 标记 failed(子进程异常退出 / 并发超限时兜底,防永久 running)。""" + if session_factory is None or job_repo_factory is None: + facts = default_factories() + session_factory = session_factory or facts["session_factory"] + job_repo_factory = job_repo_factory or facts["job_repo_factory"] + try: + with session_factory() as session: + repo = job_repo_factory(session) + job = repo.get(job_id) + if job is not None and job.status in (JobStatus.QUEUED, JobStatus.RUNNING): + job.status = JobStatus.FAILED + job.error = error + job.finished_at = datetime.now() + repo.update(job) + session.commit() + except Exception: # noqa: BLE001 —— 兜底标记失败自身异常不向上抛 + pass + + +def _run_in_subprocess( + job_id: str, + *, + session_factory: Callable | None = None, + job_repo_factory: Callable | None = None, +) -> None: + """在独立 python 进程执行 Job:子进程 RLIMIT 上限(防 OOM 整机), + 父进程等待;子进程异常退出(如被杀)时把 Job 补记 failed。""" + if not _acquire_slot(job_id): + _mark_failed( + job_id, + "系统繁忙:并发研究任务已达上限,请稍后重试", + session_factory=session_factory, + job_repo_factory=job_repo_factory, + ) + return + env = dict(os.environ) + env["QLIB_JOB_MEM_LIMIT_GB"] = str(_job_memory_limit_gb()) + rc = -1 + already_marked = False + try: + proc = subprocess.Popen( + _job_worker_cmd(job_id), + cwd=BACKEND_ROOT, + env=env, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + with _active_lock: + _active_jobs[job_id] = proc + rc = proc.wait() + except Exception as exc: # noqa: BLE001 + _mark_failed( + job_id, + f"研究子进程启动失败: {type(exc).__name__}: {exc}", + session_factory=session_factory, + job_repo_factory=job_repo_factory, + ) + already_marked = True + finally: + _release_slot(job_id) + if rc != 0 and not already_marked: + _mark_failed( + job_id, + f"研究子进程异常退出(code={rc}):任务未完成", + session_factory=session_factory, + job_repo_factory=job_repo_factory, + ) + + +def mark_stale_jobs_failed( + *, + session_factory: Callable | None = None, + job_repo_factory: Callable | None = None, + reason: str | None = None, +) -> int: + """服务启动调用:把上次进程异常退出遗留的 queued/running Job 标记 failed。 + + 返回处理数量。防「进程被杀后 Job 永久 running」(AGENT.md §20 状态机闭环)。 + """ + facts = default_factories() + sf = session_factory or facts["session_factory"] + jr = job_repo_factory or facts["job_repo_factory"] + counted = 0 + with sf() as session: + repo = jr(session) + for status in (JobStatus.QUEUED, JobStatus.RUNNING): + for job in repo.list_by_status(status): + job.status = JobStatus.FAILED + job.error = reason or "服务重启:上次未完成任务被中断" + job.finished_at = datetime.now() + repo.update(job) + counted += 1 + session.commit() + return counted diff --git a/backend/app/cli/run_job.py b/backend/app/cli/run_job.py new file mode 100644 index 0000000..f00deaf --- /dev/null +++ b/backend/app/cli/run_job.py @@ -0,0 +1,47 @@ +"""独立进程执行研究 Job —— 内存隔离的执行体(内存优化专项)。 + +由 job_executor._run_in_subprocess 以 `python -m app.cli.run_job ` 拉起 +(cwd=backend/),入口先施加 RLIMIT_AS 上限再导入重依赖,随后执行与 API 内 +execute_job 完全相同的状态机 / Experiment 归档逻辑。 + +环境变量: + QLIB_JOB_MEM_LIMIT_GB —— 子进程虚拟地址空间上限(GB);父进程 spawn 时写入, + 手动直接运行本命令而未设置时回退到 Settings(config job.max_memory_gb)。 +""" + +from __future__ import annotations + +import os +import sys + + +def _apply_memory_limit() -> None: + """设置 RLIMIT_AS(Linux)。达到上限时分配抛 MemoryError → Job 归档 failed, + 而不是让内核 OOM 杀掉整机(8G Pi5 上保护同机其它服务)。""" + try: + gb = int(os.environ.get("QLIB_JOB_MEM_LIMIT_GB") or "") + except (TypeError, ValueError): + from app.core.config import get_settings + + gb = get_settings().job_memory_limit_gb + limit = gb * 1024**3 + import resource + + resource.setrlimit(resource.RLIMIT_AS, (limit, limit)) + + +def main(argv: list[str] | None = None) -> int: + args = list(sys.argv[1:] if argv is None else argv) + if len(args) < 1: + print("用法: python -m app.cli.run_job ", file=sys.stderr) + return 2 + _apply_memory_limit() + # 先设内存上限再导入重依赖(pandas / sqlalchemy 等) + from app.application.services.job_executor import default_factories, execute_job + + execute_job(args[0], **default_factories()) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/backend/app/core/config.py b/backend/app/core/config.py index ceca378..405db12 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -72,6 +72,9 @@ class Settings: llm_base_url: str | None llm_model: str storage: dict[str, Path] + job_mode: str + job_memory_limit_gb: int + job_max_concurrent: int config_path: Path = CONFIG_PATH env_path: Path = ENV_PATH project_root: Path = PROJECT_ROOT @@ -118,6 +121,19 @@ def get_settings() -> Settings: "app/infrastructure/persistence/migrations" ) + # Job 执行:环境变量 JOB_MODE / QLIB_JOB_MEM_LIMIT_GB 可覆盖 config.yaml(测试用) + job_cfg = _deep(cfg, "job") or {} + try: + job_memory_limit_gb = int( + os.environ.get("QLIB_JOB_MEM_LIMIT_GB") or job_cfg.get("max_memory_gb") or 6 + ) + except (TypeError, ValueError): + job_memory_limit_gb = 6 + try: + job_max_concurrent = int(os.environ.get("QLIB_JOB_MAX_CONCURRENT") or job_cfg.get("max_concurrent_jobs") or 2) + except (TypeError, ValueError): + job_max_concurrent = 2 + return Settings( app_name=str(_deep(cfg, "app.name") or "qlib-platform"), app_version=str(_deep(cfg, "app.version") or "0.1.0"), @@ -134,4 +150,7 @@ def get_settings() -> Settings: llm_base_url=os.environ.get(llm_url_env) or (agent_llm.get("base_url") or None), llm_model=os.environ.get(llm_model_env) or agent_llm.get("model") or "qwen-plus", storage=_resolve_storage_dirs(cfg), + job_mode=os.environ.get("JOB_MODE") or job_cfg.get("mode") or "subprocess", + job_memory_limit_gb=job_memory_limit_gb, + job_max_concurrent=job_max_concurrent, ) diff --git a/backend/app/domain/repositories/jobs.py b/backend/app/domain/repositories/jobs.py index 5fe036c..3f1f87b 100644 --- a/backend/app/domain/repositories/jobs.py +++ b/backend/app/domain/repositories/jobs.py @@ -16,6 +16,9 @@ class JobRepository(Protocol): def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]: ... + def list_by_status(self, status: str, limit: int = 100) -> list[JobRecord]: + """按状态查询(服务启动清理残留 queued/running 用)。""" + class ExperimentRepository(Protocol): def save(self, experiment: ExperimentRecord) -> ExperimentRecord: ... diff --git a/backend/app/domain/repositories/market.py b/backend/app/domain/repositories/market.py index bb124ff..b24ed0c 100644 --- a/backend/app/domain/repositories/market.py +++ b/backend/app/domain/repositories/market.py @@ -6,7 +6,7 @@ from __future__ import annotations -from collections.abc import Sequence +from collections.abc import Iterator, Sequence from datetime import date from typing import Protocol @@ -45,6 +45,20 @@ class DailyBarRepository(Protocol): def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]: """批量区间查询(研究服务装配面板用,避免逐只查询)。""" + def stream_range_many_columns( + self, + symbols: Sequence[str], + start: date, + end: date, + columns: Sequence[str], + ) -> Iterator[tuple]: + """流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。 + + 研究装配大数据面板专用:只 SELECT 所需列并在 SQL 侧转 REAL, + 避免 ORM 对象 / Decimal 全量物化(内存大头,见内存优化专项)。 + 实现可选——ResearchService 会对缺失该方法的老实现回退到 get_range_many。 + """ + def latest_date(self, symbol: str) -> date | None: """断点续传用:该股票本地已有数据的最新交易日。""" diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py index ee3fd56..3a9287e 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py @@ -42,6 +42,15 @@ class SqlAlchemyJobRepository: for r in self._session.scalars(stmt).all() ] + def list_by_status(self, status: str, limit: int = 100) -> list[JobRecord]: + rows = self._session.scalars( + select(JobModel) + .where(JobModel.status == status) + .order_by(JobModel.created_at) + .limit(limit) + ).all() + return [JobRecord.model_validate(r, from_attributes=True) for r in rows] + class SqlAlchemyExperimentRepository: def __init__(self, session: Session) -> None: diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py index 886c7a3..c6f2bb0 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py @@ -8,11 +8,11 @@ Repository 以 domain.entities 类型进出(AGENT.md §10)。 from __future__ import annotations -from collections.abc import Sequence +from collections.abc import Iterator, Sequence from datetime import date from typing import Any -from sqlalchemy import select +from sqlalchemy import Float, String, cast, select from sqlalchemy.orm import Session from app.domain.entities.market import ( @@ -32,6 +32,9 @@ from app.infrastructure.persistence.sqlalchemy.models.market import ( TradingCalendarModel, ) +# 日线数值列白名单(研究面板只需这些;symbol/trade_date 恒返回) +BAR_FLOAT_COLUMNS = ("open", "high", "low", "close", "volume", "amount") + # 实体类型 → (ORM Model, 幂等键列) _TABLE = { Stock: (StockModel, ["symbol"]), @@ -167,6 +170,41 @@ class SqlAlchemyDailyBarRepository: ).all() return [DailyBar.model_validate(r, from_attributes=True) for r in rows] + def stream_range_many_columns( + self, + symbols: Sequence[str], + start: date, + end: date, + columns: Sequence[str], + ) -> Iterator[tuple]: + """流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。 + + 内存优化:与 get_range_many 不同,不实例化 ORM 对象 / Decimal, + 只 SELECT 所需列并在 SQL 侧 CAST 为 REAL,适合一次装配几十万~几百万行面板。 + """ + cols = list(columns) + unknown = [c for c in cols if c not in BAR_FLOAT_COLUMNS] + if unknown: + raise ValueError(f"不支持的行情列: {unknown}(可用: {BAR_FLOAT_COLUMNS})") + numeric_expr = [cast(getattr(StockDailyModel, c), Float) for c in cols] + stmt = ( + select(StockDailyModel.symbol, cast(StockDailyModel.trade_date, String), *numeric_expr) + .where( + StockDailyModel.symbol.in_(list(symbols)), + StockDailyModel.trade_date >= start, + StockDailyModel.trade_date <= end, + ) + .order_by(StockDailyModel.symbol, StockDailyModel.trade_date) + .execution_options(yield_per=20000) + ) + result = self._session.execute(stmt) + while True: + chunk = result.fetchmany(20000) + if not chunk: + break + for row in chunk: + yield tuple(row) + def latest_date(self, symbol: str) -> date | None: return self._session.scalar( select(StockDailyModel.trade_date) diff --git a/backend/app/main.py b/backend/app/main.py index 6f95932..ddb87a6 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1,22 +1,54 @@ """FastAPI 应用入口。 启动:cd backend && uv run uvicorn app.main:app --reload + +启动时(lifespan)清理上次进程异常退出遗留的 queued/running Job → +failed,防止「进程被杀后 Job 永久 running」(AGENT.md §20 状态机闭环)。 +测试可用环境变量 QLIB_SKIP_STALE_JOB_CLEANUP=1 关闭(不触碰开发库)。 """ from __future__ import annotations +import logging +import os +from contextlib import asynccontextmanager + from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from app.api.router import api_router from app.core.config import get_settings +logger = logging.getLogger("uvicorn.error") + settings = get_settings() + +def _cleanup_stale_jobs() -> None: + """把上次进程异常退出遗留的 queued/running Job 标记 failed。""" + if os.environ.get("QLIB_SKIP_STALE_JOB_CLEANUP") == "1": + return + try: + from app.application.services.job_executor import mark_stale_jobs_failed + + n = mark_stale_jobs_failed(reason="服务重启:上次未完成任务被中断") + if n: + logger.info("启动清理:%d 个遗留 Job 已标记 failed", n) + except Exception: # noqa: BLE001 —— 启动清理失败不阻断服务 + logger.exception("启动清理残留 Job 失败") + + +@asynccontextmanager +async def lifespan(_app: FastAPI): + _cleanup_stale_jobs() + yield + + app = FastAPI( title=settings.app_name, version=settings.app_version, description="A股个人量化研究平台 API(研究引擎 / Tushare 数据源)", + lifespan=lifespan, ) # 开发期 CORS:允许任意来源(个人本地平台,无 Cookie 凭据,allow_credentials=False)。 diff --git a/backend/app/quant/engine.py b/backend/app/quant/engine.py index 62f454b..7490268 100644 --- a/backend/app/quant/engine.py +++ b/backend/app/quant/engine.py @@ -10,6 +10,7 @@ from typing import Protocol import pandas as pd from app.domain.entities.research import BacktestResult, FactorTestReport, ResearchSpec +from app.quant.factors import FactorError, get_factor from app.quant.local_engine import ( TopKBacktestRunner, build_factor_panels, @@ -17,12 +18,30 @@ from app.quant.local_engine import ( run_spec_factor_test, ) +# LocalEngine 路径恒需 close(TopK 收盘撮合 / 前瞻收益) +_CLOSE = {"close"} + + +def factor_required_columns(spec: ResearchSpec) -> set[str]: + """spec 因子计算 + 回测撮合所需的行情数值列(含 close)。""" + needed = set(_CLOSE) + for fs in spec.factors: + try: + defn, _fn = get_factor(fs.name) + except FactorError: + continue # 未知因子由执行期统一报错 + needed.update(defn.requires) + return needed + class QuantEngine(Protocol): """研究引擎端口:因子面板构建 / 因子测试 / 回测。""" name: str + def required_columns(self, spec: ResearchSpec) -> set[str]: + """执行该 spec 所需的行情数值列(数据装配按此裁剪,控制内存)。""" + def run_factor_test( self, daily: pd.DataFrame, spec: ResearchSpec, horizon_days: int = 21 ) -> FactorTestReport: ... @@ -35,6 +54,9 @@ class LocalEngine: name = "local" + def required_columns(self, spec: ResearchSpec) -> set[str]: + return factor_required_columns(spec) + def run_factor_test( self, daily: pd.DataFrame, spec: ResearchSpec, horizon_days: int = 21 ) -> FactorTestReport: diff --git a/backend/app/quant/qlib_adapter/engine.py b/backend/app/quant/qlib_adapter/engine.py index 0caf36f..45c34da 100644 --- a/backend/app/quant/qlib_adapter/engine.py +++ b/backend/app/quant/qlib_adapter/engine.py @@ -39,10 +39,21 @@ class QlibEngine(QuantEngine): name = "qlib" + # qlib 落盘需要 OHLCV(vwap/factor 由本地合成,不来自行情表) + _DUMP_COLUMNS = {"open", "high", "low", "close", "volume", "amount"} + def __init__(self, qlib_dir: Path | None = None) -> None: # 默认落盘到 data/qlib(与 storage.qlib_dir 一致);可注入临时目录便于测试 self.qlib_dir = qlib_dir or _default_qlib_dir() + def required_columns(self, spec: ResearchSpec) -> set[str]: + # 回测需把全字段落盘成 qlib 数据集;因子测试走共享实现,只需 close+因子字段 + if spec.type == "backtest": + return set(self._DUMP_COLUMNS) + from app.quant.engine import factor_required_columns + + return factor_required_columns(spec) + def run_factor_test( self, daily: pd.DataFrame, spec: ResearchSpec, horizon_days: int = 21 ) -> FactorTestReport: diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index a4e7a35..e5d3a03 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -2,10 +2,15 @@ 本层是业务入口:API / Agent 只能调用这里的用例(AGENT.md §16/§17), 禁止直接拼接引擎配置。数据一律经 Repository 获取(防未来函数由查询层保证)。 + +内存优化:大数据面板优先走 Repository 的流式列裁剪查询 +(stream_range_many_columns,SQL 侧转 REAL、分批拉取),避免 ORM 对象 / +Decimal 全量物化;老实现回退到 get_range_many 逐实体路径。 """ from __future__ import annotations +from collections.abc import Iterable from datetime import date, timedelta import pandas as pd @@ -23,6 +28,9 @@ from app.domain.repositories.market import ( ) from app.quant.engine import QuantEngine +# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值) +_FRAME_CHUNK_ROWS = 50_000 + def filter_stocks(stocks: list[Stock], universe: UniverseSpec, as_of: date) -> list[Stock]: """按股票池口径过滤(名称含 ST 判定 —— 名称快照为当日口径,属历史可追溯数据)。""" @@ -55,6 +63,31 @@ def bars_to_daily_df(bars) -> pd.DataFrame: return df +def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFrame: + """把流式 (symbol, trade_date_iso, *float_cols) 分批拼成 float 长表。 + + 全程只保留分批 DataFrame + 最终一份结果,避免整批 tuple/Decimal 同时驻留。 + 返回列:symbol, trade_date(datetime64), *columns(float64)。 + """ + cols = ["symbol", "trade_date", *columns] + buf: list[tuple] = [] + pieces: list[pd.DataFrame] = [] + for row in rows: + buf.append(row) + if len(buf) >= _FRAME_CHUNK_ROWS: + pieces.append(pd.DataFrame(buf, columns=cols)) + buf = [] + if buf: + pieces.append(pd.DataFrame(buf, columns=cols)) + if not pieces: + return pd.DataFrame() + df = pd.concat(pieces, ignore_index=True) + df["trade_date"] = pd.to_datetime(df["trade_date"]) + for col in columns: # NULL → NaN,统一 float64 + df[col] = pd.to_numeric(df[col], errors="coerce") + return df + + class ResearchService: """研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。""" @@ -89,9 +122,22 @@ class ResearchService: stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start) if not stocks: return pd.DataFrame() + symbols = [s.symbol for s in stocks] + + # 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV) + required = self._engine.required_columns(spec) + streamer = getattr(self._daily_repo, "stream_range_many_columns", None) + if streamer is not None: + try: + return _frame_from_stream( + streamer(symbols, data_start, end, sorted(required)), sorted(required) + ) + except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现) + pass + # 旧路径:逐实体(供内存 / Fake 仓储等实现使用) get_many = getattr(self._daily_repo, "get_range_many", None) if get_many is not None: - bars = list(get_many([s.symbol for s in stocks], data_start, end)) + bars = list(get_many(symbols, data_start, end)) else: # 兜底:逐只查询 bars = [] for s in stocks: diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py new file mode 100644 index 0000000..fb747f5 --- /dev/null +++ b/backend/tests/conftest.py @@ -0,0 +1,16 @@ +"""全局测试配置。 + +- JOB_MODE=local:单测/CI 里 Job 在本进程执行(不 spawn 子进程、不依赖安装态) +- QLIB_SKIP_STALE_JOB_CLEANUP=1:TestClient 启动 lifespan 不清理残留 Job + (避免连接/写入开发库 data/quant.db) + +必须在任何 app / config 模块首次导入前生效,故放在 conftest 模块顶层。 +""" + +from __future__ import annotations + +import os + +# 强制(不是 setdefault):测试绝不 spawn 研究子进程 / 不触碰开发库 +os.environ["JOB_MODE"] = "local" +os.environ["QLIB_SKIP_STALE_JOB_CLEANUP"] = "1" diff --git a/backend/tests/test_jobs_experiments.py b/backend/tests/test_jobs_experiments.py index 0f7370a..cb3da60 100644 --- a/backend/tests/test_jobs_experiments.py +++ b/backend/tests/test_jobs_experiments.py @@ -2,9 +2,11 @@ from __future__ import annotations +import sys from datetime import date, datetime import pytest +from app.application.services import job_executor as je from app.application.services.job_executor import execute_job from app.domain.entities.research import ( JobRecord, @@ -147,6 +149,64 @@ class TestJobExecutor: assert done.error +class TestJobResilience: + """内存隔离专项:启动残留清理 + 子进程异常退出兜底。""" + + def test_mark_stale_jobs_failed(self, sf) -> None: + repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731 + rows = [ + ("JOB-STALE-1", JobStatus.RUNNING), + ("JOB-STALE-2", JobStatus.QUEUED), + ("JOB-OK-3", JobStatus.SUCCESS), + ] + for jid, status in rows: + job = JobRecord( + id=jid, + kind="backtest", + spec_json=_spec().model_dump_json(), + status=status, + created_at=datetime.now(), + ) + with sf() as session: + repo_f(session).create(job) + session.commit() + + n = je.mark_stale_jobs_failed(session_factory=sf, job_repo_factory=repo_f) + assert n == 2 + with sf() as session: + stale = repo_f(session).get("JOB-STALE-1") + queued = repo_f(session).get("JOB-STALE-2") + ok = repo_f(session).get("JOB-OK-3") + assert stale is not None and stale.status == JobStatus.FAILED + assert stale.finished_at is not None and "重启" in (stale.error or "") + assert queued is not None and queued.status == JobStatus.FAILED + assert ok is not None and ok.status == JobStatus.SUCCESS + + def test_subprocess_abnormal_exit_marks_failed(self, sf, monkeypatch) -> None: + repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731 + job = JobRecord( + id="JOB-SUB-1", + kind="backtest", + spec_json=_spec().model_dump_json(), + status=JobStatus.QUEUED, + created_at=datetime.now(), + ) + with sf() as session: + repo_f(session).create(job) + session.commit() + monkeypatch.setattr(je, "_job_mode", lambda: "subprocess") + monkeypatch.setattr( + je, "_job_worker_cmd", lambda job_id: [sys.executable, "-c", "import sys; sys.exit(7)"] + ) + je.run_job_background("JOB-SUB-1", session_factory=sf, job_repo_factory=repo_f) + + with sf() as session: + done = repo_f(session).get("JOB-SUB-1") + assert done is not None + assert done.status == JobStatus.FAILED + assert done.error is not None and "code=7" in done.error + + @pytest.fixture() def seeded_api_db(tmp_path, monkeypatch) -> None: """把全局 SessionLocal 指向 tmp 种子库(5 只股票 × 300 交易日), diff --git a/backend/tests/test_quant_engine.py b/backend/tests/test_quant_engine.py index ca118a8..0358609 100644 --- a/backend/tests/test_quant_engine.py +++ b/backend/tests/test_quant_engine.py @@ -71,6 +71,19 @@ class TestBacktestMain: assert all(p.value <= 0 for p in res.drawdown) assert res.summary.max_drawdown_pct <= 0 + def test_required_columns_follow_factor_dependencies(self) -> None: + """数据装配按引擎所需列裁剪:momentum 只取 close;量比因子需 volume。""" + engine = LocalEngine() + momentum = _spec() # momentum_20 + assert engine.required_columns(momentum) == {"close"} + volume = _spec().model_copy( + update={"factors": [FactorSpec(name="volume_ratio_5_60")]} + ) + assert engine.required_columns(volume) == {"close", "volume"} + unknown = _spec().model_copy(update={"factors": [FactorSpec(name="no_such")]}) + # 未知因子不参与列裁剪,交由执行期统一报错 + assert engine.required_columns(unknown) == {"close"} + def _limit_up_scenario_daily() -> pd.DataFrame: """Y 在 2024-07-01 相对前一交易日跳涨 10.5%(主板涨停不可追),且动量高于 X。""" diff --git a/backend/tests/test_repositories.py b/backend/tests/test_repositories.py index a0c08c3..395e162 100644 --- a/backend/tests/test_repositories.py +++ b/backend/tests/test_repositories.py @@ -100,6 +100,55 @@ class TestDailyBarRepository: assert repo.latest_date("600519.SH") == date(2024, 1, 4) assert repo.latest_date("000001.SZ") is None + def test_stream_range_many_columns_subset_order_and_null(self, session: Session) -> None: + """流式列裁剪:只返回所需数值列、SQL 侧转 float、按 symbol/trade_date 升序。""" + repo = SqlAlchemyDailyBarRepository(session) + bars = [ + self._bar("2024-01-02"), + self._bar("2024-01-03"), + self._bar("2024-01-04"), + ] + other = [ + DailyBar( + symbol="000001.SZ", + trade_date=d.trade_date, + close=Decimal("9"), + volume=Decimal("1"), + ) + for d in bars + ] + repo.upsert_many([*bars, *other]) + session.commit() + + rows = list( + repo.stream_range_many_columns( + ["600519.SH"], date(2024, 1, 2), date(2024, 1, 4), ["close", "volume"] + ) + ) + assert rows == [ + ("600519.SH", "2024-01-02", 100.5, 10000.0), + ("600519.SH", "2024-01-03", 100.5, 10000.0), + ("600519.SH", "2024-01-04", 100.5, 10000.0), + ] + + # NULL 数值 → None;白名单外列报错 + null_bar = self._bar("2024-01-02").model_copy(update={"volume": None}) + repo.upsert_many([null_bar]) + session.commit() + rows2 = list( + repo.stream_range_many_columns( + ["600519.SH"], date(2024, 1, 2), date(2024, 1, 2), ["volume"] + ) + ) + assert rows2 == [("600519.SH", "2024-01-02", None)] + + with pytest.raises(ValueError): + list( + repo.stream_range_many_columns( + ["600519.SH"], date(2024, 1, 2), date(2024, 1, 4), ["close", "nope"] + ) + ) + class TestFinancialRepository: def _fin(self, announce: str, report: str = "2024-06-30") -> FinancialIndicator: diff --git a/config.yaml b/config.yaml index 51652da..283089e 100644 --- a/config.yaml +++ b/config.yaml @@ -36,8 +36,15 @@ storage: qlib_dir: "data/qlib" job: - # 第一阶段异步任务模式:local(FastAPI BackgroundTasks 级);Phase 复杂后再引入队列 - mode: "local" + # 研究任务执行模式: + # subprocess —— 独立子进程执行(内存隔离 + RLIMIT 上限,防研究任务 OOM 拖垮 API 服务) + # local —— 本进程内执行(开发 / 单测,无隔离) + # 可用环境变量 JOB_MODE 覆盖(测试强制 local)。 + mode: "subprocess" + # subprocess 模式下子进程虚拟地址空间上限(GB);达到上限任务以 MemoryError 失败,不会 OOM 整机 + max_memory_gb: 6 + # 并发研究子进程上限(超出直接失败并提示稍后再试) + max_concurrent_jobs: 2 agent: llm: