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