feat(backend): Phase 1 数据层 — Domain / Provider / Failover 审计 + 持久化 + 同步 CLI
- domain:市场数据实体(Stock / 交易日历 / 日线 / 复权 / 财务含 announce_date)+ Repository 与 MarketDataProvider Protocol
- 数据源:TushareProvider(归一化、重试、鉴权错误归类)、SinaProvider(备用,明确前复权口径与能力边界)、FailoverProvider + SyncLog 审计(禁止静默切换)
- 持久化:SQLAlchemy 2.x Models + Repository 实现(按业务键幂等 upsert、as_of_date 防未来函数过滤)+ Alembic 迁移
- CLI:uv run python -m app.cli.sync {basic|calendar|daily|financial|verify},支持 --resume 断点续传
- 真实 Tushare 验证:stock 5556 / 交易日历 366 / daily+factor 242 / 财务 55;sync_log 审计完整
- 测试:38 passed(domain / provider / failover / repository / 未来函数 / 迁移),ruff clean
This commit is contained in:
@@ -0,0 +1,64 @@
|
||||
"""领域实体测试:字段约束与「防未来函数」可见性判断。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from app.domain.entities.market import (
|
||||
DailyBar,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
SyncLog,
|
||||
)
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
||||
class TestStock:
|
||||
def test_symbol_pattern_enforced(self) -> None:
|
||||
Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))
|
||||
with pytest.raises(ValidationError):
|
||||
Stock(symbol="600519", name="x", list_date=date(2001, 1, 1))
|
||||
with pytest.raises(ValidationError):
|
||||
Stock(symbol="sh600519", name="x", list_date=date(2001, 1, 1))
|
||||
|
||||
|
||||
class TestDailyBar:
|
||||
def test_is_complete(self) -> None:
|
||||
bar = DailyBar(
|
||||
symbol="600519.SH",
|
||||
trade_date=date(2024, 1, 2),
|
||||
open=Decimal("100"),
|
||||
high=Decimal("101"),
|
||||
low=Decimal("99"),
|
||||
close=Decimal("100.5"),
|
||||
volume=Decimal("10000"),
|
||||
amount=Decimal("1000000"),
|
||||
)
|
||||
assert bar.is_complete
|
||||
assert not DailyBar(symbol="600519.SH", trade_date=date(2024, 1, 2)).is_complete
|
||||
|
||||
|
||||
class TestFinancialIndicator:
|
||||
def _fin(self, announce: date) -> FinancialIndicator:
|
||||
return FinancialIndicator(
|
||||
symbol="600519.SH",
|
||||
report_date=date(2024, 6, 30),
|
||||
announce_date=announce,
|
||||
eps=Decimal("1.2"),
|
||||
)
|
||||
|
||||
def test_announced_by_after_announce(self) -> None:
|
||||
fin = self._fin(date(2024, 8, 31))
|
||||
# 公告日当天已可见;公告前不可见
|
||||
assert fin.announced_by(date(2024, 8, 31))
|
||||
assert not fin.announced_by(date(2024, 8, 30))
|
||||
assert not fin.announced_by(date(2024, 6, 30)) # 报告期不代表公开
|
||||
|
||||
|
||||
class TestSyncLog:
|
||||
def test_defaults(self) -> None:
|
||||
log = SyncLog(source="tushare", api="daily", success=True, row_count=5)
|
||||
assert log.request_time is not None
|
||||
assert log.data_start is None
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Failover 审计测试:主源失败 → 备用源兜底,每次尝试留 SyncLog。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from app.domain.entities.market import SyncLog
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceError,
|
||||
DataSourceNotSupported,
|
||||
)
|
||||
from app.infrastructure.data_sources.failover import FailoverProvider
|
||||
|
||||
|
||||
class PrimaryProvider:
|
||||
name = "tushare"
|
||||
|
||||
def __init__(self, *, fail: bool = False, fail_message: str = "boom") -> None:
|
||||
self.fail = fail
|
||||
self.fail_message = fail_message
|
||||
|
||||
def get_daily(self, symbol, start, end):
|
||||
if self.fail:
|
||||
raise DataSourceError(self.fail_message)
|
||||
return ["bar-ok"]
|
||||
|
||||
|
||||
class FallbackProvider:
|
||||
name = "sina"
|
||||
|
||||
def __init__(self, *, fail: bool = False, unsupported: bool = False) -> None:
|
||||
self.fail = fail
|
||||
self.unsupported = unsupported
|
||||
|
||||
def get_daily(self, symbol, start, end):
|
||||
if self.unsupported:
|
||||
raise DataSourceNotSupported("新浪不支持日线兜底")
|
||||
if self.fail:
|
||||
raise DataSourceError("sina down")
|
||||
return ["bar-sina"]
|
||||
|
||||
|
||||
def _logs() -> list[SyncLog]:
|
||||
collected: list[SyncLog] = []
|
||||
return collected
|
||||
|
||||
|
||||
def _make(primary, fallback, collected) -> FailoverProvider:
|
||||
return FailoverProvider(primary, fallback, audit=collected.append)
|
||||
|
||||
|
||||
class TestFailover:
|
||||
def test_primary_success_no_fallback(self) -> None:
|
||||
collected: list[SyncLog] = []
|
||||
f = _make(PrimaryProvider(), FallbackProvider(), collected)
|
||||
assert f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5)) == ["bar-ok"]
|
||||
assert len(collected) == 1
|
||||
assert collected[0].source == "tushare"
|
||||
assert collected[0].success is True
|
||||
assert collected[0].row_count == 1
|
||||
assert collected[0].data_start == date(2024, 1, 1)
|
||||
|
||||
def test_primary_fails_fallback_succeeds(self) -> None:
|
||||
collected: list[SyncLog] = []
|
||||
f = _make(PrimaryProvider(fail=True), FallbackProvider(), collected)
|
||||
assert f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5)) == ["bar-sina"]
|
||||
assert [log.source for log in collected] == ["tushare", "sina"]
|
||||
assert collected[0].success is False
|
||||
assert "boom" in (collected[0].failure_reason or "")
|
||||
assert collected[1].success is True
|
||||
|
||||
def test_fallback_unsupported_raises_with_audit(self) -> None:
|
||||
collected: list[SyncLog] = []
|
||||
f = _make(PrimaryProvider(fail=True), FallbackProvider(unsupported=True), collected)
|
||||
with pytest.raises(DataSourceError, match="备用源不支持"):
|
||||
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
|
||||
assert len(collected) == 2
|
||||
assert collected[1].success is False
|
||||
|
||||
def test_both_fail_raises_with_audit(self) -> None:
|
||||
collected: list[SyncLog] = []
|
||||
f = _make(PrimaryProvider(fail=True), FallbackProvider(fail=True), collected)
|
||||
with pytest.raises(DataSourceError, match="主备数据源均失败"):
|
||||
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
|
||||
assert len(collected) == 2
|
||||
|
||||
def test_no_fallback_raises(self) -> None:
|
||||
collected: list[SyncLog] = []
|
||||
f = FailoverProvider(PrimaryProvider(fail=True), None, audit=collected.append)
|
||||
with pytest.raises(DataSourceError, match="无备用源"):
|
||||
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
|
||||
assert len(collected) == 1
|
||||
@@ -0,0 +1,62 @@
|
||||
"""Alembic 迁移测试:全新数据库 upgrade head 后应包含全部 Phase 1 表。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
BACKEND_ROOT = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
def _alembic_config(db_path: Path) -> Config:
|
||||
cfg = Config(str(BACKEND_ROOT / "alembic.ini"))
|
||||
cfg.set_main_option("script_location", "app/infrastructure/persistence/migrations")
|
||||
cfg.set_main_option("sqlalchemy.url", f"sqlite:///{db_path}")
|
||||
return cfg
|
||||
|
||||
|
||||
def test_upgrade_head_creates_phase1_tables(tmp_path) -> None:
|
||||
db_path = tmp_path / "fresh.db"
|
||||
command.upgrade(_alembic_config(db_path), "head")
|
||||
|
||||
con = sqlite3.connect(db_path)
|
||||
try:
|
||||
tables = {
|
||||
row[0]
|
||||
for row in con.execute("select name from sqlite_master where type='table'").fetchall()
|
||||
}
|
||||
finally:
|
||||
con.close()
|
||||
|
||||
expected = {
|
||||
"stock",
|
||||
"stock_daily",
|
||||
"adjust_factor",
|
||||
"trading_calendar",
|
||||
"financial_indicator",
|
||||
"sync_log",
|
||||
"alembic_version",
|
||||
}
|
||||
assert expected <= tables
|
||||
|
||||
# 关键防未来函数列存在
|
||||
con = sqlite3.connect(db_path)
|
||||
try:
|
||||
fin_cols = {
|
||||
row[1] for row in con.execute("pragma table_info(financial_indicator)").fetchall()
|
||||
}
|
||||
daily_cols = {row[1] for row in con.execute("pragma table_info(stock_daily)").fetchall()}
|
||||
finally:
|
||||
con.close()
|
||||
assert {"report_date", "announce_date"} <= fin_cols
|
||||
assert {"symbol", "trade_date", "close"} <= daily_cols
|
||||
|
||||
|
||||
def test_upgrade_head_idempotent(tmp_path) -> None:
|
||||
db_path = tmp_path / "again.db"
|
||||
cfg = _alembic_config(db_path)
|
||||
command.upgrade(cfg, "head")
|
||||
command.upgrade(cfg, "head") # 二次执行不报错
|
||||
@@ -0,0 +1,191 @@
|
||||
"""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
|
||||
|
||||
|
||||
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
|
||||
@@ -0,0 +1,77 @@
|
||||
"""新浪 Provider 测试:代码转换、JSONP 解析、能力边界(不触网)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceError,
|
||||
DataSourceNotSupported,
|
||||
)
|
||||
from app.infrastructure.data_sources.sina import SinaProvider, _extract_jsonp, _to_sina_symbol
|
||||
|
||||
_KLINE_OK = (
|
||||
'var data=[{"day":"2024-08-30","open":"1700.0","high":"1720.0","low":"1690.0",'
|
||||
'"close":"1710.0","volume":"20000"},{"day":"2024-08-31","open":"1710.0",'
|
||||
'"high":"1725.0","low":"1705.0","close":"1720.0","volume":"18000"}]'
|
||||
)
|
||||
|
||||
|
||||
class _FakeResp:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return _KLINE_OK.encode("utf-8")
|
||||
|
||||
|
||||
class _BoomResp(_FakeResp):
|
||||
def read(self):
|
||||
raise OSError("socket timeout")
|
||||
|
||||
|
||||
class TestSymbolMap:
|
||||
def test_mapping(self) -> None:
|
||||
assert _to_sina_symbol("600519.SH") == "sh600519"
|
||||
assert _to_sina_symbol("000001.SZ") == "sz000001"
|
||||
assert _to_sina_symbol("830001.BJ") == "bj830001"
|
||||
|
||||
|
||||
class TestJsonp:
|
||||
def test_extract(self) -> None:
|
||||
rows = _extract_jsonp(_KLINE_OK)
|
||||
assert len(rows) == 2
|
||||
assert rows[0]["close"] == "1710.0"
|
||||
|
||||
def test_bad_payload_raises(self) -> None:
|
||||
with pytest.raises(DataSourceError, match="无法解析"):
|
||||
_extract_jsonp("not jsonp")
|
||||
|
||||
|
||||
class TestGetDaily:
|
||||
def test_ok_with_date_filter(self) -> None:
|
||||
provider = SinaProvider(urlopen=lambda _url, **kw: _FakeResp())
|
||||
bars = provider.get_daily("600519.SH", date(2024, 8, 30), date(2024, 8, 30))
|
||||
assert len(bars) == 1
|
||||
assert bars[0].close is not None
|
||||
assert bars[0].symbol == "600519.SH"
|
||||
|
||||
def test_network_error_wrapped(self) -> None:
|
||||
provider = SinaProvider(urlopen=lambda _url, **kw: _BoomResp())
|
||||
with pytest.raises(DataSourceError, match="sina 请求失败"):
|
||||
provider.get_daily("600519.SH", date(2024, 8, 1), date(2024, 8, 31))
|
||||
|
||||
|
||||
class TestCapabilities:
|
||||
def test_not_supported(self) -> None:
|
||||
provider = SinaProvider()
|
||||
with pytest.raises(DataSourceNotSupported):
|
||||
provider.get_stock_basic()
|
||||
with pytest.raises(DataSourceNotSupported):
|
||||
provider.get_adjust_factor("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
|
||||
with pytest.raises(DataSourceNotSupported):
|
||||
provider.get_financial("600519.SH")
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Tushare Provider 测试:归一化、重试与鉴权错误归类(用 Fake pro,不触网)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceAuthenticationError,
|
||||
DataSourceError,
|
||||
)
|
||||
from app.infrastructure.data_sources.tushare import TushareProvider
|
||||
|
||||
|
||||
class FakePro:
|
||||
"""模拟 tushare.pro 客户端:方法返回 records(list[dict]) 或抛错。"""
|
||||
|
||||
def __init__(self, *, payload=None, error: Exception | None = None) -> None:
|
||||
self.payload = payload or []
|
||||
self.error = error
|
||||
self.calls: list[str] = []
|
||||
|
||||
def __getattr__(self, api: str):
|
||||
def _run(**kwargs):
|
||||
self.calls.append(api)
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.payload
|
||||
|
||||
return _run
|
||||
|
||||
|
||||
def _pro(payload=None, error=None, retries: int = 2) -> TushareProvider:
|
||||
return TushareProvider(
|
||||
token="fake-token", pro=FakePro(payload=payload, error=error), max_retries=retries
|
||||
)
|
||||
|
||||
|
||||
class TestNormalize:
|
||||
def test_stock_records(self) -> None:
|
||||
stocks = TushareProvider.normalize_stock(
|
||||
[
|
||||
{
|
||||
"ts_code": "600519.SH",
|
||||
"symbol": "600519",
|
||||
"name": "贵州茅台",
|
||||
"area": "贵州",
|
||||
"industry": "白酒",
|
||||
"list_date": "20010827",
|
||||
}
|
||||
]
|
||||
)
|
||||
assert stocks[0].symbol == "600519.SH"
|
||||
assert stocks[0].list_date == date(2001, 8, 27)
|
||||
assert stocks[0].delist_date is None
|
||||
|
||||
def test_daily_volume_amount_scaled(self) -> None:
|
||||
bars = TushareProvider.normalize_daily(
|
||||
[
|
||||
{
|
||||
"ts_code": "600519.SH",
|
||||
"trade_date": "20240102",
|
||||
"open": "100.0",
|
||||
"high": "101.5",
|
||||
"low": "99.0",
|
||||
"close": "100.5",
|
||||
"vol": "10000.0",
|
||||
"amount": "1010000.0",
|
||||
}
|
||||
]
|
||||
)
|
||||
bar = bars[0]
|
||||
assert bar.trade_date == date(2024, 1, 2)
|
||||
assert bar.volume == Decimal("1000000") # 手 → 股(×100)
|
||||
assert bar.amount == Decimal("1010000000") # 千元 → 元(×1000)
|
||||
|
||||
def test_daily_nan_dropped(self) -> None:
|
||||
bars = TushareProvider.normalize_daily(
|
||||
[
|
||||
{
|
||||
"ts_code": "600519.SH",
|
||||
"trade_date": "20240102",
|
||||
"open": None,
|
||||
"high": float("nan"),
|
||||
"close": "10.0",
|
||||
"vol": None,
|
||||
"amount": None,
|
||||
}
|
||||
]
|
||||
)
|
||||
bar = bars[0]
|
||||
assert bar.open is None
|
||||
assert bar.high is None
|
||||
assert bar.close == Decimal("10.0")
|
||||
|
||||
def test_financial_maps_announce_date(self) -> None:
|
||||
rows = TushareProvider.normalize_financial(
|
||||
[
|
||||
{
|
||||
"ts_code": "600519.SH",
|
||||
"end_date": "20240630",
|
||||
"ann_date": "20240831",
|
||||
"eps": "1.23",
|
||||
"roe": "15.5",
|
||||
}
|
||||
]
|
||||
)
|
||||
fin = rows[0]
|
||||
assert fin.report_date == date(2024, 6, 30)
|
||||
assert fin.announce_date == date(2024, 8, 31)
|
||||
assert fin.eps == Decimal("1.23")
|
||||
|
||||
|
||||
class TestCall:
|
||||
def test_empty_result_returns_empty_list(self) -> None:
|
||||
provider = _pro(payload=[])
|
||||
assert provider.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 31)) == []
|
||||
|
||||
def test_success_records_returned(self) -> None:
|
||||
payload = [
|
||||
{
|
||||
"ts_code": "600519.SH",
|
||||
"trade_date": "20240102",
|
||||
"open": "100",
|
||||
"close": "101",
|
||||
"vol": "1",
|
||||
"amount": "1",
|
||||
}
|
||||
]
|
||||
provider = _pro(payload=payload)
|
||||
bars = provider.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
|
||||
assert len(bars) == 1
|
||||
assert bars[0].symbol == "600519.SH"
|
||||
|
||||
def test_transient_error_retries_then_raises(self) -> None:
|
||||
provider = _pro(error=RuntimeError("network down"), retries=2)
|
||||
with pytest.raises(DataSourceError, match="重试 2 次仍失败"):
|
||||
provider.get_stock_basic()
|
||||
assert len(provider._pro.calls) == 2 # noqa: SLF001 —— 测试探针
|
||||
|
||||
def test_permission_error_raises_immediately(self) -> None:
|
||||
provider = _pro(error=RuntimeError("抱歉,您没有访问该接口的权限,请升级积分"), retries=3)
|
||||
with pytest.raises(DataSourceAuthenticationError):
|
||||
provider.get_stock_basic()
|
||||
assert len(provider._pro.calls) == 1 # noqa: SLF001
|
||||
|
||||
def test_missing_token_rejected(self) -> None:
|
||||
with pytest.raises(DataSourceAuthenticationError, match="TUSHARE_TOKEN"):
|
||||
TushareProvider(token="")
|
||||
Reference in New Issue
Block a user