This commit is contained in:
2026-09-09 15:07:58 +08:00
parent d656c05b3d
commit 71a0f6e404
31 changed files with 3657 additions and 9 deletions

View File

@@ -0,0 +1,65 @@
"""复权换算 adjust_barsbfq/qfq/hfq 乘数、阶梯因子向前沿用、名义量不缩放。"""
from __future__ import annotations
from datetime import datetime
import pytest
from app.api._deps import adjust_bars
from app.domain import Bar
def _bars(*rows) -> list[Bar]:
# (date, open, close)
return [Bar(ts=datetime(y, m, d), open=o, high=o, low=o, close=c, volume=100.0)
for (y, m, d), o, c in rows]
def _factors(*pairs):
return [(datetime(y, m, d), f) for (y, m, d), f in pairs]
def test_bfq_to_bfq_identity():
bars = _bars(((2024, 1, 2), 10.0, 11.0))
out = adjust_bars(bars, _factors(((2024, 1, 2), 2.0)), "bfq", "bfq")
assert out[0].close == 11.0
def test_qfq_normalizes_by_latest_factor():
# 因子 1.0 -> 2.0中途除权qfq = f(t)/f(latest)
bars = _bars(
((2024, 1, 2), 10.0, 10.0), # f=1.0
((2024, 6, 3), 5.0, 5.0), # f=2.0(除权日,价格腰斩)
)
factors = _factors(((2024, 1, 2), 1.0), ((2024, 6, 3), 2.0))
out = adjust_bars(bars, factors, "bfq", "qfq")
# 除权前按 1.0/2.0 缩放 -> 5.0;除权后 2.0/2.0 -> 原价
assert out[0].close == pytest.approx(5.0)
assert out[1].close == pytest.approx(5.0)
def test_hfq_scales_by_factor():
bars = _bars(((2024, 1, 2), 10.0, 10.0), ((2024, 6, 3), 5.0, 5.0))
factors = _factors(((2024, 1, 2), 1.0), ((2024, 6, 3), 2.0))
out = adjust_bars(bars, factors, "bfq", "hfq")
assert out[0].close == pytest.approx(10.0) # f=1
assert out[1].close == pytest.approx(10.0) # 5 * 2
def test_factor_step_forward_fill():
"""因子是阶梯函数:变化点之间的日期向前沿用最近因子。"""
bars = _bars(((2024, 2, 1), 8.0, 8.0)) # 在 1/2 与 6/3 之间 -> 沿用 1.0
factors = _factors(((2024, 1, 2), 1.0), ((2024, 6, 3), 2.0))
out = adjust_bars(bars, factors, "bfq", "hfq")
assert out[0].close == pytest.approx(8.0) # 8 * 1.0
def test_nominal_columns_not_scaled():
"""成交额/换手率是名义量,不随复权缩放。"""
bars = [Bar(ts=datetime(2024, 1, 2), open=10, high=10, low=10, close=10,
volume=100.0, amount=1_000_000.0, turnover=1.5)]
out = adjust_bars(bars, _factors(((2024, 1, 2), 4.0)), "bfq", "hfq")
assert out[0].close == pytest.approx(40.0)
assert out[0].amount == 1_000_000.0
assert out[0].turnover == 1.5
assert out[0].volume == 100.0

View File

