"""`app.cli.restore_experiment_from_job` 测试:从 Job 结果副本重建归档。 为什么要测:删除归档是正常功能,但**历史归档在 `job.result_json` 里另存了一份完整结果** (完整存档上线前的双写遗留),所以「删了能不能救回来」是有确定答案的: 历史归档能按原 id 重建,新归档(`job.result_json IS NULL`)不能。 这个工具承载该结论,必须验证真写库路径与全部拒绝路径,而不是只做 dry-run。 测试全部在 conftest 强制的 /tmp sqlite 上跑,不触碰真实 MySQL。 """ from __future__ import annotations import json from datetime import date, datetime from unittest.mock import patch from app.cli import restore_experiment_from_job as rj from app.domain.entities.research import ( BacktestResult, BacktestSummary, CurvePoint, ) from app.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker def _payload(final_equity: float = 1_100_000.0) -> str: """一份最小可用的回测结果 JSON(与归档里存的结构同源)。""" result = BacktestResult( summary=BacktestSummary( start=date(2024, 1, 2), end=date(2024, 6, 28), initial_capital=1_000_000.0, final_equity=final_equity, total_return_pct=10.0, annual_return_pct=21.5, sharpe=1.2, max_drawdown_pct=-8.5, volatility_pct=18.0, win_rate_pct=55.0, total_trades=12, avg_turnover_pct=30.0, ), equity_curve=[CurvePoint(date=date(2024, 1, 2), value=1_000_000.0)], drawdown=[], monthly_returns=[], yearly_returns=[], positions=[], trades=[], turnover_pct=30.0, ) return json.dumps(result.model_dump(mode="json"), ensure_ascii=False) def _make_session(tmp_path): engine = create_engine(f"sqlite:///{tmp_path}/restore.db", future=True) Base.metadata.create_all(engine) return sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) def _seed_job(session, *, job_id="JOB-TEST", exp_id="EXP-TEST", result_json=None): session.add( JobModel( id=job_id, kind="backtest", status="success", spec_json=json.dumps({"type": "backtest", "factors": []}), result_json=result_json, experiment_id=exp_id, created_at=datetime(2026, 1, 5, 9, 0, 0), finished_at=datetime(2026, 1, 5, 9, 0, 1), ) ) session.commit() def _run(session_factory, argv): with patch.object(rj, "SessionLocal", session_factory): return rj.main(argv) def test_dry_run_writes_nothing(tmp_path, capsys) -> None: factory = _make_session(tmp_path) with factory() as s: _seed_job(s, result_json=_payload()) code = _run(factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234"]) out = capsys.readouterr().out assert code == 0 assert "(dry-run)" in out with factory() as s: assert s.get(ExperimentModel, "EXP-TEST") is None def test_restores_archive_from_job_copy(tmp_path, capsys) -> None: """真写库路径:原 id、逐字复制、摘要按归档口径重算、data_version 留空。""" factory = _make_session(tmp_path) payload = _payload() with factory() as s: _seed_job(s, result_json=payload) code = _run( factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"] ) assert code == 0 assert "已重建" in capsys.readouterr().out with factory() as s: row = s.get(ExperimentModel, "EXP-TEST") assert row is not None assert row.kind == "backtest" assert row.result_json == payload # 逐字复制,不重新计算 assert row.job_id == "JOB-TEST" assert row.code_version == "abc1234" # 历史归档当年没有数据指纹:留空而不是补今天的(伪造复现依据) assert row.data_version is None assert row.created_at == datetime(2026, 1, 5, 9, 0, 1) assert row.summary_text == "总收益 10.00% · 年化 21.50% · 回撤 -8.50%" def test_refuses_when_archive_already_exists(tmp_path, capsys) -> None: factory = _make_session(tmp_path) with factory() as s: _seed_job(s, result_json=_payload()) s.add( ExperimentModel( id="EXP-TEST", kind="backtest", spec_json="{}", result_json=_payload(1.0), summary_text="原有", code_version="old", data_version=None, job_id="JOB-TEST", created_at=datetime(2026, 1, 5, 9, 0, 1), ) ) s.commit() code = _run( factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"] ) assert code == 1 assert "已存在" in capsys.readouterr().err with factory() as s: assert s.get(ExperimentModel, "EXP-TEST").summary_text == "原有" # 未被改动 def test_force_overwrites_existing_archive(tmp_path) -> None: factory = _make_session(tmp_path) payload = _payload() with factory() as s: _seed_job(s, result_json=payload) s.add( ExperimentModel( id="EXP-TEST", kind="backtest", spec_json="{}", result_json="{}", summary_text="旧的", code_version="old", data_version="d20260101", job_id="JOB-TEST", created_at=datetime(2026, 1, 5, 9, 0, 1), ) ) s.commit() code = _run( factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply", "--force"], ) assert code == 0 with factory() as s: row = s.get(ExperimentModel, "EXP-TEST") assert row.result_json == payload assert row.code_version == "abc1234" assert row.data_version is None # 覆盖后也留空,不继承旧指纹 def test_refuses_new_archive_without_job_copy(tmp_path, capsys) -> None: """完整存档上线后结果只存归档一份 → 删了不可恢复,工具必须如实拒绝。""" factory = _make_session(tmp_path) with factory() as s: _seed_job(s, result_json=None) code = _run( factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"] ) assert code == 1 err = capsys.readouterr().err assert "result_json 为空" in err assert "不可恢复" in err def test_requires_explicit_code_version(tmp_path, capsys) -> None: """不猜版本:缺失时直接拒绝(退出码 2),避免写进一个编造的复现依据。""" factory = _make_session(tmp_path) with factory() as s: _seed_job(s, result_json=_payload()) code = _run(factory, ["--job-id", "JOB-TEST", "--apply"]) assert code == 2 assert "code-version" in capsys.readouterr().err def test_empty_code_version_is_allowed_as_unknown(tmp_path) -> None: """显式传空串=如实表示"版本未知",此时写 NULL 而不是空字符串。""" factory = _make_session(tmp_path) with factory() as s: _seed_job(s, result_json=_payload()) code = _run( factory, ["--job-id", "JOB-TEST", "--code-version", "", "--apply"] ) assert code == 0 with factory() as s: assert s.get(ExperimentModel, "EXP-TEST").code_version is None def test_rejects_unknown_job_and_job_without_experiment_id(tmp_path, capsys) -> None: factory = _make_session(tmp_path) with factory() as s: _seed_job(s, job_id="JOB-NOEXP", exp_id=None, result_json=_payload()) assert _run(factory, ["--job-id", "JOB-MISSING", "--code-version", "x"]) == 1 assert "不存在" in capsys.readouterr().err assert _run(factory, ["--job-id", "JOB-NOEXP", "--code-version", "x"]) == 1 assert "experiment_id" in capsys.readouterr().err