diff options
| author | Somhairle H. Marisol <[email protected]> | 2026-09-17 14:32:37 +0800 |
|---|---|---|
| committer | Somhairle H. Marisol <[email protected]> | 2026-09-17 14:32:37 +0800 |
| commit | 5c0ba37eda80d39e6ceca59bb1d5f4942f858995 (patch) | |
| tree | 948723f9cedf7ccb0707fa6ee516bd30fe20fd10 /tests/worker/test_fetch_cli.py | |
| download | strategy-lab-5c0ba37eda80d39e6ceca59bb1d5f4942f858995.tar.gz | |
chore: establish Strategy Lab source baseline (development, not release)
Diffstat (limited to 'tests/worker/test_fetch_cli.py')
| -rw-r--r-- | tests/worker/test_fetch_cli.py | 144 |
1 files changed, 144 insertions, 0 deletions
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 |
