Files
sanguo_vnpy_v2/tests/factor/test_universe.py
T

91 lines
3.5 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):
rows = []
# 600000:60根(2017-10~2017-12) + 评估窗内 3 根
for i in range(60):
rows.append(_row("600000", "SSE", f"2017-10-{(i % 28) + 1:02d}", 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-10 的 bar 也在(因子窗口预热需要)
assert str(df["datetime"].min()).startswith("2017-10")
# 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())