feat: track screener candidates across five trading days
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user