159 lines
5.8 KiB
Python
159 lines
5.8 KiB
Python
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()
|