Tortoise/sidecar/storage.py
2026-06-23 10:37:31 +08:00

210 lines
7.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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,
]