diff --git a/backend/app/infrastructure/data_sources/sina.py b/backend/app/infrastructure/data_sources/sina.py index 7a09e0c..ed68b86 100644 --- a/backend/app/infrastructure/data_sources/sina.py +++ b/backend/app/infrastructure/data_sources/sina.py @@ -69,10 +69,15 @@ def _to_sina_symbol(symbol: str) -> str: def _extract_jsonp(payload: str) -> list[dict[str, Any]]: - match = re.search(r"=\s*(\[.*\])\s*$", payload.strip(), flags=re.DOTALL) - if not match: + """从 JSONP 中提取数组:容忍前导注释 / var data=([...]) 包裹 / 尾部杂字符。 + + 直接取首个 '[' 与末个 ']' 之间的内容(行情数组为扁平结构,无嵌套数组)。 + """ + start = payload.find("[") + end = payload.rfind("]") + if start == -1 or end <= start: raise DataSourceError("新浪行情返回格式无法解析") - return json.loads(match.group(1)) + return json.loads(payload[start : end + 1]) def _d(value) -> Decimal | None: diff --git a/backend/tests/test_jobs_experiments.py b/backend/tests/test_jobs_experiments.py index fa15837..0f7370a 100644 --- a/backend/tests/test_jobs_experiments.py +++ b/backend/tests/test_jobs_experiments.py @@ -11,7 +11,12 @@ from app.domain.entities.research import ( JobStatus, ResearchSpec, ) +from app.infrastructure.persistence.sqlalchemy import session as sess_mod from app.infrastructure.persistence.sqlalchemy.base import Base +from app.infrastructure.persistence.sqlalchemy.models.market import ( + StockDailyModel, + StockModel, +) from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( SqlAlchemyExperimentRepository, SqlAlchemyJobRepository, @@ -142,8 +147,42 @@ class TestJobExecutor: assert done.error +@pytest.fixture() +def seeded_api_db(tmp_path, monkeypatch) -> None: + """把全局 SessionLocal 指向 tmp 种子库(5 只股票 × 300 交易日), + Job 后台执行跑小数据,避免依赖真实全市场库(779 万行)导致过慢。""" + from sqlalchemy.orm import sessionmaker + + engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True) + Base.metadata.create_all(engine) + sf = sessionmaker(bind=engine, expire_on_commit=False) + drifts = {s: 0.003 - 0.0015 * i for i, s in enumerate(_SYMS)} + daily = synthetic_daily(drifts, n=320) + with sf() as session: + for i, sym in enumerate(_SYMS): + session.add(StockModel(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 1, 1))) + rows = [ + { + "symbol": r.symbol, + "trade_date": r.trade_date, + "source": "tushare", + "adjust": "none", + "open": float(r.open), + "high": float(r.high), + "low": float(r.low), + "close": float(r.close), + "volume": float(r.volume), + "amount": float(r.amount), + } + for r in daily.itertuples() + ] + session.execute(StockDailyModel.__table__.insert(), rows) + session.commit() + monkeypatch.setattr(sess_mod, "SessionLocal", sf) + + class TestJobsApi: - def test_submit_then_query(self) -> None: + def test_submit_then_query(self, seeded_api_db) -> None: with TestClient(app) as client: resp = client.post("/api/jobs", json=_spec().model_dump(mode="json")) assert resp.status_code == 200 @@ -161,17 +200,25 @@ class TestJobsApi: assert state["spec"]["factors"][0]["name"] == "momentum_20" assert state["experiment_id"] - def test_unknown_job_404(self) -> None: + def test_unknown_job_404(self, seeded_api_db) -> None: with TestClient(app) as client: assert client.get("/api/jobs/JOB-NOPE").status_code == 404 - def test_experiments_list_and_detail(self) -> None: + def test_experiments_list_and_detail(self, seeded_api_db) -> None: with TestClient(app) as client: - resp = client.get("/api/experiments") + # 先跑一个 job 生成归档,再验证列表与详情 + resp = client.post("/api/jobs", json=_spec().model_dump(mode="json")) assert resp.status_code == 200 - exps = resp.json() - if exps: # 本机真实库中可能有历史实验;detail 可读即通过 - first = exps[0] - detail = client.get(f"/api/experiments/{first['id']}") - assert detail.status_code == 200 - assert "result" in detail.json() + job_id = resp.json()["job_id"] + for _ in range(30): + st = client.get(f"/api/jobs/{job_id}").json()["status"] + if st in (JobStatus.SUCCESS, JobStatus.FAILED): + break + assert st == JobStatus.SUCCESS + + exps = client.get("/api/experiments").json() + assert exps, "job 成功后应有 Experiment 归档" + first = exps[0] + detail = client.get(f"/api/experiments/{first['id']}") + assert detail.status_code == 200 + assert "result" in detail.json() diff --git a/backend/tests/test_sina_provider.py b/backend/tests/test_sina_provider.py index 6c3775f..8f7560d 100644 --- a/backend/tests/test_sina_provider.py +++ b/backend/tests/test_sina_provider.py @@ -154,3 +154,18 @@ class TestFormatParity: # 与 Tushare 同 schema:必备字段齐全 assert bar.symbol == "600519.SH" assert bar.close is not None + + +class TestJsonpParens: + def test_real_world_paren_wrapped(self) -> None: + payload = 'var data=([{"day":"2024-08-30","close":"1710.0"}])' + rows = _extract_jsonp(payload) + assert rows[0]["close"] == "1710.0" + + def test_with_leading_comment_and_trailing_semicolon(self) -> None: + payload = ( + "/**/\n" + 'var data=([{"day":"2024-08-30","close":"1710.0"}]);' + ) + rows = _extract_jsonp(payload) + assert rows[0]["close"] == "1710.0"