diff options
Diffstat (limited to 'tests')
| -rw-r--r-- | tests/__init__.py | 0 | ||||
| -rw-r--r-- | tests/fixtures_synth.py | 39 | ||||
| -rw-r--r-- | tests/worker/test_backtest.py | 270 | ||||
| -rw-r--r-- | tests/worker/test_data.py | 336 | ||||
| -rw-r--r-- | tests/worker/test_fetch_cli.py | 144 | ||||
| -rw-r--r-- | tests/worker/test_normalize.py | 92 | ||||
| -rw-r--r-- | tests/worker/test_rules_regression.py | 201 |
7 files changed, 1082 insertions, 0 deletions
diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/tests/__init__.py diff --git a/tests/fixtures_synth.py b/tests/fixtures_synth.py new file mode 100644 index 0000000..69b305b --- /dev/null +++ b/tests/fixtures_synth.py @@ -0,0 +1,39 @@ +"""Deterministic synthetic OHLCV fixtures for worker tests. + +Every object here carries ``_synthetic: true`` and must never be used as a +production fallback. Values are simple deterministic series so accounting + assertions in backtest tests can be hand-verified. +""" +import pandas as pd + +_SYNTHETIC = {"_synthetic": True} + + +def synthetic_daily(symbol: str = "SH#600000", start="2024-01-02", days=20, + base=10.0, volume=1_000_000, extra_fields=True) -> pd.DataFrame: + dates = pd.bdate_range(start, periods=days) + rows = [] + price = base + for i, d in enumerate(dates): + o = round(price, 2) + c = round(price * 1.02, 2) if i % 2 == 0 else round(price * 0.98, 2) + h = round(max(o, c) * 1.01, 2) + l = round(min(o, c) * 0.99, 2) + row = { + "date": d.strftime("%Y-%m-%d"), + "symbol": symbol, + "open": o, "high": h, "low": l, "close": c, + "volume": volume + i * 1000, + } + if extra_fields: + # raw provider-style extra fields preserved verbatim + row["amount"] = round((o + h + l + c) / 4 * (volume + i * 1000), 2) + row["turnover"] = round(0.5 + i * 0.01, 3) + rows.append(row) + price = c + df = pd.DataFrame(rows) + df.attrs["synthetic"] = True + return df + + +SYNTH = _SYNTHETIC diff --git a/tests/worker/test_backtest.py b/tests/worker/test_backtest.py new file mode 100644 index 0000000..30f8e46 --- /dev/null +++ b/tests/worker/test_backtest.py @@ -0,0 +1,270 @@ +"""Backtrader runner tests — deterministic synthetic accounting (_synthetic). + +Hand-computable expectations: price series closes 10.20, 10.00, 10.40, 10.20... +commission 0.0003, slippage 0.001 of price, next-bar-open fills. +""" +import json +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from fixtures_synth import synthetic_daily, SYNTH # noqa: E402 +from worker.backtest import run_backtest # noqa: E402 + +STRAT = ''' +import backtrader as bt + +class Strategy(bt.Strategy): + params = (("size", 100),) + + def __init__(self): + self.bought = False + + def next(self): + if not self.bought and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=self.p.size) + self.bought = True +''' + +STRAT_NAMED = ''' +import backtrader as bt + +class Strategy(bt.Strategy): + def __init__(self): + self.a = self.getdatabyname("SH#600000") + self.b = self.getdatabyname("BJ#TD001") + self.done = False + + def next(self): + if not self.done and len(self) > 2: + self.buy(data=self.a, size=100) + self.done = True +''' + + +def write_dataset(tmp_path): + d = tmp_path / "data" + d.mkdir() + dfa = synthetic_daily("SH#600000", base=10.0, extra_fields=False) + dfb = synthetic_daily("BJ#TD001", base=5.0, extra_fields=False) + for name, df in (("SH#600000", dfa), ("BJ#TD001", dfb)): + p = d / f"{name.lower().replace('#','_')}.csv" + df.to_csv(p, index=False) + manifest = { + "id": "00000000-0000-0000-0000-000000000000", + "hash": "fix-hash", + "_synthetic": True, + "objects": [ + {"instrument": {"symbol": "SH#600000", "market": "cn", "asset_type": "stock"}, + "path": f"{name}" , "object_hash": "h", "row_count": 20, "columns": ["date"]}, + ], + "warnings": [], + } + manifest["objects"] = [ + {"instrument": {"symbol": "SH#600000", "market": "cn", "asset_type": "stock"}, + "path": "sh_600000.csv", "object_hash": "h-a", "row_count": 20}, + {"instrument": {"symbol": "BJ#TD001", "market": "cn", "asset_type": "stock"}, + "path": "bj_td001.csv", "object_hash": "h-b", "row_count": 20}, + ] + return d, manifest + + +def make_request(tmp_path, code, benchmark=None, params=None): + d, manifest = write_dataset(tmp_path) + return { + "_synthetic": True, + "code": code, + "config": { + "capital": 100000.0, "commission": 0.0003, "slippage": 0.001, + "benchmark_symbol": benchmark, + "parameters": params or {}, + }, + "dataset_manifest": manifest, + "data_root": str(d), + }, d + + +def test_synthetic_labeled_runs_next_bar_fill(tmp_path): + req, _ = make_request(tmp_path, STRAT) + res = run_backtest(req) + assert res["engine"]["name"] == "backtrader" + assert res["_synthetic"] is True + assert res["trades"], "expected at least one trade" + t = res["trades"][0] + # signal bar index 3; fills at next-bar open (10.20 * 1.001 slippage), NOT signal-bar close + assert t["symbol"] == "SH#600000" + assert t["side"] == "buy" + assert t["quantity"] == 100 + assert abs(t["price"] - 10.20 * 1.001) < 1e-9 + assert abs(t["commission"] - t["value"] * 0.0003) < 1e-6 + + +def test_equity_cash_accounting(tmp_path): + req, _ = make_request(tmp_path, STRAT) + res = run_backtest(req) + eq = res["equity"] + assert len(eq) == 20 + first = eq[0] + assert first["cash"] == 100000.0 + assert first["equity"] == 100000.0 + last = eq[-1] + filled = 100 * 10.40 * 1.001 + comm = filled * 0.0003 + # after buy: cash reduced; equity = cash + 100 * final close (10.20? compute from fixture) + assert last["cash"] < 100000.0 + expected_equity = last["cash"] + 100 * last["closes"]["SH#600000"] + assert abs(last["equity"] - expected_equity) < 1e-6 + + +def test_metrics_no_nan_nulls_and_fields(tmp_path): + req, _ = make_request(tmp_path, STRAT) + res = run_backtest(req) + m = res["metrics"] + for key in ("total_return", "annual_return", "max_drawdown", "trade_count", "final_equity"): + assert key in m + assert m[key] is None or isinstance(m[key], (int, float)) + if isinstance(m[key], float): + assert not (m[key] != m[key] or m[key] in (float("inf"), float("-inf"))) + # max drawdown is a non-positive fraction or null + assert m["max_drawdown"] is None or m["max_drawdown"] <= 0 + assert m["trade_count"] >= 0 # closed round-trips; buy-alone runs have 0 + assert res["data_manifest_hash"] == "fix-hash" + assert res["elapsed_ms"] >= 0 + assert isinstance(res["logs"], list) and res["logs"] + assert len(res["equity"]) == 20 + + +def test_named_feeds_present_no_cross_lookahead(tmp_path): + req, _ = make_request(tmp_path, STRAT_NAMED) + res = run_backtest(req) + assert res["trades"] + assert res["trades"][0]["symbol"] == "SH#600000" + # the second feed was never consumed for first-symbol pricing + assert all(t["symbol"] != "BJ#TD001" for t in res["trades"]) + + +def test_benchmark_series_included(tmp_path): + req, _ = make_request(tmp_path, STRAT, benchmark="BJ#TD001") + res = run_backtest(req) + assert all("benchmark" in e and e["benchmark"] is not None for e in res["equity"]) + + +def test_missing_strategy_class_fails_clean(tmp_path): + req, _ = make_request(tmp_path, "import backtrader as bt\n\nclass Foo(bt.Strategy):\n pass\n") + res = run_backtest(req) + assert res["status"] == "failed" + assert "Strategy" in res["error"]["message"] + + +def test_syntax_error_fails_clean(tmp_path): + req, _ = make_request(tmp_path, "def broken(:\n") + res = run_backtest(req) + assert res["status"] == "failed" + assert res["error"]["code"] == "strategy_syntax" + + +def test_missing_data_object_fails(tmp_path): + req, _ = make_request(tmp_path, STRAT) + req["dataset_manifest"]["objects"][0]["path"] = "nope.csv" + res = run_backtest(req) + assert res["status"] == "failed" + assert res["error"]["code"] == "data_missing" + + +def test_lookahead_signal_uses_prior_close_not_same_day(tmp_path): + # Strategy trades only on the last bar; a legal next-bar fill must not exist yet. + strat = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def next(self): + if len(self) == 20: + self.buy(data=self.getdatabyname("SH#600000"), size=100) +''' + req, _ = make_request(tmp_path, strat) + res = run_backtest(req) + # order placed on final bar; notification/fill cannot occur after data end -> no trade + assert not res["trades"] + + +def test_strategy_module_imports_visible_in_methods(tmp_path): + # Regression: module-level imports must remain visible inside __init__ and + # next. Transport-level isolation is Docker, NOT restricted Python globals. + strat = ''' +import os +import math +import backtrader as bt + +class Strategy(bt.Strategy): + def __init__(self): + self.foo = os.sep # os imported at module level, used in a method + + def next(self): + c = self.getdatabyname("SH#600000").close + if len(self) > 2 and abs(math.copysign(1.0, c[0] - c[-1])) == 1.0: + self.foo = math.sqrt(abs(c[0])) +''' + req, _ = make_request(tmp_path, strat) + res = run_backtest(req) + assert res["status"] == "succeeded", res.get("error") + + +def test_strategy_genuine_nameerror_still_fails(tmp_path): + # a genuinely undefined name must still surface honestly as runtime_error + strat = ''' +import backtrader as bt + +class Strategy(bt.Strategy): + def next(self): + undefined_variable_xyz.bar() +''' + req, _ = make_request(tmp_path, strat) + res = run_backtest(req) + assert res["status"] == "failed" + assert res["error"]["code"] == "runtime_error" + + + +def test_fill_value_is_executed_turnover_not_cost_basis(tmp_path): + """RED/GREEN accounting regression: fills must record actual turnover + abs(ex.size * ex.price). Backtrader's ex.value for SELL orders reports the + position cost basis, NOT sale proceeds (real QA showed identical 659.66 for + a real buy and sell at different prices). Equity/cash are not affected: + broker cash and equity already use executed price and commission.""" + strat = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): + self.done = False + def next(self): + if len(self) == 3: + self.buy(data=self.getdatabyname("SH#600000"), size=100) + elif len(self) == 10: + self.sell(data=self.getdatabyname("SH#600000"), size=100) + self.done = True +''' + req, _ = make_request(tmp_path, strat) + res = run_backtest(req) + assert res["status"] == "succeeded", res.get("error") + sides = [(t["side"], t) for t in res["trades"]] + buys = [t for s, t in sides if s == "buy"] + sells = [t for s, t in sides if s == "sell"] + assert buys and sells, f"expected a buy AND a sell fill, got {sides}" + # discount-adjusted commission is charged on the turnover, not on cost basis + for t in res["trades"]: + assert abs(t["value"] - abs(t["quantity"] * t["price"])) < 1e-9, \ + f"fill value must be quantity*price turnover: {t}" + assert abs(t["commission"] - t["value"] * 0.0003) < 1e-6, \ + f"commission follows turnover: {t}" + # honest check: buy and sell execute at different prices so their turnover + # differs (cost-basis bug reported identical values for both sides) + assert buys and sells + assert abs(buys[0]["price"] - sells[0]["price"]) > 1e-9, \ + f"buy/sell fill prices must differ: {buys[0]['price']} vs {sells[0]['price']}" + assert abs(buys[0]["value"] - sells[0]["value"]) > 1e-9, \ + "distinct prices must yield distinct fill values (cost-basis bug regression)" + # trade_count remains closed round trips, not fills + assert res["metrics"]["trade_count"] == 1 diff --git a/tests/worker/test_data.py b/tests/worker/test_data.py new file mode 100644 index 0000000..d6eeae1 --- /dev/null +++ b/tests/worker/test_data.py @@ -0,0 +1,336 @@ +"""Data adapter / instrument-search tests — mocked provider transport, no network.""" +import json +import sys +from pathlib import Path + +import pandas as pd +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +import worker.data as data # noqa: E402 +from fixtures_synth import synthetic_daily, SYNTH # noqa: E402 + + +# ---- live-captured provider reply shapes (see docs/search-fix.md probes) ---- + +def _em_payload(rows): + return {"QuotationCodeTable": {"Data": rows, "Status": 0, "Message": "成功", + "TotalCount": len(rows)}, + } + + +def _em_row(code, name, classify, mktnum, sec_type_name): + return {"Code": code, "Name": name, "Classify": classify, "MktNum": mktnum, + "SecurityTypeName": sec_type_name, "MarketType": mktnum, + "QuoteID": f"{mktnum}.{code}"} + + +def test_search_stock_real_shape(monkeypatch): + """q=600000 → real-shape eastmoney AStock row classifies as stock, identity validated.""" + body = _em_payload([_em_row("600000", "浦发银行", "AStock", "1", "沪A")]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: body) + + def boom(*a, **k): + raise RuntimeError("tencent not called in this test") + monkeypatch.setattr(data, "_http_text", boom) + res = data.search_instruments("600000") + assert res["status"] == "ready" + item = res["items"][0] + assert item["symbol"] == "600000" and item["name"] == "浦发银行" + assert item["asset_type"] == "stock" and item["market"] == "cn" + assert item["canonical_symbol"] == "SH#600000" + assert item["source"] == "eastmoney" + + +def test_search_etf_real_shape_with_tencent_confirmation(monkeypatch): + """q=510300: eastmoney Classify=Fund only counts as ETF when tencent confirms.""" + em = _em_payload([_em_row("510300", "沪深300ETF华泰柏瑞", "Fund", "1", "基金")]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + tx = 'v_hint="sh~510300~\\u6caa\\u6df1300ETF\\u534e\\u6cf0\\u67cf\\u745e~hs300etfhtbr~ETF"' + monkeypatch.setattr(data, "_http_text", lambda *a, **k: tx) + res = data.search_instruments("510300") + etfs = [i for i in res["items"] if i["symbol"] == "510300"] + assert etfs and all(i["asset_type"] == "etf" for i in etfs) + assert {i["source"] for i in etfs} == {"eastmoney", "tencent"} + assert any(i["canonical_symbol"] == "SH#510300" for i in etfs) + + +def test_search_etf_ambiguous_fund_excluded_without_tencent_confirmation(monkeypatch): + """Eastmoney Classify=Fund alone cannot distinguish ETF from LOF (probed): + without a tencent ETF confirmation the item must not surface, ever.""" + em = _em_payload([_em_row("160706", "沪深300LOF", "Fund", "0", "基金")]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + monkeypatch.setattr( + data, "_http_text", + lambda *a, **k: 'v_hint="sz~160706~\\u6caa\\u6df1300LOF~hs300lof~LOF"') + res = data.search_instruments("160706") + assert res["status"] == "ready" + assert res["items"] == [] # LOF must never surface as an etf + + +def test_search_lof_not_present_when_eastmoney_only(monkeypatch): + em = _em_payload([_em_row("160706", "沪深300LOF", "Fund", "0", "基金")]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + + def boom(*a, **k): + raise RuntimeError("tencent down") + monkeypatch.setattr(data, "_http_text", boom) + res = data.search_instruments("160706") + # tencent failed but eastmoney succeeded → still ready; LOF suppressed + assert res["status"] == "ready" + assert res["items"] == [] + assert res["providers"]["tencent"].startswith("failed") + + +def test_search_index_from_chinese_query_and_class_not_code(monkeypatch): + """q=沪深300 → indexes come from provider class (Classify=Index / ZS tag); + a numeric code like 000300 is NEVER guessed to be an index by shape.""" + em = _em_payload([ + _em_row("000300", "沪深300", "Index", "1", "指数"), + _em_row("399300", "沪深300", "Index", "0", "指数"), + _em_row("510300", "沪深300ETF华泰柏瑞", "Fund", "1", "基金"), + ]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + tx = ('v_hint="sh~000300~\\u6caa\\u6df1300~hs300~ZS^sz~399300~\\u6caa\\u6df1300~hs300~ZS' + '^sh~510300~\\u6caa\\u6df1300ETF\\u534e\\u6cf0\\u67cf\\u745e~hs300etfhtbr~ETF^' + 'sz~160706~\\u6caa\\u6df1300LOF~hs300lof~LOF"') + monkeypatch.setattr(data, "_http_text", lambda *a, **k: tx) + res = data.search_instruments("沪深300") + by_key = {(i["canonical_symbol"], i["asset_type"], i["source"]) for i in res["items"]} + assert ("SH#000300", "index", "eastmoney") in by_key + assert ("SZ#399300", "index", "eastmoney") in by_key + assert ("SH#000300", "index", "tencent") in by_key + assert all(i["asset_type"] in ("stock", "etf", "index") for i in res["items"]) + assert not any("160706" == i["symbol"] and i["asset_type"] == "etf" for i in res["items"]) + + +def test_search_numeric_code_not_auto_index(monkeypatch): + """000300 with a provider stock-ish class (hypothetical) must be honored as + whatever the provider states — classification is never inferred.""" + em = _em_payload([_em_row("000300", "某标的", "AStock", "0", "深A")]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + monkeypatch.setattr( + data, "_http_text", + lambda *a, **k: 'v_hint="sz~000300~\\u67d0\\u6807\\u7684~m?~GP-A"') + res = data.search_instruments("000300") + assert all(i["asset_type"] == "stock" for i in res["items"]) + + +def test_search_excludes_unsupported_classes_real_rows(monkeypatch): + """Bond/HK/OTC/LOF rows and tencent LOF tags are dropped, no guessing.""" + em = _em_payload([ + _em_row("160706", "20天津38", "Bond", "1", "债券"), + _em_row("00700", "腾讯控股", "HK", "116", "港股"), + _em_row("160706", "嘉实沪深300ETF联接A", "OTCFUND", "0", "场外基金"), + _em_row("600000", "浦发银行", "AStock", "1", "沪A"), + ]) + monkeypatch.setattr(data, "_http_json", lambda *a, **k: em) + monkeypatch.setattr( + data, "_http_text", + lambda *a, **k: 'v_hint="sz~160706~\\u6caa\\u6df1300LOF~hs300lof~LOF"') + res = data.search_instruments("x") + assert res["status"] == "ready" + symbols = [i["symbol"] for i in res["items"]] + assert symbols == ["600000"] # bond/HK/OTCFUND dropped; tencent LOF tag excluded and eastmoney Fund lacks ETF confirmation + assert res["items"][0]["asset_type"] == "stock" + + +def test_search_chinese_name_query_matches_provider(monkeypatch): + em = _em_payload([_em_row("600000", "浦发银行", "AStock", "1", "沪A")]) + captured = {} + monkeypatch.setattr(data, "_http_json", lambda url, params, **k: captured.update(params) or em) + + def boom(*a, **k): + raise RuntimeError("tencent down") + monkeypatch.setattr(data, "_http_text", boom) + res = data.search_instruments("浦发银行") + assert captured["input"] == "浦发银行" # provider receives the raw CJK query + assert res["items"][0]["name"] == "浦发银行" + + +def test_search_provider_failure_is_explicit_failed_status(monkeypatch): + def boom(*a, **k): + raise RuntimeError("network down") + monkeypatch.setattr(data, "_http_json", boom) + monkeypatch.setattr(data, "_http_text", boom) + res = data.search_instruments("600000") + assert res["status"] == "failed" + assert res["items"] == [] + assert res["error"]["code"] == "providers_unavailable" + assert "eastmoney" in res["error"]["message"] + assert "tencent" in res["error"]["message"] + + +def test_search_provider_timeout_is_explicit_failed_status(monkeypatch): + import requests as _req + def timeout(*a, **k): + raise _req.exceptions.Timeout("read timed out") + monkeypatch.setattr(data, "_http_json", timeout) + monkeypatch.setattr(data, "_http_text", timeout) + res = data.search_instruments("600000") + assert res["status"] == "failed" + assert "Timeout" in res["error"]["message"] + + +def test_search_partial_failure_is_ready_with_warning(monkeypatch): + def boom(*a, **k): + raise RuntimeError("tencent down") + monkeypatch.setattr( + data, "_http_json", + lambda *a, **k: _em_payload([_em_row("600000", "浦发银行", "AStock", "1", "沪A")])) + monkeypatch.setattr(data, "_http_text", boom) + res = data.search_instruments("600000") + assert res["status"] == "ready" + assert res["providers"]["tencent"].startswith("failed") + assert res["items"][0]["symbol"] == "600000" + + +def test_suggest_limits_are_bounded(): + import inspect + src = inspect.getsource(data) + assert "SUGGEST_TIMEOUT_SECS = 4.0" in src # per-source timeout applies to the actual calls + # the catalog paths are gone: every query must be a bounded provider call + assert "_stock_catalog" not in src + assert "_catalog_frame" not in src + + +def test_limit_is_passed_to_provider_bounded(monkeypatch): + seen = {} + rows = [_em_row(str(600000 + i), f"n{i}", "AStock", "1", "沪A") for i in range(8)] + monkeypatch.setattr(data, "_http_json", + lambda url, params, **k: seen.update(params) or _em_payload(rows[:params["count"]])) + monkeypatch.setattr(data, "_http_text", lambda *a, **k: "") + res = data.search_instruments("600", limit=4) + assert seen["count"] == 5 # bounded floor 5, well under the old full catalogs + assert len(res["items"]) <= 4 + + +def test_fetch_rejects_unknown_asset_type(): + with pytest.raises(data.DataError): + data.fetch_source({"symbol": "X#1", "market": "cn", "asset_type": "crypto"}, + "2024-01-01", "2024-02-01", "daily", "none", ["open", "close"]) + + +# ---- sina listed-ETF adapter (recovery-01: eastmoney down, tencent has no +# ETF adapter; live evidence in docs/recovery-01-plan.md). The stock-like frame +# below is a SYNTHETIC test fixture mimicking the live-captured sina reply shape. ---- + +def _sina_synth_frame(): + """Synthetic fixture in the exact live-captured shape of + ak.fund_etf_hist_sina(symbol='sz159399') (recovered live 2026-09-17: + 381 rows 2025-02-27..2026-09-16, columns date/open/.../amount/postVol/postAmt).""" + rows = [ + {"date": "2025-12-30", "open": 1.001, "high": 1.006, "low": 0.999, + "close": 1.006, "volume": 695982432, "amount": 696910071.0, + "postVol": float("nan"), "postAmt": float("nan")}, + {"date": "2026-01-02", "open": 1.007, "high": 1.008, "low": 0.994, + "close": 0.997, "volume": 505082425, "amount": 506417752.0, + "postVol": float("nan"), "postAmt": float("nan")}, + {"date": "2026-09-16", "open": 1.001, "high": 1.001, "low": 0.982, + "close": 0.992, "volume": 144526400, "amount": 142914764.0, + "postVol": 300.0, "postAmt": 298.0}, + ] + df = pd.DataFrame(rows) + df.attrs["synthetic"] = True + return df + + +def test_sina_etf_serves_159399_with_honest_provenance(monkeypatch): + inst = {"symbol": "159399", "market": "cn", "asset_type": "etf"} + monkeypatch.setattr(data, "ak", + type("M", (), {"fund_etf_hist_sina": + staticmethod(lambda **kw: _sina_synth_frame())})()) + res = data.fetch_source(inst, "2025-12-31", "2026-09-17", "daily", "none") + df, endpoint, params = res + assert res.provider == "sina" + assert endpoint == "fund_etf_hist_sina" + assert params == {"symbol": "sz159399"} + assert set(df["symbol"]) == {"SZ#159399"} + dates = df["date"].tolist() + assert dates[0] >= "2025-12-31" and dates[-1] <= "2026-09-17" + assert "2026-09-17" not in dates # no fabricated end-date bar + warns = " ".join(df.attrs.get("source_warnings", [])) + assert "股" in warns and "unadjusted" in warns + assert "159399" not in warns # numbers only, identity preserved + + +def test_sina_rejects_adjustment_and_range_semantics(monkeypatch): + inst = {"symbol": "159399", "market": "cn", "asset_type": "etf"} + monkeypatch.setattr(data, "ak", + type("M", (), {"fund_etf_hist_sina": + staticmethod(lambda **kw: _sina_synth_frame())})()) + with pytest.raises(data.DataError) as e: + data.fetch_source(inst, "2025-12-31", "2026-09-17", "daily", "qfq", source="sina") + assert e.value.code == "unsupported_adjustment" + + +def test_sina_rejects_stock_and_index(): + for asset in ("stock", "index"): + inst = {"symbol": "600000" if asset == "stock" else "000300", + "market": "cn", "asset_type": asset} + with pytest.raises(data.DataError) as e: + data.fetch_source(inst, "2025-12-31", "2026-09-17", "daily", "none", + source="sina") + assert e.value.code == "source_unavailable" + + +def test_auto_etf_falls_back_from_eastmoney_to_sina_honestly(monkeypatch): + inst = {"symbol": "159399", "market": "cn", "asset_type": "etf"} + + def dis(*a, **k): + raise ConnectionError("RemoteDisconnected('Remote end closed connection')") + monkeypatch.setattr(data, "ak", + type("M", (), {"fund_etf_hist_em": staticmethod(dis), + "fund_etf_hist_sina": + staticmethod(lambda **kw: _sina_synth_frame())})()) + res, warns = data.fetch_source_with_warnings( + inst, "2025-12-31", "2026-09-17", "daily", "none", source="auto") + assert res.provider == "sina" # actual serving provider, never disguised + joined = " ".join(warns) + assert "provider_fallback" in joined and "eastmoney" in joined and "sina" in joined + assert set(res.df["symbol"]) == {"SZ#159399"} + + +def test_auto_etf_fails_when_no_provider_is_viable(monkeypatch): + inst = {"symbol": "159399", "market": "cn", "asset_type": "etf"} + + def dis(*a, **k): + raise ConnectionError("RemoteDisconnected") + monkeypatch.setattr(data, "ak", + type("M", (), {"fund_etf_hist_em": staticmethod(dis), + "fund_etf_hist_sina": staticmethod(dis)})()) + with pytest.raises(data.DataError) as e: + data.fetch_source(inst, "2025-12-31", "2026-09-17", "daily", "none", source="auto") + assert e.value.code == "provider_unavailable" + assert "eastmoney" in str(e.value) and "sina" in str(e.value) + + +def test_tencent_etf_rejection_message_unchanged(): + inst = {"symbol": "159399", "market": "cn", "asset_type": "etf"} + with pytest.raises(data.DataError) as e: + data.fetch_source(inst, "2025-12-31", "2026-09-17", "daily", "none", + source="tencent") + assert e.value.code == "source_unavailable" + assert "tencent source has no listed-ETF daily adapter" in str(e.value) + + +def test_sina_registered_in_explicit_source_surface(): + assert "sina" in data.SUPPORTED_SOURCES + + +def test_fetch_columns_subset_requested(): + req = {"symbol": "600000", "market": "cn", "asset_type": "stock"} + called = {} + + def fake_hist(symbol, **kw): + called["params"] = kw + return synthetic_daily("SH#600000") + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setattr(data, "ak", type("M", (), {"stock_zh_a_hist": staticmethod(fake_hist)})()) + df, endpoint, params = data.fetch_source(req, "2024-01-01", "2024-02-01", + "daily", "none", ["open", "close", "volume"]) + assert endpoint == "stock_zh_a_hist" + assert all(c in df.columns for c in ("date", "symbol", "open", "close", "volume")) + monkeypatch.undo() diff --git a/tests/worker/test_fetch_cli.py b/tests/worker/test_fetch_cli.py new file mode 100644 index 0000000..11ddd11 --- /dev/null +++ b/tests/worker/test_fetch_cli.py @@ -0,0 +1,144 @@ +"""Worker CLI fetch tests: fallback path, source labeling, raw retention — no network.""" +import json +import sys +from pathlib import Path + +import pandas as pd +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from worker import data as data_mod, main as main_mod # noqa: E402 +from fixtures_synth import synthetic_daily # noqa: E402 + + +def _fake_tx_frame(): + df = pd.DataFrame({ + "date": ["2024-06-27", "2024-06-28"], + "open": [8.21, 8.22], "close": [8.22, 8.23], "high": [8.24, 8.25], + "low": [8.20, 8.21], "volume": [21000000, 22066700], + }) + return df + + +def test_auto_falls_back_to_tencent_with_honest_warning(monkeypatch, tmp_path): + def em_fail(*a, **k): + raise data_mod.DataError("provider_unavailable", "another connect error") + monkeypatch.setattr(data_mod, "_fetch_eastmoney", em_fail) + + def tx_ok(instrument, start, end, frequency, adjustment): + df = _fake_tx_frame().assign(symbol="SH#600000") + return data_mod.FetchResult(df, "stock_zh_a_hist_tx", + {"symbol": "sh600000"}, df, "tencent") + monkeypatch.setattr(data_mod, "_fetch_tencent", tx_ok) + res, warnings = data_mod.fetch_source_with_warnings( + {"symbol": "600000", "market": "cn", "asset_type": "stock"}, + "2024-01-01", "2024-06-30", "daily", "none") + assert res.provider == "tencent" + assert res[2]["symbol"] == "sh600000" # SAME symbol, no substitution + assert any("provider_fallback" in x and "eastmoney" in x for x in warnings) + + +def test_exact_source_error_not_fallback(monkeypatch): + monkeypatch.setattr(data_mod, "_fetch_eastmoney", + lambda *a, **k: (_ for _ in ()).throw(RuntimeError("timed out"))) + with pytest.raises(RuntimeError): + data_mod.fetch_source({"symbol": "600000", "market": "cn", + "asset_type": "stock"}, "2024-01-01", "2024-02-01", + "daily", "none", source="eastmoney") + + +def test_tencent_rejects_etf(monkeypatch): + with pytest.raises(data_mod.DataError): + data_mod.fetch_source({"symbol": "510300", "market": "cn", "asset_type": "etf"}, + "2024-01-01", "2024-06-30", "daily", "none", source="tencent") + + +def test_bad_source_rejected(): + with pytest.raises(data_mod.DataError): + data_mod.fetch_source({"symbol": "600000", "market": "cn", "asset_type": "stock"}, + "2024-01-01", "2024-02-01", "daily", "none", source="sina") + + +def test_fetch_cli_records_provider_and_raw(tmp_path, monkeypatch): + def fake_fetch(inst, start, end, frequency, adjustment, fields=None, source="auto"): + df = _fake_tx_frame().assign(symbol="SH#600000", 成交额=1.0) + import akshare as _ak + return (data_mod.FetchResult(df, "stock_zh_a_hist_tx", {"symbol": "sh600000"}, + df, "tencent"), + ["provider_fallback: eastmoney attempt failed; served by tencent"]) + + monkeypatch.setattr(data_mod, "fetch_source_with_warnings", fake_fetch) + req = tmp_path / "request.json" + req.write_text(json.dumps({ + "instruments": [{"symbol": "600000", "market": "cn", "asset_type": "stock", + "name": "浦发银行"}], + "start_date": "2024-06-01", "end_date": "2024-06-30", + "frequency": "daily", "adjustment": "none", + "fields": ["open", "high", "low", "close", "volume", "amount"], + "source": "auto"})) + out = tmp_path / "output" + rc = main_mod.run_fetch(str(req), str(out)) + assert rc == 0 + result = json.loads((out / "result.json").read_text()) + assert result["status"] == "ready" + o = result["manifest"]["objects"][0] + # provider is the ACTUAL serving source in manifest + cache identity + assert o["provider"] == "tencent" + assert o["endpoint"] == "stock_zh_a_hist_tx" + assert any("provider_fallback" in w for w in result["warnings"] + o["warnings"]) + # raw provider response retained BEFORE normalized transform, as an + # immutable content-addressed object next to the normalized CSV + raw_name = f"raw_600000_{o['raw_object_hash'][:12]}.json" + raw = json.loads((out / "objects" / raw_name).read_text()) + assert raw["endpoint"] == "stock_zh_a_hist_tx" and raw["data"] + assert len(o["raw_object_hash"]) == 64 + # normalized object separate from raw + norm = pd.read_csv(out / o["path"]) + assert {"open", "high", "low", "close", "volume"} <= set(norm.columns) + assert o["object_hash"] != o["raw_object_hash"] + + +def test_fetch_cli_error_path_writes_result(tmp_path): + req = tmp_path / "request.json" + req.write_text(json.dumps({"instruments": [], "frequency": "daily", + "adjustment": "none", "start_date": "2024-01-01", + "end_date": "2024-02-01"})) + out = tmp_path / "output" + rc = main_mod.run_fetch(str(req), str(out)) + assert rc == 1 + r = json.loads((out / "result.json").read_text()) + assert r["status"] == "failed" and r["errors"] + + +def test_search_cli_envelope_contract(monkeypatch): + """worker.main search prints an explicit JSON envelope; failed status is a + real failure (non-zero exit), never an empty success.""" + def boom(q, limit): + return {"source": "provider_suggest", "status": "failed", "items": [], + "error": {"code": "providers_unavailable", "message": "all providers failed"}} + monkeypatch.setattr(data_mod, "search_instruments", boom) + assert main_mod.run_search("600000", 50) == 1 + + envelope_out: dict = {} + + class _C: + def write(self, s): + envelope_out["buf"] = envelope_out.get("buf", "") + s + import sys as _sys + orig = _sys.stdout + _sys.stdout = _C() # type: ignore[assignment] + try: + monkeypatch.setattr(data_mod, "search_instruments", lambda q, limit: { + "source": "provider_suggest", "status": "ready", + "items": [{"symbol": "600000", "name": "浦发银行", "asset_type": "stock", + "market": "cn", "canonical_symbol": "SH#600000", + "currency": "CNY", "source": "eastmoney"}], + "error": None}) + assert main_mod.run_search("600000", 50) == 0 + finally: + _sys.stdout = orig + env = json.loads(envelope_out["buf"]) + assert env["status"] == "ready" + assert env["items"][0]["source"] == "eastmoney" # per-item provenance preserved diff --git a/tests/worker/test_normalize.py b/tests/worker/test_normalize.py new file mode 100644 index 0000000..a4c9a38 --- /dev/null +++ b/tests/worker/test_normalize.py @@ -0,0 +1,92 @@ +"""Normalization tests — synthetic fixtures only (_synthetic).""" +import math +import sys +import types +from pathlib import Path + +import pandas as pd +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from fixtures_synth import synthetic_daily, SYNTH # noqa: E402 +from worker.normalize import normalize_frame, build_manifest_entry, \ + compute_object_hash, coverage_summary, NormalizeError # noqa: E402 + + +def test_synthetic_marker_present(): + assert SYNTH["_synthetic"] is True + df = synthetic_daily() + assert df.attrs["synthetic"] is True + + +def test_normalize_keeps_raw_extra_fields_and_precision(): + df = synthetic_daily() + out = normalize_frame(df, symbol="SH#600000") + assert "_synthetic" not in out.columns + # canonical columns present + for col in ("date", "open", "high", "low", "close", "volume"): + assert col in out.columns + # raw extra fields preserved + assert "amount" in out.columns and "turnover" in out.columns + # precision kept, no rounding drift + assert out["close"].iloc[0] == df["close"].iloc[0] + assert out["volume"].iloc[3] == df["volume"].iloc[3] + # symbol never substituted across instruments + assert out["symbol"].iloc[5] == "SH#600000" + assert out["date"].notna().all() + + +def test_normalize_rejects_missing_ohlcv(): + df = synthetic_daily() + df.loc[4, "close"] = math.nan + with pytest.raises(NormalizeError): + normalize_frame(df, symbol="SH#600000") + + +def test_normalize_rejects_bad_dates(): + df = synthetic_daily() + df["date"] = df["date"].astype(object) + df.loc[2, "date"] = "not-a-date" + with pytest.raises(NormalizeError): + normalize_frame(df, symbol="SH#600000") + + +def test_normalize_sorts_and_dedups_dates(): + df = synthetic_daily() + df = pd.concat([df.iloc[5:], df.iloc[:5]]).reset_index(drop=True) + out = normalize_frame(df, symbol="SH#600000") + assert list(out["date"]) == sorted(out["date"]) + assert len(out) == 20 + + +def test_object_hash_stable_and_content_sensitive(): + df = synthetic_daily() + h1 = compute_object_hash(df) + h2 = compute_object_hash(normalize_frame(df, symbol="SH#600000")) + assert h1 == h2 and len(h1) == 64 + df2 = synthetic_daily(base=11.0) + assert compute_object_hash(df2) != h1 + + +def test_coverage_summary_and_manifest_entry(): + df = synthetic_daily(days=10, extra_fields=False) + cov = coverage_summary(df) + assert cov["actual_start"] == df["date"].iloc[0] + entry = build_manifest_entry( + df, instrument={"symbol": "SH#600000", "market": "cn", "asset_type": "stock"}, + provider="akshare", endpoint="stock_zh_a_hist", params={"period": "daily"}, + adjustment="none", requested_start="2024-01-01", requested_end="2024-12-31", + raw_object_hash="rawhash", warnings=[], + schema_version="1", normalization_version="1", + path="objects/ab/cd.json", fetched_at="2026-09-16T00:00:00Z", + akshare_version="1.18.95", + ) + assert entry["immutable"] is True + assert entry["row_count"] == 10 + assert entry["adjustment"] == "none" + assert entry["actual_start"] <= entry["actual_end"] + assert "potential symbol substitution" not in entry["warnings"] + for internal in ("path",): + assert internal in entry # server strips internal-only fields downstream diff --git a/tests/worker/test_rules_regression.py b/tests/worker/test_rules_regression.py new file mode 100644 index 0000000..d59ca0a --- /dev/null +++ b/tests/worker/test_rules_regression.py @@ -0,0 +1,201 @@ +"""Regression tests: broker rules, benchmark normalization, drawdown zero.""" +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from fixtures_synth import synthetic_daily # noqa: E402 +from worker.backtest import run_backtest # noqa: E402 + + +def build_manifest(tmp_path, specs): + d = tmp_path / "data" + d.mkdir(exist_ok=True) + objects = [] + for sym, code, asset, base in specs: + df = synthetic_daily(sym, base=base, extra_fields=False) # code param unused + slug = sym.lower().replace("#", "_") + path = f"{slug}.csv" + df.to_csv(d / path, index=False) + objects.append({"instrument": {"symbol": sym, "market": "cn", "asset_type": asset}, + "path": path, "row_count": len(df)}) + manifest = {"hash": "fix-hash", "_synthetic": True, "objects": objects, "warnings": []} + return d, manifest + + +def make_request(tmp_path, code, specs, benchmark=None, params=None): + d, manifest = build_manifest(tmp_path, specs) + return { + "_synthetic": True, + "code": code, + "config": {"capital": 100000.0, "commission": 0.0003, "slippage": 0.001, + "benchmark_symbol": benchmark, "parameters": params or {}}, + "dataset_manifest": manifest, + "data_root": str(d), + } + + +BUY_STOCK_ONLY = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.done = False + def next(self): + if not self.done and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=250) + self.done = True +''' + +BUY_INDEX = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.done = False + def next(self): + if not self.done and len(self) > 2: + self.buy(data=self.getdatabyname("SH000300"), size=100) + self.done = True +''' + +BUY_BM = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.done = False + def next(self): + if not self.done and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=100) + self.done = True +''' + +SELL_OPEN = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.state = 0 + def next(self): + if self.state == 0 and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=200) + self.state = 1 + elif self.state == 1: + self.sell(data=self.getdatabyname("SH#600000"), size=150) + self.state = 2 +''' + +SHORT = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.state = 0 + def next(self): + if self.state == 0 and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=100) + self.state = 1 + elif self.state == 1: + self.sell(data=self.getdatabyname("SH#600000"), size=400) + self.state = 2 +''' + +BUY250_TEST = ''' +import backtrader as bt +class Strategy(bt.Strategy): + def __init__(self): self.done = False + def next(self): + if not self.done and len(self) > 2: + self.buy(data=self.getdatabyname("SH#600000"), size=250) + self.done = True +''' + + +def test_zero_drawdown_is_zero_not_null(tmp_path): + req = make_request(tmp_path, SELL_OPEN, [("SH#600000", "sh_600000", "stock", 10.0)]) + res = run_backtest(req) + assert res["status"] == "succeeded" + m = res["metrics"] + req2 = make_request(tmp_path / "flat" if False else tmp_path, + 'import backtrader as bt\nclass Strategy(bt.Strategy):\n pass\n', + [("SH#600000", "sh_600000", "stock", 10.0)]) + res2 = run_backtest(req2) + assert res2["metrics"]["max_drawdown"] == 0.0 + assert res2["metrics"]["total_return"] == 0.0 + assert res2["metrics"]["trade_count"] == 0 + # initial capital and per-bar positions recorded + assert res2["initial_cash"] == 100000.0 + assert all("positions" in e for e in res2["equity"]) + + +def test_benchmark_normalized_to_initial_capital(tmp_path): + # stock 10 -> wobbles; benchmark BJ feed base differs so raw price != equity scale + req = make_request( + tmp_path, SELL_OPEN, + [("SH#600000", "sh_600000", "stock", 10.0), + ("SH000300", "sh000300", "index", 3800.0)], + benchmark="SH000300") + res = run_backtest(req) + assert all(e["benchmark"] is not None for e in res["equity"]) + # raw price never mixed into the equity scale + assert res["equity"][0]["benchmark"] == res["initial_cash"] + assert res["equity"][-1]["closes"]["SH000300"] > 3000 # raw close available separately + assert any("normalized to initial capital" in w for w in res["warnings"]) + assert res["orders"] + + +def test_index_order_rejected_broker_level(tmp_path): + specs = [("SH#600000", "sh_600000", "stock", 10.0), + ("SH000300", "sh000300", "index", 3800.0)] + req = make_request(tmp_path, BUY_INDEX, specs) + res = run_backtest(req) + assert res["status"] == "succeeded" + index_fills = [t for t in res["trades"] if t["symbol"] == "SH000300"] + assert not index_fills + assert not [t for t in res["trades"] if t["symbol"] == "SH#600000"] # strategy only bought index + user_orders = [o for o in res["orders"] if o.get("date")] # real order notifications + rejected = [o for o in user_orders if o["symbol"] == "SH000300"] + assert rejected and rejected[-1]["order_state"] == "Rejected" + rejects = [o for o in res["orders"] if o.get("order_state") == "rejected"] + assert rejects and "nontradable" in rejects[0]["reject_reason"] + assert any("nontradable" in w for w in res["warnings"]) + + +def test_index_only_as_benchmark_accepted(tmp_path): + specs = [("SH#600000", "sh_600000", "stock", 10.0), + ("SH000300", "sh000300", "index", 3800.0)] + req = make_request(tmp_path, BUY_BM, specs, benchmark="SH000300") + res = run_backtest(req) + assert res["status"] == "succeeded" + assert res["trades"][0]["symbol"] == "SH#600000" + assert res["equity"][-1]["benchmark"] is not None + + +def test_lot_size_rounded_down_to_100(tmp_path): + req = make_request(tmp_path, BUY250_TEST, [("SH#600000", "sh_600000", "stock", 10.0)]) + res = run_backtest(req) + assert res["status"] == "succeeded" + assert res["trades"], "rejection would leave no fill; rounding should keep 200" + q = res["trades"][0]["quantity"] + assert q == 200, f"expected 250 -> 200 (2 lots), got {q}" + + +def test_no_naked_short_rejected(tmp_path): + req = make_request(tmp_path, SHORT, [("SH#600000", "sh_600000", "stock", 10.0)]) + res = run_backtest(req) + short_fills = [t for t in res["trades"] if t["side"] == "sell" and t["quantity"] > 100] + # position was 100 (bought) but sell attempted 400 -> trimmed or rejected, not shorted + pos_last = res["equity"][-1]["positions"]["SH#600000"] + assert pos_last >= 0 + assert any("no naked short" in w or "would exceed" in w for w in res["warnings"]) or \ + not short_fills or short_fills[0]["quantity"] <= 100 + + +def test_t1_samebar_buy_sell_guarded(tmp_path): + # sell submitted on the same bar as buy (position still 0 at submit) must be trimmed + req = make_request(tmp_path, SELL_OPEN, [("SH#600000", "sh_600000", "stock", 10.0)]) + res = run_backtest(req) + assert res["status"] == "succeeded" + positions = [e["positions"]["SH#600000"] for e in res["equity"]] + assert min(positions) >= 0 # never negative (no short anywhere in series) + + +def test_cash_plus_position_reconciles_equity(tmp_path): + req = make_request(tmp_path, BUY_STOCK_ONLY, [("SH#600000", "sh_600000", "stock", 10.0)]) + res = run_backtest(req) + last = res["equity"][-1] + held = {s: p for s, p in last["positions"].items() if p} + expected = last["cash"] + sum(p * last["closes"][s] for s, p in held.items()) + assert abs(last["equity"] - expected) < 1e-6 |
