Files
xiaobaifupan/tests/test_strategy_tracking.py
T

159 lines
5.8 KiB
Python

from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from database import ReviewDatabase
from strategy_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()