@@ -0,0 +1,48 @@
"""登录 IP 限速auth_api._login_rate_limited窗口内超限拦截、过期滑出、独立 IP 隔离。"""
from __future__ import annotations
import pytest
import app.auth_api as auth_api
from app.auth_api import _LOGIN_MAX_PER_WINDOW, _LOGIN_WINDOW, _login_attempts, _login_rate_limited
class _FakeClock:
"""替换 auth_api 命名空间里的 time不影响全局 time 模块)。"""
def __init__(self):
self.now = 1000.0
def monotonic(self) -> float:
return self.now
@pytest.fixture()
def clock(monkeypatch):
c = _FakeClock()
monkeypatch.setattr(auth_api, "time", c)
_login_attempts.clear()
yield c
_login_attempts.clear()
def test_under_limit_passes(clock):
for _ in range(_LOGIN_MAX_PER_WINDOW):
assert _login_rate_limited("1.2.3.4") is False
assert _login_rate_limited("1.2.3.4") is True # 第 16 次被拦
def test_window_slides(clock):
for _ in range(_LOGIN_MAX_PER_WINDOW):
_login_rate_limited("1.2.3.4")
assert _login_rate_limited("1.2.3.4") is True
# 窗口滑过:最早的尝试过期出窗,重新放行
clock.now += _LOGIN_WINDOW + 0.1
assert _login_rate_limited("1.2.3.4") is False
def test_ips_isolated(clock):
for _ in range(_LOGIN_MAX_PER_WINDOW):
_login_rate_limited("1.1.1.1")
assert _login_rate_limited("1.1.1.1") is True
assert _login_rate_limited("2.2.2.2") is False # 另一 IP 不受牵连

View File

@@ -0,0 +1,94 @@
"""cache.py 本地层(不碰 RedisTTL 过期、容量淘汰、熔断冷却恢复。"""
from __future__ import annotations
import pytest
import app.cache as cache_mod
from app import cache
class _FakeClock:
"""替换 cache 命名空间里的 time不影响全局 time 模块)。"""
def __init__(self):
self.now = 1000.0
def monotonic(self) -> float:
return self.now
def time(self) -> float:
return self.now
@pytest.fixture(autouse=True)
def _reset_state():
"""每个用例干净的本地缓存状态。"""
cache._local_store.clear()
cache._local_bytes = 0
cache._disabled_until = 0.0
yield
cache._local_store.clear()
cache._local_bytes = 0
cache._disabled_until = 0.0
def test_local_set_get_roundtrip():
cache.local_set("k", '{"a":1}', ttl=60)
assert cache.local_get("k") == '{"a":1}'
def test_local_get_miss():
assert cache.local_get("nope") is None
def test_local_expiry(monkeypatch):
clock = _FakeClock()
monkeypatch.setattr(cache_mod, "time", clock)
cache.local_set("k", "v", ttl=10)
clock.now = 1005.0
assert cache.local_get("k") == "v"
clock.now = 1101.0 # 过期
assert cache.local_get("k") is None
assert "k" not in cache._local_store # 过期读取顺手清理
def test_local_ttl_capped_at_120s():
"""本地层恒 ≤120s多进程部署时最多比 Redis 多陈旧 120s 的约定)。"""
base = cache.time.monotonic()
cache.local_set("k", "v", ttl=99999)
ent = cache._local_store["k"]
assert ent[0] - base <= 120.0 + 5 # 相对当前 monotonic 的上限(留误差余量)
def test_local_entries_cap_evicts_oldest():
for i in range(cache._LOCAL_MAX_ENTRIES + 5):
cache.local_set(f"k{i}", "v", ttl=60)
assert len(cache._local_store) <= cache._LOCAL_MAX_ENTRIES
# 先插入的(最旧)被近似 LRU 淘汰
assert cache.local_get("k0") is None
assert cache.local_get(f"k{cache._LOCAL_MAX_ENTRIES + 4}") == "v"
def test_local_overwrite_releases_bytes():
cache.local_set("k", "x" * 1000, ttl=60)
before = cache._local_bytes
cache.local_set("k", "y", ttl=60)
assert cache._local_bytes < before
assert cache.local_get("k") == "y"
def test_bail_cooldown_recovers(monkeypatch):
"""熔断 60s期间 _client 为 None到期自动放行Redis 未配置时也返回 None但不熔断"""
clock = _FakeClock()
monkeypatch.setattr(cache_mod, "time", clock)
cache._bail()
assert cache._disabled_until > clock.now
clock.now += 59.0
assert clock.monotonic() < cache._disabled_until # 仍在熔断期
clock.now += 2.0 # 越过 60s 冷却
assert clock.monotonic() >= cache._disabled_until # 恢复探测资格
def test_digest_stable_and_distinct():
assert cache.digest("a", 1, None) == cache.digest("a", 1, None)
assert cache.digest("a", 1) != cache.digest("a", 2)

