- financial 默认增量:按 A 股披露节奏判断已最新并跳过;--full 强制全量重拉 - Tushare fina_indicator 增加报告期窗口与 100 条/请求自动分页(修复老报告期静默截断) - 新浪兜底收紧为校验兜底:两源重叠历史一致才导入缺失键,行标记 source=sina; 财务可比字段取 eps/销售毛利率(ROE 两端口径不同不作依据),日线只比较最近重叠交易日 - CLI 输出逐只进度与导入内容描述(来源/行数/报告期与公告区间),失败股票留待重跑 - financial_indicator 增 source 列(迁移 d3f6c9a21b04);新增一致性/分页/服务测试
271 lines
9.9 KiB
Python
271 lines
9.9 KiB
Python
"""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
|