tradingSystem/scripts/test_core_units.py

191 lines
7.3 KiB
Python
Raw Permalink 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.

# -*- coding: utf-8 -*-
"""
核心纯逻辑模块单测 (实机运行, 零外部依赖)
==========================================
运行: 在 tradingSystem 仓库根目录执行 python scripts/test_core_units.py
约定: 全部断言通过输出 "ALL PASS (n cases)" 并退出码 0; 任一失败输出明细并退出码 1。
覆盖: sizer 批次拆分/一手合并/组合约束/风险披露; cushion 摊薄成本/垫子状态/保垫触发/核销次序/补仓档。
"""
import os
import sys
import traceback
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from app.core.sizer import lot_qty, split_batches, check_caps, risk_exposure, risk_warnings # noqa: E402
from app.core.cushion import (PositionCost, cushion_state, trim_trigger, # noqa: E402
sell_allocation, dca_stage)
RESULTS = []
def case(name):
def deco(fn):
RESULTS.append((name, fn))
return fn
return deco
# ---------------------------------------------------------------- sizer
@case("批次拆分·常规 (12万目标, 10元股 → 6000/3000/3000 股)")
def _():
r = split_batches(120_000, 10.0)
assert r["ok"] and r["scheme"] == (0.5, 0.25, 0.25), r
assert [b["qty"] for b in r["batches"]] == [6000, 3000, 3000], r
@case("批次拆分·高价股自动合并 (12万目标, 700元股 → 合并为单批 100 股)")
def _():
# 50%=6万/700=85股 不足一手 → 60/40: 40%=4.8万=68股 仍不足 → 100%: 171→100股
r = split_batches(120_000, 700.0)
assert r["ok"] and r["scheme"] == (1.0,), r
assert len(r["batches"]) == 1 and r["batches"][0]["qty"] == 100, r
@case("批次拆分·极端高价买不起 (12万目标, 2000元股 → 放弃并给原因)")
def _():
r = split_batches(120_000, 2000.0)
assert not r["ok"] and "买不足一手" in r["reason"], r
@case("批次拆分·非法输入")
def _():
assert not split_batches(120_000, 0)["ok"]
assert not split_batches(0, 10.0)["ok"]
assert lot_qty(9_999, 100.0) == 0 and lot_qty(10_000, 100.0) == 100
@case("组合约束·单股上限拦截 (200万规模, 8%上限, 已持14万再买4万)")
def _():
ctx = dict(scale=2_000_000, portfolio_cap=0.60, stock_cap=0.08, max_names=15,
portfolio_mv=800_000, names_count=8, stock_mv=140_000, is_new_name=False,
sector=None, sector_max_names=4, sector_max_ratio=0.40)
v = check_caps(ts_code="600000.SH", add_amount=40_000, ctx=ctx)
assert len(v) == 1 and v[0].startswith("STOCK_CAP"), v
v2 = check_caps(ts_code="600000.SH", add_amount=10_000, ctx=ctx) # 14+1=15万 < 16万上限
assert v2 == [], v2
@case("组合约束·总仓+持仓数+行业同时拦截")
def _():
ctx = dict(scale=2_000_000, portfolio_cap=0.60, stock_cap=0.08, max_names=10,
portfolio_mv=1_180_000, names_count=10, stock_mv=0, is_new_name=True,
sector="半导体", sector_names=4, sector_mv=470_000,
sector_max_names=4, sector_max_ratio=0.40)
v = check_caps(ts_code="688000.SH", add_amount=60_000, ctx=ctx)
kinds = {x.split(":")[0] for x in v}
assert kinds == {"PORTFOLIO_CAP", "MAX_NAMES", "SECTOR_NAMES", "SECTOR_RATIO"}, v
@case("组合约束·行业数据源未配置时行业约束跳过")
def _():
ctx = dict(scale=2_000_000, portfolio_cap=0.60, stock_cap=0.08, max_names=15,
portfolio_mv=100_000, names_count=2, stock_mv=0, is_new_name=True,
sector=None, sector_names=99, sector_mv=9_999_999,
sector_max_names=4, sector_max_ratio=0.40)
assert check_caps(ts_code="000001.SZ", add_amount=60_000, ctx=ctx) == []
@case("风险披露·敞口计算与告警")
def _():
e = risk_exposure(6000, 10.0, 9.5) # 6000×0.5 = 3000 元
assert abs(e - 3000) < 1e-6, e
w = risk_warnings(entry_exposure=25_000, portfolio_exposure=0, scale=2_000_000)
assert len(w) == 1 and w[0].startswith("RISK_ENTRY"), w # 1.25% > 1%
w2 = risk_warnings(entry_exposure=-1, portfolio_exposure=0, scale=2_000_000)
assert w2 and w2[0].startswith("RISK_UNKNOWN"), w2
# ---------------------------------------------------------------- cushion
@case("摊薄成本·买入/卖出/T利润全链 (卖高于成本与做T都摊低成本)")
def _():
pc = PositionCost()
pc.buy(6000, 10.0) # 均价 10
assert abs(pc.avg_cost - 10.0) < 1e-9
pc.buy(3000, 11.0) # (60000+33000)/9000 = 10.333
assert abs(pc.avg_cost - 93000 / 9000) < 1e-9
pc.sell(3000, 12.0) # 净成本 93000-36000=57000, 6000股 → 9.5
assert abs(pc.avg_cost - 9.5) < 1e-9
pc.add_t_profit(1200) # (57000-1200)/6000 = 9.3
assert abs(pc.avg_cost - 9.3) < 1e-9
assert abs(pc.cushion(10.23) - 0.1) < 1e-3
pc.sell(6000, 10.0) # 清仓即结账
assert pc.avg_cost is None and pc.qty == 0
@case("摊薄成本·本金全部收回后成本归零 (垫子视为极厚)")
def _():
pc = PositionCost()
pc.buy(1000, 10.0)
pc.buy(1000, 10.0)
pc.sell(1000, 21.0) # 收回 21000 > 总投入 20000, 剩 1000 股净成本为负 → 0
assert pc.avg_cost == 0.0 and pc.cushion(10.0) == 9.99
@case("垫子状态机·边界 (solid=3%)")
def _():
assert cushion_state(None) == "NONE"
assert cushion_state(-0.001) == "NONE"
assert cushion_state(0.0) == "THIN"
assert cushion_state(0.0299) == "THIN"
assert cushion_state(0.03) == "SOLID"
@case("保垫减仓·触发边界 (峰值6%回吐一半)")
def _():
assert trim_trigger(0.08, 0.04) is True # 回吐 0.04 = 峰值一半
assert trim_trigger(0.08, 0.0401) is False
assert trim_trigger(0.05, -0.01) is False # 峰值不足 6% 永不触发
assert trim_trigger(0.06, 0.03) is True
@case("卖出核销次序·T0 → ADD(新→旧) → DCA → FILL → BASE")
def _():
lots = [
{"lot_id": "L_base", "lot_type": "BASE", "qty": 6000, "open_date": 20260701},
{"lot_id": "L_fill", "lot_type": "FILL", "qty": 3000, "open_date": 20260703},
{"lot_id": "L_add1", "lot_type": "ADD", "qty": 2000, "open_date": 20260710},
{"lot_id": "L_add2", "lot_type": "ADD", "qty": 1000, "open_date": 20260720},
{"lot_id": "L_t0", "lot_type": "T0", "qty": 1000, "open_date": 20260727},
]
alloc = sell_allocation(lots, 4500)
assert [a["lot_id"] for a in alloc] == ["L_t0", "L_add2", "L_add1", "L_fill"], alloc
assert [a["qty"] for a in alloc] == [1000, 1000, 2000, 500], alloc
try:
sell_allocation(lots, 99_999)
assert False, "超量卖出未拦截"
except ValueError:
pass
@case("补仓评估档·-8%/-15% 两档")
def _():
assert dca_stage(-0.05) == 0
assert dca_stage(-0.08) == 1
assert dca_stage(-0.12) == 1
assert dca_stage(-0.15) == 2
assert dca_stage(-0.30) == 2
# ---------------------------------------------------------------- runner
def main():
passed, failed = 0, 0
for name, fn in RESULTS:
try:
fn()
print(f" PASS {name}")
passed += 1
except Exception:
print(f" FAIL {name}")
traceback.print_exc()
failed += 1
print("-" * 60)
if failed:
print(f"FAILED: {failed} / {passed + failed}")
sys.exit(1)
print(f"ALL PASS ({passed} cases)")
if __name__ == "__main__":
main()