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:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user