Files
qlib/backend/tests/test_repositories.py
T
Simon 195f5d41f4 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 干净。
2026-09-06 22:12:44 +08:00

241 lines
8.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Repository 集成测试:临时 SQLite 上的幂等 upsert / 查询 / 防未来函数过滤。"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pytest
from app.domain.entities.market import (
AdjustFactor,
DailyBar,
FinancialIndicator,
Stock,
SyncLog,
TradingCalendar,
)
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.models.market import (
AdjustFactorModel,
FinancialIndicatorModel,
StockDailyModel,
StockModel,
SyncLogModel,
TradingCalendarModel,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyAdjustFactorRepository,
SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository,
SqlAlchemySyncLogRepository,
SqlAlchemyTradingCalendarRepository,
)
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
@pytest.fixture()
def session(tmp_path) -> Session:
engine = create_engine(f"sqlite:///{tmp_path / 'repo.db'}", future=True)
Base.metadata.create_all(engine)
with Session(engine) as session:
yield session
def _count(session, model) -> int:
return session.scalar(select(func.count()).select_from(model))
class TestStockRepository:
def test_upsert_idempotent_and_update(self, session: Session) -> None:
repo = SqlAlchemyStockRepository(session)
s1 = Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))
s2 = Stock(symbol="000001.SZ", name="平安银行", list_date=date(1991, 4, 3))
assert repo.upsert_many([s1, s2]) == 2
session.commit()
assert _count(session, StockModel) == 2
# 幂等:再次 upsert 不新增
repo.upsert_many([s1, s2])
session.commit()
assert _count(session, StockModel) == 2
# 更新既有记录
renamed = s1.model_copy(update={"name": "贵州茅台(更新)"})
repo.upsert_many([renamed])
session.commit()
got = repo.get_by_symbol("600519.SH")
assert got is not None
assert got.name == "贵州茅台(更新)"
class TestDailyBarRepository:
def _bar(self, day: str) -> DailyBar:
return DailyBar(
symbol="600519.SH",
trade_date=date.fromisoformat(day),
open=Decimal("100"),
high=Decimal("101"),
low=Decimal("99"),
close=Decimal("100.5"),
volume=Decimal("10000"),
amount=Decimal("1000000"),
)
def test_upsert_and_get_range(self, session: Session) -> None:
repo = SqlAlchemyDailyBarRepository(session)
bars = [self._bar("2024-01-02"), self._bar("2024-01-03"), self._bar("2024-01-04")]
repo.upsert_many(bars)
session.commit()
assert _count(session, StockDailyModel) == 3
repo.upsert_many([self._bar("2024-01-03")]) # 幂等
session.commit()
assert _count(session, StockDailyModel) == 3
got = repo.get_range("600519.SH", date(2024, 1, 3), date(2024, 1, 4))
assert [b.trade_date.isoformat() for b in got] == ["2024-01-03", "2024-01-04"]
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:
return FinancialIndicator(
symbol="600519.SH",
report_date=date.fromisoformat(report),
announce_date=date.fromisoformat(announce),
eps=Decimal("1.2"),
)
def test_list_announced_blocks_future(self, session: Session) -> None:
repo = SqlAlchemyFinancialRepository(session)
repo.upsert_many(
[
self._fin("2024-08-15"),
self._fin("2024-08-31"),
self._fin("2024-09-20"),
self._fin("2024-10-30", report="2024-09-30"),
] # Q3 财报
)
session.commit()
# as_of=2024-08-31:只能看到 08-15 与 08-31 两条公告
visible = repo.list_announced("600519.SH", as_of_date=date(2024, 8, 31))
assert len(visible) == 2
assert all(f.announce_date <= date(2024, 8, 31) for f in visible)
assert [f.announce_date.day for f in visible] == [15, 31]
# 报告期约束:只看 Q3 及以后(report_date >= 2024-09-01)
narrowed = repo.list_announced(
"600519.SH", as_of_date=date(2024, 12, 31), report_start=date(2024, 9, 1)
)
assert len(narrowed) == 1
assert narrowed[0].announce_date == date(2024, 10, 30)
def test_upsert_batch_duplicate_key_takes_latest(self, session: Session) -> None:
"""同一批内出现重复幂等键(数据源偶发)不得冲突,后值覆盖。"""
repo = SqlAlchemyFinancialRepository(session)
first = self._fin("2024-08-15")
later = self._fin("2024-08-15").model_copy(update={"eps": Decimal("9.99")})
repo.upsert_many([first, later])
session.commit()
assert _count(session, FinancialIndicatorModel) == 1
got = repo.list_announced("600519.SH", as_of_date=date(2024, 12, 31))
assert len(got) == 1
assert got[0].eps == Decimal("9.99")
class TestSyncLogRepository:
def test_add_and_recent(self, session: Session) -> None:
repo = SqlAlchemySyncLogRepository(session)
repo.add(SyncLog(source="tushare", api="daily", success=True, row_count=3))
repo.add(SyncLog(source="sina", api="daily", success=False, failure_reason="timeout"))
session.commit()
assert _count(session, SyncLogModel) == 2
recent = repo.recent(source="tushare", limit=10)
assert len(recent) == 1
assert recent[0].source == "tushare"
assert recent[0].row_count == 3
class TestOtherRepos:
def test_calendar_and_factor(self, session: Session) -> None:
cal = SqlAlchemyTradingCalendarRepository(session)
cal.upsert_many(
[
TradingCalendar(calendar_date=date(2024, 1, 2)),
TradingCalendar(calendar_date=date(2024, 1, 3), is_open=False),
]
)
session.commit()
assert cal.is_open(date(2024, 1, 2))
assert not cal.is_open(date(2024, 1, 3))
assert len(cal.list_range(date(2024, 1, 1), date(2024, 1, 5))) == 2
assert _count(session, TradingCalendarModel) == 2
adj = SqlAlchemyAdjustFactorRepository(session)
adj.upsert_many(
[
AdjustFactor(
symbol="600519.SH", trade_date=date(2024, 1, 2), factor=Decimal("12.3456")
)
]
)
session.commit()
factors = adj.get_range("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
assert len(factors) == 1
assert float(factors[0].factor) == pytest.approx(12.3456)
assert _count(session, AdjustFactorModel) == 1