93 lines
3.6 KiB
Python
93 lines
3.6 KiB
Python
"""股票池加载:前缀过滤/时间缓冲/vwap/bar_idx/显式symbols/limit抽样."""
|
|
import sqlite3
|
|
import sys, os
|
|
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
|
|
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0")))
|
|
|
|
import pytest
|
|
from sanguo_factor.universe import load_universe_bars, evaluation_filter, WARMUP_BARS
|
|
|
|
_DDL = """
|
|
CREATE TABLE dbbardata(
|
|
symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT,
|
|
volume REAL, turnover REAL, open_interest REAL,
|
|
open_price REAL, high_price REAL, low_price REAL, close_price REAL)
|
|
"""
|
|
|
|
|
|
def _mk_db(tmp_path, rows):
|
|
db = str(tmp_path / "qt.db")
|
|
conn = sqlite3.connect(db)
|
|
conn.execute(_DDL)
|
|
conn.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
|
|
conn.commit(); conn.close()
|
|
return db
|
|
|
|
|
|
def _row(sym, ex, day, close, volume=100.0, turnover=None):
|
|
return (sym, ex, f"{day} 00:00:00", "d", volume,
|
|
turnover if turnover is not None else close * volume, 0,
|
|
close, close, close, close)
|
|
|
|
|
|
@pytest.fixture()
|
|
def db(tmp_path):
|
|
from datetime import date, timedelta
|
|
rows = []
|
|
# 600000:130根预热(2017-08~2017-12,唯一日期) + 评估窗内 3 根
|
|
for i in range(130):
|
|
day = (date(2017, 8, 1) + timedelta(days=i)).isoformat()
|
|
rows.append(_row("600000", "SSE", day, 10.0 + i * 0.01))
|
|
for d in ("2018-01-02", "2018-01-03", "2018-01-04"):
|
|
rows.append(_row("600000", "SSE", d, 11.0))
|
|
# 000001:只有 2 根(次新,bar_idx<WARMUP 应被 evaluation_filter 排除)
|
|
for d in ("2018-01-02", "2018-01-03"):
|
|
rows.append(_row("000001", "SZSE", d, 5.0))
|
|
# 应剔除:科创68/北交8开头/ETF 510300
|
|
rows.append(_row("688001", "SSE", "2018-01-02", 20.0))
|
|
rows.append(_row("830001", "BJSE", "2018-01-02", 3.0))
|
|
rows.append(_row("510300", "SSE", "2018-01-02", 4.0))
|
|
# 非日线 interval 应忽略
|
|
rows.append(("600000", "SSE", "2018-01-02 09:35:00", "15m", 1, 1, 0, 1, 1, 1, 1))
|
|
return _mk_db(tmp_path, rows)
|
|
|
|
|
|
def test_load_filters_prefix_and_interval(db):
|
|
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
|
syms = df["vt_symbol"].unique().to_list()
|
|
assert set(syms) == {"600000.SSE", "000001.SZSE"}
|
|
|
|
|
|
def test_load_lookback_buffer_and_vwap(db):
|
|
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
|
# lookback 生效:2017-08 的 bar 也在(因子窗口预热需要)
|
|
assert str(df["datetime"].min()).startswith("2017-08")
|
|
# vwap = turnover/volume
|
|
row = df.filter(df["vt_symbol"] == "600000.SSE").sort("datetime").row(0, named=True)
|
|
assert row["vwap"] == pytest.approx(10.0)
|
|
|
|
|
|
def test_bar_idx_per_symbol(db):
|
|
df = load_universe_bars(db, "2018-01-01", "2018-01-31").sort(["vt_symbol", "datetime"])
|
|
idx = df.filter(df["vt_symbol"] == "000001.SZSE")["bar_idx"].to_list()
|
|
assert idx == [0, 1]
|
|
|
|
|
|
def test_evaluation_filter_warmup(db):
|
|
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
|
ev = evaluation_filter(df, "2018-01-01", "2018-01-31")
|
|
# 000001 只有2根 < WARMUP_BARS → 全排除;600000 保留
|
|
assert set(ev["vt_symbol"].unique().to_list()) == {"600000.SSE"}
|
|
assert ev.height == 3
|
|
|
|
|
|
def test_explicit_symbols(db):
|
|
df = load_universe_bars(db, "2018-01-01", "2018-01-31", symbols=["000001"])
|
|
assert set(df["vt_symbol"].unique().to_list()) == {"000001.SZSE"}
|
|
|
|
|
|
def test_limit_deterministic(db):
|
|
df1 = load_universe_bars(db, "2018-01-01", "2018-01-31", limit=1)
|
|
df2 = load_universe_bars(db, "2018-01-01", "2018-01-31", limit=1)
|
|
assert set(df1["vt_symbol"].unique()) == set(df2["vt_symbol"].unique())
|