210 lines
7.6 KiB
Python
210 lines
7.6 KiB
Python
"""DuckDB 本地缓存存储。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from datetime import datetime, timedelta, timezone
|
||
from pathlib import Path
|
||
from typing import Any, List, Optional
|
||
|
||
from sidecar.models import KlineBar, Snapshot
|
||
|
||
|
||
class DuckDBStore:
|
||
"""DuckDB 行情缓存仓储。"""
|
||
|
||
def __init__(self, db_path: str) -> None:
|
||
"""
|
||
功能说明:创建 DuckDB 缓存仓储。
|
||
参数说明:db_path 为 DuckDB 文件路径。
|
||
返回值说明:无返回值。
|
||
注意事项:DuckDB 依赖在初始化连接时才导入,便于测试解析纯 Python 代码。
|
||
"""
|
||
self.db_path = db_path
|
||
Path(db_path).parent.mkdir(parents=True, exist_ok=True)
|
||
import duckdb
|
||
|
||
self.conn = duckdb.connect(db_path)
|
||
self.initialize()
|
||
|
||
def initialize(self) -> None:
|
||
"""
|
||
功能说明:初始化 sidecar 所需 schema 与表。
|
||
参数说明:无。
|
||
返回值说明:无返回值。
|
||
注意事项:该方法可重复执行,不会清空已有缓存数据。
|
||
"""
|
||
self.conn.execute("CREATE SCHEMA IF NOT EXISTS md")
|
||
self.conn.execute(
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS md.kline_bars (
|
||
symbol TEXT NOT NULL,
|
||
period TEXT NOT NULL,
|
||
trade_date DATE NOT NULL,
|
||
open DOUBLE,
|
||
high DOUBLE,
|
||
low DOUBLE,
|
||
close DOUBLE,
|
||
volume DOUBLE,
|
||
amount DOUBLE,
|
||
source TEXT NOT NULL,
|
||
fetched_at TIMESTAMP NOT NULL,
|
||
PRIMARY KEY (symbol, period, trade_date)
|
||
)
|
||
"""
|
||
)
|
||
self.conn.execute(
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS md.snapshots (
|
||
symbol TEXT NOT NULL,
|
||
name TEXT,
|
||
trade_time TIMESTAMP,
|
||
price DOUBLE,
|
||
previous_close DOUBLE,
|
||
open DOUBLE,
|
||
high DOUBLE,
|
||
low DOUBLE,
|
||
volume DOUBLE,
|
||
amount DOUBLE,
|
||
change DOUBLE,
|
||
change_percent DOUBLE,
|
||
turnover_rate DOUBLE,
|
||
pe_ttm DOUBLE,
|
||
pe_static DOUBLE,
|
||
pb DOUBLE,
|
||
market_cap DOUBLE,
|
||
float_market_cap DOUBLE,
|
||
limit_up DOUBLE,
|
||
limit_down DOUBLE,
|
||
source TEXT NOT NULL,
|
||
fetched_at TIMESTAMP NOT NULL,
|
||
expires_at TIMESTAMP NOT NULL,
|
||
PRIMARY KEY (symbol)
|
||
)
|
||
"""
|
||
)
|
||
|
||
def save_kline(self, bars: List[KlineBar]) -> None:
|
||
"""
|
||
功能说明:批量写入 K 线缓存。
|
||
参数说明:bars 为统一 KlineBar 列表。
|
||
返回值说明:无返回值。
|
||
注意事项:相同 symbol、period、trade_date 的记录会被覆盖。
|
||
"""
|
||
if not bars:
|
||
return
|
||
now = datetime.now(timezone.utc).replace(tzinfo=None)
|
||
self.conn.execute("BEGIN TRANSACTION")
|
||
try:
|
||
for bar in bars:
|
||
self.conn.execute(
|
||
"DELETE FROM md.kline_bars WHERE symbol = ? AND period = ? AND trade_date = ?",
|
||
[bar.symbol, bar.period, bar.trade_date],
|
||
)
|
||
self.conn.execute(
|
||
"INSERT INTO md.kline_bars VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||
[
|
||
bar.symbol,
|
||
bar.period,
|
||
bar.trade_date,
|
||
bar.open,
|
||
bar.high,
|
||
bar.low,
|
||
bar.close,
|
||
bar.volume,
|
||
bar.amount,
|
||
bar.source,
|
||
now,
|
||
],
|
||
)
|
||
self.conn.execute("COMMIT")
|
||
except Exception:
|
||
self.conn.execute("ROLLBACK")
|
||
raise
|
||
|
||
def load_kline(self, symbol: str, period: str, limit: int) -> List[KlineBar]:
|
||
"""
|
||
功能说明:读取本地 K 线缓存。
|
||
参数说明:symbol 为标准股票代码,period 为周期,limit 为最大条数。
|
||
返回值说明:按交易日期升序返回 KlineBar 列表。
|
||
注意事项:内部先取最近 limit 条,再恢复为升序,方便策略直接消费。
|
||
"""
|
||
rows = self.conn.execute(
|
||
"""
|
||
SELECT symbol, period, trade_date, open, high, low, close, volume, amount, source
|
||
FROM (
|
||
SELECT * FROM md.kline_bars
|
||
WHERE symbol = ? AND period = ?
|
||
ORDER BY trade_date DESC
|
||
LIMIT ?
|
||
)
|
||
ORDER BY trade_date ASC
|
||
""",
|
||
[symbol, period, limit],
|
||
).fetchall()
|
||
return [KlineBar(*row) for row in rows]
|
||
|
||
def save_snapshot(self, snapshot: Snapshot, ttl_seconds: int) -> None:
|
||
"""
|
||
功能说明:写入行情快照缓存。
|
||
参数说明:snapshot 为统一快照,ttl_seconds 为缓存有效秒数。
|
||
返回值说明:无返回值。
|
||
注意事项:同一 symbol 只保留最新一条快照。
|
||
"""
|
||
now = datetime.now(timezone.utc).replace(tzinfo=None)
|
||
expires_at = now + timedelta(seconds=ttl_seconds)
|
||
self.conn.execute("DELETE FROM md.snapshots WHERE symbol = ?", [snapshot.symbol])
|
||
self.conn.execute(
|
||
"INSERT INTO md.snapshots VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||
self._snapshot_row(snapshot) + [now, expires_at],
|
||
)
|
||
|
||
def load_snapshot(self, symbol: str) -> Optional[Snapshot]:
|
||
"""
|
||
功能说明:读取未过期的行情快照缓存。
|
||
参数说明:symbol 为标准股票代码。
|
||
返回值说明:命中返回 Snapshot,未命中或已过期返回 None。
|
||
注意事项:过期判断使用当前 UTC 时间。
|
||
"""
|
||
now = datetime.now(timezone.utc).replace(tzinfo=None)
|
||
row = self.conn.execute(
|
||
"""
|
||
SELECT symbol, name, trade_time, price, previous_close, open, high, low,
|
||
volume, amount, change, change_percent, turnover_rate, pe_ttm,
|
||
pe_static, pb, market_cap, float_market_cap, limit_up, limit_down, source
|
||
FROM md.snapshots
|
||
WHERE symbol = ? AND expires_at > ?
|
||
""",
|
||
[symbol, now],
|
||
).fetchone()
|
||
return Snapshot(*row) if row else None
|
||
|
||
def _snapshot_row(self, snapshot: Snapshot) -> List[Any]:
|
||
"""
|
||
功能说明:把 Snapshot 转换为数据库行字段。
|
||
参数说明:snapshot 为统一快照模型。
|
||
返回值说明:返回与 md.snapshots 前 21 列一致的列表。
|
||
注意事项:该私有函数不包含 fetched_at 与 expires_at。
|
||
"""
|
||
return [
|
||
snapshot.symbol,
|
||
snapshot.name,
|
||
snapshot.trade_time,
|
||
snapshot.price,
|
||
snapshot.previous_close,
|
||
snapshot.open,
|
||
snapshot.high,
|
||
snapshot.low,
|
||
snapshot.volume,
|
||
snapshot.amount,
|
||
snapshot.change,
|
||
snapshot.change_percent,
|
||
snapshot.turnover_rate,
|
||
snapshot.pe_ttm,
|
||
snapshot.pe_static,
|
||
snapshot.pb,
|
||
snapshot.market_cap,
|
||
snapshot.float_market_cap,
|
||
snapshot.limit_up,
|
||
snapshot.limit_down,
|
||
snapshot.source,
|
||
] |