feat: track screener candidates across five trading days

This commit is contained in:
leefer
2026-07-22 23:55:33 +08:00
parent 50dbb8a3ff
commit 5ae2791dd3
9 changed files with 510 additions and 1 deletions
+2
View File
@@ -9,9 +9,11 @@ class ApiAccessPolicyTests(unittest.TestCase):
def test_member_workspaces_are_consistently_protected(self):
cases = {
("GET", "/api/screener/setup"): "member",
("GET", "/api/screener/tracking"): "member",
("GET", "/api/mentors/messages"): "member",
("GET", "/api/heaven/setup"): "member",
("POST", "/api/screener/run"): "member",
("POST", "/api/screener/tracking/refresh"): "member",
("POST", "/api/mentors/chat"): "member",
("POST", "/api/heaven/interpret"): "member",
("DELETE", "/api/screener/strategies/42"): "member",
+104
View File
@@ -0,0 +1,104 @@
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)
if __name__ == "__main__":
unittest.main()