View File

@@ -0,0 +1,60 @@
"""事件回测纯函数:入场/出场索引语义、汇总统计(含空样本与分年)。"""
from __future__ import annotations
from datetime import datetime
import pytest
from app.backtest.events import _entry_exit_indices, _stats_block
from app.schemas import EventBacktestSpec
def _spec(holding_days: int = 5) -> EventBacktestSpec:
return EventBacktestSpec.model_validate({
"holding_days": holding_days,
"entry": {"indicator": []},
})
def test_entry_exit_next_day_and_hold():
spec = _spec(holding_days=5)
assert _entry_exit_indices(10, spec, n=100) == (11, 16) # 次日入,持有 5 日出
def test_entry_exit_out_of_range_none():
spec = _spec(holding_days=5)
# 出场索引越界exit_i >= n
assert _entry_exit_indices(94, spec, n=100) is None
assert _entry_exit_indices(99, spec, n=100) is None
def test_entry_exit_boundary_exact_fit():
spec = _spec(holding_days=5)
# exit_i == n-1 恰好可用
assert _entry_exit_indices(93, spec, n=100) == (94, 99)
def test_stats_block_empty():
s = _stats_block([])
assert s["samples"] == 0 and s["stocks"] == 0
assert s["by_year"] == []
assert s["win_rate"] == 0.0
def test_stats_block_aggregates():
trades = [
{"ts_code": "000001.SZ", "ret_pct": 10.0,
"entry_date": datetime(2024, 1, 5), "entry_price": 10, "exit_price": 11},
{"ts_code": "000001.SZ", "ret_pct": -4.0,
"entry_date": datetime(2024, 3, 6), "entry_price": 10, "exit_price": 9.6},
{"ts_code": "600519.SH", "ret_pct": 2.0,
"entry_date": datetime(2023, 5, 10), "entry_price": 10, "exit_price": 10.2},
]
s = _stats_block(trades)
assert s["samples"] == 3
assert s["stocks"] == 2
assert s["mean_pct"] == pytest.approx((10.0 - 4.0 + 2.0) / 3, abs=1e-3) # 统计块保留 3 位小数
assert s["win_rate"] == round(2 / 3 * 100, 2)
assert [y["year"] for y in s["by_year"]] == [2023, 2024] # 分年升序
assert s["by_year"][1]["samples"] == 2
assert s["max_pct"] == 10.0 and s["min_pct"] == -4.0

View File

