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 @@
|
||||
"""命令行工具(数据同步等)。用法:uv run python -m app.cli.sync ..."""
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Phase 1 数据同步 CLI(Tushare 首选 → SQLite)。
|
||||
|
||||
用法(cd backend):
|
||||
uv run python -m app.cli.sync basic
|
||||
uv run python -m app.cli.sync calendar --start 20240101 --end 20241231
|
||||
uv run python -m app.cli.sync daily --symbols 600519.SH,000001.SZ --start 20240101
|
||||
uv run python -m app.cli.sync daily --all --start 20240101 # 全市场
|
||||
uv run python -m app.cli.sync financial --all
|
||||
uv run python -m app.cli.sync verify --symbol 600519.SH # 新浪交叉验证
|
||||
|
||||
本模块是组装层(composition root):在此装配 Provider / Repository / Session,
|
||||
业务层代码仍只依赖抽象(domain.repositories / domain.providers)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from datetime import date, datetime, timedelta
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.infrastructure.data_sources.errors import DataSourceError
|
||||
from app.infrastructure.data_sources.failover import FailoverProvider
|
||||
from app.infrastructure.data_sources.sina import SinaProvider
|
||||
from app.infrastructure.data_sources.tushare import TushareProvider
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import StockModel
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyAdjustFactorRepository,
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyFinancialRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
SqlAlchemySyncLogRepository,
|
||||
SqlAlchemyTradingCalendarRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
_DATE_FMT = "%Y%m%d"
|
||||
|
||||
|
||||
def _parse_day(text: str) -> date:
|
||||
return datetime.strptime(text, _DATE_FMT).date()
|
||||
|
||||
|
||||
def _failover_provider(session):
|
||||
"""Tushare 单源经 FailoverProvider 包装:每次尝试写 sync_log(AGENT.md §7 审计)。"""
|
||||
audit_repo = SqlAlchemySyncLogRepository(session)
|
||||
primary = TushareProvider(token=get_settings().tushare_token)
|
||||
return FailoverProvider(primary, fallback=None, audit=audit_repo.add)
|
||||
|
||||
|
||||
def _session_ctx():
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
def cmd_basic(args) -> int:
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
stocks = provider.get_stock_basic()
|
||||
repo = SqlAlchemyStockRepository(session)
|
||||
touched = repo.upsert_many(stocks)
|
||||
session.commit()
|
||||
print(f"[basic] 拉取 {len(stocks)} 只,落库 {touched} 条")
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_calendar(args) -> int:
|
||||
start = _parse_day(args.start)
|
||||
end = _parse_day(args.end)
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
days = provider.get_trade_cal(start, end)
|
||||
repo = SqlAlchemyTradingCalendarRepository(session)
|
||||
touched = repo.upsert_many(days)
|
||||
session.commit()
|
||||
open_days = sum(1 for d in days if d.is_open)
|
||||
print(f"[calendar] {start}~{end} 共 {len(days)} 条(交易日 {open_days}),落库 {touched} 条")
|
||||
return 0
|
||||
|
||||
|
||||
def _symbols_of(args) -> list[str]:
|
||||
if getattr(args, "all", False):
|
||||
with _session_ctx() as session:
|
||||
symbols = list(session.scalars(select(StockModel.symbol).order_by(StockModel.symbol)))
|
||||
if not symbols:
|
||||
print("[error] stock 表为空,请先运行:python -m app.cli.sync basic")
|
||||
sys.exit(2)
|
||||
return symbols
|
||||
return [s.strip() for s in args.symbols.split(",") if s.strip()]
|
||||
|
||||
|
||||
def cmd_daily(args) -> int:
|
||||
symbols = _symbols_of(args)
|
||||
start = _parse_day(args.start) if args.start else date(2005, 1, 1)
|
||||
end = _parse_day(args.end) if args.end else date.today()
|
||||
total = 0
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
bar_repo = SqlAlchemyDailyBarRepository(session)
|
||||
factor_repo = SqlAlchemyAdjustFactorRepository(session)
|
||||
for i, symbol in enumerate(symbols, start=1):
|
||||
begin = start
|
||||
if args.resume:
|
||||
latest = bar_repo.latest_date(symbol)
|
||||
if latest is not None:
|
||||
begin = max(begin, latest + timedelta(days=1))
|
||||
if begin > end:
|
||||
continue
|
||||
try:
|
||||
bars = provider.get_daily(symbol, begin, end)
|
||||
factors = provider.get_adjust_factor(symbol, begin, end)
|
||||
bar_repo.upsert_many(bars)
|
||||
factor_repo.upsert_many(factors)
|
||||
total += len(bars)
|
||||
if i % 100 == 0:
|
||||
session.commit()
|
||||
print(f" ... {i}/{len(symbols)} {symbol} 累计 {total} 根")
|
||||
except DataSourceError as exc:
|
||||
print(f" [warn] {symbol} 拉取失败: {exc}", file=sys.stderr)
|
||||
session.commit()
|
||||
print(f"[daily] {len(symbols)} 只股票合计写入 {total} 根日线(含复权因子)")
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_financial(args) -> int:
|
||||
symbols = _symbols_of(args)
|
||||
total = 0
|
||||
with _session_ctx() as session:
|
||||
provider = _failover_provider(session)
|
||||
fin_repo = SqlAlchemyFinancialRepository(session)
|
||||
for symbol in symbols:
|
||||
rows = provider.get_financial(symbol)
|
||||
fin_repo.upsert_many(rows)
|
||||
total += len(rows)
|
||||
session.commit()
|
||||
print(f"[financial] {len(symbols)} 只股票合计写入 {total} 条财务指标快照")
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_verify(args) -> int:
|
||||
"""新浪交叉验证:取新浪最新前复权收盘,与本地最新交易日对照。
|
||||
|
||||
注意:新浪为前复权口径,数值不直接等于本地不复权收盘,
|
||||
本命令仅用于确认新浪可用性 / 最新交易日,不把新浪数据并入主库。
|
||||
"""
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
)
|
||||
|
||||
sina = SinaProvider()
|
||||
end = date.today()
|
||||
start = end - timedelta(days=20)
|
||||
try:
|
||||
bars = sina.get_daily(args.symbol, start, end)
|
||||
except DataSourceError as exc:
|
||||
print(f"[verify] 新浪不可用: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
if not bars:
|
||||
print(f"[verify] 新浪最近无数据({args.symbol})")
|
||||
return 1
|
||||
latest = max(bars, key=lambda b: b.trade_date)
|
||||
with _session_ctx() as session:
|
||||
local = SqlAlchemyDailyBarRepository(session).latest_date(args.symbol)
|
||||
print(
|
||||
f"[verify] {args.symbol}: 新浪最新 {latest.trade_date} 收盘(前复权) {latest.close};"
|
||||
f"本地最新交易日 {local}"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(prog="app.cli.sync", description="Tushare 数据同步 CLI")
|
||||
sub = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
p_basic = sub.add_parser("basic", help="同步股票基础信息")
|
||||
p_basic.set_defaults(func=cmd_basic)
|
||||
|
||||
p_cal = sub.add_parser("calendar", help="同步交易日历")
|
||||
p_cal.add_argument("--start", required=True, help="YYYYMMDD")
|
||||
p_cal.add_argument("--end", required=True, help="YYYYMMDD")
|
||||
p_cal.set_defaults(func=cmd_calendar)
|
||||
|
||||
p_daily = sub.add_parser("daily", help="同步日线与复权因子")
|
||||
p_daily.add_argument("--symbols", default="", help="600519.SH,000001.SZ")
|
||||
p_daily.add_argument("--all", action="store_true", help="遍历 stock 表全部股票")
|
||||
p_daily.add_argument("--start", default="", help="YYYYMMDD(默认 20050101)")
|
||||
p_daily.add_argument("--end", default="", help="YYYYMMDD(默认今天)")
|
||||
p_daily.add_argument("--resume", action="store_true", help="从本地最新交易日续传")
|
||||
p_daily.set_defaults(func=cmd_daily)
|
||||
|
||||
p_fin = sub.add_parser("financial", help="同步财务指标快照")
|
||||
p_fin.add_argument("--symbols", default="")
|
||||
p_fin.add_argument("--all", action="store_true")
|
||||
p_fin.set_defaults(func=cmd_financial)
|
||||
|
||||
p_verify = sub.add_parser("verify", help="新浪交叉验证最新行情")
|
||||
p_verify.add_argument("--symbol", required=True)
|
||||
p_verify.set_defaults(func=cmd_verify)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
try:
|
||||
return args.func(args)
|
||||
except DataSourceError as exc:
|
||||
print(f"[error] {exc}", file=sys.stderr)
|
||||
return 1
|
||||
except KeyboardInterrupt:
|
||||
print("\n[interrupt] 已中止", file=sys.stderr)
|
||||
return 130
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,104 @@
|
||||
"""市场数据领域实体(Phase 1)。
|
||||
|
||||
约定(AGENT.md §8/§9):
|
||||
- 行情时间用 trade_date;财务数据同时区分 report_date(报告期)与 announce_date(公告日)
|
||||
- 禁止以 report_date 作可见性依据 —— 只允许 announce_date 已过的数据进入研究
|
||||
- 复权一律通过独立 AdjustFactor 表达,不在此层偷偷改前/后复权口径
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# 常见精度:价格 4 位小数;成交量(股) 2 位;金额(元) 2 位
|
||||
PRICE_PLACES = Decimal("0.0001")
|
||||
AMOUNT_PLACES = Decimal("0.01")
|
||||
|
||||
|
||||
class Stock(BaseModel):
|
||||
"""A 股基础信息。symbol 统一为 Tushare 风格,如 600519.SH。"""
|
||||
|
||||
model_config = ConfigDict(str_strip_whitespace=True)
|
||||
|
||||
symbol: str = Field(pattern=r"^\d{6}\.(SH|SZ|BJ)$", description="如 600519.SH")
|
||||
name: str
|
||||
industry: str | None = None
|
||||
area: str | None = None
|
||||
market: str | None = Field(default=None, description="主板/创业板/科创板/北交所")
|
||||
exchange: str | None = None
|
||||
list_date: date
|
||||
delist_date: date | None = None
|
||||
status: str = Field(default="L", description="L 上市 / D 退市 / P 暂停")
|
||||
|
||||
|
||||
class TradingCalendar(BaseModel):
|
||||
"""交易日历。"""
|
||||
|
||||
calendar_date: date
|
||||
is_open: bool = True
|
||||
|
||||
|
||||
class DailyBar(BaseModel):
|
||||
"""不复权日线。复权请使用 AdjustFactor 在消费侧显式计算。"""
|
||||
|
||||
symbol: str
|
||||
trade_date: date
|
||||
open: Decimal | None = None
|
||||
high: Decimal | None = None
|
||||
low: Decimal | None = None
|
||||
close: Decimal | None = None
|
||||
volume: Decimal | None = Field(default=None, description="成交量(股)")
|
||||
amount: Decimal | None = Field(default=None, description="成交额(元)")
|
||||
|
||||
@property
|
||||
def is_complete(self) -> bool:
|
||||
"""基础行情字段是否齐全(供校验器使用)。"""
|
||||
return all(
|
||||
v is not None
|
||||
for v in (self.open, self.high, self.low, self.close, self.volume, self.amount)
|
||||
)
|
||||
|
||||
|
||||
class AdjustFactor(BaseModel):
|
||||
"""复权因子。因子原始口径由数据源决定,必须与数据源文档一致地存取。"""
|
||||
|
||||
symbol: str
|
||||
trade_date: date
|
||||
factor: Decimal
|
||||
|
||||
|
||||
class FinancialIndicator(BaseModel):
|
||||
"""核心财务指标(快照)。
|
||||
|
||||
可见性红线:研究侧查询一律按 announce_date <= as_of_date 过滤,
|
||||
report_date 只表示报告所属期间,不代表公开时间。
|
||||
"""
|
||||
|
||||
symbol: str
|
||||
report_date: date
|
||||
announce_date: date
|
||||
eps: Decimal | None = None
|
||||
roe: Decimal | None = None
|
||||
total_revenue: Decimal | None = None
|
||||
net_profit: Decimal | None = None
|
||||
gross_margin: Decimal | None = None
|
||||
|
||||
def announced_by(self, as_of_date: date) -> bool:
|
||||
"""as_of_date(含当日)是否已可见。防未来函数的核心判断。"""
|
||||
return self.announce_date <= as_of_date
|
||||
|
||||
|
||||
class SyncLog(BaseModel):
|
||||
"""数据拉取审计记录(AGENT.md §7:来源必须可追踪,禁止静默切换)。"""
|
||||
|
||||
source: str
|
||||
api: str
|
||||
request_time: datetime = Field(default_factory=datetime.utcnow)
|
||||
success: bool
|
||||
failure_reason: str | None = None
|
||||
row_count: int = 0
|
||||
data_start: date | None = None
|
||||
data_end: date | None = None
|
||||
@@ -0,0 +1,40 @@
|
||||
"""MarketDataProvider(数据源抽象)。
|
||||
|
||||
业务层只依赖本 Protocol(AGENT.md §6),禁止在业务代码中 import
|
||||
tushare / 新浪实现。数据源一律返回 domain.entities 中的归一化实体。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
TradingCalendar,
|
||||
)
|
||||
|
||||
|
||||
class MarketDataProvider(Protocol):
|
||||
"""统一市场数据源接口。
|
||||
|
||||
实现约定:
|
||||
- get_daily 返回**不复权**行情(复权经 AdjustFactor 显式计算,禁止静默改口径)
|
||||
- get_financial 返回带 announce_date 的指标,供上层按 as_of 过滤
|
||||
- 实现不得抛出裸连接异常以外的噪音;业务错误应转为 DataSourceError
|
||||
"""
|
||||
|
||||
name: str
|
||||
|
||||
def get_stock_basic(self) -> list[Stock]: ...
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list[TradingCalendar]: ...
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list[AdjustFactor]: ...
|
||||
|
||||
def get_financial(self, symbol: str) -> list[FinancialIndicator]: ...
|
||||
@@ -0,0 +1,70 @@
|
||||
"""Repository Protocol(Phase 1 数据层)。
|
||||
|
||||
业务层只依赖这些 Protocol;具体实现位于 infrastructure/persistence。
|
||||
实体一律以 domain.entities 类型进出,禁止把 ORM Model 泄漏到上层。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
SyncLog,
|
||||
TradingCalendar,
|
||||
)
|
||||
|
||||
|
||||
class StockRepository(Protocol):
|
||||
def get_by_symbol(self, symbol: str) -> Stock | None: ...
|
||||
|
||||
def list(self) -> list[Stock]: ...
|
||||
|
||||
def upsert_many(self, stocks: Sequence[Stock]) -> int:
|
||||
"""批量写入,以 symbol 为幂等键,返回写入/更新的行数。"""
|
||||
|
||||
|
||||
class TradingCalendarRepository(Protocol):
|
||||
def upsert_many(self, days: Sequence[TradingCalendar]) -> int: ...
|
||||
|
||||
def list_range(self, start: date, end: date) -> list[TradingCalendar]: ...
|
||||
|
||||
def is_open(self, day: date) -> bool: ...
|
||||
|
||||
|
||||
class DailyBarRepository(Protocol):
|
||||
def upsert_many(self, bars: Sequence[DailyBar]) -> int: ...
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def latest_date(self, symbol: str) -> date | None:
|
||||
"""断点续传用:该股票本地已有数据的最新交易日。"""
|
||||
|
||||
|
||||
class AdjustFactorRepository(Protocol):
|
||||
def upsert_many(self, factors: Sequence[AdjustFactor]) -> int: ...
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[AdjustFactor]: ...
|
||||
|
||||
|
||||
class FinancialRepository(Protocol):
|
||||
def upsert_many(self, rows: Sequence[FinancialIndicator]) -> int: ...
|
||||
|
||||
def list_announced(
|
||||
self,
|
||||
symbol: str,
|
||||
as_of_date: date,
|
||||
report_start: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。"""
|
||||
|
||||
|
||||
class SyncLogRepository(Protocol):
|
||||
def add(self, log: SyncLog) -> SyncLog: ...
|
||||
|
||||
def recent(self, source: str | None = None, limit: int = 20) -> list[SyncLog]: ...
|
||||
@@ -0,0 +1,7 @@
|
||||
"""数据源基础设施:MarketDataProvider 的具体实现与 Failover。
|
||||
|
||||
业务层不直接 import 本目录(AGENT.md §5/§6)——统一经
|
||||
MarketDataProvider(domain.providers)注入;唯一例外是组装处的依赖装配。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -0,0 +1,13 @@
|
||||
"""数据源异常类型。"""
|
||||
|
||||
|
||||
class DataSourceError(Exception):
|
||||
"""数据源通用失败(网络、限频、解析等)。"""
|
||||
|
||||
|
||||
class DataSourceAuthenticationError(DataSourceError):
|
||||
"""凭证无效 / 权限不足(如 Tushare token 无该接口权限)。"""
|
||||
|
||||
|
||||
class DataSourceNotSupported(DataSourceError):
|
||||
"""该数据源不提供此能力(如新浪无复权因子),用于 Failover 判定。"""
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Tushare → Sina Failover 包装(AGENT.md §7)。
|
||||
|
||||
规则:
|
||||
- 优先 primary;primary 抛错时才尝试 fallback(避免对空结果做无谓兜底请求)
|
||||
- fallback 不支持该 API(DataSourceNotSupported)或自身失败 → 抛 DataSourceError
|
||||
- 每次尝试都写 SyncLog(source / 成功与否 / 行数 / 区间),禁止静默切换
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from datetime import date
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.market import SyncLog
|
||||
from app.domain.providers import MarketDataProvider
|
||||
from app.infrastructure.data_sources.errors import DataSourceError, DataSourceNotSupported
|
||||
|
||||
|
||||
class FailoverProvider:
|
||||
"""以 primary 为主、fallback 为辅的 MarketDataProvider 实现。"""
|
||||
|
||||
name = "failover"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
primary: MarketDataProvider,
|
||||
fallback: MarketDataProvider | None = None,
|
||||
*,
|
||||
audit: Callable[[SyncLog], None] | None = None,
|
||||
) -> None:
|
||||
self.primary = primary
|
||||
self.fallback = fallback
|
||||
self._audit = audit or (lambda _log: None)
|
||||
|
||||
# ---- 各 API 代理 ----
|
||||
|
||||
def get_stock_basic(self) -> list:
|
||||
return self._with_failover(
|
||||
"get_stock_basic",
|
||||
primary_call=lambda: self.primary.get_stock_basic(),
|
||||
fallback_call=lambda: self.fallback.get_stock_basic(),
|
||||
)
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_trade_cal",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_trade_cal(start, end),
|
||||
fallback_call=lambda: self.fallback.get_trade_cal(start, end),
|
||||
)
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_daily",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_daily(symbol, start, end),
|
||||
fallback_call=lambda: self.fallback.get_daily(symbol, start, end),
|
||||
)
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list:
|
||||
return self._with_failover(
|
||||
"get_adjust_factor",
|
||||
start=start,
|
||||
end=end,
|
||||
primary_call=lambda: self.primary.get_adjust_factor(symbol, start, end),
|
||||
fallback_call=lambda: self.fallback.get_adjust_factor(symbol, start, end),
|
||||
)
|
||||
|
||||
def get_financial(self, symbol: str) -> list:
|
||||
return self._with_failover(
|
||||
"get_financial",
|
||||
primary_call=lambda: self.primary.get_financial(symbol),
|
||||
fallback_call=lambda: self.fallback.get_financial(symbol),
|
||||
)
|
||||
|
||||
# ---- 内部 ----
|
||||
|
||||
def _with_failover(
|
||||
self,
|
||||
api: str,
|
||||
*,
|
||||
primary_call: Callable[[], list],
|
||||
fallback_call: Callable[[], list] | None = None,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> list:
|
||||
try:
|
||||
rows = primary_call()
|
||||
except Exception as exc: # noqa: BLE001 —— 统一走审计
|
||||
self._log(
|
||||
source=self.primary.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return self._try_fallback(api, fallback_call, start=start, end=end, primary_error=exc)
|
||||
self._log(
|
||||
source=self.primary.name,
|
||||
api=api,
|
||||
success=True,
|
||||
row_count=_len(rows),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return rows
|
||||
|
||||
def _try_fallback(self, api, fallback_call, *, start, end, primary_error):
|
||||
if fallback_call is None or self.fallback is None:
|
||||
raise DataSourceError(
|
||||
f"{self.primary.name}.{api} 失败且无备用源: {primary_error}"
|
||||
) from primary_error
|
||||
try:
|
||||
rows = fallback_call()
|
||||
except DataSourceNotSupported as exc:
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=f"不支持: {exc}",
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
raise DataSourceError(
|
||||
f"{self.primary.name}.{api} 失败,备用源不支持: {primary_error}"
|
||||
) from primary_error
|
||||
except Exception as exc: # noqa: BLE001
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=False,
|
||||
reason=str(exc),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
raise DataSourceError(
|
||||
f"主备数据源均失败: primary[{self.primary.name}]={primary_error} "
|
||||
f"fallback[{self.fallback.name}]={exc}"
|
||||
) from exc
|
||||
self._log(
|
||||
source=self.fallback.name,
|
||||
api=api,
|
||||
success=True,
|
||||
row_count=_len(rows),
|
||||
start=start,
|
||||
end=end,
|
||||
)
|
||||
return rows
|
||||
|
||||
def _log(
|
||||
self,
|
||||
*,
|
||||
source: str,
|
||||
api: str,
|
||||
success: bool,
|
||||
reason: str | None = None,
|
||||
row_count: int = 0,
|
||||
start: date | None = None,
|
||||
end: date | None = None,
|
||||
) -> None:
|
||||
self._audit(
|
||||
SyncLog(
|
||||
source=source,
|
||||
api=api,
|
||||
success=success,
|
||||
failure_reason=reason,
|
||||
row_count=row_count,
|
||||
data_start=start,
|
||||
data_end=end,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _len(rows: Any) -> int:
|
||||
return len(rows) if rows is not None else 0
|
||||
@@ -0,0 +1,101 @@
|
||||
"""新浪财经 Provider —— 备用数据源。
|
||||
|
||||
能力边界(AGENT.md §5.2):
|
||||
- 新浪日 K 接口返回**前复权**数据,口径与 Tushare 不复权不同,
|
||||
因此本 Provider 只用于「缺失/不可用时的行情参考与交叉验证」,
|
||||
不得把结果直接并入不复权主时序库(禁止静默混口径)。
|
||||
- 新浪不提供复权因子 / 财务指标 → 相应方法抛 DataSourceNotSupported。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.market import DailyBar
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceError,
|
||||
DataSourceNotSupported,
|
||||
)
|
||||
|
||||
_KLINE_JSONP = (
|
||||
"https://quotes.sina.cn/cn/api/jsonp_v2.php/var%20data=/CN_MarketDataService"
|
||||
".getKLineData?symbol={sina_symbol}&scale=240&ma=no&datalen={datalen}"
|
||||
)
|
||||
|
||||
|
||||
def _to_sina_symbol(symbol: str) -> str:
|
||||
"""600519.SH -> sh600519;000001.SZ -> sz000001。"""
|
||||
code, _, suffix = symbol.partition(".")
|
||||
prefix = {"SH": "sh", "SZ": "sz", "BJ": "bj"}.get(suffix.upper(), "sh")
|
||||
return f"{prefix}{code}"
|
||||
|
||||
|
||||
def _extract_jsonp(payload: str) -> list[dict[str, Any]]:
|
||||
match = re.search(r"=\s*(\[.*\])\s*$", payload.strip(), flags=re.DOTALL)
|
||||
if not match:
|
||||
raise DataSourceError("新浪行情返回格式无法解析")
|
||||
return json.loads(match.group(1))
|
||||
|
||||
|
||||
class SinaProvider:
|
||||
"""新浪财经备用数据源(仅日线参考 / 交叉验证)。"""
|
||||
|
||||
name = "sina"
|
||||
|
||||
def __init__(self, *, timeout: float = 10.0, urlopen=urllib.request.urlopen) -> None:
|
||||
self._timeout = timeout
|
||||
self._urlopen = urlopen
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date, datalen: int = 320) -> list[DailyBar]:
|
||||
"""拉取前复权日 K(新浪仅支持最近 datalen 个自然日窗口)。"""
|
||||
url = _KLINE_JSONP.format(sina_symbol=_to_sina_symbol(symbol), datalen=datalen)
|
||||
try:
|
||||
with self._urlopen(url, timeout=self._timeout) as resp:
|
||||
payload = resp.read().decode("utf-8", errors="replace")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise DataSourceError(f"sina 请求失败: {exc}") from exc
|
||||
|
||||
bars: list[DailyBar] = []
|
||||
for rec in _extract_jsonp(payload):
|
||||
day = datetime.strptime(rec["day"], "%Y-%m-%d").date()
|
||||
if day < start or day > end:
|
||||
continue
|
||||
bars.append(
|
||||
DailyBar(
|
||||
symbol=symbol,
|
||||
trade_date=day,
|
||||
open=_d(rec.get("open")),
|
||||
high=_d(rec.get("high")),
|
||||
low=_d(rec.get("low")),
|
||||
close=_d(rec.get("close")),
|
||||
volume=_d(rec.get("volume")),
|
||||
)
|
||||
)
|
||||
return bars
|
||||
|
||||
def get_stock_basic(self):
|
||||
raise DataSourceNotSupported("新浪不提供股票基础信息列表")
|
||||
|
||||
def get_trade_cal(self, start, end):
|
||||
raise DataSourceNotSupported("新浪不提供交易日历")
|
||||
|
||||
def get_adjust_factor(self, symbol, start, end):
|
||||
raise DataSourceNotSupported("新浪不提供复权因子(返回数据为前复权口径)")
|
||||
|
||||
def get_financial(self, symbol):
|
||||
raise DataSourceNotSupported("新浪不提供财务指标")
|
||||
|
||||
|
||||
def _d(value) -> Decimal | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return Decimal(str(value))
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
@@ -0,0 +1,219 @@
|
||||
"""Tushare Provider —— 首选数据源实现。
|
||||
|
||||
依赖注入:pro 客户端(tushare.pro.client 或测试 Fake)。真实运行时惰性加载
|
||||
tushare 库(pyproject optional:uv sync --extra datasource-tushare)。
|
||||
归一化函数只依赖 list[dict],便于无 pandas 环境下单测。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
TradingCalendar,
|
||||
)
|
||||
from app.infrastructure.data_sources.errors import (
|
||||
DataSourceAuthenticationError,
|
||||
DataSourceError,
|
||||
)
|
||||
|
||||
_TS_DATE = "%Y%m%d"
|
||||
|
||||
|
||||
def _to_date(value: str | None) -> date | None:
|
||||
if value is None or value == "":
|
||||
return None
|
||||
return datetime.strptime(str(value)[:10], _TS_DATE).date()
|
||||
|
||||
|
||||
def _to_decimal(value) -> Decimal | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
num = float(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
if num != num: # NaN
|
||||
return None
|
||||
return Decimal(str(num))
|
||||
|
||||
|
||||
class TushareProvider:
|
||||
"""封装 Tushare Pro(ts.pro_api)。所有输出已归一化为领域实体。"""
|
||||
|
||||
name = "tushare"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token: str = "",
|
||||
*,
|
||||
pro: object | None = None,
|
||||
max_retries: int = 3,
|
||||
) -> None:
|
||||
self._pro = pro if pro is not None else _build_pro(token)
|
||||
self._max_retries = max_retries
|
||||
|
||||
# ---- 归一化(纯函数,输入 list[dict],可单测) ----
|
||||
|
||||
@staticmethod
|
||||
def normalize_stock(records: list[dict[str, Any]]) -> list[Stock]:
|
||||
stocks: list[Stock] = []
|
||||
for rec in records:
|
||||
stocks.append(
|
||||
Stock(
|
||||
symbol=str(rec.get("ts_code") or rec.get("symbol") or ""),
|
||||
name=str(rec.get("name") or ""),
|
||||
industry=rec.get("industry"),
|
||||
area=rec.get("area"),
|
||||
market=rec.get("market"),
|
||||
exchange=rec.get("exchange"),
|
||||
list_date=_to_date(rec.get("list_date")) or date.min,
|
||||
delist_date=_to_date(rec.get("delist_date")),
|
||||
status=str(rec.get("status") or "L"),
|
||||
)
|
||||
)
|
||||
return stocks
|
||||
|
||||
@staticmethod
|
||||
def normalize_calendar(records: list[dict[str, Any]]) -> list[TradingCalendar]:
|
||||
return [
|
||||
TradingCalendar(
|
||||
calendar_date=_to_date(rec.get("cal_date")) or date.min,
|
||||
is_open=bool(rec.get("is_open")),
|
||||
)
|
||||
for rec in records
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def normalize_daily(records: list[dict[str, Any]]) -> list[DailyBar]:
|
||||
bars: list[DailyBar] = []
|
||||
for rec in records:
|
||||
vol = _to_decimal(rec.get("vol"))
|
||||
amount = _to_decimal(rec.get("amount"))
|
||||
bars.append(
|
||||
DailyBar(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
open=_to_decimal(rec.get("open")),
|
||||
high=_to_decimal(rec.get("high")),
|
||||
low=_to_decimal(rec.get("low")),
|
||||
close=_to_decimal(rec.get("close")),
|
||||
volume=vol * 100 if vol is not None else None,
|
||||
amount=amount * 1000 if amount is not None else None,
|
||||
)
|
||||
)
|
||||
return bars
|
||||
|
||||
@staticmethod
|
||||
def normalize_adj_factor(records: list[dict[str, Any]]) -> list[AdjustFactor]:
|
||||
return [
|
||||
AdjustFactor(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
trade_date=_to_date(rec.get("trade_date")) or date.min,
|
||||
factor=_to_decimal(rec.get("adj_factor")) or Decimal(1),
|
||||
)
|
||||
for rec in records
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def normalize_financial(records: list[dict[str, Any]]) -> list[FinancialIndicator]:
|
||||
rows: list[FinancialIndicator] = []
|
||||
for rec in records:
|
||||
rows.append(
|
||||
FinancialIndicator(
|
||||
symbol=str(rec.get("ts_code") or ""),
|
||||
report_date=_to_date(rec.get("end_date")) or date.min,
|
||||
announce_date=_to_date(rec.get("ann_date")) or date.min,
|
||||
eps=_to_decimal(rec.get("eps")),
|
||||
roe=_to_decimal(rec.get("roe")),
|
||||
net_profit=_to_decimal(rec.get("n_income_attr_p")),
|
||||
gross_margin=_to_decimal(rec.get("grossprofit_margin")),
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
# ---- 接口调用 ----
|
||||
|
||||
def get_stock_basic(self) -> list[Stock]:
|
||||
records = self._call(
|
||||
"stock_basic",
|
||||
fields="ts_code,symbol,name,area,industry,market,exchange,list_date,delist_date,status",
|
||||
)
|
||||
return self.normalize_stock(records)
|
||||
|
||||
def get_trade_cal(self, start: date, end: date) -> list[TradingCalendar]:
|
||||
records = self._call(
|
||||
"trade_cal",
|
||||
exchange="SSE",
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
is_open="",
|
||||
)
|
||||
return self.normalize_calendar(records)
|
||||
|
||||
def get_daily(self, symbol: str, start: date, end: date) -> list[DailyBar]:
|
||||
records = self._call(
|
||||
"daily",
|
||||
ts_code=symbol,
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
)
|
||||
return self.normalize_daily(records)
|
||||
|
||||
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list[AdjustFactor]:
|
||||
records = self._call(
|
||||
"adj_factor",
|
||||
ts_code=symbol,
|
||||
start_date=start.strftime(_TS_DATE),
|
||||
end_date=end.strftime(_TS_DATE),
|
||||
)
|
||||
return self.normalize_adj_factor(records)
|
||||
|
||||
def get_financial(self, symbol: str) -> list[FinancialIndicator]:
|
||||
records = self._call("fina_indicator", ts_code=symbol)
|
||||
return self.normalize_financial(records)
|
||||
|
||||
# ---- 内部 ----
|
||||
|
||||
def _call(self, api: str, **kwargs) -> list[dict[str, Any]]:
|
||||
last_error: Exception | None = None
|
||||
for _ in range(self._max_retries):
|
||||
try:
|
||||
fn = getattr(self._pro, api)
|
||||
result = fn(**kwargs)
|
||||
if result is None:
|
||||
return []
|
||||
if hasattr(result, "to_dict"):
|
||||
return result.to_dict("records")
|
||||
if isinstance(result, list):
|
||||
return result
|
||||
return []
|
||||
except Exception as exc: # noqa: BLE001 —— tushare 异常无统一类型,逐一归类
|
||||
last_error = exc
|
||||
msg = str(exc)
|
||||
if "权限" in msg or "积分" in msg or "token" in msg.lower():
|
||||
raise DataSourceAuthenticationError(msg) from exc
|
||||
raise DataSourceError(
|
||||
f"tushare.{api} 重试 {self._max_retries} 次仍失败: {last_error}"
|
||||
) from last_error
|
||||
|
||||
|
||||
def _build_pro(token: str):
|
||||
if not token:
|
||||
raise DataSourceAuthenticationError(
|
||||
"缺少 TUSHARE_TOKEN:请 cp .env.example .env 并填入 Tushare Pro token"
|
||||
)
|
||||
try:
|
||||
ts = importlib.import_module("tushare")
|
||||
except ImportError as exc: # pragma: no cover —— 环境相关
|
||||
raise DataSourceError(
|
||||
"未安装 tushare 客户端:cd backend && uv sync --extra datasource-tushare"
|
||||
) from exc
|
||||
return ts.pro_api(token)
|
||||
@@ -11,6 +11,7 @@ from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from app.core.config import get_settings
|
||||
from app.infrastructure.persistence.sqlalchemy import models as _models # noqa: F401 —— 注册全部表
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from sqlalchemy import engine_from_config, pool
|
||||
|
||||
@@ -19,7 +20,11 @@ config = context.config
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
config.set_main_option("sqlalchemy.url", get_settings().database_url)
|
||||
# alembic.ini 中显式 sqlalchemy.url 优先(测试/运维可注入);否则用应用配置
|
||||
_db_url = config.get_main_option("sqlalchemy.url")
|
||||
if not _db_url:
|
||||
_db_url = get_settings().database_url
|
||||
config.set_main_option("sqlalchemy.url", _db_url)
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
+178
@@ -0,0 +1,178 @@
|
||||
"""phase1 market data tables
|
||||
|
||||
Revision ID: e4d188250fb2
|
||||
Revises:
|
||||
Create Date: 2026-09-06 16:58:13.904265
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "e4d188250fb2"
|
||||
down_revision: str | None = None
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table(
|
||||
"adjust_factor",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("factor", sa.Numeric(precision=20, scale=6), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "trade_date", name="uq_adj_symbol_date"),
|
||||
)
|
||||
with op.batch_alter_table("adjust_factor", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_adjust_factor_symbol"), ["symbol"], unique=False)
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_adjust_factor_trade_date"), ["trade_date"], unique=False
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"financial_indicator",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("report_date", sa.Date(), nullable=False),
|
||||
sa.Column("announce_date", sa.Date(), nullable=False),
|
||||
sa.Column("eps", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("roe", sa.Numeric(precision=10, scale=4), nullable=True),
|
||||
sa.Column("total_revenue", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("net_profit", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("gross_margin", sa.Numeric(precision=10, scale=4), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "report_date", "announce_date", name="uq_fin_sym_rep_ann"),
|
||||
)
|
||||
with op.batch_alter_table("financial_indicator", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_financial_indicator_announce_date"), ["announce_date"], unique=False
|
||||
)
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_financial_indicator_report_date"), ["report_date"], unique=False
|
||||
)
|
||||
batch_op.create_index(batch_op.f("ix_financial_indicator_symbol"), ["symbol"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"stock",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("industry", sa.String(length=64), nullable=True),
|
||||
sa.Column("area", sa.String(length=32), nullable=True),
|
||||
sa.Column("market", sa.String(length=16), nullable=True),
|
||||
sa.Column("exchange", sa.String(length=8), nullable=True),
|
||||
sa.Column("list_date", sa.Date(), nullable=False),
|
||||
sa.Column("delist_date", sa.Date(), nullable=True),
|
||||
sa.Column("status", sa.String(length=8), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("stock", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_stock_symbol"), ["symbol"], unique=True)
|
||||
|
||||
op.create_table(
|
||||
"stock_daily",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("symbol", sa.String(length=12), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("open", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("high", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("low", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("close", sa.Numeric(precision=12, scale=4), nullable=True),
|
||||
sa.Column("volume", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.Column("amount", sa.Numeric(precision=24, scale=2), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("symbol", "trade_date", name="uq_daily_symbol_date"),
|
||||
)
|
||||
with op.batch_alter_table("stock_daily", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_stock_daily_symbol"), ["symbol"], unique=False)
|
||||
batch_op.create_index(batch_op.f("ix_stock_daily_trade_date"), ["trade_date"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"sync_log",
|
||||
sa.Column(
|
||||
"id",
|
||||
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
|
||||
autoincrement=True,
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("source", sa.String(length=16), nullable=False),
|
||||
sa.Column("api", sa.String(length=32), nullable=False),
|
||||
sa.Column("request_time", sa.DateTime(), nullable=False),
|
||||
sa.Column("success", sa.Boolean(), nullable=False),
|
||||
sa.Column("failure_reason", sa.String(length=500), nullable=True),
|
||||
sa.Column("row_count", sa.Integer(), nullable=False),
|
||||
sa.Column("data_start", sa.Date(), nullable=True),
|
||||
sa.Column("data_end", sa.Date(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("sync_log", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_sync_log_source"), ["source"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"trading_calendar",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("calendar_date", sa.Date(), nullable=False),
|
||||
sa.Column("is_open", sa.Boolean(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("trading_calendar", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_trading_calendar_calendar_date"), ["calendar_date"], unique=True
|
||||
)
|
||||
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
with op.batch_alter_table("trading_calendar", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_trading_calendar_calendar_date"))
|
||||
|
||||
op.drop_table("trading_calendar")
|
||||
with op.batch_alter_table("sync_log", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_sync_log_source"))
|
||||
|
||||
op.drop_table("sync_log")
|
||||
with op.batch_alter_table("stock_daily", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_stock_daily_trade_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_stock_daily_symbol"))
|
||||
|
||||
op.drop_table("stock_daily")
|
||||
with op.batch_alter_table("stock", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_stock_symbol"))
|
||||
|
||||
op.drop_table("stock")
|
||||
with op.batch_alter_table("financial_indicator", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_symbol"))
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_report_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_financial_indicator_announce_date"))
|
||||
|
||||
op.drop_table("financial_indicator")
|
||||
with op.batch_alter_table("adjust_factor", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_adjust_factor_trade_date"))
|
||||
batch_op.drop_index(batch_op.f("ix_adjust_factor_symbol"))
|
||||
|
||||
op.drop_table("adjust_factor")
|
||||
# ### end Alembic commands ###
|
||||
@@ -3,3 +3,12 @@
|
||||
新增表流程(AGENT.md §12):Model → Alembic Migration → Test。
|
||||
模型统一继承 infra.persistence.sqlalchemy.base.Base。
|
||||
"""
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F401
|
||||
AdjustFactorModel,
|
||||
FinancialIndicatorModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
SyncLogModel,
|
||||
TradingCalendarModel,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Phase 1 市场数据表模型(SQLAlchemy 2.x 声明式)。
|
||||
|
||||
列名与 domain.entities.market 字段一一对应,便于 Repository 双向映射。
|
||||
Decimal 字段用 Numeric:SQLite 以浮点近似存储,未来 MySQL 下精确。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
Date,
|
||||
DateTime,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
# SQLite 只对 INTEGER PRIMARY KEY 自增;MySQL 下用 BIGINT
|
||||
PK_INT = BigInteger().with_variant(Integer, "sqlite")
|
||||
|
||||
SYMBOL_LEN = 12
|
||||
|
||||
|
||||
class StockModel(Base):
|
||||
__tablename__ = "stock"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), unique=True, index=True)
|
||||
name: Mapped[str] = mapped_column(String(64))
|
||||
industry: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
area: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
market: Mapped[str | None] = mapped_column(String(16), nullable=True)
|
||||
exchange: Mapped[str | None] = mapped_column(String(8), nullable=True)
|
||||
list_date: Mapped[date] = mapped_column(Date)
|
||||
delist_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(8), default="L")
|
||||
|
||||
|
||||
class TradingCalendarModel(Base):
|
||||
__tablename__ = "trading_calendar"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
calendar_date: Mapped[date] = mapped_column(Date, unique=True, index=True)
|
||||
is_open: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
|
||||
|
||||
class StockDailyModel(Base):
|
||||
"""不复权日线。"""
|
||||
|
||||
__tablename__ = "stock_daily"
|
||||
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_daily_symbol_date"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
open: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
high: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
low: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
close: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
volume: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
amount: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
|
||||
|
||||
class AdjustFactorModel(Base):
|
||||
__tablename__ = "adjust_factor"
|
||||
__table_args__ = (UniqueConstraint("symbol", "trade_date", name="uq_adj_symbol_date"),)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
trade_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
factor: Mapped[Decimal] = mapped_column(Numeric(20, 6))
|
||||
|
||||
|
||||
class FinancialIndicatorModel(Base):
|
||||
"""财务指标快照 —— report_date(报告期) 与 announce_date(公告日) 并存。"""
|
||||
|
||||
__tablename__ = "financial_indicator"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("symbol", "report_date", "announce_date", name="uq_fin_sym_rep_ann"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(SYMBOL_LEN), index=True)
|
||||
report_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
announce_date: Mapped[date] = mapped_column(Date, index=True)
|
||||
eps: Mapped[Decimal | None] = mapped_column(Numeric(12, 4), nullable=True)
|
||||
roe: Mapped[Decimal | None] = mapped_column(Numeric(10, 4), nullable=True)
|
||||
total_revenue: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
net_profit: Mapped[Decimal | None] = mapped_column(Numeric(24, 2), nullable=True)
|
||||
gross_margin: Mapped[Decimal | None] = mapped_column(Numeric(10, 4), nullable=True)
|
||||
|
||||
|
||||
class SyncLogModel(Base):
|
||||
__tablename__ = "sync_log"
|
||||
|
||||
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
|
||||
source: Mapped[str] = mapped_column(String(16), index=True)
|
||||
api: Mapped[str] = mapped_column(String(32))
|
||||
request_time: Mapped[datetime] = mapped_column(DateTime)
|
||||
success: Mapped[bool] = mapped_column(Boolean)
|
||||
failure_reason: Mapped[str | None] = mapped_column(String(500), nullable=True)
|
||||
row_count: Mapped[int] = mapped_column(default=0)
|
||||
data_start: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
data_end: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
@@ -0,0 +1,215 @@
|
||||
"""domain.repositories.market 的 SQLAlchemy 实现。
|
||||
|
||||
约定:本目录是唯一允许把 ORM 与业务实体互转的地方;
|
||||
Repository 以 domain.entities 类型进出(AGENT.md §10)。
|
||||
幂等键写在 __table_args__ 的 UniqueConstraint 上,upsert 先查后写,
|
||||
与 SQLite / MySQL 方言无关(未来切库不改业务层)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.entities.market import (
|
||||
AdjustFactor,
|
||||
DailyBar,
|
||||
FinancialIndicator,
|
||||
Stock,
|
||||
SyncLog,
|
||||
TradingCalendar,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
AdjustFactorModel,
|
||||
FinancialIndicatorModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
SyncLogModel,
|
||||
TradingCalendarModel,
|
||||
)
|
||||
|
||||
# 实体类型 → (ORM Model, 幂等键列)
|
||||
_TABLE = {
|
||||
Stock: (StockModel, ["symbol"]),
|
||||
TradingCalendar: (TradingCalendarModel, ["calendar_date"]),
|
||||
DailyBar: (StockDailyModel, ["symbol", "trade_date"]),
|
||||
AdjustFactor: (AdjustFactorModel, ["symbol", "trade_date"]),
|
||||
FinancialIndicator: (FinancialIndicatorModel, ["symbol", "report_date", "announce_date"]),
|
||||
SyncLog: (SyncLogModel, ["id"]),
|
||||
}
|
||||
|
||||
_ENTITY_TO_MODEL = {entity: model for entity, (model, _keys) in _TABLE.items()}
|
||||
|
||||
|
||||
def _fields_of(entity) -> dict:
|
||||
"""实体字段 → ORM 列名(模型列名与实体字段一致)。"""
|
||||
return {k: v for k, v in entity.model_dump().items() if k != "id"}
|
||||
|
||||
|
||||
def _upsert_by_business_key(
|
||||
session: Session,
|
||||
entity_cls,
|
||||
entities: Sequence,
|
||||
) -> int:
|
||||
"""按业务幂等键查重后 insert/update,返回触及行数(新增+更新)。
|
||||
|
||||
同一批内出现重复键(数据源偶发)时:先 flush 使前面已 add 的行可见,
|
||||
再按「后出现者覆盖」更新为最新值,避免 UNIQUE 冲突。
|
||||
"""
|
||||
model_cls, key_cols = _TABLE[entity_cls]
|
||||
seen: set[tuple] = set()
|
||||
touched = 0
|
||||
for ent in entities:
|
||||
values = _fields_of(ent)
|
||||
key = tuple(values[k] for k in key_cols)
|
||||
if key in seen:
|
||||
session.flush() # 让本批内先前新增的行进入 select 视野
|
||||
else:
|
||||
seen.add(key)
|
||||
filters = [getattr(model_cls, k) == values[k] for k in key_cols]
|
||||
row = session.scalars(select(model_cls).where(*filters)).first()
|
||||
if row is None:
|
||||
session.add(model_cls(**values))
|
||||
else:
|
||||
for col, val in values.items():
|
||||
setattr(row, col, val)
|
||||
touched += 1
|
||||
return touched
|
||||
|
||||
|
||||
class SqlAlchemyStockRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def get_by_symbol(self, symbol: str) -> Stock | None:
|
||||
row = self._session.scalars(select(StockModel).where(StockModel.symbol == symbol)).first()
|
||||
return Stock.model_validate(row.__dict__, from_attributes=True) if row else None
|
||||
|
||||
def list(self) -> list[Stock]:
|
||||
rows = self._session.scalars(select(StockModel).order_by(StockModel.symbol)).all()
|
||||
return [Stock.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def upsert_many(self, stocks: Sequence[Stock]) -> int:
|
||||
return _upsert_by_business_key(self._session, Stock, stocks)
|
||||
|
||||
|
||||
class SqlAlchemyTradingCalendarRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def upsert_many(self, days: Sequence[TradingCalendar]) -> int:
|
||||
return _upsert_by_business_key(self._session, TradingCalendar, days)
|
||||
|
||||
def list_range(self, start: date, end: date) -> list[TradingCalendar]:
|
||||
rows = self._session.scalars(
|
||||
select(TradingCalendarModel)
|
||||
.where(
|
||||
TradingCalendarModel.calendar_date >= start,
|
||||
TradingCalendarModel.calendar_date <= end,
|
||||
)
|
||||
.order_by(TradingCalendarModel.calendar_date)
|
||||
).all()
|
||||
return [TradingCalendar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def is_open(self, day: date) -> bool:
|
||||
row = self._session.scalars(
|
||||
select(TradingCalendarModel).where(TradingCalendarModel.calendar_date == day)
|
||||
).first()
|
||||
return bool(row.is_open) if row else False
|
||||
|
||||
|
||||
class SqlAlchemyDailyBarRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def upsert_many(self, bars: Sequence[DailyBar]) -> int:
|
||||
return _upsert_by_business_key(self._session, DailyBar, bars)
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]:
|
||||
rows = self._session.scalars(
|
||||
select(StockDailyModel)
|
||||
.where(
|
||||
StockDailyModel.symbol == symbol,
|
||||
StockDailyModel.trade_date >= start,
|
||||
StockDailyModel.trade_date <= end,
|
||||
)
|
||||
.order_by(StockDailyModel.trade_date)
|
||||
).all()
|
||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def latest_date(self, symbol: str) -> date | None:
|
||||
return self._session.scalar(
|
||||
select(StockDailyModel.trade_date)
|
||||
.where(StockDailyModel.symbol == symbol)
|
||||
.order_by(StockDailyModel.trade_date.desc())
|
||||
.limit(1)
|
||||
)
|
||||
|
||||
|
||||
class SqlAlchemyAdjustFactorRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def upsert_many(self, factors: Sequence[AdjustFactor]) -> int:
|
||||
return _upsert_by_business_key(self._session, AdjustFactor, factors)
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[AdjustFactor]:
|
||||
rows = self._session.scalars(
|
||||
select(AdjustFactorModel)
|
||||
.where(
|
||||
AdjustFactorModel.symbol == symbol,
|
||||
AdjustFactorModel.trade_date >= start,
|
||||
AdjustFactorModel.trade_date <= end,
|
||||
)
|
||||
.order_by(AdjustFactorModel.trade_date)
|
||||
).all()
|
||||
return [AdjustFactor.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
|
||||
class SqlAlchemyFinancialRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def upsert_many(self, rows: Sequence[FinancialIndicator]) -> int:
|
||||
return _upsert_by_business_key(self._session, FinancialIndicator, rows)
|
||||
|
||||
def list_announced(
|
||||
self,
|
||||
symbol: str,
|
||||
as_of_date: date,
|
||||
report_start: date | None = None,
|
||||
) -> list[FinancialIndicator]:
|
||||
"""只返回 announce_date <= as_of_date —— 防未来函数红线实现。"""
|
||||
stmt = (
|
||||
select(FinancialIndicatorModel)
|
||||
.where(
|
||||
FinancialIndicatorModel.symbol == symbol,
|
||||
FinancialIndicatorModel.announce_date <= as_of_date,
|
||||
)
|
||||
.order_by(FinancialIndicatorModel.announce_date)
|
||||
)
|
||||
if report_start is not None:
|
||||
stmt = stmt.where(FinancialIndicatorModel.report_date >= report_start)
|
||||
rows = self._session.scalars(stmt).all()
|
||||
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
|
||||
class SqlAlchemySyncLogRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def add(self, log: SyncLog) -> SyncLog:
|
||||
model = SyncLogModel(**log.model_dump())
|
||||
self._session.add(model)
|
||||
self._session.flush()
|
||||
return SyncLog.model_validate(model, from_attributes=True)
|
||||
|
||||
def recent(self, source: str | None = None, limit: int = 20) -> list[SyncLog]:
|
||||
stmt = select(SyncLogModel).order_by(SyncLogModel.id.desc()).limit(limit)
|
||||
if source is not None:
|
||||
stmt = stmt.where(SyncLogModel.source == source)
|
||||
rows = self._session.scalars(stmt).all()
|
||||
return [SyncLog.model_validate(r, from_attributes=True) for r in rows]
|
||||
Reference in New Issue
Block a user