summaryrefslogtreecommitdiff
path: root/tests/worker/test_fetch_cli.py
diff options
context:
space:
mode:
authorSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
committerSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
commit5c0ba37eda80d39e6ceca59bb1d5f4942f858995 (patch)
tree948723f9cedf7ccb0707fa6ee516bd30fe20fd10 /tests/worker/test_fetch_cli.py
downloadstrategy-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.py144
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