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:
@@ -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