perf(backend): 内存优化三项——全市场研究不再占满 8G
1) 数据装配流式+列裁剪:Repository 新增 stream_range_many_columns(只 SELECT 所需列、SQL 侧转 REAL、yield_per 分批),引擎按 required_columns 取数 (LocalEngine 仅 close+因子字段),消除 ORM/Decimal 全量物化; 2) 研究 Job 独立子进程执行(job.mode=subprocess):python -m app.cli.run_job 在子进程内设 RLIMIT_AS 上限,OOM 归档 failed 而非拖垮 API worker; 子进程异常退出由父进程补记 failed;并发上限 2; 3) 服务启动清理:残留 queued/running Job 标记 failed(防永久 running)。 实测同款全市场回测:uvicorn worker RSS 稳定 ~220MB,任务峰值内存由 4.1GB+ 降至 ~470MB,24s 完成并归档(此前 43s 未完成即 OOM)。 新增/更新测试 96 passed,ruff 干净。
This commit is contained in:
@@ -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}
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
|
||||
@@ -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 <job_id>),
|
||||
子进程内施加 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
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
"""独立进程执行研究 Job —— 内存隔离的执行体(内存优化专项)。
|
||||
|
||||
由 job_executor._run_in_subprocess 以 `python -m app.cli.run_job <JOB_ID>` 拉起
|
||||
(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 <JOB_ID>", 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())
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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: ...
|
||||
|
||||
@@ -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:
|
||||
"""断点续传用:该股票本地已有数据的最新交易日。"""
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)。
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
@@ -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 交易日),
|
||||
|
||||
@@ -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。"""
|
||||
|
||||
@@ -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:
|
||||
|
||||
+9
-2
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user