@@ -0,0 +1,85 @@
"""sync_utils 纯函数:取值清洗、日期格式化、限频重试语义。"""
from __future__ import annotations
import math
from datetime import timedelta
import pytest
from app.data import sync_utils
def test_s_clean():
assert sync_utils.s_clean(" 平安银行 ") == "平安银行"
assert sync_utils.s_clean("") is None
assert sync_utils.s_clean(" ") is None
assert sync_utils.s_clean(None) is None
assert sync_utils.s_clean(float("nan")) is None # pandas NaN
def test_f_clean():
assert sync_utils.f_clean("3.14") == 3.14
assert sync_utils.f_clean(2) == 2.0
assert sync_utils.f_clean(float("nan")) is None
assert sync_utils.f_clean(None) is None
assert sync_utils.f_clean("abc") is None
assert sync_utils.f_clean(math.inf) == math.inf # inf 非 NaN原样保留
def test_d8_iso():
assert sync_utils.d8_iso("20240102") == "2024-01-02"
assert sync_utils.d8_iso(20240102) == "2024-01-02"
assert sync_utils.d8_iso(None) is None
assert sync_utils.d8_iso("") is None
def test_fresh():
now = sync_utils.utcnow()
assert sync_utils.fresh(now, days=7) is True
assert sync_utils.fresh(now - timedelta(days=8), days=7) is False
assert sync_utils.fresh(None, days=7) is False
def test_call_retry_passes_through_args():
calls = []
def fn(a, b=0):
calls.append((a, b))
return a + b
assert sync_utils.call_retry(fn, 1, b=2) == 3
assert calls == [(1, 2)]
def test_call_retry_rate_limit_retries_once(monkeypatch):
"""「每分钟」级频率超限等 62s 重试一次;重试成功则返回结果。"""
monkeypatch.setattr(sync_utils.time, "sleep", lambda s: None)
calls = []
def fn():
calls.append(1)
if len(calls) == 1:
raise RuntimeError("抱歉您每分钟最多访问该接口5次")
return "ok"
assert sync_utils.call_retry(fn) == "ok"
assert len(calls) == 2
def test_call_retry_hourly_limit_raises(monkeypatch):
"""小时级限频不重试,直接抛出。"""
monkeypatch.setattr(sync_utils.time, "sleep", lambda s: None)
def fn():
raise RuntimeError("您每小时最多访问该接口10次")
with pytest.raises(RuntimeError):
sync_utils.call_retry(fn)
def test_call_retry_other_errors_raise():
def fn():
raise ValueError("数据源故障")
with pytest.raises(ValueError):
sync_utils.call_retry(fn)

View File

