from __future__ import annotations import tempfile import unittest from pathlib import Path from database import ReviewDatabase from backend.features.screener.tracking import StrategyTrackingService class StrategyTrackingTests(unittest.TestCase): def setUp(self) -> None: self.temp = tempfile.TemporaryDirectory() self.database = ReviewDatabase(Path(self.temp.name) / "review.db") self.owner = self.database.create_user("track_owner", "salt", "hash") self.other = self.database.create_user("track_other", "salt", "hash") self.service = StrategyTrackingService(self.database) self.run_id = self.database.save_screener_run( self.owner["id"], "20260710", "repair", "修复策略", {}, {"meta": {}} ) def tearDown(self) -> None: self.temp.cleanup() def test_run_is_recorded_once_and_private_to_owner(self): candidates = [ { "ts_code": "002141.SZ", "code": "002141", "name": "贤丰控股", "sector": "油气开采", "price": 10, } ] self.assertEqual( self.service.record_run( self.owner["id"], self.run_id, "20260710", "修复策略", candidates ), 1, ) self.service.record_run( self.owner["id"], self.run_id, "20260710", "修复策略", candidates ) self.assertEqual(len(self.database.list_strategy_tracks(self.owner["id"])), 1) self.assertEqual(self.database.list_strategy_tracks(self.other["id"]), []) def test_tracking_uses_next_five_trading_bars(self): self.service.record_run( self.owner["id"], self.run_id, "20260710", "修复策略", [{"ts_code": "002141.SZ", "code": "002141", "name": "贤丰控股", "sector": "油气开采", "price": 10}], ) rows = [] values = ( ("20260713", 10.2, 10.8, 9.8, 10.5), ("20260714", 10.5, 11.0, 10.1, 10.8), ("20260715", 10.8, 11.5, 10.4, 11.2), ("20260716", 11.2, 11.4, 9.5, 9.8), ("20260717", 9.8, 10.3, 9.0, 10.0), ("20260720", 10.0, 20.0, 1.0, 19.0), ) for trade_date, open_price, high, low, close in values: rows.append( { "trade_date": trade_date, "ts_code": "002141.SZ", "open": open_price, "high": high, "low": low, "close": close, } ) self.database.upsert_daily_bars(rows) payload = self.service.list_tracking(self.owner["id"]) item = payload["batches"][0]["items"][0] self.assertEqual(item["status"], "已完成") self.assertEqual(item["t1_open"], 2.0) self.assertEqual(item["t1_close"], 5.0) self.assertEqual(item["t3_close"], 12.0) self.assertEqual(item["t5_close"], 0.0) self.assertEqual(item["max_gain"], 15.0) self.assertEqual(item["max_drawdown"], -10.0) self.assertEqual(payload["summary"]["t1_win_rate"], 100.0) self.assertEqual(payload["summary"]["t5_win_rate"], 0.0) def test_partial_tracking_reports_available_days(self): metrics = self.service.calculate_metrics( 20, [ {"open": 20, "high": 21, "low": 19, "close": 20.5}, {"open": 20.5, "high": 22, "low": 20, "close": 21}, ], ) self.assertEqual(metrics["status"], "跟踪中 2/5") self.assertIsNone(metrics["t3_close"]) self.assertEqual(metrics["max_gain"], 10.0) self.assertEqual(metrics["max_drawdown"], -5.0) def test_candidate_is_added_manually_and_can_be_removed_by_owner(self): run_id = self.database.save_screener_run( self.owner["id"], "20260711", "repair", "手动跟踪策略", {}, { "meta": {}, "candidates": [{ "ts_code": "600000.SH", "code": "600000", "name": "浦发银行", "sector": "银行", "price": 12.5, }], }, ) result = self.service.add_candidate(self.owner["id"], run_id, "600000") self.assertEqual(result["added"], 1) tracks = self.database.list_strategy_tracks(self.owner["id"]) self.assertEqual(len(tracks), 1) self.assertEqual(tracks[0]["code"], "600000") self.assertEqual(self.database.list_strategy_tracks(self.other["id"]), []) with self.assertRaises(ValueError): self.service.add_candidate(self.other["id"], run_id, "600000") removed = self.service.remove_candidate(self.owner["id"], tracks[0]["id"]) self.assertTrue(removed["deleted"]) self.assertEqual(removed["tracking"]["batches"], []) def test_shared_automatic_run_can_be_added_to_private_tracking(self): run_id = self.database.save_screener_run( 0, "20260711", "repair", "系统盘后策略", {}, { "meta": {}, "candidates": [{ "ts_code": "600000.SH", "code": "600000", "name": "浦发银行", "sector": "银行", "price": 12.5, }], }, ) result = self.service.add_candidate(self.other["id"], run_id, "600000") self.assertEqual(result["added"], 1) self.assertEqual(len(self.database.list_strategy_tracks(self.other["id"])), 1) self.assertEqual(self.database.list_strategy_tracks(self.owner["id"]), []) if __name__ == "__main__": unittest.main()