Files
qlib/backend/tests/test_repositories.py
T
Simon 442999f701 feat(data): 财务/日线同步增量 + 新浪「两边一致」校验兜底 + 逐只进度
- financial 默认增量:按 A 股披露节奏判断已最新并跳过;--full 强制全量重拉
- Tushare fina_indicator 增加报告期窗口与 100 条/请求自动分页(修复老报告期静默截断)
- 新浪兜底收紧为校验兜底:两源重叠历史一致才导入缺失键,行标记 source=sina;
  财务可比字段取 eps/销售毛利率(ROE 两端口径不同不作依据),日线只比较最近重叠交易日
- CLI 输出逐只进度与导入内容描述(来源/行数/报告期与公告区间),失败股票留待重跑
- financial_indicator 增 source 列(迁移 d3f6c9a21b04);新增一致性/分页/服务测试
2026-09-08 21:48:09 +08:00

271 lines
9.9 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")
def test_source_marker_and_incremental_queries(self, session: Session) -> None:
"""source 来源标记(默认 tushare)+ 增量判断查询(has_report_period/list_symbol)。"""
repo = SqlAlchemyFinancialRepository(session)
repo.upsert_many(
[
FinancialIndicator(
symbol="600519.SH",
report_date=date(2024, 6, 30),
announce_date=date(2024, 8, 31),
eps=Decimal("33.19"),
), # source 默认 tushare
FinancialIndicator(
symbol="600519.SH",
report_date=date(2024, 9, 30),
announce_date=date(2024, 10, 30),
source="sina",
eps=Decimal("48.42"),
),
]
)
session.commit()
rows = repo.list_symbol("600519.SH")
assert {r.source for r in rows} == {"tushare", "sina"}
assert rows[0].source == "tushare" # 按 announce_date 升序
assert repo.has_report_period("600519.SH", date(2024, 9, 30))
assert not repo.has_report_period("600519.SH", date(2024, 12, 31))
assert repo.list_symbol("000001.SZ") == []
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