303 lines
13 KiB
Python
303 lines
13 KiB
Python
from __future__ import annotations
|
|
|
|
import sqlite3
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
|
|
from database import ReviewDatabase
|
|
from server import DashboardService, RequestHandler
|
|
|
|
|
|
FORMULA = {"all": [{"field": "change", "operator": ">", "value": 0}]}
|
|
|
|
|
|
class AccountDataBoundaryTests(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.database = ReviewDatabase(Path(self.temp.name) / "review.db")
|
|
self.first = self.database.create_user("first_user", "salt", "hash")
|
|
self.second = self.database.create_user("second_user", "salt", "hash")
|
|
|
|
def tearDown(self) -> None:
|
|
self.temp.cleanup()
|
|
|
|
def test_custom_strategies_are_private_and_builtins_are_shared(self):
|
|
builtin_id = self.database.save_screener_strategy(
|
|
None, "共享策略", "", ["repair"], FORMULA, builtin=True
|
|
)
|
|
private_id = self.database.save_screener_strategy(
|
|
self.first["id"], "甲的策略", "", ["repair"], FORMULA
|
|
)
|
|
|
|
first_names = {item["name"] for item in self.database.list_screener_strategies(self.first["id"])}
|
|
second_names = {item["name"] for item in self.database.list_screener_strategies(self.second["id"])}
|
|
self.assertEqual(first_names, {"共享策略", "甲的策略"})
|
|
self.assertEqual(second_names, {"共享策略"})
|
|
with self.assertRaises(ValueError):
|
|
self.database.delete_screener_strategy(self.second["id"], private_id)
|
|
with self.assertRaises(ValueError):
|
|
self.database.delete_screener_strategy(self.first["id"], builtin_id)
|
|
self.assertTrue(self.database.delete_screener_strategy(self.first["id"], private_id))
|
|
|
|
def test_screener_runs_are_private(self):
|
|
self.database.save_screener_run(
|
|
self.first["id"], "20260721", "repair", "甲的策略", FORMULA,
|
|
{"candidates": [{"code": "002141"}], "meta": {}},
|
|
)
|
|
self.assertEqual(
|
|
self.database.latest_screener_run(self.first["id"], "20260722")["candidates"][0]["code"],
|
|
"002141",
|
|
)
|
|
self.assertIsNone(self.database.latest_screener_run(self.second["id"], "20260722"))
|
|
|
|
def test_latest_screener_runs_are_isolated_by_mode_and_user(self):
|
|
expected = {
|
|
"smart": "600001",
|
|
"curated": "600002",
|
|
"quant": "600003",
|
|
}
|
|
for mode, code in expected.items():
|
|
self.database.save_screener_run(
|
|
self.first["id"], "20260721", "repair", f"{mode}-strategy", FORMULA,
|
|
{"candidates": [{"code": code}], "meta": {}}, mode,
|
|
)
|
|
|
|
results = self.database.latest_screener_runs(self.first["id"], "20260722")
|
|
self.assertEqual(set(results), set(expected))
|
|
for mode, code in expected.items():
|
|
self.assertEqual(results[mode]["meta"]["mode"], mode)
|
|
self.assertEqual(results[mode]["candidates"][0]["code"], code)
|
|
self.assertEqual(
|
|
self.database.latest_screener_run(
|
|
self.first["id"], "20260722", mode
|
|
)["candidates"][0]["code"],
|
|
code,
|
|
)
|
|
self.assertEqual(
|
|
self.database.latest_screener_runs(self.second["id"], "20260722"), {}
|
|
)
|
|
|
|
def test_latest_screener_context_runs_keep_each_stage_and_strategy(self):
|
|
runs = [
|
|
("smart", "repair", "Repair", "600001"),
|
|
("smart", "repair", "Repair", "600002"),
|
|
("smart", "retreat", "Retreat", "600003"),
|
|
("curated", "repair", "Dividend", "600004"),
|
|
("curated", "repair", "Momentum", "600005"),
|
|
("quant", "repair", "Custom quant", "600006"),
|
|
("quant", "repair", "Custom quant", "600007"),
|
|
]
|
|
for mode, regime, strategy, code in runs:
|
|
self.database.save_screener_run(
|
|
self.first["id"], "20260722", regime, strategy, FORMULA,
|
|
{"candidates": [{"code": code}], "meta": {}}, mode,
|
|
)
|
|
|
|
results = self.database.latest_screener_context_runs(
|
|
self.first["id"], "20260722"
|
|
)
|
|
by_context = {
|
|
(
|
|
item["meta"]["mode"],
|
|
item["meta"]["regime"] if item["meta"]["mode"] == "smart" else "",
|
|
item["meta"]["strategy_name"] if item["meta"]["mode"] != "quant" else "",
|
|
): item["candidates"][0]["code"]
|
|
for item in results
|
|
}
|
|
|
|
self.assertEqual(by_context, {
|
|
("smart", "repair", "Repair"): "600002",
|
|
("smart", "retreat", "Retreat"): "600003",
|
|
("curated", "", "Dividend"): "600004",
|
|
("curated", "", "Momentum"): "600005",
|
|
("quant", "", ""): "600007",
|
|
})
|
|
self.assertEqual(
|
|
self.database.latest_screener_context_runs(
|
|
self.second["id"], "20260722"
|
|
),
|
|
[],
|
|
)
|
|
|
|
def test_mentor_messages_are_scoped_by_user_mentor_and_date(self):
|
|
self.database.save_mentor_exchange(
|
|
self.first["id"], "mentor-a", "20260721", "怎么看?", "先看承接。", "20260721"
|
|
)
|
|
self.assertEqual(
|
|
[item["role"] for item in self.database.list_mentor_messages(
|
|
self.first["id"], "mentor-a", "20260721"
|
|
)],
|
|
["user", "assistant"],
|
|
)
|
|
self.assertEqual(
|
|
self.database.list_mentor_messages(self.second["id"], "mentor-a", "20260721"), []
|
|
)
|
|
self.assertEqual(
|
|
self.database.list_mentor_messages(self.first["id"], "mentor-b", "20260721"), []
|
|
)
|
|
self.assertEqual(
|
|
self.database.list_mentor_messages(self.first["id"], "mentor-a", "20260722"), []
|
|
)
|
|
self.assertEqual(
|
|
self.database.delete_mentor_messages(self.second["id"], "mentor-a", "20260721"), 0
|
|
)
|
|
self.assertEqual(
|
|
self.database.delete_mentor_messages(self.first["id"], "mentor-a", "20260721"), 2
|
|
)
|
|
|
|
def test_mentor_preferences_are_scoped_by_user(self):
|
|
self.database.save_mentor_preferences(
|
|
self.first["id"], ["mentor-b", "mentor-a"], {"mentor-b"}
|
|
)
|
|
self.database.save_mentor_preferences(
|
|
self.second["id"], ["mentor-a", "mentor-b"], set()
|
|
)
|
|
|
|
first = self.database.list_mentor_preferences(self.first["id"])
|
|
second = self.database.list_mentor_preferences(self.second["id"])
|
|
self.assertEqual([item["mentor_id"] for item in first], ["mentor-b", "mentor-a"])
|
|
self.assertTrue(first[0]["pinned"])
|
|
self.assertEqual([item["mentor_id"] for item in second], ["mentor-a", "mentor-b"])
|
|
self.assertFalse(any(item["pinned"] for item in second))
|
|
|
|
def test_latest_data_snapshot_skips_demo_and_future_records(self):
|
|
self.database.save_data_snapshot(
|
|
"stock_detail", "002141:20260718", "tushare", {"marker": "real"}
|
|
)
|
|
self.database.save_data_snapshot(
|
|
"stock_detail", "002141:20260719", "demo", {"marker": "demo"}
|
|
)
|
|
self.database.save_data_snapshot(
|
|
"stock_detail", "002141:20260723", "tushare", {"marker": "future"}
|
|
)
|
|
payload = self.database.get_latest_data_snapshot(
|
|
"stock_detail", "002141:", "002141:20260722", exclude_source="demo"
|
|
)
|
|
self.assertEqual(payload["marker"], "real")
|
|
|
|
def test_stock_detail_never_falls_back_to_demo_data(self):
|
|
self.database.save_data_snapshot(
|
|
"stock_detail",
|
|
"002141:20260721",
|
|
"demo",
|
|
{"meta": {"source": "demo"}, "stock": {"code": "002141"}, "prices": [{}]},
|
|
)
|
|
service = DashboardService.__new__(DashboardService)
|
|
service.database = self.database
|
|
service._system_credentials = {"tushare_token": ""}
|
|
service._request_context = SimpleNamespace(user_id=self.first["id"])
|
|
|
|
with self.assertRaisesRegex(ValueError, "暂无 002141 的真实行情数据"):
|
|
service.get_stock_detail("002141", "2026-07-22")
|
|
|
|
|
|
class LegacyStrategyMigrationTests(unittest.TestCase):
|
|
def test_legacy_custom_strategy_and_run_move_to_first_admin(self):
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
database_path = Path(directory) / "legacy.db"
|
|
connection = sqlite3.connect(database_path)
|
|
try:
|
|
connection.executescript(
|
|
"""
|
|
CREATE TABLE users (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
username TEXT NOT NULL UNIQUE,
|
|
password_salt TEXT NOT NULL,
|
|
password_hash TEXT NOT NULL,
|
|
created_at TEXT NOT NULL,
|
|
updated_at TEXT NOT NULL
|
|
);
|
|
INSERT INTO users VALUES (1, 'legacy_admin', 'salt', 'hash', '2026-01-01', '2026-01-01');
|
|
CREATE TABLE screener_strategies (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
name TEXT NOT NULL,
|
|
description TEXT NOT NULL DEFAULT '',
|
|
regimes TEXT NOT NULL,
|
|
formula TEXT NOT NULL,
|
|
builtin INTEGER NOT NULL DEFAULT 0,
|
|
created_at TEXT NOT NULL,
|
|
updated_at TEXT NOT NULL
|
|
);
|
|
INSERT INTO screener_strategies VALUES
|
|
(1, '旧策略', '', '["repair"]', '{"all":[]}', 0, '2026-01-01', '2026-01-01'),
|
|
(2, '旧内置', '', '["repair"]', '{"all":[]}', 1, '2026-01-01', '2026-01-01');
|
|
CREATE TABLE screener_runs (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
trade_date TEXT NOT NULL,
|
|
regime TEXT NOT NULL,
|
|
strategy_name TEXT NOT NULL,
|
|
formula TEXT NOT NULL,
|
|
result TEXT NOT NULL,
|
|
created_at TEXT NOT NULL
|
|
);
|
|
INSERT INTO screener_runs VALUES
|
|
(1, '20260721', 'repair', '旧策略', '{}', '{"meta":{}}', '2026-07-21'),
|
|
(2, '20260721', 'repair', '旧精选策略',
|
|
'{"meta":{"library":"curated"}}', '{"meta":{}}', '2026-07-21'),
|
|
(3, '20260721', 'repair', '自定义量化公式',
|
|
'{"meta":{"library":"custom","category":"量化公式"}}',
|
|
'{"meta":{}}', '2026-07-21');
|
|
"""
|
|
)
|
|
connection.commit()
|
|
finally:
|
|
connection.close()
|
|
|
|
migrated = ReviewDatabase(database_path)
|
|
strategies = migrated.list_screener_strategies(1)
|
|
owners = {item["name"]: item["user_id"] for item in strategies}
|
|
self.assertEqual(owners["旧策略"], 1)
|
|
self.assertIsNone(owners["旧内置"])
|
|
self.assertIsNotNone(migrated.latest_screener_run(1, "20260722"))
|
|
self.assertEqual(
|
|
set(migrated.latest_screener_runs(1, "20260722")),
|
|
{"smart", "curated", "quant"},
|
|
)
|
|
|
|
|
|
class PublicKnowledgePermissionTests(unittest.TestCase):
|
|
@staticmethod
|
|
def handler(path: str, method_name: str):
|
|
handler = RequestHandler.__new__(RequestHandler)
|
|
handler.path = path
|
|
handler.require_auth = lambda: True
|
|
handler.require_csrf = lambda: True
|
|
handler.require_admin = lambda: False
|
|
handler.require_member = lambda: True
|
|
setattr(handler, method_name, lambda: (_ for _ in ()).throw(AssertionError("mutation ran")))
|
|
return handler
|
|
|
|
def test_regular_user_cannot_change_reason_or_seat_alias(self):
|
|
RequestHandler.do_POST(self.handler("/api/reasons", "save_reason"))
|
|
RequestHandler.do_POST(self.handler("/api/seat-aliases", "save_seat_alias"))
|
|
|
|
def test_regular_user_cannot_change_sector_phase(self):
|
|
RequestHandler.do_POST(
|
|
self.handler("/api/heaven/sector-phases", "save_sector_phase_override")
|
|
)
|
|
RequestHandler.do_DELETE(
|
|
self.handler("/api/heaven/sector-phases/%E6%B2%B9%E6%B0%94", "send_json")
|
|
)
|
|
|
|
def test_static_shell_bypasses_the_api_access_registry(self):
|
|
handler = RequestHandler.__new__(RequestHandler)
|
|
handler.path = "/"
|
|
handler.require_auth = lambda: (_ for _ in ()).throw(
|
|
AssertionError("static request required authentication")
|
|
)
|
|
handler.require_access = lambda *_: (_ for _ in ()).throw(
|
|
AssertionError("static request entered the API registry")
|
|
)
|
|
served: list[str] = []
|
|
handler.serve_static = served.append
|
|
|
|
RequestHandler.do_GET(handler)
|
|
|
|
self.assertEqual(served, ["/"])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|