@@ -0,0 +1,101 @@
"""交割单解析器:四类真实导出格式 + 边界(转账/配号/利息跳过、费用合计去重、日期多格式)。
移植自 scripts/test_trades_parser.py已删不碰数据库直接调 app.trades.parse_statement。
"""
from __future__ import annotations
import io
from datetime import datetime as dt
import pytest
from fastapi import HTTPException
from openpyxl import Workbook
from app.trades import parse_statement
def test_tdx_gbk_tabs():
"""通达信式GBK + 制表符 + 标题行在前;跳过银行转存/利息归本。"""
tdx = (
"交割单\n"
"股东账号: A123456789 起始日期: 20240102 终止日期: 20240105 币种: 人民币\n"
"\t交割日期\t业务名称\t证券代码\t证券名称\t成交价格\t成交数量\t成交金额\t手续费\t印花税\t过户费\t发生金额\t资金余额\t合同号\n"
"\t20240102\t证券买入\t600519\t贵州茅台\t1680.00\t100\t168000.00\t5.00\t0.00\t1.68\t-168006.68\t200000.00\t1000001\n"
"\t20240102\t银行转存\t\t\t\t\t\t\t\t\t50000.00\t250000.00\t\n"
"\t20240103\t证券卖出\t600519\t贵州茅台\t1700.50\t100\t170050.00\t5.00\t170.05\t1.70\t169873.25\t419873.25\t1000002\n"
"\t20240105\t利息归本\t\t\t\t\t\t\t\t\t1.25\t419874.50\t\n"
)
r = parse_statement(tdx.encode("gbk"), "交割单.txt")
assert len(r.trades) == 2
assert r.skipped_other == 2
t0, t1 = r.trades[0], r.trades[1]
assert (t0.trade_date.isoformat(), t0.ts_code) == ("2024-01-02", "600519.SH")
assert t0.direction == "buy" and abs(t0.fee - 6.68) < 1e-9
assert t1.direction == "sell" and abs(t1.fee - 176.75) < 1e-9
assert t0.amount == 168000.0
def test_hengsheng_csv_fee_total():
"""恒生柜台式UTF-8 CSV费用合计列不与其他费用列重复累加。"""
hs = (
"序号,交收日期,证券代码,证券名称,交易类别,成交价格,成交数量,证券余额,成交金额,资金发生数,资金余额,流水序号,业务标志,业务名称,发生金额,后资金额,货币类别,费用合计,净佣金,规费,印花税,过户费,合同号\n"
"1,2024-06-07,000858,五粮液,证券买入,132.50,200,200,26500.00,-26505.80,73494.20,1,0101,证券买入,-26505.80,73494.20,人民币,5.80,4.20,1.60,0.00,0.00,66778001\n"
"2,2024-06-07,,,\t,,,,5120.00,78614.20,2,2041,银行转存,5120.00,78614.20,人民币,0,0,0,0,0,\n"
"3,2024-06-10,000858,五粮液,证券卖出,135.00,200,0,27000.00,26975.30,105589.50,3,0102,证券卖出,26975.30,105589.50,人民币,24.70,4.20,1.60,18.90,0.00,66779001\n"
)
r = parse_statement(hs.encode("utf-8"), "hsi.csv")
assert len(r.trades) == 2
assert abs(r.trades[1].fee - 24.70) < 1e-9
assert r.trades[0].ts_code == "000858.SZ"
assert r.trades[0].trade_date.isoformat() == "2024-06-07"
def test_html_pseudo_xls():
"""HTML 伪 .xls同花顺导出常见真身千分位金额、斜杠日期、创业板后缀。"""
html = """<html><head><meta charset="gbk"></head><body>
<table>
<tr><td>客户姓名</td><td>测试</td></tr>
<tr><td>成交日期</td><td>业务名称</td><td>证券代码</td><td>证券名称</td><td>成交价格</td><td>成交数量</td><td>成交金额</td><td>手续费</td></tr>
<tr><td>2024/03/15</td><td>证券买入</td><td>300750</td><td>宁德时代</td><td>182.30</td><td>300</td><td>54,690.00</td><td>16.41</td></tr>
<tr><td>2024/03/18</td><td>证券卖出</td><td>300750</td><td>宁德时代</td><td>185.00</td><td>300</td><td>55,500.00</td><td>5.55</td></tr>
</table></body></html>"""
r = parse_statement(html.encode("gbk"), "jiaogedan.xls")
assert len(r.trades) == 2
assert r.trades[0].amount == 54690.0
assert r.trades[0].ts_code == "300750.SZ"
assert r.trades[1].trade_date.isoformat() == "2024-03-18"
def test_amount_sign_direction():
"""无业务名称列(招商式):发生金额正负判方向。"""
zh = (
"证券名称,成交日期,成交价格,成交数量,发生金额,资金余额,合同编号\n"
"贵州茅台,20240102,1680.00,100,-168005.00,200000.00,SZ1000001\n"
"贵州茅台,20240103,1700.50,100,170049.50,370049.50,SZ1000002\n"
)
r = parse_statement(zh.encode("utf-8"), "zszs.csv")
assert len(r.trades) == 2
assert (r.trades[0].direction, r.trades[1].direction) == ("buy", "sell")
def test_xlsx_openpyxl():
"""xlsxopenpyxl 内存构造datetime 日期、科创板后缀、佣金+过户费合计。"""
wb = Workbook()
ws = wb.active
ws.append(["对账单", None, None])
ws.append(["成交日期", "业务名称", "证券代码", "证券名称", "成交均价", "成交股数", "成交金额", "佣金", "过户费"])
ws.append([dt(2024, 2, 28, 14, 35, 0), "证券买入", "688981", "中芯国际", 52.80, 200, 10560.00, 2.50, 1.06])
ws.append([dt(2024, 3, 1, 9, 31, 0), "证券卖出", "688981", "中芯国际", 54.10, 200, 10820.00, 2.50, 1.06])
buf = io.BytesIO()
wb.save(buf)
r = parse_statement(buf.getvalue(), "sm.xlsx")
assert len(r.trades) == 2
assert r.trades[0].trade_date.isoformat() == "2024-02-28"
assert r.trades[0].ts_code == "688981.SH"
assert abs(r.trades[0].fee - 3.56) < 1e-9
def test_garbage_input_422():
with pytest.raises(HTTPException) as ei:
parse_statement("随便一串不是交割单的文字,1,2,3".encode("utf-8"), "x.csv")
assert ei.value.status_code == 422